diff --git a/README.md b/README.md index 7dc78ce..dd7b157 100644 --- a/README.md +++ b/README.md @@ -300,6 +300,34 @@ python training/train_retriever.py \ Do not publish the demo-derived checkpoint as a research result. Build workflow-disjoint train, calibration, and hidden test splits first. +## Paper experiments + +The CLINC150 and BANKING77 results reported in the FlowRoute paper are +reproduced by `experiments/run_benchmark.py`. The script downloads both +datasets from their official repositories, verifies their SHA-256 digests, +pins the `all-MiniLM-L6-v2` model revision, keeps the test splits out of every +tuning decision, and records complete per-seed results. + +```bash +cd experiments +python -m venv .venv +.venv/bin/pip install --index-url https://download.pytorch.org/whl/cpu torch==2.8.0 +.venv/bin/pip install -r requirements.txt +HF_HOME=cache .venv/bin/python run_benchmark.py +``` + +`experiments/results/benchmark_results.json` is the record the paper's tables +are derived from: per-seed metrics, selected thresholds, model revision, source +URLs, checksums, and the environment. `experiments/results/ranking_summary.csv` +is a compact view of the ranking table. + +Note that the scoring rule benchmarked in the paper is not the one used by +`flowroute.backends.HuggingFaceRetriever`. The paper embeds each workflow +description and each approved example separately and takes the maximum example +similarity; the shipped retriever embeds one concatenated capability string per +contract. The two are not interchangeable, and the paper's numbers correspond +to the former. + ## Repository map ```text diff --git a/experiments/README.md b/experiments/README.md new file mode 100644 index 0000000..c77781b --- /dev/null +++ b/experiments/README.md @@ -0,0 +1,37 @@ +# FlowRoute paper benchmark + +This directory reproduces every empirical number in the FlowRoute paper. The +benchmark downloads the official CLINC150 and BANKING77 files, verifies their +SHA-256 digests, and evaluates the pinned `all-MiniLM-L6-v2` model revision. + +The public test splits are never used to select the description/example mixing +weight or abstention thresholds. CLINC150 uses its official validation split +and combines its official out-of-scope training and validation sets for +threshold selection. Both calibration constraints use the upper endpoints of +two-sided 95% Wilson intervals rather than point estimates. +Because BANKING77 has no official validation split, the script deterministically +reserves 20 training examples per intent before selecting workflow examples. + +## Run + +Python 3.12 was used for the paper. Create an environment and install the pinned +dependencies: + +```bash +python -m venv .venv +.venv/bin/pip install --index-url https://download.pytorch.org/whl/cpu torch==2.8.0 +.venv/bin/pip install -r requirements.txt +HF_HOME=cache .venv/bin/python run_benchmark.py +``` + +The script writes: + +- `results/benchmark_results.json`: complete per-seed metrics, thresholds, + environment, model revision, source URLs, and checksums. +- `results/ranking_summary.csv`: compact ranking table. +- `environment-lock.txt`: the complete package-version record from the + reported run. It is an environment record, not a cross-platform installer. + +Downloaded dataset files and model cache files are intentionally excluded from +the source bundle. The original licenses apply: CC BY 3.0 for CLINC150, CC BY +4.0 for BANKING77, and Apache 2.0 for `all-MiniLM-L6-v2`. diff --git a/experiments/environment-lock.txt b/experiments/environment-lock.txt new file mode 100644 index 0000000..fc61cb3 --- /dev/null +++ b/experiments/environment-lock.txt @@ -0,0 +1,34 @@ +# Exact `pip freeze` record from the reported Linux x86_64 benchmark run. +Jinja2==3.1.6 +MarkupSafe==3.0.3 +PyYAML==6.0.3 +certifi==2026.7.22 +charset-normalizer==3.5.1 +cloudpickle==3.1.2 +filelock==3.32.3 +fsspec==2026.7.0 +hf-xet==1.6.0 +huggingface_hub==0.36.2 +idna==3.19 +joblib==1.6.0 +mpmath==1.3.0 +narwhals==2.26.0 +networkx==3.6.1 +numpy==2.5.3 +packaging==26.3 +pillow==12.3.0 +regex==2026.9.10 +requests==2.34.2 +safetensors==0.8.0 +scikit-learn==1.9.1 +scipy==1.18.1 +sentence-transformers==5.1.2 +setuptools==78.1.0 +sympy==1.14.0 +threadpoolctl==3.6.0 +tokenizers==0.22.2 +torch==2.8.0+cpu +tqdm==4.70.1 +transformers==4.57.6 +typing_extensions==4.16.0 +urllib3==2.7.0 diff --git a/experiments/requirements.txt b/experiments/requirements.txt new file mode 100644 index 0000000..58f8851 --- /dev/null +++ b/experiments/requirements.txt @@ -0,0 +1,4 @@ +numpy==2.5.3 +scikit-learn==1.9.1 +sentence-transformers==5.1.2 +torch==2.8.0 diff --git a/experiments/results/benchmark_results.json b/experiments/results/benchmark_results.json new file mode 100644 index 0000000..49da9e2 --- /dev/null +++ b/experiments/results/benchmark_results.json @@ -0,0 +1,805 @@ +{ + "data": { + "banking77_test.csv": { + "sha256": "d12d6e3bc4c3103966ae786dc435913c0c563dfa328f5a3646d0e62cfeeb474d", + "url": "https://raw.githubusercontent.com/PolyAI-LDN/task-specific-datasets/master/banking_data/test.csv" + }, + "banking77_train.csv": { + "sha256": "b06e26ac675513959a63135f11b94ea7786ed02da65db93a5650d8838cbc664b", + "url": "https://raw.githubusercontent.com/PolyAI-LDN/task-specific-datasets/master/banking_data/train.csv" + }, + "clinc150_full.json": { + "sha256": "36923c3705a59e08fe9c3883d8bc2dd966ef93e22cb78ac41171782a698d56e0", + "url": "https://raw.githubusercontent.com/clinc/oos-eval/master/data/data_full.json" + } + }, + "datasets": [ + { + "ablation": { + "0": { + "alpha": { + "mean": 1.0, + "sd": 0.0 + }, + "recall_at_5": { + "mean": 0.8846666666666667, + "sd": 0.0 + }, + "top1": { + "mean": 0.6537777777777778, + "sd": 0.0 + } + }, + "1": { + "alpha": { + "mean": 0.58, + "sd": 0.04472135954999579 + }, + "recall_at_5": { + "mean": 0.9428444444444445, + "sd": 0.006145157684141472 + }, + "top1": { + "mean": 0.7440888888888889, + "sd": 0.013314431045826693 + } + }, + "10": { + "alpha": { + "mean": 0.32, + "sd": 0.08366600265340757 + }, + "recall_at_5": { + "mean": 0.9799555555555555, + "sd": 0.0013462338871876591 + }, + "top1": { + "mean": 0.8501333333333333, + "sd": 0.0045259198642136735 + } + }, + "3": { + "alpha": { + "mean": 0.38, + "sd": 0.08366600265340757 + }, + "recall_at_5": { + "mean": 0.9665777777777779, + "sd": 0.004514995590826805 + }, + "top1": { + "mean": 0.7989333333333334, + "sd": 0.005271399654834336 + } + }, + "5": { + "alpha": { + "mean": 0.38, + "sd": 0.08366600265340757 + }, + "recall_at_5": { + "mean": 0.9732, + "sd": 0.0031055505863605594 + }, + "top1": { + "mean": 0.818888888888889, + "sd": 0.007573508083534212 + } + } + }, + "dataset": "CLINC150", + "labels": 150, + "lexical_10_examples": { + "recall_at_5": { + "mean": 0.9475111111111112, + "sd": 0.004313057979647678 + }, + "top1": { + "mean": 0.8070666666666668, + "sd": 0.003549995652927268 + } + }, + "ood_test_examples": 1000, + "ood_validation_examples": 200, + "reference_confusions": { + "examples_per_workflow": 10, + "pairs": [ + { + "count": 15, + "gold": "time", + "predicted": "timezone" + }, + { + "count": 15, + "gold": "todo_list_update", + "predicted": "todo_list" + }, + { + "count": 13, + "gold": "shopping_list_update", + "predicted": "shopping_list" + }, + { + "count": 13, + "gold": "reminder_update", + "predicted": "reminder" + }, + { + "count": 12, + "gold": "change_user_name", + "predicted": "change_ai_name" + }, + { + "count": 10, + "gold": "change_user_name", + "predicted": "what_is_your_name" + }, + { + "count": 10, + "gold": "user_name", + "predicted": "what_is_your_name" + }, + { + "count": 10, + "gold": "share_location", + "predicted": "current_location" + } + ], + "seed": 37 + }, + "seed_results": [ + { + "examples": { + "0": { + "alpha": 1.0, + "recall_at_5": 0.8846666666666667, + "top1": 0.6537777777777778 + }, + "1": { + "alpha": 0.6, + "recall_at_5": 0.9444444444444444, + "top1": 0.7533333333333333 + }, + "10": { + "alpha": 0.2, + "recall_at_5": 0.9808888888888889, + "top1": 0.8557777777777777 + }, + "3": { + "alpha": 0.4, + "recall_at_5": 0.9684444444444444, + "top1": 0.8053333333333333 + }, + "5": { + "alpha": 0.3, + "recall_at_5": 0.9746666666666667, + "top1": 0.8295555555555556 + } + }, + "seed": 11 + }, + { + "examples": { + "0": { + "alpha": 1.0, + "recall_at_5": 0.8846666666666667, + "top1": 0.6537777777777778 + }, + "1": { + "alpha": 0.6, + "recall_at_5": 0.9324444444444444, + "top1": 0.7282222222222222 + }, + "10": { + "alpha": 0.4, + "recall_at_5": 0.9806666666666667, + "top1": 0.8486666666666667 + }, + "3": { + "alpha": 0.4, + "recall_at_5": 0.964, + "top1": 0.798 + }, + "5": { + "alpha": 0.5, + "recall_at_5": 0.9686666666666667, + "top1": 0.8095555555555556 + } + }, + "seed": 23 + }, + { + "examples": { + "0": { + "alpha": 1.0, + "recall_at_5": 0.8846666666666667, + "top1": 0.6537777777777778 + }, + "1": { + "alpha": 0.6, + "recall_at_5": 0.9462222222222222, + "top1": 0.7615555555555555 + }, + "10": { + "alpha": 0.4, + "recall_at_5": 0.978, + "top1": 0.8455555555555555 + }, + "3": { + "alpha": 0.5, + "recall_at_5": 0.9693333333333334, + "top1": 0.8024444444444444 + }, + "5": { + "alpha": 0.4, + "recall_at_5": 0.976, + "top1": 0.8171111111111111 + } + }, + "seed": 37 + }, + { + "examples": { + "0": { + "alpha": 1.0, + "recall_at_5": 0.8846666666666667, + "top1": 0.6537777777777778 + }, + "1": { + "alpha": 0.6, + "recall_at_5": 0.9482222222222222, + "top1": 0.7366666666666667 + }, + "10": { + "alpha": 0.3, + "recall_at_5": 0.9791111111111112, + "top1": 0.854 + }, + "3": { + "alpha": 0.3, + "recall_at_5": 0.9711111111111111, + "top1": 0.7973333333333333 + }, + "5": { + "alpha": 0.3, + "recall_at_5": 0.9753333333333334, + "top1": 0.8226666666666667 + } + }, + "seed": 53 + }, + { + "examples": { + "0": { + "alpha": 1.0, + "recall_at_5": 0.8846666666666667, + "top1": 0.6537777777777778 + }, + "1": { + "alpha": 0.5, + "recall_at_5": 0.9428888888888889, + "top1": 0.7406666666666667 + }, + "10": { + "alpha": 0.3, + "recall_at_5": 0.9811111111111112, + "top1": 0.8466666666666667 + }, + "3": { + "alpha": 0.3, + "recall_at_5": 0.96, + "top1": 0.7915555555555556 + }, + "5": { + "alpha": 0.4, + "recall_at_5": 0.9713333333333334, + "top1": 0.8155555555555556 + } + }, + "seed": 71 + } + ], + "seeds": [ + 11, + 23, + 37, + 53, + 71 + ], + "selective_routing": { + "seed_results": [ + { + "alpha": 0.2, + "margin_threshold": 0.055, + "score_threshold": 0.535, + "seed": 11, + "test": { + "accepted_id": 3047, + "accepted_ood": 35, + "id_coverage": 0.6771111111111111, + "ood_fpr": 0.035, + "selective_error": 0.05382343288480473, + "wrong_accepted_id": 164 + }, + "validation": { + "accepted_id": 1988, + "accepted_ood": 3, + "id_coverage": 0.6626666666666666, + "ood_fpr": 0.015, + "ood_fpr_wilson_upper_95": 0.04316572879269026, + "selective_error": 0.039235412474849095, + "selective_error_wilson_upper_95": 0.048696674268041445, + "wrong_accepted_id": 78 + } + }, + { + "alpha": 0.4, + "margin_threshold": 0.065, + "score_threshold": 0.47000000000000003, + "seed": 23, + "test": { + "accepted_id": 2909, + "accepted_ood": 37, + "id_coverage": 0.6464444444444445, + "ood_fpr": 0.037, + "selective_error": 0.04090752836026126, + "wrong_accepted_id": 119 + }, + "validation": { + "accepted_id": 1855, + "accepted_ood": 3, + "id_coverage": 0.6183333333333333, + "ood_fpr": 0.015, + "ood_fpr_wilson_upper_95": 0.04316572879269026, + "selective_error": 0.03827493261455526, + "selective_error_wilson_upper_95": 0.04800303825871906, + "wrong_accepted_id": 71 + } + }, + { + "alpha": 0.4, + "margin_threshold": 0.04, + "score_threshold": 0.505, + "seed": 37, + "test": { + "accepted_id": 2871, + "accepted_ood": 26, + "id_coverage": 0.638, + "ood_fpr": 0.026, + "selective_error": 0.054684778822709855, + "wrong_accepted_id": 157 + }, + "validation": { + "accepted_id": 1925, + "accepted_ood": 3, + "id_coverage": 0.6416666666666667, + "ood_fpr": 0.015, + "ood_fpr_wilson_upper_95": 0.04316572879269026, + "selective_error": 0.03896103896103896, + "selective_error_wilson_upper_95": 0.048563379564802375, + "wrong_accepted_id": 75 + } + }, + { + "alpha": 0.3, + "margin_threshold": 0.05, + "score_threshold": 0.53, + "seed": 53, + "test": { + "accepted_id": 2889, + "accepted_ood": 23, + "id_coverage": 0.642, + "ood_fpr": 0.023, + "selective_error": 0.04915195569401177, + "wrong_accepted_id": 142 + }, + "validation": { + "accepted_id": 1912, + "accepted_ood": 2, + "id_coverage": 0.6373333333333333, + "ood_fpr": 0.01, + "ood_fpr_wilson_upper_95": 0.0357217617161768, + "selective_error": 0.0397489539748954, + "selective_error_wilson_upper_95": 0.049468642990548276, + "wrong_accepted_id": 76 + } + }, + { + "alpha": 0.3, + "margin_threshold": 0.085, + "score_threshold": 0.505, + "seed": 71, + "test": { + "accepted_id": 2687, + "accepted_ood": 25, + "id_coverage": 0.5971111111111111, + "ood_fpr": 0.025, + "selective_error": 0.03126163007071083, + "wrong_accepted_id": 84 + }, + "validation": { + "accepted_id": 1771, + "accepted_ood": 3, + "id_coverage": 0.5903333333333334, + "ood_fpr": 0.015, + "ood_fpr_wilson_upper_95": 0.04316572879269026, + "selective_error": 0.024844720496894408, + "selective_error_wilson_upper_95": 0.03318720358989378, + "wrong_accepted_id": 44 + } + } + ], + "target_validation_ood_fpr": 0.05, + "target_validation_selective_error": 0.05, + "test": { + "id_coverage": { + "mean": 0.6401333333333333, + "sd": 0.028575047389870285 + }, + "ood_fpr": { + "mean": 0.029199999999999997, + "sd": 0.0063403469936589435 + }, + "selective_error": { + "mean": 0.045965865166499684, + "sd": 0.009870578752526052 + } + } + }, + "test_examples": 4500, + "validation_examples": 3000 + }, + { + "ablation": { + "0": { + "alpha": { + "mean": 1.0, + "sd": 0.0 + }, + "recall_at_5": { + "mean": 0.8389610389610389, + "sd": 0.0 + }, + "top1": { + "mean": 0.6133116883116884, + "sd": 0.0 + } + }, + "1": { + "alpha": { + "mean": 0.6799999999999999, + "sd": 0.04472135954999579 + }, + "recall_at_5": { + "mean": 0.899155844155844, + "sd": 0.007574752561660528 + }, + "top1": { + "mean": 0.6756493506493506, + "sd": 0.008152519174993315 + } + }, + "10": { + "alpha": { + "mean": 0.12000000000000002, + "sd": 0.044721359549995794 + }, + "recall_at_5": { + "mean": 0.9687662337662338, + "sd": 0.001385112273227392 + }, + "top1": { + "mean": 0.831038961038961, + "sd": 0.0053625513823539125 + } + }, + "3": { + "alpha": { + "mean": 0.42000000000000004, + "sd": 0.08366600265340755 + }, + "recall_at_5": { + "mean": 0.9370129870129871, + "sd": 0.0032790600449227373 + }, + "top1": { + "mean": 0.7368181818181818, + "sd": 0.004791574237499559 + } + }, + "5": { + "alpha": { + "mean": 0.32, + "sd": 0.04472135954999581 + }, + "recall_at_5": { + "mean": 0.9527922077922077, + "sd": 0.0031023062819717675 + }, + "top1": { + "mean": 0.7834415584415584, + "sd": 0.006775531413866621 + } + } + }, + "dataset": "BANKING77", + "labels": 77, + "lexical_10_examples": { + "recall_at_5": { + "mean": 0.9344155844155845, + "sd": 0.003862135471873467 + }, + "top1": { + "mean": 0.7235064935064935, + "sd": 0.012326122422674071 + } + }, + "ood_test_examples": 0, + "ood_validation_examples": 0, + "reference_confusions": { + "examples_per_workflow": 10, + "pairs": [ + { + "count": 14, + "gold": "top_up_reverted", + "predicted": "pending_top_up" + }, + { + "count": 14, + "gold": "why_verify_identity", + "predicted": "verify_my_identity" + }, + { + "count": 10, + "gold": "exchange_via_app", + "predicted": "fiat_currency_support" + }, + { + "count": 8, + "gold": "top_up_by_bank_transfer_charge", + "predicted": "transfer_fee_charged" + }, + { + "count": 8, + "gold": "unable_to_verify_identity", + "predicted": "verify_my_identity" + }, + { + "count": 7, + "gold": "top_up_reverted", + "predicted": "top_up_failed" + }, + { + "count": 7, + "gold": "transfer_timing", + "predicted": "balance_not_updated_after_bank_transfer" + }, + { + "count": 7, + "gold": "beneficiary_not_allowed", + "predicted": "failed_transfer" + } + ], + "seed": 37 + }, + "seed_results": [ + { + "examples": { + "0": { + "alpha": 1.0, + "recall_at_5": 0.8389610389610389, + "top1": 0.6133116883116884 + }, + "1": { + "alpha": 0.7, + "recall_at_5": 0.8918831168831168, + "top1": 0.6743506493506494 + }, + "10": { + "alpha": 0.1, + "recall_at_5": 0.9701298701298702, + "top1": 0.8311688311688312 + }, + "3": { + "alpha": 0.3, + "recall_at_5": 0.9327922077922078, + "top1": 0.736038961038961 + }, + "5": { + "alpha": 0.4, + "recall_at_5": 0.950974025974026, + "top1": 0.7724025974025974 + } + }, + "seed": 11 + }, + { + "examples": { + "0": { + "alpha": 1.0, + "recall_at_5": 0.8389610389610389, + "top1": 0.6133116883116884 + }, + "1": { + "alpha": 0.7, + "recall_at_5": 0.8951298701298701, + "top1": 0.6814935064935065 + }, + "10": { + "alpha": 0.1, + "recall_at_5": 0.9672077922077922, + "top1": 0.8262987012987013 + }, + "3": { + "alpha": 0.5, + "recall_at_5": 0.936038961038961, + "top1": 0.7396103896103896 + }, + "5": { + "alpha": 0.3, + "recall_at_5": 0.9571428571428572, + "top1": 0.7827922077922078 + } + }, + "seed": 23 + }, + { + "examples": { + "0": { + "alpha": 1.0, + "recall_at_5": 0.8389610389610389, + "top1": 0.6133116883116884 + }, + "1": { + "alpha": 0.6, + "recall_at_5": 0.910064935064935, + "top1": 0.6746753246753247 + }, + "10": { + "alpha": 0.1, + "recall_at_5": 0.9675324675324676, + "top1": 0.8386363636363636 + }, + "3": { + "alpha": 0.4, + "recall_at_5": 0.9366883116883117, + "top1": 0.7435064935064936 + }, + "5": { + "alpha": 0.3, + "recall_at_5": 0.9525974025974026, + "top1": 0.788961038961039 + } + }, + "seed": 37 + }, + { + "examples": { + "0": { + "alpha": 1.0, + "recall_at_5": 0.8389610389610389, + "top1": 0.6133116883116884 + }, + "1": { + "alpha": 0.7, + "recall_at_5": 0.8948051948051948, + "top1": 0.6633116883116883 + }, + "10": { + "alpha": 0.1, + "recall_at_5": 0.9688311688311688, + "top1": 0.8334415584415584 + }, + "3": { + "alpha": 0.4, + "recall_at_5": 0.9376623376623376, + "top1": 0.7318181818181818 + }, + "5": { + "alpha": 0.3, + "recall_at_5": 0.9542207792207792, + "top1": 0.788961038961039 + } + }, + "seed": 53 + }, + { + "examples": { + "0": { + "alpha": 1.0, + "recall_at_5": 0.8389610389610389, + "top1": 0.6133116883116884 + }, + "1": { + "alpha": 0.7, + "recall_at_5": 0.9038961038961039, + "top1": 0.6844155844155844 + }, + "10": { + "alpha": 0.2, + "recall_at_5": 0.9701298701298702, + "top1": 0.8256493506493506 + }, + "3": { + "alpha": 0.5, + "recall_at_5": 0.9418831168831169, + "top1": 0.7331168831168832 + }, + "5": { + "alpha": 0.3, + "recall_at_5": 0.949025974025974, + "top1": 0.7840909090909091 + } + }, + "seed": 71 + } + ], + "seeds": [ + 11, + 23, + 37, + 53, + 71 + ], + "test_examples": 3080, + "validation_examples": 1540 + } + ], + "environment": { + "cpu": "AMD EPYC 9V74 80-Core Processor", + "numpy": "2.5.3", + "platform": "Linux-6.18.35-x86_64-with-glibc2.39", + "python": "3.12.14", + "scikit_learn": "1.9.1", + "sentence_transformers": "5.1.2", + "torch": "2.8.0+cpu" + }, + "evaluation_seconds_excluding_initial_query_encoding": 90.92143702507019, + "experiment": "FlowRoute public-data benchmark", + "latency": { + "batch_size": 1, + "catalog_workflows": 150, + "examples_per_workflow": 10, + "iterations": 500, + "mean_ms": 11.637435477999999, + "median_ms": 11.536568500000001, + "p95_ms": 13.47621255, + "p99_ms": 14.88310988, + "threads": 1 + }, + "method": { + "alpha_grid": [ + 0.0, + 0.1, + 0.2, + 0.3, + 0.4, + 0.5, + 0.6, + 0.7, + 0.8, + 0.9, + 1.0 + ], + "example_counts": [ + 0, + 1, + 3, + 5, + 10 + ], + "seeds": [ + 11, + 23, + 37, + 53, + 71 + ], + "selection": "alpha and route thresholds selected on validation data only; both validation risk constraints use upper endpoints of two-sided 95% Wilson intervals" + }, + "model": { + "id": "sentence-transformers/all-MiniLM-L6-v2", + "parameters": 22713216, + "revision": "1110a243fdf4706b3f48f1d95db1a4f5529b4d41" + } +} diff --git a/experiments/results/ranking_summary.csv b/experiments/results/ranking_summary.csv new file mode 100644 index 0000000..e5d467d --- /dev/null +++ b/experiments/results/ranking_summary.csv @@ -0,0 +1,13 @@ +dataset,system,top1_mean,top1_sd,recall_at_5_mean,recall_at_5_sd +CLINC150,word-char TF-IDF (10 examples),0.8070666666666668,0.003549995652927268,0.9475111111111112,0.004313057979647678 +CLINC150,dense example-aware (0 examples),0.6537777777777778,0.0,0.8846666666666667,0.0 +CLINC150,dense example-aware (1 examples),0.7440888888888889,0.013314431045826693,0.9428444444444445,0.006145157684141472 +CLINC150,dense example-aware (3 examples),0.7989333333333334,0.005271399654834336,0.9665777777777779,0.004514995590826805 +CLINC150,dense example-aware (5 examples),0.818888888888889,0.007573508083534212,0.9732,0.0031055505863605594 +CLINC150,dense example-aware (10 examples),0.8501333333333333,0.0045259198642136735,0.9799555555555555,0.0013462338871876591 +BANKING77,word-char TF-IDF (10 examples),0.7235064935064935,0.012326122422674071,0.9344155844155845,0.003862135471873467 +BANKING77,dense example-aware (0 examples),0.6133116883116884,0.0,0.8389610389610389,0.0 +BANKING77,dense example-aware (1 examples),0.6756493506493506,0.008152519174993315,0.899155844155844,0.007574752561660528 +BANKING77,dense example-aware (3 examples),0.7368181818181818,0.004791574237499559,0.9370129870129871,0.0032790600449227373 +BANKING77,dense example-aware (5 examples),0.7834415584415584,0.006775531413866621,0.9527922077922077,0.0031023062819717675 +BANKING77,dense example-aware (10 examples),0.831038961038961,0.0053625513823539125,0.9687662337662338,0.001385112273227392 diff --git a/experiments/run_benchmark.py b/experiments/run_benchmark.py new file mode 100644 index 0000000..ceea183 --- /dev/null +++ b/experiments/run_benchmark.py @@ -0,0 +1,767 @@ +#!/usr/bin/env python3 +"""Reproduce the FlowRoute benchmark reported in the paper. + +The experiment treats each intent label as a runtime workflow. It compares a +lexical index, description-only dense routing, and example-aware dense routing. +No generative model is called and no benchmark test example is used to choose +hyperparameters or decision thresholds. +""" + +from __future__ import annotations + +import argparse +import csv +import hashlib +import json +import math +import os +import platform +import random +import statistics +import sys +import time +import urllib.request +from collections import Counter, defaultdict +from dataclasses import dataclass +from pathlib import Path +from typing import Any + +import numpy as np +import sklearn +import torch +from sentence_transformers import SentenceTransformer +from sklearn.feature_extraction.text import TfidfVectorizer + + +MODEL_ID = "sentence-transformers/all-MiniLM-L6-v2" +MODEL_REVISION = "1110a243fdf4706b3f48f1d95db1a4f5529b4d41" +SEEDS = (11, 23, 37, 53, 71) +EXAMPLE_COUNTS = (0, 1, 3, 5, 10) +ALPHAS = tuple(value / 10 for value in range(11)) +TARGET_SELECTIVE_ERROR = 0.05 +TARGET_OOD_FPR = 0.05 + +DATASETS = { + "clinc150_full.json": { + "url": "https://raw.githubusercontent.com/clinc/oos-eval/master/data/data_full.json", + "sha256": "36923c3705a59e08fe9c3883d8bc2dd966ef93e22cb78ac41171782a698d56e0", + }, + "banking77_train.csv": { + "url": "https://raw.githubusercontent.com/PolyAI-LDN/task-specific-datasets/master/banking_data/train.csv", + "sha256": "b06e26ac675513959a63135f11b94ea7786ed02da65db93a5650d8838cbc664b", + }, + "banking77_test.csv": { + "url": "https://raw.githubusercontent.com/PolyAI-LDN/task-specific-datasets/master/banking_data/test.csv", + "sha256": "d12d6e3bc4c3103966ae786dc435913c0c563dfa328f5a3646d0e62cfeeb474d", + }, +} + + +@dataclass(frozen=True) +class DatasetSplit: + name: str + train_by_label: dict[str, list[str]] + validation: list[tuple[str, str]] + test: list[tuple[str, str]] + ood_validation: list[str] + ood_test: list[str] + + +def sha256(path: Path) -> str: + digest = hashlib.sha256() + with path.open("rb") as handle: + for block in iter(lambda: handle.read(1024 * 1024), b""): + digest.update(block) + return digest.hexdigest() + + +def ensure_data(data_dir: Path) -> None: + data_dir.mkdir(parents=True, exist_ok=True) + for filename, metadata in DATASETS.items(): + target = data_dir / filename + if not target.exists(): + print(f"Downloading {filename}", file=sys.stderr) + urllib.request.urlretrieve(metadata["url"], target) + actual = sha256(target) + if actual != metadata["sha256"]: + raise RuntimeError( + f"Checksum mismatch for {target}: expected {metadata['sha256']}, got {actual}" + ) + + +def load_clinc(data_dir: Path) -> DatasetSplit: + raw = json.loads((data_dir / "clinc150_full.json").read_text(encoding="utf-8")) + train_by_label: dict[str, list[str]] = defaultdict(list) + for text, label in raw["train"]: + train_by_label[label].append(text) + return DatasetSplit( + name="CLINC150", + train_by_label=dict(train_by_label), + validation=[tuple(item) for item in raw["val"]], + test=[tuple(item) for item in raw["test"]], + ood_validation=[text for text, _ in [*raw["oos_train"], *raw["oos_val"]]], + ood_test=[text for text, _ in raw["oos_test"]], + ) + + +def load_banking(data_dir: Path) -> DatasetSplit: + train_by_label: dict[str, list[str]] = defaultdict(list) + with (data_dir / "banking77_train.csv").open( + encoding="utf-8", newline="" + ) as handle: + for row in csv.DictReader(handle): + train_by_label[row["category"]].append(row["text"]) + with (data_dir / "banking77_test.csv").open(encoding="utf-8", newline="") as handle: + test = [(row["text"], row["category"]) for row in csv.DictReader(handle)] + + validation: list[tuple[str, str]] = [] + exemplar_pool: dict[str, list[str]] = {} + for label, examples in sorted(train_by_label.items()): + label_seed = int(hashlib.sha256(label.encode("utf-8")).hexdigest()[:8], 16) + shuffled = list(examples) + random.Random(20260912 + label_seed).shuffle(shuffled) + validation.extend((text, label) for text in shuffled[:20]) + exemplar_pool[label] = shuffled[20:] + + return DatasetSplit( + name="BANKING77", + train_by_label=exemplar_pool, + validation=validation, + test=test, + ood_validation=[], + ood_test=[], + ) + + +def workflow_description(label: str) -> str: + return f"Workflow for {label.replace('_', ' ')}." + + +def select_examples( + train_by_label: dict[str, list[str]], labels: list[str], count: int, seed: int +) -> dict[str, list[str]]: + selected: dict[str, list[str]] = {} + for label in labels: + label_seed = int(hashlib.sha256(label.encode("utf-8")).hexdigest()[:8], 16) + values = list(train_by_label[label]) + random.Random(seed + label_seed).shuffle(values) + selected[label] = values[:count] + return selected + + +def dense_components( + model: SentenceTransformer, + query_embeddings: np.ndarray, + labels: list[str], + examples: dict[str, list[str]], +) -> tuple[np.ndarray, np.ndarray | None, np.ndarray, np.ndarray | None]: + descriptions = [workflow_description(label) for label in labels] + description_embeddings = model.encode( + descriptions, + batch_size=128, + normalize_embeddings=True, + convert_to_numpy=True, + show_progress_bar=False, + ) + description_scores = query_embeddings @ description_embeddings.T + count = len(examples[labels[0]]) if labels else 0 + if count == 0: + return description_scores, None, description_embeddings, None + + flat_examples = [text for label in labels for text in examples[label]] + example_embeddings = model.encode( + flat_examples, + batch_size=128, + normalize_embeddings=True, + convert_to_numpy=True, + show_progress_bar=False, + ).reshape(len(labels), count, -1) + example_scores = np.einsum("qd,lnd->qln", query_embeddings, example_embeddings).max( + axis=2 + ) + return ( + description_scores, + example_scores, + description_embeddings, + example_embeddings, + ) + + +def combine_scores( + description_scores: np.ndarray, example_scores: np.ndarray | None, alpha: float +) -> np.ndarray: + if example_scores is None: + return description_scores + return (alpha * description_scores) + ((1.0 - alpha) * example_scores) + + +def label_indices( + rows: list[tuple[str, str]], label_to_index: dict[str, int] +) -> np.ndarray: + return np.asarray([label_to_index[label] for _, label in rows], dtype=np.int64) + + +def ranking_metrics(scores: np.ndarray, gold: np.ndarray) -> dict[str, float]: + order = np.argsort(-scores, axis=1, kind="stable") + top1 = float(np.mean(order[:, 0] == gold)) + top5 = float( + np.mean(np.any(order[:, : min(5, scores.shape[1])] == gold[:, None], axis=1)) + ) + return {"top1": top1, "recall_at_5": top5} + + +def common_confusions( + scores: np.ndarray, gold: np.ndarray, labels: list[str], limit: int = 8 +) -> list[dict[str, Any]]: + predictions = np.argmax(scores, axis=1) + pairs = Counter( + (labels[int(expected)], labels[int(predicted)]) + for expected, predicted in zip(gold, predictions, strict=True) + if expected != predicted + ) + return [ + {"gold": gold_label, "predicted": predicted_label, "count": count} + for (gold_label, predicted_label), count in pairs.most_common(limit) + ] + + +def tune_alpha( + description_scores: np.ndarray, + example_scores: np.ndarray | None, + gold: np.ndarray, +) -> float: + if example_scores is None: + return 1.0 + candidates: list[tuple[float, float]] = [] + for alpha in ALPHAS: + metrics = ranking_metrics( + combine_scores(description_scores, example_scores, alpha), gold + ) + candidates.append((metrics["top1"], alpha)) + # Prefer more description weight when validation accuracy ties. + return max(candidates, key=lambda item: (item[0], item[1]))[1] + + +def lexical_scores( + queries: list[str], labels: list[str], examples: dict[str, list[str]] +) -> np.ndarray: + documents = [ + " ".join([workflow_description(label), *examples[label]]) for label in labels + ] + word = TfidfVectorizer( + lowercase=True, + strip_accents="unicode", + ngram_range=(1, 2), + sublinear_tf=True, + min_df=1, + ) + char = TfidfVectorizer( + lowercase=True, + strip_accents="unicode", + analyzer="char_wb", + ngram_range=(3, 5), + sublinear_tf=True, + min_df=1, + ) + word_documents = word.fit_transform(documents) + char_documents = char.fit_transform(documents) + word_scores = (word.transform(queries) @ word_documents.T).toarray() + char_scores = (char.transform(queries) @ char_documents.T).toarray() + return np.clip((0.65 * word_scores) + (0.35 * char_scores), 0.0, 1.0) + + +def top_score_and_margin( + scores: np.ndarray, +) -> tuple[np.ndarray, np.ndarray, np.ndarray]: + order = np.argsort(-scores, axis=1, kind="stable") + best_index = order[:, 0] + row_ids = np.arange(scores.shape[0]) + best_score = scores[row_ids, best_index] + second_score = ( + scores[row_ids, order[:, 1]] + if scores.shape[1] > 1 + else np.zeros_like(best_score) + ) + return best_index, best_score, best_score - second_score + + +def decision_metrics( + id_scores: np.ndarray, + id_gold: np.ndarray, + ood_scores: np.ndarray, + score_threshold: float, + margin_threshold: float, +) -> dict[str, float | int]: + prediction, best_score, margin = top_score_and_margin(id_scores) + accepted = (best_score >= score_threshold) & (margin >= margin_threshold) + accepted_count = int(accepted.sum()) + wrong_count = int((accepted & (prediction != id_gold)).sum()) + id_coverage = accepted_count / len(id_gold) + selective_error = wrong_count / accepted_count if accepted_count else 0.0 + + if len(ood_scores): + _, ood_best_score, ood_margin = top_score_and_margin(ood_scores) + ood_accepted = (ood_best_score >= score_threshold) & ( + ood_margin >= margin_threshold + ) + ood_false_positive_count = int(ood_accepted.sum()) + ood_fpr = ood_false_positive_count / len(ood_scores) + else: + ood_false_positive_count = 0 + ood_fpr = math.nan + + return { + "id_coverage": id_coverage, + "selective_error": selective_error, + "accepted_id": accepted_count, + "wrong_accepted_id": wrong_count, + "ood_fpr": ood_fpr, + "accepted_ood": ood_false_positive_count, + } + + +def tune_decision_thresholds( + id_scores: np.ndarray, + id_gold: np.ndarray, + ood_scores: np.ndarray, +) -> tuple[float, float, dict[str, float | int]]: + id_prediction, id_best_score, id_margin = top_score_and_margin(id_scores) + id_wrong = id_prediction != id_gold + _, ood_best_score, ood_margin = top_score_and_margin(ood_scores) + best: tuple[float, float, float, float, float, dict[str, float | int]] | None = None + for score_threshold in np.linspace(0.0, 1.0, 201): + for margin_threshold in np.linspace(0.0, 0.5, 101): + id_accepted = (id_best_score >= score_threshold) & ( + id_margin >= margin_threshold + ) + accepted_id = int(id_accepted.sum()) + wrong_accepted_id = int((id_accepted & id_wrong).sum()) + ood_accepted = (ood_best_score >= score_threshold) & ( + ood_margin >= margin_threshold + ) + accepted_ood = int(ood_accepted.sum()) + metrics: dict[str, float | int] = { + "id_coverage": accepted_id / len(id_gold), + "selective_error": wrong_accepted_id / accepted_id + if accepted_id + else 0.0, + "accepted_id": accepted_id, + "wrong_accepted_id": wrong_accepted_id, + "ood_fpr": accepted_ood / len(ood_scores), + "accepted_ood": accepted_ood, + "selective_error_wilson_upper_95": wilson_interval( + wrong_accepted_id, accepted_id + )[1], + "ood_fpr_wilson_upper_95": wilson_interval( + accepted_ood, len(ood_scores) + )[1], + } + if metrics["accepted_id"] < 100: + continue + if metrics["selective_error_wilson_upper_95"] > TARGET_SELECTIVE_ERROR: + continue + if metrics["ood_fpr_wilson_upper_95"] > TARGET_OOD_FPR: + continue + candidate = ( + float(metrics["id_coverage"]), + -float(metrics["selective_error"]), + -float(metrics["ood_fpr"]), + -float(score_threshold), + -float(margin_threshold), + metrics, + ) + if best is None or candidate[:5] > best[:5]: + best = candidate + if best is None: + raise RuntimeError("No threshold pair satisfied the validation constraints") + return -best[3], -best[4], best[5] + + +def mean_sd(values: list[float]) -> dict[str, float]: + return { + "mean": float(statistics.fmean(values)), + "sd": float(statistics.stdev(values)) if len(values) > 1 else 0.0, + } + + +def wilson_interval( + successes: int, total: int, z: float = 1.959963984540054 +) -> list[float]: + if total == 0: + return [math.nan, math.nan] + proportion = successes / total + denominator = 1 + (z * z / total) + centre = (proportion + (z * z / (2 * total))) / denominator + radius = ( + z + * math.sqrt( + (proportion * (1 - proportion) / total) + (z * z / (4 * total * total)) + ) + / denominator + ) + return [centre - radius, centre + radius] + + +def evaluate_dataset( + model: SentenceTransformer, + split: DatasetSplit, + all_query_embeddings: dict[str, np.ndarray], +) -> dict[str, Any]: + labels = sorted(split.train_by_label) + label_to_index = {label: index for index, label in enumerate(labels)} + test_texts = [text for text, _ in split.test] + validation_gold = label_indices(split.validation, label_to_index) + test_gold = label_indices(split.test, label_to_index) + validation_embeddings = all_query_embeddings[f"{split.name}:validation"] + test_embeddings = all_query_embeddings[f"{split.name}:test"] + + seed_results: list[dict[str, Any]] = [] + ablations: dict[int, list[dict[str, float]]] = { + count: [] for count in EXAMPLE_COUNTS + } + reference_confusions: list[dict[str, Any]] = [] + + for seed in SEEDS: + seed_record: dict[str, Any] = {"seed": seed, "examples": {}} + for count in EXAMPLE_COUNTS: + examples = select_examples(split.train_by_label, labels, count, seed) + val_desc, val_ex, desc_embeddings, example_embeddings = dense_components( + model, validation_embeddings, labels, examples + ) + alpha = tune_alpha(val_desc, val_ex, validation_gold) + test_desc = test_embeddings @ desc_embeddings.T + test_ex = ( + np.einsum("qd,lnd->qln", test_embeddings, example_embeddings).max( + axis=2 + ) + if example_embeddings is not None + else None + ) + test_scores = combine_scores(test_desc, test_ex, alpha) + metrics = ranking_metrics(test_scores, test_gold) + metrics["alpha"] = alpha + seed_record["examples"][str(count)] = metrics + ablations[count].append(metrics) + if seed == 37 and count == 10: + reference_confusions = common_confusions(test_scores, test_gold, labels) + seed_results.append(seed_record) + + aggregate_ablations: dict[str, Any] = {} + for count, records in ablations.items(): + aggregate_ablations[str(count)] = { + "top1": mean_sd([record["top1"] for record in records]), + "recall_at_5": mean_sd([record["recall_at_5"] for record in records]), + "alpha": mean_sd([record["alpha"] for record in records]), + } + + lexical_seed_metrics: list[dict[str, float]] = [] + for seed in SEEDS: + examples = select_examples(split.train_by_label, labels, 10, seed) + scores = lexical_scores(test_texts, labels, examples) + lexical_seed_metrics.append(ranking_metrics(scores, test_gold)) + + result: dict[str, Any] = { + "dataset": split.name, + "labels": len(labels), + "validation_examples": len(split.validation), + "test_examples": len(split.test), + "ood_validation_examples": len(split.ood_validation), + "ood_test_examples": len(split.ood_test), + "seeds": list(SEEDS), + "seed_results": seed_results, + "ablation": aggregate_ablations, + "lexical_10_examples": { + "top1": mean_sd([item["top1"] for item in lexical_seed_metrics]), + "recall_at_5": mean_sd( + [item["recall_at_5"] for item in lexical_seed_metrics] + ), + }, + "reference_confusions": { + "seed": 37, + "examples_per_workflow": 10, + "pairs": reference_confusions, + }, + } + + if split.ood_validation and split.ood_test: + ood_validation_embeddings = all_query_embeddings[f"{split.name}:ood_validation"] + ood_test_embeddings = all_query_embeddings[f"{split.name}:ood_test"] + selective_results: list[dict[str, Any]] = [] + for seed in SEEDS: + examples = select_examples(split.train_by_label, labels, 10, seed) + val_desc, val_ex, desc_embeddings, example_embeddings = dense_components( + model, validation_embeddings, labels, examples + ) + alpha = tune_alpha(val_desc, val_ex, validation_gold) + validation_scores = combine_scores(val_desc, val_ex, alpha) + ood_val_desc = ood_validation_embeddings @ desc_embeddings.T + ood_val_ex = np.einsum( + "qd,lnd->qln", ood_validation_embeddings, example_embeddings + ).max(axis=2) + ood_validation_scores = combine_scores(ood_val_desc, ood_val_ex, alpha) + score_threshold, margin_threshold, validation_metrics = ( + tune_decision_thresholds( + validation_scores, validation_gold, ood_validation_scores + ) + ) + + test_desc = test_embeddings @ desc_embeddings.T + test_ex = np.einsum("qd,lnd->qln", test_embeddings, example_embeddings).max( + axis=2 + ) + test_scores = combine_scores(test_desc, test_ex, alpha) + ood_test_desc = ood_test_embeddings @ desc_embeddings.T + ood_test_ex = np.einsum( + "qd,lnd->qln", ood_test_embeddings, example_embeddings + ).max(axis=2) + ood_test_scores = combine_scores(ood_test_desc, ood_test_ex, alpha) + test_metrics = decision_metrics( + test_scores, + test_gold, + ood_test_scores, + score_threshold, + margin_threshold, + ) + selective_results.append( + { + "seed": seed, + "alpha": alpha, + "score_threshold": score_threshold, + "margin_threshold": margin_threshold, + "validation": validation_metrics, + "test": test_metrics, + } + ) + + result["selective_routing"] = { + "target_validation_selective_error": TARGET_SELECTIVE_ERROR, + "target_validation_ood_fpr": TARGET_OOD_FPR, + "seed_results": selective_results, + "test": { + "id_coverage": mean_sd( + [float(item["test"]["id_coverage"]) for item in selective_results] + ), + "selective_error": mean_sd( + [ + float(item["test"]["selective_error"]) + for item in selective_results + ] + ), + "ood_fpr": mean_sd( + [float(item["test"]["ood_fpr"]) for item in selective_results] + ), + }, + } + return result + + +def benchmark_latency( + model: SentenceTransformer, + split: DatasetSplit, + seed: int, + count: int = 10, + iterations: int = 500, +) -> dict[str, Any]: + torch.set_num_threads(1) + labels = sorted(split.train_by_label) + examples = select_examples(split.train_by_label, labels, count, seed) + descriptions = model.encode( + [workflow_description(label) for label in labels], + normalize_embeddings=True, + convert_to_numpy=True, + show_progress_bar=False, + ) + example_embeddings = model.encode( + [text for label in labels for text in examples[label]], + batch_size=128, + normalize_embeddings=True, + convert_to_numpy=True, + show_progress_bar=False, + ).reshape(len(labels), count, -1) + queries = [text for text, _ in split.test] + alpha = 0.4 + + def route_once(query: str) -> None: + query_embedding = model.encode( + [query], + normalize_embeddings=True, + convert_to_numpy=True, + show_progress_bar=False, + )[0] + description_scores = descriptions @ query_embedding + example_scores = np.einsum( + "lnd,d->ln", example_embeddings, query_embedding + ).max(axis=1) + scores = (alpha * description_scores) + ((1.0 - alpha) * example_scores) + int(np.argmax(scores)) + + for query in queries[:50]: + route_once(query) + samples: list[float] = [] + for index in range(iterations): + start = time.perf_counter_ns() + route_once(queries[index % len(queries)]) + samples.append((time.perf_counter_ns() - start) / 1_000_000) + return { + "threads": torch.get_num_threads(), + "iterations": iterations, + "batch_size": 1, + "catalog_workflows": len(labels), + "examples_per_workflow": count, + "median_ms": float(np.median(samples)), + "p95_ms": float(np.percentile(samples, 95)), + "p99_ms": float(np.percentile(samples, 99)), + "mean_ms": float(np.mean(samples)), + } + + +def cpu_model_name() -> str: + try: + for line in Path("/proc/cpuinfo").read_text(encoding="utf-8").splitlines(): + if line.lower().startswith("model name"): + return line.split(":", 1)[1].strip() + except OSError: + pass + return platform.processor() or "unknown" + + +def write_summary_csv(results: dict[str, Any], path: Path) -> None: + rows: list[dict[str, Any]] = [] + for dataset in results["datasets"]: + rows.append( + { + "dataset": dataset["dataset"], + "system": "word-char TF-IDF (10 examples)", + "top1_mean": dataset["lexical_10_examples"]["top1"]["mean"], + "top1_sd": dataset["lexical_10_examples"]["top1"]["sd"], + "recall_at_5_mean": dataset["lexical_10_examples"]["recall_at_5"][ + "mean" + ], + "recall_at_5_sd": dataset["lexical_10_examples"]["recall_at_5"]["sd"], + } + ) + for count in EXAMPLE_COUNTS: + record = dataset["ablation"][str(count)] + rows.append( + { + "dataset": dataset["dataset"], + "system": f"dense example-aware ({count} examples)", + "top1_mean": record["top1"]["mean"], + "top1_sd": record["top1"]["sd"], + "recall_at_5_mean": record["recall_at_5"]["mean"], + "recall_at_5_sd": record["recall_at_5"]["sd"], + } + ) + with path.open("w", encoding="utf-8", newline="") as handle: + writer = csv.DictWriter(handle, fieldnames=list(rows[0])) + writer.writeheader() + writer.writerows(rows) + + +def main() -> int: + parser = argparse.ArgumentParser() + parser.add_argument("--data-dir", type=Path, default=Path(__file__).parent / "data") + parser.add_argument( + "--output-dir", type=Path, default=Path(__file__).parent / "results" + ) + parser.add_argument( + "--cache-dir", type=Path, default=Path(__file__).parent / "cache" + ) + args = parser.parse_args() + + ensure_data(args.data_dir) + args.output_dir.mkdir(parents=True, exist_ok=True) + args.cache_dir.mkdir(parents=True, exist_ok=True) + os.environ.setdefault("HF_HOME", str(args.cache_dir.resolve())) + + torch.manual_seed(20260912) + np.random.seed(20260912) + model = SentenceTransformer(MODEL_ID, revision=MODEL_REVISION) + model.eval() + parameter_count = sum(parameter.numel() for parameter in model.parameters()) + + datasets = [load_clinc(args.data_dir), load_banking(args.data_dir)] + query_embeddings: dict[str, np.ndarray] = {} + for split in datasets: + groups = { + "validation": [text for text, _ in split.validation], + "test": [text for text, _ in split.test], + "ood_validation": split.ood_validation, + "ood_test": split.ood_test, + } + for group_name, texts in groups.items(): + if texts: + query_embeddings[f"{split.name}:{group_name}"] = model.encode( + texts, + batch_size=128, + normalize_embeddings=True, + convert_to_numpy=True, + show_progress_bar=False, + ) + + started = time.time() + dataset_results = [ + evaluate_dataset(model, split, query_embeddings) for split in datasets + ] + latency = benchmark_latency(model, datasets[0], seed=SEEDS[0]) + elapsed = time.time() - started + + results: dict[str, Any] = { + "experiment": "FlowRoute public-data benchmark", + "model": { + "id": MODEL_ID, + "revision": MODEL_REVISION, + "parameters": parameter_count, + }, + "method": { + "seeds": list(SEEDS), + "example_counts": list(EXAMPLE_COUNTS), + "alpha_grid": list(ALPHAS), + "selection": ( + "alpha and route thresholds selected on validation data only; both validation " + "risk constraints use upper endpoints of two-sided 95% Wilson intervals" + ), + }, + "datasets": dataset_results, + "latency": latency, + "environment": { + "python": platform.python_version(), + "platform": platform.platform(), + "cpu": cpu_model_name(), + "torch": torch.__version__, + "numpy": np.__version__, + "scikit_learn": sklearn.__version__, + "sentence_transformers": __import__("sentence_transformers").__version__, + }, + "data": { + filename: { + "url": metadata["url"], + "sha256": sha256(args.data_dir / filename), + } + for filename, metadata in DATASETS.items() + }, + "evaluation_seconds_excluding_initial_query_encoding": elapsed, + } + result_path = args.output_dir / "benchmark_results.json" + result_path.write_text( + json.dumps(results, indent=2, sort_keys=True) + "\n", encoding="utf-8" + ) + write_summary_csv(results, args.output_dir / "ranking_summary.csv") + print(f"Wrote {result_path}") + for dataset in dataset_results: + dense = dataset["ablation"]["10"] + print( + f"{dataset['dataset']}: " + f"top-1={dense['top1']['mean']:.4f} +/- {dense['top1']['sd']:.4f}, " + f"R@5={dense['recall_at_5']['mean']:.4f} +/- " + f"{dense['recall_at_5']['sd']:.4f}" + ) + print( + f"CPU latency: median={latency['median_ms']:.2f} ms, " + f"p95={latency['p95_ms']:.2f} ms" + ) + return 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/paper/FlowRoute_Ali_Norouzi.pdf b/paper/FlowRoute_Ali_Norouzi.pdf new file mode 100644 index 0000000..56066a4 Binary files /dev/null and b/paper/FlowRoute_Ali_Norouzi.pdf differ