Files
esh-pfi-infrastructure/services/intern-decision-serve/tests/test_engine.py
T
vh 21d16d7ad8 fix(intern-decision-serve): code-review fixes
- a failure while building the response (a non-finite number included) is a 500 inside the
  envelope, never a 422 or a render crash outside it
- an engine ValueError keeps its message but is released and raised unchained, like an OOM
- the prompt is built (and the model's own validation run) before the forward
- the row cap is counted before any ordering is built
- /health reads a device name cached at load, so it makes no driver call off the inference thread
2026-09-30 09:30:40 -07:00

347 lines
14 KiB
Python

"""TorchEngine against a fake torch and a fake checkpoint (no GPU, no model).
Contract: intern-decision-serve.contract.md INV-3, INV-4, INV-7, INV-8."""
import hashlib
import logging
import textwrap
import weakref
import pytest
from intern_decision_serve.config import REVISION, Settings
from intern_decision_serve.engine import TorchEngine
from intern_decision_serve.errors import OutOfMemory, ScoringFailed
TOKEN = "t" * 40
REQUEST = {"state": "s", "questions": {"q": {"type": "choice", "instructions": "Q?", "criteria": {"a": "A", "b": "B"}}}}
# ---------------------------------------------------------------------------------------------
# A fake torch: records what was called, in order.
# ---------------------------------------------------------------------------------------------
class FakeTorch:
__version__ = "2.10.0+fake"
events: list = []
class nn:
class Module:
def __call__(self, *a, **k):
return self.forward(*a, **k)
class cuda:
class OutOfMemoryError(RuntimeError):
pass
available = True
capability = (12, 0)
arch_list = ["sm_90", "sm_120"]
reserved = 0
watched: list = []
empties: list = []
@classmethod
def is_available(cls):
return cls.available
@classmethod
def get_device_capability(cls, _i):
return cls.capability
@classmethod
def get_arch_list(cls):
return cls.arch_list
@classmethod
def get_device_properties(cls, _i):
return type("P", (), {"total_memory": 96 * 2**30})
@classmethod
def set_per_process_memory_fraction(cls, fraction, device):
FakeTorch.events.append(("cap", round(fraction, 6), device))
@classmethod
def memory_reserved(cls, _i):
return cls.reserved
@classmethod
def memory_allocated(cls, _i):
return cls.reserved
@classmethod
def empty_cache(cls):
cls.empties.append(all(ref() is None for ref in cls.watched))
FakeTorch.events.append(("empty_cache",))
@classmethod
def get_device_name(cls, _i):
return "Fake RTX"
@pytest.fixture(autouse=True)
def reset_fake_torch():
FakeTorch.events = []
c = FakeTorch.cuda
c.available, c.capability, c.arch_list, c.reserved = True, (12, 0), ["sm_90", "sm_120"], 0
c.watched, c.empties = [], []
yield
# ---------------------------------------------------------------------------------------------
# A fake checkpoint: snapshots/<REVISION>/inference.py defining a DecisionEngine shaped like the real one.
# ---------------------------------------------------------------------------------------------
FAKE_INFERENCE = textwrap.dedent('''
from __future__ import annotations # as in the real one: dataclasses then look the module up in sys.modules
import json
from dataclasses import dataclass
MODEL_NAME = "Intern-Decision-4B"
EVENTS = []
WARMUP_SHIFT = {"after_swap": 0.0}
TOKEN_SKEW = {"n": 0}
@dataclass(frozen=True) # the real inference.py defines one: needs sys.modules at exec time
class Compiled:
messages: list
def validate_request(request):
return {"state": request["state"], "questions": request["questions"]}
def compile_row(row):
return Compiled([{"role": "user", "content": json.dumps(row, sort_keys=True)}])
class Tokenizer:
def apply_chat_template(self, messages, tokenize, add_generation_prompt, enable_thinking, add_vision_id):
assert (tokenize, add_generation_prompt, enable_thinking, add_vision_id) == (False, False, False, True)
return "<|im_start|>" + messages[0]["content"]
def __call__(self, text, add_special_tokens):
return {"input_ids": list(range(len(text.split()) + TOKEN_SKEW["n"]))}
class Param:
class device:
type = "cuda"
class Inner:
def __init__(self):
self.visual = "the vision tower"
class Model:
def __init__(self):
self.model = Inner()
def parameters(self):
yield Param()
class Backend:
def __init__(self):
self.tokenizer = Tokenizer()
self.model = Model()
class DecisionEngine:
def __init__(self, checkpoint=None, *, max_length=8192, device="cuda", **kw):
EVENTS.append(("construct", checkpoint, max_length, device))
self.backend = Backend()
self.tokenizer = self.backend.tokenizer
self.temperature = 1.99241824
def predict(self, request):
row = validate_request(request)
swapped = self.backend.model.model.visual != "the vision tower"
p = 0.75 + (WARMUP_SHIFT["after_swap"] if swapped else 0.0)
text = self.tokenizer.apply_chat_template(compile_row(row).messages, tokenize=False,
add_generation_prompt=False, enable_thinking=False, add_vision_id=True)
answers = {f: {"type": "choice", "probabilities": dict(zip(q["criteria"], [p, 1 - p])),
"confidence": p, "choice": list(q["criteria"])[0], "source": "local",
"decision": list(q["criteria"])[0]} for f, q in row["questions"].items()}
return {"answers": answers, "usage": {"input_tokens": len(text.split())},
"timing": {"inference_ms": 3.0}, "calibration": {"method": "temperature-scaling",
"temperature": self.temperature}, "model": MODEL_NAME, "backend": "hf"}
''')
def fake_checkpoint(tmp_path, source=FAKE_INFERENCE, revision=REVISION):
snap = tmp_path / "hub" / "models--internlm--Intern-Decision-4B" / "snapshots" / revision
snap.mkdir(parents=True)
(snap / "inference.py").write_text(source)
return snap, hashlib.sha256(source.encode()).hexdigest()
def load(tmp_path, source=FAKE_INFERENCE, revision=REVISION, sha=None, **settings):
snap, real_sha = fake_checkpoint(tmp_path, source, revision)
s = Settings(api_token=TOKEN, checkpoint=str(snap), **settings)
return TorchEngine.load(s, torch=FakeTorch, inference_sha256=sha or real_sha)
# ---------------------------------------------------------------------------------------------
# Load (INV-3, INV-7, INV-8)
# ---------------------------------------------------------------------------------------------
def test_load_serves_the_checkpoint_and_hashes_the_rendered_prompt(tmp_path):
engine = load(tmp_path, vram_cap_gib=12.0, max_tokens=6000)
response, sha = engine.predict(REQUEST)
assert response["answers"]["q"]["probabilities"] == {"a": 0.75, "b": 0.25}
text = "<|im_start|>" + __import__("json").dumps(REQUEST, sort_keys=True)
assert sha == hashlib.sha256(text.encode()).hexdigest()
assert engine.metadata["revision"] == REVISION and engine.metadata["vision_tower"] == "removed"
assert engine.health()["device_name"] == "Fake RTX"
def test_the_cap_is_applied_before_the_engine_is_constructed(tmp_path):
engine = load(tmp_path, vram_cap_gib=12.0)
inference = engine._inference
assert FakeTorch.events[0] == ("cap", 0.125, 0)
assert inference.EVENTS[0][0] == "construct" and inference.EVENTS[0][2:] == (8192, "cuda")
def test_a_cap_larger_than_the_card_is_refused(tmp_path):
with pytest.raises(ValueError, match="VRAM_CAP_GIB"):
load(tmp_path, vram_cap_gib=200.0)
def test_a_changed_inference_py_is_refused_before_anything_loads(tmp_path):
with pytest.raises(RuntimeError, match="inference.py"):
load(tmp_path, sha="0" * 64)
assert FakeTorch.events == []
def test_a_checkpoint_that_is_not_the_pinned_snapshot_is_refused(tmp_path):
with pytest.raises(RuntimeError, match=REVISION):
load(tmp_path, revision="f" * 40)
@pytest.mark.parametrize("problem", ["no_cuda", "no_arch"])
def test_a_card_torch_cannot_drive_is_refused(tmp_path, problem):
if problem == "no_cuda":
FakeTorch.cuda.available = False
else:
FakeTorch.cuda.arch_list = ["sm_80", "sm_90"]
with pytest.raises(RuntimeError):
load(tmp_path)
def test_the_vision_tower_is_replaced_by_a_stub_that_refuses_to_run(tmp_path):
engine = load(tmp_path)
stub = engine._engine.backend.model.model.visual
assert stub != "the vision tower"
with pytest.raises(RuntimeError, match="vision tower"):
stub(None)
def test_keep_vision_leaves_the_tower_alone(tmp_path):
engine = load(tmp_path, keep_vision=True)
assert engine._engine.backend.model.model.visual == "the vision tower"
assert engine.metadata["vision_tower"] == "loaded"
def test_startup_refuses_when_removing_the_vision_tower_changes_the_warm_up(tmp_path):
source = FAKE_INFERENCE.replace('WARMUP_SHIFT = {"after_swap": 0.0}', 'WARMUP_SHIFT = {"after_swap": 0.001}')
with pytest.raises(RuntimeError, match="INV-7"):
load(tmp_path, source)
def test_startup_refuses_a_prompt_hash_the_model_never_saw(tmp_path):
source = FAKE_INFERENCE.replace('TOKEN_SKEW = {"n": 0}', 'TOKEN_SKEW = {"n": 1}')
with pytest.raises(RuntimeError, match="INV-8"):
load(tmp_path, source)
def test_the_baseline_is_taken_after_warm_up_and_a_release(tmp_path):
FakeTorch.cuda.reserved = 9 * 2**30
engine = load(tmp_path, release_slack_mib=256)
assert engine._release_above == 9 * 2**30 + 256 * 2**20
assert ("empty_cache",) in FakeTorch.events
# ---------------------------------------------------------------------------------------------
# The guard around every call (INV-4), on an engine built directly
# ---------------------------------------------------------------------------------------------
class Tensor:
pass
class StubEngine:
def __init__(self, fn):
self.predict = fn
def direct(fn, release_above=None):
import types
inference = types.SimpleNamespace(validate_request=lambda r: r,
compile_row=lambda row: types.SimpleNamespace(messages=[]))
tokenizer = types.SimpleNamespace(apply_chat_template=lambda *a, **k: "prompt")
return TorchEngine(FakeTorch, StubEngine(fn), inference, tokenizer, metadata={},
settings=Settings(api_token=TOKEN), release_above_bytes=release_above)
def failing(exc_factory):
def predict(_request):
kv = Tensor() # stands in for the failed forward's activations
FakeTorch.cuda.watched.append(weakref.ref(kv))
raise exc_factory()
return predict
def test_an_oom_frees_the_failed_call_before_emptying_the_cache_and_is_unchained():
engine = direct(failing(lambda: FakeTorch.cuda.OutOfMemoryError("CUDA out of memory. Tried to allocate 2 GiB.\nmore")))
with pytest.raises(OutOfMemory) as info:
engine.predict(REQUEST)
assert str(info.value) == "CUDA out of memory. Tried to allocate 2 GiB."
assert info.value.__cause__ is None and info.value.__context__ is None
assert FakeTorch.cuda.watched[0]() is None and FakeTorch.cuda.empties == [True]
def test_an_empty_oom_message_still_reads_as_an_oom():
with pytest.raises(OutOfMemory, match="CUDA out of memory"):
direct(failing(lambda: FakeTorch.cuda.OutOfMemoryError(""))).predict(REQUEST)
def test_a_runtime_error_saying_out_of_memory_is_an_oom():
with pytest.raises(OutOfMemory):
direct(failing(lambda: RuntimeError("CUBLAS_STATUS_ALLOC_FAILED: out of memory"))).predict(REQUEST)
def test_any_other_failure_is_logged_released_and_raised_unchained_as_scoring_failed(caplog):
engine = direct(failing(lambda: KeyError("q")))
with caplog.at_level(logging.ERROR), pytest.raises(ScoringFailed) as info:
engine.predict(REQUEST)
assert "KeyError" in str(info.value) and info.value.__context__ is None
assert "Traceback" in caplog.text
assert FakeTorch.cuda.empties == [True]
def test_a_value_error_keeps_its_message_but_is_released_and_raised_unchained():
"""A ValueError can come from inside the forward too (a transformers shape check): it must not pin
the failed forward's tensors through a chained traceback (review 2026-09-30)."""
engine = direct(failing(lambda: ValueError("Example has 9000 tokens, above 8192; truncation is forbidden")))
with pytest.raises(ValueError, match="truncation is forbidden") as info:
engine.predict(REQUEST)
assert info.value.__cause__ is None and info.value.__context__ is None
assert FakeTorch.cuda.watched[0]() is None and FakeTorch.cuda.empties == [True]
def test_a_request_the_prompt_builder_rejects_never_reaches_the_model():
import types
called = []
inference = types.SimpleNamespace(validate_request=lambda r: (_ for _ in ()).throw(ValueError("Supply 1-16 questions.")),
compile_row=lambda row: None)
engine = TorchEngine(FakeTorch, StubEngine(lambda r: called.append(r)), inference, tokenizer=None, metadata={},
settings=Settings(api_token=TOKEN))
with pytest.raises(ValueError, match="1-16 questions"):
engine.predict(REQUEST)
assert called == []
def test_health_does_not_ask_the_driver_for_the_device_name_again(tmp_path):
engine = load(tmp_path)
FakeTorch.cuda.get_device_name = classmethod(lambda cls, _i: (_ for _ in ()).throw(AssertionError("driver call")))
try:
assert engine.health()["device_name"] == "Fake RTX"
finally:
FakeTorch.cuda.get_device_name = classmethod(lambda cls, _i: "Fake RTX")
@pytest.mark.parametrize("reserved, released", [(2 * 2**30, True), (2**30, False), (2**30 - 1, False)])
def test_a_burst_over_the_baseline_plus_slack_is_released(reserved, released):
def ok(_request):
FakeTorch.cuda.reserved = reserved
return {"answers": {}}
direct(ok, release_above=2**30).predict(REQUEST)
assert (FakeTorch.cuda.empties != []) is released