Checkpoint frontend UI/UX overhaul and ingestion embed benchmark work
This commit is contained in:
@@ -0,0 +1,157 @@
|
||||
import csv
|
||||
import json
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
|
||||
from ingestion.embed.benchmark_local import (
|
||||
BenchmarkCase,
|
||||
candidate_chunks,
|
||||
load_cases,
|
||||
run_benchmark,
|
||||
)
|
||||
from ingestion.embed.local_bge_m3 import BgeM3Local
|
||||
|
||||
|
||||
def _write_csv(path: Path, rows: list[dict]) -> None:
|
||||
with path.open("w", encoding="utf-8", newline="") as handle:
|
||||
writer = csv.DictWriter(handle, fieldnames=list(rows[0]))
|
||||
writer.writeheader()
|
||||
writer.writerows(rows)
|
||||
|
||||
|
||||
def test_load_cases_uses_verified_aliases_and_skips_policy_and_multidrug(tmp_path):
|
||||
entities = tmp_path / "entities.json"
|
||||
entities.write_text(
|
||||
json.dumps(
|
||||
{
|
||||
"entities": [
|
||||
{"drug_id": "paracetamol", "aliases": ["Paracetamol"]}
|
||||
]
|
||||
}
|
||||
),
|
||||
encoding="utf-8",
|
||||
)
|
||||
golden = tmp_path / "golden.csv"
|
||||
_write_csv(
|
||||
golden,
|
||||
[
|
||||
{
|
||||
"id": "1",
|
||||
"cau_hoi": "Liều Paracetamol?",
|
||||
"thuoc_ky_vong": "Paracetamol",
|
||||
"thuoc_tinh_ky_vong": "lieu_dung",
|
||||
},
|
||||
{
|
||||
"id": "2",
|
||||
"cau_hoi": "Uống gì?",
|
||||
"thuoc_ky_vong": "",
|
||||
"thuoc_tinh_ky_vong": "",
|
||||
},
|
||||
{
|
||||
"id": "3",
|
||||
"cau_hoi": "So sánh",
|
||||
"thuoc_ky_vong": "Paracetamol; Ibuprofen",
|
||||
"thuoc_tinh_ky_vong": "lieu_dung",
|
||||
},
|
||||
],
|
||||
)
|
||||
|
||||
assert load_cases(golden, entities) == [
|
||||
BenchmarkCase(
|
||||
case_id="1",
|
||||
query="Liều Paracetamol?",
|
||||
drug_id="paracetamol",
|
||||
expected_section="lieu_luong_va_cach_dung",
|
||||
)
|
||||
]
|
||||
|
||||
|
||||
def test_candidate_chunks_fails_when_ground_truth_is_missing(tmp_path):
|
||||
path = tmp_path / "chunks.jsonl"
|
||||
path.write_text(
|
||||
json.dumps(
|
||||
{
|
||||
"drug_id": "drug",
|
||||
"section_key": "chi_dinh",
|
||||
"chunk_kind": "prose",
|
||||
"text": "text",
|
||||
}
|
||||
)
|
||||
+ "\n",
|
||||
encoding="utf-8",
|
||||
)
|
||||
cases = [BenchmarkCase("1", "dose", "drug", "lieu_luong_va_cach_dung")]
|
||||
|
||||
with pytest.raises(ValueError, match="no canonical target"):
|
||||
candidate_chunks(path, cases)
|
||||
|
||||
|
||||
def test_load_cases_rejects_an_alias_ambiguous_in_the_verified_catalog(tmp_path):
|
||||
entities = tmp_path / "entities.json"
|
||||
entities.write_text(
|
||||
json.dumps(
|
||||
{
|
||||
"entities": [
|
||||
{"drug_id": "drug_a", "aliases": ["Shared"]},
|
||||
{"drug_id": "drug_b", "aliases": ["Shared"]},
|
||||
]
|
||||
}
|
||||
),
|
||||
encoding="utf-8",
|
||||
)
|
||||
golden = tmp_path / "golden.csv"
|
||||
_write_csv(
|
||||
golden,
|
||||
[
|
||||
{
|
||||
"id": "1",
|
||||
"cau_hoi": "Liều Shared?",
|
||||
"thuoc_ky_vong": "Shared",
|
||||
"thuoc_tinh_ky_vong": "lieu_dung",
|
||||
}
|
||||
],
|
||||
)
|
||||
|
||||
with pytest.raises(ValueError, match="is ambiguous"):
|
||||
load_cases(golden, entities)
|
||||
|
||||
|
||||
def test_benchmark_filters_by_drug_and_reports_first_relevant_rank():
|
||||
vectors = {
|
||||
"query": [1.0, 0.0],
|
||||
"wrong section": [0.9, 0.1],
|
||||
"correct section": [0.8, 0.2],
|
||||
"other drug": [1.0, 0.0],
|
||||
}
|
||||
provider = BgeM3Local(encoder=lambda texts: [vectors[text] for text in texts])
|
||||
provider._check_dimensions = lambda _values: None
|
||||
cases = [BenchmarkCase("1", "query", "drug", "target")]
|
||||
chunks = [
|
||||
{
|
||||
"chunk_id": "wrong",
|
||||
"drug_id": "drug",
|
||||
"section_key": "other",
|
||||
"text": "wrong section",
|
||||
},
|
||||
{
|
||||
"chunk_id": "right",
|
||||
"drug_id": "drug",
|
||||
"section_key": "target",
|
||||
"text": "correct section",
|
||||
},
|
||||
{
|
||||
"chunk_id": "leak",
|
||||
"drug_id": "other",
|
||||
"section_key": "target",
|
||||
"text": "other drug",
|
||||
},
|
||||
]
|
||||
|
||||
report = run_benchmark(provider, cases, chunks)
|
||||
|
||||
assert report.case_count == 1
|
||||
assert report.cases[0].first_relevant_rank == 2
|
||||
assert report.hit_at_1 == 0.0
|
||||
assert report.hit_at_3 == 1.0
|
||||
assert report.mrr == 0.5
|
||||
Reference in New Issue
Block a user