augustoFranke's picture
GLiNER2.5-Decide Core ML multifunction package (L64-L512) with bucket router
cb96101 verified
Raw History Blame
2.26 kB
"""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)]})