From 84eed8f55f6c3fe39c4a59fdb3c81ea1f3382075 Mon Sep 17 00:00:00 2001 From: 4ment Date: Wed, 23 Sep 2026 19:58:46 +1000 Subject: [PATCH 1/5] Add search functionality and related structures - Introduced `search.hpp` with definitions for query handling, search options, and search reports. - Implemented `Searcher` class for executing searches against a record store. - Added `QueryField` and `QueryRecord` structures to represent search queries. - Implemented `ParseQueryField` function for parsing query strings. - Added `SearchOptions` and `SearchReport` structures to manage search configurations and results. - Implemented `PrintSearchReport` and `SearchReportJson` functions for reporting search results. - Added `GroupHitsByCluster` function to label search hits by their respective clusters. - Extended `SignatureTable` in `signature.cpp` to accommodate dynamic dictionary sizes. - Created comprehensive unit tests for search functionality in `search_test.cpp`, covering various scenarios and edge cases. --- README.md | 7 + bench/scale/search_latency.json | 338 +++++++++++++ bench/scale/search_latency.py | 128 +++++ docs/commands/index.md | 3 + docs/commands/search.md | 202 ++++++++ mkdocs.yml | 1 + src/cpplink/app.cpp | 149 ++++++ src/cpplink/comparison.cpp | 22 + src/cpplink/comparison.hpp | 27 + src/cpplink/dictionary.cpp | 13 + src/cpplink/dictionary.hpp | 15 + src/cpplink/score.cpp | 8 +- src/cpplink/search.cpp | 842 ++++++++++++++++++++++++++++++++ src/cpplink/search.hpp | 218 +++++++++ src/cpplink/signature.cpp | 12 + src/cpplink/signature.hpp | 4 + tests/search_test.cpp | 571 ++++++++++++++++++++++ 17 files changed, 2559 insertions(+), 1 deletion(-) create mode 100644 bench/scale/search_latency.json create mode 100644 bench/scale/search_latency.py create mode 100644 docs/commands/search.md create mode 100644 src/cpplink/search.cpp create mode 100644 src/cpplink/search.hpp create mode 100644 tests/search_test.cpp diff --git a/README.md b/README.md index c21ce20..afd2efc 100644 --- a/README.md +++ b/README.md @@ -109,6 +109,13 @@ print(linker.last_cluster.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. +A model is also a search index. `cpplink search` takes a query record and returns the records that score highest against it, top-k over the same weight `predict` writes, and it answers in tens of milliseconds at four million records because interning makes a string level a question about a *value* rather than about a row. + +```sh +cpplink search --schema schema.json --model model.json sample.parquet \ + --field last_name=zolnerowich --field dob=1979-08-06 --expected-matches 1 +``` + 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/bench/scale/search_latency.json b/bench/scale/search_latency.json new file mode 100644 index 0000000..ad1b127 --- /dev/null +++ b/bench/scale/search_latency.json @@ -0,0 +1,338 @@ +[ + { + "rows": 250000, + "query": "exact", + "threads": 1, + "walk_ms": 0.19808299999999998, + "gather_ms": 10.897167, + "values_walked": 9000, + "rescored": 377, + "tabulated": 1, + "evaluated": 1, + "constant": 7, + "top": "r0", + "top_weight": 11.717544504910038 + }, + { + "rows": 250000, + "query": "exact", + "threads": 8, + "walk_ms": 0.31670899999999996, + "gather_ms": 2.720917, + "values_walked": 9000, + "rescored": 2957, + "tabulated": 1, + "evaluated": 1, + "constant": 7, + "top": "r0", + "top_weight": 11.717544504910038 + }, + { + "rows": 250000, + "query": "full", + "threads": 1, + "walk_ms": 58.307625, + "gather_ms": 14.936166, + "values_walked": 1274236, + "rescored": 234, + "tabulated": 6, + "evaluated": 1, + "constant": 2, + "top": "r0", + "top_weight": 70.46587527914643 + }, + { + "rows": 250000, + "query": "full", + "threads": 8, + "walk_ms": 13.3905, + "gather_ms": 3.3235, + "values_walked": 1274236, + "rescored": 1122, + "tabulated": 6, + "evaluated": 1, + "constant": 2, + "top": "r0", + "top_weight": 70.46587527914643 + }, + { + "rows": 250000, + "query": "fuzzy", + "threads": 1, + "walk_ms": 4.767792, + "gather_ms": 7.853333, + "values_walked": 100787, + "rescored": 53869, + "tabulated": 1, + "evaluated": 0, + "constant": 8, + "top": "r0", + "top_weight": 5.36650263108208 + }, + { + "rows": 250000, + "query": "fuzzy", + "threads": 8, + "walk_ms": 1.330916, + "gather_ms": 5.376417, + "values_walked": 100787, + "rescored": 245007, + "tabulated": 1, + "evaluated": 0, + "constant": 8, + "top": "r0", + "top_weight": 5.36650263108208 + }, + { + "rows": 1000000, + "query": "exact", + "threads": 1, + "walk_ms": 0.42812500000000003, + "gather_ms": 44.227959, + "values_walked": 9000, + "rescored": 747, + "tabulated": 1, + "evaluated": 1, + "constant": 7, + "top": "r0", + "top_weight": 10.207822943323324 + }, + { + "rows": 1000000, + "query": "exact", + "threads": 8, + "walk_ms": 0.558833, + "gather_ms": 10.704084000000002, + "values_walked": 9000, + "rescored": 3339, + "tabulated": 1, + "evaluated": 1, + "constant": 7, + "top": "r0", + "top_weight": 10.207822943323324 + }, + { + "rows": 1000000, + "query": "full", + "threads": 1, + "walk_ms": 215.632667, + "gather_ms": 58.280333, + "values_walked": 4796686, + "rescored": 272, + "tabulated": 6, + "evaluated": 1, + "constant": 2, + "top": "r0", + "top_weight": 72.45608783602324 + }, + { + "rows": 1000000, + "query": "full", + "threads": 8, + "walk_ms": 46.683583, + "gather_ms": 13.764042, + "values_walked": 4796686, + "rescored": 1414, + "tabulated": 6, + "evaluated": 1, + "constant": 2, + "top": "r0", + "top_weight": 72.45608783602324 + }, + { + "rows": 1000000, + "query": "fuzzy", + "threads": 1, + "walk_ms": 7.484375, + "gather_ms": 28.290792, + "values_walked": 149192, + "rescored": 54029, + "tabulated": 1, + "evaluated": 0, + "constant": 8, + "top": "r0", + "top_weight": 2.4799602682808395 + }, + { + "rows": 1000000, + "query": "fuzzy", + "threads": 8, + "walk_ms": 2.4970830000000004, + "gather_ms": 13.482959, + "values_walked": 149192, + "rescored": 396874, + "tabulated": 1, + "evaluated": 0, + "constant": 8, + "top": "r0", + "top_weight": 2.4799602682808395 + }, + { + "rows": 4000000, + "query": "exact", + "threads": 1, + "walk_ms": 1.325875, + "gather_ms": 174.675042, + "values_walked": 9000, + "rescored": 861, + "tabulated": 1, + "evaluated": 1, + "constant": 7, + "top": "r0", + "top_weight": 8.068591826542518 + }, + { + "rows": 4000000, + "query": "exact", + "threads": 8, + "walk_ms": 1.632416, + "gather_ms": 38.276083, + "values_walked": 9000, + "rescored": 4795, + "tabulated": 1, + "evaluated": 1, + "constant": 7, + "top": "r0", + "top_weight": 8.068591826542518 + }, + { + "rows": 4000000, + "query": "full", + "threads": 1, + "walk_ms": 861.101667, + "gather_ms": 242.876541, + "values_walked": 18718486, + "rescored": 324, + "tabulated": 6, + "evaluated": 1, + "constant": 2, + "top": "r0", + "top_weight": 73.44513500940663 + }, + { + "rows": 4000000, + "query": "full", + "threads": 8, + "walk_ms": 200.91, + "gather_ms": 58.472625, + "values_walked": 18718486, + "rescored": 1961, + "tabulated": 6, + "evaluated": 1, + "constant": 2, + "top": "r0", + "top_weight": 73.44513500940663 + }, + { + "rows": 4000000, + "query": "fuzzy", + "threads": 1, + "walk_ms": 12.12425, + "gather_ms": 109.033125, + "values_walked": 237310, + "rescored": 54046, + "tabulated": 1, + "evaluated": 0, + "constant": 8, + "top": "r0", + "top_weight": 0.6526652113186122 + }, + { + "rows": 4000000, + "query": "fuzzy", + "threads": 8, + "walk_ms": 4.594708000000001, + "gather_ms": 33.168375, + "values_walked": 237310, + "rescored": 307665, + "tabulated": 1, + "evaluated": 0, + "constant": 8, + "top": "r0", + "top_weight": 0.6526652113186122 + }, + { + "rows": 20000000, + "query": "exact", + "threads": 1, + "walk_ms": 13.007166, + "gather_ms": 848.898125, + "values_walked": 9000, + "rescored": 886, + "tabulated": 1, + "evaluated": 1, + "constant": 7, + "top": "r0", + "top_weight": 5.719266274137542 + }, + { + "rows": 20000000, + "query": "exact", + "threads": 8, + "walk_ms": 8.869084, + "gather_ms": 186.34825, + "values_walked": 9000, + "rescored": 8753, + "tabulated": 1, + "evaluated": 1, + "constant": 7, + "top": "r0", + "top_weight": 5.719266274137542 + }, + { + "rows": 20000000, + "query": "full", + "threads": 1, + "walk_ms": 4194.0, + "gather_ms": 1308.0, + "values_walked": 91645846, + "rescored": 636, + "tabulated": 6, + "evaluated": 1, + "constant": 2, + "top": "r0", + "top_weight": 75.98 + }, + { + "rows": 20000000, + "query": "full", + "threads": 8, + "walk_ms": 901.0, + "gather_ms": 293.0, + "values_walked": 91645846, + "rescored": 2249, + "tabulated": 6, + "evaluated": 1, + "constant": 2, + "top": "r0", + "top_weight": 75.98 + }, + { + "rows": 20000000, + "query": "fuzzy", + "threads": 1, + "walk_ms": 42.3645, + "gather_ms": 555.802167, + "values_walked": 657528, + "rescored": 54123, + "tabulated": 1, + "evaluated": 0, + "constant": 8, + "top": "r0", + "top_weight": -1.5300124376597202 + }, + { + "rows": 20000000, + "query": "fuzzy", + "threads": 8, + "walk_ms": 28.6, + "gather_ms": 137.6, + "values_walked": 657528, + "rescored": 356775, + "tabulated": 1, + "evaluated": 0, + "constant": 8, + "top": "r0", + "top_weight": -1.5300124376597202 + } +] diff --git a/bench/scale/search_latency.py b/bench/scale/search_latency.py new file mode 100644 index 0000000..562f94b --- /dev/null +++ b/bench/scale/search_latency.py @@ -0,0 +1,128 @@ +# Copyright 2026 Mathieu Fourment +# SPDX-License-Identifier: MIT +"""Per-query latency of `cpplink search`, swept over store size and thread count. + +The question the sweep answers is which of the two phases dominates, because that +is what decides whether the next step is an inverted index or a q-gram index over +the dictionary. Phase one walks each queried column's dictionary once, so it grows +with the number of distinct values; phase two gathers a level per row, so it grows +with the rows. They are reported separately. + +Three queries are run at every size, because a query's cost is a property of the +columns it names and not of the store: + + exact a near-unique column with no fuzzy level reached (postcode, dob) + fuzzy one name column, which is a dictionary walk with a metric in it + full a whole record, which is every column the schema compares + +Usage: search_latency.py [ ...] +It writes /search_latency.json and prints the table. The workdir holds +s.parquet and s.model.json for every size; `gen-sample` writes the +first and `estimate` the second. + +The recorded sweep is search_latency.json beside this file: 250k, 1M, 4M and 20M, +one thread and eight. The 20M rows there are single runs rather than the median of +five, because the store is 4.5 GB and reloading it per repeat is most of the wall +clock. +""" + +import json +import pathlib +import statistics +import subprocess +import sys + +ROOT = pathlib.Path(__file__).resolve().parents[2] +CPPLINK = ROOT / "build" / "cpplink" +SCHEMA = ROOT / "examples" / "sample_schema.json" + +# The query record is row r0 of every `gen-sample` file, which is generated from +# the same seed at every size, so the three queries name the same person +# throughout and the only thing changing is how much data it is asked about. +QUERIES = { + "exact": ["dob=1979-08-06", "postcode=4508"], + "fuzzy": ["last_name=chirdloackki"], + "full": [ + "first_name=zoudcut", + "last_name=chirdloackki", + "gender=F", + "dob=1979-08-06", + "email=zoudcut.chirdloackki3970@example.com", + "phone=0404342668", + "postcode=4508", + ], +} +REPEATS = 5 + + +def run(work, rows, name, fields, threads): + path = work / f"s{rows}.parquet" + model = work / f"s{rows}.model.json" + command = [ + str(CPPLINK), + "search", + "--schema", + str(SCHEMA), + "--model", + str(model), + str(path), + "--threads", + str(threads), + "-k", + "10", + "--json", + ] + for field in fields: + command += ["--field", field] + walks, gathers = [], [] + report = {} + for _ in range(REPEATS): + result = subprocess.run(command, capture_output=True, text=True, check=True) + report = json.loads(result.stdout) + walks.append(report["walk_seconds"]) + gathers.append(report["gather_seconds"]) + return { + "rows": rows, + "query": name, + "threads": threads, + "walk_ms": 1000 * statistics.median(walks), + "gather_ms": 1000 * statistics.median(gathers), + "values_walked": report["values_walked"], + "rescored": report["rescored"], + "tabulated": report["tabulated"], + "evaluated": report["evaluated"], + "constant": report["constant"], + "top": report["hits"][0]["id"] if report["hits"] else None, + "top_weight": report["hits"][0]["match_weight"] if report["hits"] else None, + } + + +def main(): + if len(sys.argv) < 3: + print(__doc__) + return 1 + work = pathlib.Path(sys.argv[1]) + sizes = [int(value) for value in sys.argv[2:]] + rows = [] + for size in sizes: + for name, fields in QUERIES.items(): + for threads in (1, 8): + rows.append(run(work, size, name, fields, threads)) + (work / "search_latency.json").write_text(json.dumps(rows, indent=2)) + + header = f"{'rows':>12} {'query':<6} {'thr':>4} {'walk ms':>9} {'gather ms':>10}" + header += f" {'total ms':>9} {'values':>12} {'rescored':>10}" + print(header) + print("-" * len(header)) + for row in rows: + print( + f"{row['rows']:>12,} {row['query']:<6} {row['threads']:>4}" + f" {row['walk_ms']:>9.2f} {row['gather_ms']:>10.2f}" + f" {row['walk_ms'] + row['gather_ms']:>9.2f}" + f" {row['values_walked']:>12,} {row['rescored']:>10,}" + ) + return 0 + + +if __name__ == "__main__": + sys.exit(main()) diff --git a/docs/commands/index.md b/docs/commands/index.md index 67e8858..5887448 100644 --- a/docs/commands/index.md +++ b/docs/commands/index.md @@ -23,6 +23,8 @@ commands: merged prediction file or a shard directory merge-predictions combine the prediction shards into one csv or parquet file rescore re-score a spilled run under a new model, without comparing again + search find the records a query record scores highest against, + top-k over the same weight predict writes gen-sample write a sample parquet file with planted duplicates options: @@ -48,6 +50,7 @@ one store and each becomes a dataset, so two files mean linking rather than dedu | [`estimate`](estimate.md) | What are `m`, `u` and `λ`? | parquet + schema | `model.json` | | [`completeness`](completeness.md) | What fraction of true matches does blocking reach, with no truth file? | parquet + schema + model | stdout | | [`predict`](predict.md) | Which pairs score above the threshold? | parquet + schema + model | one prediction file, or one shard per thread | +| [`search`](search.md) | Which records score highest against *this* query record? | parquet + schema + model | stdout | | [`rescore`](rescore.md) | What would a different model have scored? | parquet + schema + model + spill | the same, one file or shards | | [`cluster`](cluster.md) | Which records are the same entity? | parquet + schema + predictions (a file or a shard directory) | `clusters.csv` | | [`merge-predictions`](merge-predictions.md) | Give me the predictions as one file something else can open. | prediction shards (+ parquet + schema for the ids) | one csv or parquet file | diff --git a/docs/commands/search.md b/docs/commands/search.md new file mode 100644 index 0000000..38ad44a --- /dev/null +++ b/docs/commands/search.md @@ -0,0 +1,202 @@ +# `search` + +**Goal:** given a query record, return the records that score highest against it. + +This is not linkage, it is retrieval. The match weight is a sum of per-column terms, which is +the shape a search engine's ranking function has, so "find this person in 20M records" is +top-k retrieval over an additive score. What makes it cheap here is interning: a string level +is a function of two *value ids*, so the level every distinct value of a column lands on +against the query is one table, computed once per query, after which a row costs a load and a +compare rather than a string metric. + +The answer is exact with respect to the model. It is the same k records, in the same order, +that scoring the query against every record one at a time would return. + +## Synopsis + +```sh +cpplink search --schema --model + --field = [--field ...] + [-k N] [--threshold BITS] [--threads N] + [--expected-matches N | --prior-weight BITS] + [--clusters ] [--explain] [--json] +``` + +| Option | Default | Meaning | +| --- | --- | --- | +| `--schema ` | — | required | +| `--model ` | — | required; the model whose `m`, `u` and levels the weight is built from | +| `--field =` | — | required, repeatable; one field of the query record | +| `-k N`, `--top N` | 10 | how many records to return | +| `--threshold BITS` | none | drop records below this match weight, however few are left | +| `--threads N` | 1 | splits both the dictionary walk and the row pass | +| `--expected-matches N` | — | the prior, as the number of records in the store you expect to be the person asked about | +| `--prior-weight BITS` | — | the same thing stated directly, in bits of prior odds | +| `--clusters ` | — | a file written by [`cluster`](cluster.md); labels each hit with its cluster | +| `--explain` | off | print the [waterfall](explain.md) behind every hit | +| `--json` | off | the whole report as one JSON object | +| `--tf-damping F` | 1.0 | scale the term-frequency adjustment | +| `--no-interactions` | off | score the plain conditionally-independent model | +| *(positional)* | — | required; the parquet file | + +A column the query does not name is **missing**, not empty. So a query carrying a name and a +date of birth against a ten-column schema is scored as a record whose other eight fields were +never recorded, which is what it is, and those comparisons land on their null level. + +A list column takes one element per repeat of its field: `--field alias=bill --field +alias=will`. A date is written `YYYY-MM-DD` and is refused in any other form. + +## What it prints + +```text +Records 1,000,000 +Comparisons 3 tabulated over their dictionary, 1 evaluated per row, 5 constant because the query is missing the column +Values walked 174,418 (1 query value the store never held) +Rows scored 81 of 1,000,000 exactly; the rest were dropped on the bracket +Prior -22.619 bits +Elapsed 0.0079 s walking the dictionaries, 0.0449 s over the rows on 1 thread + +Record weight posterior +r0 38.684 1.000000 +r259138 -7.916 0.004124 +``` + +The three counts on the `Comparisons` line are the three things a comparison can cost. + +- **constant** is free. The query does not carry the column, so every record lands on the + same level and nothing is evaluated at all. Most of a real query's schema is here, because + a caller who knows ten fields about someone is not searching for them. +- **tabulated** is one walk over the column's dictionary, then one byte read per record. This + is where every string metric in the query runs, and it runs once per *distinct value* + rather than once per record. +- **evaluated** is the pair path's own evaluation, once per record. A comparison over a date, + a number, a list or a coordinate pair is not a function of a single value id, so it gets no + table. + +`Rows scored` is how many records reached the exact weight. The rest were dropped on +`BaseWeight + Δ_max`, the same admissible bracket [`predict`](predict.md) uses, so they could +not have entered the answer whatever the term-frequency tables said. + +## The prior is not the model's, and it is the one thing you have to choose + +λ is the match rate over the **pair** space, which is the question deduplication asks. A query +against N records asks a different one, roughly "I expect this person to be in here about +once", and the two differ by orders of magnitude. By default `search` inherits the model's +prior, so a hit scores exactly what [`predict`](predict.md) would have given that pair. +`--expected-matches 1` states the search question instead: + +```sh +cpplink search --schema s.json --model model.json data.parquet \ + --field last_name=zolnerowich --field postcode=4508 --expected-matches 1 +``` + +The prior is a constant added to every hit, so it moves the posterior and where a sensible +threshold sits, and it never changes the order. + +## Why the weight is worth more than a ranking + +A search engine returns a list. This returns a list with a number attached, and the number is +calibrated, so "nobody in here is this person" is expressible rather than being an +empty-looking list of bad matches. That is what `--threshold` is for: zero bits is even odds, +so a hit below it is evidence *against*. + +And a hit explains itself through the same ledger as any other pair: + +```sh +cpplink search ... -k 1 --explain +``` + +```text +Comparison Level bits tf running +------------------------------------------------------------------------ +(prior) lambda -22.62 -22.62 +last_name jaro_winkler >= 0.92 +15.81 -6.80 +first_name exact +9.31 +3.35 +5.86 +gender null +0.36 +6.22 +dob exact +14.07 +0.01 +20.30 +email null +2.79 +23.09 +phone null +2.77 +25.87 +postcode exact +12.90 -0.08 +38.68 +location null +0.00 +38.68 +address null +0.00 +38.68 +------------------------------------------------------------------------ +Match weight 38.684 bits posterior 1.000000000 + +Term frequency, for the comparisons that moved the weight: + first_name zoudcut 115 rows 1.15e-04 +3.35 bits + dob 3504d 48 rows 4.80e-05 +0.01 bits + postcode 4508 115 rows 1.15e-04 -0.08 bits +``` + +This is [`explain`](explain.md)'s waterfall, over the same pair the search scored, from the +same scorer. Nothing here recomputes a bit of the weight. + +## Entities rather than rows + +A deduplicated store holds several records of the same person, so the top ten records can be +three people. `--clusters `, given the file [`cluster`](cluster.md) wrote, labels each +hit with its cluster and marks the ones that are another record of an entity already listed. + +Grouping at the end rather than indexing one record per cluster is deliberate: indexing a +representative would shrink the scan by the duplicate rate and lose every cluster whose +members disagree on the columns the query happens to carry. + +## What it costs + +Per query on the `gen-sample` sample against the nine-comparison sample schema. `exact` names a +date and a postcode, `fuzzy` one name column, `full` seven columns including the email address. + +| Records | Query | Walk | Gather | Total, 1 thread | Total, 8 threads | Values walked | +| ---: | --- | ---: | ---: | ---: | ---: | ---: | +| 250,000 | exact | 0.2 ms | 10.9 ms | 11.1 ms | 3.0 ms | 9,000 | +| 250,000 | fuzzy | 4.8 ms | 7.9 ms | 12.6 ms | 6.7 ms | 100,787 | +| 250,000 | full | 58.3 ms | 14.9 ms | 73.2 ms | 16.7 ms | 1,274,236 | +| 1,000,000 | exact | 0.4 ms | 44.2 ms | 44.7 ms | 11.3 ms | 9,000 | +| 1,000,000 | fuzzy | 7.5 ms | 28.3 ms | 35.8 ms | 16.0 ms | 149,192 | +| 1,000,000 | full | 215.6 ms | 58.3 ms | 273.9 ms | 60.5 ms | 4,796,686 | +| 4,000,000 | exact | 1.3 ms | 174.7 ms | 176.0 ms | 39.9 ms | 9,000 | +| 4,000,000 | fuzzy | 12.1 ms | 109.0 ms | 121.2 ms | 37.8 ms | 237,310 | +| 4,000,000 | full | 861.1 ms | 242.9 ms | 1,104 ms | 259.4 ms | 18,718,486 | +| 20,000,000 | exact | 13.0 ms | 848.9 ms | 862 ms | 195 ms | 9,000 | +| 20,000,000 | fuzzy | 42.4 ms | 555.8 ms | 598 ms | 166 ms | 657,528 | +| 20,000,000 | full | 4,194 ms | 1,308 ms | 5,502 ms | 1,194 ms | 91,645,846 | + +The two phases scale with different things, which is the point of reporting them apart. + +- The **row pass** grows with the records, at 27 to 65 ns each depending on how many + comparisons the query engages, and that rate holds from 250k to 20M. +- The **dictionary walk** grows with the queried columns' *distinct values*, at a flat 46 ns + each. On a near-unique column with a fuzzy level — an email address — that is the whole of + the latency: at 20M the `full` query walks 91.6M values and spends 76% of itself there, + while the same query without the email address is answered in 598 ms. + +So the thing to watch is not the size of the store but whether the query names a near-unique +column that a fuzzy level reads. Both phases split 4–5× over eight threads. + +The measurements are `bench/scale/search_latency.json`, produced by +`bench/scale/search_latency.py`, which is how to reproduce them on other data. The load is not +in these numbers: `search` is a scan over a resident store, and the sense in which that is a +per-query cost is that the process stays up. + +## Limits + +!!! warning "One input" + `search` reads a single parquet file. The query is appended to the store as a row, and + over more than one input that row would have to be a dataset of its own, which the levels + reading which input a record came from were bound before it existed. Linking a file of + queries against a store is the throughput question, and that is what + [`--mode link`](../linking.md) already answers. + +Two more, both narrow: + +- The query record is held as a row on the store for the length of the query, so one search + runs at a time against one store. Threads split the work *within* a query. +- A query value the store's dictionary never held is adopted for the length of the query, so + it compares fuzzily against everything as it should. It gets no entry in the + `list_contains` alias map, so a nickname the store never saw is not looked up as one. + +## See also + +- [`predict`](predict.md) — the same weight over candidate pairs rather than one query +- [`explain`](explain.md) — the waterfall `--explain` prints +- [`cluster`](cluster.md) — what `--clusters` reads diff --git a/mkdocs.yml b/mkdocs.yml index 07e9299..55f40f7 100644 --- a/mkdocs.yml +++ b/mkdocs.yml @@ -57,6 +57,7 @@ nav: - completeness: commands/completeness.md - estimate: commands/estimate.md - predict: commands/predict.md + - search: commands/search.md - rescore: commands/rescore.md - cluster: commands/cluster.md - merge-predictions: commands/merge-predictions.md diff --git a/src/cpplink/app.cpp b/src/cpplink/app.cpp index 0be8670..baf3d0e 100644 --- a/src/cpplink/app.cpp +++ b/src/cpplink/app.cpp @@ -39,6 +39,7 @@ #include "cpplink/sample_data.hpp" #include "cpplink/schema.hpp" #include "cpplink/score.hpp" +#include "cpplink/search.hpp" #include "cpplink/simplify.hpp" #include "cpplink/waterfall.hpp" @@ -100,6 +101,8 @@ void PrintUsage(std::ostream& out) { "parquet file\n" << " rescore re-score a spilled run under a new model, without " "comparing again\n" + << " search find the records a query record scores highest against,\n" + " top-k over the same weight predict writes\n" << " gen-sample write a sample parquet file with planted duplicates\n" << "\n" << "options:\n" @@ -1845,6 +1848,151 @@ int RunGenSample(const std::vector& args, std::ostream& out, return 0; } +// Top-k retrieval over the match weight. There is no blocking here and no plan: +// the query is compared against every record, which is affordable because the +// comparison is a table lookup per column rather than a string metric per row. +int RunSearch(const std::vector& args, std::ostream& out, + std::ostream& err) { + std::string schema_path; + std::string model_path; + std::string cluster_path; + std::vector data_paths; + QueryRecord query; + SearchOptions options; + ScoreOptions score; + std::string value; + bool explain = false; + bool as_json = false; + bool expected_given = false; + double expected = 0.0; + 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] == "--field") { + if (!TakeValue(args, &i, &value, err)) return 1; + QueryField field; + std::string error; + if (!ParseQueryField(value, &field, &error)) { + err << "cpplink search: " << error << "\n"; + return 1; + } + query.fields.push_back(field); + } else if (args[i] == "-k" || args[i] == "--top") { + if (!TakeValue(args, &i, &value, err)) return 1; + options.k = std::stoull(value); + } else if (args[i] == "--threshold") { + if (!TakeValue(args, &i, &value, err)) return 1; + options.threshold = std::stod(value); + } else if (args[i] == "--threads") { + if (!TakeValue(args, &i, &value, err)) return 1; + options.threads = static_cast(std::stoul(value)); + } else if (args[i] == "--expected-matches") { + // How many records of the person asked about the store is expected to + // hold. It becomes bits once the record count is known. + if (!TakeValue(args, &i, &value, err)) return 1; + options.override_prior = true; + expected_given = true; + expected = std::stod(value); + } else if (args[i] == "--prior-weight") { + if (!TakeValue(args, &i, &value, err)) return 1; + options.override_prior = true; + expected_given = false; + options.prior_weight = 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] == "--no-interactions") { + score.use_interactions = false; + } else if (args[i] == "--clusters") { + if (!TakeValue(args, &i, &cluster_path, err)) return 1; + } else if (args[i] == "--explain") { + explain = true; + } else if (args[i] == "--json") { + as_json = true; + } else if (!args[i].empty() && args[i][0] == '-') { + err << "cpplink search: unknown option '" << args[i] << "'\n"; + return 1; + } else { + data_paths.push_back(args[i]); + } + } + if (schema_path.empty() || data_paths.empty() || model_path.empty()) { + err << "cpplink search: --schema , --model and a " + "parquet file are required\n"; + return 1; + } + if (query.fields.empty()) { + err << "cpplink search: give at least one --field =\n"; + return 1; + } + Schema schema; + std::string error; + if (!LoadSchemaFor(schema_path, data_paths, &schema, &error)) { + err << "cpplink: " << error << "\n"; + return 1; + } + if (schema.comparisons.empty()) { + err << "cpplink search: the schema declares no \"comparisons\"\n"; + return 1; + } + RecordStore store(schema); + if (!LoadParquetFiles(data_paths, schema, &store, nullptr, &error)) { + err << "cpplink: " << error << "\n"; + return 1; + } + ComparisonSet comparisons; + if (!comparisons.Bind(schema, store, &error)) { + err << "cpplink: " << error << "\n"; + return 1; + } + Model model; + if (!LoadModel(model_path, &model, &error)) { + err << "cpplink: " << error << "\n"; + return 1; + } + Scorer scorer; + if (!scorer.Bind(model, comparisons, store, score, &error)) { + err << "cpplink search: " << error << "\n"; + return 1; + } + + Searcher searcher; + if (!searcher.Bind(&store, &comparisons, &scorer, &error)) { + err << "cpplink search: " << error << "\n"; + return 1; + } + if (expected_given) { + options.prior_weight = PriorWeightForExpected(expected, store.NumRecords()); + } + SearchReport report; + if (!searcher.Search(query, options, &report, &error)) { + err << "cpplink search: " << error << "\n"; + return 1; + } + if (!cluster_path.empty() && + !GroupHitsByCluster(store, cluster_path, &report, &error)) { + err << "cpplink search: " << error << "\n"; + return 1; + } + if (as_json) { + out << SearchReportJson(report) << "\n"; + return 0; + } + PrintSearchReport(report, out); + if (explain) { + // The store row is the first side throughout, so a term-frequency move + // reads the frequency of a value the table was built over. + for (const SearchHit& hit : report.hits) { + out << "\n"; + PrintPairWaterfall(store, comparisons, scorer, hit.row, searcher.QueryRow(), + out); + } + } + return 0; +} + } // namespace int Run(const std::vector& args, std::ostream& out, std::ostream& err) { @@ -1877,6 +2025,7 @@ int Run(const std::vector& args, std::ostream& out, std::ostream& e if (first == "rescore") return RunRescore(rest, out, err); if (first == "cluster") return RunCluster(rest, out, err); if (first == "merge-predictions") return RunMergeEdges(rest, out, err); + if (first == "search") return RunSearch(rest, out, err); if (first == "gen-sample") return RunGenSample(rest, out, err); err << "cpplink: unknown command '" << first << "'\n"; diff --git a/src/cpplink/comparison.cpp b/src/cpplink/comparison.cpp index 2a867e5..dabb608 100644 --- a/src/cpplink/comparison.cpp +++ b/src/cpplink/comparison.cpp @@ -70,6 +70,7 @@ bool ComparisonSet::Bind(const Schema& schema, const RecordStore& store, std::string* error, bool use_signatures, bool use_ladders) { bound_.clear(); tables_.clear(); + table_dicts_.clear(); width_ = 0; dataset_starts_ = store.dataset_starts(); // Several comparisons can read the same column, and the signatures belong to @@ -149,6 +150,7 @@ bool ComparisonSet::Bind(const Schema& schema, const RecordStore& store, if (entry.first == dict) return entry.second; } tables_.push_back(std::make_unique()); + table_dicts_.push_back(dict); tables_.back()->Build(*dict); built.emplace_back(dict, tables_.back().get()); return tables_.back().get(); @@ -228,6 +230,12 @@ bool ComparisonSet::Bind(const Schema& schema, const RecordStore& store, return true; } +void ComparisonSet::ResizeTables() { + for (size_t i = 0; i < tables_.size(); ++i) { + tables_[i]->Extend(*table_dicts_[i]); + } +} + uint64_t ComparisonSet::SignatureBytes() const { uint64_t bytes = 0; for (const auto& table : tables_) bytes += table->BytesUsed(); @@ -510,6 +518,20 @@ uint8_t ComparisonSet::LevelForValues(size_t comparison, uint32_t left, return static_cast(levels.size() - 1); } +bool ComparisonSet::StringLevelForValues(size_t comparison, size_t level, uint32_t left, + uint32_t right) const { + const BoundComparison& bound = bound_[comparison]; + const LevelSpec& spec = bound.spec->levels[level]; + if (spec.type == LevelType::kLevenshtein || spec.type == LevelType::kJaroWinkler) { + // One rung of the run, which is what the level would be asked alone. The + // screen is the run's, looser than this rung's threshold and so still + // admissible for it. + return StringRunLevel(bound, level, level + 1, left, right) == level; + } + if (spec.type == LevelType::kNull || spec.type == LevelType::kElse) return false; + return StringLevelFires(bound, spec, left, right); +} + bool ComparisonSet::LevelFires(const BoundComparison& comparison, size_t index, uint64_t a, uint64_t b) const { const LevelSpec& level = comparison.spec->levels[index]; diff --git a/src/cpplink/comparison.hpp b/src/cpplink/comparison.hpp index 14e6a0f..9d1c683 100644 --- a/src/cpplink/comparison.hpp +++ b/src/cpplink/comparison.hpp @@ -95,6 +95,16 @@ class ComparisonSet { // self-joined once instead of the pair stream being walked again. uint8_t LevelForValues(size_t comparison, uint32_t left, uint32_t right) const; + // Whether one string level fires for two values of the column *that level* + // reads. `LevelForValues` answers for a whole comparison and so cannot serve + // one whose levels read different columns -- an address against its username + // is the shape -- because the two ids it is given belong to one dictionary. + // Asked a level at a time, the question is well posed again, and a search + // tabulates such a comparison as one table per level. False for the null and + // else levels, which are not a value's business. + bool StringLevelForValues(size_t comparison, size_t level, uint32_t left, + uint32_t right) const; + // Whether a level could fire for this pair, decided without evaluating one // string metric: exact for the cheap level types, and the signature bounds // for the fuzzy ones. False means "certainly not"; true means "maybe". @@ -108,6 +118,19 @@ class ComparisonSet { // sampled. bool IsNullValue(size_t comparison, uint64_t row) const; + // Brings every per-value table this set holds to the size of the dictionary + // it was built from. + // + // A signature table is indexed by value id, so a dictionary that has grown + // past it would be read off the end. Nothing in a run grows one: this exists + // for the search path, where a query carries values the store never saw and + // adopts them into the column's dictionary for the length of the query. The + // membership levels' alias map needs no such call, because it is consulted + // only below the `alias_size` it was built at, and a value the map does not + // reach simply has no alias -- which is what a value the list column never + // held would get anyway. + void ResizeTables(); + size_t Size() const { return bound_.size(); } const BoundComparison& at(size_t index) const { return bound_[index]; } uint8_t Width() const { return width_; } @@ -155,6 +178,10 @@ class ComparisonSet { // ComparisonSet cannot be copied into one whose comparisons point at another's // tables. std::vector> tables_; + // The dictionary each table was built from, in the same order, so a table + // can be brought back to its dictionary's size without asking the caller + // which one it belongs to. + std::vector table_dicts_; // One per list_contains comparison, addressed through BoundComparison. Held // by pointer for the same reason the signature tables are: the vector may // grow, and a BoundComparison holds the data() of one of these. diff --git a/src/cpplink/dictionary.cpp b/src/cpplink/dictionary.cpp index 3f2c354..0e0171d 100644 --- a/src/cpplink/dictionary.cpp +++ b/src/cpplink/dictionary.cpp @@ -40,6 +40,19 @@ uint32_t Dictionary::Intern(std::string_view value) { return id; } +uint32_t Dictionary::Adopt(std::string_view value) { + const std::string_view stored = Store(value); + const uint32_t id = static_cast(entries_.size()); + entries_.push_back(stored); + text_bytes_ += stored.size(); + return id; +} + +void Dictionary::Truncate(uint32_t size) { + if (size >= entries_.size()) return; + entries_.resize(size); +} + uint64_t Dictionary::BytesUsed() const { uint64_t bytes = 0; // Chunks are fully allocated whether or not they are fully used. diff --git a/src/cpplink/dictionary.hpp b/src/cpplink/dictionary.hpp index 610950e..b2ce024 100644 --- a/src/cpplink/dictionary.hpp +++ b/src/cpplink/dictionary.hpp @@ -40,6 +40,21 @@ class Dictionary { // The index is only needed while loading; dropping it frees roughly half. void ReleaseIndex(); + // Appends a value without consulting the index, which a finalized dictionary + // no longer has, and without checking whether the text is already here. + // + // This is the search path and nothing else: a query carries a value the store + // may never have seen, and a fuzzy level needs its text and its signature, + // which only a value id addresses. The caller has already looked the text up + // -- `Adopt` cannot, having no index -- so an adopted id is either a value the + // dictionary genuinely lacks or a deliberate duplicate. + uint32_t Adopt(std::string_view value); + + // Drops every value from `size` on, undoing the adoptions one query made. The + // text stays in the arena, which is a few bytes a query and not worth a + // free-list to reclaim. + void Truncate(uint32_t size); + private: static constexpr size_t kChunkBytes = 1u << 20; diff --git a/src/cpplink/score.cpp b/src/cpplink/score.cpp index 79b61b5..2f2eb7d 100644 --- a/src/cpplink/score.cpp +++ b/src/cpplink/score.cpp @@ -43,7 +43,13 @@ uint32_t TermFrequencyAdjustment::Frequency(uint64_t row) const { if (dates != nullptr) { const int32_t value = dates->values[row]; if (value == kNullDate) return 0; - return dates->tf[static_cast(value - dates->tf_origin)]; + const size_t at = static_cast(value - dates->tf_origin); + // The table is dense over the range the *loaded* rows cover, and a search + // appends a query row that may carry a date outside it. Nothing the file + // holds can be out of range, so this costs a predicted compare and is + // only ever taken by a value the table was not built over. + if (at >= dates->tf.size()) return 0; + return dates->tf[at]; } if (booleans != nullptr) { const int8_t value = booleans->values[row]; diff --git a/src/cpplink/search.cpp b/src/cpplink/search.cpp new file mode 100644 index 0000000..ee4d561 --- /dev/null +++ b/src/cpplink/search.cpp @@ -0,0 +1,842 @@ +// Copyright 2026 Mathieu Fourment +// SPDX-License-Identifier: MIT + +#include "cpplink/search.hpp" + +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include + +#include + +#include "cpplink/arrow_c.hpp" +#include "cpplink/batch_loader.hpp" +#include "cpplink/derive.hpp" +#include "cpplink/format.hpp" +#include "cpplink/pair_stream.hpp" +#include "cpplink/parquet_io.hpp" + +namespace cpplink { +namespace { + +double Seconds(std::chrono::steady_clock::time_point from, + std::chrono::steady_clock::time_point to) { + return std::chrono::duration(to - from).count(); +} + +// Howard Hinnant's days_from_civil, the inverse of the conversion `derive` uses +// to take a stored date apart. A query types a date; the store holds days. +int32_t DaysFromCivil(int64_t year, int64_t month, int64_t day) { + year -= month <= 2; + const int64_t era = (year >= 0 ? year : year - 399) / 400; + const int64_t yoe = year - era * 400; + const int64_t doy = (153 * (month + (month > 2 ? -3 : 9)) + 2) / 5 + day - 1; + const int64_t doe = yoe * 365 + yoe / 4 - yoe / 100 + doy; + return static_cast(era * 146097 + doe - 719468); +} + +// A date as the file would have held it. Only the ISO form is accepted, because +// a query that means one thing to the store and another to the reader is worse +// than a query that is refused. +bool ParseDate(const std::string& text, int32_t* days) { + if (text.size() != 10 || text[4] != '-' || text[7] != '-') return false; + for (size_t i = 0; i < text.size(); ++i) { + if (i == 4 || i == 7) continue; + if (text[i] < '0' || text[i] > '9') return false; + } + const int64_t year = std::stol(text.substr(0, 4)); + const int64_t month = std::stol(text.substr(5, 2)); + const int64_t day = std::stol(text.substr(8, 2)); + if (month < 1 || month > 12 || day < 1 || day > 31) return false; + *days = DaysFromCivil(year, month, day); + return true; +} + +bool ParseBoolean(const std::string& text, int8_t* value) { + std::string lower; + lower.reserve(text.size()); + for (const char ch : text) lower.push_back(static_cast(std::tolower(ch))); + if (lower == "1" || lower == "true" || lower == "t" || lower == "yes" || + lower == "y") { + *value = 1; + return true; + } + if (lower == "0" || lower == "false" || lower == "f" || lower == "no" || + lower == "n") { + *value = 0; + return true; + } + return false; +} + +// The id a dictionary already holds this text under, by a walk over its values. +// +// The load-time index is released when the store is finalized and rebuilding it +// would cost more memory than the dictionary itself, so the lookup is a scan: +// a length test and a memcmp per value, which on the columns a fuzzy level reads +// is a fraction of the walk that follows it anyway. +bool FindValue(const Dictionary& dict, uint32_t limit, std::string_view text, + uint32_t* id) { + for (uint32_t v = 0; v < limit; ++v) { + if (dict.Value(v) == text) { + *id = v; + return true; + } + } + return false; +} + +// A hit is better than another when it scores higher, and where two tie, when it +// sits at the lower row. Without the second half the answer would depend on which +// thread saw which row first, and "the same k rows in the same order as scoring +// every pair" would not be a property anything could assert. +bool Better(const SearchHit& left, const SearchHit& right) { + if (left.weight != right.weight) return left.weight > right.weight; + return left.row < right.row; +} + +void Offer(std::vector* best, size_t k, const SearchHit& hit) { + // A max-heap under "better" keeps the worst kept hit at the front, which is + // the one a new hit has to beat and the one the bound is tested against. + if (best->size() < k) { + best->push_back(hit); + std::push_heap(best->begin(), best->end(), Better); + return; + } + if (!Better(hit, best->front())) return; + std::pop_heap(best->begin(), best->end(), Better); + best->back() = hit; + std::push_heap(best->begin(), best->end(), Better); +} + +// Whether every level of a comparison is a predicate over two values of one +// string column -- not necessarily the *same* column for every level, which is +// the address-against-its-username shape. Such a comparison can be tabulated: +// each level's verdict against the query is one byte per value id, and a row +// then costs a load and a compare instead of a string metric. +// +// A comparison reading a date, a number, a list or a coordinate pair is not of +// this shape and keeps the pair path's own evaluation against the query row. +bool TabulatableLevels(const BoundComparison& bound) { + if (bound.slots.empty()) return false; + if (bound.dates != nullptr || bound.booleans != nullptr || bound.lists != nullptr || + bound.numbers != nullptr) { + return false; + } + for (const LevelSpec& level : bound.spec->levels) { + switch (level.type) { + case LevelType::kNull: + case LevelType::kElse: + case LevelType::kExact: + case LevelType::kLevenshtein: + case LevelType::kJaroWinkler: + break; + default: + return false; + } + } + return true; +} + +// And whether it reads one column throughout, in which case the whole +// comparison collapses to a single table and a row costs one lookup rather than +// one per level. +bool TabulatableAsOne(const BoundComparison& bound) { + return bound.slots.size() == 1 && TabulatableLevels(bound); +} + +// The level a comparison of this shape lands on where no level fires: its else, +// which every comparison is required to end in. +uint8_t OtherwiseLevel(const BoundComparison& bound) { + return static_cast(bound.spec->levels.size() - 1); +} + +// A dictionary walk split over threads. It is trivially parallel -- every value +// is an independent question about the query -- and on a near-unique column it is +// the whole of a query's latency, so it is the first thing worth splitting. +template +void Walk(uint32_t values, unsigned threads, Body&& body) { + if (threads <= 1 || values <= 4096) { + body(0u, values); + return; + } + std::vector workers; + workers.reserve(threads); + const uint32_t span = (values + threads - 1) / threads; + for (unsigned t = 0; t < threads; ++t) { + const uint32_t from = std::min(values, span * t); + const uint32_t to = std::min(values, from + span); + if (from >= to) break; + workers.emplace_back([&body, from, to] { body(from, to); }); + } + for (std::thread& worker : workers) worker.join(); +} + +// The level index of the comparison's null level, or -1 where it declares none. +int NullLevel(const BoundComparison& bound) { + for (size_t i = 0; i < bound.spec->levels.size(); ++i) { + if (bound.spec->levels[i].type == LevelType::kNull) return static_cast(i); + } + return -1; +} + +} // namespace + +void QueryRecord::Set(const std::string& column, const std::string& value) { + fields.push_back({column, value}); +} + +const std::string* QueryRecord::Find(const std::string& column) const { + for (const QueryField& field : fields) { + if (field.column == column) return &field.value; + } + return nullptr; +} + +bool ParseQueryField(const std::string& text, QueryField* field, std::string* error) { + const size_t split = text.find('='); + if (split == std::string::npos || split == 0) { + *error = "'" + text + "' is not ="; + return false; + } + field->column = text.substr(0, split); + field->value = text.substr(split + 1); + return true; +} + +double PriorWeightForExpected(double expected, uint64_t records) { + if (records == 0) return 0.0; + const double n = static_cast(records); + // Clamped away from both ends: zero expected matches is minus infinity, and + // expecting every record to be the query's subject is plus infinity. + const double p = std::min(std::max(expected / n, 1e-300), 1.0 - 1e-12); + return std::log2(p / (1.0 - p)); +} + +Searcher::~Searcher() { Uninstall(); } + +bool Searcher::Bind(RecordStore* store, ComparisonSet* comparisons, const Scorer* scorer, + std::string* error) { + if (store->NumDatasets() > 1) { + *error = + "search reads one input: a query is a record the store does not hold, " + "so it is appended as a row, and over more than one input that row " + "would have to be a dataset of its own -- which the levels reading " + "which input a row came from were bound before it existed"; + return false; + } + store_ = store; + comparisons_ = comparisons; + scorer_ = scorer; + query_row_ = store->NumRecords(); + return true; +} + +void Searcher::Uninstall() { + if (!installed_ && written_ == 0 && !wrote_id_) return; + // Only as far as the install got: a query refused half way through has + // appended to the columns before the offending one and to none after it, and + // popping a column that was never written to would take a record's value off + // the store. + const size_t columns = written_; + for (size_t i = 0; i < columns; ++i) { + Column& column = store_->mutable_column(i); + if (auto* typed = std::get_if(&column)) { + typed->ids.pop_back(); + typed->dict.Truncate(dictionary_sizes_[i]); + if (!typed->tf.empty()) typed->tf.resize(typed->dict.Size()); + } else if (auto* typed = std::get_if(&column)) { + typed->ids.resize(typed->offsets[query_row_]); + typed->offsets.pop_back(); + typed->dict.Truncate(dictionary_sizes_[i]); + if (!typed->tf.empty()) typed->tf.resize(typed->dict.Size()); + } else if (auto* typed = std::get_if(&column)) { + typed->values.pop_back(); + } else if (auto* typed = std::get_if(&column)) { + typed->values.pop_back(); + } else if (auto* typed = std::get_if(&column)) { + typed->values.pop_back(); + } + } + if (wrote_id_) { + IdColumn& ids = store_->mutable_ids(); + ids.offsets.pop_back(); + ids.text.resize(ids.offsets.back()); + } + comparisons_->ResizeTables(); + plans_.clear(); + written_ = 0; + wrote_id_ = false; + installed_ = false; +} + +bool Searcher::Install(const QueryRecord& query, std::string* error) { + const Schema& schema = store_->schema(); + const size_t columns = store_->NumColumns(); + dictionary_sizes_.assign(columns, 0); + written_ = 0; + wrote_id_ = false; + + // Every column's text, so a derived column can be computed from the source + // the query gave rather than having to be given itself. A derived column the + // query names explicitly keeps what it was given. + std::vector given(columns, nullptr); + std::unordered_map by_name; + for (size_t i = 0; i < columns; ++i) { + by_name.emplace(schema.columns[i].name, i); + given[i] = query.Find(schema.columns[i].name); + } + std::vector derived(columns); + for (size_t i = 0; i < columns; ++i) { + const ColumnSpec& spec = schema.columns[i]; + if (!spec.IsDerived() || given[i] != nullptr) continue; + const auto source = by_name.find(spec.derive.from); + if (source == by_name.end()) continue; + const std::string* text = given[source->second]; + if (text == nullptr || text->empty()) continue; + if (schema.columns[source->second].type == ColumnType::kDate) { + int32_t days = 0; + if (!ParseDate(*text, &days)) continue; + ApplyDateTransforms(spec.derive.transforms, days, &derived[i]); + } else { + ApplyTransforms(spec.derive.transforms, *text, &derived[i]); + } + if (!derived[i].empty()) given[i] = &derived[i]; + } + + size_t adopted = 0; + for (size_t i = 0; i < columns; ++i) { + const ColumnSpec& spec = schema.columns[i]; + Column& column = store_->mutable_column(i); + const std::string* text = given[i]; + const bool missing = text == nullptr || text->empty(); + + if (auto* typed = std::get_if(&column)) { + dictionary_sizes_[i] = typed->dict.Size(); + uint32_t id = kNullId; + if (!missing) { + if (!FindValue(typed->dict, dictionary_sizes_[i], *text, &id)) { + // A value the store never saw still has to be a value id: a + // fuzzy level reads its text and its signature, and both are + // addressed by id. It agrees with nothing exactly, which is + // right, because nothing here holds it. + id = typed->dict.Adopt(*text); + if (!typed->tf.empty()) typed->tf.push_back(0); + ++adopted; + } + } + typed->ids.push_back(id); + } else if (auto* typed = std::get_if(&column)) { + dictionary_sizes_[i] = typed->dict.Size(); + // A list column takes one element per repeat of the field, which is + // unambiguous where a delimiter inside the values would not be. + std::vector cell; + for (const QueryField& field : query.fields) { + if (field.column != spec.name || field.value.empty()) continue; + uint32_t id = kNullId; + if (!FindValue(typed->dict, dictionary_sizes_[i], field.value, &id)) { + id = typed->dict.Adopt(field.value); + if (!typed->tf.empty()) typed->tf.push_back(0); + ++adopted; + } + cell.push_back(id); + } + // The cells are sorted and deduplicated at load, and every level over + // them is a linear merge that assumes it. + std::sort(cell.begin(), cell.end()); + cell.erase(std::unique(cell.begin(), cell.end()), cell.end()); + typed->ids.insert(typed->ids.end(), cell.begin(), cell.end()); + typed->offsets.push_back(typed->ids.size()); + } else if (auto* typed = std::get_if(&column)) { + int32_t days = kNullDate; + if (!missing && !ParseDate(*text, &days)) { + *error = "column \"" + spec.name + "\" is a date and '" + *text + + "' is not one: write it as YYYY-MM-DD"; + return false; + } + typed->values.push_back(days); + } else if (auto* typed = std::get_if(&column)) { + double value = std::nan(""); + if (!missing) { + try { + value = std::stod(*text); + } catch (const std::exception&) { + *error = "column \"" + spec.name + "\" is a number and '" + *text + + "' is not one"; + return false; + } + } + typed->values.push_back(value); + } else if (auto* typed = std::get_if(&column)) { + int8_t value = kNullBoolean; + if (!missing && !ParseBoolean(*text, &value)) { + *error = "column \"" + spec.name + "\" is a boolean and '" + *text + + "' is not one"; + return false; + } + typed->values.push_back(value); + } + written_ = i + 1; + } + store_->mutable_ids().Append("query"); + wrote_id_ = true; + query_row_ = store_->NumRecords(); + installed_ = true; + adopted_ = adopted; + // Signatures are indexed by value id, so any dictionary the query grew has to + // be followed before a level reads one. + comparisons_->ResizeTables(); + return true; +} + +void Searcher::BuildPlans(SearchReport* report) { + const size_t count = comparisons_->Size(); + plans_.assign(count, Plan()); + const unsigned threads = ResolveThreads(report->threads); + for (size_t c = 0; c < count; ++c) { + const BoundComparison& bound = comparisons_->at(c); + Plan& plan = plans_[c]; + plan.shift = bound.shift; + const int null_level = NullLevel(bound); + // A comparison the query says nothing about is the same level for every + // row, whatever its shape: the null level is the first the walk reaches + // and it fires on the query's side alone. That is most of a query, since + // a caller who knows ten fields about the person is not searching, and it + // is why a nine-comparison schema costs three or four of them. + if (null_level == 0 && comparisons_->IsNullValue(c, query_row_)) { + plan.kind = Plan::Kind::kConstant; + plan.constant = 0; + ++report->constant; + continue; + } + plan.has_null = null_level >= 0; + plan.null_level = null_level >= 0 ? static_cast(null_level) : 0; + if (!TabulatableLevels(bound)) { + plan.kind = Plan::Kind::kEvaluate; + ++report->evaluated; + continue; + } + if (!TabulatableAsOne(bound)) { + // The levels read different columns, so each gets its own table over + // its own dictionary and a row walks them in order, exactly as the + // pair path walks the levels themselves. + plan.kind = Plan::Kind::kLevels; + plan.otherwise = OtherwiseLevel(bound); + for (size_t i = 0; i < bound.spec->levels.size(); ++i) { + const LevelType type = bound.spec->levels[i].type; + if (type == LevelType::kNull || type == LevelType::kElse) continue; + Plan::LevelTable table; + table.level = static_cast(i); + const BoundComparison::StringSlot& slot = + bound.slots[bound.spec->levels[i].column]; + table.strings = slot.strings; + const uint32_t self = slot.strings->ids[query_row_]; + const uint32_t values = slot.strings->dict.Size(); + table.fires.assign(values, 0); + if (self != kNullId) { + const size_t level = i; + Walk(values, threads, [&](uint32_t from, uint32_t to) { + for (uint32_t v = from; v < to; ++v) { + table.fires[v] = + comparisons_->StringLevelForValues(c, level, self, v) ? 1 + : 0; + } + }); + report->values_walked += values; + } + plan.levels.push_back(std::move(table)); + } + ++report->tabled; + continue; + } + const StringColumn& strings = *bound.slots.front().strings; + const uint32_t self = strings.ids[query_row_]; + // A column the query leaves out whose comparison declares no null level + // at the top is still a constant, since no level a value reaches can fire + // against a value that is not there. + if (self == kNullId) { + plan.kind = Plan::Kind::kConstant; + plan.constant = null_level >= 0 + ? static_cast(null_level) + : comparisons_->LevelForValues(c, kNullId, kNullId); + ++report->constant; + continue; + } + plan.kind = Plan::Kind::kTable; + plan.strings = &strings; + if (null_level < 0) { + plan.null_level = comparisons_->LevelForValues(c, self, kNullId); + } + const uint32_t values = strings.dict.Size(); + plan.table.resize(values); + // This is where every string metric of the query runs: once per distinct + // value of the column rather than once per row, which is the whole of the + // economy and the reason the row pass below touches no character. + Walk(values, threads, [&](uint32_t from, uint32_t to) { + for (uint32_t v = from; v < to; ++v) { + plan.table[v] = comparisons_->LevelForValues(c, self, v); + } + }); + report->values_walked += values; + ++report->tabled; + } +} + +void Searcher::ScanRange(uint64_t begin, uint64_t end, const SearchOptions& options, + double shift, std::vector* best, + uint64_t* rescored) const { + uint32_t constant = 0; + for (const Plan& plan : plans_) { + if (plan.kind == Plan::Kind::kConstant) { + constant |= static_cast(plan.constant) << plan.shift; + } + } + for (uint64_t row = begin; row < end; ++row) { + uint32_t gamma = constant; + for (size_t c = 0; c < plans_.size(); ++c) { + const Plan& plan = plans_[c]; + uint8_t level = 0; + switch (plan.kind) { + case Plan::Kind::kConstant: + continue; + case Plan::Kind::kTable: { + const uint32_t value = plan.strings->ids[row]; + level = value == kNullId ? plan.null_level : plan.table[value]; + break; + } + case Plan::Kind::kLevels: { + // The comparison is null wherever *any* of its columns is, + // which is not a question one slot's table can answer. + if (plan.has_null && comparisons_->IsNullValue(c, row)) { + level = plan.null_level; + break; + } + level = plan.otherwise; + for (const Plan::LevelTable& table : plan.levels) { + const uint32_t value = table.strings->ids[row]; + if (value != kNullId && table.fires[value] != 0) { + level = table.level; + break; + } + } + break; + } + case Plan::Kind::kEvaluate: + level = comparisons_->EvaluateOne(c, row, query_row_); + break; + } + gamma |= static_cast(level) << plan.shift; + } + const double bound = + scorer_->BaseWeight(gamma) + scorer_->DeltaMax(gamma) + shift; + if (bound < options.threshold) continue; + if (best->size() == options.k && bound < best->front().weight) continue; + // The store's row is the first argument throughout: an exact level's + // adjustment reads the frequency of `a`'s value, and the query's own + // value was never counted into the table. + SearchHit hit; + hit.row = row; + hit.gamma = gamma; + hit.weight = scorer_->Weight(gamma, row, query_row_) + shift; + ++*rescored; + if (hit.weight < options.threshold) continue; + Offer(best, options.k, hit); + } +} + +bool Searcher::Search(const QueryRecord& query, const SearchOptions& options, + SearchReport* report, std::string* error) { + if (store_ == nullptr) { + *error = "the searcher is not bound to a store"; + return false; + } + if (options.k == 0) { + *error = "search needs a k of at least one"; + return false; + } + Uninstall(); + *report = SearchReport(); + report->threads = ResolveThreads(options.threads); + report->records = store_->NumRecords(); + report->model_prior = scorer_->PriorWeight(); + report->prior = options.override_prior ? options.prior_weight : report->model_prior; + + const auto started = std::chrono::steady_clock::now(); + if (!Install(query, error)) { + Uninstall(); + return false; + } + report->values_adopted = adopted_; + BuildPlans(report); + const auto walked = std::chrono::steady_clock::now(); + report->walk_seconds = Seconds(started, walked); + + const double shift = report->prior - report->model_prior; + const uint64_t rows = store_->NumRecords(); + std::vector> per_thread(report->threads); + std::vector rescored(report->threads, 0); + if (report->threads > 1 && rows > 0) { + std::vector workers; + workers.reserve(report->threads); + const uint64_t span = (rows + report->threads - 1) / report->threads; + for (unsigned t = 0; t < report->threads; ++t) { + const uint64_t from = std::min(rows, span * t); + const uint64_t to = std::min(rows, from + span); + if (from >= to) break; + workers.emplace_back([&, t, from, to] { + ScanRange(from, to, options, shift, &per_thread[t], &rescored[t]); + }); + } + for (std::thread& worker : workers) worker.join(); + } else { + ScanRange(0, rows, options, shift, &per_thread[0], &rescored[0]); + } + report->gather_seconds = Seconds(walked, std::chrono::steady_clock::now()); + + std::vector merged; + for (size_t t = 0; t < per_thread.size(); ++t) { + report->rescored += rescored[t]; + merged.insert(merged.end(), per_thread[t].begin(), per_thread[t].end()); + } + std::sort(merged.begin(), merged.end(), Better); + if (merged.size() > options.k) merged.resize(options.k); + for (SearchHit& hit : merged) { + hit.id = std::string(store_->ids().Get(hit.row)); + hit.dataset = store_->NumDatasets() > 1 + ? store_->DatasetName(store_->DatasetOf(hit.row)) + : std::string(); + hit.probability = ProbabilityForWeight(hit.weight); + } + report->hits = std::move(merged); + return true; +} + +void PrintSearchReport(const SearchReport& report, std::ostream& out) { + out << "Records " << WithThousands(report.records) << "\n" + << "Comparisons " << report.tabled << " tabulated over their dictionary, " + << report.evaluated << " evaluated per row, " << report.constant + << " constant because the query is missing the column\n" + << "Values walked " << WithThousands(report.values_walked); + if (report.values_adopted > 0) { + out << " (" << report.values_adopted << " query value" + << (report.values_adopted == 1 ? "" : "s") << " the store never held)"; + } + out << "\n" + << "Rows scored " << WithThousands(report.rescored) << " of " + << WithThousands(report.records) << " exactly; the rest were dropped on the " + << "bracket\n" + << "Prior " << std::fixed << std::setprecision(3) << report.prior + << " bits"; + if (report.prior != report.model_prior) { + out << " (the model's lambda says " << report.model_prior << ")"; + } + out << "\n" + << "Elapsed " << std::setprecision(4) << report.walk_seconds + << " s walking the dictionaries, " << report.gather_seconds + << " s over the rows on " << report.threads + << (report.threads == 1 ? " thread" : " threads") << "\n\n"; + + if (report.hits.empty()) { + out << "Nothing scores above the threshold: no record here is this one.\n"; + return; + } + out << std::left << std::setw(26) << "Record" << std::right << std::setw(12) + << "weight" << std::setw(14) << "posterior"; + if (report.clusters > 0) out << " cluster"; + out << "\n"; + for (const SearchHit& hit : report.hits) { + std::string name = hit.id; + if (!hit.dataset.empty()) name = hit.dataset + ":" + name; + out << std::left << std::setw(26) << Truncate(name, 25) << std::right + << std::setw(12) << std::fixed << std::setprecision(3) << hit.weight + << std::setw(14) << std::setprecision(6) << hit.probability; + if (report.clusters > 0) { + out << " " << hit.cluster; + if (!hit.cluster_best) out << " (same cluster)"; + } + out << "\n"; + } + out << "\nA weight is bits of evidence for the query and the record being the " + "same\nperson, and the posterior is what that means under the prior above. " + "Zero bits\nis even odds, so a hit below it is evidence against.\n"; +} + +std::string SearchReportJson(const SearchReport& report) { + nlohmann::json out; + out["records"] = report.records; + out["tabulated"] = report.tabled; + out["evaluated"] = report.evaluated; + out["constant"] = report.constant; + out["values_walked"] = report.values_walked; + out["values_adopted"] = report.values_adopted; + out["rescored"] = report.rescored; + out["prior"] = report.prior; + out["model_prior"] = report.model_prior; + out["walk_seconds"] = report.walk_seconds; + out["gather_seconds"] = report.gather_seconds; + out["threads"] = report.threads; + nlohmann::json hits = nlohmann::json::array(); + for (const SearchHit& hit : report.hits) { + nlohmann::json one; + one["row"] = hit.row; + one["id"] = hit.id; + if (!hit.dataset.empty()) one["dataset"] = hit.dataset; + one["gamma"] = hit.gamma; + one["match_weight"] = hit.weight; + one["match_probability"] = hit.probability; + if (!hit.cluster.empty()) { + one["cluster_id"] = hit.cluster; + one["cluster_best"] = hit.cluster_best; + } + hits.push_back(one); + } + out["hits"] = hits; + return out.dump(); +} + +namespace { + +// The cluster file as a map from record id to cluster id, read from either shape +// `cluster --out` writes. Only the ids the hits name are kept, so a 20M-row +// cluster file costs the hits and not the file. +bool ReadClusterCsv(const std::string& path, + std::unordered_map* clusters, + std::string* error) { + std::ifstream file(path); + if (!file) { + *error = "cannot open " + path; + return false; + } + std::string line; + if (!std::getline(file, line)) { + *error = path + " is empty"; + return false; + } + if (!line.empty() && line.back() == '\r') line.pop_back(); + std::vector header; + for (size_t at = 0; at <= line.size();) { + const size_t comma = std::min(line.find(',', at), line.size()); + header.push_back(line.substr(at, comma - at)); + at = comma + 1; + } + size_t id_at = header.size(); + size_t cluster_at = header.size(); + for (size_t i = 0; i < header.size(); ++i) { + if (header[i] == "unique_id") id_at = i; + if (header[i] == "cluster_id") cluster_at = i; + } + if (id_at == header.size() || cluster_at == header.size()) { + *error = path + " has no \"unique_id\" and \"cluster_id\" columns, so it is " + + "not a cpplink cluster file"; + return false; + } + while (std::getline(file, line)) { + if (!line.empty() && line.back() == '\r') line.pop_back(); + if (line.empty()) continue; + std::vector cells; + for (size_t at = 0; at <= line.size();) { + const size_t comma = std::min(line.find(',', at), line.size()); + cells.push_back(line.substr(at, comma - at)); + at = comma + 1; + } + if (cells.size() <= std::max(id_at, cluster_at)) continue; + clusters->emplace(cells[id_at], cells[cluster_at]); + } + return true; +} + +bool ReadClusterParquet(const std::string& path, + std::unordered_map* clusters, + std::string* error) { + ArrowArrayStream stream; + stream.release = nullptr; + if (!OpenParquetStream(path, {"unique_id", "cluster_id"}, &stream, error)) { + return false; + } + ArrowSchema schema; + schema.release = nullptr; + const auto release = [&] { + if (schema.release != nullptr) schema.release(&schema); + if (stream.release != nullptr) stream.release(&stream); + }; + if (stream.get_schema(&stream, &schema) != 0) { + release(); + *error = "cannot read the schema of " + path; + return false; + } + const int id_at = FieldIndex(schema, "unique_id"); + const int cluster_at = FieldIndex(schema, "cluster_id"); + if (id_at < 0 || cluster_at < 0) { + release(); + *error = path + " has no \"unique_id\" and \"cluster_id\" columns, so it is " + + "not a cpplink cluster file"; + return false; + } + while (true) { + ArrowArray batch; + batch.release = nullptr; + if (stream.get_next(&stream, &batch) != 0) { + release(); + *error = "cannot read the next batch of " + path; + return false; + } + if (batch.release == nullptr) break; + TextReader ids, names; + if (!ids.Bind(schema, batch, id_at, error) || + !names.Bind(schema, batch, cluster_at, error)) { + batch.release(&batch); + release(); + return false; + } + char left[24], right[24]; + for (int64_t row = 0; row < batch.length; ++row) { + std::string_view id; + std::string_view cluster; + if (!ids.At(row, &id, &left)) continue; + if (!names.At(row, &cluster, &right)) continue; + clusters->emplace(std::string(id), std::string(cluster)); + } + batch.release(&batch); + } + release(); + return true; +} + +} // namespace + +bool GroupHitsByCluster(const RecordStore& store, const std::string& cluster_path, + SearchReport* report, std::string* error) { + (void)store; + std::unordered_map clusters; + const bool parquet = + cluster_path.size() > 8 && + cluster_path.compare(cluster_path.size() - 8, 8, ".parquet") == 0; + if (parquet) { + if (!ReadClusterParquet(cluster_path, &clusters, error)) return false; + } else { + if (!ReadClusterCsv(cluster_path, &clusters, error)) return false; + } + // The hits are already in descending weight, so the first hit of a cluster is + // its best row and everything after it is another row of the same entity. + std::unordered_map seen; + for (SearchHit& hit : report->hits) { + const auto found = clusters.find(hit.id); + if (found == clusters.end()) continue; + hit.cluster = found->second; + hit.cluster_best = seen.emplace(hit.cluster, 1).second; + } + report->clusters = seen.size(); + return true; +} + +} // namespace cpplink diff --git a/src/cpplink/search.hpp b/src/cpplink/search.hpp new file mode 100644 index 0000000..091c67d --- /dev/null +++ b/src/cpplink/search.hpp @@ -0,0 +1,218 @@ +// Copyright 2026 Mathieu Fourment +// SPDX-License-Identifier: MIT + +#pragma once + +#include +#include +#include +#include +#include + +#include "cpplink/comparison.hpp" +#include "cpplink/explain.hpp" +#include "cpplink/record_store.hpp" +#include "cpplink/schema.hpp" +#include "cpplink/score.hpp" + +namespace cpplink { + +// Top-k retrieval over the match weight: one query record against a whole store. +// +// The weight is a sum of per-column terms, which is the shape a search engine's +// ranking function has, so this is top-k retrieval over an additive score rather +// than a linkage problem. What makes it cheap is the interning invariant: a +// string level is a function of two *value ids*, so the level every distinct +// value of a column lands on against the query is one table, computed once per +// query, and a row's pattern is then one gather per column and one lookup into +// the tabulated base weight. No string metric runs per row. +// +// The result is exact with respect to the model. Ranking is on +// `BaseWeight + DeltaMax`, which brackets every pair that pattern can produce, +// and the exact weight is computed only for the rows that bound admits -- the +// same admissibility `predict` rests on, used to skip rows instead of pairs. + +// One column of a query, as text: the same form the file the store was loaded +// from holds, because a query value has to be interned against that column's +// dictionary to mean anything. +struct QueryField { + std::string column; + std::string value; +}; + +// What a caller knows about the record being looked for. +// +// A column the query does not name is *missing*, not empty, which is the null +// level -- so a query carrying a name and a date of birth against a store of ten +// columns is scored as a record whose other eight fields were not recorded, +// which is what it is. +struct QueryRecord { + std::vector fields; + + void Set(const std::string& column, const std::string& value); + // The text given for a column, or null where the query does not name it. + const std::string* Find(const std::string& column) const; +}; + +// Parses `column=value`, which is how the command line takes one field. +bool ParseQueryField(const std::string& text, QueryField* field, std::string* error); + +// The prior odds in bits for "I expect this store to hold `expected` records of +// the person I am asking about". +// +// The model's own prior is lambda, the match rate over the *pair* space, which is +// the question deduplication asks. A query against N records asks a different one +// and the two differ by orders of magnitude, so the search prior is a per-query +// option rather than something inherited. It shifts every hit by the same +// constant, so it changes the posterior and the threshold and never the order. +double PriorWeightForExpected(double expected, uint64_t records); + +struct SearchOptions { + size_t k = 10; + // Rows scoring below this are not returned however few hits there are, which + // is what makes "nobody in here is this person" expressible. + double threshold = -std::numeric_limits::infinity(); + unsigned threads = 1; + // Replaces the model's prior with this many bits. Off by default, so a search + // scores exactly what `predict` would score for the same pair. + bool override_prior = false; + double prior_weight = 0.0; +}; + +struct SearchHit { + uint64_t row = 0; + std::string id; + std::string dataset; + uint32_t gamma = 0; + double weight = 0.0; + double probability = 0.0; + // The cluster the row belongs to, where the caller gave a cluster file, and + // whether this hit is the best row of it. + std::string cluster; + bool cluster_best = true; +}; + +struct SearchReport { + std::vector hits; + uint64_t records = 0; + // Comparisons answered from a level table, and those evaluated per row + // because their shape has no table: a comparison reading a date, a list, a + // coordinate pair or two string columns at once is not a function of one + // value id, so it keeps the pair path's own evaluation. + size_t tabled = 0; + size_t evaluated = 0; + // Comparisons over a column the query does not carry, which are the same + // level for every row and cost nothing at all. + size_t constant = 0; + // Distinct values phase one walked, summed over the tabled comparisons. + uint64_t values_walked = 0; + // Values the query carried that the store's dictionaries had never seen. + size_t values_adopted = 0; + // Rows whose exact term-frequency-adjusted weight had to be computed. The + // rest were dropped on the bracket alone. + uint64_t rescored = 0; + double prior = 0.0; // the prior in force, in bits + double model_prior = 0.0; // what the model's own lambda says + double walk_seconds = 0.0; + double gather_seconds = 0.0; + unsigned threads = 1; + // Set where a cluster file was read: clusters the hits fall in. + size_t clusters = 0; +}; + +void PrintSearchReport(const SearchReport& report, std::ostream& out); +std::string SearchReportJson(const SearchReport& report); + +// A query held against a store for as long as it takes to answer it and explain +// the answer. +// +// The query is installed as one row past the store's loaded rows. That is the one +// place the store is written to after load, and it is deliberate: a query that is +// a row is scored by the same `Evaluate`, bounded by the same brackets and +// explained by the same `BuildPairWaterfall` as any pair, with no second +// implementation of any of it to drift from the first. The row is written before +// any worker starts and removed when the next query replaces it or the searcher +// is destroyed, so nothing concurrent ever observes the store change and the +// invariant that earns its keep -- no locking in the hot path -- is untouched. +// +// One searcher answers one query at a time. Two searchers over one store, or two +// queries at once, are not supported and would be a second writer. +class Searcher { + public: + Searcher() = default; + ~Searcher(); + Searcher(const Searcher&) = delete; + Searcher& operator=(const Searcher&) = delete; + + // The store and the comparison set are written to, so both are taken + // mutable; the scorer is not. Fails where the store holds more than one + // input: a query row would have to be a dataset of its own, and the levels + // that read which input a row came from were bound against a boundary list + // that does not know about it. + bool Bind(RecordStore* store, ComparisonSet* comparisons, const Scorer* scorer, + std::string* error); + + bool Search(const QueryRecord& query, const SearchOptions& options, + SearchReport* report, std::string* error); + + // The row the current query occupies, valid until the next `Search` or the + // searcher's destruction. This is what makes a hit explainable: the pair is + // `(hit.row, QueryRow())` and every report that takes a pair takes it. + uint64_t QueryRow() const { return query_row_; } + bool HasQuery() const { return installed_; } + + private: + // What one comparison does per row: read a level out of a table indexed by + // the column's value id, take a level no row can change, or evaluate the + // comparison against the query row the way the pair path does. + struct Plan { + enum class Kind : uint8_t { kConstant, kTable, kLevels, kEvaluate }; + // One string level of a kLevels comparison: whether it fires against the + // query, by value id of the column *that level* reads. + struct LevelTable { + uint8_t level = 0; + const StringColumn* strings = nullptr; + std::vector fires; + }; + Kind kind = Kind::kEvaluate; + uint8_t shift = 0; + uint8_t constant = 0; // kConstant: the level every row lands on + uint8_t null_level = 0; // the level a row missing the value takes + bool has_null = false; // whether the comparison declares one + uint8_t otherwise = 0; // kLevels: the level a row reaching none takes + const StringColumn* strings = nullptr; + std::vector table; // kTable: the level, by value id + std::vector levels; + }; + + void Uninstall(); + bool Install(const QueryRecord& query, std::string* error); + void BuildPlans(SearchReport* report); + void ScanRange(uint64_t begin, uint64_t end, const SearchOptions& options, + double shift, std::vector* best, uint64_t* rescored) const; + + RecordStore* store_ = nullptr; + ComparisonSet* comparisons_ = nullptr; + const Scorer* scorer_ = nullptr; + uint64_t query_row_ = 0; + bool installed_ = false; + size_t adopted_ = 0; + // How much of the query row was written before the install returned, so a + // refusal takes back exactly what it added and no more. + size_t written_ = 0; + bool wrote_id_ = false; + // What each string column's dictionary held before the query adopted + // anything, so the adoptions can be taken back. + std::vector dictionary_sizes_; + std::vector plans_; +}; + +// Reads a cluster file -- what `cluster` writes -- and labels each hit with the +// cluster its row belongs to, marking the best-scoring row of each. The hits are +// rows and a caller looking for a person wants entities, which is the store's +// clusters; grouping at the end rather than indexing one row per cluster is what +// keeps the recall of a cluster whose members disagree. +bool GroupHitsByCluster(const RecordStore& store, const std::string& cluster_path, + SearchReport* report, std::string* error); + +} // namespace cpplink diff --git a/src/cpplink/signature.cpp b/src/cpplink/signature.cpp index 6d8a1cd..f282c03 100644 --- a/src/cpplink/signature.cpp +++ b/src/cpplink/signature.cpp @@ -24,4 +24,16 @@ void SignatureTable::Build(const Dictionary& dict) { } } +void SignatureTable::Extend(const Dictionary& dict) { + const uint32_t size = dict.Size(); + const uint32_t have = static_cast(mask_.size()); + mask_.resize(size); + length_.resize(size); + for (uint32_t id = have; id < size; ++id) { + const std::string_view value = dict.Value(id); + mask_[id] = CharacterMask(value); + length_[id] = static_cast(value.size()); + } +} + } // namespace cpplink diff --git a/src/cpplink/signature.hpp b/src/cpplink/signature.hpp index 1129f90..6c7852f 100644 --- a/src/cpplink/signature.hpp +++ b/src/cpplink/signature.hpp @@ -28,6 +28,10 @@ uint64_t CharacterMask(std::string_view value); class SignatureTable { public: void Build(const Dictionary& dict); + // Brings the table to the dictionary's current size, computing only what is + // new. A table is indexed by value id, so a dictionary a query has grown + // would be read past the end of one that has not followed it. + void Extend(const Dictionary& dict); bool Empty() const { return mask_.empty(); } uint64_t Mask(uint32_t id) const { return mask_[id]; } diff --git a/tests/search_test.cpp b/tests/search_test.cpp new file mode 100644 index 0000000..ac4d86a --- /dev/null +++ b/tests/search_test.cpp @@ -0,0 +1,571 @@ +// Copyright 2026 Mathieu Fourment +// SPDX-License-Identifier: MIT + +#include "cpplink/search.hpp" + +#include +#include +#include +#include +#include +#include + +#include + +#include "cpplink/comparison.hpp" +#include "cpplink/explain.hpp" +#include "cpplink/model.hpp" +#include "cpplink/record_store.hpp" +#include "cpplink/schema.hpp" +#include "cpplink/score.hpp" + +namespace { + +constexpr uint64_t kRecords = 500; + +// Four comparisons of three different shapes: two the searcher can tabulate over +// a dictionary, one over a derived column it has to compute for the query, and a +// date, which is not a function of a value id and so keeps the pair path's own +// evaluation. A search that gets the same answer either way has to get both right. +const char* const kSchemaJson = R"({ + "unique_id": "id", + "columns": [ + {"name": "surname", "type": "string"}, + {"name": "city", "type": "string"}, + {"name": "dob", "type": "date"}, + {"name": "surname_key", "derive": {"from": "surname", "transform": "soundex"}} + ], + "comparisons": [ + {"name": "surname", "columns": ["surname"], "term_frequency": true, + "levels": [{"type": "null"}, {"type": "exact"}, + {"type": "jaro_winkler", "threshold": 0.85}, {"type": "else"}]}, + {"name": "city", "columns": ["city"], + "levels": [{"type": "exact"}, {"type": "else"}]}, + {"name": "dob", "columns": ["dob"], + "levels": [{"type": "null"}, {"type": "exact"}, + {"type": "date_within", "threshold": 400}, {"type": "else"}]}, + {"name": "surname_key", "columns": ["surname_key"], + "levels": [{"type": "null"}, {"type": "exact"}, {"type": "else"}]} + ], + "blocking": [{"type": "exact_value", "column": "city"}] +})"; + +cpplink::ModelComparison Comparison(const std::string& name, bool tf, + const std::vector>& rates) { + cpplink::ModelComparison comparison; + comparison.name = name; + comparison.term_frequency = tf; + for (const auto& entry : rates) { + cpplink::ModelLevel level; + level.m = entry.first; + level.u = entry.second; + level.m_estimated = true; + comparison.levels.push_back(level); + } + return comparison; +} + +// The same ordering the searcher uses: the best weight first, and where two tie, +// the lower row. +struct Scored { + uint64_t row = 0; + uint32_t gamma = 0; + double weight = 0.0; +}; + +bool BetterThan(const Scored& left, const Scored& right) { + if (left.weight != right.weight) return left.weight > right.weight; + return left.row < right.row; +} + +class SearchFixture : public ::testing::Test { + protected: + void SetUp() override { + std::string error; + ASSERT_TRUE(cpplink::ParseSchema(kSchemaJson, &schema_, &error)) << error; + store_ = std::make_unique(schema_); + + auto& surname = std::get(store_->mutable_column(0)); + auto& city = std::get(store_->mutable_column(1)); + auto& dob = std::get(store_->mutable_column(2)); + // A skewed surname distribution, so the term-frequency adjustment has a + // spread and the bracket is not degenerate; names one edit apart, so the + // fuzzy level fires on something. + std::vector names; + for (int i = 0; i < 30; ++i) { + names.push_back(surname.dict.Intern("sander" + std::to_string(i % 10) + + (i < 10 ? "" : "s"))); + } + std::vector cities; + for (int i = 0; i < 5; ++i) { + cities.push_back(city.dict.Intern("city" + std::to_string(i))); + } + for (uint64_t row = 0; row < kRecords; ++row) { + const size_t pick = static_cast((row * row) % 30) / 2; + // Every seventh record is missing its surname, so the null level and + // the derived column's empty result are both exercised. + surname.ids.push_back(row % 7 == 3 ? cpplink::kNullId : names[pick]); + city.ids.push_back(cities[(row * 11) % 5]); + dob.values.push_back(row % 11 == 5 + ? cpplink::kNullDate + : static_cast(4000 + (row * 37) % 900)); + store_->mutable_ids().Append("r" + std::to_string(row)); + } + store_->set_num_records(kRecords); + store_->Finalize(); + + ASSERT_TRUE(comparisons_.Bind(schema_, *store_, &error)) << error; + model_.lambda = 0.02; + model_.records = kRecords; + model_.comparisons = { + Comparison("surname", true, + {{1e-9, 1e-9}, {0.75, 0.03}, {0.15, 0.05}, {0.10, 0.92}}), + Comparison("city", false, {{0.7, 0.2}, {0.3, 0.8}}), + Comparison("dob", false, + {{1e-9, 1e-9}, {0.6, 0.002}, {0.25, 0.4}, {0.15, 0.598}}), + Comparison("surname_key", false, {{1e-9, 1e-9}, {0.8, 0.06}, {0.2, 0.94}})}; + cpplink::ScoreOptions score; + ASSERT_TRUE(scorer_.Bind(model_, comparisons_, *store_, score, &error)) << error; + } + + // Every row scored the expensive way against the query the searcher installed, + // which is what the top-k answer has to agree with. + std::vector ScoreEveryRow(uint64_t query_row) const { + std::vector all; + all.reserve(kRecords); + for (uint64_t row = 0; row < kRecords; ++row) { + Scored one; + one.row = row; + one.gamma = comparisons_.Evaluate(row, query_row); + one.weight = scorer_.Weight(one.gamma, row, query_row); + all.push_back(one); + } + std::sort(all.begin(), all.end(), BetterThan); + return all; + } + + cpplink::Schema schema_; + std::unique_ptr store_; + cpplink::ComparisonSet comparisons_; + cpplink::Model model_; + cpplink::Scorer scorer_; +}; + +cpplink::QueryRecord Query() { + cpplink::QueryRecord query; + query.Set("surname", "sander3s"); + query.Set("city", "city2"); + query.Set("dob", "1981-01-24"); + return query; +} + +// The claim the whole design rests on: the k rows the two-phase search returns are +// the k rows scoring every pair exactly would return, in the same order. The +// bound it prunes on is admissible or this fails. +TEST_F(SearchFixture, MatchesScoringEveryRow) { + cpplink::Searcher searcher; + std::string error; + ASSERT_TRUE(searcher.Bind(store_.get(), &comparisons_, &scorer_, &error)) << error; + + cpplink::SearchOptions options; + options.k = 12; + cpplink::SearchReport report; + ASSERT_TRUE(searcher.Search(Query(), options, &report, &error)) << error; + + const std::vector all = ScoreEveryRow(searcher.QueryRow()); + ASSERT_EQ(report.hits.size(), options.k); + for (size_t i = 0; i < report.hits.size(); ++i) { + EXPECT_EQ(report.hits[i].row, all[i].row) << "at rank " << i; + EXPECT_EQ(report.hits[i].gamma, all[i].gamma) << "at rank " << i; + EXPECT_DOUBLE_EQ(report.hits[i].weight, all[i].weight) << "at rank " << i; + EXPECT_EQ(report.hits[i].id, "r" + std::to_string(all[i].row)); + } + // Three comparisons are a function of one value id and are answered from a + // table; the date is not, and keeps the pair path's evaluation. + EXPECT_EQ(report.tabled, 3u); + EXPECT_EQ(report.evaluated, 1u); + EXPECT_EQ(report.constant, 0u); + // The bracket is what makes the exact weight optional, so most rows never + // reach it. + EXPECT_LT(report.rescored, kRecords); +} + +// And the level every row lands on is the level the pair path assigns it, not +// only for the winners: a table that disagreed anywhere would be a different +// model, whether or not it changed the top of the list. +TEST_F(SearchFixture, EveryRowsPatternIsThePairPaths) { + cpplink::Searcher searcher; + std::string error; + ASSERT_TRUE(searcher.Bind(store_.get(), &comparisons_, &scorer_, &error)) << error; + + cpplink::SearchOptions options; + options.k = kRecords; + cpplink::SearchReport report; + ASSERT_TRUE(searcher.Search(Query(), options, &report, &error)) << error; + ASSERT_EQ(report.hits.size(), kRecords); + for (const cpplink::SearchHit& hit : report.hits) { + EXPECT_EQ(hit.gamma, comparisons_.Evaluate(hit.row, searcher.QueryRow())) + << "row " << hit.row; + } + // Every row was scored exactly, so nothing was pruned and nothing was missed. + EXPECT_EQ(report.rescored, kRecords); +} + +// A record already in the store is its own best answer, which is the sanity check +// that the query's values reach the dictionary ids the store holds. +TEST_F(SearchFixture, AStoredRecordFindsItself) { + const auto& surname = std::get(store_->column(0)); + const auto& city = std::get(store_->column(1)); + uint64_t subject = 0; + while (surname.ids[subject] == cpplink::kNullId) ++subject; + + cpplink::QueryRecord query; + query.Set("surname", std::string(surname.dict.Value(surname.ids[subject]))); + query.Set("city", std::string(city.dict.Value(city.ids[subject]))); + + cpplink::Searcher searcher; + std::string error; + ASSERT_TRUE(searcher.Bind(store_.get(), &comparisons_, &scorer_, &error)) << error; + cpplink::SearchOptions options; + options.k = 5; + cpplink::SearchReport report; + ASSERT_TRUE(searcher.Search(query, options, &report, &error)) << error; + ASSERT_FALSE(report.hits.empty()); + // The surname agrees exactly and so does its key and the city, which is the + // best any record can do against this query. + EXPECT_EQ(comparisons_.Evaluate(subject, searcher.QueryRow()), + report.hits.front().gamma); + EXPECT_GT(report.hits.front().weight, 0.0); + EXPECT_EQ(report.values_adopted, 0u); +} + +// A value the store never held is still a value: the fuzzy level needs its text +// and its signature, both of which only a value id addresses. +TEST_F(SearchFixture, AValueTheStoreNeverHeldIsStillCompared) { + const auto& surname = std::get(store_->column(0)); + const uint32_t before = surname.dict.Size(); + + cpplink::QueryRecord query; + query.Set("surname", "sander3x"); // one edit from "sander3", which is here + query.Set("city", "city2"); + + std::string error; + { + cpplink::Searcher searcher; + ASSERT_TRUE(searcher.Bind(store_.get(), &comparisons_, &scorer_, &error)) + << error; + cpplink::SearchOptions options; + options.k = 5; + cpplink::SearchReport report; + ASSERT_TRUE(searcher.Search(query, options, &report, &error)) << error; + // The surname is new; its soundex key is not, because a key is what two + // spellings of a name have in common. + EXPECT_EQ(report.values_adopted, 1u); + ASSERT_FALSE(report.hits.empty()); + // Nothing agrees exactly, but the fuzzy level does, which is the whole + // reason a query value has to become a value id at all. + const uint8_t level = comparisons_.LevelOf(report.hits.front().gamma, 0); + EXPECT_EQ(level, 2u); + } + // And the store is what it was: the adoptions are taken back with the row. + EXPECT_EQ(surname.dict.Size(), before); + EXPECT_EQ(surname.ids.size(), kRecords); + EXPECT_EQ(surname.tf.size(), before); + EXPECT_EQ(store_->NumRecords(), kRecords); +} + +// The query row is removed when the searcher goes, and the pairs the store +// already held score what they scored before it arrived. +TEST_F(SearchFixture, TheStoreIsLeftAsItWasFound) { + const uint32_t gamma_before = comparisons_.Evaluate(3, 17); + const double weight_before = scorer_.Weight(gamma_before, 3, 17); + { + cpplink::Searcher searcher; + std::string error; + ASSERT_TRUE(searcher.Bind(store_.get(), &comparisons_, &scorer_, &error)) + << error; + cpplink::SearchOptions options; + cpplink::SearchReport report; + ASSERT_TRUE(searcher.Search(Query(), options, &report, &error)) << error; + } + EXPECT_EQ(comparisons_.Evaluate(3, 17), gamma_before); + EXPECT_DOUBLE_EQ(scorer_.Weight(gamma_before, 3, 17), weight_before); + const auto& ids = store_->ids(); + EXPECT_EQ(ids.offsets.size(), kRecords + 1); + EXPECT_EQ(ids.Get(kRecords - 1), "r" + std::to_string(kRecords - 1)); +} + +// Splitting the rows over threads is a partition of the same scan, so it is the +// same answer or it is a bug. +TEST_F(SearchFixture, ThreadsChangeNothing) { + std::string error; + cpplink::SearchReport one; + cpplink::SearchReport many; + { + cpplink::Searcher searcher; + ASSERT_TRUE(searcher.Bind(store_.get(), &comparisons_, &scorer_, &error)) + << error; + cpplink::SearchOptions options; + options.k = 9; + options.threads = 1; + ASSERT_TRUE(searcher.Search(Query(), options, &one, &error)) << error; + options.threads = 4; + ASSERT_TRUE(searcher.Search(Query(), options, &many, &error)) << error; + } + ASSERT_EQ(one.hits.size(), many.hits.size()); + for (size_t i = 0; i < one.hits.size(); ++i) { + EXPECT_EQ(one.hits[i].row, many.hits[i].row) << "at rank " << i; + EXPECT_DOUBLE_EQ(one.hits[i].weight, many.hits[i].weight) << "at rank " << i; + } +} + +// The prior is the one thing search does not inherit from the model, and it moves +// every hit by the same constant: the posterior changes, the ranking does not. +TEST_F(SearchFixture, ThePriorShiftsEveryHitAndReordersNothing) { + cpplink::Searcher searcher; + std::string error; + ASSERT_TRUE(searcher.Bind(store_.get(), &comparisons_, &scorer_, &error)) << error; + + cpplink::SearchOptions options; + options.k = 8; + cpplink::SearchReport plain; + ASSERT_TRUE(searcher.Search(Query(), options, &plain, &error)) << error; + + options.override_prior = true; + options.prior_weight = cpplink::PriorWeightForExpected(1.0, kRecords); + cpplink::SearchReport shifted; + ASSERT_TRUE(searcher.Search(Query(), options, &shifted, &error)) << error; + + ASSERT_EQ(plain.hits.size(), shifted.hits.size()); + const double move = shifted.prior - plain.prior; + EXPECT_NE(move, 0.0); + for (size_t i = 0; i < plain.hits.size(); ++i) { + EXPECT_EQ(plain.hits[i].row, shifted.hits[i].row) << "at rank " << i; + EXPECT_NEAR(shifted.hits[i].weight, plain.hits[i].weight + move, 1e-9); + } + EXPECT_DOUBLE_EQ(plain.prior, scorer_.PriorWeight()); +} + +// A hit explains itself through the report every other pair goes through, which +// is what keeps the explanation and the answer one calculation. +TEST_F(SearchFixture, AHitIsExplainedByTheWaterfall) { + cpplink::Searcher searcher; + std::string error; + ASSERT_TRUE(searcher.Bind(store_.get(), &comparisons_, &scorer_, &error)) << error; + cpplink::SearchOptions options; + options.k = 4; + cpplink::SearchReport report; + ASSERT_TRUE(searcher.Search(Query(), options, &report, &error)) << error; + ASSERT_FALSE(report.hits.empty()); + for (const cpplink::SearchHit& hit : report.hits) { + const cpplink::PairWaterfall waterfall = cpplink::BuildPairWaterfall( + *store_, comparisons_, scorer_, hit.row, searcher.QueryRow(), &model_); + EXPECT_EQ(waterfall.gamma, hit.gamma); + EXPECT_DOUBLE_EQ(waterfall.weight, hit.weight); + EXPECT_EQ(waterfall.steps.size(), comparisons_.Size()); + } +} + +// A column the query does not name is missing, not empty, so every row lands on +// that comparison's null level whatever it holds. +TEST_F(SearchFixture, AnUnnamedColumnIsTheNullLevel) { + cpplink::QueryRecord query; + query.Set("city", "city1"); + + cpplink::Searcher searcher; + std::string error; + ASSERT_TRUE(searcher.Bind(store_.get(), &comparisons_, &scorer_, &error)) << error; + cpplink::SearchOptions options; + options.k = 6; + cpplink::SearchReport report; + ASSERT_TRUE(searcher.Search(query, options, &report, &error)) << error; + ASSERT_FALSE(report.hits.empty()); + for (const cpplink::SearchHit& hit : report.hits) { + EXPECT_EQ(comparisons_.LevelOf(hit.gamma, 0), 0u); // surname: null + EXPECT_EQ(comparisons_.LevelOf(hit.gamma, 2), 0u); // dob: null + EXPECT_EQ(comparisons_.LevelOf(hit.gamma, 3), 0u); // the derived key: null + } + // Three of the four comparisons cost nothing at all: the query does not carry + // their columns, so no row can move them off the null level. + EXPECT_EQ(report.constant, 3u); + EXPECT_EQ(report.tabled, 1u); + EXPECT_EQ(report.evaluated, 0u); +} + +// A derived column is filled in for the query from the source the query gave, +// exactly as the loader fills it in for a row. +TEST_F(SearchFixture, ADerivedColumnIsComputedForTheQuery) { + cpplink::QueryRecord query; + query.Set("surname", "saunders0"); // not in the store; the same soundex as + // "sander0s", which is + cpplink::Searcher searcher; + std::string error; + ASSERT_TRUE(searcher.Bind(store_.get(), &comparisons_, &scorer_, &error)) << error; + cpplink::SearchOptions options; + options.k = 5; + cpplink::SearchReport report; + ASSERT_TRUE(searcher.Search(query, options, &report, &error)) << error; + + const auto& key = std::get(store_->column(3)); + ASSERT_NE(key.ids[searcher.QueryRow()], cpplink::kNullId); + EXPECT_FALSE(key.dict.Value(key.ids[searcher.QueryRow()]).empty()); +} + +// The threshold is what makes "nobody here is this person" expressible: an empty +// answer rather than k bad ones. +TEST_F(SearchFixture, AThresholdCanRefuseEveryRow) { + cpplink::Searcher searcher; + std::string error; + ASSERT_TRUE(searcher.Bind(store_.get(), &comparisons_, &scorer_, &error)) << error; + cpplink::SearchOptions options; + options.k = 10; + options.threshold = 1e6; + cpplink::SearchReport report; + ASSERT_TRUE(searcher.Search(Query(), options, &report, &error)) << error; + EXPECT_TRUE(report.hits.empty()); + std::ostringstream text; + cpplink::PrintSearchReport(report, text); + EXPECT_NE(text.str().find("no record here is this one"), std::string::npos); +} + +TEST_F(SearchFixture, ADateThatIsNotOneIsRefused) { + cpplink::QueryRecord query; + query.Set("dob", "24/01/1981"); + cpplink::Searcher searcher; + std::string error; + ASSERT_TRUE(searcher.Bind(store_.get(), &comparisons_, &scorer_, &error)) << error; + cpplink::SearchOptions options; + cpplink::SearchReport report; + EXPECT_FALSE(searcher.Search(query, options, &report, &error)); + EXPECT_NE(error.find("YYYY-MM-DD"), std::string::npos); + // And the refusal leaves nothing behind, so the next query starts clean. + // Both sides of the column that refused: the ones already written are taken + // back, and the ones after it were never written and must not be popped. + EXPECT_EQ(store_->NumRecords(), kRecords); + EXPECT_EQ(std::get(store_->column(0)).ids.size(), kRecords); + EXPECT_EQ(std::get(store_->column(3)).ids.size(), kRecords); + EXPECT_EQ(store_->ids().offsets.size(), kRecords + 1); + EXPECT_EQ(store_->ids().Get(kRecords - 1), "r" + std::to_string(kRecords - 1)); + // And the next query works, which it would not if a column had lost a row. + cpplink::SearchOptions good; + good.k = 3; + ASSERT_TRUE(searcher.Search(Query(), good, &report, &error)) << error; + EXPECT_EQ(report.hits.size(), 3u); +} + +// A comparison whose levels read *different* columns -- splink's email, an +// address and the username derived from it -- is not a function of one value id, +// so it is tabulated a level at a time instead. The levels still have to be +// walked in the schema's order, and null still has to pre-empt them. +const char* const kEmailSchemaJson = R"({ + "unique_id": "id", + "columns": [ + {"name": "email", "type": "string"}, + {"name": "email_username", + "derive": {"from": "email", "transform": "email_username"}} + ], + "comparisons": [ + {"name": "email", "columns": ["email", "email_username"], + "term_frequency": true, + "levels": [{"type": "null"}, {"type": "exact"}, + {"type": "exact", "column": "email_username"}, + {"type": "jaro_winkler", "threshold": 0.93}, + {"type": "jaro_winkler", "threshold": 0.93, + "column": "email_username"}, + {"type": "else"}]} + ], + "blocking": [{"type": "exact_value", "column": "email"}] +})"; + +TEST(SearchEmailTest, AComparisonReadingTwoColumnsIsTabulatedALevelAtATime) { + cpplink::Schema schema; + std::string error; + ASSERT_TRUE(cpplink::ParseSchema(kEmailSchemaJson, &schema, &error)) << error; + cpplink::RecordStore store(schema); + auto& email = std::get(store.mutable_column(0)); + constexpr uint64_t kRows = 300; + for (uint64_t row = 0; row < kRows; ++row) { + // Four shapes: the username with another domain, a near-miss username, + // an unrelated address, and no address at all. + std::string value; + switch (row % 4) { + case 0: + value = "user" + std::to_string(row / 4) + "@example.com"; + break; + case 1: + value = "user" + std::to_string(row / 4) + "@other.org"; + break; + case 2: + value = "user" + std::to_string(row / 4) + "x@example.com"; + break; + default: + break; + } + email.ids.push_back(value.empty() ? cpplink::kNullId : email.dict.Intern(value)); + store.mutable_ids().Append("r" + std::to_string(row)); + } + store.set_num_records(kRows); + store.Finalize(); + + cpplink::ComparisonSet comparisons; + ASSERT_TRUE(comparisons.Bind(schema, store, &error)) << error; + cpplink::Model model; + model.lambda = 0.01; + model.records = kRows; + model.comparisons = {Comparison( + "email", true, + {{1e-9, 1e-9}, {0.5, 1e-5}, {0.2, 1e-4}, {0.1, 1e-3}, {0.1, 1e-2}, {0.1, 0.98}})}; + cpplink::Scorer scorer; + cpplink::ScoreOptions score; + ASSERT_TRUE(scorer.Bind(model, comparisons, store, score, &error)) << error; + + cpplink::QueryRecord query; + query.Set("email", "user7@example.com"); + + cpplink::Searcher searcher; + ASSERT_TRUE(searcher.Bind(&store, &comparisons, &scorer, &error)) << error; + cpplink::SearchOptions options; + options.k = kRows; + cpplink::SearchReport report; + ASSERT_TRUE(searcher.Search(query, options, &report, &error)) << error; + EXPECT_EQ(report.tabled, 1u); + EXPECT_EQ(report.evaluated, 0u); + ASSERT_EQ(report.hits.size(), kRows); + // The level-at-a-time table is the pair path or it is a different model. + for (const cpplink::SearchHit& hit : report.hits) { + EXPECT_EQ(hit.gamma, comparisons.Evaluate(hit.row, searcher.QueryRow())) + << "row " << hit.row; + } + // And it reaches every one of the shapes: the address exactly (row 28), the + // username where the domain differs (row 29), and null (row 31). + const auto level = [&](uint64_t row) { + for (const cpplink::SearchHit& hit : report.hits) { + if (hit.row == row) return comparisons.LevelOf(hit.gamma, 0); + } + return static_cast(255); + }; + EXPECT_EQ(level(28), 1u); // user7@example.com + EXPECT_EQ(level(29), 2u); // user7@other.org: the username alone + EXPECT_EQ(level(31), 0u); // no address at all +} + +TEST(SearchQueryTest, AFieldIsColumnEqualsValue) { + cpplink::QueryField field; + std::string error; + ASSERT_TRUE(cpplink::ParseQueryField("first_name=john smith", &field, &error)); + EXPECT_EQ(field.column, "first_name"); + EXPECT_EQ(field.value, "john smith"); + EXPECT_FALSE(cpplink::ParseQueryField("first_name", &field, &error)); + EXPECT_FALSE(cpplink::ParseQueryField("=john", &field, &error)); +} + +// The prior a query faces is not the pair space's: one expected match in 500 +// records is a very different number from lambda, and it is meant to be. +TEST(SearchQueryTest, ThePriorIsOddsOverRecordsNotOverPairs) { + EXPECT_NEAR(cpplink::PriorWeightForExpected(1.0, 1000), std::log2(1.0 / 999.0), + 1e-12); + EXPECT_GT(cpplink::PriorWeightForExpected(10.0, 1000), + cpplink::PriorWeightForExpected(1.0, 1000)); +} + +} // namespace From b3a6521dd492558a0a487917686952c648ae8af2 Mon Sep 17 00:00:00 2001 From: 4ment Date: Wed, 23 Sep 2026 20:55:51 +1000 Subject: [PATCH 2/5] Add search functionality to the python bindings --- docs/commands/cluster.md | 7 + docs/commands/explain.md | 1 + docs/commands/predict.md | 8 ++ docs/commands/search.md | 21 ++- docs/getting-started.md | 53 +++++++- docs/index.md | 12 ++ docs/model.md | 10 ++ docs/python.md | 41 +++++- python/bindings/common.hpp | 21 +++ python/bindings/module.cpp | 1 + python/bindings/search.cpp | 195 +++++++++++++++++++++++++++ python/bindings/store.cpp | 5 + python/cpplink/__init__.py | 69 ++++++++++ python/tests/test_reports.py | 14 ++ python/tests/test_search.py | 249 +++++++++++++++++++++++++++++++++++ src/cpplink/record_store.cpp | 22 ++++ src/cpplink/record_store.hpp | 21 +++ tests/search_test.cpp | 38 ++++++ 18 files changed, 785 insertions(+), 3 deletions(-) create mode 100644 python/bindings/search.cpp create mode 100644 python/tests/test_search.py diff --git a/docs/commands/cluster.md b/docs/commands/cluster.md index 3649ed3..0210a9f 100644 --- a/docs/commands/cluster.md +++ b/docs/commands/cluster.md @@ -235,3 +235,10 @@ a merged file alike (a merged file adds the id index, 4 bytes a record). The pas predictions and 1.8M records took 0.016 s against 36 s to produce them, which is why the lock-free CAS version the design mentions stays unbuilt. + +## See also + +- [`merge-predictions`](merge-predictions.md) — the shard directory as one file +- [`search`](search.md) — `--clusters` takes the file this writes and labels each hit with the + cluster it belongs to, so a query answers with entities rather than records +- [The cluster viewer](../viewer.md) — the partition as a page diff --git a/docs/commands/explain.md b/docs/commands/explain.md index cbe9d13..465ade6 100644 --- a/docs/commands/explain.md +++ b/docs/commands/explain.md @@ -222,6 +222,7 @@ Nothing in it recomputes a bit of the weight, and it does not run this binary. ## See also - [`predict`](predict.md) — the same scoring, over every candidate pair +- [`search`](search.md) — the same ledger, for each hit of a query record, under `--explain` - [The model](../model.md) — where `m`, `u`, `λ` and the weights come from ### The packed pattern diff --git a/docs/commands/predict.md b/docs/commands/predict.md index 402bf8d..f0dd58c 100644 --- a/docs/commands/predict.md +++ b/docs/commands/predict.md @@ -272,3 +272,11 @@ the bracket a fuzzy level admits is much wider than an exact one's — the check On the 1M synthetic sample the adjustment moves 47% of the weights, by a median of 0.6 bits and by more than 5 bits on 1,128 of them, and changes no decision at all: the posterior there saturates so hard that nothing near the threshold exists to move. + +## See also + +- [`rescore`](rescore.md) — replay a spill under a new model without comparing again +- [`cluster`](cluster.md) — join the predictions into duplicate clusters +- [`search`](search.md) — the same weight and the same bracket, over one query record instead + of every candidate pair +- [The model](../model.md) — where the weights, the bracket and the zones come from diff --git a/docs/commands/search.md b/docs/commands/search.md index 38ad44a..df3342a 100644 --- a/docs/commands/search.md +++ b/docs/commands/search.md @@ -168,7 +168,7 @@ The two phases scale with different things, which is the point of reporting them - The **dictionary walk** grows with the queried columns' *distinct values*, at a flat 46 ns each. On a near-unique column with a fuzzy level — an email address — that is the whole of the latency: at 20M the `full` query walks 91.6M values and spends 76% of itself there, - while the same query without the email address is answered in 598 ms. + while the `fuzzy` query, one name column, walks 657k and is answered in 598 ms. So the thing to watch is not the size of the store but whether the query names a near-unique column that a fuzzy level reads. Both phases split 4–5× over eight threads. @@ -195,8 +195,27 @@ Two more, both narrow: it compares fuzzily against everything as it should. It gets no entry in the `list_contains` alias map, so a nickname the store never saw is not looked up as one. +## From Python + +The resident session behind [`Linker`](../python.md) is what a search service is: one store +loaded once, a query answered in milliseconds. + +```python +hits = linker.search( + {"last_name": "zolnerowich", "dob": "1979-08-06", "postcode": "4508"}, + model, k=10, expected_matches=1, +) +for hit in hits: + print(hit.id, hit.match_weight, hit.match_probability) +``` + +A list or tuple value is spread over the repeats a list column takes one element per, `None` is +the same as omitting the column, and `explain=True` fills `SearchResult.waterfalls` with the +ledger behind each hit. See [From Python](../python.md#8-search-for-a-record). + ## See also - [`predict`](predict.md) — the same weight over candidate pairs rather than one query - [`explain`](explain.md) — the waterfall `--explain` prints - [`cluster`](cluster.md) — what `--clusters` reads +- [The model](../model.md) — where the weight, the prior and the bracket come from diff --git a/docs/getting-started.md b/docs/getting-started.md index dc85096..da43ade 100644 --- a/docs/getting-started.md +++ b/docs/getting-started.md @@ -45,7 +45,7 @@ ctest --preset asan On Windows the `windows` preset builds with Visual Studio 2026 against the conda-forge Arrow, which lives under `%CONDA_PREFIX%\Library`, and the build and test presets of the same name select the `Release` configuration, since the Visual Studio generator holds every configuration in one tree. The GitHub Actions workflow runs `release`, `asan` and `tsan` on Linux and macOS and `windows` on Windows, then the Python suite and the format and lint checks. -## The whole pipeline in seven commands +## The whole pipeline in eight commands Everything below runs against the synthetic sample cpplink can write for itself. The numbers shown are from a real run on 1.8M rows with eight threads. @@ -173,6 +173,57 @@ bounds still admit and drops the pair if even that sum cannot clear the threshol Edges carry their weight, so re-clustering at a higher `--threshold` is a re-read and no re-scoring. See [`cluster`](commands/cluster.md). +### 8. Ask the model about one record + +Everything above builds a model of what a matching pair looks like. That model is also a search +index: [`search`](commands/search.md) takes a query record and returns the records scoring +highest against it. + +```sh +./build/cpplink search --schema examples/sample_schema.json --model model.json \ + examples/sample.parquet -k 5 \ + --field last_name=chirdloackki --field dob=1979-08-06 \ + --field postcode=4508 --expected-matches 1 +``` + +```text +Records 1,800,000 +Comparisons 2 tabulated over their dictionary, 1 evaluated per row, 6 constant because the query is missing the column +Values walked 182,320 +Rows scored 436 of 1,800,000 exactly; the rest were dropped on the bracket +Prior -20.780 bits (the model's lambda says -23.475) +Elapsed 0.0080 s walking the dictionaries, 0.0762 s over the rows on 1 thread + +Record weight posterior +r0 27.785 1.000000 +r259138 -3.045 0.108036 +r1482739 -3.045 0.108036 +r632071 -7.288 0.006357 +r1048594 -7.288 0.006357 +``` + +That is one record of the sample asked for as a query, and it comes back at the top with 27.8 +bits while the runners-up sit below even odds. The three counts on the `Comparisons` line are +the three things a comparison can cost: **constant** is free, the query not carrying the column; +**tabulated** is one walk over that column's dictionary and then a byte read per record; +**evaluated** is the pair path's own evaluation per record, which a date or a coordinate pair +needs because it is not a function of a single value id. + +No blocking and no candidate pairs: the query is compared against every record, which is +affordable because interning makes a string level a question about a *value*. Every string +metric in the query runs once per **distinct value** of a column rather than once per record, +after which a record costs a load and a compare. A column the query does not name is *missing* +and costs nothing at all. + +The answer is exact with respect to the model — the same records, in the same order, as scoring +the query against every record one at a time — because the ranking is on the same admissible +bracket `predict` prunes with. + +`--expected-matches 1` is the prior stated as the question a *search* asks. λ is the match rate +over the pair space, which is what deduplication faces; a query against N records asks something +else, and the two differ by orders of magnitude. `--explain` prints the same waterfall +[`explain`](commands/explain.md) does, for each hit. + ## Where the time goes On this run, at 1.8M rows: diff --git a/docs/index.md b/docs/index.md index 1828e12..bb0ba91 100644 --- a/docs/index.md +++ b/docs/index.md @@ -60,6 +60,9 @@ machine, and what it ran out of was scratch space rather than memory. 5 Score TF-adjusted, bound-pruned cpplink predict cpplink rescore 6 Cluster union–find over the predictions cpplink cluster + + Search one query record against the store, cpplink search + top-k over the same weight ``` Every data command takes one or more parquet files. One file deduplicates, two link, and @@ -85,6 +88,7 @@ Four invariants hold the design together: - [Blocking](blocking.md) — the four pair sources, exact costing, and measured recall. - [Linking two files](linking.md): what changes when the pair space is a cross-product. - [Commands](commands/index.md) — what each subcommand is for and how to read its output. +- [`search`](commands/search.md) — a model is also a search index: one query record, top-k. - [The schema file](reference/schema.md): every field of the JSON that configures a run. ## Status @@ -101,6 +105,14 @@ So is everything built on top of them: the pair-global ceiling, [`estimate --interactions`](commands/estimate.md#relaxing-conditional-independence), which relaxes conditional independence inside the scoring model. +**A model is also a search index.** [`search`](commands/search.md) takes a query record and +returns the records scoring highest against it, top-k over the same weight `predict` writes and +exact with respect to the model. It needs no blocking and no candidate pairs: interning makes a +string level a question about a *value*, so every string metric in the query runs once per +distinct value of a column rather than once per record, after which a record costs a load and a +compare. At the 20M target a query naming one name column is answered in 598 ms on one thread +and 166 ms on eight. + Two things are worth stating plainly. The **approximate-nearest-neighbour source was retired by measurement rather than built**. diff --git a/docs/model.md b/docs/model.md index 3af1ed1..3550518 100644 --- a/docs/model.md +++ b/docs/model.md @@ -131,6 +131,16 @@ whose emails differ is worth **−18.06 bits**. The prior contributes \(\log_2(\lambda/(1-\lambda)) = -23.6\) bits at \(\lambda = 7.7\times10^{-8}\), which is the hurdle every pair starts behind. +!!! note "λ is the prior of a *pair*, and a query asks a different question" + λ is the match rate over the pair space, which is what deduplication faces: of all + \(\binom{N}{2}\) pairs, how many are the same person. A query record against N records + faces something else — roughly "I expect this person to be in here about once" — and the + two differ by orders of magnitude. [`search`](commands/search.md) therefore takes the prior + as an option (`--expected-matches`) rather than inheriting it, and defaults to the model's + so that a hit scores exactly what `predict` would have given that pair. The prior is a + constant added to every pair, so it moves the posterior and where a threshold sits, and + never the order. + !!! warning "Conditional independence is doing real work" Record data violates it — forename and sex are correlated, postcode and address more so. Splink carries the same exposure. Correlated comparisons double-count evidence and push diff --git a/docs/python.md b/docs/python.md index a1adae2..ec337a4 100644 --- a/docs/python.md +++ b/docs/python.md @@ -213,6 +213,45 @@ Which member names a cluster depends on the order the edges arrived, so two runs The numpy arrays view the C++ vectors and keep the assignment alive for as long as they do. A `Linker` holds every column, so its `cluster` costs no second load but more memory than the command, which loads the id column alone; `cpplink.cluster_file(schema_path, files, predictions)` is that lighter path. +### 8. Search for a record + +A resident `Linker` is a search service: one store loaded once, a query answered in milliseconds. + +```python +hits = linker.search( + {"last_name": "zolnerowich", "dob": "1979-08-06", "postcode": "4508"}, + model, + k=10, + expected_matches=1, +) +for hit in hits: + print(hit.id, hit.match_weight, hit.match_probability) +hits.report.walk_seconds, hits.report.gather_seconds # what each phase cost +``` + +The query maps column names to values, as text, in the form the file holds them; a date is `YYYY-MM-DD` and is refused in any other form. +A list or a tuple is spread over the repeats a list column takes one element per, and `None` is the same as leaving the column out. +A column the query does not name is **missing**, not empty, so its comparison lands on its null level and costs nothing at all. + +The answer is exact with respect to the model: the same `k` records, in the same order, as scoring the query against every record. +`SearchResult` iterates its hits, indexes them and carries the printed report as `text`, the numbers as `report` and the command's `--json` as `json()`. + +`expected_matches` is the prior stated as the number of records here you expect to be the person asked about, which is the question a search asks; λ is the match rate over the *pair* space, which is the question deduplication asks, and the two differ by orders of magnitude. +Without it the model's own prior applies and a hit scores exactly what `predict` would have given that pair. +The prior is a constant added to every hit, so it moves the posterior and where a threshold sits, and never the order. + +```python +hits = linker.search({"last_name": "zolnerowich"}, model, k=3, explain=True) +print(hits.waterfalls[0]) # the same ledger `explain` prints, for that pair +``` + +`explain=True` builds the ledger behind each hit during the search, which is when it has to happen: the query is a record of the store only while the search runs, so a result that outlived it would be a result nothing could explain. +`clusters=` takes the file `cluster` wrote and labels each hit with the cluster it belongs to, marking the ones that are another record of an entity already listed. + +Two limits, both narrow. +`search` reads one input and raises over several, because the query row would have to be a dataset of its own. +And one search runs at a time against one store; `threads` splits the work *within* a query rather than running several. + ### And then ```python @@ -247,7 +286,7 @@ That text is the run's, not a recomputation. ## What it costs, and what it does not -The `Linker` releases the GIL for every stage that does work: the load, `estimate`, `predict`, `cluster`, `rescore`, `profile`, `levels`, `simplify`, `recall` and `completeness`. +The `Linker` releases the GIL for every stage that does work: the load, `estimate`, `predict`, `cluster`, `rescore`, `profile`, `levels`, `simplify`, `recall`, `completeness` and `search`. 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. diff --git a/python/bindings/common.hpp b/python/bindings/common.hpp index 024dedb..d342ac4 100644 --- a/python/bindings/common.hpp +++ b/python/bindings/common.hpp @@ -16,10 +16,12 @@ #include "cpplink/blocking.hpp" #include "cpplink/cluster.hpp" #include "cpplink/comparison.hpp" +#include "cpplink/explain.hpp" #include "cpplink/parquet_loader.hpp" #include "cpplink/predict.hpp" #include "cpplink/record_store.hpp" #include "cpplink/schema.hpp" +#include "cpplink/search.hpp" namespace cpplink { namespace python { @@ -77,6 +79,14 @@ class Session { const Schema& schema() const { return schema_; } const RecordStore& store() const { return *store_; } + // The store and the bound comparisons as a `Searcher` needs them: it makes + // the query a row, which is a write, and grows the signature tables with any + // dictionary that row extends. Nothing else here is allowed to write, and + // the store reserved room for exactly this row when it was finalized, so the + // borrowed columns a prediction or cluster table hands to Python keep + // pointing at what they pointed at. + RecordStore* mutable_store() { return store_.get(); } + ComparisonSet* mutable_comparisons(); const LoadStats& stats() const { return stats_; } const std::vector& paths() const { return paths_; } PairMode mode() const { return mode_; } @@ -158,6 +168,16 @@ class ArrowTable : public std::enable_shared_from_this { std::vector names_; }; +// One search's hits, each with the ledger behind it where the caller asked for +// one. The waterfalls are built while the query is still a row, because the row +// goes as soon as the search returns and every report that explains a pair takes +// a pair of rows. +struct SearchOutcome { + SearchReport report; + std::vector waterfalls; + std::string text; +}; + // The predictions of one run: `dataset_a, id_a, dataset_b, id_b` (the datasets // only over several inputs), `gamma`, `match_weight`, `match_probability`, the // merged file's own columns. @@ -205,6 +225,7 @@ void BindSchema(py::module_& m); void BindModel(py::module_& m); SessionClass BindSession(py::module_& m); void BindStages(py::module_& m, SessionClass* session); +void BindSearch(py::module_& m, SessionClass* session); void BindDiagnostics(py::module_& m, SessionClass* session); } // namespace python diff --git a/python/bindings/module.cpp b/python/bindings/module.cpp index f25daff..6ae6692 100644 --- a/python/bindings/module.cpp +++ b/python/bindings/module.cpp @@ -171,5 +171,6 @@ PYBIND11_MODULE(_cpplink, m) { cpplink::python::BindTables(m); cpplink::python::SessionClass session = cpplink::python::BindSession(m); cpplink::python::BindStages(m, &session); + cpplink::python::BindSearch(m, &session); cpplink::python::BindDiagnostics(m, &session); } diff --git a/python/bindings/search.cpp b/python/bindings/search.cpp new file mode 100644 index 0000000..e998f3a --- /dev/null +++ b/python/bindings/search.cpp @@ -0,0 +1,195 @@ +// Copyright 2026 Mathieu Fourment +// SPDX-License-Identifier: MIT + +#include "cpplink/search.hpp" + +#include +#include +#include +#include + +#include +#include + +#include "bindings/common.hpp" +#include "cpplink/explain.hpp" +#include "cpplink/model.hpp" +#include "cpplink/neighbourhood.hpp" +#include "cpplink/pipeline.hpp" +#include "cpplink/score.hpp" + +namespace cpplink { +namespace python { + +namespace { + +// The query as Python gives it: a mapping of column to value, with a list or a +// tuple spread over the repeats a list column takes one element per. +QueryRecord RecordFrom(const py::dict& fields) { + QueryRecord query; + for (const auto& item : fields) { + const std::string column = py::cast(py::str(item.first)); + const py::handle value = item.second; + if (value.is_none()) continue; + if (py::isinstance(value) || py::isinstance(value)) { + for (const py::handle element : value) { + query.Set(column, py::cast(py::str(element))); + } + continue; + } + query.Set(column, py::cast(py::str(value))); + } + return query; +} + +// One search, start to finish, with the query row installed for exactly as long +// as it takes to answer and explain it. +// +// The ledgers are built here rather than handed back as rows to explain later, +// because the query stops being a row when the searcher goes and every report +// that explains a pair takes two rows. A result that outlived its query row +// would be a result nothing could explain. +SearchOutcome SearchStage(Session& session, const py::dict& fields, + const py::object& model_object, size_t k, + const py::object& threshold, const py::object& probability, + unsigned threads, const py::object& expected_matches, + const py::object& prior_weight, bool explain, + const std::string& clusters, double tf_damping, bool fuzzy_tf, + const BallOptions& ball, bool interactions) { + const Model& model = model_object.cast(); + const QueryRecord query = RecordFrom(fields); + Check(!query.fields.empty(), "search: give at least one column and value"); + + ComparisonSet* comparisons = session.mutable_comparisons(); + RecordStore* store = session.mutable_store(); + + ScoreOptions score; + score.tf_damping = tf_damping; + score.use_interactions = interactions; + BallTables balls; + std::ostringstream ball_text; + if (fuzzy_tf) { + py::gil_scoped_release release; + BuildBallTables(*comparisons, *store, ball, &balls, ball_text); + } + Scorer scorer; + std::string error; + Check(scorer.Bind(model, *comparisons, *store, score, &error, + fuzzy_tf ? &balls : nullptr), + "search: " + error); + + SearchOptions options; + options.k = k; + options.threads = threads; + if (!threshold.is_none() || !probability.is_none()) { + options.threshold = ThresholdFrom(threshold, probability, "search"); + } + if (!expected_matches.is_none() && !prior_weight.is_none()) { + throw Error("search: give expected_matches or prior_weight, not both"); + } + if (!expected_matches.is_none()) { + options.override_prior = true; + options.prior_weight = PriorWeightForExpected(py::cast(expected_matches), + store->NumRecords()); + } else if (!prior_weight.is_none()) { + options.override_prior = true; + options.prior_weight = py::cast(prior_weight); + } + + SearchOutcome outcome; + Searcher searcher; + Check(searcher.Bind(store, comparisons, &scorer, &error), "search: " + error); + { + py::gil_scoped_release release; + if (!searcher.Search(query, options, &outcome.report, &error)) { + // The searcher takes the query row back on destruction, so nothing + // is left behind whichever way this goes. + error = "search: " + error; + } else { + error.clear(); + if (explain) { + // The store's row is the first side throughout: an exact level's + // term-frequency move reads the frequency of `a`'s value, and the + // query's own value was never counted into the table. + for (const SearchHit& hit : outcome.report.hits) { + outcome.waterfalls.push_back( + BuildPairWaterfall(*store, *comparisons, scorer, hit.row, + searcher.QueryRow(), &model)); + } + } + } + } + Check(error.empty(), error); + if (!clusters.empty()) { + Check(GroupHitsByCluster(*store, clusters, &outcome.report, &error), + "search: " + error); + } + outcome.text = + CaptureText([&](std::ostream& out) { PrintSearchReport(outcome.report, out); }); + return outcome; +} + +} // namespace + +void BindSearch(py::module_& m, SessionClass* session) { + py::class_(m, "SearchHit", "One record a query scored against.") + .def_readonly("row", &SearchHit::row) + .def_readonly("id", &SearchHit::id) + .def_readonly("dataset", &SearchHit::dataset) + .def_readonly("gamma", &SearchHit::gamma) + .def_readonly("match_weight", &SearchHit::weight) + .def_readonly("match_probability", &SearchHit::probability) + .def_readonly("cluster", &SearchHit::cluster) + .def_readonly("cluster_best", &SearchHit::cluster_best) + .def("__repr__", [](const SearchHit& hit) { + return "SearchHit(id=" + hit.id + + ", match_weight=" + std::to_string(hit.weight) + ")"; + }); + + py::class_(m, "SearchReport", + "What one search did: the hits, and what each phase cost.") + .def_readonly("hits", &SearchReport::hits) + .def_readonly("records", &SearchReport::records) + .def_readonly("tabulated", &SearchReport::tabled) + .def_readonly("evaluated", &SearchReport::evaluated) + .def_readonly("constant", &SearchReport::constant) + .def_readonly("values_walked", &SearchReport::values_walked) + .def_readonly("values_adopted", &SearchReport::values_adopted) + .def_readonly("rescored", &SearchReport::rescored) + .def_readonly("prior", &SearchReport::prior) + .def_readonly("model_prior", &SearchReport::model_prior) + .def_readonly("walk_seconds", &SearchReport::walk_seconds) + .def_readonly("gather_seconds", &SearchReport::gather_seconds) + .def_readonly("threads", &SearchReport::threads) + .def_readonly("clusters", &SearchReport::clusters); + + py::class_( + m, "SearchResult", + "The hits of one query, with the ledger behind each where `explain` asked\n" + "for one. Iterating the result iterates the hits.") + .def_property_readonly("hits", + [](const SearchOutcome& o) { return o.report.hits; }) + .def_readonly("report", &SearchOutcome::report) + .def_readonly("waterfalls", &SearchOutcome::waterfalls) + .def_readonly("text", &SearchOutcome::text) + .def("json", [](const SearchOutcome& o) { return SearchReportJson(o.report); }) + .def("__len__", [](const SearchOutcome& o) { return o.report.hits.size(); }) + .def("__getitem__", + [](const SearchOutcome& o, size_t index) { + if (index >= o.report.hits.size()) throw py::index_error(); + return o.report.hits[index]; + }) + .def("__repr__", [](const SearchOutcome& o) { return o.text; }); + + session->def("search", &SearchStage, py::arg("fields"), py::arg("model"), + py::arg("k") = 10, py::arg("threshold") = py::none(), + py::arg("probability") = py::none(), py::arg("threads") = 1, + py::arg("expected_matches") = py::none(), + py::arg("prior_weight") = py::none(), py::arg("explain") = false, + py::arg("clusters") = "", py::arg("tf_damping") = 1.0, + py::arg("fuzzy_tf") = false, py::arg("ball") = BallOptions{}, + py::arg("interactions") = true); +} + +} // namespace python +} // namespace cpplink diff --git a/python/bindings/store.cpp b/python/bindings/store.cpp index 9628f1c..bb260dd 100644 --- a/python/bindings/store.cpp +++ b/python/bindings/store.cpp @@ -156,6 +156,11 @@ const ComparisonSet& Session::Comparisons() { return *comparisons_; } +ComparisonSet* Session::mutable_comparisons() { + Comparisons(); + return comparisons_.get(); +} + std::unique_ptr Session::BindComparisons(bool use_signatures, bool use_ladders) { Check(!schema_.comparisons.empty(), "the schema declares no \"comparisons\""); diff --git a/python/cpplink/__init__.py b/python/cpplink/__init__.py index 613c6c0..5e8ea6a 100644 --- a/python/cpplink/__init__.py +++ b/python/cpplink/__init__.py @@ -54,6 +54,9 @@ RecallResult, RescoreReport, Schema, + SearchHit, + SearchReport, + SearchResult, SimplifyOptions, SimplifyReport, __version__, @@ -101,6 +104,9 @@ "RecallResult", "RescoreReport", "Schema", + "SearchHit", + "SearchReport", + "SearchResult", "SimplifyOptions", "SimplifyReport", "__version__", @@ -822,6 +828,69 @@ def explain( interactions=interactions, ) + def search( + self, + record: Mapping[str, Any], + model: Model | PathLike, + *, + k: int = 10, + threshold: float | None = None, + probability: float | None = None, + threads: int = 1, + expected_matches: float | None = None, + prior_weight: float | None = None, + explain: bool = False, + clusters: PathLike | None = None, + tf_damping: float = 1.0, + fuzzy_tf: bool = False, + ball_budget: int | None = None, + interactions: bool = True, + ) -> SearchResult: + """The ``k`` records that score highest against a query record. + + ``record`` maps column names to values; a list or tuple is spread over + the repeats a list column takes one element per, and ``None`` is the + same as leaving the column out. A column the query does not name is + *missing*, not empty, so its comparison lands on its null level. + + The answer is exact with respect to the model: the same ``k`` records, + in the same order, as scoring the query against every record. The + result iterates its hits and carries the report as + :attr:`SearchResult.report`; ``explain=True`` also builds the ledger + behind each hit into :attr:`SearchResult.waterfalls`, which has to + happen during the search because the query is a record of the store + only while the search runs. + + ``expected_matches`` is the prior stated as the number of records here + expected to be the person asked about, which is the question a search + asks; without it the model's own prior applies and a hit scores what + :meth:`predict` would have given that pair. ``clusters`` is a file + :meth:`cluster` wrote, and labels each hit with the cluster it belongs + to. + + Reads one input; over several it raises, for the reason + :attr:`SearchResult` documents. + """ + ball = BallOptions() + if ball_budget is not None: + ball.budget = ball_budget + return self._session.search( + dict(record), + _model(model), + k=k, + threshold=threshold, + probability=probability, + threads=threads, + expected_matches=expected_matches, + prior_weight=prior_weight, + explain=explain, + clusters=_path(clusters), + tf_damping=tf_damping, + fuzzy_tf=fuzzy_tf, + ball=ball, + interactions=interactions, + ) + def explain_blocking(self, count: bool = False) -> BlockingReport: """Price every blocking source without enumerating a pair; ``count`` also enumerates the deduplicated union.""" diff --git a/python/tests/test_reports.py b/python/tests/test_reports.py index 2d672fb..01ad7a8 100644 --- a/python/tests/test_reports.py +++ b/python/tests/test_reports.py @@ -38,6 +38,9 @@ "ClusterQuality": 9, "MergeReport": 10, "Explanation": 8, + "SearchResult": 3, + "SearchReport": 14, + "SearchHit": 8, "PairWaterfall": 16, "WaterfallStep": 11, "BlockingReport": 7, @@ -56,6 +59,13 @@ } +def _first_surname(sample) -> str: + """A value the store holds, so the search finds something.""" + pq = pytest.importorskip("pyarrow.parquet") + table = pq.read_table(str(sample.parquet), columns=["last_name"]) + return str(table.column("last_name")[0].as_py()) + + @pytest.fixture(scope="module") def reports(sample, tmp_path_factory: pytest.TempPathFactory) -> dict[str, object]: root = tmp_path_factory.mktemp("reports") @@ -69,6 +79,7 @@ def reports(sample, tmp_path_factory: pytest.TempPathFactory) -> dict[str, objec linker.cluster(sample.predictions, truth=sample.truth) cluster = linker.last_cluster explanation = linker.explain(linker.id_of(0), linker.id_of(1), model=model) + hits = linker.search({"last_name": _first_surname(sample)}, model, k=3) recall = linker.recall(sample.truth, why=True) levels, _ = linker.levels() simplify, _ = linker.simplify(model) @@ -95,6 +106,9 @@ def reports(sample, tmp_path_factory: pytest.TempPathFactory) -> dict[str, objec "ClusterReport": cluster.report, "ClusterQuality": cluster.quality, "Explanation": explanation, + "SearchResult": hits, + "SearchReport": hits.report, + "SearchHit": hits.hits[0], "PairWaterfall": explanation.waterfall, "WaterfallStep": explanation.waterfall.steps[0], "BlockingReport": linker.explain_blocking(), diff --git a/python/tests/test_search.py b/python/tests/test_search.py new file mode 100644 index 0000000..31f7dc8 --- /dev/null +++ b/python/tests/test_search.py @@ -0,0 +1,249 @@ +# Copyright 2026 Mathieu Fourment +# SPDX-License-Identifier: MIT +"""`Linker.search`: top-k retrieval over the weight, from the resident session. + +The query is a record of the store only while the search runs, so what this +checks above all is that nothing is left behind afterwards -- neither a row, nor +a moved buffer that a zero-copy table was pointing at. +""" + +from __future__ import annotations + +import json +import math + +import pytest +from conftest import needs_core_parquet + +import cpplink + + +def _value(sample, row: int, column: str) -> str: + """One cell of the sample file, as the text a caller would type.""" + pq = pytest.importorskip("pyarrow.parquet") + table = pq.read_table(str(sample.parquet), columns=[column]) + return str(table.column(column)[row].as_py()) + + +@needs_core_parquet +def test_a_record_finds_itself(sample) -> None: + query = { + "last_name": _value(sample, 0, "last_name"), + "dob": _value(sample, 0, "dob"), + "postcode": _value(sample, 0, "postcode"), + } + hits = sample.linker.search(query, sample.model, k=5) + assert len(hits) == 5 + assert hits[0].id == sample.linker.id_of(0) + assert hits[0].match_weight > 0 + # Descending weight, and the probability is the weight through the logistic. + weights = [hit.match_weight for hit in hits] + assert weights == sorted(weights, reverse=True) + assert hits[0].match_probability == pytest.approx( + cpplink.probability_for_weight(hits[0].match_weight) + ) + + +@needs_core_parquet +def test_the_report_counts_the_three_kinds_of_comparison(sample) -> None: + result = sample.linker.search( + {"last_name": _value(sample, 0, "last_name")}, sample.model, k=3 + ) + report = result.report + assert report.records == sample.linker.records + # Every comparison is one of the three, and the ones the query says nothing + # about are free. + total = report.tabulated + report.evaluated + report.constant + assert total == len(sample.schema.to_dict()["comparisons"]) + assert report.constant > 0 + assert report.values_walked > 0 + assert report.rescored <= report.records + assert report.threads == 1 + assert report.prior == report.model_prior + + +@needs_core_parquet +def test_explain_gives_the_ledger_behind_every_hit(sample) -> None: + result = sample.linker.search( + {"last_name": _value(sample, 0, "last_name")}, + sample.model, + k=4, + explain=True, + ) + assert len(result.waterfalls) == len(result.hits) + for hit, waterfall in zip(result.hits, result.waterfalls, strict=True): + assert waterfall.gamma == hit.gamma + assert waterfall.weight == pytest.approx(hit.match_weight) + # The store's record is the first side, the query the second. + assert waterfall.row_a == hit.row + assert waterfall.id_a == hit.id + assert "Match weight" in waterfall.text + # Without it there is nothing to explain, which is the default. + plain = sample.linker.search( + {"last_name": _value(sample, 0, "last_name")}, sample.model, k=4 + ) + assert plain.waterfalls == [] + + +@needs_core_parquet +def test_the_prior_shifts_every_hit_and_reorders_nothing(sample) -> None: + query = {"last_name": _value(sample, 0, "last_name")} + plain = sample.linker.search(query, sample.model, k=6) + shifted = sample.linker.search(query, sample.model, k=6, expected_matches=1) + assert shifted.report.prior != plain.report.prior + move = shifted.report.prior - plain.report.prior + assert [hit.id for hit in shifted] == [hit.id for hit in plain] + for before, after in zip(plain, shifted, strict=True): + assert after.match_weight == pytest.approx(before.match_weight + move) + # And it is the odds of one record in the store being the one asked about. + n = sample.linker.records + assert shifted.report.prior == pytest.approx(math.log2(1.0 / (n - 1.0))) + + +@needs_core_parquet +def test_a_threshold_can_refuse_every_record(sample) -> None: + result = sample.linker.search( + {"last_name": _value(sample, 0, "last_name")}, + sample.model, + k=10, + threshold=1e6, + ) + assert len(result) == 0 + assert "no record here is this one" in result.text + + +@needs_core_parquet +def test_threads_change_nothing(sample) -> None: + query = { + "last_name": _value(sample, 0, "last_name"), + "dob": _value(sample, 0, "dob"), + } + one = sample.linker.search(query, sample.model, k=8, threads=1) + many = sample.linker.search(query, sample.model, k=8, threads=4) + assert [(h.row, h.match_weight) for h in one] == [ + (h.row, h.match_weight) for h in many + ] + assert many.report.threads == 4 + + +@needs_core_parquet +def test_a_column_the_query_omits_is_missing_not_empty(sample) -> None: + # None is the same as leaving the column out, and both leave that + # comparison on its null level for every record, which is free. + named = sample.linker.search( + {"last_name": _value(sample, 0, "last_name"), "postcode": None}, + sample.model, + k=3, + ) + omitted = sample.linker.search( + {"last_name": _value(sample, 0, "last_name")}, sample.model, k=3 + ) + assert [h.id for h in named] == [h.id for h in omitted] + assert named.report.constant == omitted.report.constant + + +@needs_core_parquet +def test_the_store_is_left_as_it_was_found(sample) -> None: + """The query is a record for the length of the search and no longer. + + The zero-copy tables the bindings hand out are borrowed pointers into the + store's id arena, so a query row that moved one would leave them dangling. + The store reserves room for it at load, and this is the test of that: a + table exported before the search still reads after it. + """ + pa = pytest.importorskip("pyarrow") + sample.linker.predict(sample.model, threshold=10) + table = pa.table(sample.linker.last_predict.table) + before = table.column("id_a").to_pylist() + records = sample.linker.records + + sample.linker.search( + {"last_name": "a-surname-this-store-has-never-held"}, sample.model, k=5 + ) + + assert sample.linker.records == records + assert table.column("id_a").to_pylist() == before + + +@needs_core_parquet +def test_a_value_the_store_never_held_is_still_compared(sample) -> None: + result = sample.linker.search( + {"last_name": "zzzzzzzzznotasurname"}, sample.model, k=3 + ) + assert result.report.values_adopted >= 1 + assert len(result) == 3 + # And the next query is unaffected by it. + again = sample.linker.search( + {"last_name": _value(sample, 0, "last_name")}, sample.model, k=3 + ) + assert again.report.values_adopted == 0 + assert again[0].id == sample.linker.id_of(0) + + +@needs_core_parquet +def test_a_list_or_tuple_is_spread_over_the_repeats(sample) -> None: + # address_tokens is the sample's list column; two elements are two fields. + result = sample.linker.search({"address_tokens": ["one", "two"]}, sample.model, k=2) + assert len(result) == 2 + + +@needs_core_parquet +def test_json_carries_the_hits(sample) -> None: + result = sample.linker.search( + {"last_name": _value(sample, 0, "last_name")}, sample.model, k=3 + ) + parsed = json.loads(result.json()) + assert parsed["records"] == sample.linker.records + assert len(parsed["hits"]) == 3 + assert parsed["hits"][0]["id"] == result[0].id + assert parsed["hits"][0]["match_weight"] == pytest.approx(result[0].match_weight) + + +@needs_core_parquet +def test_an_empty_query_and_a_bad_date_are_refused(sample) -> None: + with pytest.raises(cpplink.Error, match="at least one column"): + sample.linker.search({}, sample.model) + with pytest.raises(cpplink.Error, match="YYYY-MM-DD"): + sample.linker.search({"dob": "24/01/1979"}, sample.model) + with pytest.raises(cpplink.Error, match="not both"): + sample.linker.search( + {"last_name": "x"}, sample.model, expected_matches=1, prior_weight=0.0 + ) + + +@needs_core_parquet +def test_link_mode_is_refused_with_the_reason(link) -> None: + model, _ = link.linker.estimate() + with pytest.raises(cpplink.Error, match="one input"): + link.linker.search({"last_name": "smith"}, model) + + +@needs_core_parquet +def test_matches_the_command_line(sample) -> None: + """The binding and `cpplink search` are one implementation, so the hits and + the printed report are the same.""" + query = { + "last_name": _value(sample, 0, "last_name"), + "dob": _value(sample, 0, "dob"), + "postcode": _value(sample, 0, "postcode"), + } + result = sample.linker.search(query, sample.model, k=5) + args = [ + "search", + "--schema", + str(sample.schema_path), + "--model", + str(sample.model_path), + str(sample.parquet), + "-k", + "5", + ] + for column, value in query.items(): + args += ["--field", f"{column}={value}"] + cli = cpplink.run(args) + assert cli.code == 0, cli.stderr + for hit in result: + assert hit.id in cli.stdout + # The two reports differ only in what the timings read. + body = result.text.split("Record ", 1)[1] + assert body == cli.stdout.split("Record ", 1)[1] diff --git a/src/cpplink/record_store.cpp b/src/cpplink/record_store.cpp index 5220956..1e8be0b 100644 --- a/src/cpplink/record_store.cpp +++ b/src/cpplink/record_store.cpp @@ -124,6 +124,28 @@ void RecordStore::Finalize() { } } } + + ReserveQueryRow(); +} + +void RecordStore::ReserveQueryRow() { + const size_t rows = static_cast(num_records_) + 1; + for (Column& column : columns_) { + if (auto* col = std::get_if(&column)) { + col->ids.reserve(rows); + if (!col->tf.empty()) col->tf.reserve(col->tf.size() + 1); + } else if (auto* col = std::get_if(&column)) { + col->offsets.reserve(rows + 1); + } else if (auto* col = std::get_if(&column)) { + col->values.reserve(rows); + } else if (auto* col = std::get_if(&column)) { + col->values.reserve(rows); + } else if (auto* col = std::get_if(&column)) { + col->values.reserve(rows); + } + } + ids_.offsets.reserve(rows + 1); + ids_.text.reserve(ids_.text.size() + kQueryIdBytes); } uint32_t RecordStore::DistinctValues(size_t index) const { diff --git a/src/cpplink/record_store.hpp b/src/cpplink/record_store.hpp index 6f7131c..e06148a 100644 --- a/src/cpplink/record_store.hpp +++ b/src/cpplink/record_store.hpp @@ -21,6 +21,10 @@ inline constexpr int32_t kNullDate = std::numeric_limits::min(); // A missing boolean. Like kNullId, it never compares equal to anything. inline constexpr int8_t kNullBoolean = -1; +// Bytes kept past the loaded ids for the one a search's query row carries. The +// id it writes is short and fixed; this is that with room to spare. +inline constexpr size_t kQueryIdBytes = 64; + // The shape of the pair space a run enumerates. // // Deduplication is the upper triangle of one input; linking is the cross-product @@ -169,6 +173,23 @@ class RecordStore { // no term frequencies. Also releases the dictionaries' load-time indexes. void Finalize(); + // Leaves room past the loaded rows for the one row a search appends. + // + // `search` makes the query a row so that every level, bound and report the + // pair path has applies to it unchanged, and removes it again afterwards. + // What that must not do is *move* anything, because the Python bindings hand + // out tables whose columns are borrowed pointers into this store -- the id + // arena above all, which every prediction and cluster table uses as its + // dictionary. Growing a vector at capacity would reallocate it and leave + // those pointers dangling, so the room is made here, once, while nothing can + // yet be pointing at anything. Finalize calls it, so it is true of every + // loaded store and no caller has to remember. + // + // A list column's flat `ids` is the one array a query can grow by more than + // one element, its cell being as long as the query cares to make it. Nothing + // borrows it. + void ReserveQueryRow(); + uint32_t DistinctValues(size_t index) const; uint64_t NullCount(size_t index) const; MemoryReport Memory() const; diff --git a/tests/search_test.cpp b/tests/search_test.cpp index ac4d86a..095de22 100644 --- a/tests/search_test.cpp +++ b/tests/search_test.cpp @@ -274,6 +274,44 @@ TEST_F(SearchFixture, AValueTheStoreNeverHeldIsStillCompared) { EXPECT_EQ(store_->NumRecords(), kRecords); } +// And nothing *moves*, which is a stronger claim than nothing changing. +// +// The Python bindings hand out tables whose columns are borrowed pointers into +// this store -- the id arena above all -- so a query row that grew a vector at +// capacity would reallocate it and leave those pointers dangling. `Finalize` +// reserves the room for exactly this row, and this is the test of it: every +// buffer a query touches is at the same address afterwards. +TEST_F(SearchFixture, AQueryRowMovesNothing) { + const char* ids_text = store_->ids().text.data(); + const uint64_t* ids_offsets = store_->ids().offsets.data(); + const uint32_t* surname = + std::get(store_->column(0)).ids.data(); + const uint32_t* surname_tf = + std::get(store_->column(0)).tf.data(); + const int32_t* dob = std::get(store_->column(2)).values.data(); + { + cpplink::Searcher searcher; + std::string error; + ASSERT_TRUE(searcher.Bind(store_.get(), &comparisons_, &scorer_, &error)) + << error; + cpplink::SearchOptions options; + cpplink::SearchReport report; + // A surname the store never held, so the dictionary and its term + // frequencies grow too, which is the case with the most to move. + cpplink::QueryRecord query; + query.Set("surname", "notasurnamethisstoreholds"); + query.Set("city", "city1"); + query.Set("dob", "1981-01-24"); + ASSERT_TRUE(searcher.Search(query, options, &report, &error)) << error; + EXPECT_EQ(report.values_adopted, 2u); + } + EXPECT_EQ(store_->ids().text.data(), ids_text); + EXPECT_EQ(store_->ids().offsets.data(), ids_offsets); + EXPECT_EQ(std::get(store_->column(0)).ids.data(), surname); + EXPECT_EQ(std::get(store_->column(0)).tf.data(), surname_tf); + EXPECT_EQ(std::get(store_->column(2)).values.data(), dob); +} + // The query row is removed when the searcher goes, and the pairs the store // already held score what they scored before it arrived. TEST_F(SearchFixture, TheStoreIsLeftAsItWasFound) { From d8662c185cb20f4a583e914dc58d7eed2218528e Mon Sep 17 00:00:00 2001 From: 4ment Date: Sun, 27 Sep 2026 08:32:09 +1000 Subject: [PATCH 3/5] Improve docs for fuzzy TF --- docs/commands/predict.md | 24 +++++++++++++++++++++--- docs/model.md | 1 + 2 files changed, 22 insertions(+), 3 deletions(-) diff --git a/docs/commands/predict.md b/docs/commands/predict.md index f0dd58c..db92cd5 100644 --- a/docs/commands/predict.md +++ b/docs/commands/predict.md @@ -250,9 +250,27 @@ A term-frequency adjustment needs an exact-match level, here and in splink. So two records sharing the misspelling "Zolnerowitch" against "Zolnerowich" get the averaged fuzzy weight, and the rarity that makes the pair convincing is thrown away. -`--fuzzy-tf` replaces the value's own frequency with the mass of its *neighbourhood* — the -share of the file that falls inside the ball the level defines — which is exactly `p_v` again -when the level is exact. +`--fuzzy-tf` replaces the value's own frequency with the mass of its *neighbourhood*: the share of the file that falls inside the ball the level defines. +For a value \(v\) of comparison \(c\) and a level \(\ell\), the ball is every value that would land on \(\ell\) against \(v\), and its mass sums their relative frequencies: + +\[ +B_\ell(v) \;=\; \{\, w : (v, w) \text{ reaches level } \ell \,\} +\qquad +M_\ell(v) \;=\; \sum_{w \in B_\ell(v)} p_w +\] + +The two records of a pair sit in different neighbourhoods, so the adjustment reads the geometric mean of their two masses: + +\[ +\Delta_{c,\ell}(a, b) \;=\; w_c \cdot \log_2 \frac{u_{c,\ell}}{\sqrt{M_\ell(v_a)\, M_\ell(v_b)}} +\;=\; w_c \left( \log_2 u_{c,\ell} \;-\; \tfrac{1}{2}\bigl(\log_2 M_\ell(v_a) + \log_2 M_\ell(v_b)\bigr) \right) +\] + +where \(u_{c,\ell}\) is the level's learned \(u\) and \(w_c\) the damping factor (`--tf-damping`). +When the level is exact, \(B_\ell(v) = \{v\}\), both values are the same, and the formula is the [exact-level adjustment](../model.md#term-frequency-adjustment) with \(M = p_v\). +On a fuzzy level the geometric mean moves half as far as either side would on its own. +A pair where either side has no mass is not adjusted. +The admissible bracket follows from the same formula: \(\Delta_{\max}\) comes from the smallest mass the level holds in the column and \(\Delta_{\min}\) from the largest. Computing it is a similarity self-join over every distinct value, which is affordable here only because values are interned: the join is over the 149k distinct surnames of the 1M sample, not its 1M rows, and it is amortised over every pair the run scores. diff --git a/docs/model.md b/docs/model.md index 3550518..bf300b3 100644 --- a/docs/model.md +++ b/docs/model.md @@ -173,6 +173,7 @@ where \(p_v\) is the value's relative frequency, \(p_{\min,c}\) is the singleton \(w_c\) is a damping factor (`--tf-damping`, 1.0 by default). The adjustment applies only on **exact** levels of comparisons declared `"term_frequency": true`, and only where both sides agree — so the value can be read from either record. +[`predict --fuzzy-tf`](commands/predict.md#term-frequency-on-the-fuzzy-levels) extends it to the fuzzy levels by replacing \(p_v\) with the mass of the value's neighbourhood under the level. TF adjustment is the one thing that breaks γ's sufficiency, which is why cpplink keeps it **out of EM and applies it at scoring**. Estimation stays exact on the pattern histogram, and From 7167714e29bbbc55e431bdfdd953f609714059aa Mon Sep 17 00:00:00 2001 From: 4ment Date: Sun, 27 Sep 2026 11:55:25 +1000 Subject: [PATCH 4/5] Add term-frequency controls to search A common exact name could rank below a typo of it, because only the exact level was term-frequency adjusted. `cpplink search` gains --fuzzy-tf and --ball-budget, and `Linker.search` takes tf_damping and fuzzy_tf. Search measures the query's own neighbourhood mass from the dictionary walk and pins it on the scorer, so an adopted value is adjusted too (it read past the end of the table before) and a stored value scores what predict would. The mass is clamped to the table's range to keep the bracket admissible. Neighbourhood tables are built only for the comparisons the query names and the Python session keeps them. Co-Authored-By: Claude Opus 5.5 (1M context) --- docs/commands/search.md | 17 ++++++- python/bindings/common.hpp | 15 ++++++ python/bindings/search.cpp | 15 +++--- python/bindings/store.cpp | 32 ++++++++++++ python/cpplink/__init__.py | 6 +++ src/cpplink/app.cpp | 18 ++++++- src/cpplink/neighbourhood.cpp | 7 ++- src/cpplink/neighbourhood.hpp | 8 ++- src/cpplink/pipeline.cpp | 5 +- src/cpplink/pipeline.hpp | 4 +- src/cpplink/score.cpp | 36 +++++++++++++- src/cpplink/score.hpp | 14 ++++++ src/cpplink/search.cpp | 37 +++++++++++++- src/cpplink/search.hpp | 14 ++++-- tests/search_test.cpp | 94 +++++++++++++++++++++++++++++++++++ 15 files changed, 301 insertions(+), 21 deletions(-) diff --git a/docs/commands/search.md b/docs/commands/search.md index df3342a..a934332 100644 --- a/docs/commands/search.md +++ b/docs/commands/search.md @@ -35,7 +35,9 @@ cpplink search --schema --model | `--clusters ` | — | a file written by [`cluster`](cluster.md); labels each hit with its cluster | | `--explain` | off | print the [waterfall](explain.md) behind every hit | | `--json` | off | the whole report as one JSON object | -| `--tf-damping F` | 1.0 | scale the term-frequency adjustment | +| `--tf-damping F` | 1.0 | scale the term-frequency adjustment; 0 turns it off | +| `--fuzzy-tf` | off | adjust the fuzzy levels too, by the mass of each value's neighbourhood, as [`predict --fuzzy-tf`](predict.md) does | +| `--ball-budget N` | 40e9 | value pairs the neighbourhood join may look at per comparison; a larger dictionary gets no fuzzy adjustment | | `--no-interactions` | off | score the plain conditionally-independent model | | *(positional)* | — | required; the parquet file | @@ -43,6 +45,19 @@ A column the query does not name is **missing**, not empty. So a query carrying date of birth against a ten-column schema is scored as a record whose other eight fields were never recorded, which is what it is, and those comparisons land on their null level. +### A common exact name can rank below a rare typo + +Term frequency moves an exact agreement by how common the value is, so agreeing on `smith` is worth far less than agreeing on a rare surname. +Without `--fuzzy-tf` only the exact level moves: a `smythe` landing on a Jaro-Winkler level scores that level's weight whatever its neighbourhood looks like, and that weight can be larger than what an exact `smith` has left. +The query `John Smith` then returns the typos first. + +`--fuzzy-tf` is the fix that keeps the model: the fuzzy level pays for a crowded neighbourhood as the exact level pays for a common value, and `smith`'s neighbourhood contains `smith`. +On a 200,000-record `gen-sample` file, the commonest first and last name returns 6 of its 10 exact matches in the top 10 by default and 9 with `--fuzzy-tf`. +The neighbourhood of the query's own value is measured from the dictionary walk, so a query value the store never held is adjusted too, and a value the store does hold scores exactly what `predict --fuzzy-tf` would give the same pair. +The join behind it is built once per run, only for the comparisons the query names, and it is the slow part: 23 s over both name columns at 200,000 records. + +`--tf-damping 0` is the blunt alternative: every agreement counts the same, which returns all 10, and a rare name no longer outweighs a common one. + A list column takes one element per repeat of its field: `--field alias=bill --field alias=will`. A date is written `YYYY-MM-DD` and is refused in any other form. diff --git a/python/bindings/common.hpp b/python/bindings/common.hpp index d342ac4..026d8b1 100644 --- a/python/bindings/common.hpp +++ b/python/bindings/common.hpp @@ -17,6 +17,7 @@ #include "cpplink/cluster.hpp" #include "cpplink/comparison.hpp" #include "cpplink/explain.hpp" +#include "cpplink/neighbourhood.hpp" #include "cpplink/parquet_loader.hpp" #include "cpplink/predict.hpp" #include "cpplink/record_store.hpp" @@ -103,6 +104,17 @@ class Session { // --no-signatures` and `--no-ladders`. std::unique_ptr BindComparisons(bool use_signatures, bool use_ladders); + // The neighbourhood masses `fuzzy_tf` reads, for the comparisons in + // `wanted`, each built over the bound comparisons the first time it is + // wanted and kept. The dictionary self-join costs seconds to minutes where a + // search costs milliseconds, so a search asks only for the comparisons its + // query names: one over a column the query leaves out is its null level for + // every row and no fuzzy level of it can fire. Nothing the join reads changes + // after load -- a query adopts values past the end of a dictionary and takes + // them back -- and every table is dropped when the budget changes, since that + // decides which columns get one. + const BallTables& Balls(const BallOptions& options, const std::vector& wanted); + // The predictions of the last `predict`, which `cluster` with no table // given reads: the run's own rows, nothing to resolve. const std::shared_ptr& last_edges() const { return last_edges_; } @@ -121,6 +133,9 @@ class Session { std::unique_ptr estimation_plan_; std::unique_ptr prediction_plan_; std::unique_ptr comparisons_; + std::unique_ptr balls_; + std::vector balls_tried_; // by comparison: built, or refused with a reason + uint64_t balls_budget_ = 0; }; // The schema text `levels` and `simplify` rewrite: the source text where the diff --git a/python/bindings/search.cpp b/python/bindings/search.cpp index e998f3a..de14325 100644 --- a/python/bindings/search.cpp +++ b/python/bindings/search.cpp @@ -3,7 +3,6 @@ #include "cpplink/search.hpp" -#include #include #include #include @@ -66,16 +65,14 @@ SearchOutcome SearchStage(Session& session, const py::dict& fields, ScoreOptions score; score.tf_damping = tf_damping; score.use_interactions = interactions; - BallTables balls; - std::ostringstream ball_text; - if (fuzzy_tf) { - py::gil_scoped_release release; - BuildBallTables(*comparisons, *store, ball, &balls, ball_text); - } + // Kept on the session: the self-join is the one part of a fuzzy search that + // does not depend on the query's values, only on which columns it names. + const BallTables* balls = + fuzzy_tf ? &session.Balls(ball, ComparisonsNamedBy(session.schema(), query)) + : nullptr; Scorer scorer; std::string error; - Check(scorer.Bind(model, *comparisons, *store, score, &error, - fuzzy_tf ? &balls : nullptr), + Check(scorer.Bind(model, *comparisons, *store, score, &error, balls), "search: " + error); SearchOptions options; diff --git a/python/bindings/store.cpp b/python/bindings/store.cpp index bb260dd..02a0357 100644 --- a/python/bindings/store.cpp +++ b/python/bindings/store.cpp @@ -161,6 +161,38 @@ ComparisonSet* Session::mutable_comparisons() { return comparisons_.get(); } +const BallTables& Session::Balls(const BallOptions& options, + const std::vector& wanted) { + ComparisonSet* comparisons = mutable_comparisons(); + const size_t count = comparisons->Size(); + if (!balls_ || balls_budget_ != options.budget) { + balls_ = std::make_unique(); + balls_->tables.resize(count); + balls_->reasons.assign(count, std::string()); + balls_tried_.assign(count, false); + balls_budget_ = options.budget; + } + for (size_t c = 0; c < count && c < wanted.size(); ++c) { + if (!wanted[c] || balls_tried_[c]) continue; + BallMassTable table; + std::string reason; + bool built = false; + { + py::gil_scoped_release release; + built = table.Build(*comparisons, c, store_->NumRecords(), options, &reason); + } + if (built) { + balls_->bytes += table.BytesUsed(); + balls_->seconds += table.Seconds(); + balls_->tables[c] = std::move(table); + } else { + balls_->reasons[c] = reason; + } + balls_tried_[c] = true; + } + return *balls_; +} + std::unique_ptr Session::BindComparisons(bool use_signatures, bool use_ladders) { Check(!schema_.comparisons.empty(), "the schema declares no \"comparisons\""); diff --git a/python/cpplink/__init__.py b/python/cpplink/__init__.py index 5e8ea6a..4926e83 100644 --- a/python/cpplink/__init__.py +++ b/python/cpplink/__init__.py @@ -868,6 +868,12 @@ def search( :meth:`cluster` wrote, and labels each hit with the cluster it belongs to. + ``tf_damping`` scales the term-frequency move and 0 turns it off. + ``fuzzy_tf`` adjusts the fuzzy levels too, by how crowded each value's + neighbourhood is, which is what keeps a typo of a common name from + outranking the name itself; the neighbourhood tables are built for the + comparisons the query names the first time they are asked for, and kept. + Reads one input; over several it raises, for the reason :attr:`SearchResult` documents. """ diff --git a/src/cpplink/app.cpp b/src/cpplink/app.cpp index baf3d0e..83fa813 100644 --- a/src/cpplink/app.cpp +++ b/src/cpplink/app.cpp @@ -1865,6 +1865,8 @@ int RunSearch(const std::vector& args, std::ostream& out, bool as_json = false; bool expected_given = false; double expected = 0.0; + bool fuzzy_tf = false; + BallOptions ball; for (size_t i = 0; i < args.size(); ++i) { if (args[i] == "--schema") { if (!TakeValue(args, &i, &schema_path, err)) return 1; @@ -1903,6 +1905,11 @@ int RunSearch(const std::vector& args, std::ostream& out, } 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] == "--clusters") { @@ -1952,8 +1959,17 @@ int RunSearch(const std::vector& args, std::ostream& out, err << "cpplink: " << error << "\n"; return 1; } + // In JSON mode `out` carries nothing but the object, so the ball-table + // report goes to `err`, as `explain` does it. Only the comparisons the query + // names get a table: the rest are their null level on every row. + BallTables balls; + if (fuzzy_tf) { + const std::vector named = ComparisonsNamedBy(schema, query); + BuildBallTables(comparisons, store, ball, &balls, as_json ? err : out, &named); + } Scorer scorer; - if (!scorer.Bind(model, comparisons, store, score, &error)) { + if (!scorer.Bind(model, comparisons, store, score, &error, + fuzzy_tf ? &balls : nullptr)) { err << "cpplink search: " << error << "\n"; return 1; } diff --git a/src/cpplink/neighbourhood.cpp b/src/cpplink/neighbourhood.cpp index 58787e1..d7afe09 100644 --- a/src/cpplink/neighbourhood.cpp +++ b/src/cpplink/neighbourhood.cpp @@ -237,12 +237,17 @@ bool BallMassTable::Build(const ComparisonSet& comparisons, size_t comparison, } void BallTables::Build(const ComparisonSet& comparisons, uint64_t records, - const BallOptions& options) { + const BallOptions& options, const std::vector* only) { const auto started = std::chrono::steady_clock::now(); tables.resize(comparisons.Size()); reasons.assign(comparisons.Size(), std::string()); for (size_t c = 0; c < comparisons.Size(); ++c) { std::string reason; + if (only != nullptr && (c >= only->size() || !(*only)[c])) { + reasons[c] = "not named by the query"; + tables[c] = BallMassTable(); + continue; + } if (!tables[c].Build(comparisons, c, records, options, &reason)) { reasons[c] = reason; tables[c] = BallMassTable(); diff --git a/src/cpplink/neighbourhood.hpp b/src/cpplink/neighbourhood.hpp index 417cd86..fc3bffa 100644 --- a/src/cpplink/neighbourhood.hpp +++ b/src/cpplink/neighbourhood.hpp @@ -70,6 +70,10 @@ class BallMassTable { double Mass(size_t level, uint32_t id) const { return static_cast(mass_[level][id]) * inverse_records_; } + // A count of records as the same probability `Mass` returns. + double MassOf(uint64_t records) const { + return static_cast(records) * inverse_records_; + } double MinMass(size_t level) const { return min_mass_[level]; } double MaxMass(size_t level) const { return max_mass_[level]; } // Exact u for the level, over ordered draws, on the same denominator the @@ -102,8 +106,10 @@ struct BallTables { double seconds = 0.0; uint64_t bytes = 0; + // `only`, where given, names the comparisons to build; the rest are left + // empty with the reason that they were not asked for. void Build(const ComparisonSet& comparisons, uint64_t records, - const BallOptions& options); + const BallOptions& options, const std::vector* only = nullptr); bool Has(size_t comparison) const { return comparison < tables.size() && !tables[comparison].Empty(); } diff --git a/src/cpplink/pipeline.cpp b/src/cpplink/pipeline.cpp index c426ef3..3fba324 100644 --- a/src/cpplink/pipeline.cpp +++ b/src/cpplink/pipeline.cpp @@ -106,8 +106,9 @@ bool ResolveEdgeOutput(const std::string& command, const std::string& out, } void BuildBallTables(const ComparisonSet& comparisons, const RecordStore& store, - const BallOptions& options, BallTables* balls, std::ostream& out) { - balls->Build(comparisons, store.NumRecords(), options); + const BallOptions& options, BallTables* balls, std::ostream& out, + const std::vector* only) { + balls->Build(comparisons, store.NumRecords(), options, only); out << "Neighbourhood masses in " << std::fixed << std::setprecision(1) << balls->seconds << " s"; if (store.NumDatasets() > 1) { diff --git a/src/cpplink/pipeline.hpp b/src/cpplink/pipeline.hpp index 628d651..e7405c9 100644 --- a/src/cpplink/pipeline.hpp +++ b/src/cpplink/pipeline.hpp @@ -85,8 +85,10 @@ bool ResolveEdgeOutput(const std::string& command, const std::string& out, // Builds the neighbourhood masses a fuzzy term-frequency adjustment needs, and // says which columns got one. A column too large for the budget keeps today's // behaviour, which is worth saying out loud rather than degrading quietly. +// `only` limits the build to the comparisons it marks, as `search` asks. void BuildBallTables(const ComparisonSet& comparisons, const RecordStore& store, - const BallOptions& options, BallTables* balls, std::ostream& out); + const BallOptions& options, BallTables* balls, std::ostream& out, + const std::vector* only = nullptr); // One record named by its id, or by `:` where the inputs share ids. // Found is a row; missing and ambiguous are each an error that says so, the diff --git a/src/cpplink/score.cpp b/src/cpplink/score.cpp index 2f2eb7d..4deee0d 100644 --- a/src/cpplink/score.cpp +++ b/src/cpplink/score.cpp @@ -60,9 +60,11 @@ uint32_t TermFrequencyAdjustment::Frequency(uint64_t row) const { } double TermFrequencyAdjustment::Mass(uint64_t row) const { + if (row == pinned_row) return pinned_mass; if (ball == nullptr || strings == nullptr) return 0.0; const uint32_t id = strings->ids[row]; - if (id == kNullId) return 0.0; + // A value adopted after the table was built has no entry in it. + if (id == kNullId || id >= ball->Values()) return 0.0; return ball->Mass(level, id); } @@ -472,6 +474,38 @@ bool Scorer::AdjustsFuzzyLevels() const { return false; } +bool Scorer::AdjustsFuzzyLevels(size_t comparison) const { + for (const TermFrequencyAdjustment& adjustment : adjustments_) { + if (adjustment.fuzzy && adjustment.comparison == comparison) return true; + } + return false; +} + +void Scorer::PinMass(uint64_t row, size_t comparison, + const std::vector& records) { + for (TermFrequencyAdjustment& adjustment : adjustments_) { + if (!adjustment.fuzzy || adjustment.comparison != comparison) continue; + const BallMassTable& table = *adjustment.ball; + const uint64_t count = + adjustment.level < records.size() ? records[adjustment.level] : 0; + // The bracket was taken from the table's rarest and commonest + // neighbourhoods, so a mass outside them would let a pair beat its own + // bound. Only a value the store never held can land outside. + const double mass = + std::clamp(table.MassOf(count), table.MinMass(adjustment.level), + table.MaxMass(adjustment.level)); + adjustment.pinned_row = row; + adjustment.pinned_mass = mass; + } +} + +void Scorer::UnpinMasses() { + for (TermFrequencyAdjustment& adjustment : adjustments_) { + adjustment.pinned_row = std::numeric_limits::max(); + adjustment.pinned_mass = 0.0; + } +} + uint32_t Scorer::FrequencyFor(size_t comparison, uint32_t gamma, uint64_t row) const { for (const TermFrequencyAdjustment& adjustment : adjustments_) { if (adjustment.comparison != comparison || adjustment.fuzzy) continue; diff --git a/src/cpplink/score.hpp b/src/cpplink/score.hpp index 2f8f4ed..5333b8a 100644 --- a/src/cpplink/score.hpp +++ b/src/cpplink/score.hpp @@ -4,6 +4,7 @@ #pragma once #include +#include #include #include @@ -58,6 +59,11 @@ struct TermFrequencyAdjustment { const DateColumn* dates = nullptr; const BooleanColumn* booleans = nullptr; const BallMassTable* ball = nullptr; + // The one row whose mass is given rather than looked up. A search query is a + // row the table was built without, and its value may be one the dictionary + // never held, so the searcher measures its neighbourhood from its own walk. + uint64_t pinned_row = std::numeric_limits::max(); + double pinned_mass = 0.0; // log2(u / p) for the pair, damped: p is the shared value's frequency on an // exact level and the geometric mean of the two neighbourhood masses on a @@ -129,6 +135,14 @@ class Scorer { // Levels of this comparison that carry an adjustment, for the report that has // to say which ones did. bool AdjustsFuzzyLevels() const; + bool AdjustsFuzzyLevels(size_t comparison) const; + // Gives `row` its own neighbourhood mass on each fuzzy level of one + // comparison: `records[level]` is how many records land on that level + // against it. The mass is clamped into the table's range so the bracket + // stays admissible without being rebuilt; a value the dictionary holds is + // inside it already and scores exactly what `predict` would give it. + void PinMass(uint64_t row, size_t comparison, const std::vector& records); + void UnpinMasses(); // The fitted two-way corrections in force, for the waterfall that has to show // them: without these rows it would explain a different sum from the one that diff --git a/src/cpplink/search.cpp b/src/cpplink/search.cpp index ee4d561..c814187 100644 --- a/src/cpplink/search.cpp +++ b/src/cpplink/search.cpp @@ -13,6 +13,7 @@ #include #include #include +#include #include #include #include @@ -203,6 +204,27 @@ const std::string* QueryRecord::Find(const std::string& column) const { return nullptr; } +std::vector ComparisonsNamedBy(const Schema& schema, const QueryRecord& query) { + std::unordered_set named; + for (const QueryField& field : query.fields) named.insert(field.column); + for (bool grew = true; grew;) { + grew = false; + for (const ColumnSpec& column : schema.columns) { + if (column.IsDerived() && named.count(column.derive.from) != 0 && + named.insert(column.name).second) { + grew = true; + } + } + } + std::vector wanted(schema.comparisons.size(), false); + for (size_t c = 0; c < schema.comparisons.size(); ++c) { + for (const std::string& column : schema.comparisons[c].columns) { + if (named.count(column) != 0) wanted[c] = true; + } + } + return wanted; +} + bool ParseQueryField(const std::string& text, QueryField* field, std::string* error) { const size_t split = text.find('='); if (split == std::string::npos || split == 0) { @@ -225,7 +247,7 @@ double PriorWeightForExpected(double expected, uint64_t records) { Searcher::~Searcher() { Uninstall(); } -bool Searcher::Bind(RecordStore* store, ComparisonSet* comparisons, const Scorer* scorer, +bool Searcher::Bind(RecordStore* store, ComparisonSet* comparisons, Scorer* scorer, std::string* error) { if (store->NumDatasets() > 1) { *error = @@ -274,6 +296,7 @@ void Searcher::Uninstall() { ids.text.resize(ids.offsets.back()); } comparisons_->ResizeTables(); + scorer_->UnpinMasses(); plans_.clear(); written_ = 0; wrote_id_ = false; @@ -489,6 +512,18 @@ void Searcher::BuildPlans(SearchReport* report) { }); report->values_walked += values; ++report->tabled; + if (scorer_->AdjustsFuzzyLevels(c)) { + // The query's neighbourhood on each level is every record whose value + // the walk just put there, which is what the ball table holds for a + // value it was built over and the only way to have it for one it was + // not. The query row itself was never counted into `tf`. + std::vector landed(bound.spec->levels.size(), 0); + const uint32_t counted = static_cast(strings.tf.size()); + for (uint32_t v = 0; v < std::min(values, counted); ++v) { + landed[plan.table[v]] += strings.tf[v]; + } + scorer_->PinMass(query_row_, c, landed); + } } } diff --git a/src/cpplink/search.hpp b/src/cpplink/search.hpp index 091c67d..2bc4c51 100644 --- a/src/cpplink/search.hpp +++ b/src/cpplink/search.hpp @@ -54,6 +54,12 @@ struct QueryRecord { const std::string* Find(const std::string& column) const; }; +// Which comparisons read a column the query gives a value for, directly or as the +// source of a column derived from it. The rest land on their null level for every +// row, so nothing measured over their dictionaries can move a score: this is what +// `--fuzzy-tf` builds neighbourhood tables for. +std::vector ComparisonsNamedBy(const Schema& schema, const QueryRecord& query); + // Parses `column=value`, which is how the command line takes one field. bool ParseQueryField(const std::string& text, QueryField* field, std::string* error); @@ -145,11 +151,13 @@ class Searcher { Searcher& operator=(const Searcher&) = delete; // The store and the comparison set are written to, so both are taken - // mutable; the scorer is not. Fails where the store holds more than one + // mutable. So is the scorer, for one thing only: under fuzzy term frequency + // the query row's neighbourhood mass is measured by the walk and pinned on it + // for as long as the query is installed. Fails where the store holds more than one // input: a query row would have to be a dataset of its own, and the levels // that read which input a row came from were bound against a boundary list // that does not know about it. - bool Bind(RecordStore* store, ComparisonSet* comparisons, const Scorer* scorer, + bool Bind(RecordStore* store, ComparisonSet* comparisons, Scorer* scorer, std::string* error); bool Search(const QueryRecord& query, const SearchOptions& options, @@ -193,7 +201,7 @@ class Searcher { RecordStore* store_ = nullptr; ComparisonSet* comparisons_ = nullptr; - const Scorer* scorer_ = nullptr; + Scorer* scorer_ = nullptr; uint64_t query_row_ = 0; bool installed_ = false; size_t adopted_ = 0; diff --git a/tests/search_test.cpp b/tests/search_test.cpp index 095de22..f782318 100644 --- a/tests/search_test.cpp +++ b/tests/search_test.cpp @@ -15,6 +15,7 @@ #include "cpplink/comparison.hpp" #include "cpplink/explain.hpp" #include "cpplink/model.hpp" +#include "cpplink/neighbourhood.hpp" #include "cpplink/record_store.hpp" #include "cpplink/schema.hpp" #include "cpplink/score.hpp" @@ -466,6 +467,99 @@ TEST_F(SearchFixture, AThresholdCanRefuseEveryRow) { EXPECT_NE(text.str().find("no record here is this one"), std::string::npos); } +// Under fuzzy term frequency a fuzzy level's move reads the neighbourhood mass +// of both sides, and the query's side is measured by the walk rather than looked +// up. For a value the store holds that has to be the table's own mass, or a +// search would score a pair differently from `predict`. +TEST_F(SearchFixture, FuzzyTermFrequencyScoresAStoredValueAsPredictWould) { + cpplink::BallTables balls; + balls.Build(comparisons_, kRecords, cpplink::BallOptions{}); + ASSERT_TRUE(balls.Has(0)); + std::string error; + cpplink::Scorer fuzzy; + ASSERT_TRUE(fuzzy.Bind(model_, comparisons_, *store_, cpplink::ScoreOptions{}, &error, + &balls)) + << error; + ASSERT_TRUE(fuzzy.AdjustsFuzzyLevels(0)); + // The same model with nothing pinned on it, which is what `predict` runs. + cpplink::Scorer reference; + ASSERT_TRUE(reference.Bind(model_, comparisons_, *store_, cpplink::ScoreOptions{}, + &error, &balls)) + << error; + + // A surname rows carry: one interned but held by no row has no mass in the + // table, which is the adopted case below rather than this one. + const auto& surname = std::get(store_->column(0)); + ASSERT_GT(surname.tf[surname.ids[2]], 0u); + cpplink::QueryRecord query; + query.Set("surname", std::string(surname.dict.Value(surname.ids[2]))); + query.Set("city", "city2"); + query.Set("dob", "1981-01-24"); + + cpplink::Searcher searcher; + ASSERT_TRUE(searcher.Bind(store_.get(), &comparisons_, &fuzzy, &error)) << error; + cpplink::SearchOptions options; + options.k = kRecords; + cpplink::SearchReport report; + ASSERT_TRUE(searcher.Search(query, options, &report, &error)) << error; + EXPECT_EQ(report.values_adopted, 0u); + ASSERT_EQ(report.hits.size(), kRecords); + size_t moved = 0; + for (const cpplink::SearchHit& hit : report.hits) { + const double expected = reference.Weight(hit.gamma, hit.row, searcher.QueryRow()); + EXPECT_DOUBLE_EQ(hit.weight, expected) << "row " << hit.row; + if (comparisons_.LevelOf(hit.gamma, 0) == 2 && + fuzzy.AdjustmentFor(0, hit.gamma, hit.row, searcher.QueryRow()) != 0.0) { + ++moved; + } + } + EXPECT_GT(moved, 0u); +} + +// And for a value the store never held, whose id is past the end of the table: +// the walk is the only place its neighbourhood can come from, and the answer is +// still the top k of scoring every row. +TEST_F(SearchFixture, FuzzyTermFrequencyMeasuresAnAdoptedValue) { + cpplink::BallTables balls; + balls.Build(comparisons_, kRecords, cpplink::BallOptions{}); + std::string error; + cpplink::Scorer fuzzy; + ASSERT_TRUE(fuzzy.Bind(model_, comparisons_, *store_, cpplink::ScoreOptions{}, &error, + &balls)) + << error; + + cpplink::QueryRecord query; + query.Set("surname", "sander3x"); + query.Set("city", "city2"); + cpplink::Searcher searcher; + ASSERT_TRUE(searcher.Bind(store_.get(), &comparisons_, &fuzzy, &error)) << error; + cpplink::SearchOptions options; + options.k = 12; + cpplink::SearchReport report; + ASSERT_TRUE(searcher.Search(query, options, &report, &error)) << error; + ASSERT_EQ(report.values_adopted, 1u); + + std::vector all; + for (uint64_t row = 0; row < kRecords; ++row) { + Scored one; + one.row = row; + one.gamma = comparisons_.Evaluate(row, searcher.QueryRow()); + one.weight = fuzzy.Weight(one.gamma, row, searcher.QueryRow()); + all.push_back(one); + } + std::sort(all.begin(), all.end(), BetterThan); + ASSERT_EQ(report.hits.size(), options.k); + for (size_t i = 0; i < report.hits.size(); ++i) { + EXPECT_EQ(report.hits[i].row, all[i].row) << "at rank " << i; + EXPECT_DOUBLE_EQ(report.hits[i].weight, all[i].weight) << "at rank " << i; + } + // The top hit lands on the fuzzy level and is moved by it: without the pin + // the query side would read no mass and the adjustment would be zero. + const cpplink::SearchHit& top = report.hits.front(); + ASSERT_EQ(comparisons_.LevelOf(top.gamma, 0), 2u); + EXPECT_NE(fuzzy.AdjustmentFor(0, top.gamma, top.row, searcher.QueryRow()), 0.0); +} + TEST_F(SearchFixture, ADateThatIsNotOneIsRefused) { cpplink::QueryRecord query; query.Set("dob", "24/01/1981"); From e4c5745a19360841b1393f08cd43a20b83fb9ab9 Mon Sep 17 00:00:00 2001 From: 4ment Date: Sun, 27 Sep 2026 13:10:45 +1000 Subject: [PATCH 5/5] Fix two portability bugs found on Linux and Windows search_test.cpp used std::log2 without including , which libc++ provides transitively and libstdc++ does not. Check(f(&error), "stage: " + error) left the order of its two arguments to the compiler, and MSVC built the message before f had written the error, so the Python binding raised 'search: ' with no reason. A Check overload taking the prefix separately defers the concatenation until after the call; the five search, rescore and simplify sites use it. --- python/bindings/common.hpp | 8 ++++++++ python/bindings/diagnostics.cpp | 2 +- python/bindings/search.cpp | 6 +++--- python/bindings/stages.cpp | 2 +- tests/search_test.cpp | 1 + 5 files changed, 14 insertions(+), 5 deletions(-) diff --git a/python/bindings/common.hpp b/python/bindings/common.hpp index 026d8b1..00205d5 100644 --- a/python/bindings/common.hpp +++ b/python/bindings/common.hpp @@ -40,6 +40,14 @@ inline void Check(bool ok, const std::string& error) { if (!ok) throw Error(error); } +// The same with a prefix. `Check(f(&error), "stage: " + error)` is wrong: the +// order the two arguments are evaluated in is unspecified, and MSVC builds the +// message before `f` has written `error`. Taking both by reference defers the +// concatenation until after the call. +inline void Check(bool ok, const char* prefix, const std::string& error) { + if (!ok) throw Error(prefix + error); +} + // Runs one of the core's `Print*(..., std::ostream&)` into a string. Every // report's `text` and `__repr__` is the table the command line prints, produced // by the same function, so the two cannot disagree. diff --git a/python/bindings/diagnostics.cpp b/python/bindings/diagnostics.cpp index e8b1d2c..4061048 100644 --- a/python/bindings/diagnostics.cpp +++ b/python/bindings/diagnostics.cpp @@ -66,7 +66,7 @@ SimplifyReport SimplifyStage(Session& session, const Model& model, "simplify: alpha must be between 0 and 1"); const ComparisonSet& comparisons = session.Comparisons(); std::string error; - Check(ModelMatches(model, comparisons, &error), "simplify: " + error); + Check(ModelMatches(model, comparisons, &error), "simplify: ", error); const BlockingPlan& plan = session.Plan(/*for_estimation=*/false); py::gil_scoped_release release; return BuildSimplify(session.store(), comparisons, plan, model, options); diff --git a/python/bindings/search.cpp b/python/bindings/search.cpp index de14325..6fcee21 100644 --- a/python/bindings/search.cpp +++ b/python/bindings/search.cpp @@ -73,7 +73,7 @@ SearchOutcome SearchStage(Session& session, const py::dict& fields, Scorer scorer; std::string error; Check(scorer.Bind(model, *comparisons, *store, score, &error, balls), - "search: " + error); + "search: ", error); SearchOptions options; options.k = k; @@ -95,7 +95,7 @@ SearchOutcome SearchStage(Session& session, const py::dict& fields, SearchOutcome outcome; Searcher searcher; - Check(searcher.Bind(store, comparisons, &scorer, &error), "search: " + error); + Check(searcher.Bind(store, comparisons, &scorer, &error), "search: ", error); { py::gil_scoped_release release; if (!searcher.Search(query, options, &outcome.report, &error)) { @@ -119,7 +119,7 @@ SearchOutcome SearchStage(Session& session, const py::dict& fields, Check(error.empty(), error); if (!clusters.empty()) { Check(GroupHitsByCluster(*store, clusters, &outcome.report, &error), - "search: " + error); + "search: ", error); } outcome.text = CaptureText([&](std::ostream& out) { PrintSearchReport(outcome.report, out); }); diff --git a/python/bindings/stages.cpp b/python/bindings/stages.cpp index 03c5a4d..fdc6ff0 100644 --- a/python/bindings/stages.cpp +++ b/python/bindings/stages.cpp @@ -213,7 +213,7 @@ RescoreOutcome RescoreStage(const py::object& self, const Model& model, const ComparisonSet& comparisons = session.Comparisons(); Scorer scorer; Check(scorer.Bind(model, comparisons, session.store(), score, &error), - "rescore: " + error); + "rescore: ", error); RescoreOutcome result; bool ok = false; { diff --git a/tests/search_test.cpp b/tests/search_test.cpp index f782318..3316192 100644 --- a/tests/search_test.cpp +++ b/tests/search_test.cpp @@ -4,6 +4,7 @@ #include "cpplink/search.hpp" #include +#include #include #include #include