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
This commit is contained in:
vh
2026-09-30 09:30:40 -07:00
parent 618390c5fa
commit 21d16d7ad8
5 changed files with 110 additions and 21 deletions
@@ -43,7 +43,12 @@ request. A violation is a 422.
- **One call** is one `predict()` request: - **One call** is one `predict()` request:
`{"state": state, "questions": {<field>: {"type": "choice", "instructions": question, "criteria": {option.id: option.description, ...}}}}`. `{"state": state, "questions": {<field>: {"type": "choice", "instructions": question, "criteria": {option.id: option.description, ...}}}}`.
The criteria keep the caller's option order. This is the format the bench measured. The criteria keep the caller's option order.
- **A one-question call is exactly the bench's native request.** Acceptance showed it
bit-identical, row for row.
- **A multi-question call differs from the bench's multifield run only in the field names.**
Those are positional here and were the decision names in the bench. Measured: 1 of 84 Wyrd
rows differs (78/84 against 77/84).
- **Field names are positional:** `q` when a call carries one question, and `q1`..`qN` in - **Field names are positional:** `q` when a call carries one question, and `q1`..`qN` in
request order when it carries several. Decision ids never reach the prompt. **Option ids do:** request order when it carries several. Decision ids never reach the prompt. **Option ids do:**
the model prints `A = <id>: <description>`. the model prints `A = <id>: <description>`.
@@ -87,12 +92,15 @@ request. A violation is a 422.
- **INV-1 pass-through.** Every number in `native`, `probabilities`, `top`, `confidence`, - **INV-1 pass-through.** Every number in `native`, `probabilities`, `top`, `confidence`,
`calibration` and `input_tokens` is what `predict()` returned. The wrapper only re-keys it. `calibration` and `input_tokens` is what `predict()` returned. The wrapper only re-keys it.
If an answer lacks one of the decision's option ids, that is a 500 `scoring_failed`, never a If an answer lacks one of the decision's option ids, that is a 500 `scoring_failed`, never a
guess. guess. So is any failure while building the response, a non-finite number included. That stays
inside the error envelope and is never a 422 or a bare 500.
- **INV-2 one model, one inference thread.** The model loads at startup on a dedicated - **INV-2 one model, one inference thread.** The model loads at startup on a dedicated
single-thread executor. The warm-ups and every later call run on that **same host thread**, single-thread executor. The warm-ups and every later call run on that **same host thread**,
never on the event loop's threadpool. A lock also serialises each request's calls, so all of a never on the event loop's threadpool. A lock also serialises each request's calls, so all of a
request's chunks run inside one hold. The app runs one worker, and `/health` answers during a request's chunks run inside one hold. The app runs one worker, and `/health` answers during a
call. call.
- `/health` reads only allocator counters and a device name cached at load, so it adds no
per-thread CUDA state.
- **Why one thread:** torch keeps CUDA state per host thread (cuBLAS handles and workspaces), - **Why one thread:** torch keeps CUDA state per host thread (cuBLAS handles and workspaces),
and part of it sits outside the VRAM cap. and part of it sits outside the VRAM cap.
- **Measured on 2026-09-30, fv-ml1 GPU 3:** anyio's 40 worker threads added 252 MiB outside the - **Measured on 2026-09-30, fv-ml1 GPU 3:** anyio's 40 worker threads added 252 MiB outside the
@@ -110,10 +118,12 @@ request. A violation is a 422.
- An OOM in any call makes the whole request a 503 `out_of_memory`. So does a RuntimeError - An OOM in any call makes the whole request a 503 `out_of_memory`. So does a RuntimeError
whose first line says "out of memory". The engine then frees the failed call's frames, whose first line says "out of memory". The engine then frees the failed call's frames,
runs `gc.collect()` and `empty_cache()`, and raises unchained. The process stays up. runs `gc.collect()` and `empty_cache()`, and raises unchained. The process stays up.
- A `ValueError` (the token limit, or one raised inside a forward) keeps its message and maps
to 422. It is released and raised unchained the same way.
- After every call, if reserved memory exceeds the post-warm-up baseline by more than - After every call, if reserved memory exceeds the post-warm-up baseline by more than
`RELEASE_SLACK_MIB` (default 512), the engine runs `empty_cache()`. `RELEASE_SLACK_MIB` (default 512), the engine runs `empty_cache()`.
- Any other failure except `ValueError` is logged with its traceback and released the same - Any other failure is logged with its traceback and released the same way, then raised
way, then raised unchained as `ScoringFailed`. unchained as `ScoringFailed`.
- **INV-5 no network.** The entry point sets `HF_HUB_OFFLINE=1` and `TRANSFORMERS_OFFLINE=1` - **INV-5 no network.** The entry point sets `HF_HUB_OFFLINE=1` and `TRANSFORMERS_OFFLINE=1`
before torch or transformers load. The weights are read from the mounted, read-only HF cache. before torch or transformers load. The weights are read from the mounted, read-only HF cache.
- **INV-6 constant-time auth.** The token is compared with `hmac.compare_digest`. It must be at - **INV-6 constant-time auth.** The token is compared with `hmac.compare_digest`. It must be at
@@ -249,7 +259,8 @@ on the host, the cap is the single knob `VRAM_CAP_GIB`.
are freed. are freed.
- A RuntimeError saying "out of memory" becomes `OutOfMemory`. - A RuntimeError saying "out of memory" becomes `OutOfMemory`.
- Another failure becomes `ScoringFailed`, unchained. - Another failure becomes `ScoringFailed`, unchained.
- `ValueError` passes through. - A `ValueError` keeps its message but is released and raised unchained.
- A request the prompt builder rejects never reaches the model.
- A burst over baseline plus slack is released, and one at or under it is left alone. - A burst over baseline plus slack is released, and one at or under it is left alone.
- **Engine load (fake).** - **Engine load (fake).**
- A wrong `inference.py` hash, or a checkpoint path that is not the pinned snapshot, refuses - A wrong `inference.py` hash, or a checkpoint path that is not the pinned snapshot, refuses
@@ -124,6 +124,19 @@ def check_request(state: State, decisions: list[Decision], workload: str | None)
raise bad(f"decision {d.id!r}: option ids must be unique") raise bad(f"decision {d.id!r}: option ids must be unique")
def ordering_count(d: Decision) -> int:
"""How many orderings `d` asks, without building them (the row cap is checked first)."""
n = len(d.options)
if d.orderings == "none":
return 1
if d.orderings == "rotations":
return n
if n > MAX_OPTIONS_FOR_ALL:
raise ApiError(422, "invalid_request", f"decision {d.id!r}: orderings 'all' allows at most "
f"{MAX_OPTIONS_FOR_ALL} options ({n} given); use 'rotations'")
return math.factorial(n)
def ordering_perms(d: Decision) -> list[tuple[int, ...]]: def ordering_perms(d: Decision) -> list[tuple[int, ...]]:
"""Index permutations of the caller's options, the caller's own order first.""" """Index permutations of the caller's options, the caller's own order first."""
n = len(d.options) n = len(d.options)
@@ -190,6 +203,19 @@ def row_result(slot: Slot, response: dict, sha: str, field: str, questions: int,
"call": {"index": index, "field": field, "questions": questions}} "call": {"index": index, "field": field, "questions": questions}}
def built(fn, *args):
"""Building the response from the model's output is server-side work: any failure there is a 500
scoring_failed, never a 422 the caller would read as their own bad input (review 2026-09-30)."""
try:
out = fn(*args)
json.dumps(out, allow_nan=False) # a NaN/inf would otherwise crash rendering OUTSIDE the envelope
return out
except ScoringFailed:
raise
except Exception as exc: # noqa: BLE001
raise ScoringFailed(f"building the response failed: {type(exc).__name__}: {exc}") from None
def combine(d: Decision, results: list[dict]) -> dict: def combine(d: Decision, results: list[dict]) -> dict:
"""semif-serve's averaging over log p (the model returns probabilities, not logits): the mean """semif-serve's averaging over log p (the model returns probabilities, not logits): the mean
per option id, renormalised. The per-ordering results ride along unchanged.""" per option id, renormalised. The per-ordering results ride along unchanged."""
@@ -259,12 +285,11 @@ def create_app(settings: Settings, engine: Any, executor: ThreadPoolExecutor | N
def planned(state: State, decisions: list[Decision], workload: str | None) -> list[list[Slot]]: def planned(state: State, decisions: list[Decision], workload: str | None) -> list[list[Slot]]:
check_request(state, decisions, workload) check_request(state, decisions, workload)
waves = plan_waves(decisions) rows = sum(map(ordering_count, decisions)) # counted before anything is built
rows = sum(map(len, waves))
if not 1 <= rows <= settings.max_decisions: if not 1 <= rows <= settings.max_decisions:
raise ApiError(422, "invalid_request", raise ApiError(422, "invalid_request",
f"this request scores {rows} rows; the limit is 1..{settings.max_decisions}") f"this request scores {rows} rows; the limit is 1..{settings.max_decisions}")
return waves return plan_waves(decisions)
def run(state: State, decisions: list[Decision], waves: list[list[Slot]]) -> tuple[list[dict], dict]: def run(state: State, decisions: list[Decision], waves: list[list[Slot]]) -> tuple[list[dict], dict]:
"""Every call of one request, back to back under the lock (INV-2); results in request order.""" """Every call of one request, back to back under the lock (INV-2); results in request order."""
@@ -285,13 +310,13 @@ def create_app(settings: Settings, engine: Any, executor: ThreadPoolExecutor | N
tokens.append(response["usage"]["input_tokens"]) tokens.append(response["usage"]["input_tokens"])
inference_ms += response["timing"]["inference_ms"] inference_ms += response["timing"]["inference_ms"]
for f, s in zip(fields, chunk): for f, s in zip(fields, chunk):
by_slot[(s.decision, s.k)] = row_result(s, response, sha, f, len(chunk), index, model) by_slot[(s.decision, s.k)] = built(row_result, s, response, sha, f, len(chunk), index, model)
out = [] out = []
for i, d in enumerate(decisions): for i, d in enumerate(decisions):
if d.orderings == "none": if d.orderings == "none":
out.append(by_slot[(i, 0)]) out.append(by_slot[(i, 0)])
else: else:
out.append(combine(d, [by_slot[(i, k)] for k in range(len(ordering_perms(d)))])) out.append(built(combine, d, [by_slot[(i, k)] for k in range(ordering_count(d))]))
timing = {"total_seconds": time.perf_counter() - started, "batch_size": len(by_slot), timing = {"total_seconds": time.perf_counter() - started, "batch_size": len(by_slot),
"calls": len(sizes), "questions_per_call": sizes, "input_tokens": tokens, "calls": len(sizes), "questions_per_call": sizes, "input_tokens": tokens,
"inference_seconds": inference_ms / 1000} "inference_seconds": inference_ms / 1000}
@@ -80,6 +80,8 @@ class TorchEngine:
self._torch, self._engine, self._inference, self._tokenizer = torch, engine, inference, tokenizer self._torch, self._engine, self._inference, self._tokenizer = torch, engine, inference, tokenizer
self.metadata, self._settings = metadata, settings self.metadata, self._settings = metadata, settings
self._release_above = release_above_bytes self._release_above = release_above_bytes
# Asked once, here, on the inference thread: /health then reads only allocator counters.
self._device_name = torch.cuda.get_device_name(0) if settings.device == "cuda" else None
@classmethod @classmethod
def load(cls, settings: Settings, *, torch: Any = None, def load(cls, settings: Settings, *, torch: Any = None,
@@ -167,7 +169,7 @@ class TorchEngine:
info = dict(self.metadata) info = dict(self.metadata)
if self._settings.device == "cuda": if self._settings.device == "cuda":
cuda = self._torch.cuda cuda = self._torch.cuda
info["device_name"] = cuda.get_device_name(0) info["device_name"] = self._device_name
info["allocated_gib"] = round(cuda.memory_allocated(0) / 2**30, 3) info["allocated_gib"] = round(cuda.memory_allocated(0) / 2**30, 3)
info["reserved_gib"] = round(cuda.memory_reserved(0) / 2**30, 3) info["reserved_gib"] = round(cuda.memory_reserved(0) / 2**30, 3)
if hasattr(cuda, "max_memory_reserved"): if hasattr(cuda, "max_memory_reserved"):
@@ -181,12 +183,15 @@ class TorchEngine:
self._torch.cuda.empty_cache() self._torch.cuda.empty_cache()
def predict(self, request: dict) -> tuple[dict, str]: def predict(self, request: dict) -> tuple[dict, str]:
"""One call: the model's own predict(), then the prompt hash (INV-8). Returns (response, sha).""" """One call: the prompt hash (INV-8; this also runs the model's own validation before any GPU
work), then the model's own predict(). Returns (response, sha)."""
sha = hashlib.sha256(self._prompt_text(request).encode()).hexdigest()
try: try:
response = self._engine.predict(request) response = self._engine.predict(request)
sha = hashlib.sha256(self._prompt_text(request).encode()).hexdigest() except ValueError as exc:
except ValueError: # Usually the token limit, found before the forward; but transformers raises ValueError
raise # the model's validation: raised before any GPU work # from inside a forward too. Either way, keep the message and release like any failure.
failure, message = ValueError, str(exc)
except self._torch.cuda.OutOfMemoryError as exc: except self._torch.cuda.OutOfMemoryError as exc:
failure, message = OutOfMemory, _first_line(exc) or "CUDA out of memory" failure, message = OutOfMemory, _first_line(exc) or "CUDA out of memory"
except Exception as exc: # noqa: BLE001 — every other failure is released and reported below except Exception as exc: # noqa: BLE001 — every other failure is released and reported below
@@ -464,3 +464,28 @@ def test_the_app_uses_the_executor_it_is_given():
assert client.post("/decide", json=ROW, headers=AUTH).status_code == 200 assert client.post("/decide", json=ROW, headers=AUTH).status_code == 200
assert engine.threads == {home} assert engine.threads == {home}
executor.shutdown() executor.shutdown()
class InfiniteEngine(FakeEngine):
"""A server-side defect: a probability of +inf, which breaks the averaging arithmetic."""
def predict(self, request):
response, sha = super().predict(request)
for answer in response["answers"].values():
first = next(iter(answer["probabilities"]))
answer["probabilities"][first] = float("inf")
return response, sha
def test_a_failure_while_building_the_response_is_500_never_a_422():
response = make_client(InfiniteEngine()).post("/decide", json={**ROW, "orderings": "rotations"}, headers=AUTH)
assert response.status_code == 500
assert response.json()["error"]["code"] == "scoring_failed"
def test_a_huge_orderings_request_is_refused_with_its_true_row_count():
engine = FakeEngine()
body = {"state": "s", "decisions": [{**d, "orderings": "rotations"} for d in decisions(3000, opts(16))]}
response = make_client(engine, max_body_bytes=16 * 2**20).post("/decide/shared", json=body, headers=AUTH)
assert response.status_code == 422
assert "48000 rows" in response.json()["error"]["message"]
assert engine.calls == []
@@ -306,12 +306,35 @@ def test_any_other_failure_is_logged_released_and_raised_unchained_as_scoring_fa
assert FakeTorch.cuda.empties == [True] assert FakeTorch.cuda.empties == [True]
def test_value_errors_pass_through_untouched(): def test_a_value_error_keeps_its_message_but_is_released_and_raised_unchained():
def bad(_request): """A ValueError can come from inside the forward too (a transformers shape check): it must not pin
raise ValueError("Example has 9000 tokens, above 8192; truncation is forbidden") the failed forward's tensors through a chained traceback (review 2026-09-30)."""
with pytest.raises(ValueError, match="truncation is forbidden"): engine = direct(failing(lambda: ValueError("Example has 9000 tokens, above 8192; truncation is forbidden")))
direct(bad).predict(REQUEST) with pytest.raises(ValueError, match="truncation is forbidden") as info:
assert FakeTorch.cuda.empties == [] 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)]) @pytest.mark.parametrize("reserved, released", [(2 * 2**30, True), (2**30, False), (2**30 - 1, False)])