"""The Core ML router must reproduce the native PyTorch model (tests/golden.json).""" import json from pathlib import Path import pytest from cases import CASES, SUPPORT_TASKS from gliner_decide_coreml import BUCKETS, DecideRouter GOLDEN = json.loads((Path(__file__).parent / "golden.json").read_text())["cases"] CONFIDENCE_TOLERANCE = 0.02 # fp16 Core ML vs fp32 PyTorch @pytest.fixture(scope="module") def router(): return DecideRouter() def entries(value): return value if isinstance(value, list) else [value] @pytest.mark.parametrize("case", CASES, ids=[c["id"] for c in CASES]) def test_matches_native(router, case): result, route = router.classify_with_route(case["text"], case["tasks"]) assert route.bucket == case["bucket"] and route.calls == 1 native = GOLDEN[case["id"]] assert list(result) == list(native) for head, expected in native.items(): got = {e["label"]: e["confidence"] for e in entries(result[head])} want = {e["label"]: e["confidence"] for e in entries(expected)} assert set(got) == set(want), f"{head}: {sorted(got)} != {sorted(want)}" for label, confidence in want.items(): assert got[label] == pytest.approx(confidence, abs=CONFIDENCE_TOLERANCE), f"{head}/{label}" def test_every_bucket_is_exercised(): assert {c["bucket"] for c in CASES} == set(BUCKETS) def test_long_text_is_chunked(router): text = " ".join([CASES[0]["text"]] * 40) # ~1,000 tokens: larger than the 512 bucket result, route = router.classify_with_route(text, SUPPORT_TASKS) assert route.bucket is None and route.chunks >= 2 and route.calls == route.chunks assert result["intent"]["label"] == "refund_request" def test_more_than_four_heads_are_split_across_calls(router): tasks = {**CASES[4]["tasks"], "language": ["english", "spanish", "german"], "sentiment": ["positive", "negative"]} result, route = router.classify_with_route(CASES[4]["text"], tasks) assert list(result) == list(tasks) and route.calls == 2 assert result["language"]["label"] == "english" def test_too_many_labels_is_rejected(router): with pytest.raises(ValueError, match="labels"): router.classify("hello", {"topic": [f"label_{i}" for i in range(33)]})