From f72c219b42e1d0729e370caeead662bfcf384dd8 Mon Sep 17 00:00:00 2001 From: 4ment Date: Thu, 10 Sep 2026 21:43:26 +1000 Subject: [PATCH 01/10] Add cluster viewer tool and HTML template for cpplink clusters --- tools/cluster_server.py | 390 +++++++++++++++++++++++ tools/cluster_view.py | 309 ++++++++++++++++++ tools/cluster_view_template.html | 521 +++++++++++++++++++++++++++++++ 3 files changed, 1220 insertions(+) create mode 100755 tools/cluster_server.py create mode 100755 tools/cluster_view.py create mode 100644 tools/cluster_view_template.html diff --git a/tools/cluster_server.py b/tools/cluster_server.py new file mode 100755 index 0000000..21f9450 --- /dev/null +++ b/tools/cluster_server.py @@ -0,0 +1,390 @@ +#!/usr/bin/env python3 +# Copyright 2026 Mathieu Fourment +# SPDX-License-Identifier: MIT +"""Serve the cluster viewer over a DuckDB cache, for runs too large to embed. + + tools/cluster_server.py --schema s.json --clusters clusters.csv \ + --edges edges --truth truth.csv data.parquet + +Same page as cluster_view.py, but the clusters are queried rather than written +into the file, so every cluster of the run is reachable instead of a sample. + +The cache is built once and reused until an input changes. It holds only what +the viewer can ever show -- the clustered records, never the singletons -- so +it is a fraction of the dataset: at 20M records with 2.9M of them clustered, +the parquet is read exactly once and every later request touches a table seven +times smaller than the file. +""" + +import argparse +import json +import os +import sys +import threading +import time +import webbrowser +from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer +from urllib.parse import urlparse, parse_qs + +import duckdb +import numpy as np +import pyarrow as pa +import pyarrow.parquet as pq + +sys.path.insert(0, os.path.dirname(os.path.abspath(__file__))) +from cluster_view import (EDGE_DTYPE, EDGE_MAGIC, cell, # noqa: E402 + read_truth, schema_columns) + +EDGE_CHUNK = 1 << 20 +CACHE_VERSION = 2 + + +def quoted(name): + return '"' + name.replace('"', '""') + '"' + + +def fingerprint(args, columns): + """What the cache was built from, so a changed input rebuilds it.""" + parts = [CACHE_VERSION, columns, args.max_rows] + for path in [args.clusters, args.truth] + list(args.data): + if path: + parts.append([os.path.abspath(path), os.path.getmtime(path), + os.path.getsize(path)]) + if args.edges: + for name in sorted(os.listdir(args.edges)): + if name.endswith(".bin"): + full = os.path.join(args.edges, name) + parts.append([full, os.path.getmtime(full), os.path.getsize(full)]) + return json.dumps(parts, sort_keys=True) + + +def load_edges(conn, directory): + """Every shard edge into `raw_edges`, a chunk at a time. + + The shards hold row indices rather than ids, which is why this is worth + doing once: the join that names them costs a pass over the record ids, and + a served request must not pay it. + """ + conn.execute("CREATE TABLE raw_edges (a UINTEGER, b UINTEGER, w DOUBLE)") + total = 0 + for name in sorted(n for n in os.listdir(directory) if n.endswith(".bin")): + with open(os.path.join(directory, name), "rb") as handle: + if handle.read(len(EDGE_MAGIC)) != EDGE_MAGIC: + raise SystemExit(f"{name}: not a cpplink edge shard") + while True: + data = handle.read(EDGE_CHUNK * EDGE_DTYPE.itemsize) + if not data: + break + block = np.frombuffer(data, dtype=EDGE_DTYPE) + chunk = pa.table({"a": block["a"], "b": block["b"], "w": block["w"]}) + conn.register("chunk", chunk) + conn.execute("INSERT INTO raw_edges SELECT a, b, w FROM chunk") + conn.unregister("chunk") + total += block.size + return total + + +def build(conn, args, id_column, columns, say): + """Fill the cache: members, their records, their edges, per-cluster stats.""" + files = [os.path.abspath(p) for p in args.data] + picked = ", ".join(quoted(c) for c in columns) + + say("reading the cluster assignment") + conn.execute(f""" + CREATE TABLE members AS + SELECT unique_id AS uid, cluster_id, cluster_size AS size + FROM read_csv('{args.clusters}', header = true, + columns = {{'unique_id': 'VARCHAR', 'cluster_id': 'VARCHAR', + 'cluster_size': 'BIGINT'}}) + WHERE cluster_size >= {args.min_size} + """) + + say("reading the clustered records") + conn.execute(f""" + CREATE TABLE records AS + SELECT m.cluster_id, d.{quoted(id_column)} AS uid, {picked}, + lower(concat_ws(' ', d.{quoted(id_column)}, {picked})) AS text + FROM read_parquet({files}) d + JOIN members m ON m.uid = d.{quoted(id_column)} + """) + + if args.edges: + say("naming the edges") + # A row index in a shard is a position in the inputs read in order, and + # file_row_number is that position inside one file, so the offsets that + # separate the inputs are the only thing to carry across. + conn.execute("CREATE TABLE rowmap (rid BIGINT, uid VARCHAR)") + offset = 0 + for path in files: + conn.execute(f""" + INSERT INTO rowmap + SELECT {offset} + file_row_number, {quoted(id_column)} + FROM read_parquet('{path}', file_row_number = true) + """) + offset += pq.ParquetFile(path).metadata.num_rows + load_edges(conn, args.edges) + conn.execute(""" + CREATE TABLE edges AS + SELECT m.cluster_id, ra.uid AS a, rb.uid AS b, e.w AS weight + FROM raw_edges e + JOIN rowmap ra ON ra.rid = e.a + JOIN rowmap rb ON rb.rid = e.b + JOIN members m ON m.uid = ra.uid + """) + conn.execute("DROP TABLE raw_edges") + conn.execute("DROP TABLE rowmap") + conn.execute("CREATE INDEX edges_by_cluster ON edges (cluster_id)") + else: + conn.execute("CREATE TABLE edges (cluster_id VARCHAR, a VARCHAR, " + "b VARCHAR, weight DOUBLE)") + + if args.truth: + say("closing the known pairs") + uids = [row[0] for row in conn.execute("SELECT uid FROM members").fetchall()] + groups = read_truth(args.truth, set(uids)) + table = pa.table({"uid": pa.array(list(groups.keys()), pa.string()), + "grp": pa.array(list(groups.values()), pa.int64())}) + conn.register("groups", table) + conn.execute("CREATE TABLE truth AS SELECT uid, grp FROM groups") + conn.unregister("groups") + else: + conn.execute("CREATE TABLE truth (uid VARCHAR, grp BIGINT)") + + say("counting what each cluster agrees on") + # count(DISTINCT x) skips nulls, which is the page's rule: a column with one + # value and some gaps still agrees. + splits = " + ".join( + f"CASE WHEN count(DISTINCT {quoted(c)}) > 1 THEN 1 ELSE 0 END" for c in columns) + conn.execute(f""" + CREATE TABLE stats AS + WITH per_cluster AS ( + SELECT cluster_id, count(*) AS size, {splits} AS discord + FROM records GROUP BY cluster_id), + per_edge AS ( + SELECT cluster_id, count(*) AS edge_count, min(weight) AS weakest + FROM edges GROUP BY cluster_id), + per_truth AS ( + SELECT r.cluster_id, + count(DISTINCT coalesce(CAST(t.grp AS VARCHAR), 'x' || r.uid)) AS entities + FROM records r LEFT JOIN truth t ON t.uid = r.uid + GROUP BY r.cluster_id) + SELECT c.cluster_id AS id, c.size, c.discord, + coalesce(e.edge_count, 0) AS edge_count, e.weakest, + coalesce(t.entities, 0) AS entities + FROM per_cluster c + LEFT JOIN per_edge e ON e.cluster_id = c.cluster_id + LEFT JOIN per_truth t ON t.cluster_id = c.cluster_id + """) + conn.execute("CREATE INDEX stats_by_id ON stats (id)") + conn.execute("CREATE INDEX records_by_cluster ON records (cluster_id)") + + +ORDERS = { + "discord": "discord DESC, size DESC, id", + "size": "size DESC, discord DESC, id", + "weakest": "weakest NULLS LAST, size DESC, id", + "id": "id", +} + +FILTERS = { + "all": "TRUE", + "split": "discord > 0", + "transitive": "size * (size - 1) / 2 > edge_count", + "mixed": "entities > 1", +} + + +class Viewer: + """The queries the page makes, over one shared read-only connection.""" + + def __init__(self, conn, columns, source, max_rows, has_truth): + self.conn = conn + self.columns = columns + self.source = source + self.max_rows = max_rows + self.has_truth = has_truth + self.lock = threading.Lock() + with self.lock: + self.totals = conn.execute( + "SELECT count(*), coalesce(sum(size), 0) FROM stats").fetchone() + + def query(self, sql, params=()): + with self.lock: + return self.conn.execute(sql, params).fetchall() + + def summary(self): + return {"columns": self.columns, "source": self.source, + "totals": {"clusters": self.totals[0], "records": self.totals[1]}, + "has_truth": self.has_truth} + + def listing(self, q, sort, only, offset, limit): + where = FILTERS.get(only, "TRUE") + params = [] + if q: + where += (" AND id IN (SELECT cluster_id FROM records " + "WHERE text LIKE ? ESCAPE '\\')") + params.append("%" + q.lower().replace("\\", "\\\\") + .replace("%", "\\%").replace("_", "\\_") + "%") + matched = self.query(f"SELECT count(*) FROM stats WHERE {where}", params)[0][0] + rows = self.query( + f"SELECT id, size, discord, edge_count, weakest, entities FROM stats " + f"WHERE {where} ORDER BY {ORDERS.get(sort, ORDERS['discord'])} " + f"LIMIT {int(limit)} OFFSET {int(offset)}", params) + return {"matched": matched, + "clusters": [{"id": r[0], "size": r[1], "discord": r[2], + "edge_count": r[3], "weakest": r[4], "entities": r[5]} + for r in rows]} + + def cluster(self, cluster_id): + head = self.query("SELECT id, size, discord, edge_count, weakest, entities " + "FROM stats WHERE id = ?", [cluster_id]) + if not head: + return None + picked = ", ".join(quoted(c) for c in self.columns) + rows = self.query( + f"SELECT uid, {picked} FROM records WHERE cluster_id = ? " + f"ORDER BY uid LIMIT {int(self.max_rows)}", [cluster_id]) + ids = [r[0] for r in rows] + groups, labels = {}, {} + if self.has_truth: + for uid, grp in self.query( + "SELECT uid, grp FROM truth WHERE uid IN " + "(SELECT uid FROM records WHERE cluster_id = ?)", [cluster_id]): + groups[uid] = grp + for uid in ids: + key = groups.get(uid, "x" + uid) + labels.setdefault(key, len(labels)) + # The whole cluster's distinct counts, not the page of rows shown. + counts = self.query( + "SELECT " + ", ".join(f"count(DISTINCT {quoted(c)})" for c in self.columns) + + " FROM records WHERE cluster_id = ?", [cluster_id])[0] + edges = self.query( + "SELECT a, b, weight FROM edges WHERE cluster_id = ? AND a IN " + "(SELECT uid FROM records WHERE cluster_id = ?)", [cluster_id, cluster_id]) + keep = set(ids) + answer = { + "id": head[0][0], "size": head[0][1], "discord": head[0][2], + "edge_count": head[0][3], "weakest": head[0][4], "entities": head[0][5], + "dist": list(counts), + "rows": [{"id": r[0], "v": [cell(v) for v in r[1:]]} for r in rows], + "edges": [[a, b, round(w, 3)] for a, b, w in edges + if a in keep and b in keep], + } + if self.has_truth: + for row in answer["rows"]: + row["t"] = labels[groups.get(row["id"], "x" + row["id"])] + return answer + + +def handler_for(viewer, page): + class Handler(BaseHTTPRequestHandler): + def log_message(self, *_): + pass + + def send_json(self, payload, status=200): + body = json.dumps(payload).encode() + self.send_response(status) + self.send_header("Content-Type", "application/json") + self.send_header("Content-Length", str(len(body))) + self.end_headers() + self.wfile.write(body) + + def do_GET(self): + url = urlparse(self.path) + query = {k: v[0] for k, v in parse_qs(url.query).items()} + try: + if url.path in ("/", "/index.html"): + body = page.encode() + self.send_response(200) + self.send_header("Content-Type", "text/html; charset=utf-8") + self.send_header("Content-Length", str(len(body))) + self.end_headers() + self.wfile.write(body) + elif url.path == "/api/summary": + self.send_json(viewer.summary()) + elif url.path == "/api/list": + self.send_json(viewer.listing( + query.get("q", ""), query.get("sort", "discord"), + query.get("only", "all"), int(query.get("offset", 0)), + min(int(query.get("limit", 100)), 500))) + elif url.path == "/api/cluster": + found = viewer.cluster(query.get("id", "")) + self.send_json(found or {"error": "no such cluster"}, + 200 if found else 404) + else: + self.send_json({"error": "not found"}, 404) + except BrokenPipeError: + pass + except Exception as trouble: # a bad query should not kill the server + self.send_json({"error": str(trouble)}, 500) + + return Handler + + +def main(): + ap = argparse.ArgumentParser(description=__doc__, + formatter_class=argparse.RawDescriptionHelpFormatter) + ap.add_argument("data", nargs="+", help="the parquet file(s) the run read, in order") + ap.add_argument("--schema", required=True) + ap.add_argument("--clusters", required=True, help="clusters.csv from cpplink cluster") + ap.add_argument("--edges", help="the shard directory, for within-cluster weights") + ap.add_argument("--truth", help="known pairs, to colour members by true entity") + ap.add_argument("--cache", default="cluster_view.duckdb") + ap.add_argument("--rebuild", action="store_true", help="rebuild the cache first") + ap.add_argument("--min-size", type=int, default=2) + ap.add_argument("--max-rows", type=int, default=200, + help="members shown per cluster; the rest are counted only") + ap.add_argument("--port", type=int, default=8770) + ap.add_argument("--open", action="store_true", help="open a browser on it") + args = ap.parse_args() + + id_column, columns = schema_columns(args.schema) + stamp = fingerprint(args, columns) + + fresh = args.rebuild or not os.path.exists(args.cache) + if not fresh: + try: + conn = duckdb.connect(args.cache, read_only=True) + fresh = conn.execute("SELECT stamp FROM meta").fetchone()[0] != stamp + conn.close() + except Exception: + fresh = True + + if fresh: + if os.path.exists(args.cache): + os.remove(args.cache) + start = time.time() + conn = duckdb.connect(args.cache) + + def say(what): + print(f" {time.time() - start:6.1f}s {what}", flush=True) + + print(f"building {args.cache}") + build(conn, args, id_column, columns, say) + conn.execute("CREATE TABLE meta (stamp VARCHAR)") + conn.execute("INSERT INTO meta VALUES (?)", [stamp]) + conn.close() + size = os.path.getsize(args.cache) / 1e6 + print(f" {time.time() - start:6.1f}s done, {size:.0f} MB") + + conn = duckdb.connect(args.cache, read_only=True) + viewer = Viewer(conn, columns, ", ".join(os.path.basename(p) for p in args.data), + args.max_rows, bool(args.truth)) + + here = os.path.dirname(os.path.abspath(__file__)) + with open(os.path.join(here, "cluster_view_template.html")) as handle: + page = handle.read().replace("__CLUSTER_DATA__", "null") + + server = ThreadingHTTPServer(("127.0.0.1", args.port), handler_for(viewer, page)) + where = f"http://127.0.0.1:{args.port}/" + print(f"{viewer.totals[0]:,} clusters over {viewer.totals[1]:,} records at {where}") + if args.open: + webbrowser.open(where) + try: + server.serve_forever() + except KeyboardInterrupt: + print("\nstopped") + + +if __name__ == "__main__": + sys.exit(main()) diff --git a/tools/cluster_view.py b/tools/cluster_view.py new file mode 100755 index 0000000..26cb7d3 --- /dev/null +++ b/tools/cluster_view.py @@ -0,0 +1,309 @@ +#!/usr/bin/env python3 +# Copyright 2026 Mathieu Fourment +# SPDX-License-Identifier: MIT +"""Build a self-contained HTML viewer for the clusters a cpplink run produced. + + cpplink cluster --schema s.json --edges edges --out clusters.csv data.parquet + tools/cluster_view.py --schema s.json --clusters clusters.csv \ + --edges edges --out clusters.html data.parquet + +Reads the cluster assignment, pulls each member's column values back out of the +parquet, and writes one HTML file holding the clusters it selected. Open it in a +browser: the left pane lists clusters, the right one shows the members side by +side with every disagreeing cell highlighted, so what a cluster has in common is +the part that is not highlighted. + +Nothing here holds a row per record. The clusters are chosen in Arrow, the +parquet is walked one row group at a time and only the row groups holding a +chosen member are materialised, and the edge shards are filtered in chunks, so +the resident cost follows the clusters embedded rather than the file's size. +""" + +import argparse +import datetime +import json +import os +import sys + +import numpy as np +import pyarrow as pa +import pyarrow.compute as pc +import pyarrow.csv as pacsv +import pyarrow.parquet as pq + +EDGE_MAGIC = b"CPPLNKE1" +EDGE_DTYPE = np.dtype([("a", " weight for the shard edges whose both ends are wanted. + + The shards are the run's whole output, which is the one thing here that can + be larger than memory, so they are read in fixed chunks and filtered down to + the embedded rows before anything is kept. + """ + found = {} + names = sorted(n for n in os.listdir(directory) + if n.startswith("shard-") and n.endswith(".bin")) + wanted = np.sort(np.asarray(wanted_rows, dtype=np.uint32)) + read = 0 + for name in names: + with open(os.path.join(directory, name), "rb") as handle: + if handle.read(len(EDGE_MAGIC)) != EDGE_MAGIC: + raise SystemExit(f"{name}: not a cpplink edge shard") + while True: + raw = handle.read(EDGE_CHUNK * EDGE_DTYPE.itemsize) + if not raw: + break + block = np.frombuffer(raw, dtype=EDGE_DTYPE) + read += block.size + keep = block[np.isin(block["a"], wanted) & np.isin(block["b"], wanted)] + for edge in keep: + found[(int(edge["a"]), int(edge["b"]))] = float(edge["w"]) + return found, read + + +def read_truth(path, wanted_ids): + """Group number per wanted id, from the known pairs. + + Being a duplicate is transitive, so the whole file has to be walked: two + embedded records can be the same entity through a third the page never + shows. It is walked as integers, though -- both columns are dictionary + encoded into one dictionary and the union-find runs over that -- so a truth + file of any size costs one entry per distinct id rather than one string. + """ + with open(path) as handle: + first = handle.readline().rstrip("\n").split(",") + header = first[:1] in (["id_a"], ["unique_id"], ["a"]) + table = pacsv.read_csv( + path, + read_options=pacsv.ReadOptions(column_names=["a", "b"], + skip_rows=1 if header else 0), + convert_options=pacsv.ConvertOptions( + column_types={"a": pa.string(), "b": pa.string()})) + left = table.column("a").combine_chunks() + right = table.column("b").combine_chunks() + pairs = len(left) + codes = pa.concat_arrays([left, right]).dictionary_encode() + index = codes.indices.to_numpy(zero_copy_only=False) + + parent = list(range(len(codes.dictionary))) + + def find(x): + while parent[x] != x: + parent[x] = parent[parent[x]] + x = parent[x] + return x + + for a, b in zip(index[:pairs].tolist(), index[pairs:].tolist()): + ra, rb = find(a), find(b) + if ra != rb: + parent[ra] = rb + + want = pa.array(sorted(wanted_ids), type=pa.string()) + at = pc.index_in(want, value_set=codes.dictionary).to_pylist() + return {uid: find(i) + for uid, i in zip(want.to_pylist(), at) if i is not None} + + +def main(): + ap = argparse.ArgumentParser(description=__doc__, + formatter_class=argparse.RawDescriptionHelpFormatter) + ap.add_argument("data", nargs="+", help="the parquet file(s) the run read, in order") + ap.add_argument("--schema", required=True) + ap.add_argument("--clusters", required=True, help="clusters.csv from cpplink cluster") + ap.add_argument("--edges", help="the shard directory, for within-cluster weights") + ap.add_argument("--truth", help="known pairs, to colour members by true entity") + ap.add_argument("--out", default="clusters.html") + ap.add_argument("--limit", type=int, default=400, help="clusters to embed (0 = all)") + ap.add_argument("--min-size", type=int, default=2) + ap.add_argument("--max-rows", type=int, default=200, + help="members embedded per cluster; the rest are counted only") + ap.add_argument("--max-size", type=int, default=0, help="0 = no ceiling") + ap.add_argument("--sort", default="size", choices=["size", "id", "random"], + help="which clusters to embed when --limit cuts the list") + ap.add_argument("--seed", type=int, default=1) + ap.add_argument("--id", action="append", default=[], + help="always embed this cluster id; repeatable") + args = ap.parse_args() + + id_column, columns = schema_columns(args.schema) + + assignment = read_assignment(args.clusters) + chosen, total_clusters = choose_clusters(assignment, args) + members = read_members(assignment, chosen) + del assignment + + sizes = {cluster: len(uids) for cluster, uids in members.items()} + if args.max_rows: + members = {c: uids[: args.max_rows] for c, uids in members.items()} + wanted_ids = {uid for cluster in members.values() for uid in cluster} + table, rows_by_id, records, groups, touched = read_records( + args.data, id_column, columns, wanted_ids) + if table is None: + raise SystemExit("none of the chosen clusters' ids are in the parquet") + values = {name: table.column(name).to_pylist() for name in columns} + by_id = {uid: i for i, uid in enumerate(table.column(id_column).to_pylist())} + + truth = read_truth(args.truth, wanted_ids) if args.truth else None + edges, edges_read = ({}, 0) + if args.edges: + edges, edges_read = read_edges( + args.edges, [rows_by_id[u] for u in wanted_ids if u in rows_by_id]) + + payload_clusters = [] + for cluster in chosen: + uids = [u for u in members.get(cluster, []) if u in by_id] + if len(uids) < args.min_size: + continue + rows, groups_seen = [], {} + for uid in uids: + at = by_id[uid] + row = {"id": uid, "v": [cell(values[name][at]) for name in columns]} + if truth is not None: + group = truth.get(uid, "\x00" + uid) + row["t"] = groups_seen.setdefault(group, len(groups_seen)) + rows.append(row) + entry = {"id": cluster, "rows": rows} + if sizes[cluster] > len(rows): + entry["more"] = sizes[cluster] - len(rows) + if args.edges: + position = {rows_by_id[u]: i for i, u in enumerate(uids) if u in rows_by_id} + found = [] + for (a, b), weight in edges.items(): + if a in position and b in position: + i, j = position[a], position[b] + found.append([min(i, j), max(i, j), round(weight, 3)]) + entry["edges"] = sorted(found) + payload_clusters.append(entry) + + payload = { + "source": ", ".join(os.path.basename(p) for p in args.data), + "generated": datetime.datetime.now().isoformat(timespec="seconds"), + "columns": columns, + "totals": {"clusters": total_clusters, "shown": len(payload_clusters)}, + "clusters": payload_clusters, + } + + here = os.path.dirname(os.path.abspath(__file__)) + with open(os.path.join(here, "cluster_view_template.html")) as handle: + page = handle.read() + blob = json.dumps(payload, separators=(",", ":")).replace(" {args.out} ({size:.1f} MB)") + print(f"{len(wanted_ids):,} records read from {touched:,} of {groups:,} row " + f"groups over {records:,} rows") + if args.edges: + print(f"{len(edges):,} within-cluster edges carried, {edges_read:,} scanned") + + +if __name__ == "__main__": + sys.exit(main()) diff --git a/tools/cluster_view_template.html b/tools/cluster_view_template.html new file mode 100644 index 0000000..e94739b --- /dev/null +++ b/tools/cluster_view_template.html @@ -0,0 +1,521 @@ + + + + + +cpplink clusters + + + +
+

cpplink clusters

+ + + j k move · / search + +
+
+
+
+ +
+ + +
+
+ + +
+
+
+
+
Pick a cluster.
+
+ + + + From 19b3473dc2f3a066e77f0698404d133098a157a3 Mon Sep 17 00:00:00 2001 From: 4ment Date: Fri, 11 Sep 2026 07:25:10 +1000 Subject: [PATCH 02/10] Viewers use predictions argument to parse prediction file --- tools/cluster_server.py | 124 +++++++++++++++++++++++---------- tools/cluster_view.py | 147 +++++++++++++++++++++++++++++++++------- 2 files changed, 211 insertions(+), 60 deletions(-) diff --git a/tools/cluster_server.py b/tools/cluster_server.py index 21f9450..5f9d00d 100755 --- a/tools/cluster_server.py +++ b/tools/cluster_server.py @@ -4,11 +4,15 @@ """Serve the cluster viewer over a DuckDB cache, for runs too large to embed. tools/cluster_server.py --schema s.json --clusters clusters.csv \ - --edges edges --truth truth.csv data.parquet + --predictions predictions.parquet --truth truth.csv data.parquet Same page as cluster_view.py, but the clusters are queried rather than written into the file, so every cluster of the run is reachable instead of a sample. +`--clusters` and `--predictions` each name one file, csv or parquet, picked by +the extension; `--predictions` also takes the shard directory, which names +records by row and so costs the pass that names them. + The cache is built once and reused until an input changes. It holds only what the viewer can ever show -- the clustered records, never the singletons -- so it is a fraction of the dataset: at 20M records with 2.9M of them clustered, @@ -36,13 +40,28 @@ read_truth, schema_columns) EDGE_CHUNK = 1 << 20 -CACHE_VERSION = 2 +CACHE_VERSION = 3 def quoted(name): return '"' + name.replace('"', '""') + '"' +def scan(path, types): + """The duckdb table function that reads one csv or parquet file. + + Both of a run's outputs are one file whose extension picks the format, so + that is what picks the reader here. The ids are pinned to VARCHAR rather than + sniffed, because a file of numeric ids would otherwise read as integers and + join against nothing. + """ + path = os.path.abspath(path) + if path.endswith(".parquet"): + return f"read_parquet('{path}')" + pinned = ", ".join(f"'{name}': 'VARCHAR'" for name in types) + return f"read_csv('{path}', header = true, types = {{{pinned}}})" + + def fingerprint(args, columns): """What the cache was built from, so a changed input rebuilds it.""" parts = [CACHE_VERSION, columns, args.max_rows] @@ -50,15 +69,20 @@ def fingerprint(args, columns): if path: parts.append([os.path.abspath(path), os.path.getmtime(path), os.path.getsize(path)]) - if args.edges: - for name in sorted(os.listdir(args.edges)): - if name.endswith(".bin"): - full = os.path.join(args.edges, name) - parts.append([full, os.path.getmtime(full), os.path.getsize(full)]) + if args.predictions: + if os.path.isdir(args.predictions): + for name in sorted(os.listdir(args.predictions)): + if name.endswith(".bin"): + full = os.path.join(args.predictions, name) + parts.append([full, os.path.getmtime(full), os.path.getsize(full)]) + else: + parts.append([os.path.abspath(args.predictions), + os.path.getmtime(args.predictions), + os.path.getsize(args.predictions)]) return json.dumps(parts, sort_keys=True) -def load_edges(conn, directory): +def load_shard_edges(conn, directory): """Every shard edge into `raw_edges`, a chunk at a time. The shards hold row indices rather than ids, which is why this is worth @@ -92,10 +116,10 @@ def build(conn, args, id_column, columns, say): say("reading the cluster assignment") conn.execute(f""" CREATE TABLE members AS - SELECT unique_id AS uid, cluster_id, cluster_size AS size - FROM read_csv('{args.clusters}', header = true, - columns = {{'unique_id': 'VARCHAR', 'cluster_id': 'VARCHAR', - 'cluster_size': 'BIGINT'}}) + SELECT CAST(unique_id AS VARCHAR) AS uid, + CAST(cluster_id AS VARCHAR) AS cluster_id, + CAST(cluster_size AS BIGINT) AS size + FROM {scan(args.clusters, ("unique_id", "cluster_id"))} WHERE cluster_size >= {args.min_size} """) @@ -108,31 +132,52 @@ def build(conn, args, id_column, columns, say): JOIN members m ON m.uid = d.{quoted(id_column)} """) - if args.edges: - say("naming the edges") - # A row index in a shard is a position in the inputs read in order, and - # file_row_number is that position inside one file, so the offsets that - # separate the inputs are the only thing to carry across. - conn.execute("CREATE TABLE rowmap (rid BIGINT, uid VARCHAR)") - offset = 0 - for path in files: + if args.predictions: + if os.path.isdir(args.predictions): + say("naming the predictions") + # A row index in a shard is a position in the inputs read in order, + # and file_row_number is that position inside one file, so the + # offsets that separate the inputs are the only thing to carry. + conn.execute("CREATE TABLE rowmap (rid BIGINT, uid VARCHAR)") + offset = 0 + for path in files: + conn.execute(f""" + INSERT INTO rowmap + SELECT {offset} + file_row_number, {quoted(id_column)} + FROM read_parquet('{path}', file_row_number = true) + """) + offset += pq.ParquetFile(path).metadata.num_rows + load_shard_edges(conn, args.predictions) + conn.execute(""" + CREATE TABLE named_edges AS + SELECT ra.uid AS a, rb.uid AS b, e.w AS weight + FROM raw_edges e + JOIN rowmap ra ON ra.rid = e.a + JOIN rowmap rb ON rb.rid = e.b + """) + conn.execute("DROP TABLE raw_edges") + conn.execute("DROP TABLE rowmap") + else: + # A merged file already names records by unique_id, which is the + # whole of what the rowmap above exists to recover. + say("reading the predictions") conn.execute(f""" - INSERT INTO rowmap - SELECT {offset} + file_row_number, {quoted(id_column)} - FROM read_parquet('{path}', file_row_number = true) + CREATE TABLE named_edges AS + SELECT CAST(id_a AS VARCHAR) AS a, CAST(id_b AS VARCHAR) AS b, + match_weight AS weight + FROM {scan(args.predictions, ("id_a", "id_b"))} """) - offset += pq.ParquetFile(path).metadata.num_rows - load_edges(conn, args.edges) + # Both ends, because clustering at a threshold above the one the run + # wrote at leaves predictions that cross two clusters or land outside + # every one, and the page shows a weight only inside a cluster. conn.execute(""" CREATE TABLE edges AS - SELECT m.cluster_id, ra.uid AS a, rb.uid AS b, e.w AS weight - FROM raw_edges e - JOIN rowmap ra ON ra.rid = e.a - JOIN rowmap rb ON rb.rid = e.b - JOIN members m ON m.uid = ra.uid + SELECT m.cluster_id, e.a, e.b, e.weight + FROM named_edges e + JOIN members m ON m.uid = e.a + JOIN members mb ON mb.uid = e.b AND mb.cluster_id = m.cluster_id """) - conn.execute("DROP TABLE raw_edges") - conn.execute("DROP TABLE rowmap") + conn.execute("DROP TABLE named_edges") conn.execute("CREATE INDEX edges_by_cluster ON edges (cluster_id)") else: conn.execute("CREATE TABLE edges (cluster_id VARCHAR, a VARCHAR, " @@ -326,14 +371,20 @@ def main(): formatter_class=argparse.RawDescriptionHelpFormatter) ap.add_argument("data", nargs="+", help="the parquet file(s) the run read, in order") ap.add_argument("--schema", required=True) - ap.add_argument("--clusters", required=True, help="clusters.csv from cpplink cluster") - ap.add_argument("--edges", help="the shard directory, for within-cluster weights") + ap.add_argument("--clusters", required=True, + help="the csv or parquet cpplink cluster wrote") + ap.add_argument("--predictions", "--edges", dest="predictions", + help="the run's predictions -- one csv or parquet file, or the " + "shard directory -- for within-cluster weights") ap.add_argument("--truth", help="known pairs, to colour members by true entity") ap.add_argument("--cache", default="cluster_view.duckdb") ap.add_argument("--rebuild", action="store_true", help="rebuild the cache first") ap.add_argument("--min-size", type=int, default=2) ap.add_argument("--max-rows", type=int, default=200, help="members shown per cluster; the rest are counted only") + ap.add_argument("--memory", default="4GB", + help="what DuckDB may use while building the cache") + ap.add_argument("--threads", type=int, default=0, help="0 = DuckDB's own default") ap.add_argument("--port", type=int, default=8770) ap.add_argument("--open", action="store_true", help="open a browser on it") args = ap.parse_args() @@ -355,6 +406,10 @@ def main(): os.remove(args.cache) start = time.time() conn = duckdb.connect(args.cache) + conn.execute(f"SET memory_limit = '{args.memory}'") + conn.execute("SET preserve_insertion_order = false") + if args.threads: + conn.execute(f"SET threads = {args.threads}") def say(what): print(f" {time.time() - start:6.1f}s {what}", flush=True) @@ -368,6 +423,7 @@ def say(what): print(f" {time.time() - start:6.1f}s done, {size:.0f} MB") conn = duckdb.connect(args.cache, read_only=True) + conn.execute(f"SET memory_limit = '{args.memory}'") viewer = Viewer(conn, columns, ", ".join(os.path.basename(p) for p in args.data), args.max_rows, bool(args.truth)) diff --git a/tools/cluster_view.py b/tools/cluster_view.py index 26cb7d3..64f8935 100755 --- a/tools/cluster_view.py +++ b/tools/cluster_view.py @@ -3,9 +3,11 @@ # SPDX-License-Identifier: MIT """Build a self-contained HTML viewer for the clusters a cpplink run produced. - cpplink cluster --schema s.json --edges edges --out clusters.csv data.parquet + cpplink predict --schema s.json --model m.json --out predictions.parquet data.parquet + cpplink cluster --schema s.json --predictions predictions.parquet \ + --out clusters.csv data.parquet tools/cluster_view.py --schema s.json --clusters clusters.csv \ - --edges edges --out clusters.html data.parquet + --predictions predictions.parquet --out clusters.html data.parquet Reads the cluster assignment, pulls each member's column values back out of the parquet, and writes one HTML file holding the clusters it selected. Open it in a @@ -13,9 +15,14 @@ side with every disagreeing cell highlighted, so what a cluster has in common is the part that is not highlighted. +A run ends with single files, so `--clusters` and `--predictions` each name one +file and its extension picks csv or parquet. `--predictions` still takes the +shard directory too, which names records by row rather than by `unique_id` and so +needs the row index a merged file makes unnecessary. + Nothing here holds a row per record. The clusters are chosen in Arrow, the parquet is walked one row group at a time and only the row groups holding a -chosen member are materialised, and the edge shards are filtered in chunks, so +chosen member are materialised, and the predictions are filtered in batches, so the resident cost follows the clusters embedded rather than the file's size. """ @@ -34,6 +41,11 @@ EDGE_MAGIC = b"CPPLNKE1" EDGE_DTYPE = np.dtype([("a", " weight for the shard edges whose both ends are wanted. +def read_shard_predictions(directory, rows_by_id): + """(id_a, id_b) -> weight for the shard edges whose both ends are wanted. The shards are the run's whole output, which is the one thing here that can be larger than memory, so they are read in fixed chunks and filtered down to - the embedded rows before anything is kept. + the embedded rows before anything is kept. A shard names records by row, so + the row index is what the filter runs over and the ids go back on afterwards. """ found = {} names = sorted(n for n in os.listdir(directory) if n.startswith("shard-") and n.endswith(".bin")) - wanted = np.sort(np.asarray(wanted_rows, dtype=np.uint32)) + if not names: + raise SystemExit(f"{directory}: no shard-*.bin files") + id_by_row = {row: uid for uid, row in rows_by_id.items()} + wanted = np.sort(np.asarray(list(id_by_row), dtype=np.uint32)) read = 0 for name in names: with open(os.path.join(directory, name), "rb") as handle: if handle.read(len(EDGE_MAGIC)) != EDGE_MAGIC: - raise SystemExit(f"{name}: not a cpplink edge shard") + raise SystemExit(f"{name}: not a cpplink prediction shard") while True: raw = handle.read(EDGE_CHUNK * EDGE_DTYPE.itemsize) if not raw: @@ -163,10 +195,59 @@ def read_edges(directory, wanted_rows): read += block.size keep = block[np.isin(block["a"], wanted) & np.isin(block["b"], wanted)] for edge in keep: - found[(int(edge["a"]), int(edge["b"]))] = float(edge["w"]) + found[(id_by_row[int(edge["a"])], + id_by_row[int(edge["b"])])] = float(edge["w"]) + return found, read + + +def prediction_batches(path): + """One merged prediction file, a batch of rows at a time.""" + if path.endswith(".parquet"): + handle = pq.ParquetFile(path) + held = handle.schema_arrow.names + if any(name not in held for name in PREDICTION_COLUMNS): + raise SystemExit(f"{path}: not a cpplink prediction file") + for batch in handle.iter_batches(batch_size=PREDICTION_BATCH, + columns=PREDICTION_COLUMNS): + yield pa.Table.from_batches([batch]) + return + reader = pacsv.open_csv( + path, convert_options=pacsv.ConvertOptions(column_types=PREDICTION_IDS)) + if any(name not in reader.schema.names for name in PREDICTION_COLUMNS): + raise SystemExit(f"{path}: not a cpplink prediction file") + for batch in reader: + yield pa.Table.from_batches([batch]) + + +def read_file_predictions(path, wanted_ids): + """(id_a, id_b) -> weight for the predictions naming two wanted records. + + A merged file names records by `unique_id`, so there is no row index to carry + and no dependence on which file a record was read from. It is still the run's + whole output, so it is read a batch at a time and cut to the embedded records + before anything is kept. + """ + wanted = pa.array(sorted(wanted_ids), type=pa.string()) + found, read = {}, 0 + for table in prediction_batches(path): + read += table.num_rows + table = as_strings(table, ("id_a", "id_b")) + keep = pc.and_(pc.is_in(table["id_a"], value_set=wanted), + pc.is_in(table["id_b"], value_set=wanted)) + table = table.filter(keep) + for a, b, weight in zip(table["id_a"].to_pylist(), table["id_b"].to_pylist(), + table["match_weight"].to_pylist()): + found[(a, b)] = float(weight) return found, read +def read_predictions(path, wanted_ids, rows_by_id): + """The run's predictions, from the merged file or from the shard directory.""" + if os.path.isdir(path): + return read_shard_predictions(path, rows_by_id) + return read_file_predictions(path, wanted_ids) + + def read_truth(path, wanted_ids): """Group number per wanted id, from the known pairs. @@ -215,8 +296,11 @@ def main(): formatter_class=argparse.RawDescriptionHelpFormatter) ap.add_argument("data", nargs="+", help="the parquet file(s) the run read, in order") ap.add_argument("--schema", required=True) - ap.add_argument("--clusters", required=True, help="clusters.csv from cpplink cluster") - ap.add_argument("--edges", help="the shard directory, for within-cluster weights") + ap.add_argument("--clusters", required=True, + help="the csv or parquet cpplink cluster wrote") + ap.add_argument("--predictions", "--edges", dest="predictions", + help="the run's predictions -- one csv or parquet file, or the " + "shard directory -- for within-cluster weights") ap.add_argument("--truth", help="known pairs, to colour members by true entity") ap.add_argument("--out", default="clusters.html") ap.add_argument("--limit", type=int, default=400, help="clusters to embed (0 = all)") @@ -250,10 +334,20 @@ def main(): by_id = {uid: i for i, uid in enumerate(table.column(id_column).to_pylist())} truth = read_truth(args.truth, wanted_ids) if args.truth else None - edges, edges_read = ({}, 0) - if args.edges: - edges, edges_read = read_edges( - args.edges, [rows_by_id[u] for u in wanted_ids if u in rows_by_id]) + predictions, predictions_read = ({}, 0) + by_cluster = {} + if args.predictions: + predictions, predictions_read = read_predictions( + args.predictions, wanted_ids, + {u: rows_by_id[u] for u in wanted_ids if u in rows_by_id}) + # A record is in one cluster, so grouping the predictions once is what + # keeps the loop below linear in them rather than one pass per cluster. + cluster_of = {uid: cluster + for cluster, uids in members.items() for uid in uids} + for (a, b), weight in predictions.items(): + cluster = cluster_of.get(a) + if cluster is not None and cluster == cluster_of.get(b): + by_cluster.setdefault(cluster, []).append((a, b, weight)) payload_clusters = [] for cluster in chosen: @@ -271,10 +365,10 @@ def main(): entry = {"id": cluster, "rows": rows} if sizes[cluster] > len(rows): entry["more"] = sizes[cluster] - len(rows) - if args.edges: - position = {rows_by_id[u]: i for i, u in enumerate(uids) if u in rows_by_id} + if args.predictions: + position = {uid: i for i, uid in enumerate(uids)} found = [] - for (a, b), weight in edges.items(): + for a, b, weight in by_cluster.get(cluster, []): if a in position and b in position: i, j = position[a], position[b] found.append([min(i, j), max(i, j), round(weight, 3)]) @@ -301,8 +395,9 @@ def main(): f"-> {args.out} ({size:.1f} MB)") print(f"{len(wanted_ids):,} records read from {touched:,} of {groups:,} row " f"groups over {records:,} rows") - if args.edges: - print(f"{len(edges):,} within-cluster edges carried, {edges_read:,} scanned") + if args.predictions: + print(f"{len(predictions):,} within-cluster predictions carried, " + f"{predictions_read:,} scanned") if __name__ == "__main__": From daa85cea76ea48780bb86a987592916d5622d2ea Mon Sep 17 00:00:00 2001 From: 4ment Date: Fri, 11 Sep 2026 07:42:03 +1000 Subject: [PATCH 03/10] Draw a network --- tools/cluster_view.py | 5 + tools/cluster_view_template.html | 202 ++++++++++++++++++++++++++++++- 2 files changed, 205 insertions(+), 2 deletions(-) diff --git a/tools/cluster_view.py b/tools/cluster_view.py index 64f8935..c5e1573 100755 --- a/tools/cluster_view.py +++ b/tools/cluster_view.py @@ -15,6 +15,11 @@ side with every disagreeing cell highlighted, so what a cluster has in common is the part that is not highlighted. +The `network` checkbox in the header draws the selected cluster's predictions as +a graph, which is where a chain shows itself as a chain. It is off by default and +the choice is remembered, because the layout is the one quadratic thing the page +does and most clusters are read without it. + A run ends with single files, so `--clusters` and `--predictions` each name one file and its extension picks csv or parquet. `--predictions` still takes the shard directory too, which names records by row rather than by `unique_id` and so diff --git a/tools/cluster_view_template.html b/tools/cluster_view_template.html index e94739b..cb1429b 100644 --- a/tools/cluster_view_template.html +++ b/tools/cluster_view_template.html @@ -151,6 +151,30 @@ .matrix td.none { color: var(--muted); opacity: 0.45; } .matrix td.weak { background: var(--warn-soft); color: var(--warn); } .matrix th.rowhead { text-align: left; font-family: var(--mono); text-transform: none; letter-spacing: 0; } +.note { color: var(--muted); font-size: 12px; } +.net { + width: 100%; height: auto; max-height: 55vh; display: block; + background: var(--panel); border: 1px solid var(--line); border-radius: 8px; +} +.net line { stroke: var(--muted); stroke-linecap: round; } +.net line.weak { stroke: var(--warn); } +.net line.hit { stroke: transparent; stroke-width: 9; } +.net g.link { cursor: pointer; } +.net g.link:hover line.wire { stroke: var(--ink); opacity: 1; } +.net g.link.on line.wire { stroke: var(--accent); stroke-width: 3; opacity: 1; } +.net circle { stroke: var(--panel); stroke-width: 1.5; cursor: pointer; } +.net text { font-family: var(--mono); text-anchor: middle; pointer-events: none; } +.net text.tag { font-size: 5.5px; fill: var(--muted); } +.net text.num { font-size: 5px; fill: var(--panel); font-weight: 600; } +.net g.node:hover circle { stroke: var(--ink); stroke-width: 2; } +tr.lit td, tr.pin td { background: var(--accent-soft); } +tr.lit td.odd, tr.pin td.odd { background: var(--warn-soft); } +tr.pin td { box-shadow: inset 0 1px 0 var(--accent), inset 0 -1px 0 var(--accent); } +.toggle { + font-size: 12px; color: var(--muted); display: flex; gap: 5px; + align-items: center; cursor: pointer; user-select: none; +} +.toggle input { margin: 0; accent-color: var(--accent); } .empty-state { color: var(--muted); padding: 40px 0; text-align: center; } kbd { font: inherit; font-size: 11px; font-family: var(--mono); border: 1px solid var(--line); @@ -164,6 +188,9 @@

cpplink clusters

j k move · / search +
@@ -247,6 +274,141 @@

cpplink clusters

}); } +// The network of one cluster. A force-directed layout is the expensive part of +// this page -- it is O(n^2) an iteration, over rows the table already holds -- +// so it runs only when the checkbox asks for it, only up to NET_MAX members, +// and its result is cached on the cluster the way the rest of the detail is. +const NET_MAX = 250; // members above which the layout is not worth it +const NET_W = 640, NET_H = 360; + +// Fruchterman-Reingold in a unit square, seeded on a phyllotaxis spiral so the +// same cluster always draws the same picture and no two nodes start on top of +// one another. Gravity is what keeps a cluster held together only transitively +// from flying apart into components the pane cannot fit. +function layout(n, links) { + const x = new Float64Array(n), y = new Float64Array(n); + for (let i = 0; i < n; i++) { + const a = i * 2.39996323, r = 0.45 * Math.sqrt((i + 0.5) / n); + x[i] = 0.5 + r * Math.cos(a); + y[i] = 0.5 + r * Math.sin(a); + } + const k = Math.sqrt(1 / n); // the edge length the layout aims at + const dx = new Float64Array(n), dy = new Float64Array(n); + const steps = n > 120 ? 160 : 300; + for (let step = 0; step < steps; step++) { + dx.fill(0); dy.fill(0); + for (let i = 0; i < n; i++) { + for (let j = i + 1; j < n; j++) { + const ex = x[i] - x[j], ey = y[i] - y[j]; + const d2 = ex * ex + ey * ey || 1e-9; + const f = k * k / d2; // k^2/d, over d again for the direction + dx[i] += ex * f; dy[i] += ey * f; + dx[j] -= ex * f; dy[j] -= ey * f; + } + } + for (const l of links) { + const ex = x[l.i] - x[l.j], ey = y[l.i] - y[l.j]; + const d = Math.hypot(ex, ey) || 1e-9; + const f = d / k; // d^2/k, over d for the direction + dx[l.i] -= ex * f; dy[l.i] -= ey * f; + dx[l.j] += ex * f; dy[l.j] += ey * f; + } + const t = 0.1 * (1 - step / steps) + 0.0005; + for (let i = 0; i < n; i++) { + dx[i] += (0.5 - x[i]) * 0.08 * k; + dy[i] += (0.5 - y[i]) * 0.08 * k; + const d = Math.hypot(dx[i], dy[i]) || 1e-9; + const move = Math.min(d, t) / d; + x[i] += dx[i] * move; y[i] += dy[i] * move; + } + } + // Fit whatever it settled on to the viewBox, keeping the aspect ratio. + let x0 = Infinity, y0 = Infinity, x1 = -Infinity, y1 = -Infinity; + for (let i = 0; i < n; i++) { + x0 = Math.min(x0, x[i]); x1 = Math.max(x1, x[i]); + y0 = Math.min(y0, y[i]); y1 = Math.max(y1, y[i]); + } + const pad = 22; + const scale = Math.min((NET_W - 2 * pad) / (x1 - x0 || 1), + (NET_H - 2 * pad) / (y1 - y0 || 1)); + const out = []; + for (let i = 0; i < n; i++) { + out.push([NET_W / 2 + (x[i] - (x0 + x1) / 2) * scale, + NET_H / 2 + (y[i] - (y0 + y1) / 2) * scale]); + } + return out; +} + +function network(c) { + if (c.net !== undefined) return c.net; + const rows = c.rows, n = rows.length; + let body; + if (n < 2 || !c.edges) { + body = ''; + } else if (n > NET_MAX) { + body = `
${n.toLocaleString()} members is more than the ` + + `${NET_MAX} this draws; the weights are below.
`; + } else { + const at = new Map(rows.map((r, i) => [r.id, i])); + const links = []; + let lo = Infinity, hi = -Infinity; + for (const [a, b, w] of c.edges) { + const i = at.get(a), j = at.get(b); + if (i === undefined || j === undefined || i === j) continue; + links.push({ i, j, w }); + lo = Math.min(lo, w); hi = Math.max(hi, w); + } + const pos = layout(n, links); + const span = hi - lo || 1; + const lines = links.map(l => { + const [xa, ya] = pos[l.i], [xb, yb] = pos[l.j]; + const width = (1 + 2.2 * (l.w - lo) / span).toFixed(2); + const dim = (0.3 + 0.5 * (l.w - lo) / span).toFixed(2); + // The visible line is a hairline at the weakest weight, which is + // nothing to aim a mouse at, so an invisible fat one carries the + // clicking and the thin one carries the reading. + const ends = `x1="${xa.toFixed(1)}" y1="${ya.toFixed(1)}" ` + + `x2="${xb.toFixed(1)}" y2="${yb.toFixed(1)}"`; + return `` + + `${esc(rows[l.i].id)} — ${esc(rows[l.j].id)}: ` + + `${l.w.toFixed(1)} bits` + + `` + + ``; + }).join(''); + const truthy = rows[0].t !== undefined; + const nodes = rows.map((r, i) => { + const [px, py] = pos[i]; + const colour = truthy ? truthColors[(r.t || 0) % truthColors.length] + : 'var(--accent)'; + const tag = n <= 14 + ? `` + + `${esc(r.id.length > 16 ? r.id.slice(0, 15) + '…' : r.id)}` + : ''; + const num = n <= 100 + ? `` + + `${i + 1}` + : ''; + return `` + + `${i + 1}. ${esc(r.id)}${num}${tag}`; + }).join(''); + const pairs = n * (n - 1) / 2; + body = `${lines}${nodes}` + + `
` + + `${links.length.toLocaleString()} of ${pairs.toLocaleString()} pairs were ` + + `scored above the threshold; a pair with no line is in this cluster ` + + `through the others. Line weight is the match weight, the thinnest ` + + `${links.length ? `(${lo.toFixed(1)} bits) ` : ''}marked. ` + + `Click a line to hold its two records lit in the table above.
`; + } + c.net = body + ? `

Network

${body}
` + : ''; + return c.net; +} + function summarise(cluster) { const n = cluster.rows.length; const stats = statsOf(cluster.rows, cluster.dist); @@ -403,7 +565,7 @@

cpplink clusters

return `${esc(name)}${d}`; })).join(''); - const body = c.rows.map(r => { + const body = c.rows.map((r, i) => { const cells = COLS.map((_, i) => { const s = view.stats[i], v = r.v[i]; if (v === null || v === '') return '—'; @@ -413,7 +575,7 @@

cpplink clusters

const colour = truthColors[(r.t || 0) % truthColors.length]; const t = truthy ? `${r.t + 1}` : ''; - return `${esc(r.id)}${t}${cells}`; + return `${esc(r.id)}${t}${cells}`; }).join(''); let matrix = ''; @@ -449,7 +611,35 @@

${esc(c.id)}

What they have in common

${chips}

Records

${head}${body}
+ ${showNet ? network(c) : ''} ${matrix}`; + // A node is a record and a line is a pair of them, so hovering either lights + // the rows it names; a line is also clickable, and that one holds. + const table = d.querySelector('tbody'); + const lit = (rows, on) => rows.forEach(r => r.classList.toggle('lit', on)); + for (const node of d.querySelectorAll('.net g.node')) { + const row = table.children[+node.dataset.i]; + node.onmouseenter = () => lit([row], true); + node.onmouseleave = () => lit([row], false); + node.onclick = () => row.scrollIntoView({ block: 'nearest' }); + } + let held = null; + for (const link of d.querySelectorAll('.net g.link')) { + const ends = [table.children[+link.dataset.i], table.children[+link.dataset.j]]; + link.onmouseenter = () => lit(ends, true); + link.onmouseleave = () => lit(ends, false); + link.onclick = () => { + if (held) { + held.link.classList.remove('on'); + held.ends.forEach(r => r.classList.remove('pin')); + } + if (held && held.link === link) { held = null; return; } + link.classList.add('on'); + ends.forEach(r => r.classList.add('pin')); + held = { link, ends }; + ends[0].scrollIntoView({ block: 'nearest' }); + }; + } d.scrollTop = 0; } @@ -485,6 +675,14 @@

${esc(c.id)}

await loadPage(false); renderList(); }); +let showNet = false; +try { showNet = localStorage.getItem('cpplink.net') === '1'; } catch (e) {} +$('#net').checked = showNet; +$('#net').addEventListener('change', () => { + showNet = $('#net').checked; + try { localStorage.setItem('cpplink.net', showNet ? '1' : '0'); } catch (e) {} + if (current) select(current.id); +}); $('#theme').onclick = () => { const now = document.documentElement.getAttribute('data-theme'); const dark = now ? now === 'dark' : matchMedia('(prefers-color-scheme: dark)').matches; From 7603aa11ac28de599da0aa0da6aeb7adedf870d7 Mon Sep 17 00:00:00 2001 From: 4ment Date: Sun, 13 Sep 2026 12:10:02 +1000 Subject: [PATCH 04/10] Add support for pairs mode in cluster viewer - Introduced a new mode for viewing pairs alongside clusters. - Updated the UI to include buttons for switching between clusters and pairs. - Modified data loading and rendering logic to handle pairs, including sorting and filtering options. - Enhanced detail view to display pair-specific information and waterfall visualizations. - Updated event handlers to manage interactions for pairs, including selection and navigation. - Adjusted data structures to accommodate pairs and their relationships. --- docs/commands/explain.md | 46 ++- src/cpplink/app.cpp | 219 ++++++++++--- src/cpplink/cluster.cpp | 33 +- src/cpplink/explain.cpp | 204 ++++++++---- src/cpplink/explain.hpp | 61 +++- src/cpplink/id_index.hpp | 47 +++ src/cpplink/waterfall.cpp | 479 +++++++++++++++++++++++++++ src/cpplink/waterfall.hpp | 65 ++++ tests/explain_batch_test.cpp | 390 ++++++++++++++++++++++ tests/explain_test.cpp | 74 ++++- tools/cluster_server.py | 204 +++++++++--- tools/cluster_view.py | 204 +++++++++++- tools/cluster_view_template.html | 540 ++++++++++++++++++++++++++----- 13 files changed, 2291 insertions(+), 275 deletions(-) create mode 100644 src/cpplink/id_index.hpp create mode 100644 src/cpplink/waterfall.cpp create mode 100644 src/cpplink/waterfall.hpp create mode 100644 tests/explain_batch_test.cpp diff --git a/docs/commands/explain.md b/docs/commands/explain.md index 4203c84..0b117e8 100644 --- a/docs/commands/explain.md +++ b/docs/commands/explain.md @@ -14,6 +14,10 @@ cpplink explain --schema --pair , cpplink explain --schema --rows , cpplink explain --schema --pair , \ --model [--threshold BITS] +cpplink explain --schema --model --json \ + --pairs +cpplink explain --schema --model \ + --predictions --out ``` | Option | Meaning | @@ -21,14 +25,22 @@ cpplink explain --schema --pair , \ | `--schema ` | required | | `--pair ,` | the two records by their `unique_id` values | | `--rows ,` | the two records by zero-based row index | +| `--pairs ` | many pairs, one `,` per line; `-` reads them from stdin | +| `--row-pairs ` | the same, with `,` row indices per line | +| `--predictions ` | every prediction in the csv or parquet `predict --out` wrote; needs `--out` | +| `--out ` | where to write the waterfalls, one wide row per prediction; `.csv` or `.parquet` | | `--model ` | optional; adds the score waterfall below the level table | | `--threshold BITS` | default 0; only affects the zone and the emit/drop line | | `--tf-damping F` | default 1.0; scale the term-frequency move, as in [`predict`](predict.md) | +| `--fuzzy-tf`, `--ball-budget N` | adjust fuzzy levels by neighbourhood mass, as in [`predict`](predict.md) | +| `--no-interactions` | score the plain model out of a file that carries two-way corrections | +| `--json` | one JSON object per pair instead of the text report; needs `--model` | | *(positional)* | required; the parquet file | -Give **exactly one** of `--pair` or `--rows`. The schema must declare `comparisons`; -`blocking` is not needed. Without `--model` there is no weight to explain, and the waterfall -is simply not printed. +Give **exactly one** of `--pair`, `--rows`, `--pairs`, `--row-pairs` or `--predictions`. The +schema must declare `comparisons`; `blocking` is not needed. Without `--model` there is no +weight to explain, and the waterfall is simply not printed. +Give `explain` the same scoring options the run used, or it explains a different weight than the one the run wrote. ## Example @@ -179,6 +191,34 @@ skips it without touching a term-frequency table. asserted over every pair in a fixture, because a report that plausibly explains a *different* calculation from the one that runs would be worse than no report at all. +## Many pairs from one load + +`--pairs` and `--row-pairs` answer one pair per input line, from a single load of the store. +That is the difference between a tool that can ask about a pair and one that cannot: at 20M records the load is the cost, and the pair is free. +With `-` the lines come from stdin and each answer is flushed as it is written, so a process can be held open and asked one pair at a time. +A line that names no record is reported on its own line and does not end the batch. + +`--json` writes the waterfall as one object per line, with the same numbers as the text report: `prior`, one entry per comparison in `steps` (`name`, `level`, `label`, the two `values`, `m`, `u`, `bits`, `tf`, `frequency` where a move was made, and the `running` total), the `interactions`, then `weight`, `probability`, the pattern `bracket`, `threshold`, `zone` and `emitted`. + +## Every prediction of a run, as a file + +`--predictions --out ` explains the whole prediction file `predict --out` wrote and writes one wide row per pair, csv or parquet by the extension: + +```text +id_a, id_b, gamma, prior, match_weight, match_probability, +bracket_low, bracket_high, threshold, zone, emitted, +_level, _bits, _tf, _frequency, one set per comparison +_x__bits one per two-way correction +``` + +`prior` plus every `_bits` and `_tf` plus every interaction column is `match_weight`, to the last bit the run wrote. +A level's label, `m` and `u` are the model's rather than the pair's, so the file does not repeat them; the model file carries them, indexed by `_level`. +Ids are resolved through a sorted index over the id column, 4 bytes a record, so a file of millions of predictions costs one load and one pass. +A prediction naming an id no record holds is skipped and counted, and so is one whose stored `gamma` the comparisons no longer produce, which means the schema changed after `predict` ran and the ledgers explain today's schema rather than the file's weights. + +This is what the cluster viewers in `tools/` draw from: `tools/cluster_view.py --waterfalls` embeds a ledger per prediction, and `tools/cluster_server.py --waterfalls` loads the file beside the predictions so a click is one lookup. +Nothing in either recomputes a bit of the weight, and neither runs this binary. + ## See also - [`predict`](predict.md) — the same scoring, over every candidate pair diff --git a/src/cpplink/app.cpp b/src/cpplink/app.cpp index d68cb4c..293d3a6 100644 --- a/src/cpplink/app.cpp +++ b/src/cpplink/app.cpp @@ -3,13 +3,17 @@ #include "cpplink/app.hpp" +#include #include #include #include #include #include #include +#include +#include +#include #include #include @@ -35,6 +39,7 @@ #include "cpplink/schema.hpp" #include "cpplink/score.hpp" #include "cpplink/simplify.hpp" +#include "cpplink/waterfall.hpp" namespace cpplink { @@ -109,9 +114,11 @@ void PrintUsage(std::ostream& out) { << " [--pair-cap N] [--threads N] [--seed N] [--json]\n" << " [--mode MODE] ...\n" << "cpplink explain --schema --pair ,\n" - << " [--rows ,] [--model ] " - "[--threshold BITS]\n" - << " [--tf-damping F] ...\n" + << " [--rows ,] [--pairs ] [--row-pairs ]\n" + << " [--predictions --out ]\n" + << " [--model ] [--threshold BITS] [--tf-damping F]\n" + << " [--fuzzy-tf] [--ball-budget N] [--no-interactions] [--json]\n" + << " ...\n" << "cpplink explain-blocking --schema [--count] " "[--mode MODE]\n" << " [--all-pairs] ...\n" @@ -608,30 +615,94 @@ bool SplitPair(const std::string& text, std::string* first, std::string* second) return true; } +void BuildBallTables(const ComparisonSet& comparisons, const RecordStore& store, + const BallOptions& options, BallTables* balls, std::ostream& out); + +// Resolves one `,` line of `explain` to two rows, by id or by row index. +bool ResolvePair(const RecordStore& store, const std::string& text, bool by_row, + uint64_t* row_a, uint64_t* row_b, std::string* error) { + std::string first; + std::string second; + if (!SplitPair(text, &first, &second)) { + *error = by_row ? "wants ," : "wants ,"; + return false; + } + if (by_row) { + char* end_a = nullptr; + char* end_b = nullptr; + *row_a = std::strtoull(first.c_str(), &end_a, 10); + *row_b = std::strtoull(second.c_str(), &end_b, 10); + if (*end_a != '\0' || *end_b != '\0') { + *error = "wants ,"; + return false; + } + if (*row_a >= store.NumRecords() || *row_b >= store.NumRecords()) { + *error = "row out of range; the file has " + + std::to_string(store.NumRecords()) + " records"; + return false; + } + return true; + } + if (!FindRowById(store, first, row_a)) { + *error = "no record with id '" + first + "'"; + return false; + } + if (!FindRowById(store, second, row_b)) { + *error = "no record with id '" + second + "'"; + return false; + } + return true; +} + int RunExplain(const std::vector& args, std::ostream& out, std::ostream& err) { std::string schema_path; std::vector data_paths; std::string pair; std::string rows; + std::string pairs_path; + bool pairs_by_row = false; std::string model_path; std::string value; ScoreOptions score; + BallOptions ball; + WaterfallOptions waterfalls; + bool fuzzy_tf = false; + bool as_json = false; for (size_t i = 0; i < args.size(); ++i) { if (args[i] == "--schema") { if (!TakeValue(args, &i, &schema_path, err)) return 1; } else if (args[i] == "--model") { if (!TakeValue(args, &i, &model_path, err)) return 1; + } else if (args[i] == "--predictions") { + if (!TakeValue(args, &i, &waterfalls.predictions_path, err)) return 1; + } else if (args[i] == "--out") { + if (!TakeValue(args, &i, &waterfalls.out_path, err)) return 1; } else if (args[i] == "--threshold") { if (!TakeValue(args, &i, &value, err)) return 1; score.threshold = std::stod(value); } else if (args[i] == "--tf-damping") { if (!TakeValue(args, &i, &value, err)) return 1; score.tf_damping = std::stod(value); + } else if (args[i] == "--fuzzy-tf") { + fuzzy_tf = true; + } else if (args[i] == "--ball-budget") { + if (!TakeValue(args, &i, &value, err)) return 1; + ball.budget = std::stoull(value); + } else if (args[i] == "--no-interactions") { + score.use_interactions = false; } else if (args[i] == "--pair") { if (!TakeValue(args, &i, &pair, err)) return 1; } else if (args[i] == "--rows") { if (!TakeValue(args, &i, &rows, err)) return 1; + } else if (args[i] == "--pairs") { + if (!TakeValue(args, &i, &pairs_path, err)) return 1; + pairs_by_row = false; + } else if (args[i] == "--row-pairs") { + if (!TakeValue(args, &i, &pairs_path, err)) return 1; + pairs_by_row = true; + } else if (args[i] == "--json") { + as_json = true; } else if (!args[i].empty() && args[i][0] == '-') { err << "cpplink explain: unknown option '" << args[i] << "'\n"; return 1; @@ -644,9 +715,28 @@ int RunExplain(const std::vector& args, std::ostream& out, "required\n"; return 1; } - if (pair.empty() == rows.empty()) { - err << "cpplink explain: give exactly one of --pair , or " - "--rows ,\n"; + const bool file_mode = !waterfalls.predictions_path.empty(); + if ((pair.empty() ? 0 : 1) + (rows.empty() ? 0 : 1) + (pairs_path.empty() ? 0 : 1) + + (file_mode ? 1 : 0) != + 1) { + err << "cpplink explain: give exactly one of --pair ,, " + "--rows ,, --pairs , --row-pairs or " + "--predictions \n"; + return 1; + } + // The JSON object and the waterfall file are the waterfall, and there is no + // waterfall without a model. + if ((as_json || file_mode) && model_path.empty()) { + err << "cpplink explain: " << (file_mode ? "--predictions" : "--json") + << " needs --model \n"; + return 1; + } + if (file_mode && waterfalls.out_path.empty()) { + err << "cpplink explain: --predictions needs --out \n"; + return 1; + } + if (!file_mode && !waterfalls.out_path.empty()) { + err << "cpplink explain: --out goes with --predictions \n"; return 1; } @@ -673,55 +763,102 @@ int RunExplain(const std::vector& args, std::ostream& out, return 1; } - uint64_t row_a = 0; - uint64_t row_b = 0; - std::string first; - std::string second; - if (!pair.empty()) { - if (!SplitPair(pair, &first, &second)) { - err << "cpplink explain: --pair wants ,\n"; + // Without a model there is no weight to explain: the levels are the whole + // story, and the waterfall is simply not printed. + Model model; + Scorer scorer; + BallTables balls; + if (!model_path.empty()) { + if (!LoadModel(model_path, &model, &error)) { + err << "cpplink: " << error << "\n"; return 1; } - if (!FindRowById(store, first, &row_a)) { - err << "cpplink explain: no record with id '" << first << "'\n"; + // In JSON mode `out` carries nothing but the objects, so the ball-table + // report goes where a tool reading them will not have to parse it. + if (fuzzy_tf) + BuildBallTables(comparisons, store, ball, &balls, as_json ? err : out); + if (!scorer.Bind(model, comparisons, store, score, &error, + fuzzy_tf ? &balls : nullptr)) { + err << "cpplink explain: " << error << "\n"; return 1; } - if (!FindRowById(store, second, &row_b)) { - err << "cpplink explain: no record with id '" << second << "'\n"; - return 1; + } + + const auto explain_one = [&](uint64_t row_a, uint64_t row_b) { + if (as_json) { + out << PairWaterfallJson(BuildPairWaterfall(store, comparisons, scorer, row_a, + row_b, &model)) + << "\n"; + return; } - } else { - if (!SplitPair(rows, &first, &second)) { - err << "cpplink explain: --rows wants ,\n"; - return 1; + PrintPairExplanation(store, comparisons, row_a, row_b, out); + if (!model_path.empty()) { + PrintPairWaterfall(store, comparisons, scorer, row_a, row_b, out); } - row_a = std::stoull(first); - row_b = std::stoull(second); - if (row_a >= store.NumRecords() || row_b >= store.NumRecords()) { - err << "cpplink explain: row out of range; the file has " - << store.NumRecords() << " records\n"; + }; + + // Every prediction a run wrote, as one wide row each: the file a viewer + // draws from, so the arithmetic never leaves the scorer. + if (file_mode) { + WaterfallReport report; + if (!WriteWaterfalls(store, comparisons, scorer, model, waterfalls, &report, + &error)) { + err << "cpplink " << error << "\n"; return 1; } + PrintWaterfallReport(report, out); + return 0; } - PrintGammaLayout(comparisons, out); - out << "\n"; - PrintPairExplanation(store, comparisons, row_a, row_b, out); - - // Without a model there is no weight to explain: the levels are the whole - // story, and the waterfall is simply not printed. - if (!model_path.empty()) { - Model model; - if (!LoadModel(model_path, &model, &error)) { - err << "cpplink: " << error << "\n"; + if (pairs_path.empty()) { + uint64_t row_a = 0; + uint64_t row_b = 0; + const bool by_row = pair.empty(); + if (!ResolvePair(store, by_row ? rows : pair, by_row, &row_a, &row_b, &error)) { + err << "cpplink explain: " << (by_row ? "--rows " : "--pair ") << error + << "\n"; return 1; } - Scorer scorer; - if (!scorer.Bind(model, comparisons, store, score, &error)) { - err << "cpplink explain: " << error << "\n"; + if (!as_json) { + PrintGammaLayout(comparisons, out); + out << "\n"; + } + explain_one(row_a, row_b); + return 0; + } + + // One pair per line, answered one line at a time and flushed, so a tool that + // holds this process open can write a pair and read its answer. The store is + // loaded once, which is the whole point: at 20M records the load is the cost, + // and the pair is free. A bad line is reported and does not end the batch. + std::ifstream file; + std::istream* in = &std::cin; + if (pairs_path != "-") { + file.open(pairs_path); + if (!file) { + err << "cpplink explain: cannot open " << pairs_path << "\n"; return 1; } - PrintPairWaterfall(store, comparisons, scorer, row_a, row_b, out); + in = &file; + } + std::string line; + while (std::getline(*in, line)) { + if (!line.empty() && line.back() == '\r') line.pop_back(); + if (line.empty()) continue; + uint64_t row_a = 0; + uint64_t row_b = 0; + if (!ResolvePair(store, line, pairs_by_row, &row_a, &row_b, &error)) { + if (as_json) { + out << nlohmann::json{{"pair", line}, {"error", error}}.dump() << "\n"; + } else { + out << "Pair " << line << ": " << error << "\n\n"; + } + out.flush(); + continue; + } + explain_one(row_a, row_b); + if (!as_json) out << "\n"; + out.flush(); } return 0; } diff --git a/src/cpplink/cluster.cpp b/src/cpplink/cluster.cpp index 04f93a7..7ee14d2 100644 --- a/src/cpplink/cluster.cpp +++ b/src/cpplink/cluster.cpp @@ -22,6 +22,7 @@ #include #include "cpplink/format.hpp" +#include "cpplink/id_index.hpp" #include "cpplink/merge_edges.hpp" #include "cpplink/predict.hpp" @@ -118,38 +119,6 @@ struct EdgeTally { double max_weight = 0.0; }; -// An id-to-row lookup over the store's id column: one `uint32` per record, kept -// in the order of the id it names. A merged edge file carries `unique_id`s -// rather than row indices, so clustering one has to map them back, and a hash -// map over 20M ids costs an order of magnitude more than the union-find it -// feeds. This is 4 bytes a record and a handful of string compares an edge. -class IdIndex { - public: - explicit IdIndex(const RecordStore& store) : ids_(store.ids()) { - order_.resize(static_cast(store.NumRecords())); - for (size_t row = 0; row < order_.size(); ++row) { - order_[row] = static_cast(row); - } - std::sort(order_.begin(), order_.end(), - [this](uint32_t a, uint32_t b) { return ids_.Get(a) < ids_.Get(b); }); - } - - bool Find(std::string_view id, uint32_t* row) const { - const auto at = - std::lower_bound(order_.begin(), order_.end(), id, - [this](uint32_t candidate, std::string_view key) { - return ids_.Get(candidate) < key; - }); - if (at == order_.end() || ids_.Get(*at) != id) return false; - *row = *at; - return true; - } - - private: - const IdColumn& ids_; - std::vector order_; -}; - bool ReadBinaryShard(const std::string& path, EdgeTally* tally, std::string* error) { std::ifstream file(path, std::ios::binary); if (!file) { diff --git a/src/cpplink/explain.cpp b/src/cpplink/explain.cpp index 799e689..2202b1a 100644 --- a/src/cpplink/explain.cpp +++ b/src/cpplink/explain.cpp @@ -8,9 +8,12 @@ #include #include #include +#include #include #include +#include + namespace cpplink { namespace { @@ -169,85 +172,119 @@ std::string Signed(double bits, int precision = 2) { } // namespace -void PrintPairWaterfall(const RecordStore& store, const ComparisonSet& comparisons, - const Scorer& scorer, uint64_t a, uint64_t b, std::ostream& out) { - const uint32_t gamma = comparisons.Evaluate(a, b); - const double prior = scorer.PriorWeight(); +PairWaterfall BuildPairWaterfall(const RecordStore& store, + const ComparisonSet& comparisons, const Scorer& scorer, + uint64_t a, uint64_t b, const Model* model) { + PairWaterfall w; + w.row_a = a; + w.row_b = b; + const IdColumn& ids = store.ids(); + if (!ids.offsets.empty()) { + w.id_a = std::string(ids.Get(a)); + w.id_b = std::string(ids.Get(b)); + } + w.records = store.NumRecords(); + w.gamma = comparisons.Evaluate(a, b); + w.prior = scorer.PriorWeight(); + + double running = w.prior; + for (size_t i = 0; i < comparisons.Size(); ++i) { + const BoundComparison& bound = comparisons.at(i); + WaterfallStep step; + step.name = bound.spec->name; + step.level = comparisons.LevelOf(w.gamma, i); + step.label = bound.spec->levels[step.level].Describe(); + step.value_a = ValueOf(bound, a); + step.value_b = ValueOf(bound, b); + step.m = std::nan(""); + step.u = std::nan(""); + if (model != nullptr && i < model->comparisons.size() && + step.level < model->comparisons[i].levels.size()) { + step.m = model->comparisons[i].levels[step.level].m; + step.u = model->comparisons[i].levels[step.level].u; + } + step.bits = scorer.LevelWeight(i, step.level); + step.tf = scorer.AdjustmentFor(i, w.gamma, a, b); + if (step.tf != 0.0) step.frequency = scorer.FrequencyFor(i, w.gamma, a); + running += step.bits + step.tf; + step.running = running; + w.steps.push_back(std::move(step)); + } + // The two-way corrections, where the model carries any. They are part of the + // sum the scorer computes, so they have to be part of the ledger that explains + // it: a waterfall missing them would total something the run never produced. + for (size_t i = 0; i < scorer.InteractionCount(); ++i) { + WaterfallInteraction term; + term.name = scorer.InteractionName(i); + term.bits = scorer.InteractionBits(i, w.gamma); + running += term.bits; + term.running = running; + w.interactions.push_back(std::move(term)); + } + w.weight = scorer.Weight(w.gamma, a, b); + w.probability = ProbabilityForWeight(w.weight); + w.bracket_low = scorer.BaseWeight(w.gamma) + scorer.DeltaMin(w.gamma); + w.bracket_high = scorer.BaseWeight(w.gamma) + scorer.DeltaMax(w.gamma); + w.threshold = scorer.threshold(); + w.zone = scorer.Classify(w.gamma); + w.emitted = w.weight >= w.threshold; + return w; +} + +void PrintPairWaterfall(const PairWaterfall& w, std::ostream& out) { out << "\n" << std::left << std::setw(18) << "Comparison" << std::setw(22) << "Level" << std::right << std::setw(10) << "bits" << std::setw(10) << "tf" << std::setw(12) << "running" << "\n"; out << std::string(72, '-') << "\n"; - double running = prior; out << std::left << std::setw(18) << "(prior)" << std::setw(22) << "lambda" - << std::right << std::setw(10) << Signed(prior) << std::setw(10) << "" - << std::setw(12) << Signed(running) << "\n"; - - for (size_t i = 0; i < comparisons.Size(); ++i) { - const BoundComparison& bound = comparisons.at(i); - const uint8_t level = comparisons.LevelOf(gamma, i); - const double bits = scorer.LevelWeight(i, level); - const double move = scorer.AdjustmentFor(i, gamma, a, b); - running += bits + move; - out << std::left << std::setw(18) << Truncate(bound.spec->name, 17) - << std::setw(22) << Truncate(bound.spec->levels[level].Describe(), 21) - << std::right << std::setw(10) << Signed(bits) << std::setw(10) - << (move != 0.0 ? Signed(move) : std::string("")) << std::setw(12) - << Signed(running) << "\n"; + << std::right << std::setw(10) << Signed(w.prior) << std::setw(10) << "" + << std::setw(12) << Signed(w.prior) << "\n"; + for (const WaterfallStep& step : w.steps) { + out << std::left << std::setw(18) << Truncate(step.name, 17) << std::setw(22) + << Truncate(step.label, 21) << std::right << std::setw(10) + << Signed(step.bits) << std::setw(10) + << (step.tf != 0.0 ? Signed(step.tf) : std::string("")) << std::setw(12) + << Signed(step.running) << "\n"; } - // The two-way corrections, where the model carries any. They are part of the - // sum the scorer computes, so they have to be part of the ledger that explains - // it: a waterfall missing them would total something the run never produced. - for (size_t i = 0; i < scorer.InteractionCount(); ++i) { - const double bits = scorer.InteractionBits(i, gamma); - running += bits; + for (const WaterfallInteraction& term : w.interactions) { out << std::left << std::setw(18) << "(interaction)" << std::setw(22) - << Truncate(scorer.InteractionName(i), 21) << std::right << std::setw(10) - << Signed(bits) << std::setw(10) << "" << std::setw(12) << Signed(running) - << "\n"; + << Truncate(term.name, 21) << std::right << std::setw(10) << Signed(term.bits) + << std::setw(10) << "" << std::setw(12) << Signed(term.running) << "\n"; } out << std::string(72, '-') << "\n"; - const double weight = scorer.Weight(gamma, a, b); - out << "Match weight " << std::fixed << std::setprecision(3) << weight - << " bits posterior " << std::setprecision(9) << ProbabilityForWeight(weight) - << "\n"; + out << "Match weight " << std::fixed << std::setprecision(3) << w.weight + << " bits posterior " << std::setprecision(9) << w.probability << "\n"; // Where the term-frequency moves came from. Without the counts the adjustment // is an unexplained number, and this is the report whose job is to explain it. bool any = false; - for (size_t i = 0; i < comparisons.Size(); ++i) { - if (!scorer.HasAdjustment(i)) continue; - const double move = scorer.AdjustmentFor(i, gamma, a, b); - if (move == 0.0) continue; + for (const WaterfallStep& step : w.steps) { + if (step.tf == 0.0) continue; if (!any) { out << "\nTerm frequency, for the comparisons that moved the weight:\n"; any = true; } - const uint32_t frequency = scorer.FrequencyFor(i, gamma, a); - const double share = - store.NumRecords() > 0 - ? static_cast(frequency) / static_cast(store.NumRecords()) - : 0.0; - out << " " << std::left << std::setw(18) - << Truncate(comparisons.at(i).spec->name, 17) << std::setw(24) - << Truncate(ValueOf(comparisons.at(i), a), 23) << std::right << std::setw(12) - << frequency << " rows" << std::setw(12) << std::scientific + const double share = w.records > 0 ? static_cast(step.frequency) / + static_cast(w.records) + : 0.0; + out << " " << std::left << std::setw(18) << Truncate(step.name, 17) + << std::setw(24) << Truncate(step.value_a, 23) << std::right << std::setw(12) + << step.frequency << " rows" << std::setw(12) << std::scientific << std::setprecision(2) << share << std::setw(10) << std::defaultfloat - << Signed(move) << " bits\n"; + << Signed(step.tf) << " bits\n"; } // The same three-way decision `predict` makes, so a pair can be traced from // here to whether it would have been emitted. - const Zone zone = scorer.Classify(gamma); - out << "\nPattern bracket " << std::fixed << std::setprecision(3) - << scorer.BaseWeight(gamma) + scorer.DeltaMin(gamma) << " to " - << scorer.BaseWeight(gamma) + scorer.DeltaMax(gamma) << " bits, against a " - << "threshold of " << scorer.threshold() << "\n"; - out << "Zone " << ZoneName(zone) << " -- "; - switch (zone) { + out << "\nPattern bracket " << std::fixed << std::setprecision(3) << w.bracket_low + << " to " << w.bracket_high << " bits, against a threshold of " << w.threshold + << "\n"; + out << "Zone " << ZoneName(w.zone) << " -- "; + switch (w.zone) { case Zone::kDrop: out << "no pair with this pattern can clear the threshold, whatever " "values it carries\n"; @@ -264,8 +301,65 @@ void PrintPairWaterfall(const RecordStore& store, const ComparisonSet& compariso out << "no evaluation can produce this pattern\n"; break; } - out << (weight >= scorer.threshold() ? "This pair would be emitted.\n" - : "This pair would not be emitted.\n"); + out << (w.emitted ? "This pair would be emitted.\n" + : "This pair would not be emitted.\n"); +} + +void PrintPairWaterfall(const RecordStore& store, const ComparisonSet& comparisons, + const Scorer& scorer, uint64_t a, uint64_t b, std::ostream& out) { + PrintPairWaterfall(BuildPairWaterfall(store, comparisons, scorer, a, b), out); +} + +namespace { + +// JSON has no NaN, and a rate the report was not given is absent rather than null. +void PutRate(nlohmann::json* item, const char* key, double value) { + if (!std::isnan(value)) (*item)[key] = value; +} + +} // namespace + +std::string PairWaterfallJson(const PairWaterfall& w) { + nlohmann::json root; + root["row_a"] = w.row_a; + root["row_b"] = w.row_b; + if (!w.id_a.empty() || !w.id_b.empty()) { + root["id_a"] = w.id_a; + root["id_b"] = w.id_b; + } + root["records"] = w.records; + root["gamma"] = w.gamma; + root["prior"] = w.prior; + root["steps"] = nlohmann::json::array(); + for (const WaterfallStep& step : w.steps) { + nlohmann::json item; + item["name"] = step.name; + item["level"] = step.level; + item["label"] = step.label; + item["values"] = {step.value_a, step.value_b}; + PutRate(&item, "m", step.m); + PutRate(&item, "u", step.u); + item["bits"] = step.bits; + item["tf"] = step.tf; + if (step.tf != 0.0) item["frequency"] = step.frequency; + item["running"] = step.running; + root["steps"].push_back(std::move(item)); + } + root["interactions"] = nlohmann::json::array(); + for (const WaterfallInteraction& term : w.interactions) { + nlohmann::json item; + item["name"] = term.name; + item["bits"] = term.bits; + item["running"] = term.running; + root["interactions"].push_back(std::move(item)); + } + root["weight"] = w.weight; + root["probability"] = w.probability; + root["bracket"] = {w.bracket_low, w.bracket_high}; + root["threshold"] = w.threshold; + root["zone"] = ZoneName(w.zone); + root["emitted"] = w.emitted; + return root.dump(); } } // namespace cpplink diff --git a/src/cpplink/explain.hpp b/src/cpplink/explain.hpp index f5188c1..74b51c7 100644 --- a/src/cpplink/explain.hpp +++ b/src/cpplink/explain.hpp @@ -6,8 +6,10 @@ #include #include #include +#include #include "cpplink/comparison.hpp" +#include "cpplink/model.hpp" #include "cpplink/record_store.hpp" #include "cpplink/score.hpp" @@ -22,15 +24,70 @@ void PrintGammaLayout(const ComparisonSet& comparisons, std::ostream& out); void PrintPairExplanation(const RecordStore& store, const ComparisonSet& comparisons, uint64_t a, uint64_t b, std::ostream& out); -// Prints the match weight as a waterfall: the prior, then each comparison's -// contribution in bits and its term-frequency move, with a running total. +// One comparison's row of the waterfall: the level it landed on, what it charged +// for it, and the values that put it there. +struct WaterfallStep { + std::string name; + uint8_t level = 0; + std::string label; + std::string value_a; + std::string value_b; + // The level's rates, where the model was given; NaN otherwise. + double m = 0.0; + double u = 0.0; + double bits = 0.0; // log2(m/u) + double tf = 0.0; // the term-frequency move, zero where the level has none + // Records carrying the shared value, on an exact adjusted level; zero elsewhere. + uint32_t frequency = 0; + double running = 0.0; // the total after this row +}; + +struct WaterfallInteraction { + std::string name; + double bits = 0.0; + double running = 0.0; +}; + +// The match weight of one pair as a ledger: the prior, then each comparison's +// contribution in bits and its term-frequency move, then the two-way corrections, +// with a running total that ends at `Scorer::Weight`. The text report and the +// JSON one are both printed from this, so they cannot explain different sums. // // A pattern says which levels fired; it does not say why the pair scored what it // scored. Two exact surname matches carry the same gamma and can differ by ten // bits, because one of them is "Smith". This is the view that shows that. +struct PairWaterfall { + uint64_t row_a = 0; + uint64_t row_b = 0; + std::string id_a; // empty where the store carries no ids + std::string id_b; + uint64_t records = 0; + uint32_t gamma = 0; + double prior = 0.0; + std::vector steps; + std::vector interactions; + double weight = 0.0; + double probability = 0.0; + double bracket_low = 0.0; // what any pair with this pattern can score + double bracket_high = 0.0; + double threshold = 0.0; + Zone zone = Zone::kCheck; + bool emitted = false; +}; + +// `model` is optional and only supplies each level's m and u for the report; the +// bits come from the scorer either way. +PairWaterfall BuildPairWaterfall(const RecordStore& store, + const ComparisonSet& comparisons, const Scorer& scorer, + uint64_t a, uint64_t b, const Model* model = nullptr); + +void PrintPairWaterfall(const PairWaterfall& waterfall, std::ostream& out); void PrintPairWaterfall(const RecordStore& store, const ComparisonSet& comparisons, const Scorer& scorer, uint64_t a, uint64_t b, std::ostream& out); +// The same ledger as one JSON object on one line, for a tool that draws it. +std::string PairWaterfallJson(const PairWaterfall& waterfall); + // Resolves a unique_id to a row by linear scan. There is no id index: ids are // almost all distinct, so an index would cost as much as the values and is only // ever needed for one-off lookups like this one. diff --git a/src/cpplink/id_index.hpp b/src/cpplink/id_index.hpp new file mode 100644 index 0000000..869816c --- /dev/null +++ b/src/cpplink/id_index.hpp @@ -0,0 +1,47 @@ +// Copyright 2026 Mathieu Fourment +// SPDX-License-Identifier: MIT + +#pragma once + +#include +#include +#include +#include + +#include "cpplink/record_store.hpp" + +namespace cpplink { + +// An id-to-row lookup over the store's id column: one `uint32` per record, kept +// in the order of the id it names. A merged prediction file carries `unique_id`s +// rather than row indices, so reading one back has to map them, and a hash map +// over 20M ids costs an order of magnitude more than the structures it feeds. +// This is 4 bytes a record and a handful of string compares a lookup. +class IdIndex { + public: + explicit IdIndex(const RecordStore& store) : ids_(store.ids()) { + order_.resize(static_cast(store.NumRecords())); + for (size_t row = 0; row < order_.size(); ++row) { + order_[row] = static_cast(row); + } + std::sort(order_.begin(), order_.end(), + [this](uint32_t a, uint32_t b) { return ids_.Get(a) < ids_.Get(b); }); + } + + bool Find(std::string_view id, uint32_t* row) const { + const auto at = + std::lower_bound(order_.begin(), order_.end(), id, + [this](uint32_t candidate, std::string_view key) { + return ids_.Get(candidate) < key; + }); + if (at == order_.end() || ids_.Get(*at) != id) return false; + *row = *at; + return true; + } + + private: + const IdColumn& ids_; + std::vector order_; +}; + +} // namespace cpplink diff --git a/src/cpplink/waterfall.cpp b/src/cpplink/waterfall.cpp new file mode 100644 index 0000000..2006b59 --- /dev/null +++ b/src/cpplink/waterfall.cpp @@ -0,0 +1,479 @@ +// Copyright 2026 Mathieu Fourment +// SPDX-License-Identifier: MIT + +#include "cpplink/waterfall.hpp" + +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include + +#include +#include +#include +#include + +#include "cpplink/explain.hpp" +#include "cpplink/format.hpp" +#include "cpplink/id_index.hpp" + +namespace cpplink { +namespace { + +// One prediction as the merged file names it. `gamma` is what the run stored, +// kept so a schema that has drifted from the file can be noticed. +struct NamedPrediction { + std::string_view id_a; + std::string_view id_b; + uint32_t gamma = 0; +}; + +using PredictionSink = std::function; + +bool ForEachCsvPrediction(const std::string& path, const PredictionSink& sink, + std::string* error) { + std::ifstream file(path); + if (!file) { + *error = "explain: could not open '" + path + "'"; + return false; + } + std::string line; + uint64_t number = 0; + while (std::getline(file, line)) { + ++number; + if (!line.empty() && line.back() == '\r') line.pop_back(); + if (line.empty()) continue; + if (number == 1 && line.rfind("id_a,", 0) == 0) continue; + NamedPrediction prediction; + double weight = 0.0; + if (!ParseEdgeCsvLine(line, &prediction.id_a, &prediction.id_b, &prediction.gamma, + &weight)) { + *error = "explain: '" + path + "' line " + std::to_string(number) + + " is not a cpplink prediction row"; + return false; + } + if (!sink(prediction, error)) return false; + } + return true; +} + +bool ForEachParquetPrediction(const std::string& path, const PredictionSink& sink, + std::string* error) { + auto input = arrow::io::ReadableFile::Open(path); + if (!input.ok()) { + *error = "explain: cannot open " + path + ": " + input.status().message(); + return false; + } + parquet::arrow::FileReaderBuilder builder; + arrow::Status status = builder.Open(*input); + if (!status.ok()) { + *error = + "explain: cannot read parquet metadata for " + path + ": " + status.message(); + return false; + } + auto reader_result = builder.Build(); + if (!reader_result.ok()) { + *error = "explain: cannot open parquet reader for " + path + ": " + + reader_result.status().message(); + return false; + } + std::unique_ptr reader = std::move(*reader_result); + std::shared_ptr schema; + status = reader->GetSchema(&schema); + if (!status.ok()) { + *error = "explain: cannot read the schema of " + path + ": " + status.message(); + return false; + } + std::vector indices; + for (const char* name : {"id_a", "id_b", "gamma"}) { + const int at = schema->GetFieldIndex(name); + if (at < 0) { + *error = "explain: '" + path + "' has no column \"" + name + + "\", so it is not a cpplink prediction file"; + return false; + } + indices.push_back(at); + } + for (int group = 0; group < reader->num_row_groups(); ++group) { + auto group_result = reader->ReadRowGroup(group, indices); + if (!group_result.ok()) { + *error = "explain: cannot read row group " + std::to_string(group) + " of " + + path + ": " + group_result.status().message(); + return false; + } + const std::shared_ptr table = *group_result; + const auto ids_a = table->GetColumnByName("id_a"); + const auto ids_b = table->GetColumnByName("id_b"); + const auto gammas = table->GetColumnByName("gamma"); + for (int chunk = 0; chunk < ids_a->num_chunks(); ++chunk) { + const auto a_array = + std::dynamic_pointer_cast(ids_a->chunk(chunk)); + const auto b_array = + std::dynamic_pointer_cast(ids_b->chunk(chunk)); + const auto gamma_array = + std::dynamic_pointer_cast(gammas->chunk(chunk)); + if (a_array == nullptr || b_array == nullptr || gamma_array == nullptr) { + *error = "explain: '" + path + + "' holds id_a, id_b or gamma in a type a cpplink prediction " + "file does not use"; + return false; + } + for (int64_t row = 0; row < a_array->length(); ++row) { + if (a_array->IsNull(row) || b_array->IsNull(row)) continue; + NamedPrediction prediction; + prediction.id_a = a_array->GetView(row); + prediction.id_b = b_array->GetView(row); + prediction.gamma = gamma_array->IsNull(row) ? 0 : gamma_array->Value(row); + if (!sink(prediction, error)) return false; + } + } + } + return true; +} + +// Column names are the comparison's name with a suffix, and an interaction's is +// the scorer's "left x right" made a legal identifier. +std::string InteractionColumn(const std::string& name) { + std::string out = name; + const size_t at = out.find(" x "); + if (at != std::string::npos) out.replace(at, 3, "_x_"); + return out + "_bits"; +} + +constexpr size_t kFixedColumns = 11; + +class WideSink { + public: + virtual ~WideSink() = default; + virtual bool Write(const PairWaterfall& w, std::string* error) = 0; + virtual bool Close(std::string* error) = 0; +}; + +class CsvWideSink : public WideSink { + public: + bool Open(const std::string& path, const std::vector& columns, + std::string* error) { + path_ = path; + file_.open(path); + if (!file_) { + *error = "explain: cannot create " + path; + return false; + } + for (size_t i = 0; i < columns.size(); ++i) { + if (i > 0) buffer_ += ','; + buffer_ += columns[i]; + } + buffer_ += '\n'; + return true; + } + + bool Write(const PairWaterfall& w, std::string* error) override { + char numbers[160]; + buffer_ += w.id_a; + buffer_ += ','; + buffer_ += w.id_b; + std::snprintf(numbers, sizeof(numbers), ",%u,%.6f,%.6f,%.9f,%.6f,%.6f,%.6f,", + w.gamma, w.prior, w.weight, w.probability, w.bracket_low, + w.bracket_high, w.threshold); + buffer_ += numbers; + buffer_ += ZoneName(w.zone); + buffer_ += w.emitted ? ",true" : ",false"; + for (const WaterfallStep& step : w.steps) { + std::snprintf(numbers, sizeof(numbers), ",%u,%.6f,%.6f,%u", + static_cast(step.level), step.bits, step.tf, + step.frequency); + buffer_ += numbers; + } + for (const WaterfallInteraction& term : w.interactions) { + std::snprintf(numbers, sizeof(numbers), ",%.6f", term.bits); + buffer_ += numbers; + } + buffer_ += '\n'; + if (buffer_.size() >= (1u << 20)) return Flush(error); + return true; + } + + bool Close(std::string* error) override { + if (!Flush(error)) return false; + file_.close(); + return true; + } + + private: + bool Flush(std::string* error) { + file_.write(buffer_.data(), static_cast(buffer_.size())); + buffer_.clear(); + if (!file_) { + *error = "explain: writing " + path_ + " failed"; + return false; + } + return true; + } + + std::string path_; + std::ofstream file_; + std::string buffer_; +}; + +class ParquetWideSink : public WideSink { + public: + ParquetWideSink(size_t batch_rows, size_t comparisons, size_t interactions) + : batch_rows_(batch_rows == 0 ? 1 : batch_rows), + levels_(comparisons), + bits_(comparisons), + tf_(comparisons), + frequency_(comparisons), + interaction_(interactions) {} + + bool Open(const std::string& path, const std::vector& columns, + std::string* error) { + path_ = path; + std::vector> fields = { + arrow::field("id_a", arrow::utf8()), + arrow::field("id_b", arrow::utf8()), + arrow::field("gamma", arrow::uint32()), + arrow::field("prior", arrow::float64()), + arrow::field("match_weight", arrow::float64()), + arrow::field("match_probability", arrow::float64()), + arrow::field("bracket_low", arrow::float64()), + arrow::field("bracket_high", arrow::float64()), + arrow::field("threshold", arrow::float64()), + arrow::field("zone", arrow::utf8()), + arrow::field("emitted", arrow::boolean()), + }; + size_t at = kFixedColumns; + for (size_t c = 0; c < levels_.size(); ++c) { + fields.push_back(arrow::field(columns[at++], arrow::uint8())); + fields.push_back(arrow::field(columns[at++], arrow::float64())); + fields.push_back(arrow::field(columns[at++], arrow::float64())); + fields.push_back(arrow::field(columns[at++], arrow::uint32())); + } + for (size_t i = 0; i < interaction_.size(); ++i) { + fields.push_back(arrow::field(columns[at++], arrow::float64())); + } + schema_ = arrow::schema(fields); + auto sink = arrow::io::FileOutputStream::Open(path); + if (!sink.ok()) { + *error = "explain: cannot create " + path + ": " + sink.status().message(); + return false; + } + auto props = parquet::WriterProperties::Builder() + .compression(parquet::Compression::SNAPPY) + ->build(); + auto writer = parquet::arrow::FileWriter::Open( + *schema_, arrow::default_memory_pool(), *sink, props); + if (!writer.ok()) { + *error = "explain: cannot open parquet writer for " + path + ": " + + writer.status().message(); + return false; + } + writer_ = std::move(*writer); + return true; + } + + bool Write(const PairWaterfall& w, std::string* error) override { + arrow::Status status = id_a_.Append(w.id_a); + status &= id_b_.Append(w.id_b); + status &= gamma_.Append(w.gamma); + status &= prior_.Append(w.prior); + status &= weight_.Append(w.weight); + status &= probability_.Append(w.probability); + status &= low_.Append(w.bracket_low); + status &= high_.Append(w.bracket_high); + status &= threshold_.Append(w.threshold); + status &= zone_.Append(ZoneName(w.zone)); + status &= emitted_.Append(w.emitted); + for (size_t c = 0; c < w.steps.size(); ++c) { + status &= levels_[c].Append(w.steps[c].level); + status &= bits_[c].Append(w.steps[c].bits); + status &= tf_[c].Append(w.steps[c].tf); + status &= frequency_[c].Append(w.steps[c].frequency); + } + for (size_t i = 0; i < w.interactions.size(); ++i) { + status &= interaction_[i].Append(w.interactions[i].bits); + } + if (!status.ok()) { + *error = "explain: building a row: " + status.message(); + return false; + } + ++pending_; + if (pending_ >= batch_rows_) return Flush(error); + return true; + } + + bool Close(std::string* error) override { + if (!Flush(error)) return false; + const arrow::Status closed = writer_->Close(); + if (!closed.ok()) { + *error = "explain: closing " + path_ + ": " + closed.message(); + return false; + } + return true; + } + + private: + bool Flush(std::string* error) { + if (pending_ == 0) return true; + std::vector> arrays; + arrays.reserve(static_cast(schema_->num_fields())); + arrow::Status status; + const auto finish = [&](arrow::ArrayBuilder* builder) { + std::shared_ptr array; + status &= builder->Finish(&array); + arrays.push_back(std::move(array)); + }; + finish(&id_a_); + finish(&id_b_); + finish(&gamma_); + finish(&prior_); + finish(&weight_); + finish(&probability_); + finish(&low_); + finish(&high_); + finish(&threshold_); + finish(&zone_); + finish(&emitted_); + for (size_t c = 0; c < levels_.size(); ++c) { + finish(&levels_[c]); + finish(&bits_[c]); + finish(&tf_[c]); + finish(&frequency_[c]); + } + for (auto& builder : interaction_) finish(&builder); + if (!status.ok()) { + *error = "explain: finishing a batch: " + status.message(); + return false; + } + const auto table = + arrow::Table::Make(schema_, arrays, static_cast(pending_)); + status = writer_->WriteTable(*table, static_cast(pending_)); + if (!status.ok()) { + *error = "explain: writing a row group to " + path_ + ": " + status.message(); + return false; + } + pending_ = 0; + return true; + } + + size_t batch_rows_; + std::shared_ptr schema_; + arrow::StringBuilder id_a_, id_b_, zone_; + arrow::UInt32Builder gamma_; + arrow::DoubleBuilder prior_, weight_, probability_, low_, high_, threshold_; + arrow::BooleanBuilder emitted_; + std::vector levels_; + std::vector bits_, tf_; + std::vector frequency_; + std::vector interaction_; + size_t pending_ = 0; + std::string path_; + std::unique_ptr writer_; +}; + +} // namespace + +std::vector WaterfallColumns(const ComparisonSet& comparisons, + const Scorer& scorer) { + std::vector columns = { + "id_a", "id_b", "gamma", + "prior", "match_weight", "match_probability", + "bracket_low", "bracket_high", "threshold", + "zone", "emitted"}; + for (size_t c = 0; c < comparisons.Size(); ++c) { + const std::string& name = comparisons.at(c).spec->name; + columns.push_back(name + "_level"); + columns.push_back(name + "_bits"); + columns.push_back(name + "_tf"); + columns.push_back(name + "_frequency"); + } + for (size_t i = 0; i < scorer.InteractionCount(); ++i) { + columns.push_back(InteractionColumn(scorer.InteractionName(i))); + } + return columns; +} + +bool WriteWaterfalls(const RecordStore& store, const ComparisonSet& comparisons, + const Scorer& scorer, const Model& model, + const WaterfallOptions& options, WaterfallReport* report, + std::string* error) { + const auto started = std::chrono::steady_clock::now(); + *report = WaterfallReport{}; + report->out_path = options.out_path; + if (!MergedFormatOf(options.out_path, &report->format)) { + *error = "explain: --out wants a .csv or .parquet file"; + return false; + } + MergeFormat in_format = MergeFormat::kCsv; + if (!MergedFormatOf(options.predictions_path, &in_format)) { + *error = "explain: --predictions wants the .csv or .parquet file predict wrote"; + return false; + } + + const std::vector columns = WaterfallColumns(comparisons, scorer); + report->columns = columns.size(); + std::unique_ptr sink; + if (report->format == MergeFormat::kParquet) { + auto parquet_sink = std::make_unique( + options.batch_rows, comparisons.Size(), scorer.InteractionCount()); + if (!parquet_sink->Open(options.out_path, columns, error)) return false; + sink = std::move(parquet_sink); + } else { + auto csv_sink = std::make_unique(); + if (!csv_sink->Open(options.out_path, columns, error)) return false; + sink = std::move(csv_sink); + } + + const IdIndex index(store); + const PredictionSink each = [&](const NamedPrediction& prediction, + std::string* trouble) { + ++report->read; + uint32_t a = 0; + uint32_t b = 0; + if (!index.Find(prediction.id_a, &a) || !index.Find(prediction.id_b, &b)) { + ++report->unresolved; + return true; + } + const PairWaterfall w = + BuildPairWaterfall(store, comparisons, scorer, a, b, &model); + if (w.gamma != prediction.gamma) ++report->pattern_changed; + if (!sink->Write(w, trouble)) return false; + ++report->written; + return true; + }; + const bool ok = in_format == MergeFormat::kParquet + ? ForEachParquetPrediction(options.predictions_path, each, error) + : ForEachCsvPrediction(options.predictions_path, each, error); + if (!ok) return false; + if (!sink->Close(error)) return false; + report->seconds = + std::chrono::duration(std::chrono::steady_clock::now() - started).count(); + return true; +} + +void PrintWaterfallReport(const WaterfallReport& report, std::ostream& out) { + out << "Read " << WithThousands(report.read) << " predictions, wrote " + << WithThousands(report.written) << " waterfalls to " << report.out_path << " (" + << (report.format == MergeFormat::kParquet ? "parquet" : "csv") << ", " + << report.columns << " columns) in " << std::fixed << std::setprecision(2) + << report.seconds << " s\n"; + if (report.unresolved > 0) { + out << " " << WithThousands(report.unresolved) + << " predictions name an id no record holds and were skipped: was this " + "file written from these inputs?\n"; + } + if (report.pattern_changed > 0) { + out << " " << WithThousands(report.pattern_changed) + << " predictions carry a gamma the comparisons no longer produce: the " + "schema has changed since predict ran, and the ledgers explain the " + "schema as it is now, not the weights in the file\n"; + } +} + +} // namespace cpplink diff --git a/src/cpplink/waterfall.hpp b/src/cpplink/waterfall.hpp new file mode 100644 index 0000000..2c44ee8 --- /dev/null +++ b/src/cpplink/waterfall.hpp @@ -0,0 +1,65 @@ +// Copyright 2026 Mathieu Fourment +// SPDX-License-Identifier: MIT + +#pragma once + +#include +#include +#include +#include + +#include "cpplink/comparison.hpp" +#include "cpplink/merge_edges.hpp" +#include "cpplink/model.hpp" +#include "cpplink/record_store.hpp" +#include "cpplink/score.hpp" + +namespace cpplink { + +// `explain --predictions --out `: the waterfall of every prediction +// a run wrote, as one wide row per pair. The file is what a viewer draws from, +// so the arithmetic stays in the scorer and a tool only reads. +// +// The columns are the pair and its totals, then four per comparison and one per +// two-way correction, named after the comparison: +// +// id_a, id_b, gamma, prior, match_weight, match_probability, +// bracket_low, bracket_high, threshold, zone, emitted, +// _level, _bits, _tf, _frequency, +// _x__bits +// +// A level's label, m and u are the model's rather than the pair's, so they are +// not repeated here: the model file carries them, indexed by the level. +struct WaterfallOptions { + std::string predictions_path; + std::string out_path; + size_t batch_rows = 65536; // rows per parquet row group +}; + +struct WaterfallReport { + MergeFormat format = MergeFormat::kCsv; + std::string out_path; + uint64_t read = 0; + uint64_t written = 0; + // Predictions naming an id no record holds. A file from another input. + uint64_t unresolved = 0; + // Predictions whose stored gamma is not what the comparisons produce now. + // The schema changed under the file; the row is written with the pattern + // the comparisons give today, and the count says the file is stale. + uint64_t pattern_changed = 0; + size_t columns = 0; + double seconds = 0.0; +}; + +bool WriteWaterfalls(const RecordStore& store, const ComparisonSet& comparisons, + const Scorer& scorer, const Model& model, + const WaterfallOptions& options, WaterfallReport* report, + std::string* error); + +void PrintWaterfallReport(const WaterfallReport& report, std::ostream& out); + +// The column names the file carries, in order, for the tools that read it. +std::vector WaterfallColumns(const ComparisonSet& comparisons, + const Scorer& scorer); + +} // namespace cpplink diff --git a/tests/explain_batch_test.cpp b/tests/explain_batch_test.cpp new file mode 100644 index 0000000..e2ec62d --- /dev/null +++ b/tests/explain_batch_test.cpp @@ -0,0 +1,390 @@ +// Copyright 2026 Mathieu Fourment +// SPDX-License-Identifier: MIT + +// `cpplink explain` answering many pairs from one load: the batch modes, the +// JSON object, and the wide file of every prediction's waterfall. + +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include + +#include +#include +#include +#include +#include +#include + +#include "cpplink/app.hpp" +#include "cpplink/model.hpp" +#include "cpplink/sample_data.hpp" + +namespace { + +constexpr const char* kSchema = R"({ + "unique_id": "id", + "columns": [ + {"name": "first_name", "type": "string"}, + {"name": "last_name", "type": "string"}, + {"name": "postcode", "type": "string"} + ], + "comparisons": [ + {"name": "last_name", "columns": ["last_name"], "term_frequency": true, "levels": [ + {"type": "null"}, {"type": "exact"}, + {"type": "jaro_winkler", "threshold": 0.85}, {"type": "else"}]}, + {"name": "first_name", "columns": ["first_name"], "levels": [ + {"type": "null"}, {"type": "exact"}, {"type": "else"}]}, + {"name": "postcode", "columns": ["postcode"], "term_frequency": true, "levels": [ + {"type": "null"}, {"type": "exact"}, {"type": "else"}]} + ], + "blocking": [{"type": "exact_value", "column": "postcode"}] +})"; + +cpplink::ModelComparison Learned(const std::string& name, bool tf, + const std::vector>& levels) { + cpplink::ModelComparison comparison; + comparison.name = name; + comparison.term_frequency = tf; + for (const auto& entry : levels) { + cpplink::ModelLevel level; + level.m = entry.first; + level.u = entry.second; + level.m_estimated = true; + comparison.levels.push_back(level); + } + return comparison; +} + +class ExplainBatch : public ::testing::Test { + protected: + void SetUp() override { + dir_ = std::filesystem::temp_directory_path() / + ("cpplink_explain_batch_" + std::to_string(::getpid())); + std::filesystem::remove_all(dir_); + std::filesystem::create_directories(dir_); + data_ = (dir_ / "sample.parquet").string(); + schema_ = (dir_ / "schema.json").string(); + model_ = (dir_ / "model.json").string(); + pairs_ = (dir_ / "pairs.txt").string(); + + cpplink::SampleOptions options; + options.rows = 500; + options.row_group_size = 100; + options.duplicate_rate = 0.2; + std::string error; + ASSERT_TRUE(cpplink::WriteSampleParquet(data_, options, &error)) << error; + std::ofstream(schema_) << kSchema; + + cpplink::Model model; + model.lambda = 0.01; + model.records = 500; + model.comparisons = { + Learned("last_name", true, + {{1e-3, 1e-3}, {0.8, 0.01}, {0.1, 0.02}, {0.099, 0.969}}), + Learned("first_name", false, {{1e-3, 1e-3}, {0.85, 0.05}, {0.149, 0.949}}), + Learned("postcode", true, {{1e-3, 1e-3}, {0.9, 0.005}, {0.099, 0.994}})}; + ASSERT_TRUE(cpplink::WriteModelJson(model, model_, &error)) << error; + } + void TearDown() override { + std::error_code ec; + std::filesystem::remove_all(dir_, ec); + } + + int Explain(std::vector extra, std::string* out, std::string* err) { + std::vector args = {"explain", "--schema", schema_, "--model", + model_, "--threshold", "2"}; + args.insert(args.end(), extra.begin(), extra.end()); + args.push_back(data_); + std::ostringstream o; + std::ostringstream e; + const int code = cpplink::Run(args, o, e); + *out = o.str(); + *err = e.str(); + return code; + } + + static std::vector Lines(const std::string& text) { + std::vector lines; + std::istringstream in(text); + std::string line; + while (std::getline(in, line)) { + if (!line.empty()) lines.push_back(nlohmann::json::parse(line)); + } + return lines; + } + + std::filesystem::path dir_; + std::string data_; + std::string schema_; + std::string model_; + std::string pairs_; +}; + +TEST_F(ExplainBatch, OnePairGivesOneObject) { + std::string out; + std::string err; + ASSERT_EQ(Explain({"--json", "--pair", "r1,r2"}, &out, &err), 0) << err; + const std::vector lines = Lines(out); + ASSERT_EQ(lines.size(), 1u) << out; + const nlohmann::json& w = lines[0]; + EXPECT_EQ(w["id_a"], "r1"); + EXPECT_EQ(w["id_b"], "r2"); + EXPECT_EQ(w["row_a"].get(), 1); + EXPECT_EQ(w["row_b"].get(), 2); + EXPECT_EQ(w["records"].get(), 500); + EXPECT_EQ(w["threshold"].get(), 2.0); + ASSERT_EQ(w["steps"].size(), 3u); + EXPECT_EQ(w["steps"][0]["name"], "last_name"); + EXPECT_TRUE(w["steps"][0].contains("m")); + // The ledger ends where the weight is. + const nlohmann::json& last = w["steps"][2]; + EXPECT_NEAR(last["running"].get(), w["weight"].get(), 1e-9); +} + +TEST_F(ExplainBatch, BatchAnswersEveryLineAndSurvivesABadOne) { + std::ofstream(pairs_) << "r1,r2\nr3,nope\n\nr5,r6\n"; + std::string out; + std::string err; + ASSERT_EQ(Explain({"--json", "--pairs", pairs_}, &out, &err), 0) << err; + const std::vector lines = Lines(out); + ASSERT_EQ(lines.size(), 3u) << out; + EXPECT_EQ(lines[0]["id_a"], "r1"); + EXPECT_EQ(lines[1]["pair"], "r3,nope"); + EXPECT_NE(lines[1]["error"].get().find("nope"), std::string::npos); + EXPECT_EQ(lines[2]["id_b"], "r6"); + + // The same pair answered alone is the same object. + std::string one; + ASSERT_EQ(Explain({"--json", "--pair", "r5,r6"}, &one, &err), 0) << err; + EXPECT_EQ(Lines(one)[0], lines[2]); +} + +TEST_F(ExplainBatch, RowPairsNameRowsAndAgreeWithIds) { + std::ofstream(pairs_) << "5,6\n999,1\n"; + std::string out; + std::string err; + ASSERT_EQ(Explain({"--json", "--row-pairs", pairs_}, &out, &err), 0) << err; + const std::vector lines = Lines(out); + ASSERT_EQ(lines.size(), 2u) << out; + EXPECT_EQ(lines[0]["id_a"], "r5"); + EXPECT_EQ(lines[0]["id_b"], "r6"); + EXPECT_NE(lines[1]["error"].get().find("out of range"), + std::string::npos); + + std::string by_id; + ASSERT_EQ(Explain({"--json", "--pair", "r5,r6"}, &by_id, &err), 0) << err; + EXPECT_EQ(Lines(by_id)[0], lines[0]); +} + +TEST_F(ExplainBatch, TextBatchPrintsOneReportPerPair) { + std::ofstream(pairs_) << "r1,r2\nr5,r6\n"; + std::string out; + std::string err; + ASSERT_EQ(Explain({"--pairs", pairs_}, &out, &err), 0) << err; + size_t reports = 0; + for (size_t at = out.find("Match weight"); at != std::string::npos; + at = out.find("Match weight", at + 1)) { + ++reports; + } + EXPECT_EQ(reports, 2u) << out; +} + +// The wide file: one row per prediction, whose columns add up to the weight the +// prediction file holds, in either format. +TEST_F(ExplainBatch, WaterfallFileExplainsEveryPrediction) { + const std::string predictions = (dir_ / "predictions.parquet").string(); + std::ostringstream out; + std::ostringstream err; + ASSERT_EQ( + cpplink::Run({"predict", "--schema", schema_, "--model", model_, "--threshold", + "2", "--out", predictions, "--threads", "2", data_}, + out, err), + 0) + << err.str(); + + std::map, std::pair> wrote; + { + auto input = arrow::io::ReadableFile::Open(predictions); + ASSERT_TRUE(input.ok()); + auto reader = parquet::arrow::OpenFile(*input, arrow::default_memory_pool()); + ASSERT_TRUE(reader.ok()); + std::shared_ptr table; + ASSERT_TRUE((*reader)->ReadTable(&table).ok()); + table = table->CombineChunks().ValueOrDie(); + const auto a = std::static_pointer_cast( + table->GetColumnByName("id_a")->chunk(0)); + const auto b = std::static_pointer_cast( + table->GetColumnByName("id_b")->chunk(0)); + const auto gamma = std::static_pointer_cast( + table->GetColumnByName("gamma")->chunk(0)); + const auto weight = std::static_pointer_cast( + table->GetColumnByName("match_weight")->chunk(0)); + for (int64_t i = 0; i < table->num_rows(); ++i) { + wrote[{a->GetString(i), b->GetString(i)}] = {gamma->Value(i), + weight->Value(i)}; + } + } + ASSERT_GT(wrote.size(), 10u) << "the fixture must produce predictions"; + + const std::string waterfalls = (dir_ / "waterfalls.parquet").string(); + std::string text; + std::string trouble; + ASSERT_EQ( + Explain({"--predictions", predictions, "--out", waterfalls}, &text, &trouble), 0) + << trouble; + EXPECT_NE(text.find("wrote " + std::to_string(wrote.size()) + " waterfalls"), + std::string::npos) + << text; + EXPECT_EQ(text.find("no longer produce"), std::string::npos) << text; + + auto input = arrow::io::ReadableFile::Open(waterfalls); + ASSERT_TRUE(input.ok()); + auto reader = parquet::arrow::OpenFile(*input, arrow::default_memory_pool()); + ASSERT_TRUE(reader.ok()); + std::shared_ptr table; + ASSERT_TRUE((*reader)->ReadTable(&table).ok()); + table = table->CombineChunks().ValueOrDie(); + ASSERT_EQ(static_cast(table->num_rows()), wrote.size()); + const std::vector names = {"last_name", "first_name", "postcode"}; + for (const std::string& name : names) { + for (const char* suffix : {"_level", "_bits", "_tf", "_frequency"}) { + ASSERT_NE(table->GetColumnByName(name + suffix), nullptr) << name << suffix; + } + } + const auto column = [&](const char* name) { + return table->GetColumnByName(name)->chunk(0); + }; + const auto a = std::static_pointer_cast(column("id_a")); + const auto b = std::static_pointer_cast(column("id_b")); + const auto gamma = std::static_pointer_cast(column("gamma")); + const auto prior = std::static_pointer_cast(column("prior")); + const auto weight = + std::static_pointer_cast(column("match_weight")); + const auto threshold = + std::static_pointer_cast(column("threshold")); + const auto emitted = std::static_pointer_cast(column("emitted")); + std::vector> bits; + std::vector> tf; + std::vector> frequency; + for (const std::string& name : names) { + bits.push_back(std::static_pointer_cast( + column((name + "_bits").c_str()))); + tf.push_back( + std::static_pointer_cast(column((name + "_tf").c_str()))); + frequency.push_back(std::static_pointer_cast( + column((name + "_frequency").c_str()))); + } + for (int64_t i = 0; i < table->num_rows(); ++i) { + const auto found = wrote.find({a->GetString(i), b->GetString(i)}); + ASSERT_NE(found, wrote.end()) << a->GetString(i) << "," << b->GetString(i); + EXPECT_EQ(gamma->Value(i), found->second.first); + // The weight is the run's, to the last bit, and the row adds up to it. + EXPECT_EQ(weight->Value(i), found->second.second); + double running = prior->Value(i); + for (size_t c = 0; c < names.size(); ++c) { + running += bits[c]->Value(i) + tf[c]->Value(i); + EXPECT_EQ(tf[c]->Value(i) != 0.0, frequency[c]->Value(i) != 0) + << names[c] << " row " << i; + } + EXPECT_NEAR(running, weight->Value(i), 1e-9); + EXPECT_EQ(threshold->Value(i), 2.0); + EXPECT_EQ(emitted->Value(i), weight->Value(i) >= 2.0); + } + + // The csv holds the same rows: a header naming every column, one line each. + const std::string csv = (dir_ / "waterfalls.csv").string(); + ASSERT_EQ(Explain({"--predictions", predictions, "--out", csv}, &text, &trouble), 0) + << trouble; + std::ifstream file(csv); + std::string line; + ASSERT_TRUE(std::getline(file, line)); + EXPECT_EQ(line.rfind("id_a,id_b,gamma,prior,match_weight,", 0), 0u) << line; + EXPECT_NE(line.find(",postcode_frequency"), std::string::npos) << line; + const size_t width = static_cast(std::count(line.begin(), line.end(), ',')); + size_t rows = 0; + while (std::getline(file, line)) { + if (line.empty()) continue; + ++rows; + // Every row has the header's width, or a reader would misalign columns. + ASSERT_EQ(static_cast(std::count(line.begin(), line.end(), ',')), width) + << line; + std::string_view id_a; + std::string_view id_b; + const size_t first = line.find(','); + const size_t second = line.find(',', first + 1); + id_a = std::string_view(line).substr(0, first); + id_b = std::string_view(line).substr(first + 1, second - first - 1); + const auto found = wrote.find({std::string(id_a), std::string(id_b)}); + ASSERT_NE(found, wrote.end()) << line; + // The csv rounds to six decimals, as the prediction csv does. + const size_t third = line.find(',', second + 1); + const size_t fourth = line.find(',', third + 1); + const size_t fifth = line.find(',', fourth + 1); + const double weight = std::stod(line.substr(fourth + 1, fifth - fourth - 1)); + EXPECT_NEAR(weight, found->second.second, 1e-6) << line; + } + EXPECT_EQ(rows, wrote.size()); +} + +TEST_F(ExplainBatch, WaterfallFileSkipsAndReportsUnknownIds) { + const std::string predictions = (dir_ / "predictions.csv").string(); + std::ofstream(predictions) << "id_a,id_b,gamma,match_weight,match_probability\n" + << "r1,r2,0,1.0,0.5\nr3,nope,0,1.0,0.5\n"; + const std::string waterfalls = (dir_ / "waterfalls.csv").string(); + std::string text; + std::string trouble; + ASSERT_EQ( + Explain({"--predictions", predictions, "--out", waterfalls}, &text, &trouble), 0) + << trouble; + EXPECT_NE(text.find("wrote 1 waterfalls"), std::string::npos) << text; + EXPECT_NE(text.find("1 predictions name an id no record holds"), std::string::npos) + << text; + + // A waterfall file needs a model and an output, and --out goes with it. + std::ostringstream out; + std::ostringstream err; + EXPECT_EQ(cpplink::Run({"explain", "--schema", schema_, "--predictions", predictions, + "--out", waterfalls, data_}, + out, err), + 1); + EXPECT_NE(err.str().find("--model"), std::string::npos); + err.str(""); + EXPECT_EQ(cpplink::Run({"explain", "--schema", schema_, "--model", model_, + "--predictions", predictions, data_}, + out, err), + 1); + EXPECT_NE(err.str().find("--out"), std::string::npos); + err.str(""); + EXPECT_EQ(cpplink::Run({"explain", "--schema", schema_, "--model", model_, "--pair", + "r1,r2", "--out", waterfalls, data_}, + out, err), + 1); + EXPECT_NE(err.str().find("--predictions"), std::string::npos); +} + +TEST_F(ExplainBatch, JsonNeedsAModelAndOnePairSource) { + std::ostringstream out; + std::ostringstream err; + EXPECT_EQ( + cpplink::Run({"explain", "--schema", schema_, "--json", "--pair", "r1,r2", data_}, + out, err), + 1); + EXPECT_NE(err.str().find("--model"), std::string::npos); + + err.str(""); + EXPECT_EQ(cpplink::Run({"explain", "--schema", schema_, "--pair", "r1,r2", "--pairs", + pairs_, data_}, + out, err), + 1); + EXPECT_NE(err.str().find("exactly one"), std::string::npos); +} + +} // namespace diff --git a/tests/explain_test.cpp b/tests/explain_test.cpp index 57cbf3a..915e0c4 100644 --- a/tests/explain_test.cpp +++ b/tests/explain_test.cpp @@ -11,6 +11,7 @@ #include #include +#include #include "cpplink/comparison.hpp" #include "cpplink/model.hpp" @@ -66,15 +67,15 @@ class ExplainFixture : public ::testing::Test { store_->Finalize(); ASSERT_TRUE(comparisons_.Bind(schema_, *store_, &error)) << error; - cpplink::Model model; - model.lambda = 0.05; - model.records = kRecords; - model.comparisons = { + model_.lambda = 0.05; + model_.records = kRecords; + model_.comparisons = { Comparison("surname", true, {{1e-9, 1e-9}, {0.9, 0.02}, {0.1, 0.98}}), Comparison("city", false, {{0.8, 0.25}, {0.2, 0.75}})}; cpplink::ScoreOptions options; options.threshold = 1.0; - ASSERT_TRUE(scorer_.Bind(model, comparisons_, *store_, options, &error)) << error; + ASSERT_TRUE(scorer_.Bind(model_, comparisons_, *store_, options, &error)) + << error; } static cpplink::ModelComparison Comparison( @@ -96,6 +97,7 @@ class ExplainFixture : public ::testing::Test { cpplink::Schema schema_; std::unique_ptr store_; cpplink::ComparisonSet comparisons_; + cpplink::Model model_; cpplink::Scorer scorer_; }; @@ -104,18 +106,66 @@ class ExplainFixture : public ::testing::Test { TEST_F(ExplainFixture, TheStepsSumToTheScore) { for (uint64_t a = 0; a < 40; ++a) { for (uint64_t b = a + 1; b < 40; ++b) { - const uint32_t gamma = comparisons_.Evaluate(a, b); - double running = scorer_.PriorWeight(); - for (size_t c = 0; c < comparisons_.Size(); ++c) { - running += scorer_.LevelWeight(c, comparisons_.LevelOf(gamma, c)); - running += scorer_.AdjustmentFor(c, gamma, a, b); + const cpplink::PairWaterfall w = cpplink::BuildPairWaterfall( + *store_, comparisons_, scorer_, a, b, &model_); + double running = w.prior; + for (const cpplink::WaterfallStep& step : w.steps) { + running += step.bits + step.tf; + ASSERT_NEAR(running, step.running, 1e-12); + } + for (const cpplink::WaterfallInteraction& term : w.interactions) { + running += term.bits; + ASSERT_NEAR(running, term.running, 1e-12); } - ASSERT_NEAR(running, scorer_.Weight(gamma, a, b), 1e-9) - << "rows " << a << " and " << b; + ASSERT_NEAR(running, w.weight, 1e-9) << "rows " << a << " and " << b; + ASSERT_NEAR(w.weight, scorer_.Weight(w.gamma, a, b), 1e-12); + ASSERT_EQ(w.gamma, comparisons_.Evaluate(a, b)); } } } +// The JSON object carries the same ledger as the text, so a tool drawing it draws +// the run's own arithmetic. The rates come from the model the report was given. +TEST_F(ExplainFixture, JsonCarriesTheSameLedger) { + const cpplink::PairWaterfall w = + cpplink::BuildPairWaterfall(*store_, comparisons_, scorer_, 3, 11, &model_); + const nlohmann::json root = nlohmann::json::parse(cpplink::PairWaterfallJson(w)); + EXPECT_EQ(root["id_a"], "r3"); + EXPECT_EQ(root["id_b"], "r11"); + EXPECT_EQ(root["gamma"].get(), w.gamma); + EXPECT_DOUBLE_EQ(root["prior"].get(), w.prior); + EXPECT_DOUBLE_EQ(root["weight"].get(), w.weight); + EXPECT_DOUBLE_EQ(root["probability"].get(), w.probability); + EXPECT_EQ(root["threshold"].get(), 1.0); + EXPECT_EQ(root["emitted"].get(), w.weight >= 1.0); + ASSERT_EQ(root["steps"].size(), 2u); + ASSERT_EQ(root["bracket"].size(), 2u); + double running = w.prior; + for (size_t i = 0; i < 2; ++i) { + const nlohmann::json& step = root["steps"][i]; + EXPECT_EQ(step["name"], w.steps[i].name); + EXPECT_EQ(step["label"], w.steps[i].label); + EXPECT_EQ(step["level"].get(), w.steps[i].level); + EXPECT_EQ(step["values"].size(), 2u); + running += step["bits"].get() + step["tf"].get(); + EXPECT_DOUBLE_EQ(step["running"].get(), running); + const cpplink::ModelLevel& learned = + model_.comparisons[i].levels[w.steps[i].level]; + EXPECT_DOUBLE_EQ(step["m"].get(), learned.m); + EXPECT_DOUBLE_EQ(step["u"].get(), learned.u); + // The frequency is reported exactly where a term-frequency move is. + EXPECT_EQ(step.contains("frequency"), step["tf"].get() != 0.0); + } + EXPECT_DOUBLE_EQ(running, w.weight); + + // Without a model the rates are absent rather than null. + const cpplink::PairWaterfall bare = + cpplink::BuildPairWaterfall(*store_, comparisons_, scorer_, 3, 11); + const nlohmann::json plain = nlohmann::json::parse(cpplink::PairWaterfallJson(bare)); + EXPECT_FALSE(plain["steps"][0].contains("m")); + EXPECT_DOUBLE_EQ(plain["weight"].get(), w.weight); +} + TEST_F(ExplainFixture, AdjustmentAppliesOnlyToAnAdjustedLevel) { // Any pair whose surnames disagree lands on `else`, where no term-frequency // move applies however rare either value is. diff --git a/tools/cluster_server.py b/tools/cluster_server.py index 5f9d00d..73d3a18 100755 --- a/tools/cluster_server.py +++ b/tools/cluster_server.py @@ -4,7 +4,8 @@ """Serve the cluster viewer over a DuckDB cache, for runs too large to embed. tools/cluster_server.py --schema s.json --clusters clusters.csv \ - --predictions predictions.parquet --truth truth.csv data.parquet + --predictions predictions.parquet --waterfalls waterfalls.parquet \ + --model m.json --truth truth.csv data.parquet Same page as cluster_view.py, but the clusters are queried rather than written into the file, so every cluster of the run is reachable instead of a sample. @@ -18,6 +19,14 @@ it is a fraction of the dataset: at 20M records with 2.9M of them clustered, the parquet is read exactly once and every later request touches a table seven times smaller than the file. + +With `--waterfalls` the page also draws the waterfall of the prediction picked. +The file is what `cpplink explain --predictions --out ` writes, one +wide row per prediction; it is loaded into the cache beside the predictions and +a click is one lookup, so the ledger is the scorer's own arithmetic and nothing +is computed or spawned here. The level labels and rates come from `--model`. +Predictions clustering kept apart -- above the write threshold, below the +clustering one -- are kept and listed as `rejected`. """ import argparse @@ -36,11 +45,11 @@ import pyarrow.parquet as pq sys.path.insert(0, os.path.dirname(os.path.abspath(__file__))) -from cluster_view import (EDGE_DTYPE, EDGE_MAGIC, cell, # noqa: E402 - read_truth, schema_columns) +from cluster_view import (EDGE_DTYPE, EDGE_MAGIC, cell, ledger, # noqa: E402 + read_model_levels, read_truth, schema_columns) EDGE_CHUNK = 1 << 20 -CACHE_VERSION = 3 +CACHE_VERSION = 5 def quoted(name): @@ -65,7 +74,7 @@ def scan(path, types): def fingerprint(args, columns): """What the cache was built from, so a changed input rebuilds it.""" parts = [CACHE_VERSION, columns, args.max_rows] - for path in [args.clusters, args.truth] + list(args.data): + for path in [args.clusters, args.truth, args.waterfalls] + list(args.data): if path: parts.append([os.path.abspath(path), os.path.getmtime(path), os.path.getsize(path)]) @@ -124,42 +133,36 @@ def build(conn, args, id_column, columns, say): """) say("reading the clustered records") - conn.execute(f""" - CREATE TABLE records AS - SELECT m.cluster_id, d.{quoted(id_column)} AS uid, {picked}, - lower(concat_ws(' ', d.{quoted(id_column)}, {picked})) AS text - FROM read_parquet({files}) d - JOIN members m ON m.uid = d.{quoted(id_column)} - """) + # A record keeps its store row (`rid`): the position in the inputs read in + # order, which is what a shard names. file_row_number is that position + # inside one file, and the offsets between the inputs are the only thing to + # carry. + parts, offset = [], 0 + for path in files: + parts.append(f""" + SELECT m.cluster_id, d.{quoted(id_column)} AS uid, {picked}, + lower(concat_ws(' ', d.{quoted(id_column)}, {picked})) AS text, + {offset} + file_row_number AS rid + FROM read_parquet('{path}', file_row_number = true) d + JOIN members m ON m.uid = d.{quoted(id_column)}""") + offset += pq.ParquetFile(path).metadata.num_rows + conn.execute("CREATE TABLE records AS " + " UNION ALL ".join(parts)) if args.predictions: if os.path.isdir(args.predictions): say("naming the predictions") - # A row index in a shard is a position in the inputs read in order, - # and file_row_number is that position inside one file, so the - # offsets that separate the inputs are the only thing to carry. - conn.execute("CREATE TABLE rowmap (rid BIGINT, uid VARCHAR)") - offset = 0 - for path in files: - conn.execute(f""" - INSERT INTO rowmap - SELECT {offset} + file_row_number, {quoted(id_column)} - FROM read_parquet('{path}', file_row_number = true) - """) - offset += pq.ParquetFile(path).metadata.num_rows + # A shard names rows, and the clustered records carry theirs. load_shard_edges(conn, args.predictions) conn.execute(""" CREATE TABLE named_edges AS SELECT ra.uid AS a, rb.uid AS b, e.w AS weight FROM raw_edges e - JOIN rowmap ra ON ra.rid = e.a - JOIN rowmap rb ON rb.rid = e.b + JOIN records ra ON ra.rid = e.a + JOIN records rb ON rb.rid = e.b """) conn.execute("DROP TABLE raw_edges") - conn.execute("DROP TABLE rowmap") else: - # A merged file already names records by unique_id, which is the - # whole of what the rowmap above exists to recover. + # A merged file already names records by unique_id. say("reading the predictions") conn.execute(f""" CREATE TABLE named_edges AS @@ -167,21 +170,47 @@ def build(conn, args, id_column, columns, say): match_weight AS weight FROM {scan(args.predictions, ("id_a", "id_b"))} """) - # Both ends, because clustering at a threshold above the one the run - # wrote at leaves predictions that cross two clusters or land outside - # every one, and the page shows a weight only inside a cluster. + # Both ends clustered, because clustering at a threshold above the one + # the run wrote at leaves predictions that cross two clusters or land + # outside every one. A pair across two clusters is kept with both + # cluster ids, since it is the prediction clustering overruled and the + # one most worth explaining; a pair with an end no cluster holds has no + # record to show and is dropped. conn.execute(""" - CREATE TABLE edges AS - SELECT m.cluster_id, e.a, e.b, e.weight + CREATE TABLE pairs AS + SELECT e.a, e.b, e.weight, ra.cluster_id AS cluster_a, + rb.cluster_id AS cluster_b FROM named_edges e - JOIN members m ON m.uid = e.a - JOIN members mb ON mb.uid = e.b AND mb.cluster_id = m.cluster_id + JOIN records ra ON ra.uid = e.a + JOIN records rb ON rb.uid = e.b """) conn.execute("DROP TABLE named_edges") + conn.execute(""" + CREATE TABLE edges AS + SELECT cluster_a AS cluster_id, a, b, weight FROM pairs + WHERE cluster_a = cluster_b + """) conn.execute("CREATE INDEX edges_by_cluster ON edges (cluster_id)") + conn.execute("CREATE INDEX pairs_by_weight ON pairs (weight)") + conn.execute("CREATE INDEX pairs_by_a ON pairs (a)") + conn.execute("CREATE INDEX pairs_by_b ON pairs (b)") + if args.waterfalls: + # Only the rows the page can ever ask for: the pairs kept above. + say("reading the waterfalls") + conn.execute(f""" + CREATE TABLE waterfalls AS + SELECT w.* REPLACE (CAST(w.id_a AS VARCHAR) AS id_a, + CAST(w.id_b AS VARCHAR) AS id_b) + FROM {scan(args.waterfalls, ("id_a", "id_b"))} w + JOIN pairs p ON p.a = CAST(w.id_a AS VARCHAR) + AND p.b = CAST(w.id_b AS VARCHAR) + """) + conn.execute("CREATE INDEX waterfalls_by_pair ON waterfalls (id_a, id_b)") else: conn.execute("CREATE TABLE edges (cluster_id VARCHAR, a VARCHAR, " "b VARCHAR, weight DOUBLE)") + conn.execute("CREATE TABLE pairs (a VARCHAR, b VARCHAR, weight DOUBLE, " + "cluster_a VARCHAR, cluster_b VARCHAR)") if args.truth: say("closing the known pairs") @@ -222,6 +251,7 @@ def build(conn, args, id_column, columns, say): """) conn.execute("CREATE INDEX stats_by_id ON stats (id)") conn.execute("CREATE INDEX records_by_cluster ON records (cluster_id)") + conn.execute("CREATE INDEX records_by_uid ON records (uid)") ORDERS = { @@ -238,20 +268,39 @@ def build(conn, args, id_column, columns, say): "mixed": "entities > 1", } +PAIR_ORDERS = { + "weakest": "weight, a, b", + "strongest": "weight DESC, a, b", + "id": "a, b", +} + +PAIR_FILTERS = { + "all": "TRUE", + "within": "cluster_a = cluster_b", + "rejected": "cluster_a <> cluster_b", +} + class Viewer: """The queries the page makes, over one shared read-only connection.""" - def __init__(self, conn, columns, source, max_rows, has_truth): + def __init__(self, conn, columns, source, max_rows, has_truth, levels=None): self.conn = conn self.columns = columns self.source = source self.max_rows = max_rows self.has_truth = has_truth + self.levels = levels # the model's level table, when waterfalls are held self.lock = threading.Lock() with self.lock: self.totals = conn.execute( "SELECT count(*), coalesce(sum(size), 0) FROM stats").fetchone() + self.pair_total = conn.execute("SELECT count(*) FROM pairs").fetchone()[0] + self.waterfall_columns = [ + r[0] for r in conn.execute( + "SELECT column_name FROM information_schema.columns " + "WHERE table_name = 'waterfalls' ORDER BY ordinal_position") + .fetchall()] def query(self, sql, params=()): with self.lock: @@ -259,8 +308,64 @@ def query(self, sql, params=()): def summary(self): return {"columns": self.columns, "source": self.source, - "totals": {"clusters": self.totals[0], "records": self.totals[1]}, - "has_truth": self.has_truth} + "totals": {"clusters": self.totals[0], "records": self.totals[1], + "pairs": self.pair_total}, + "has_truth": self.has_truth, + "has_model": bool(self.waterfall_columns)} + + def waterfall(self, a, b): + """The ledger of one prediction, from the wide row the run's explain wrote.""" + if not self.waterfall_columns or self.levels is None: + return None + rows = self.query("SELECT * FROM waterfalls WHERE id_a = ? AND id_b = ?", [a, b]) + if not rows: + return None + return ledger(dict(zip(self.waterfall_columns, rows[0])), self.levels) + + def pairs(self, q, sort, only, offset, limit): + where = PAIR_FILTERS.get(only, "TRUE") + params = [] + if q: + # Either record's text, since a pair is read by either of its ends. + like = ("%" + q.lower().replace("\\", "\\\\") + .replace("%", "\\%").replace("_", "\\_") + "%") + where += (" AND (a IN (SELECT uid FROM records WHERE text LIKE ? " + "ESCAPE '\\') OR b IN (SELECT uid FROM records WHERE text " + "LIKE ? ESCAPE '\\'))") + params += [like, like] + matched = self.query(f"SELECT count(*) FROM pairs WHERE {where}", params)[0][0] + rows = self.query( + f"SELECT a, b, weight, cluster_a, cluster_b FROM pairs WHERE {where} " + f"ORDER BY {PAIR_ORDERS.get(sort, PAIR_ORDERS['weakest'])} " + f"LIMIT {int(limit)} OFFSET {int(offset)}", params) + return {"matched": matched, + "pairs": [{"a": r[0], "b": r[1], "w": round(r[2], 3), + "ca": r[3], "cb": r[4]} for r in rows]} + + def pair(self, a, b): + head = self.query("SELECT a, b, weight, cluster_a, cluster_b FROM pairs " + "WHERE a = ? AND b = ?", [a, b]) + if not head: + return None + picked = ", ".join(quoted(c) for c in self.columns) + rows = {r[0]: r for r in self.query( + f"SELECT uid, {picked} FROM records WHERE uid IN (?, ?)", [a, b])} + if a not in rows or b not in rows: + return None + answer = { + "a": a, "b": b, "w": round(head[0][2], 3), + "ca": head[0][3], "cb": head[0][4], + "rows": [{"id": uid, "v": [cell(v) for v in rows[uid][1:]]} + for uid in (a, b)], + } + if self.has_truth: + groups = dict(self.query( + "SELECT uid, grp FROM truth WHERE uid IN (?, ?)", [a, b])) + same = a in groups and b in groups and groups[a] == groups[b] + answer["rows"][0]["t"] = 0 + answer["rows"][1]["t"] = 0 if same else 1 + answer["waterfall"] = self.waterfall(a, b) + return answer def listing(self, q, sort, only, offset, limit): where = FILTERS.get(only, "TRUE") @@ -356,6 +461,15 @@ def do_GET(self): found = viewer.cluster(query.get("id", "")) self.send_json(found or {"error": "no such cluster"}, 200 if found else 404) + elif url.path == "/api/pairs": + self.send_json(viewer.pairs( + query.get("q", ""), query.get("sort", "weakest"), + query.get("only", "all"), int(query.get("offset", 0)), + min(int(query.get("limit", 100)), 500))) + elif url.path == "/api/pair": + found = viewer.pair(query.get("a", ""), query.get("b", "")) + self.send_json(found or {"error": "no such prediction"}, + 200 if found else 404) else: self.send_json({"error": "not found"}, 404) except BrokenPipeError: @@ -377,6 +491,11 @@ def main(): help="the run's predictions -- one csv or parquet file, or the " "shard directory -- for within-cluster weights") ap.add_argument("--truth", help="known pairs, to colour members by true entity") + ap.add_argument("--waterfalls", help="the csv or parquet cpplink explain " + "--predictions wrote; draws each " + "prediction's waterfall") + ap.add_argument("--model", help="the model the run scored with, for the level " + "labels and rates the waterfall shows") ap.add_argument("--cache", default="cluster_view.duckdb") ap.add_argument("--rebuild", action="store_true", help="rebuild the cache first") ap.add_argument("--min-size", type=int, default=2) @@ -388,6 +507,10 @@ def main(): ap.add_argument("--port", type=int, default=8770) ap.add_argument("--open", action="store_true", help="open a browser on it") args = ap.parse_args() + if bool(args.waterfalls) != bool(args.model): + ap.error("--waterfalls and --model go together") + if args.waterfalls and not args.predictions: + ap.error("--waterfalls needs --predictions") id_column, columns = schema_columns(args.schema) stamp = fingerprint(args, columns) @@ -424,8 +547,9 @@ def say(what): conn = duckdb.connect(args.cache, read_only=True) conn.execute(f"SET memory_limit = '{args.memory}'") + levels = read_model_levels(args.model) if args.model else None viewer = Viewer(conn, columns, ", ".join(os.path.basename(p) for p in args.data), - args.max_rows, bool(args.truth)) + args.max_rows, bool(args.truth), levels) here = os.path.dirname(os.path.abspath(__file__)) with open(os.path.join(here, "cluster_view_template.html")) as handle: diff --git a/tools/cluster_view.py b/tools/cluster_view.py index c5e1573..aa0f7ad 100755 --- a/tools/cluster_view.py +++ b/tools/cluster_view.py @@ -6,8 +6,11 @@ cpplink predict --schema s.json --model m.json --out predictions.parquet data.parquet cpplink cluster --schema s.json --predictions predictions.parquet \ --out clusters.csv data.parquet + cpplink explain --schema s.json --model m.json --predictions predictions.parquet \ + --out waterfalls.parquet data.parquet tools/cluster_view.py --schema s.json --clusters clusters.csv \ - --predictions predictions.parquet --out clusters.html data.parquet + --predictions predictions.parquet --waterfalls waterfalls.parquet \ + --model m.json --out clusters.html data.parquet Reads the cluster assignment, pulls each member's column values back out of the parquet, and writes one HTML file holding the clusters it selected. Open it in a @@ -20,6 +23,17 @@ the choice is remembered, because the layout is the one quadratic thing the page does and most clusters are read without it. +With `--waterfalls`, every embedded prediction also carries its waterfall: the +prior, then what each comparison charged and its term-frequency move, ending at +the match weight. The page lists the pairs beside the clusters and draws the +ledger for the one picked. The file is what `cpplink explain --predictions + --out ` writes, one wide row per prediction, and nothing here +recomputes a bit of it: the chart is the scorer's own arithmetic, read back. The +level labels and rates come from `--model`, since they are the model's and not +the pair's. Predictions the run made between two records that clustering then +put in different clusters are kept as well, since a pair scored above the write +threshold and below the clustering one is the pair most worth reading. + A run ends with single files, so `--clusters` and `--predictions` each name one file and its extension picks csv or parquet. `--predictions` still takes the shard directory too, which names records by row rather than by `unique_id` and so @@ -296,6 +310,137 @@ def find(x): for uid, i in zip(want.to_pylist(), at) if i is not None} +WATERFALL_FIXED = ["id_a", "id_b", "gamma", "prior", "match_weight", "match_probability", + "bracket_low", "bracket_high", "threshold", "zone", "emitted"] + + +def read_model_levels(path): + """What a ledger needs from the model file: labels and rates per level. + + A waterfall row names the level each comparison landed on; its label, m and + u are the model's rather than the pair's, so the file does not repeat them + and they are read from here once. + """ + with open(path) as handle: + model = json.load(handle) + comparisons = [{"name": c["name"], + "levels": [{"label": lvl.get("label", str(i)), + "m": lvl.get("m"), "u": lvl.get("u")} + for i, lvl in enumerate(c["levels"])]} + for c in model["comparisons"]] + interactions = [f'{t["left"]} x {t["right"]}' for t in model.get("interactions", [])] + return {"records": model.get("records"), "comparisons": comparisons, + "interactions": interactions} + + +def waterfall_batches(path): + """One waterfall file, a batch of rows at a time, whichever format it is.""" + if path.endswith(".parquet"): + handle = pq.ParquetFile(path) + held = handle.schema_arrow.names + if any(name not in held for name in WATERFALL_FIXED): + raise SystemExit(f"{path}: not a cpplink waterfall file") + for batch in handle.iter_batches(batch_size=PREDICTION_BATCH): + yield pa.Table.from_batches([batch]) + return + reader = pacsv.open_csv( + path, convert_options=pacsv.ConvertOptions(column_types=PREDICTION_IDS)) + if any(name not in reader.schema.names for name in WATERFALL_FIXED): + raise SystemExit(f"{path}: not a cpplink waterfall file") + for batch in reader: + yield pa.Table.from_batches([batch]) + + +def read_waterfalls(path, wanted_ids): + """(id_a, id_b) -> the wide row, for the pairs naming two wanted records.""" + wanted = pa.array(sorted(wanted_ids), type=pa.string()) + found, read = {}, 0 + for table in waterfall_batches(path): + read += table.num_rows + table = as_strings(table, ("id_a", "id_b")) + keep = pc.and_(pc.is_in(table["id_a"], value_set=wanted), + pc.is_in(table["id_b"], value_set=wanted)) + for row in table.filter(keep).to_pylist(): + found[(row["id_a"], row["id_b"])] = row + return found, read + + +def ledger(row, levels): + """A wide row expanded to the ledger `cpplink explain --json` writes. + + The same shape for both viewers, so the page has one renderer: the prior, + one step per comparison with a running total, the two-way corrections, then + the totals and the pattern's bracket and zone. + """ + running = row["prior"] + steps = [] + for c in levels["comparisons"]: + name = c["name"] + level = row[f"{name}_level"] + lvl = c["levels"][level] if level < len(c["levels"]) else {"label": str(level)} + bits, tf = row[f"{name}_bits"], row[f"{name}_tf"] + running += bits + tf + steps.append({"name": name, "level": level, "label": lvl["label"], + "m": lvl.get("m"), "u": lvl.get("u"), "bits": bits, "tf": tf, + "frequency": row[f"{name}_frequency"], "running": running}) + interactions = [] + for name in levels["interactions"]: + bits = row[name.replace(" x ", "_x_") + "_bits"] + running += bits + interactions.append({"name": name, "bits": bits, "running": running}) + return {"prior": row["prior"], "records": levels["records"], + "threshold": row["threshold"], "steps": steps, + "interactions": interactions, "weight": row["match_weight"], + "probability": row["match_probability"], + "bracket": [row["bracket_low"], row["bracket_high"]], + "zone": row["zone"], "emitted": bool(row["emitted"])} + + +class Ledger: + """The waterfalls compacted for embedding. + + A level's label, rates and bits are the model's, not the pair's, so they are + kept once per (comparison, level) and a pair stores the level it hit, its + term-frequency move and the frequency behind it. The page expands that back + to the full ledger before drawing, so both viewers draw the same shape. + """ + + def __init__(self, levels): + self.levels = levels + self.model = None + self.comparisons = [{"name": c["name"], "levels": []} + for c in levels["comparisons"]] + self.count = 0 + + def add(self, row): + w = ledger(row, self.levels) + if self.model is None: + self.model = {"prior": w["prior"], "records": w["records"], + "threshold": w["threshold"], + "interactions": self.levels["interactions"]} + steps = [] + for c, step in enumerate(w["steps"]): + table = self.comparisons[c]["levels"] + while len(table) <= step["level"]: + table.append(None) + if table[step["level"]] is None: + table[step["level"]] = {"label": step["label"], "bits": step["bits"], + "m": step["m"], "u": step["u"]} + entry = [step["level"]] + if step["tf"]: + entry += [step["tf"], step["frequency"]] + steps.append(entry) + self.count += 1 + return {"s": steps, "i": [t["bits"] for t in w["interactions"]], + "w": w["weight"], "p": w["probability"], "b": w["bracket"], + "z": w["zone"]} + + def summary(self): + if self.model is None: + return None + return dict(self.model, comparisons=self.comparisons) + + def main(): ap = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter) @@ -318,7 +463,14 @@ def main(): ap.add_argument("--seed", type=int, default=1) ap.add_argument("--id", action="append", default=[], help="always embed this cluster id; repeatable") + ap.add_argument("--waterfalls", help="the csv or parquet cpplink explain " + "--predictions wrote; embeds each " + "prediction's waterfall") + ap.add_argument("--model", help="the model the run scored with, for the level " + "labels and rates the waterfall shows") args = ap.parse_args() + if bool(args.waterfalls) != bool(args.model): + ap.error("--waterfalls and --model go together") id_column, columns = schema_columns(args.schema) @@ -341,18 +493,27 @@ def main(): truth = read_truth(args.truth, wanted_ids) if args.truth else None predictions, predictions_read = ({}, 0) by_cluster = {} + rejected = [] if args.predictions: predictions, predictions_read = read_predictions( args.predictions, wanted_ids, {u: rows_by_id[u] for u in wanted_ids if u in rows_by_id}) # A record is in one cluster, so grouping the predictions once is what # keeps the loop below linear in them rather than one pass per cluster. + # A prediction whose ends clustering kept apart is kept too: it was + # scored above the write threshold and below the clustering one, and + # that is the pair a reader most wants explained. cluster_of = {uid: cluster for cluster, uids in members.items() for uid in uids} for (a, b), weight in predictions.items(): cluster = cluster_of.get(a) - if cluster is not None and cluster == cluster_of.get(b): + other = cluster_of.get(b) + if cluster is None or other is None or a not in by_id or b not in by_id: + continue + if cluster == other: by_cluster.setdefault(cluster, []).append((a, b, weight)) + else: + rejected.append((a, b, weight, cluster, other)) payload_clusters = [] for cluster in chosen: @@ -366,6 +527,10 @@ def main(): if truth is not None: group = truth.get(uid, "\x00" + uid) row["t"] = groups_seen.setdefault(group, len(groups_seen)) + # The label is per cluster; the group is what says "same entity" + # for a pair the clustering split across two of them. + if uid in truth: + row["g"] = truth[uid] rows.append(row) entry = {"id": cluster, "rows": rows} if sizes[cluster] > len(rows): @@ -380,12 +545,38 @@ def main(): entry["edges"] = sorted(found) payload_clusters.append(entry) + payload_rejected = [{"a": a, "b": b, "w": round(weight, 3), "ca": ca, "cb": cb} + for a, b, weight, ca, cb in sorted(rejected)] + + # The waterfalls of the embedded predictions, read from the file the run's + # explain wrote and compacted against the model's level table. + ledgers = None + explained = waterfalls_read = 0 + if args.waterfalls and args.predictions: + ledgers = Ledger(read_model_levels(args.model)) + found, waterfalls_read = read_waterfalls(args.waterfalls, wanted_ids) + for entry in payload_clusters: + uid_at = [r["id"] for r in entry["rows"]] + for edge in entry.get("edges", []): + row = found.get((uid_at[edge[0]], uid_at[edge[1]])) + if row is not None: + edge.append(ledgers.add(row)) + explained += 1 + for pair in payload_rejected: + row = found.get((pair["a"], pair["b"])) + if row is not None: + pair["wf"] = ledgers.add(row) + explained += 1 + payload = { "source": ", ".join(os.path.basename(p) for p in args.data), "generated": datetime.datetime.now().isoformat(timespec="seconds"), "columns": columns, - "totals": {"clusters": total_clusters, "shown": len(payload_clusters)}, + "totals": {"clusters": total_clusters, "shown": len(payload_clusters), + "records": records}, "clusters": payload_clusters, + "rejected": payload_rejected, + "model": ledgers.summary() if ledgers else None, } here = os.path.dirname(os.path.abspath(__file__)) @@ -401,8 +592,11 @@ def main(): print(f"{len(wanted_ids):,} records read from {touched:,} of {groups:,} row " f"groups over {records:,} rows") if args.predictions: - print(f"{len(predictions):,} within-cluster predictions carried, " - f"{predictions_read:,} scanned") + print(f"{sum(len(c.get('edges', [])) for c in payload_clusters):,} " + f"within-cluster predictions carried and {len(payload_rejected):,} " + f"across clusters, {predictions_read:,} scanned") + if args.waterfalls: + print(f"{explained:,} waterfalls embedded, {waterfalls_read:,} scanned") if __name__ == "__main__": diff --git a/tools/cluster_view_template.html b/tools/cluster_view_template.html index cb1429b..5a75bbf 100644 --- a/tools/cluster_view_template.html +++ b/tools/cluster_view_template.html @@ -90,6 +90,13 @@ background: var(--panel); display: flex; flex-direction: column; min-height: 0; } .controls { padding: 10px; border-bottom: 1px solid var(--line-soft); display: grid; gap: 8px; } +.modes { display: flex; border: 1px solid var(--line); border-radius: 6px; overflow: hidden; } +.modes button { + flex: 1; font: inherit; font-size: 12px; padding: 4px 0; border: none; cursor: pointer; + background: var(--bg); color: var(--muted); +} +.modes button.on { background: var(--accent-soft); color: var(--accent); font-weight: 600; } +.modes button:disabled { opacity: 0.45; cursor: default; } .controls input, .controls select { font: inherit; font-size: 13px; width: 100%; padding: 5px 8px; color: var(--ink); background: var(--bg); border: 1px solid var(--line); border-radius: 6px; @@ -106,6 +113,7 @@ .item.more { justify-content: center; color: var(--muted); font-size: 12px; } .item.more .cid { flex: none; } .item .cid { font-family: var(--mono); font-size: 12px; flex: 1; overflow: hidden; text-overflow: ellipsis; white-space: nowrap; } +.item .cid .sep { color: var(--muted); padding: 0 3px; } .item .sz { font-size: 11px; color: var(--muted); font-variant-numeric: tabular-nums; } .pill { font-size: 10px; padding: 1px 5px; border-radius: 999px; white-space: nowrap; @@ -152,6 +160,27 @@ .matrix td.weak { background: var(--warn-soft); color: var(--warn); } .matrix th.rowhead { text-align: left; font-family: var(--mono); text-transform: none; letter-spacing: 0; } .note { color: var(--muted); font-size: 12px; } +.links { display: flex; gap: 8px; flex-wrap: wrap; margin-bottom: 14px; } +.links button.ghost { font-family: var(--mono); } +.matrix td.pair, .matrix td.pair:hover { cursor: pointer; } +.matrix td.pair:hover { background: var(--sel); color: var(--ink); } +.wf { width: 100%; height: auto; display: block; background: var(--panel); border: 1px solid var(--line); border-radius: 8px; } +.wf text { font-family: var(--mono); font-size: 11px; fill: var(--ink); } +.wf text.dim { fill: var(--muted); } +.wf text.tick { font-size: 10px; fill: var(--muted); text-anchor: middle; } +.wf text.val { font-size: 10.5px; } +.wf text.name { font-weight: 600; } +.wf text.lvl { fill: var(--muted); text-anchor: end; font-size: 10px; } +.wf line.grid { stroke: var(--line-soft); } +.wf line.zero { stroke: var(--muted); stroke-dasharray: 3 3; } +.wf line.thr { stroke: var(--accent); stroke-dasharray: 4 3; stroke-width: 1.5; } +.wf line.step { stroke: var(--line); } +.wf rect.up { fill: var(--accent); } +.wf rect.down { fill: var(--bad); } +.wf rect.tf { fill: var(--warn); } +.wf rect.total { fill: var(--ink); opacity: 0.85; } +.wf rect.prior { fill: var(--muted); } +.wf g.bar:hover rect { opacity: 0.7; } .net { width: 100%; height: auto; max-height: 55vh; display: block; background: var(--panel); border: 1px solid var(--line); border-radius: 8px; @@ -196,29 +225,23 @@

cpplink clusters

+
+ + +
- +
- +
-
Pick a cluster.
+
Pick a cluster or a pair.
From 1a4b0a1ca454ac0479235303e0e84b07278b6eac Mon Sep 17 00:00:00 2001 From: 4ment Date: Wed, 16 Sep 2026 21:36:57 +1000 Subject: [PATCH 05/10] Move pairs tab to right panel --- tools/cluster_server.py | 66 ++--- tools/cluster_view.py | 22 +- tools/cluster_view_template.html | 469 +++++++++++++++++++------------ 3 files changed, 338 insertions(+), 219 deletions(-) diff --git a/tools/cluster_server.py b/tools/cluster_server.py index 73d3a18..a584f36 100755 --- a/tools/cluster_server.py +++ b/tools/cluster_server.py @@ -26,7 +26,7 @@ a click is one lookup, so the ledger is the scorer's own arithmetic and nothing is computed or spawned here. The level labels and rates come from `--model`. Predictions clustering kept apart -- above the write threshold, below the -clustering one -- are kept and listed as `rejected`. +clustering one -- are kept and listed under both of their clusters as `rejected`. """ import argparse @@ -49,7 +49,7 @@ read_model_levels, read_truth, schema_columns) EDGE_CHUNK = 1 << 20 -CACHE_VERSION = 5 +CACHE_VERSION = 6 def quoted(name): @@ -191,7 +191,9 @@ def build(conn, args, id_column, columns, say): WHERE cluster_a = cluster_b """) conn.execute("CREATE INDEX edges_by_cluster ON edges (cluster_id)") - conn.execute("CREATE INDEX pairs_by_weight ON pairs (weight)") + # A rejected pair is asked for from either of its clusters. + conn.execute("CREATE INDEX pairs_by_cluster_a ON pairs (cluster_a)") + conn.execute("CREATE INDEX pairs_by_cluster_b ON pairs (cluster_b)") conn.execute("CREATE INDEX pairs_by_a ON pairs (a)") conn.execute("CREATE INDEX pairs_by_b ON pairs (b)") if args.waterfalls: @@ -268,17 +270,7 @@ def build(conn, args, id_column, columns, say): "mixed": "entities > 1", } -PAIR_ORDERS = { - "weakest": "weight, a, b", - "strongest": "weight DESC, a, b", - "id": "a, b", -} - -PAIR_FILTERS = { - "all": "TRUE", - "within": "cluster_a = cluster_b", - "rejected": "cluster_a <> cluster_b", -} +REJECTED_LIMIT = 1000 # rejected predictions listed per cluster, weakest first class Viewer: @@ -322,25 +314,27 @@ def waterfall(self, a, b): return None return ledger(dict(zip(self.waterfall_columns, rows[0])), self.levels) - def pairs(self, q, sort, only, offset, limit): - where = PAIR_FILTERS.get(only, "TRUE") - params = [] - if q: - # Either record's text, since a pair is read by either of its ends. - like = ("%" + q.lower().replace("\\", "\\\\") - .replace("%", "\\%").replace("_", "\\_") + "%") - where += (" AND (a IN (SELECT uid FROM records WHERE text LIKE ? " - "ESCAPE '\\') OR b IN (SELECT uid FROM records WHERE text " - "LIKE ? ESCAPE '\\'))") - params += [like, like] - matched = self.query(f"SELECT count(*) FROM pairs WHERE {where}", params)[0][0] + def rejected(self, cluster_id): + """The predictions clustering overruled that touch one cluster. + + Each names the other cluster its second record went to, and the truth's + verdict where there is one, so the page can list them beside the + cluster's own edges without a lookup per row. + """ rows = self.query( - f"SELECT a, b, weight, cluster_a, cluster_b FROM pairs WHERE {where} " - f"ORDER BY {PAIR_ORDERS.get(sort, PAIR_ORDERS['weakest'])} " - f"LIMIT {int(limit)} OFFSET {int(offset)}", params) - return {"matched": matched, - "pairs": [{"a": r[0], "b": r[1], "w": round(r[2], 3), - "ca": r[3], "cb": r[4]} for r in rows]} + "SELECT p.a, p.b, p.weight, p.cluster_a, p.cluster_b, ta.grp, tb.grp " + "FROM pairs p LEFT JOIN truth ta ON ta.uid = p.a " + "LEFT JOIN truth tb ON tb.uid = p.b " + "WHERE (p.cluster_a = ? OR p.cluster_b = ?) AND p.cluster_a <> p.cluster_b " + f"ORDER BY p.weight, p.a, p.b LIMIT {REJECTED_LIMIT + 1}", + [cluster_id, cluster_id]) + more = max(0, len(rows) - REJECTED_LIMIT) + found = [] + for a, b, w, ca, cb, ga, gb in rows[:REJECTED_LIMIT]: + same = (ga is not None and ga == gb) if self.has_truth else None + found.append({"a": a, "b": b, "w": round(w, 3), + "other": cb if ca == cluster_id else ca, "same": same}) + return found, more def pair(self, a, b): head = self.query("SELECT a, b, weight, cluster_a, cluster_b FROM pairs " @@ -423,6 +417,9 @@ def cluster(self, cluster_id): if self.has_truth: for row in answer["rows"]: row["t"] = labels[groups.get(row["id"], "x" + row["id"])] + answer["rejected"], more = self.rejected(cluster_id) + if more: + answer["rejected_more"] = more return answer @@ -461,11 +458,6 @@ def do_GET(self): found = viewer.cluster(query.get("id", "")) self.send_json(found or {"error": "no such cluster"}, 200 if found else 404) - elif url.path == "/api/pairs": - self.send_json(viewer.pairs( - query.get("q", ""), query.get("sort", "weakest"), - query.get("only", "all"), int(query.get("offset", 0)), - min(int(query.get("limit", 100)), 500))) elif url.path == "/api/pair": found = viewer.pair(query.get("a", ""), query.get("b", "")) self.send_json(found or {"error": "no such prediction"}, diff --git a/tools/cluster_view.py b/tools/cluster_view.py index aa0f7ad..33dce91 100755 --- a/tools/cluster_view.py +++ b/tools/cluster_view.py @@ -14,9 +14,11 @@ Reads the cluster assignment, pulls each member's column values back out of the parquet, and writes one HTML file holding the clusters it selected. Open it in a -browser: the left pane lists clusters, the right one shows the members side by -side with every disagreeing cell highlighted, so what a cluster has in common is -the part that is not highlighted. +browser: the left pane lists clusters and the right one has two tabs for the +cluster picked. The Cluster tab shows the members side by side with every +disagreeing cell highlighted, so what a cluster has in common is the part that is +not highlighted; the Pairs tab lists every prediction touching the cluster and +draws the one picked. The `network` checkbox in the header draws the selected cluster's predictions as a graph, which is where a chain shows itself as a chain. It is off by default and @@ -25,13 +27,13 @@ With `--waterfalls`, every embedded prediction also carries its waterfall: the prior, then what each comparison charged and its term-frequency move, ending at -the match weight. The page lists the pairs beside the clusters and draws the -ledger for the one picked. The file is what `cpplink explain --predictions - --out ` writes, one wide row per prediction, and nothing here -recomputes a bit of it: the chart is the scorer's own arithmetic, read back. The -level labels and rates come from `--model`, since they are the model's and not -the pair's. Predictions the run made between two records that clustering then -put in different clusters are kept as well, since a pair scored above the write +the match weight, drawn under the pairs table for the pair picked. The file is +what `cpplink explain --predictions --out ` writes, one wide row +per prediction, and nothing here recomputes a bit of it: the chart is the +scorer's own arithmetic, read back. The level labels and rates come from +`--model`, since they are the model's and not the pair's. Predictions the run +made between two records that clustering then put in different clusters are +kept as well and listed under both clusters, since a pair scored above the write threshold and below the clustering one is the pair most worth reading. A run ends with single files, so `--clusters` and `--predictions` each name one diff --git a/tools/cluster_view_template.html b/tools/cluster_view_template.html index 5a75bbf..874083c 100644 --- a/tools/cluster_view_template.html +++ b/tools/cluster_view_template.html @@ -90,13 +90,6 @@ background: var(--panel); display: flex; flex-direction: column; min-height: 0; } .controls { padding: 10px; border-bottom: 1px solid var(--line-soft); display: grid; gap: 8px; } -.modes { display: flex; border: 1px solid var(--line); border-radius: 6px; overflow: hidden; } -.modes button { - flex: 1; font: inherit; font-size: 12px; padding: 4px 0; border: none; cursor: pointer; - background: var(--bg); color: var(--muted); -} -.modes button.on { background: var(--accent-soft); color: var(--accent); font-weight: 600; } -.modes button:disabled { opacity: 0.45; cursor: default; } .controls input, .controls select { font: inherit; font-size: 13px; width: 100%; padding: 5px 8px; color: var(--ink); background: var(--bg); border: 1px solid var(--line); border-radius: 6px; @@ -123,7 +116,35 @@ .pill.ok { background: var(--accent-soft); color: var(--accent); } #detail { flex: 1; overflow: auto; padding: 16px 20px 60px; min-width: 0; } #detail h2 { font-family: var(--mono); font-size: 15px; font-weight: 600; margin: 0 0 2px; } -.sub { color: var(--muted); font-size: 12px; margin-bottom: 16px; } +.sub { color: var(--muted); font-size: 12px; margin-bottom: 12px; } +.tabs { display: flex; border-bottom: 1px solid var(--line); margin-bottom: 18px; } +.tabs button { + font: inherit; font-size: 13px; padding: 7px 14px; border: none; background: none; + border-bottom: 2px solid transparent; margin-bottom: -1px; color: var(--muted); cursor: pointer; +} +.tabs button:hover { color: var(--ink); } +.tabs button.on { color: var(--accent); border-bottom-color: var(--accent); font-weight: 600; } +.tabs .count { + font-size: 11px; font-weight: 400; color: var(--muted); background: var(--line-soft); + border-radius: 999px; padding: 0 6px; margin-left: 5px; font-variant-numeric: tabular-nums; +} +.tabs .count.warn { background: var(--warn-soft); color: var(--warn); } +.pairlist { max-height: 38vh; overflow: auto; } +table.pairs tr.row { cursor: pointer; } +table.pairs tr.row:hover td { background: var(--sel); } +table.pairs tr.row.pin td { background: var(--accent-soft); box-shadow: none; } +table.pairs td.n { color: var(--muted); font-variant-numeric: tabular-nums; text-align: right; } +table.pairs td.rej { color: var(--warn); } +table.pairs td.rej button.ghost { font-family: var(--mono); font-size: 11px; padding: 1px 6px; } +table.pairs td.same { color: var(--accent); } +table.pairs td.differ { color: var(--bad); } +th.sortable { cursor: pointer; user-select: none; } +th.sortable:hover { color: var(--ink); } +th.sortable.on { color: var(--ink); } +th.sortable.on::after { content: " \2193"; } +th.sortable.on.desc::after { content: " \2191"; } +.pairhead { display: flex; align-items: baseline; gap: 10px; flex-wrap: wrap; margin: 18px 0 4px; } +.pairhead h3 { font-family: var(--mono); font-size: 14px; font-weight: 600; margin: 0; } .section { margin-bottom: 22px; } .section > h3 { font-size: 11px; text-transform: uppercase; letter-spacing: 0.06em; color: var(--muted); @@ -160,7 +181,7 @@ .matrix td.weak { background: var(--warn-soft); color: var(--warn); } .matrix th.rowhead { text-align: left; font-family: var(--mono); text-transform: none; letter-spacing: 0; } .note { color: var(--muted); font-size: 12px; } -.links { display: flex; gap: 8px; flex-wrap: wrap; margin-bottom: 14px; } +.links { display: flex; gap: 8px; flex-wrap: wrap; margin-bottom: 12px; } .links button.ghost { font-family: var(--mono); } .matrix td.pair, .matrix td.pair:hover { cursor: pointer; } .matrix td.pair:hover { background: var(--sel); color: var(--ink); } @@ -216,7 +237,7 @@

cpplink clusters

- j k move · / search + j k cluster · n p pair · / search @@ -225,10 +246,6 @@

cpplink clusters

-
- - -
@@ -241,7 +258,7 @@

cpplink clusters

-
Pick a cluster or a pair.
+
Pick a cluster.
From 35257098bbd3f886a8192b1cf7088add6c7c3c03 Mon Sep 17 00:00:00 2001 From: 4ment Date: Thu, 17 Sep 2026 06:53:36 +1000 Subject: [PATCH 06/10] Add threshold filtering and sorting options to cluster viewer --- tools/cluster_server.py | 94 ++++++++++--- tools/cluster_view.py | 6 + tools/cluster_view_template.html | 229 +++++++++++++++++++++++++++---- 3 files changed, 281 insertions(+), 48 deletions(-) diff --git a/tools/cluster_server.py b/tools/cluster_server.py index a584f36..18806af 100755 --- a/tools/cluster_server.py +++ b/tools/cluster_server.py @@ -73,7 +73,8 @@ def scan(path, types): def fingerprint(args, columns): """What the cache was built from, so a changed input rebuilds it.""" - parts = [CACHE_VERSION, columns, args.max_rows] + parts = [CACHE_VERSION, columns, args.max_rows, args.min_size, args.max_size, + args.threshold] for path in [args.clusters, args.truth, args.waterfalls] + list(args.data): if path: parts.append([os.path.abspath(path), os.path.getmtime(path), @@ -123,13 +124,14 @@ def build(conn, args, id_column, columns, say): picked = ", ".join(quoted(c) for c in columns) say("reading the cluster assignment") + ceiling = f" AND cluster_size <= {args.max_size}" if args.max_size else "" conn.execute(f""" CREATE TABLE members AS SELECT CAST(unique_id AS VARCHAR) AS uid, CAST(cluster_id AS VARCHAR) AS cluster_id, CAST(cluster_size AS BIGINT) AS size FROM {scan(args.clusters, ("unique_id", "cluster_id"))} - WHERE cluster_size >= {args.min_size} + WHERE cluster_size >= {args.min_size}{ceiling} """) say("reading the clustered records") @@ -175,14 +177,17 @@ def build(conn, args, id_column, columns, say): # outside every one. A pair across two clusters is kept with both # cluster ids, since it is the prediction clustering overruled and the # one most worth explaining; a pair with an end no cluster holds has no - # record to show and is dropped. - conn.execute(""" + # record to show and is dropped. `--threshold` drops the predictions + # below it here, so the cache never holds them. + floor = (f" WHERE e.weight >= {args.threshold}" + if args.threshold is not None else "") + conn.execute(f""" CREATE TABLE pairs AS SELECT e.a, e.b, e.weight, ra.cluster_id AS cluster_a, rb.cluster_id AS cluster_b FROM named_edges e JOIN records ra ON ra.uid = e.a - JOIN records rb ON rb.uid = e.b + JOIN records rb ON rb.uid = e.b{floor} """) conn.execute("DROP TABLE named_edges") conn.execute(""" @@ -256,13 +261,27 @@ def build(conn, args, id_column, columns, say): conn.execute("CREATE INDEX records_by_uid ON records (uid)") +# The list's order: the chosen key in either direction, a cluster with no +# value for it (no predictions, so no weakest edge) last either way, and the +# ties broken the same way whichever direction the key runs. ORDERS = { - "discord": "discord DESC, size DESC, id", - "size": "size DESC, discord DESC, id", - "weakest": "weakest NULLS LAST, size DESC, id", - "id": "id", + "discord": ("discord", "size DESC, id"), + "size": ("size", "discord DESC, id"), + "weakest": ("weakest", "size DESC, id"), + "id": ("id", ""), } + +def order_clause(sort, desc): + key, ties = ORDERS.get(sort, ORDERS["discord"]) + clause = f"{key} {'DESC' if desc else 'ASC'} NULLS LAST" + return clause + (", " + ties if ties else "") + + +def like_pattern(text): + return "%" + text.lower().replace("\\", "\\\\").replace("%", "\\%").replace("_", "\\_") + "%" + + FILTERS = { "all": "TRUE", "split": "discord > 0", @@ -361,23 +380,51 @@ def pair(self, a, b): answer["waterfall"] = self.waterfall(a, b) return answer - def listing(self, q, sort, only, offset, limit): + def listing(self, q, field, column, sort, desc, only, offset, limit): + """One page of the cluster list. + + `field` says what `q` is read against: `any` is a substring of the + concatenated record, `cluster` of the cluster id, `col` of the named + column, and `id` a comma-separated list of record ids matched whole, in + which case the answer also says which of them named a record. + """ where = FILTERS.get(only, "TRUE") params = [] - if q: + found = None + if q and field == "id": + ids = list(dict.fromkeys(s.strip() for s in q.split(",") if s.strip())) + marks = ", ".join("?" * len(ids)) + where += (f" AND id IN (SELECT cluster_id FROM records " + f"WHERE CAST(uid AS VARCHAR) IN ({marks}))") + params.extend(ids) + found = [r[0] for r in self.query( + f"SELECT DISTINCT CAST(uid AS VARCHAR) FROM records " + f"WHERE CAST(uid AS VARCHAR) IN ({marks})", ids)] + elif q and field == "cluster": + where += " AND lower(id) LIKE ? ESCAPE '\\'" + params.append(like_pattern(q)) + elif q and field == "col": + if column not in self.columns: + raise ValueError(f"no column {column!r}") + where += (f" AND id IN (SELECT cluster_id FROM records " + f"WHERE lower(CAST({quoted(column)} AS VARCHAR)) LIKE ? ESCAPE '\\')") + params.append(like_pattern(q)) + elif q: where += (" AND id IN (SELECT cluster_id FROM records " "WHERE text LIKE ? ESCAPE '\\')") - params.append("%" + q.lower().replace("\\", "\\\\") - .replace("%", "\\%").replace("_", "\\_") + "%") + params.append(like_pattern(q)) matched = self.query(f"SELECT count(*) FROM stats WHERE {where}", params)[0][0] rows = self.query( f"SELECT id, size, discord, edge_count, weakest, entities FROM stats " - f"WHERE {where} ORDER BY {ORDERS.get(sort, ORDERS['discord'])} " + f"WHERE {where} ORDER BY {order_clause(sort, desc)} " f"LIMIT {int(limit)} OFFSET {int(offset)}", params) - return {"matched": matched, - "clusters": [{"id": r[0], "size": r[1], "discord": r[2], - "edge_count": r[3], "weakest": r[4], "entities": r[5]} - for r in rows]} + answer = {"matched": matched, + "clusters": [{"id": r[0], "size": r[1], "discord": r[2], + "edge_count": r[3], "weakest": r[4], "entities": r[5]} + for r in rows]} + if found is not None: + answer["found"] = found + return answer def cluster(self, cluster_id): head = self.query("SELECT id, size, discord, edge_count, weakest, entities " @@ -451,9 +498,10 @@ def do_GET(self): self.send_json(viewer.summary()) elif url.path == "/api/list": self.send_json(viewer.listing( - query.get("q", ""), query.get("sort", "discord"), - query.get("only", "all"), int(query.get("offset", 0)), - min(int(query.get("limit", 100)), 500))) + query.get("q", ""), query.get("field", "any"), + query.get("column", ""), query.get("sort", "discord"), + query.get("dir", "desc") == "desc", query.get("only", "all"), + int(query.get("offset", 0)), min(int(query.get("limit", 100)), 500))) elif url.path == "/api/cluster": found = viewer.cluster(query.get("id", "")) self.send_json(found or {"error": "no such cluster"}, @@ -491,6 +539,10 @@ def main(): ap.add_argument("--cache", default="cluster_view.duckdb") ap.add_argument("--rebuild", action="store_true", help="rebuild the cache first") ap.add_argument("--min-size", type=int, default=2) + ap.add_argument("--max-size", type=int, default=0, help="0 = no ceiling") + ap.add_argument("--threshold", type=float, + help="keep only the predictions whose match_weight is at " + "least this; the default keeps every one the run wrote") ap.add_argument("--max-rows", type=int, default=200, help="members shown per cluster; the rest are counted only") ap.add_argument("--memory", default="4GB", diff --git a/tools/cluster_view.py b/tools/cluster_view.py index 33dce91..26b9001 100755 --- a/tools/cluster_view.py +++ b/tools/cluster_view.py @@ -460,6 +460,9 @@ def main(): ap.add_argument("--max-rows", type=int, default=200, help="members embedded per cluster; the rest are counted only") ap.add_argument("--max-size", type=int, default=0, help="0 = no ceiling") + ap.add_argument("--threshold", type=float, + help="keep only the predictions whose match_weight is at " + "least this; the default keeps every one the run wrote") ap.add_argument("--sort", default="size", choices=["size", "id", "random"], help="which clusters to embed when --limit cuts the list") ap.add_argument("--seed", type=int, default=1) @@ -500,6 +503,9 @@ def main(): predictions, predictions_read = read_predictions( args.predictions, wanted_ids, {u: rows_by_id[u] for u in wanted_ids if u in rows_by_id}) + if args.threshold is not None: + predictions = {pair: weight for pair, weight in predictions.items() + if weight >= args.threshold} # A record is in one cluster, so grouping the predictions once is what # keeps the loop below linear in them rather than one pass per cluster. # A prediction whose ends clustering kept apart is kept too: it was diff --git a/tools/cluster_view_template.html b/tools/cluster_view_template.html index 874083c..ffd6cf5 100644 --- a/tools/cluster_view_template.html +++ b/tools/cluster_view_template.html @@ -95,6 +95,11 @@ background: var(--bg); border: 1px solid var(--line); border-radius: 6px; } .controls .row { display: flex; gap: 8px; align-items: center; } +.controls .row select, .controls .row input { min-width: 0; } +.controls #field { width: 118px; flex: 0 0 118px; } +.controls #dir { flex: none; padding: 4px 8px; font-family: var(--mono); } +.controls #qnote { font-size: 11px; } +.controls #qnote .missing { color: var(--warn); } .controls label { font-size: 11px; color: var(--muted); text-transform: uppercase; letter-spacing: 0.04em; } #list { overflow-y: auto; flex: 1; } .item { @@ -219,6 +224,8 @@ .net g.node:hover circle { stroke: var(--ink); stroke-width: 2; } tr.lit td, tr.pin td { background: var(--accent-soft); } tr.lit td.odd, tr.pin td.odd { background: var(--warn-soft); } +/* what the search matched: a record the id search named, or a cell a column search hit */ +tr.hit td.id, td.val.hit { box-shadow: inset 0 -2px 0 var(--accent); } tr.pin td { box-shadow: inset 0 1px 0 var(--accent), inset 0 -1px 0 var(--accent); } .toggle { font-size: 12px; color: var(--muted); display: flex; gap: 5px; @@ -246,10 +253,15 @@

cpplink clusters

-
- +
+ + +
+
+
@@ -286,11 +298,18 @@

cpplink clusters

let tab = 'cluster'; // which tab of the detail pane is open: cluster or pairs let pick = null; // {a, b}: the pair open on the pairs tab let pairSort = { key: 'w', desc: false }; // the pairs table's order, kept across clusters +let sortDesc = true; // the list's direction; each sort key has its own default +let search = { field: 'any', q: '', ids: [], col: -1 }; // the search as last parsed let details = new Map(); let token = 0; // guards against a slow response for an old query -const SORTS = [['discord', 'most disagreement'], ['size', 'size'], - ['weakest', 'weakest edge'], ['id', 'id']]; +// Sort keys with the direction each starts in: the fragile clusters lead on +// edge strength, the big and the split ones on size and disagreement. +const SORTS = [['discord', 'most disagreement', true], ['size', 'size', true], + ['weakest', 'edge strength', false], ['id', 'cluster id', false]]; +// The field a search reads: everything, a record id (a comma-separated list of +// them, matched whole), the cluster id, or one column, each as a substring. +const FIELDS = [['any', 'any field'], ['id', 'record id'], ['cluster', 'cluster id']]; const ONLYS = [['all', 'every cluster'], ['split', 'columns that disagree'], ['transitive', 'transitive-only links'], ['mixed', 'mixed against truth']]; const pairKey = (a, b) => a + '\u0000' + b; @@ -478,48 +497,141 @@

cpplink clusters

if (reset) { entries = []; exhausted = false; matched = 0; } if (exhausted) return; const mine = ++token; - const query = { q: $('#q').value.trim(), sort: $('#sort').value, - only: $('#only').value, offset: entries.length, limit: PAGE }; + search = parseSearch(); + const query = { field: search.field, q: search.q, sort: $('#sort').value, + dir: sortDesc ? 'desc' : 'asc', only: $('#only').value, + offset: entries.length, limit: PAGE }; + if (search.field === 'col') query.column = COLS[search.col]; + let found = null; // id search: the ids that name a record if (LIVE) { const page = await api('/api/list', query); if (mine !== token) return; entries = entries.concat(page.clusters); matched = page.matched; exhausted = entries.length >= matched; + found = page.found || null; } else { const all = EMBEDDED.clusters.filter(c => passes(c, query)); matched = all.length; - sortLocal(all, query.sort); + sortLocal(all, query.sort, sortDesc); entries = all.slice(0, entries.length + PAGE); exhausted = entries.length >= matched; + if (search.field === 'id') found = search.ids.filter(id => EMBEDDED.records.has(id)); } + searchNote(found); +} + +// The search box against the field selector: the text lowered for a substring +// match, or split on commas into whole ids when the field is the record id. +function parseSearch() { + const field = $('#field').value; + const text = $('#q').value.trim(); + const out = { field: 'any', q: text.toLowerCase(), ids: [], col: -1 }; + if (!text) return out; + if (field === 'id') { + out.field = 'id'; + out.ids = [...new Set(text.split(',').map(s => s.trim()).filter(s => s))]; + out.q = out.ids.join(','); + } else if (field === 'cluster') { + out.field = 'cluster'; + } else if (field.startsWith('col:')) { + out.field = 'col'; + out.col = +field.slice(4); + } + return out; +} + +// Under the box, for an id search: how many of the ids named a record, and +// which did not, since "nothing matches" on five ids hides which one is wrong. +function searchNote(found) { + const note = $('#qnote'); + if (search.field !== 'id' || !search.ids.length || found === null) { + note.hidden = true; + return; + } + const have = new Set(found); + const missing = search.ids.filter(id => !have.has(id)); + const n = search.ids.length; + let text = `${n} id${n === 1 ? '' : 's'} · ${have.size} found`; + if (missing.length) { + text += ` · missing: ${missing.map(esc).join(', ')}`; + } + note.innerHTML = text; + note.hidden = false; +} + +// `field:text` typed into the box picks the field, so a keyboard user need not +// reach for the selector; only a prefix that names a field counts, so a value +// holding a colon is left alone. +function shortcut() { + const box = $('#q'); + const m = box.value.match(/^\s*([^:]+?)\s*:(.*)$/); + if (!m) return; + const name = m[1].toLowerCase(); + const alias = { any: 'any', id: 'id', 'record id': 'id', cluster: 'cluster', + 'cluster id': 'cluster' }; + let value = alias[name]; + if (value === undefined) { + const i = COLS.findIndex(c => c.toLowerCase() === name); + if (i < 0) return; + value = `col:${i}`; + } + $('#field').value = value; + box.value = m[2].replace(/^\s+/, ''); + placeholder(); +} + +function placeholder() { + const field = $('#field').value; + $('#q').placeholder = field === 'id' ? 'record ids, comma separated' + : field === 'cluster' ? 'cluster id contains' + : field.startsWith('col:') ? `${COLS[+field.slice(4)]} contains` + : 'search any value'; } function passes(c, query) { if (query.only === 'split' && !c.discord) return false; if (query.only === 'transitive' && !(c.size * (c.size - 1) / 2 > c.edge_count)) return false; if (query.only === 'mixed' && !(c.entities > 1)) return false; - return !query.q || c.hay.includes(query.q.toLowerCase()); + if (!query.q) return true; + switch (query.field) { + case 'id': return c.rows.some(r => search.ids.includes(r.id)); + case 'cluster': return c.id.toLowerCase().includes(query.q); + case 'col': return c.rows.some(r => cellHit(r.v[search.col], query.q)); + default: return c.hay.includes(query.q); + } +} + +const cellHit = (v, q) => v !== null && v !== undefined && String(v).toLowerCase().includes(q); + +function compare(ka, kb) { + for (let i = 0; i < ka.length; i++) { + if (ka[i] < kb[i]) return -1; + if (ka[i] > kb[i]) return 1; + } + return 0; } function sortBy(list, key) { - list.sort((a, b) => { - const ka = key(a), kb = key(b); - for (let i = 0; i < ka.length; i++) { - if (ka[i] < kb[i]) return -1; - if (ka[i] > kb[i]) return 1; - } - return 0; - }); + list.sort((a, b) => compare(key(a), key(b))); } -function sortLocal(list, sort) { - sortBy(list, { - discord: c => [-c.discord, -c.size], - size: c => [-c.size, -c.discord], - weakest: c => [c.weakest === null ? Infinity : c.weakest, -c.size], - id: c => [c.id], - }[sort]); +// The list's order: the chosen key in the chosen direction, a cluster with no +// value for it (no predictions, so no weakest edge) last either way, and the +// ties broken the same way whichever direction the key runs. +function sortLocal(list, sort, desc) { + const primary = { discord: c => c.discord, size: c => c.size, + weakest: c => c.weakest, id: c => c.id }[sort]; + const ties = { discord: c => [-c.size, c.id], size: c => [-c.discord, c.id], + weakest: c => [-c.size, c.id], id: c => [] }[sort]; + const sign = desc ? -1 : 1; + list.sort((a, b) => { + const pa = primary(a), pb = primary(b); + const na = pa === null || pa === undefined, nb = pb === null || pb === undefined; + if (na !== nb) return na ? 1 : -1; + if (!na && pa !== pb) return (pa < pb ? -1 : 1) * sign; + return compare(ties(a), ties(b)); + }); } async function detail(id) { @@ -604,7 +716,13 @@

cpplink clusters

if (e.entities > 1) pills += `${e.entities} entities`; if (e.discord) pills += `${e.discord} split`; else pills += `all agree`; + // The sort key's value is shown so the order can be read without + // opening anything; size and the split count are already there, so + // only the weakest edge needs adding. + const key = $('#sort').value === 'weakest' && e.weakest !== null && e.weakest !== undefined + ? `${e.weakest.toFixed(1)} bits` : ''; el.innerHTML = `${esc(e.id)}${pills}` + + (key ? `${key}` : '') + `${e.size.toLocaleString()}`; el.onclick = () => select(e.id); list.appendChild(el); @@ -625,12 +743,14 @@

cpplink clusters

} function fillSelect(el, options) { - el.innerHTML = options.map(([v, label]) => ``).join(''); + el.innerHTML = options.map(([v, label]) => + ``).join(''); } async function rebuild(reset) { await loadPage(reset !== false); renderList(); + remember(); if (!entries.some(e => e.id === current)) { await select(entries.length ? entries[0].id : null); } @@ -652,12 +772,15 @@

cpplink clusters

if (v === null || v === '') return '—'; // Two rows have no majority, so both sides of a difference are marked. const odd = s.distinct > 1 && (opts && opts.both ? true : v !== s.top); - return `${esc(v)}`; + const hit = search.field === 'col' && search.col === i && cellHit(v, search.q); + return `${esc(v)}`; }).join(''); const colour = truthColors[(r.t || 0) % truthColors.length]; const t = truthy ? `${r.t + 1}` : ''; - return `${esc(r.id)}${t}${cells}`; + const hit = search.field === 'id' && search.ids.includes(r.id); + return `${esc(r.id)}` + + `${t}${cells}`; }).join(''); return `
${head}` + `${body}
`; @@ -665,12 +788,49 @@

cpplink clusters

// The address names what is open, so a cluster or a pair can be linked to: // #cluster=, or #pair=
, for a pair open on its cluster's pairs tab. +// The search and the order go in the query string, defaults left out, so a +// search is a link too. function remember() { const hash = !current ? '' : tab === 'pairs' && pick ? `#pair=${encodeURIComponent(pick.a)},${encodeURIComponent(pick.b)}` : `#cluster=${encodeURIComponent(current)}`; - history.replaceState(null, '', location.pathname + location.search + hash); + const params = new URLSearchParams(); + const field = $('#field').value, q = $('#q').value.trim(); + if (q) { + if (field !== 'any') params.set('field', field === 'id' || field === 'cluster' + ? field : COLS[+field.slice(4)]); + params.set('q', q); + } + const sort = $('#sort').value; + const [, , defaultDesc] = SORTS.find(([k]) => k === sort); + if (sort !== SORTS[0][0]) params.set('sort', sort); + if (sortDesc !== defaultDesc) params.set('dir', sortDesc ? 'desc' : 'asc'); + if ($('#only').value !== 'all') params.set('only', $('#only').value); + const query = params.toString(); + try { + history.replaceState(null, '', location.pathname + (query ? '?' + query : '') + hash); + } catch (e) {} // a file:// page may refuse a new query string; the hash still works +} + +// The controls from the query string, before the first list is drawn. +function recall() { + const params = new URLSearchParams(location.search); + const field = params.get('field'); + if (field === 'id' || field === 'cluster') $('#field').value = field; + else if (field !== null) { + const i = COLS.indexOf(field); + if (i >= 0) $('#field').value = `col:${i}`; + } + $('#q').value = params.get('q') || ''; + const sort = SORTS.find(([k]) => k === params.get('sort')); + if (sort) $('#sort').value = sort[0]; + sortDesc = SORTS.find(([k]) => k === $('#sort').value)[2]; + if (params.get('dir') === 'asc') sortDesc = false; + if (params.get('dir') === 'desc') sortDesc = true; + if (ONLYS.some(([k]) => k === params.get('only'))) $('#only').value = params.get('only'); + placeholder(); + direction(); } // Open one cluster in the detail pane: its heading, the two tabs, and both @@ -1109,11 +1269,24 @@

${esc(c.id)}

let typing = null; $('#q').addEventListener('input', () => { + shortcut(); clearTimeout(typing); typing = setTimeout(() => rebuild(true), 200); }); -$('#sort').addEventListener('change', () => rebuild(true)); +$('#field').addEventListener('change', () => { placeholder(); rebuild(true); }); +$('#sort').addEventListener('change', () => { + sortDesc = SORTS.find(([k]) => k === $('#sort').value)[2]; + direction(); + rebuild(true); +}); +$('#dir').addEventListener('click', () => { sortDesc = !sortDesc; direction(); rebuild(true); }); $('#only').addEventListener('change', () => rebuild(true)); + +function direction() { + const b = $('#dir'); + b.textContent = sortDesc ? '\u2193' : '\u2191'; + b.title = sortDesc ? 'descending; click for ascending' : 'ascending; click for descending'; +} $('#list').addEventListener('scroll', async e => { const el = e.target; if (exhausted || el.scrollTop + el.clientHeight < el.scrollHeight - 200) return; @@ -1191,8 +1364,10 @@

${esc(c.id)}

TOTALS.pairs = predictions; } try { tab = localStorage.getItem('cpplink.tab') === 'pairs' ? 'pairs' : 'cluster'; } catch (e) {} + fillSelect($('#field'), FIELDS.concat(COLS.map((name, i) => [`col:${i}`, name]))); fillSelect($('#sort'), SORTS); fillSelect($('#only'), ONLYS); + recall(); const hash = decodeURIComponent(location.hash.slice(1)); let open = null; if (hash.startsWith('pair=')) { From 62eccb3dac27a65909b62a85cd36b6075b3611e7 Mon Sep 17 00:00:00 2001 From: 4ment Date: Thu, 17 Sep 2026 07:19:52 +1000 Subject: [PATCH 07/10] Enhance search functionality in cluster viewer to highlight matched rows --- tools/cluster_view_template.html | 43 +++++++++++++++++++++++++------- 1 file changed, 34 insertions(+), 9 deletions(-) diff --git a/tools/cluster_view_template.html b/tools/cluster_view_template.html index ffd6cf5..e34c221 100644 --- a/tools/cluster_view_template.html +++ b/tools/cluster_view_template.html @@ -224,8 +224,10 @@ .net g.node:hover circle { stroke: var(--ink); stroke-width: 2; } tr.lit td, tr.pin td { background: var(--accent-soft); } tr.lit td.odd, tr.pin td.odd { background: var(--warn-soft); } -/* what the search matched: a record the id search named, or a cell a column search hit */ -tr.hit td.id, td.val.hit { box-shadow: inset 0 -2px 0 var(--accent); } +/* what the search matched: the row holding it, tinted whole so it reads as one + row (a disagreeing cell keeps its orange text), and the cell or id the query is in */ +tr.hit td, tr.hit td.odd { background: var(--accent-soft); } +td.id.hit, td.val.hit { box-shadow: inset 0 -2px 0 var(--accent); } tr.pin td { box-shadow: inset 0 1px 0 var(--accent), inset 0 -1px 0 var(--accent); } .toggle { font-size: 12px; color: var(--muted); display: flex; gap: 5px; @@ -604,6 +606,22 @@

cpplink clusters

const cellHit = (v, q) => v !== null && v !== undefined && String(v).toLowerCase().includes(q); +// What the search matched in one row: its id, and the columns the query is a +// substring of. A cluster-id search names no row. +function rowHits(r) { + const hits = { id: false, cols: new Set() }; + if (!search.q) return hits; + switch (search.field) { + case 'id': hits.id = search.ids.includes(r.id); break; + case 'cluster': break; + case 'col': if (cellHit(r.v[search.col], search.q)) hits.cols.add(search.col); break; + default: + hits.id = cellHit(r.id, search.q); + r.v.forEach((v, i) => { if (cellHit(v, search.q)) hits.cols.add(i); }); + } + return hits; +} + function compare(ka, kb) { for (let i = 0; i < ka.length; i++) { if (ka[i] < kb[i]) return -1; @@ -751,9 +769,11 @@

cpplink clusters

await loadPage(reset !== false); renderList(); remember(); - if (!entries.some(e => e.id === current)) { - await select(entries.length ? entries[0].id : null); - } + // The open cluster is redrawn either way: a search that still matches it + // may match different rows, and the table marks the rows it matched. + const keep = entries.some(e => e.id === current); + await select(keep ? current : entries.length ? entries[0].id : null, + keep ? { tab, pair: pick } : undefined); } // The rows of a cluster, or of a pair, as one table: a cell differing from the @@ -767,20 +787,22 @@

cpplink clusters

return `${esc(name)}${d}`; })).join(''); const body = rows.map((r, i) => { + const hits = rowHits(r); const cells = COLS.map((_, i) => { const s = stats[i], v = r.v[i]; if (v === null || v === '') return '—'; // Two rows have no majority, so both sides of a difference are marked. const odd = s.distinct > 1 && (opts && opts.both ? true : v !== s.top); - const hit = search.field === 'col' && search.col === i && cellHit(v, search.q); + const hit = hits.cols.has(i); return `${esc(v)}`; }).join(''); const colour = truthColors[(r.t || 0) % truthColors.length]; const t = truthy ? `${r.t + 1}` : ''; - const hit = search.field === 'id' && search.ids.includes(r.id); - return `${esc(r.id)}` + - `${t}${cells}`; + // The row the search matched is lit, so it is the one the eye lands on. + const hit = hits.id || hits.cols.size > 0; + return `` + + `${esc(r.id)}${t}${cells}`; }).join(''); return `
${head}` + `${body}
`; @@ -971,6 +993,9 @@

${esc(c.id)}

renderPairTable(c); showTab(tab); d.scrollTop = 0; + // A cluster opened by a search scrolls to the first row it matched. + const hit = tab === 'cluster' && table.querySelector('tr.hit'); + if (hit) hit.scrollIntoView({ block: 'nearest' }); } function showTab(name) { From caea42e12f74cc87c7b884c84e1bba2140d0c941 Mon Sep 17 00:00:00 2001 From: 4ment Date: Fri, 18 Sep 2026 19:42:40 +1000 Subject: [PATCH 08/10] Add threshold filtering for predictions in cluster processing --- tools/cluster_server.py | 189 ++++++++++++++++++++++++++++++++++++---- 1 file changed, 173 insertions(+), 16 deletions(-) diff --git a/tools/cluster_server.py b/tools/cluster_server.py index 18806af..9a91102 100755 --- a/tools/cluster_server.py +++ b/tools/cluster_server.py @@ -27,6 +27,13 @@ is computed or spawned here. The level labels and rates come from `--model`. Predictions clustering kept apart -- above the write threshold, below the clustering one -- are kept and listed under both of their clusters as `rejected`. + +`--threshold` keeps only the predictions at or above it. With `--clusters` the +file is taken as what `cpplink cluster --threshold` wrote at that threshold and +read as it stands, which is the fast path; without one the predictions are +clustered here, by the same union-find over the same predictions in the same +order, so the cache holds exactly the partition that command writes, named by +the same representatives. """ import argparse @@ -42,6 +49,8 @@ import duckdb import numpy as np import pyarrow as pa +import pyarrow.compute as pc +import pyarrow.csv as pacsv import pyarrow.parquet as pq sys.path.insert(0, os.path.dirname(os.path.abspath(__file__))) @@ -49,7 +58,7 @@ read_model_levels, read_truth, schema_columns) EDGE_CHUNK = 1 << 20 -CACHE_VERSION = 6 +CACHE_VERSION = 7 def quoted(name): @@ -118,21 +127,155 @@ def load_shard_edges(conn, directory): return total +def read_shard_edges(directory, floor): + """Every shard edge at or above `floor` as one (a, b, w) table of rows, in + the order `cpplink cluster` reads them: shard by shard, sorted by name.""" + blocks = [] + for name in sorted(n for n in os.listdir(directory) if n.endswith(".bin")): + with open(os.path.join(directory, name), "rb") as handle: + if handle.read(len(EDGE_MAGIC)) != EDGE_MAGIC: + raise SystemExit(f"{name}: not a cpplink edge shard") + block = np.frombuffer(handle.read(), dtype=EDGE_DTYPE) + blocks.append(block[block["w"] >= floor]) + block = np.concatenate(blocks) if blocks else np.empty(0, dtype=EDGE_DTYPE) + return pa.table({"a": block["a"], "b": block["b"], "w": block["w"]}) + + +def read_file_edges(path, floor): + """Every prediction of a merged file at or above `floor` as one (a, b, w) + table of ids, in file order, which DuckDB would not promise once the build + connection stops preserving it.""" + if path.endswith(".parquet"): + table = pq.read_table(path, columns=["id_a", "id_b", "match_weight"]) + else: + table = pacsv.read_csv(path, convert_options=pacsv.ConvertOptions( + include_columns=["id_a", "id_b", "match_weight"], + column_types={"id_a": pa.string(), "id_b": pa.string(), + "match_weight": pa.float64()})) + table = table.filter(pc.greater_equal(table["match_weight"], floor)) + return pa.table({"a": table["id_a"].cast(pa.string()), + "b": table["id_b"].cast(pa.string()), "w": table["match_weight"]}) + + +def id_rows(conn, files, id_column, keys, by): + """`(uid, rid)` of the records `keys` names, by uid or by rid. + + One pass over the id column, which is the index `cpplink cluster` builds to + read a merged file back, and what a shard needs to name its rows. + """ + conn.register("keys", pa.table({"key": keys})) + parts, offset = [], 0 + for path in files: + parts.append(f""" + SELECT CAST({quoted(id_column)} AS VARCHAR) AS uid, + {offset} + file_row_number AS rid + FROM read_parquet('{path}', file_row_number = true)""") + offset += pq.ParquetFile(path).metadata.num_rows + found = conn.execute(f""" + SELECT uid, rid FROM ({" UNION ALL ".join(parts)}) + WHERE {by} IN (SELECT key FROM keys) + """).to_arrow_table() + conn.unregister("keys") + return found + + +def union_find(edges, count): + """The binary's union-find, edge for edge: union by rank, a tie keeping the + first end's root, so each cluster comes out named by the representative + `cpplink cluster` names it by. Returns the root of every index.""" + parent = list(range(count)) + rank = [0] * count + + def find(x): + while parent[x] != x: + parent[x] = parent[parent[x]] + x = parent[x] + return x + + for a, b in edges: + ra, rb = find(a), find(b) + if ra == rb: + continue + if rank[ra] < rank[rb]: + ra, rb = rb, ra + parent[rb] = ra + if rank[ra] == rank[rb]: + rank[ra] += 1 + return [find(x) for x in range(count)] + + +def recluster(conn, args, files, id_column, say): + """`members` and `named_edges` from the predictions at or above the + threshold, as `cpplink cluster --threshold` would write them. + + A prediction naming a record no input holds is skipped as the binary skips + it, which is why the id pass comes before the union-find rather than after. + """ + say("reading the predictions at the threshold") + if os.path.isdir(args.predictions): + edges = read_shard_edges(args.predictions, args.threshold) + keys = pc.unique(pa.concat_arrays([edges["a"].combine_chunks(), + edges["b"].combine_chunks()])) + names = id_rows(conn, files, id_column, keys, "rid") + by_rid = dict(zip(names["rid"].to_pylist(), names["uid"].to_pylist())) + a = [by_rid.get(r) for r in edges["a"].to_pylist()] + b = [by_rid.get(r) for r in edges["b"].to_pylist()] + else: + edges = read_file_edges(args.predictions, args.threshold) + keys = pc.unique(pa.concat_arrays([edges["a"].combine_chunks(), + edges["b"].combine_chunks()])) + names = id_rows(conn, files, id_column, keys, "uid") + known = set(names["uid"].to_pylist()) + a = [u if u in known else None for u in edges["a"].to_pylist()] + b = [u if u in known else None for u in edges["b"].to_pylist()] + weights = edges["w"].to_pylist() + kept = [(x, y, w) for x, y, w in zip(a, b, weights) if x is not None and y is not None] + + say("clustering them") + index = {} + for x, y, _ in kept: + index.setdefault(x, len(index)) + index.setdefault(y, len(index)) + root = union_find(((index[x], index[y]) for x, y, _ in kept), len(index)) + uids = list(index) + size = {} + for r in root: + size[r] = size.get(r, 0) + 1 + rows = [(uids[i], uids[root[i]], size[root[i]]) for i in range(len(uids)) + if size[root[i]] >= args.min_size + and (not args.max_size or size[root[i]] <= args.max_size)] + conn.register("assignment", pa.table({ + "uid": pa.array([r[0] for r in rows], pa.string()), + "cluster_id": pa.array([r[1] for r in rows], pa.string()), + "size": pa.array([r[2] for r in rows], pa.int64())})) + conn.execute("CREATE TABLE members AS SELECT uid, cluster_id, size FROM assignment") + conn.unregister("assignment") + conn.register("kept", pa.table({ + "a": pa.array([k[0] for k in kept], pa.string()), + "b": pa.array([k[1] for k in kept], pa.string()), + "weight": pa.array([k[2] for k in kept], pa.float64())})) + conn.execute("CREATE TABLE named_edges AS SELECT a, b, weight FROM kept") + conn.unregister("kept") + + def build(conn, args, id_column, columns, say): """Fill the cache: members, their records, their edges, per-cluster stats.""" files = [os.path.abspath(p) for p in args.data] picked = ", ".join(quoted(c) for c in columns) - say("reading the cluster assignment") - ceiling = f" AND cluster_size <= {args.max_size}" if args.max_size else "" - conn.execute(f""" - CREATE TABLE members AS - SELECT CAST(unique_id AS VARCHAR) AS uid, - CAST(cluster_id AS VARCHAR) AS cluster_id, - CAST(cluster_size AS BIGINT) AS size - FROM {scan(args.clusters, ("unique_id", "cluster_id"))} - WHERE cluster_size >= {args.min_size}{ceiling} - """) + if args.threshold is not None and not args.clusters: + recluster(conn, args, files, id_column, say) + else: + say("reading the cluster assignment") + ceiling = f" AND cluster_size <= {args.max_size}" if args.max_size else "" + conn.execute(f""" + CREATE TABLE members AS + SELECT CAST(unique_id AS VARCHAR) AS uid, + CAST(cluster_id AS VARCHAR) AS cluster_id, + CAST(cluster_size AS BIGINT) AS size + FROM {scan(args.clusters, ("unique_id", "cluster_id"))} + WHERE cluster_size >= {args.min_size}{ceiling} + """) say("reading the clustered records") # A record keeps its store row (`rid`): the position in the inputs read in @@ -151,7 +294,9 @@ def build(conn, args, id_column, columns, say): conn.execute("CREATE TABLE records AS " + " UNION ALL ".join(parts)) if args.predictions: - if os.path.isdir(args.predictions): + if args.threshold is not None and not args.clusters: + pass # `named_edges` is what `recluster` clustered + elif os.path.isdir(args.predictions): say("naming the predictions") # A shard names rows, and the clustered records carry theirs. load_shard_edges(conn, args.predictions) @@ -178,7 +323,9 @@ def build(conn, args, id_column, columns, say): # cluster ids, since it is the prediction clustering overruled and the # one most worth explaining; a pair with an end no cluster holds has no # record to show and is dropped. `--threshold` drops the predictions - # below it here, so the cache never holds them. + # below it here, so the cache never holds them; when the clusters were + # computed here they are the components of these very predictions, so + # every pair lands inside one cluster and none is rejected. floor = (f" WHERE e.weight >= {args.threshold}" if args.threshold is not None else "") conn.execute(f""" @@ -525,7 +672,7 @@ def main(): formatter_class=argparse.RawDescriptionHelpFormatter) ap.add_argument("data", nargs="+", help="the parquet file(s) the run read, in order") ap.add_argument("--schema", required=True) - ap.add_argument("--clusters", required=True, + ap.add_argument("--clusters", help="the csv or parquet cpplink cluster wrote") ap.add_argument("--predictions", "--edges", dest="predictions", help="the run's predictions -- one csv or parquet file, or the " @@ -541,8 +688,10 @@ def main(): ap.add_argument("--min-size", type=int, default=2) ap.add_argument("--max-size", type=int, default=0, help="0 = no ceiling") ap.add_argument("--threshold", type=float, - help="keep only the predictions whose match_weight is at " - "least this; the default keeps every one the run wrote") + help="keep only the predictions at or above this match_weight; " + "with --clusters that file is read as the clustering at " + "this threshold, without one the predictions are clustered " + "here exactly as cpplink cluster --threshold would") ap.add_argument("--max-rows", type=int, default=200, help="members shown per cluster; the rest are counted only") ap.add_argument("--memory", default="4GB", @@ -555,6 +704,14 @@ def main(): ap.error("--waterfalls and --model go together") if args.waterfalls and not args.predictions: ap.error("--waterfalls needs --predictions") + if args.threshold is not None: + if not args.predictions: + ap.error("--threshold prunes the predictions, so it needs --predictions") + if not args.clusters and args.min_size < 2: + ap.error("--threshold lists only what a prediction reaches, so a " + "singleton is never shown; --min-size must be at least 2") + elif not args.clusters: + ap.error("--clusters is needed without --threshold") id_column, columns = schema_columns(args.schema) stamp = fingerprint(args, columns) From 3bb4b13ec6a25e98bd05b61fdc14560888055fb3 Mon Sep 17 00:00:00 2001 From: 4ment Date: Fri, 18 Sep 2026 20:10:50 +1000 Subject: [PATCH 09/10] Refactor ID resolution logic in pair handling and enhance dataset support in prediction processing --- src/cpplink/app.cpp | 71 +++++++++--------------------------- src/cpplink/explain.cpp | 1 + src/cpplink/explain.hpp | 5 --- src/cpplink/id_index.hpp | 34 ------------------ src/cpplink/waterfall.cpp | 75 ++++++++++++++++++++++++++++++++------- 5 files changed, 81 insertions(+), 105 deletions(-) diff --git a/src/cpplink/app.cpp b/src/cpplink/app.cpp index 46791bb..bba0f7b 100644 --- a/src/cpplink/app.cpp +++ b/src/cpplink/app.cpp @@ -722,15 +722,23 @@ bool ResolvePair(const RecordStore& store, const std::string& text, bool by_row, } return true; } - if (!FindRowById(store, first, row_a)) { - *error = "no record with id '" + first + "'"; - return false; - } - if (!FindRowById(store, second, row_b)) { - *error = "no record with id '" + second + "'"; - return false; - } - return true; + const auto resolve = [&](const std::string& text, uint64_t* row) { + const IdLookup found = FindRowById(store, text, row); + if (found == IdLookup::kMissing) { + *error = "no record with id '" + text + "'"; + return false; + } + if (found == IdLookup::kAmbiguous) { + *error = "more than one record has id '" + text + + "'; name its input as :, where the datasets are"; + for (size_t d = 0; d < store.NumDatasets(); ++d) { + *error += (d == 0 ? " " : ", ") + store.DatasetName(d); + } + return false; + } + return true; + }; + return resolve(first, row_a) && resolve(second, row_b); } int RunExplain(const std::vector& args, std::ostream& out, @@ -842,51 +850,6 @@ int RunExplain(const std::vector& args, std::ostream& out, return 1; } - uint64_t row_a = 0; - uint64_t row_b = 0; - std::string first; - std::string second; - if (!pair.empty()) { - if (!SplitPair(pair, &first, &second)) { - err << "cpplink explain: --pair wants ,\n"; - return 1; - } - auto resolve = [&](const std::string& text, uint64_t* row) { - const IdLookup found = FindRowById(store, text, row); - if (found == IdLookup::kMissing) { - err << "cpplink explain: no record with id '" << text << "'\n"; - return false; - } - if (found == IdLookup::kAmbiguous) { - err << "cpplink explain: more than one record has id '" << text - << "'; name its input as :, where the datasets are"; - for (size_t d = 0; d < store.NumDatasets(); ++d) { - err << (d == 0 ? " " : ", ") << store.DatasetName(d); - } - err << "\n"; - return false; - } - return true; - }; - if (!resolve(first, &row_a) || !resolve(second, &row_b)) return 1; - } else { - if (!SplitPair(rows, &first, &second)) { - err << "cpplink explain: --rows wants ,\n"; - return 1; - } - row_a = std::stoull(first); - row_b = std::stoull(second); - if (row_a >= store.NumRecords() || row_b >= store.NumRecords()) { - err << "cpplink explain: row out of range; the file has " - << store.NumRecords() << " records\n"; - return 1; - } - } - - PrintGammaLayout(comparisons, out); - out << "\n"; - PrintPairExplanation(store, comparisons, row_a, row_b, out); - // Without a model there is no weight to explain: the levels are the whole // story, and the waterfall is simply not printed. Model model; diff --git a/src/cpplink/explain.cpp b/src/cpplink/explain.cpp index ab5efe8..4dd1401 100644 --- a/src/cpplink/explain.cpp +++ b/src/cpplink/explain.cpp @@ -13,6 +13,7 @@ #include #include + #include "cpplink/id_index.hpp" namespace cpplink { diff --git a/src/cpplink/explain.hpp b/src/cpplink/explain.hpp index 74b51c7..dd96e59 100644 --- a/src/cpplink/explain.hpp +++ b/src/cpplink/explain.hpp @@ -88,9 +88,4 @@ void PrintPairWaterfall(const RecordStore& store, const ComparisonSet& compariso // The same ledger as one JSON object on one line, for a tool that draws it. std::string PairWaterfallJson(const PairWaterfall& waterfall); -// Resolves a unique_id to a row by linear scan. There is no id index: ids are -// almost all distinct, so an index would cost as much as the values and is only -// ever needed for one-off lookups like this one. -bool FindRowById(const RecordStore& store, const std::string& id, uint64_t* row); - } // namespace cpplink diff --git a/src/cpplink/id_index.hpp b/src/cpplink/id_index.hpp index ee0ca48..55aabee 100644 --- a/src/cpplink/id_index.hpp +++ b/src/cpplink/id_index.hpp @@ -3,8 +3,6 @@ #pragma once -#include -#include #include #include #include @@ -14,38 +12,6 @@ namespace cpplink { -// An id-to-row lookup over the store's id column: one `uint32` per record, kept -// in the order of the id it names. A merged prediction file carries `unique_id`s -// rather than row indices, so reading one back has to map them, and a hash map -// over 20M ids costs an order of magnitude more than the structures it feeds. -// This is 4 bytes a record and a handful of string compares a lookup. -class IdIndex { - public: - explicit IdIndex(const RecordStore& store) : ids_(store.ids()) { - order_.resize(static_cast(store.NumRecords())); - for (size_t row = 0; row < order_.size(); ++row) { - order_[row] = static_cast(row); - } - std::sort(order_.begin(), order_.end(), - [this](uint32_t a, uint32_t b) { return ids_.Get(a) < ids_.Get(b); }); - } - - bool Find(std::string_view id, uint32_t* row) const { - const auto at = - std::lower_bound(order_.begin(), order_.end(), id, - [this](uint32_t candidate, std::string_view key) { - return ids_.Get(candidate) < key; - }); - if (at == order_.end() || ids_.Get(*at) != id) return false; - *row = *at; - return true; - } - - private: - const IdColumn& ids_; - std::vector order_; -}; - // What looking an id up can come back with. An id is unique within its input and // not necessarily across inputs, so an unqualified id can name one record, none, // or one per input that holds it -- and the last is not a hit on any of them. diff --git a/src/cpplink/waterfall.cpp b/src/cpplink/waterfall.cpp index 2006b59..dc7bde7 100644 --- a/src/cpplink/waterfall.cpp +++ b/src/cpplink/waterfall.cpp @@ -27,10 +27,13 @@ namespace cpplink { namespace { -// One prediction as the merged file names it. `gamma` is what the run stored, -// kept so a schema that has drifted from the file can be noticed. +// One prediction as the merged file names it: a record is its id, qualified by +// its dataset where the file carries one. `gamma` is what the run stored, kept +// so a schema that has drifted from the file can be noticed. struct NamedPrediction { + std::string_view dataset_a; std::string_view id_a; + std::string_view dataset_b; std::string_view id_b; uint32_t gamma = 0; }; @@ -46,19 +49,24 @@ bool ForEachCsvPrediction(const std::string& path, const PredictionSink& sink, } std::string line; uint64_t number = 0; + EdgeCsvLayout layout; while (std::getline(file, line)) { ++number; if (!line.empty() && line.back() == '\r') line.pop_back(); if (line.empty()) continue; - if (number == 1 && line.rfind("id_a,", 0) == 0) continue; - NamedPrediction prediction; - double weight = 0.0; - if (!ParseEdgeCsvLine(line, &prediction.id_a, &prediction.id_b, &prediction.gamma, - &weight)) { + if (number == 1 && ParseEdgeCsvHeader(line, &layout)) continue; + EdgeRow edge; + if (!ParseEdgeCsvLine(line, layout, &edge)) { *error = "explain: '" + path + "' line " + std::to_string(number) + " is not a cpplink prediction row"; return false; } + NamedPrediction prediction; + prediction.dataset_a = edge.dataset_a; + prediction.id_a = edge.id_a; + prediction.dataset_b = edge.dataset_b; + prediction.id_b = edge.id_b; + prediction.gamma = edge.gamma; if (!sink(prediction, error)) return false; } return true; @@ -101,6 +109,17 @@ bool ForEachParquetPrediction(const std::string& path, const PredictionSink& sin } indices.push_back(at); } + const bool datasets = schema->GetFieldIndex("dataset_a") >= 0; + if (datasets) { + for (const char* name : {"dataset_a", "dataset_b"}) { + const int at = schema->GetFieldIndex(name); + if (at < 0) { + *error = "explain: '" + path + "' has dataset_a but no " + name; + return false; + } + indices.push_back(at); + } + } for (int group = 0; group < reader->num_row_groups(); ++group) { auto group_result = reader->ReadRowGroup(group, indices); if (!group_result.ok()) { @@ -112,6 +131,8 @@ bool ForEachParquetPrediction(const std::string& path, const PredictionSink& sin const auto ids_a = table->GetColumnByName("id_a"); const auto ids_b = table->GetColumnByName("id_b"); const auto gammas = table->GetColumnByName("gamma"); + const auto sets_a = datasets ? table->GetColumnByName("dataset_a") : nullptr; + const auto sets_b = datasets ? table->GetColumnByName("dataset_b") : nullptr; for (int chunk = 0; chunk < ids_a->num_chunks(); ++chunk) { const auto a_array = std::dynamic_pointer_cast(ids_a->chunk(chunk)); @@ -119,15 +140,31 @@ bool ForEachParquetPrediction(const std::string& path, const PredictionSink& sin std::dynamic_pointer_cast(ids_b->chunk(chunk)); const auto gamma_array = std::dynamic_pointer_cast(gammas->chunk(chunk)); - if (a_array == nullptr || b_array == nullptr || gamma_array == nullptr) { + std::shared_ptr set_a; + std::shared_ptr set_b; + if (datasets) { + set_a = + std::dynamic_pointer_cast(sets_a->chunk(chunk)); + set_b = + std::dynamic_pointer_cast(sets_b->chunk(chunk)); + } + if (a_array == nullptr || b_array == nullptr || gamma_array == nullptr || + (datasets && (set_a == nullptr || set_b == nullptr))) { *error = "explain: '" + path + - "' holds id_a, id_b or gamma in a type a cpplink prediction " - "file does not use"; + "' holds id_a, id_b, dataset_a, dataset_b or gamma in a type " + "a cpplink prediction file does not use"; return false; } for (int64_t row = 0; row < a_array->length(); ++row) { - if (a_array->IsNull(row) || b_array->IsNull(row)) continue; + if (a_array->IsNull(row) || b_array->IsNull(row) || + (datasets && (set_a->IsNull(row) || set_b->IsNull(row)))) { + continue; + } NamedPrediction prediction; + if (datasets) { + prediction.dataset_a = set_a->GetView(row); + prediction.dataset_b = set_b->GetView(row); + } prediction.id_a = a_array->GetView(row); prediction.id_b = b_array->GetView(row); prediction.gamma = gamma_array->IsNull(row) ? 0 : gamma_array->Value(row); @@ -431,12 +468,26 @@ bool WriteWaterfalls(const RecordStore& store, const ComparisonSet& comparisons, } const IdIndex index(store); + if (!index.Unique(error)) { + *error = "explain: " + *error; + return false; + } + // With a dataset the lookup is exact; without one it is answered only where + // a single record carries the id, as `cluster` reads the same file. + const auto resolve = [&](std::string_view dataset, std::string_view id, + uint32_t* row) { + if (dataset.empty()) return index.Find(id, row) == IdLookup::kFound; + size_t which = 0; + return store.DatasetIndex(dataset, &which) && + index.Find(which, id, row) == IdLookup::kFound; + }; const PredictionSink each = [&](const NamedPrediction& prediction, std::string* trouble) { ++report->read; uint32_t a = 0; uint32_t b = 0; - if (!index.Find(prediction.id_a, &a) || !index.Find(prediction.id_b, &b)) { + if (!resolve(prediction.dataset_a, prediction.id_a, &a) || + !resolve(prediction.dataset_b, prediction.id_b, &b)) { ++report->unresolved; return true; } From de7d5712206b1f9a80adca67609fe16d19aed581 Mon Sep 17 00:00:00 2001 From: 4ment Date: Mon, 21 Sep 2026 20:53:35 +1000 Subject: [PATCH 10/10] Remove cluster_view.py script as it is no longer needed for HTML viewer generation of cpplink clusters. --- .gitignore | 1 + README.md | 1 + docs/commands/cluster.md | 4 + docs/commands/explain.md | 4 +- docs/python.md | 6 + docs/viewer.md | 67 ++ environment.yml | 2 + mkdocs.yml | 1 + pyproject.toml | 9 +- python/cpplink_viewer/__init__.py | 13 + python/cpplink_viewer/__main__.py | 10 + python/cpplink_viewer/cache.py | 531 ++++++++++++ python/cpplink_viewer/inputs.py | 178 ++++ python/cpplink_viewer/server.py | 506 ++++++++++++ .../cpplink_viewer/viewer.html | 182 +---- python/tests/test_viewer.py | 318 ++++++++ tools/cluster_server.py | 771 ------------------ tools/cluster_view.py | 611 -------------- 18 files changed, 1668 insertions(+), 1547 deletions(-) create mode 100644 docs/viewer.md create mode 100644 python/cpplink_viewer/__init__.py create mode 100644 python/cpplink_viewer/__main__.py create mode 100644 python/cpplink_viewer/cache.py create mode 100644 python/cpplink_viewer/inputs.py create mode 100644 python/cpplink_viewer/server.py rename tools/cluster_view_template.html => python/cpplink_viewer/viewer.html (87%) create mode 100644 python/tests/test_viewer.py delete mode 100755 tools/cluster_server.py delete mode 100755 tools/cluster_view.py diff --git a/.gitignore b/.gitignore index 3df4381..b29e834 100644 --- a/.gitignore +++ b/.gitignore @@ -13,6 +13,7 @@ predictions/ predictions.csv *.shards/ clusters.csv +*.duckdb # Local working notes, not part of the repo CLAUDE.md diff --git a/README.md b/README.md index e612b9e..981b282 100644 --- a/README.md +++ b/README.md @@ -108,6 +108,7 @@ result = linker.cluster("predictions.parquet", truth="sample.truth.csv") print(result.quality) ``` +`cpplink-viewer` (`pip install "cpplink[viewer]"`) then serves the clusters as a page: each cluster's members side by side with every disagreeing cell highlighted, and every prediction behind it as the ledger the scorer produced. Between `init` and `estimate` sit the diagnostics that cost seconds and decide the quality of the result: `profile` says what each column can be worth before any model exists, `levels` places the fuzzy thresholds from the data, `explain-blocking` prices every source without enumerating a pair, and `recall` measures what blocking reaches. See [Getting started](https://4ment.github.io/cpplink/getting-started/) for the whole pipeline with its output explained, [Commands](https://4ment.github.io/cpplink/commands/) for every command at a glance, the [schema reference](https://4ment.github.io/cpplink/reference/schema/) for every field, and [From Python](https://4ment.github.io/cpplink/python/) for the package. diff --git a/docs/commands/cluster.md b/docs/commands/cluster.md index be6dd39..3649ed3 100644 --- a/docs/commands/cluster.md +++ b/docs/commands/cluster.md @@ -223,6 +223,10 @@ The `cluster_id` is the representative record's own `unique_id`, so the output s record the others collapse onto. Only records in a cluster of at least `--min-size` are written; singletons are excluded by default. +## Reading the clusters + +[The cluster viewer](../viewer.md) serves this file beside the predictions and their waterfalls: the members of each cluster side by side with every disagreeing cell highlighted, and every prediction touching it, including the ones clustering at a higher threshold overruled. + ## Cost The union–find is `uint32` parent plus `uint8` rank — **5 bytes a record**, 100 MB at 20M rows — diff --git a/docs/commands/explain.md b/docs/commands/explain.md index 0b117e8..cbe9d13 100644 --- a/docs/commands/explain.md +++ b/docs/commands/explain.md @@ -216,8 +216,8 @@ A level's label, `m` and `u` are the model's rather than the pair's, so the file Ids are resolved through a sorted index over the id column, 4 bytes a record, so a file of millions of predictions costs one load and one pass. A prediction naming an id no record holds is skipped and counted, and so is one whose stored `gamma` the comparisons no longer produce, which means the schema changed after `predict` ran and the ledgers explain today's schema rather than the file's weights. -This is what the cluster viewers in `tools/` draw from: `tools/cluster_view.py --waterfalls` embeds a ledger per prediction, and `tools/cluster_server.py --waterfalls` loads the file beside the predictions so a click is one lookup. -Nothing in either recomputes a bit of the weight, and neither runs this binary. +This is what [the cluster viewer](../viewer.md) draws from: `cpplink-viewer --waterfalls` loads the file beside the predictions so a click is one lookup. +Nothing in it recomputes a bit of the weight, and it does not run this binary. ## See also diff --git a/docs/python.md b/docs/python.md index 848d1da..7cf624f 100644 --- a/docs/python.md +++ b/docs/python.md @@ -168,6 +168,12 @@ None of them calls back into Python, so another thread can run while they do. Errors the core reports as `false` and a message become `cpplink.Error`, a `RuntimeError` carrying the core's own text. +## The cluster viewer + +`cpplink-viewer`, the `cpplink_viewer` package in the same wheel, serves the clusters a run produced as a page; `pip install "cpplink[viewer]"` adds DuckDB, its one dependency beyond the binding's. +It reads the files the pipeline wrote and never the compiled module, so it also runs from a checkout with nothing built. +See [The cluster viewer](viewer.md). + ## Not in this version - **In-memory input.** The `Linker` reads parquet files, as the command does. A pyarrow, polars or pandas table through `__arrow_c_stream__` is the next step, and needs the loader split so a table can be appended without a file. diff --git a/docs/viewer.md b/docs/viewer.md new file mode 100644 index 0000000..fe7256a --- /dev/null +++ b/docs/viewer.md @@ -0,0 +1,67 @@ +# The cluster viewer + +`cpplink-viewer` serves a page over the clusters a run produced: a list of clusters on the left, and for the one picked, its members side by side with every disagreeing cell highlighted, and every prediction touching it with the one picked drawn as the ledger the scorer produced. +It reads what the pipeline wrote and runs none of it: the clusters and the predictions come from the files, the waterfalls from `explain --predictions`, and nothing in the viewer recomputes a bit of a weight. + +```sh +cpplink predict --schema schema.json --model model.json --out predictions.parquet sample.parquet +cpplink cluster --schema schema.json --predictions predictions.parquet --out clusters.csv sample.parquet +cpplink explain --schema schema.json --model model.json --predictions predictions.parquet --out waterfalls.parquet sample.parquet +cpplink-viewer --schema schema.json --clusters clusters.csv --predictions predictions.parquet \ + --waterfalls waterfalls.parquet --model model.json --truth sample.truth.csv --open sample.parquet +``` + +The page is served at `http://127.0.0.1:8770/` (`--port` moves it, `--open` opens a browser on it). + +## Install + +The viewer is the `cpplink_viewer` package, shipped in the same wheel as the binding, with DuckDB as its one extra dependency: + +```sh +pip install "cpplink[viewer]" +``` + +It never imports the compiled module, so it also runs from a checkout with nothing built, against a run the command line produced: + +```sh +pip install duckdb pyarrow numpy +PYTHONPATH=python python -m cpplink_viewer --schema schema.json --clusters clusters.csv sample.parquet +``` + +## What it reads + +| Option | File | Written by | +|---|---|---| +| `data` | the parquet file(s) the run read, in order | you | +| `--schema` | the schema, for the id column and the columns to show | `init` | +| `--clusters` | `unique_id, cluster_id, cluster_size`, csv or parquet | `cluster --out` | +| `--predictions` | one csv or parquet file, or the shard directory | `predict --out` | +| `--waterfalls` | one wide row per prediction, csv or parquet | `explain --predictions --out` | +| `--model` | the level labels and rates the waterfall shows | `estimate --out` | +| `--truth` | known pairs, to colour members by true entity | you, or `gen-sample --truth` | + +Only `--schema` and the data are required, plus one of `--clusters` or `--threshold`. +Without `--predictions` the page shows the members and nothing about why they are together; without `--waterfalls` and `--model` it shows each prediction's weight but not the ledger behind it. +A shard directory names records by row, so it costs one pass over the id column that a merged file does not. + +## The cache + +The first run builds a DuckDB file (`--cache`, `cpplink_viewer.duckdb` by default) holding only what the page can ever show: the clustered records, never the singletons, their predictions, the waterfalls of those predictions, and one row of statistics per cluster. +At 20M records with 2.9M of them clustered the parquet is read exactly once and every later request touches a table seven times smaller than the file. +The cache is keyed on every input's path, size and modification time and on the options that shape it, so a changed input rebuilds it and an unchanged one is reused; `--rebuild` forces it. +`--memory` and `--threads` are DuckDB's limits while building. + +## Thresholds and rejected predictions + +`--threshold` keeps only the predictions at or above it. +With `--clusters`, the file is taken as what `cluster --threshold` wrote at that threshold and read as it stands. +Without one, the predictions are clustered here by the same union-find in the same order, so the cache holds exactly the partition that command would write, named by the same representatives; the test suite holds the two to the same file. + +A prediction whose two records clustering put in different clusters, above the write threshold and below the clustering one, is the prediction most worth reading, so it is kept and listed under both of its clusters as `rejected`, each naming the other cluster its second record went to. +`--min-size` and `--max-size` bound the clusters listed; `--max-rows` bounds the members shown per cluster, the rest being counted only. + +## The list + +The list is sorted by disagreement (columns whose values differ within the cluster), size, weakest prediction or id, and filtered to every cluster, the split ones, the ones held together only transitively (fewer predictions than pairs), or the ones mixing entities against the truth file. +The search box reads any value, one column, the cluster id, or a comma-separated list of record ids matched whole, in which case it says which of them named a record. +The `network` checkbox draws the selected cluster's predictions as a graph, which is where a chain shows itself as a chain. diff --git a/environment.yml b/environment.yml index 150addd..f1a660a 100644 --- a/environment.yml +++ b/environment.yml @@ -21,3 +21,5 @@ dependencies: - numpy - pytest - ruff + # The viewer's cache; the `viewer` extra of the wheel. + - python-duckdb diff --git a/mkdocs.yml b/mkdocs.yml index 7f453d8..07e9299 100644 --- a/mkdocs.yml +++ b/mkdocs.yml @@ -36,6 +36,7 @@ nav: - Home: index.md - Getting started: getting-started.md - From Python: python.md + - The cluster viewer: viewer.md - Concepts: - The model: model.md - Comparisons: comparisons.md diff --git a/pyproject.toml b/pyproject.toml index 67e2b05..aae8c5a 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -14,6 +14,11 @@ dependencies = ["numpy", "pyarrow"] [project.optional-dependencies] test = ["pytest"] +# `cpplink-viewer` serves the clusters of a run over a DuckDB cache. +viewer = ["duckdb"] + +[project.scripts] +cpplink-viewer = "cpplink_viewer.server:main" [project.urls] Homepage = "https://github.com/4ment/cpplink" @@ -26,7 +31,7 @@ metadata.version.input = "CMakeLists.txt" metadata.version.regex = 'project\(cpplink VERSION (?P[0-9.]+)' cmake.args = ["-DCPPLINK_BUILD_PYTHON=ON"] cmake.build-type = "Release" -wheel.packages = ["python/cpplink"] +wheel.packages = ["python/cpplink", "python/cpplink_viewer"] # Keep the extension's build out of build/, which is the CLI's. build-dir = "build/python" @@ -62,4 +67,4 @@ indent-style = "space" line-ending = "auto" [tool.ruff.lint.isort] -known-first-party = ["cpplink"] +known-first-party = ["cpplink", "cpplink_viewer"] diff --git a/python/cpplink_viewer/__init__.py b/python/cpplink_viewer/__init__.py new file mode 100644 index 0000000..1cfbb94 --- /dev/null +++ b/python/cpplink_viewer/__init__.py @@ -0,0 +1,13 @@ +# Copyright 2026 Mathieu Fourment +# SPDX-License-Identifier: MIT +"""``cpplink_viewer``: the cluster viewer, served over a DuckDB cache of a run. + +It reads what a run wrote and never the compiled module, so it runs from a +checkout with no build (``PYTHONPATH=python python -m cpplink_viewer``) as well +as from the wheel (``cpplink-viewer``). DuckDB is the ``viewer`` extra. +""" + +from .cache import Options, open_cache +from .server import Viewer, main, make_server + +__all__ = ["Options", "Viewer", "main", "make_server", "open_cache"] diff --git a/python/cpplink_viewer/__main__.py b/python/cpplink_viewer/__main__.py new file mode 100644 index 0000000..de4eba5 --- /dev/null +++ b/python/cpplink_viewer/__main__.py @@ -0,0 +1,10 @@ +# Copyright 2026 Mathieu Fourment +# SPDX-License-Identifier: MIT +"""``python -m cpplink_viewer [options] data.parquet``: the viewer, served.""" + +import sys + +from .server import main + +if __name__ == "__main__": + sys.exit(main()) diff --git a/python/cpplink_viewer/cache.py b/python/cpplink_viewer/cache.py new file mode 100644 index 0000000..c64862b --- /dev/null +++ b/python/cpplink_viewer/cache.py @@ -0,0 +1,531 @@ +# Copyright 2026 Mathieu Fourment +# SPDX-License-Identifier: MIT +"""The DuckDB cache the viewer serves from, built once per set of inputs. + +It holds only what the page can ever show: the clustered records, never the +singletons, so it is a fraction of the dataset. At 20M records with 2.9M of them +clustered the parquet is read exactly once and every later request touches a +table seven times smaller than the file. +""" + +from __future__ import annotations + +import json +import os +import time +from dataclasses import dataclass + +import numpy as np +import pyarrow as pa +import pyarrow.compute as pc +import pyarrow.csv as pacsv +import pyarrow.parquet as pq + +from .inputs import EDGE_DTYPE, EDGE_MAGIC, read_truth, schema_columns + +EDGE_CHUNK = 1 << 20 +CACHE_VERSION = 8 + + +@dataclass +class Options: + """What the cache is built from and how; the command line fills one.""" + + data: list[str] + schema: str + clusters: str | None = None + predictions: str | None = None + truth: str | None = None + waterfalls: str | None = None + model: str | None = None + cache: str = "cpplink_viewer.duckdb" + rebuild: bool = False + min_size: int = 2 + max_size: int = 0 + threshold: float | None = None + max_rows: int = 200 + memory: str = "4GB" + threads: int = 0 + + def check(self): + """The reason these options cannot build a cache, or None.""" + if bool(self.waterfalls) != bool(self.model): + return "--waterfalls and --model go together" + if self.waterfalls and not self.predictions: + return "--waterfalls needs --predictions" + if self.threshold is not None: + if not self.predictions: + return "--threshold prunes the predictions, so it needs --predictions" + if not self.clusters and self.min_size < 2: + return ( + "--threshold lists only what a prediction reaches, so a " + "singleton is never shown; --min-size must be at least 2" + ) + elif not self.clusters: + return "--clusters is needed without --threshold" + return None + + +def connect(path, read_only=False): + """A DuckDB connection, naming the extra to install when there is no DuckDB.""" + try: + import duckdb + except ImportError as missing: + raise SystemExit( + "the viewer needs duckdb: pip install duckdb, or pip install cpplink[viewer]" + ) from missing + return duckdb.connect(path, read_only=read_only) + + +def quoted(name): + return '"' + name.replace('"', '""') + '"' + + +def scan(path, types): + """The duckdb table function that reads one csv or parquet file. + + Both of a run's outputs are one file whose extension picks the format, so + that is what picks the reader here. The ids are pinned to VARCHAR rather than + sniffed, because a file of numeric ids would otherwise read as integers and + join against nothing. + """ + path = os.path.abspath(path) + if path.endswith(".parquet"): + return f"read_parquet('{path}')" + pinned = ", ".join(f"'{name}': 'VARCHAR'" for name in types) + return f"read_csv('{path}', header = true, types = {{{pinned}}})" + + +def fingerprint(opts, columns): + """What the cache was built from, so a changed input rebuilds it.""" + parts = [ + CACHE_VERSION, + columns, + opts.max_rows, + opts.min_size, + opts.max_size, + opts.threshold, + ] + for path in [opts.clusters, opts.truth, opts.waterfalls] + list(opts.data): + if path: + parts.append( + [os.path.abspath(path), os.path.getmtime(path), os.path.getsize(path)] + ) + if opts.predictions: + if os.path.isdir(opts.predictions): + for name in sorted(os.listdir(opts.predictions)): + if name.endswith(".bin"): + full = os.path.join(opts.predictions, name) + parts.append([full, os.path.getmtime(full), os.path.getsize(full)]) + else: + parts.append( + [ + os.path.abspath(opts.predictions), + os.path.getmtime(opts.predictions), + os.path.getsize(opts.predictions), + ] + ) + return json.dumps(parts, sort_keys=True) + + +def shard_blocks(directory): + """Each shard of a directory as one array of its records, sorted by name, + which is the order `cpplink cluster` reads them in.""" + for name in sorted(n for n in os.listdir(directory) if n.endswith(".bin")): + with open(os.path.join(directory, name), "rb") as handle: + if handle.read(len(EDGE_MAGIC)) != EDGE_MAGIC: + raise SystemExit(f"{name}: not a cpplink prediction shard") + while True: + data = handle.read(EDGE_CHUNK * EDGE_DTYPE.itemsize) + if not data: + break + yield np.frombuffer(data, dtype=EDGE_DTYPE) + + +def load_shard_edges(conn, directory): + """Every shard edge into `raw_edges`, a chunk at a time. + + The shards hold row indices rather than ids, which is why this is worth + doing once: the join that names them costs a pass over the record ids, and + a served request must not pay it. + """ + conn.execute("CREATE TABLE raw_edges (a UINTEGER, b UINTEGER, w DOUBLE)") + total = 0 + for block in shard_blocks(directory): + chunk = pa.table({"a": block["a"], "b": block["b"], "w": block["w"]}) + conn.register("chunk", chunk) + conn.execute("INSERT INTO raw_edges SELECT a, b, w FROM chunk") + conn.unregister("chunk") + total += block.size + return total + + +def read_shard_edges(directory, floor): + """Every shard edge at or above `floor` as one (a, b, w) table of rows, in + the order `cpplink cluster` reads them: shard by shard, sorted by name.""" + blocks = [block[block["w"] >= floor] for block in shard_blocks(directory)] + block = np.concatenate(blocks) if blocks else np.empty(0, dtype=EDGE_DTYPE) + return pa.table({"a": block["a"], "b": block["b"], "w": block["w"]}) + + +def read_file_edges(path, floor): + """Every prediction of a merged file at or above `floor` as one (a, b, w) + table of ids, in file order, which DuckDB would not promise once the build + connection stops preserving it.""" + if path.endswith(".parquet"): + table = pq.read_table(path, columns=["id_a", "id_b", "match_weight"]) + else: + table = pacsv.read_csv( + path, + convert_options=pacsv.ConvertOptions( + include_columns=["id_a", "id_b", "match_weight"], + column_types={ + "id_a": pa.string(), + "id_b": pa.string(), + "match_weight": pa.float64(), + }, + ), + ) + table = table.filter(pc.greater_equal(table["match_weight"], floor)) + return pa.table( + { + "a": table["id_a"].cast(pa.string()), + "b": table["id_b"].cast(pa.string()), + "w": table["match_weight"], + } + ) + + +def id_rows(conn, files, id_column, keys, by): + """`(uid, rid)` of the records `keys` names, by uid or by rid. + + One pass over the id column, which is the index `cpplink cluster` builds to + read a merged file back, and what a shard needs to name its rows. + """ + conn.register("keys", pa.table({"key": keys})) + parts, offset = [], 0 + for path in files: + parts.append(f""" + SELECT CAST({quoted(id_column)} AS VARCHAR) AS uid, + {offset} + file_row_number AS rid + FROM read_parquet('{path}', file_row_number = true)""") + offset += pq.ParquetFile(path).metadata.num_rows + found = conn.execute(f""" + SELECT uid, rid FROM ({" UNION ALL ".join(parts)}) + WHERE {by} IN (SELECT key FROM keys) + """).to_arrow_table() + conn.unregister("keys") + return found + + +def union_find(edges, count): + """The core's union-find, edge for edge: union by rank, a tie keeping the + first end's root, so each cluster comes out named by the representative + `cpplink cluster` names it by. Returns the root of every index. + + This is the one piece of the core the viewer reimplements, and a test holds + it to the partition `cpplink cluster --threshold` writes. + """ + parent = list(range(count)) + rank = [0] * count + + def find(x): + while parent[x] != x: + parent[x] = parent[parent[x]] + x = parent[x] + return x + + for a, b in edges: + ra, rb = find(a), find(b) + if ra == rb: + continue + if rank[ra] < rank[rb]: + ra, rb = rb, ra + parent[rb] = ra + if rank[ra] == rank[rb]: + rank[ra] += 1 + return [find(x) for x in range(count)] + + +def recluster(conn, opts, files, id_column, say): + """`members` and `named_edges` from the predictions at or above the + threshold, as `cpplink cluster --threshold` would write them. + + A prediction naming a record no input holds is skipped as the core skips + it, which is why the id pass comes before the union-find rather than after. + """ + say("reading the predictions at the threshold") + if os.path.isdir(opts.predictions): + edges = read_shard_edges(opts.predictions, opts.threshold) + keys = pc.unique( + pa.concat_arrays([edges["a"].combine_chunks(), edges["b"].combine_chunks()]) + ) + names = id_rows(conn, files, id_column, keys, "rid") + by_rid = dict( + zip(names["rid"].to_pylist(), names["uid"].to_pylist(), strict=True) + ) + a = [by_rid.get(r) for r in edges["a"].to_pylist()] + b = [by_rid.get(r) for r in edges["b"].to_pylist()] + else: + edges = read_file_edges(opts.predictions, opts.threshold) + keys = pc.unique( + pa.concat_arrays([edges["a"].combine_chunks(), edges["b"].combine_chunks()]) + ) + names = id_rows(conn, files, id_column, keys, "uid") + known = set(names["uid"].to_pylist()) + a = [u if u in known else None for u in edges["a"].to_pylist()] + b = [u if u in known else None for u in edges["b"].to_pylist()] + weights = edges["w"].to_pylist() + kept = [ + (x, y, w) + for x, y, w in zip(a, b, weights, strict=True) + if x is not None and y is not None + ] + + say("clustering them") + index = {} + for x, y, _ in kept: + index.setdefault(x, len(index)) + index.setdefault(y, len(index)) + root = union_find(((index[x], index[y]) for x, y, _ in kept), len(index)) + uids = list(index) + size = {} + for r in root: + size[r] = size.get(r, 0) + 1 + rows = [ + (uids[i], uids[root[i]], size[root[i]]) + for i in range(len(uids)) + if size[root[i]] >= opts.min_size + and (not opts.max_size or size[root[i]] <= opts.max_size) + ] + conn.register( + "assignment", + pa.table( + { + "uid": pa.array([r[0] for r in rows], pa.string()), + "cluster_id": pa.array([r[1] for r in rows], pa.string()), + "size": pa.array([r[2] for r in rows], pa.int64()), + } + ), + ) + conn.execute("CREATE TABLE members AS SELECT uid, cluster_id, size FROM assignment") + conn.unregister("assignment") + conn.register( + "kept", + pa.table( + { + "a": pa.array([k[0] for k in kept], pa.string()), + "b": pa.array([k[1] for k in kept], pa.string()), + "weight": pa.array([k[2] for k in kept], pa.float64()), + } + ), + ) + conn.execute("CREATE TABLE named_edges AS SELECT a, b, weight FROM kept") + conn.unregister("kept") + + +def build(conn, opts, id_column, columns, say): + """Fill the cache: members, their records, their edges, per-cluster stats.""" + files = [os.path.abspath(p) for p in opts.data] + picked = ", ".join(quoted(c) for c in columns) + + if opts.threshold is not None and not opts.clusters: + recluster(conn, opts, files, id_column, say) + else: + say("reading the cluster assignment") + ceiling = f" AND cluster_size <= {opts.max_size}" if opts.max_size else "" + conn.execute(f""" + CREATE TABLE members AS + SELECT CAST(unique_id AS VARCHAR) AS uid, + CAST(cluster_id AS VARCHAR) AS cluster_id, + CAST(cluster_size AS BIGINT) AS size + FROM {scan(opts.clusters, ("unique_id", "cluster_id"))} + WHERE cluster_size >= {opts.min_size}{ceiling} + """) + + say("reading the clustered records") + # A record keeps its store row (`rid`): the position in the inputs read in + # order, which is what a shard names. file_row_number is that position + # inside one file, and the offsets between the inputs are the only thing to + # carry. + parts, offset = [], 0 + for path in files: + parts.append(f""" + SELECT m.cluster_id, d.{quoted(id_column)} AS uid, {picked}, + lower(concat_ws(' ', d.{quoted(id_column)}, {picked})) AS text, + {offset} + file_row_number AS rid + FROM read_parquet('{path}', file_row_number = true) d + JOIN members m ON m.uid = d.{quoted(id_column)}""") + offset += pq.ParquetFile(path).metadata.num_rows + conn.execute("CREATE TABLE records AS " + " UNION ALL ".join(parts)) + + if opts.predictions: + if opts.threshold is not None and not opts.clusters: + pass # `named_edges` is what `recluster` clustered + elif os.path.isdir(opts.predictions): + say("naming the predictions") + # A shard names rows, and the clustered records carry theirs. + load_shard_edges(conn, opts.predictions) + conn.execute(""" + CREATE TABLE named_edges AS + SELECT ra.uid AS a, rb.uid AS b, e.w AS weight + FROM raw_edges e + JOIN records ra ON ra.rid = e.a + JOIN records rb ON rb.rid = e.b + """) + conn.execute("DROP TABLE raw_edges") + else: + # A merged file already names records by unique_id. + say("reading the predictions") + conn.execute(f""" + CREATE TABLE named_edges AS + SELECT CAST(id_a AS VARCHAR) AS a, CAST(id_b AS VARCHAR) AS b, + match_weight AS weight + FROM {scan(opts.predictions, ("id_a", "id_b"))} + """) + # Both ends clustered, because clustering at a threshold above the one + # the run wrote at leaves predictions that cross two clusters or land + # outside every one. A pair across two clusters is kept with both + # cluster ids, since it is the prediction clustering overruled and the + # one most worth explaining; a pair with an end no cluster holds has no + # record to show and is dropped. `--threshold` drops the predictions + # below it here, so the cache never holds them; when the clusters were + # computed here they are the components of these very predictions, so + # every pair lands inside one cluster and none is rejected. + floor = ( + f" WHERE e.weight >= {opts.threshold}" if opts.threshold is not None else "" + ) + conn.execute(f""" + CREATE TABLE pairs AS + SELECT e.a, e.b, e.weight, ra.cluster_id AS cluster_a, + rb.cluster_id AS cluster_b + FROM named_edges e + JOIN records ra ON ra.uid = e.a + JOIN records rb ON rb.uid = e.b{floor} + """) + conn.execute("DROP TABLE named_edges") + conn.execute(""" + CREATE TABLE edges AS + SELECT cluster_a AS cluster_id, a, b, weight FROM pairs + WHERE cluster_a = cluster_b + """) + conn.execute("CREATE INDEX edges_by_cluster ON edges (cluster_id)") + # A rejected pair is asked for from either of its clusters. + conn.execute("CREATE INDEX pairs_by_cluster_a ON pairs (cluster_a)") + conn.execute("CREATE INDEX pairs_by_cluster_b ON pairs (cluster_b)") + conn.execute("CREATE INDEX pairs_by_a ON pairs (a)") + conn.execute("CREATE INDEX pairs_by_b ON pairs (b)") + if opts.waterfalls: + # Only the rows the page can ever ask for: the pairs kept above. + say("reading the waterfalls") + conn.execute(f""" + CREATE TABLE waterfalls AS + SELECT w.* REPLACE (CAST(w.id_a AS VARCHAR) AS id_a, + CAST(w.id_b AS VARCHAR) AS id_b) + FROM {scan(opts.waterfalls, ("id_a", "id_b"))} w + JOIN pairs p ON p.a = CAST(w.id_a AS VARCHAR) + AND p.b = CAST(w.id_b AS VARCHAR) + """) + conn.execute("CREATE INDEX waterfalls_by_pair ON waterfalls (id_a, id_b)") + else: + conn.execute( + "CREATE TABLE edges (cluster_id VARCHAR, a VARCHAR, b VARCHAR, weight DOUBLE)" + ) + conn.execute( + "CREATE TABLE pairs (a VARCHAR, b VARCHAR, weight DOUBLE, " + "cluster_a VARCHAR, cluster_b VARCHAR)" + ) + + if opts.truth: + say("closing the known pairs") + uids = [row[0] for row in conn.execute("SELECT uid FROM members").fetchall()] + groups = read_truth(opts.truth, set(uids)) + table = pa.table( + { + "uid": pa.array(list(groups.keys()), pa.string()), + "grp": pa.array(list(groups.values()), pa.int64()), + } + ) + conn.register("groups", table) + conn.execute("CREATE TABLE truth AS SELECT uid, grp FROM groups") + conn.unregister("groups") + else: + conn.execute("CREATE TABLE truth (uid VARCHAR, grp BIGINT)") + + say("counting what each cluster agrees on") + # count(DISTINCT x) skips nulls, which is the page's rule: a column with one + # value and some gaps still agrees. + splits = " + ".join( + f"CASE WHEN count(DISTINCT {quoted(c)}) > 1 THEN 1 ELSE 0 END" for c in columns + ) + conn.execute(f""" + CREATE TABLE stats AS + WITH per_cluster AS ( + SELECT cluster_id, count(*) AS size, {splits} AS discord + FROM records GROUP BY cluster_id), + per_edge AS ( + SELECT cluster_id, count(*) AS edge_count, min(weight) AS weakest + FROM edges GROUP BY cluster_id), + per_truth AS ( + SELECT r.cluster_id, + count(DISTINCT coalesce(CAST(t.grp AS VARCHAR), 'x' || r.uid)) + AS entities + FROM records r LEFT JOIN truth t ON t.uid = r.uid + GROUP BY r.cluster_id) + SELECT c.cluster_id AS id, c.size, c.discord, + coalesce(e.edge_count, 0) AS edge_count, e.weakest, + coalesce(t.entities, 0) AS entities + FROM per_cluster c + LEFT JOIN per_edge e ON e.cluster_id = c.cluster_id + LEFT JOIN per_truth t ON t.cluster_id = c.cluster_id + """) + conn.execute("CREATE INDEX stats_by_id ON stats (id)") + conn.execute("CREATE INDEX records_by_cluster ON records (cluster_id)") + conn.execute("CREATE INDEX records_by_uid ON records (uid)") + + +def announce(line): + """The default progress logger: one line, flushed, so a server started in + the background still shows what it is doing.""" + print(line, flush=True) + + +def open_cache(opts, log=announce): + """A read-only connection to the cache for `opts`, built first if it is + missing, stale, or `opts.rebuild` says so. Returns `(conn, id_column, columns)`.""" + id_column, columns = schema_columns(opts.schema) + stamp = fingerprint(opts, columns) + + fresh = opts.rebuild or not os.path.exists(opts.cache) + if not fresh: + try: + conn = connect(opts.cache, read_only=True) + fresh = conn.execute("SELECT stamp FROM meta").fetchone()[0] != stamp + conn.close() + except Exception: + fresh = True + + if fresh: + if os.path.exists(opts.cache): + os.remove(opts.cache) + start = time.time() + conn = connect(opts.cache) + conn.execute(f"SET memory_limit = '{opts.memory}'") + conn.execute("SET preserve_insertion_order = false") + if opts.threads: + conn.execute(f"SET threads = {opts.threads}") + + def say(what): + log(f" {time.time() - start:6.1f}s {what}") + + log(f"building {opts.cache}") + build(conn, opts, id_column, columns, say) + conn.execute("CREATE TABLE meta (stamp VARCHAR)") + conn.execute("INSERT INTO meta VALUES (?)", [stamp]) + conn.close() + size = os.path.getsize(opts.cache) / 1e6 + log(f" {time.time() - start:6.1f}s done, {size:.0f} MB") + + conn = connect(opts.cache, read_only=True) + conn.execute(f"SET memory_limit = '{opts.memory}'") + return conn, id_column, columns diff --git a/python/cpplink_viewer/inputs.py b/python/cpplink_viewer/inputs.py new file mode 100644 index 0000000..77755f2 --- /dev/null +++ b/python/cpplink_viewer/inputs.py @@ -0,0 +1,178 @@ +# Copyright 2026 Mathieu Fourment +# SPDX-License-Identifier: MIT +"""The run's files as the viewer reads them: schema, shards, truth, model, waterfalls. + +Nothing here imports the compiled module. The viewer reads what a run wrote and +so runs against a checkout with no build, or against a wheel with no Arrow. +""" + +from __future__ import annotations + +import json + +import numpy as np +import pyarrow as pa +import pyarrow.compute as pc +import pyarrow.csv as pacsv + +# The binary shard `predict` writes: `kEdgeMagic` and one `kEdgeBytes` record per +# prediction, (row a, row b, packed gamma, match weight). A test reads a shard the +# core wrote through this, so the two cannot drift silently. +EDGE_MAGIC = b"CPPLNKE1" +EDGE_DTYPE = np.dtype([("a", " --out ` writes, one +wide row per prediction; it is loaded into the cache beside the predictions and +a click is one lookup, so the ledger is the scorer's own arithmetic and nothing +is computed or spawned here. The level labels and rates come from `--model`. +Predictions clustering kept apart, above the write threshold and below the +clustering one, are kept and listed under both of their clusters as `rejected`. + +`--threshold` keeps only the predictions at or above it. With `--clusters` the +file is taken as what `cpplink cluster --threshold` wrote at that threshold and +read as it stands, which is the fast path; without one the predictions are +clustered here, by the same union-find over the same predictions in the same +order, so the cache holds exactly the partition that command writes, named by +the same representatives. +""" + +from __future__ import annotations + +import argparse +import json +import os +import sys +import threading +import webbrowser +from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer +from importlib import resources +from urllib.parse import parse_qs, urlparse + +from .cache import Options, announce, open_cache, quoted +from .inputs import cell, ledger, read_model_levels + +# The list's order: the chosen key in either direction, a cluster with no +# value for it (no predictions, so no weakest edge) last either way, and the +# ties broken the same way whichever direction the key runs. +ORDERS = { + "discord": ("discord", "size DESC, id"), + "size": ("size", "discord DESC, id"), + "weakest": ("weakest", "size DESC, id"), + "id": ("id", ""), +} + +FILTERS = { + "all": "TRUE", + "split": "discord > 0", + "transitive": "size * (size - 1) / 2 > edge_count", + "mixed": "entities > 1", +} + +REJECTED_LIMIT = 1000 # rejected predictions listed per cluster, weakest first +LIST_LIMIT = 500 # clusters a list request may ask for at once + + +def order_clause(sort, desc): + key, ties = ORDERS.get(sort, ORDERS["discord"]) + clause = f"{key} {'DESC' if desc else 'ASC'} NULLS LAST" + return clause + (", " + ties if ties else "") + + +def like_pattern(text): + escaped = text.lower().replace("\\", "\\\\").replace("%", "\\%").replace("_", "\\_") + return "%" + escaped + "%" + + +class Viewer: + """The queries the page makes, over one shared read-only connection.""" + + def __init__(self, conn, columns, source, max_rows, has_truth, levels=None): + self.conn = conn + self.columns = columns + self.source = source + self.max_rows = max_rows + self.has_truth = has_truth + self.levels = levels # the model's level table, when waterfalls are held + self.lock = threading.Lock() + with self.lock: + self.totals = conn.execute( + "SELECT count(*), coalesce(sum(size), 0) FROM stats" + ).fetchone() + self.pair_total = conn.execute("SELECT count(*) FROM pairs").fetchone()[0] + self.waterfall_columns = [ + r[0] + for r in conn.execute( + "SELECT column_name FROM information_schema.columns " + "WHERE table_name = 'waterfalls' ORDER BY ordinal_position" + ).fetchall() + ] + + @classmethod + def open(cls, opts, log=announce): + """A viewer over the cache `opts` names, built first if it has to be.""" + reason = opts.check() + if reason: + raise SystemExit(reason) + conn, _, columns = open_cache(opts, log) + levels = read_model_levels(opts.model) if opts.model else None + source = ", ".join(os.path.basename(p) for p in opts.data) + return cls(conn, columns, source, opts.max_rows, bool(opts.truth), levels) + + def query(self, sql, params=()): + with self.lock: + return self.conn.execute(sql, params).fetchall() + + def summary(self): + return { + "columns": self.columns, + "source": self.source, + "totals": { + "clusters": self.totals[0], + "records": self.totals[1], + "pairs": self.pair_total, + }, + "has_truth": self.has_truth, + "has_model": bool(self.waterfall_columns), + } + + def waterfall(self, a, b): + """The ledger of one prediction, from the wide row the run's explain wrote.""" + if not self.waterfall_columns or self.levels is None: + return None + rows = self.query("SELECT * FROM waterfalls WHERE id_a = ? AND id_b = ?", [a, b]) + if not rows: + return None + row = dict(zip(self.waterfall_columns, rows[0], strict=True)) + return ledger(row, self.levels) + + def rejected(self, cluster_id): + """The predictions clustering overruled that touch one cluster. + + Each names the other cluster its second record went to, and the truth's + verdict where there is one, so the page can list them beside the + cluster's own edges without a lookup per row. + """ + rows = self.query( + "SELECT p.a, p.b, p.weight, p.cluster_a, p.cluster_b, ta.grp, tb.grp " + "FROM pairs p LEFT JOIN truth ta ON ta.uid = p.a " + "LEFT JOIN truth tb ON tb.uid = p.b " + "WHERE (p.cluster_a = ? OR p.cluster_b = ?) AND p.cluster_a <> p.cluster_b " + f"ORDER BY p.weight, p.a, p.b LIMIT {REJECTED_LIMIT + 1}", + [cluster_id, cluster_id], + ) + more = max(0, len(rows) - REJECTED_LIMIT) + found = [] + for a, b, w, ca, cb, ga, gb in rows[:REJECTED_LIMIT]: + same = (ga is not None and ga == gb) if self.has_truth else None + found.append( + { + "a": a, + "b": b, + "w": round(w, 3), + "other": cb if ca == cluster_id else ca, + "same": same, + } + ) + return found, more + + def pair(self, a, b): + head = self.query( + "SELECT a, b, weight, cluster_a, cluster_b FROM pairs WHERE a = ? AND b = ?", + [a, b], + ) + if not head: + return None + picked = ", ".join(quoted(c) for c in self.columns) + rows = { + r[0]: r + for r in self.query( + f"SELECT uid, {picked} FROM records WHERE uid IN (?, ?)", [a, b] + ) + } + if a not in rows or b not in rows: + return None + answer = { + "a": a, + "b": b, + "w": round(head[0][2], 3), + "ca": head[0][3], + "cb": head[0][4], + "rows": [ + {"id": uid, "v": [cell(v) for v in rows[uid][1:]]} for uid in (a, b) + ], + } + if self.has_truth: + groups = dict( + self.query("SELECT uid, grp FROM truth WHERE uid IN (?, ?)", [a, b]) + ) + same = a in groups and b in groups and groups[a] == groups[b] + answer["rows"][0]["t"] = 0 + answer["rows"][1]["t"] = 0 if same else 1 + answer["waterfall"] = self.waterfall(a, b) + return answer + + def listing(self, q, field, column, sort, desc, only, offset, limit): + """One page of the cluster list. + + `field` says what `q` is read against: `any` is a substring of the + concatenated record, `cluster` of the cluster id, `col` of the named + column, and `id` a comma-separated list of record ids matched whole, in + which case the answer also says which of them named a record. + """ + where = FILTERS.get(only, "TRUE") + params = [] + found = None + if q and field == "id": + ids = list(dict.fromkeys(s.strip() for s in q.split(",") if s.strip())) + marks = ", ".join("?" * len(ids)) + where += ( + f" AND id IN (SELECT cluster_id FROM records " + f"WHERE CAST(uid AS VARCHAR) IN ({marks}))" + ) + params.extend(ids) + found = [ + r[0] + for r in self.query( + f"SELECT DISTINCT CAST(uid AS VARCHAR) FROM records " + f"WHERE CAST(uid AS VARCHAR) IN ({marks})", + ids, + ) + ] + elif q and field == "cluster": + where += " AND lower(id) LIKE ? ESCAPE '\\'" + params.append(like_pattern(q)) + elif q and field == "col": + if column not in self.columns: + raise ValueError(f"no column {column!r}") + where += ( + f" AND id IN (SELECT cluster_id FROM records " + f"WHERE lower(CAST({quoted(column)} AS VARCHAR)) LIKE ? ESCAPE '\\')" + ) + params.append(like_pattern(q)) + elif q: + where += ( + " AND id IN (SELECT cluster_id FROM records " + "WHERE text LIKE ? ESCAPE '\\')" + ) + params.append(like_pattern(q)) + matched = self.query(f"SELECT count(*) FROM stats WHERE {where}", params)[0][0] + rows = self.query( + f"SELECT id, size, discord, edge_count, weakest, entities FROM stats " + f"WHERE {where} ORDER BY {order_clause(sort, desc)} " + f"LIMIT {int(limit)} OFFSET {int(offset)}", + params, + ) + answer = { + "matched": matched, + "clusters": [ + { + "id": r[0], + "size": r[1], + "discord": r[2], + "edge_count": r[3], + "weakest": r[4], + "entities": r[5], + } + for r in rows + ], + } + if found is not None: + answer["found"] = found + return answer + + def cluster(self, cluster_id): + head = self.query( + "SELECT id, size, discord, edge_count, weakest, entities " + "FROM stats WHERE id = ?", + [cluster_id], + ) + if not head: + return None + picked = ", ".join(quoted(c) for c in self.columns) + rows = self.query( + f"SELECT uid, {picked} FROM records WHERE cluster_id = ? " + f"ORDER BY uid LIMIT {int(self.max_rows)}", + [cluster_id], + ) + ids = [r[0] for r in rows] + groups, labels = {}, {} + if self.has_truth: + for uid, grp in self.query( + "SELECT uid, grp FROM truth WHERE uid IN " + "(SELECT uid FROM records WHERE cluster_id = ?)", + [cluster_id], + ): + groups[uid] = grp + for uid in ids: + key = groups.get(uid, "x" + uid) + labels.setdefault(key, len(labels)) + # The whole cluster's distinct counts, not the page of rows shown. + counts = self.query( + "SELECT " + + ", ".join(f"count(DISTINCT {quoted(c)})" for c in self.columns) + + " FROM records WHERE cluster_id = ?", + [cluster_id], + )[0] + edges = self.query( + "SELECT a, b, weight FROM edges WHERE cluster_id = ? AND a IN " + "(SELECT uid FROM records WHERE cluster_id = ?)", + [cluster_id, cluster_id], + ) + keep = set(ids) + answer = { + "id": head[0][0], + "size": head[0][1], + "discord": head[0][2], + "edge_count": head[0][3], + "weakest": head[0][4], + "entities": head[0][5], + "dist": list(counts), + "rows": [{"id": r[0], "v": [cell(v) for v in r[1:]]} for r in rows], + "edges": [ + [a, b, round(w, 3)] for a, b, w in edges if a in keep and b in keep + ], + } + if self.has_truth: + for row in answer["rows"]: + row["t"] = labels[groups.get(row["id"], "x" + row["id"])] + answer["rejected"], more = self.rejected(cluster_id) + if more: + answer["rejected_more"] = more + return answer + + +def page_html(): + """The viewer page, shipped beside this module.""" + return ( + resources.files(__package__).joinpath("viewer.html").read_text(encoding="utf-8") + ) + + +def handler_for(viewer, page): + class Handler(BaseHTTPRequestHandler): + def log_message(self, *_): + pass + + def send_json(self, payload, status=200): + body = json.dumps(payload).encode() + self.send_response(status) + self.send_header("Content-Type", "application/json") + self.send_header("Content-Length", str(len(body))) + self.end_headers() + self.wfile.write(body) + + def do_GET(self): + url = urlparse(self.path) + query = {k: v[0] for k, v in parse_qs(url.query).items()} + try: + if url.path in ("/", "/index.html"): + body = page.encode() + self.send_response(200) + self.send_header("Content-Type", "text/html; charset=utf-8") + self.send_header("Content-Length", str(len(body))) + self.end_headers() + self.wfile.write(body) + elif url.path == "/api/summary": + self.send_json(viewer.summary()) + elif url.path == "/api/list": + self.send_json( + viewer.listing( + query.get("q", ""), + query.get("field", "any"), + query.get("column", ""), + query.get("sort", "discord"), + query.get("dir", "desc") == "desc", + query.get("only", "all"), + int(query.get("offset", 0)), + min(int(query.get("limit", 100)), LIST_LIMIT), + ) + ) + elif url.path == "/api/cluster": + found = viewer.cluster(query.get("id", "")) + self.send_json( + found or {"error": "no such cluster"}, 200 if found else 404 + ) + elif url.path == "/api/pair": + found = viewer.pair(query.get("a", ""), query.get("b", "")) + self.send_json( + found or {"error": "no such prediction"}, 200 if found else 404 + ) + else: + self.send_json({"error": "not found"}, 404) + except BrokenPipeError: + pass + except Exception as trouble: # a bad query should not kill the server + self.send_json({"error": str(trouble)}, 500) + + return Handler + + +def make_server(viewer, port=0, host="127.0.0.1"): + """An HTTP server over `viewer`, bound and not yet serving; port 0 picks a + free one, which is what a test wants.""" + return ThreadingHTTPServer((host, port), handler_for(viewer, page_html())) + + +def parse_args(argv=None): + ap = argparse.ArgumentParser( + prog="cpplink-viewer", + description=__doc__, + formatter_class=argparse.RawDescriptionHelpFormatter, + ) + ap.add_argument("data", nargs="+", help="the parquet file(s) the run read, in order") + ap.add_argument("--schema", required=True) + ap.add_argument("--clusters", help="the csv or parquet cpplink cluster wrote") + ap.add_argument( + "--predictions", + help="the run's predictions -- one csv or parquet file, or the shard " + "directory -- for within-cluster weights", + ) + ap.add_argument("--truth", help="known pairs, to colour members by true entity") + ap.add_argument( + "--waterfalls", + help="the csv or parquet cpplink explain --predictions wrote; draws each " + "prediction's waterfall", + ) + ap.add_argument( + "--model", + help="the model the run scored with, for the level labels and rates the " + "waterfall shows", + ) + ap.add_argument("--cache", default=Options.cache) + ap.add_argument("--rebuild", action="store_true", help="rebuild the cache first") + ap.add_argument("--min-size", type=int, default=Options.min_size) + ap.add_argument( + "--max-size", type=int, default=Options.max_size, help="0 = no ceiling" + ) + ap.add_argument( + "--threshold", + type=float, + help="keep only the predictions at or above this match_weight; with " + "--clusters that file is read as the clustering at this threshold, " + "without one the predictions are clustered here exactly as " + "cpplink cluster --threshold would", + ) + ap.add_argument( + "--max-rows", + type=int, + default=Options.max_rows, + help="members shown per cluster; the rest are counted only", + ) + ap.add_argument( + "--memory", default=Options.memory, help="what DuckDB may use while building" + ) + ap.add_argument( + "--threads", type=int, default=Options.threads, help="0 = DuckDB's own default" + ) + ap.add_argument("--port", type=int, default=8770) + ap.add_argument("--open", action="store_true", help="open a browser on it") + args = ap.parse_args(argv) + opts = Options( + data=args.data, + schema=args.schema, + clusters=args.clusters, + predictions=args.predictions, + truth=args.truth, + waterfalls=args.waterfalls, + model=args.model, + cache=args.cache, + rebuild=args.rebuild, + min_size=args.min_size, + max_size=args.max_size, + threshold=args.threshold, + max_rows=args.max_rows, + memory=args.memory, + threads=args.threads, + ) + reason = opts.check() + if reason: + ap.error(reason) + return opts, args.port, args.open + + +def main(argv=None): + opts, port, open_browser = parse_args(argv) + viewer = Viewer.open(opts) + server = make_server(viewer, port) + where = f"http://127.0.0.1:{server.server_address[1]}/" + clusters, records = viewer.totals + announce(f"{clusters:,} clusters over {records:,} records at {where}") + if open_browser: + webbrowser.open(where) + try: + server.serve_forever() + except KeyboardInterrupt: + print("\nstopped") + return 0 + + +if __name__ == "__main__": + sys.exit(main()) diff --git a/tools/cluster_view_template.html b/python/cpplink_viewer/viewer.html similarity index 87% rename from tools/cluster_view_template.html rename to python/cpplink_viewer/viewer.html index e34c221..4384a90 100644 --- a/tools/cluster_view_template.html +++ b/python/cpplink_viewer/viewer.html @@ -274,16 +274,9 @@

cpplink clusters

Pick a cluster.