"""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//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