LFM2.5-350M-RLCD / tests /test_engine.py
notnotsamuel's picture
Release unchanged-weight LFM2.5 parallel constrained inference and benchmarks
a99edfc verified
Raw
History Blame Contribute Delete
3.71 kB
import copy
import json
import pytest
import jsonschema
import torch
from rlcd.engine import Engine, fork_cache, validate_schema
from rlcd.benchmark import evaluate
from rlcd.tasks import ROUTING
@pytest.fixture(scope="module")
def engine():
import os
return Engine(os.environ.get("RLCD_TEST_DEVICE", "mps"), "float16")
@torch.inference_mode()
def test_hybrid_cache_matches_full_forward_and_is_independent(engine):
prefix = engine.encode(engine.prompt("Route north west with express plus and insurance.", ROUTING))
suffixes = [engine.encode(' "route": "north west"\n'), engine.encode(' "service": "express plus"\n')]
base = engine.model(engine.tensor([prefix]), use_cache=True).past_key_values
snapshot = copy.deepcopy(base)
fork = fork_cache(base, 2)
assert sum(hasattr(l, "keys") for l in fork.layers) == 6
assert sum(hasattr(l, "conv_states") for l in fork.layers) == 10
width = max(map(len, suffixes))
ids = engine.tensor([s + [0] * (width - len(s)) for s in suffixes])
mask = engine.tensor([[1] * (len(prefix) + len(s)) + [0] * (width - len(s)) for s in suffixes])
batched = engine.model(ids, attention_mask=mask, past_key_values=fork).logits
errors = []
for i, suffix in enumerate(suffixes):
full = engine.model(engine.tensor([prefix + suffix]), use_cache=False).logits[0, len(prefix):]
cached = batched[i, :len(suffix)]
errors.append((cached - full).abs().max().item())
# FP16 backend ordering can alter logits slightly; compare distributions too.
torch.testing.assert_close(cached.float().softmax(-1), full.float().softmax(-1), atol=0.015, rtol=0.08)
assert torch.equal(cached.argmax(-1), full.argmax(-1))
for old, current in zip(snapshot.layers, base.layers):
if hasattr(old, "keys"):
assert torch.equal(old.keys, current.keys)
assert torch.equal(old.values, current.values)
else:
assert torch.equal(old.conv_states[0], current.conv_states[0])
print("cache/full max logit error:", errors)
@torch.inference_mode()
def test_candidate_scores_equal_uncached_reference(engine):
context = "Route: north west. Service: express plus. No insurance."
result = engine.constrained(context, ROUTING)
prefix = engine.encode(engine.prompt(context, ROUTING))
for name, choices in result["scores"].items():
for choice in choices:
suffix = engine.encode(" " + json.dumps(name) + ": ")
value = engine.encode(json.dumps(choice["value"]) + "\n")
tokens = prefix + suffix + value
start = len(prefix) + len(suffix)
logits = engine.model(engine.tensor([tokens]), use_cache=False).logits[0, start-1:start+len(value)-1].float()
score = logits.log_softmax(-1).gather(1, engine.tensor(value)[:, None]).sum().item()
assert abs(score - choice["log_likelihood"]) < 0.15
parsed = json.loads(result["text"])
assert parsed["service"] in ROUTING["properties"]["service"]["enum"]
assert type(parsed["insured"]) is bool
def test_strict_validation():
expected = {"route":"north west", "service":"express", "insured":False}
assert evaluate(json.dumps(expected), ROUTING, expected)["exact_match"]
assert not evaluate('```json\n' + json.dumps(expected) + '\n```', ROUTING, expected)["syntax_valid"]
assert not evaluate(json.dumps({**expected, "insured":0}), ROUTING, expected)["schema_compliant"]
with pytest.raises(ValueError):
validate_schema({**ROUTING, "allOf": [{}]})
with pytest.raises(jsonschema.SchemaError):
validate_schema({**ROUTING, "properties": {"bad": {"type":"string", "enum":"abc"}}})