feat(intern-decision-serve): Intern-Decision-4B behind semif-serve's HTTP surface
Contract, service and tests (fake engine, no GPU). Scores through the checkpoint's own inference.py (DecisionEngine.predict, sha256-pinned); maps semif decisions onto Jev choice questions, packs /decide/shared into calls of at most 16, runs orderings in waves, and keeps semif's error mapping, admission, body limit and hard VRAM cap. Deltas from semif-serve are listed in the contract.
This commit is contained in:
@@ -0,0 +1,6 @@
|
||||
*
|
||||
!pyproject.toml
|
||||
!uv.lock
|
||||
!src/
|
||||
**/__pycache__
|
||||
src/*.egg-info
|
||||
@@ -0,0 +1,4 @@
|
||||
.venv/
|
||||
.pytest_cache/
|
||||
__pycache__/
|
||||
*.egg-info/
|
||||
@@ -0,0 +1,51 @@
|
||||
# syntax=docker/dockerfile:1
|
||||
# intern-decision-serve: Intern-Decision-4B, scored by the checkpoint's own inference.py, behind
|
||||
# semif-serve's HTTP surface. Contract: intern-decision-serve.contract.md.
|
||||
# docker build -t intern-decision-serve:<version> .
|
||||
# Weights AND inference.py are NOT in the image: both are read from the mounted HF cache at the
|
||||
# pinned revision, offline (INV-5); inference.py's sha256 is checked before it is imported (INV-3).
|
||||
|
||||
# The dependency manifest with the service's own version blanked to 0.0.0, so a version bump
|
||||
# leaves these two files byte-identical and the ~4 GB torch/CUDA layer below stays cached.
|
||||
FROM python:3.12-slim-bookworm AS deps
|
||||
WORKDIR /deps
|
||||
COPY pyproject.toml uv.lock ./
|
||||
RUN python - <<'PY'
|
||||
import re, pathlib
|
||||
p = pathlib.Path("pyproject.toml")
|
||||
p.write_text(re.sub(r'(?m)^version = "[^"]+"', 'version = "0.0.0"', p.read_text(), count=1))
|
||||
l = pathlib.Path("uv.lock")
|
||||
l.write_text(re.sub(r'(name = "intern-decision-serve"\nversion = )"[^"]+"', r'\1"0.0.0"', l.read_text(), count=1))
|
||||
PY
|
||||
|
||||
FROM python:3.12-slim-bookworm
|
||||
# The model plus Qwen3.5's fast kernels (fla + causal-conv1d): the bench stack.
|
||||
ARG EXTRAS="--extra model --extra fast"
|
||||
COPY --from=ghcr.io/astral-sh/uv:0.6.9 /uv /bin/uv
|
||||
ENV UV_COMPILE_BYTECODE=1 UV_LINK_MODE=copy UV_PYTHON_DOWNLOADS=never
|
||||
# The fast extra needs a C compiler AT RUNTIME: triton builds its CUDA driver shim on first use,
|
||||
# and without gcc the warm-up dies with "Failed to find C compiler" (semif-serve, 2026-09-27).
|
||||
RUN apt-get update && apt-get install -y --no-install-recommends ca-certificates \
|
||||
&& if echo "$EXTRAS" | grep -q -- '--extra fast'; then \
|
||||
apt-get install -y --no-install-recommends gcc libc6-dev; fi \
|
||||
&& rm -rf /var/lib/apt/lists/*
|
||||
WORKDIR /app
|
||||
COPY --from=deps /deps/pyproject.toml /deps/uv.lock ./
|
||||
RUN --mount=type=cache,target=/root/.cache/uv \
|
||||
uv sync --frozen --no-dev $EXTRAS --no-install-project
|
||||
COPY pyproject.toml uv.lock ./
|
||||
COPY src ./src
|
||||
RUN uv sync --frozen --no-dev $EXTRAS --no-editable --no-cache
|
||||
RUN groupadd --system --gid 10001 intern \
|
||||
&& useradd --system --uid 10001 --gid 10001 --no-create-home --shell /usr/sbin/nologin intern
|
||||
USER intern
|
||||
ENV PATH=/app/.venv/bin:$PATH \
|
||||
HF_HOME=/hf \
|
||||
HF_HUB_OFFLINE=1 \
|
||||
TRANSFORMERS_OFFLINE=1 \
|
||||
HF_HUB_DISABLE_TELEMETRY=1 \
|
||||
TRITON_CACHE_DIR=/tmp/triton-cache \
|
||||
NVIDIA_DRIVER_CAPABILITIES=compute,utility
|
||||
EXPOSE 8000
|
||||
# One worker (INV-2): the model and the inference lock live in this one process.
|
||||
CMD ["uvicorn", "intern_decision_serve.main:app_from_env", "--factory", "--host", "0.0.0.0", "--port", "8000", "--workers", "1"]
|
||||
@@ -0,0 +1,271 @@
|
||||
---
|
||||
title: intern-decision-serve
|
||||
kind: module-contract
|
||||
status: draft
|
||||
owner: infra-ops
|
||||
created: 2026-09-30
|
||||
replaces: semif-serve 0.1.4 (services/semif-serve/semif-serve.contract.md), external surface kept
|
||||
depends_on:
|
||||
- internlm/Intern-Decision-4B at revision 0e5e6aa7d6d750e2b1504ba11a8136cb58aeb3cd, BF16 (Apache-2.0)
|
||||
- that snapshot's own inference.py (DecisionEngine.predict), sha256 c904e2c67ca0775621a22375ee373d2ba30b52117cda870c6c9ef74143b29863
|
||||
- torch 2.10.0+cu128, transformers 5.17.0, flash-linear-attention 0.5.2, causal-conv1d 1.7.0 (the 2026-09-30 bench stack)
|
||||
---
|
||||
|
||||
# intern-decision-serve: Intern-Decision-4B behind semif-serve's HTTP surface
|
||||
|
||||
## Purpose
|
||||
|
||||
Prime, 2026-09-30: "replace semif with intern-decision now". The bench
|
||||
(`docs/pfi/jev-candidates-bench-2026-09-30.md`) picked Intern-Decision-4B on its own runtime.
|
||||
This service loads the model once and scores every request with the checkpoint's **own**
|
||||
`inference.py` (`DecisionEngine.predict`). It re-implements neither the prompt nor the readout.
|
||||
It keeps semif-serve's external surface, so a caller written for semif-serve works unchanged.
|
||||
It only maps semif-shaped requests onto the model's Jev request schema and maps the answers back.
|
||||
Every deliberate difference is listed under **Deltas from semif-serve**.
|
||||
|
||||
## Endpoints (same as semif-serve)
|
||||
|
||||
Every POST takes and returns JSON and needs `Authorization: Bearer <token>`. `GET /health` is open.
|
||||
|
||||
| method | path | body | success |
|
||||
|---|---|---|---|
|
||||
| GET | `/health` | none | 200 `{status: "ok", model, vram_cap_gib, max_tokens, max_decisions, max_questions_per_call: 16, chunking, workloads: []}` |
|
||||
| POST | `/decide` | `{id, state, question, options[2..16], orderings?, workload?}` | 200 one decision result |
|
||||
| POST | `/decide/shared` | `{state, decisions: [{id, question, options, orderings?}], workload?}` | 200 `{results: [...], timing: {...}}`, results in request order |
|
||||
|
||||
`options` items are `{id, description}`. Validation is SemIf's, re-stated here because SemIf is
|
||||
gone. `id` and `question` are nonempty strings. `state` is a nonempty string, object or array,
|
||||
and must be finite JSON. There are 2..16 options, each with a string `id` and a string
|
||||
`description`, and option ids are unique within a decision. Decision ids are unique within a
|
||||
request. A violation is a 422.
|
||||
|
||||
## Mapping onto the model (the seam)
|
||||
|
||||
- **One call** is one `predict()` request:
|
||||
`{"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.
|
||||
- **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:**
|
||||
the model prints `A = <id>: <description>`.
|
||||
- **`/decide`** is one call with one question.
|
||||
- **`/decide/shared`** is packed into calls of **at most 16 questions**, the model's own limit
|
||||
(`validate_request`). The packing is greedy and in request order: questions 1–16, then 17–32,
|
||||
and so on. Every call of a request runs back to back under the inference lock. `/health` reports
|
||||
`max_questions_per_call: 16` and the rule in `chunking`.
|
||||
- **Orderings.** A decision may set `orderings: "none" | "rotations" | "all"`, with semif's
|
||||
meaning. `rotations` gives the n cyclic shifts, the caller's order first. `all` gives the n!
|
||||
permutations, the caller's order first, and is 422 above 4 options. Ordering k of every decision
|
||||
that has more than k orderings forms **wave k**. Each wave is packed as above. So wave 0 is the
|
||||
request exactly as written, and no prompt ever holds two orderings of the same decision. Every
|
||||
ordering counts toward `MAX_DECISIONS`.
|
||||
- **The result for one ordering** (and for every plain decision):
|
||||
```
|
||||
{id, option_ids, probabilities, top, confidence, calibration, native,
|
||||
input_tokens, prompt_sha256, prompt_version, model, readout, probability_status,
|
||||
call: {index, field, questions}}
|
||||
```
|
||||
- `native` is the answer object `predict()` returned for that field, unchanged.
|
||||
- `probabilities` is `native.probabilities` listed in `option_ids` order.
|
||||
- `top` is `native.choice` and `confidence` is `native.confidence`.
|
||||
- `calibration` is the `calibration` object from `predict()`.
|
||||
- `input_tokens` and `prompt_sha256` describe the whole call.
|
||||
- `/decide` results (plain) also carry `total_seconds` and `forward_seconds`. `forward_seconds`
|
||||
is `predict()`'s `timing.inference_ms` / 1000.
|
||||
- **Averaged result:** semif's shape, `{id, option_ids, combined: {method, orderings, probabilities, top, agreement, spread}, orderings: [...]}`.
|
||||
- The per-ordering ids are `<id>#o<k>`.
|
||||
- `combined.probabilities` renormalises the per-option mean of `log p`. A probability of
|
||||
exactly 0 is floored at 1e-300 before the log.
|
||||
- `combined.top` is the first maximum in the caller's order.
|
||||
- `agreement` is the share of orderings whose `top` equals `combined.top`.
|
||||
- `spread` holds each option's min and max `probabilities` across orderings.
|
||||
- Temperature scaling is one monotone transform per ordering, so `combined.top` and
|
||||
`agreement` are what combining T=1 scores would give.
|
||||
- **`/decide/shared` timing:** `{total_seconds, batch_size (orderings scored), calls, questions_per_call: [...], input_tokens: [...], inference_seconds}`.
|
||||
|
||||
## Invariants
|
||||
|
||||
- **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.
|
||||
If an answer lacks one of the decision's option ids, that is a 500 `scoring_failed`, never a
|
||||
guess.
|
||||
- **INV-2 one model, one inference at a time.** The model loads at startup, and a process-wide
|
||||
lock serialises every request's calls (all of a request's chunks run inside one hold). The app
|
||||
runs one worker. Calls run off the event loop, so `/health` answers during one.
|
||||
- **INV-3 fail-closed startup.** Before the service serves, all of these must hold:
|
||||
- `inference.py` in the checkpoint hashes to the pinned sha256 (it is executed code, loaded
|
||||
from a data mount);
|
||||
- the checkpoint path ends in `snapshots/<pinned revision>`;
|
||||
- the model sits on CUDA, and torch's arch list has the card's `sm_XY`;
|
||||
- one warm-up decision scores.
|
||||
|
||||
`device=cpu` is allowed only when set explicitly.
|
||||
- **INV-4 VRAM cap.** `VRAM_CAP_GIB`, when set, becomes `torch.cuda.set_per_process_memory_fraction`
|
||||
**before** the weights load.
|
||||
- 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,
|
||||
runs `gc.collect()` and `empty_cache()`, and raises unchained. The process stays up.
|
||||
- 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()`.
|
||||
- Any other failure except `ValueError` is logged with its traceback and released the same
|
||||
way, then raised unchained as `ScoringFailed`.
|
||||
- **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.
|
||||
- **INV-6 constant-time auth.** The token is compared with `hmac.compare_digest`. It must be at
|
||||
least 32 visible ASCII characters (33–126), or startup refuses it.
|
||||
- **INV-7 the text-only model is the same model.**
|
||||
- The service takes no images, and neither did semif. So after the first warm-up the engine
|
||||
replaces the vision tower (`model.model.visual`, 0.62 GiB) with a stub that raises if it is
|
||||
ever called.
|
||||
- It then scores the warm-up again. It refuses to start unless the answer is bit-identical to
|
||||
the first one.
|
||||
- `KEEP_VISION=1` keeps the tower.
|
||||
- **INV-8 honest prompt hash.** `prompt_sha256` is the sha256 of the chat-template text rendered
|
||||
from `inference.compile_row(...)` with the same arguments `HFBackend.encode` uses. At startup,
|
||||
that text must tokenise to exactly the `input_tokens` `predict()` reports for the warm-up.
|
||||
Otherwise the service refuses to start rather than hash a prompt the model never saw.
|
||||
|
||||
## Limits and errors
|
||||
|
||||
- `MAX_TOKENS` (default 8192, the model's own default) is `DecisionEngine(max_length=...)`, and
|
||||
applies to a **whole call**: the state plus all its questions. A longer call is a 422. It is
|
||||
never truncated; the model raises.
|
||||
- `MAX_DECISIONS` (default 64) caps the orderings scored per request, counted after expansion.
|
||||
A request must score 1..max of them, else 422.
|
||||
- The body may be at most `MAX_BODY_BYTES` (default 1 MiB), else 413. This is checked before
|
||||
each chunk is kept. A declared `Content-Length` is trusted only as ASCII digits.
|
||||
- **Admission:** at most `MAX_QUEUE` (default 32) POSTs may be queued or scoring at once. The
|
||||
next one gets 429 `busy` before its body is read.
|
||||
- `workload`: there is no per-workload table, and the model's own calibration always applies.
|
||||
So any non-null `workload` is a 422, exactly as semif-serve behaved with its deployed empty
|
||||
table.
|
||||
|
||||
| status | code | when |
|
||||
|---|---|---|
|
||||
| 401 | `unauthorized` | missing or wrong bearer |
|
||||
| 413 | `request_too_large` | body over the limit |
|
||||
| 422 | `invalid_request` | bad JSON or shape; a SemIf-rule violation; a model `ValueError` (token limit, reserved `<decision>` marker in the input, ...); a `workload`; `all` over 4 options; a row count outside 1..`MAX_DECISIONS` |
|
||||
| 429 | `busy` | `MAX_QUEUE` requests in progress |
|
||||
| 503 | `out_of_memory` | CUDA OOM in any call of the request |
|
||||
| 500 | `scoring_failed` | any other model failure, including one while building the response |
|
||||
|
||||
The error body is `{error: {code, message}}`.
|
||||
|
||||
## Configuration (env, prefix `INTERN_DECISION_`)
|
||||
|
||||
- `API_TOKEN` is required.
|
||||
- `CHECKPOINT` defaults to
|
||||
`/hf/hub/models--internlm--Intern-Decision-4B/snapshots/<revision>`.
|
||||
- `DEVICE` defaults to `cuda`.
|
||||
- `VRAM_CAP_GIB` has no default: unset means uncapped, and when set it must be finite and > 0.
|
||||
- `MAX_TOKENS`, `MAX_DECISIONS`, `MAX_BODY_BYTES`, `MAX_QUEUE` and `RELEASE_SLACK_MIB` are
|
||||
integers; each must be ≥ 1, except the slack, which must be ≥ 0.
|
||||
- `KEEP_VISION` is `0` or `1`.
|
||||
|
||||
A bad value is refused at startup with a `ValueError` naming the variable. In the stack's `.env`
|
||||
on the host, the cap is the single knob `VRAM_CAP_GIB`.
|
||||
|
||||
## Deltas from semif-serve (deliberate)
|
||||
|
||||
1. **Prompt and model.** The prompt is Intern-Decision's own Jev prompt.
|
||||
- **Option ids are shown to the model** (`A = <id>: <description>`); SemIf showed only the
|
||||
descriptions. An option id is therefore part of the question, so give options meaningful
|
||||
or neutral ids.
|
||||
- In the bench negative control, the ids-in-prompt cue made the top stay on 10/144 rows
|
||||
after the descriptions moved, against SemIf's 14/144.
|
||||
2. **Shared requests are one prompt, not independent rows.** The questions in a call are asked
|
||||
together.
|
||||
- A decision's answer can depend on the other questions in its call and on their order. In
|
||||
the bench, Wyrd scored 79/84 asked one decision at a time and 77/84 with a turn's 4
|
||||
decisions in one prompt.
|
||||
- SemIf only shared a KV prefix, so each of its rows was independent.
|
||||
- Calls hold at most 16 questions, and the chunk boundaries follow request order.
|
||||
3. **`probabilities` are temperature-scaled** by the checkpoint's shipped calibration
|
||||
(T = 1.99241824).
|
||||
- `probability_status` says so, and `calibration` carries the method and T.
|
||||
- SemIf's were raw softmax, labelled uncalibrated. The argmax is the same either way.
|
||||
- This is the vendor's calibration on the vendor's data, not ours.
|
||||
4. **`option_logits` does not exist.** `predict()` does not expose logits. T=1 scores could only
|
||||
be derived up to a constant, which would be a derived number, not logits.
|
||||
5. **New fields:** `top`, `confidence`, `calibration`, `native`, `call`.
|
||||
6. **`prompt_sha256` and `input_tokens` describe the whole call**, shared by every decision in
|
||||
it. SemIf's were per row.
|
||||
7. **`prompt_version`** names the pinned `inference.py`. **`readout`** and **`model`** describe
|
||||
Intern-Decision.
|
||||
8. **`workload` is always a 422, and `/health.workloads` is `[]`.** There is no per-workload
|
||||
temperature table. This matches the deployed semif, whose table was empty.
|
||||
9. **`/health`**: `semif_commit` is gone; the model's pins live in `model`. It adds
|
||||
`max_questions_per_call` and `chunking`.
|
||||
10. **`/decide/shared` timing:** SemIf's prefix-cache fields cannot exist, because there is no
|
||||
prefix cache: `prefix_tokens`, `prefill_seconds`, `replicate_seconds`,
|
||||
`suffix_forward_seconds`, `true_suffix_tokens`, `padded_suffix_tokens` and `encode_seconds`.
|
||||
The timing adds `calls`, `questions_per_call`, `input_tokens` and `inference_seconds`.
|
||||
`total_seconds` and `batch_size` keep their meaning.
|
||||
11. **`MAX_TOKENS` is per call** (state plus up to 16 questions) and defaults to 8192. SemIf's
|
||||
was per row and defaulted to 4096.
|
||||
12. **Orderings run in waves**, one forward pass per wave per chunk. SemIf batched every
|
||||
ordering in one shared forward.
|
||||
13. **Ties.** Per-ordering `top` uses the model's argmax, whose exact tie goes to the smaller
|
||||
option id string. SemIf used the first maximum in the ordering. `combined.top` keeps semif's
|
||||
rule.
|
||||
14. **Images.** The vision tower is dropped (INV-7). The surface never took images.
|
||||
|
||||
## Tests (TDD, fake engine: no torch, no model)
|
||||
|
||||
- **Auth.** A POST without the right bearer is 401 and never reaches the engine. `/health`
|
||||
needs no auth. A short or non-visible-ASCII token is refused at startup.
|
||||
- **Mapping.**
|
||||
- `/decide` sends one call with field `q` and the criteria in the caller's order.
|
||||
- `/decide/shared` sends `q1..qN` over the shared state.
|
||||
- 17 questions become 2 calls (16 + 1), and 40 become 3 (16 + 16 + 8). Results come back in
|
||||
request order, and `call.index` and `call.field` are right.
|
||||
- Decision ids never appear in a call.
|
||||
- **Result.** `probabilities` follows `option_ids`. `top`, `confidence`, `calibration`, `native`
|
||||
and `input_tokens` are passed through. An answer missing an option id is a 500.
|
||||
- **Orderings.**
|
||||
- `rotations` puts each option in each position once, and the waves never repeat a decision
|
||||
within a call.
|
||||
- `all` is n!, and 422 above 4 options.
|
||||
- A position bias cancels exactly.
|
||||
- `agreement` and `spread` are computed from the orderings.
|
||||
- A mixed request keeps plain results unchanged, and wave 0 equals the plain request.
|
||||
- Orderings count toward the cap.
|
||||
- **Validation (422).** Fewer than 2 or more than 16 options, duplicate option ids, duplicate
|
||||
decision ids, an empty id or question, an empty or non-finite state, a `workload`, a row count
|
||||
outside 1..max, malformed JSON, and a model `ValueError`.
|
||||
- **Limits.** A body over the limit is 413. A queue past `MAX_QUEUE` is 429 before the body is
|
||||
read.
|
||||
- **Engine failures.** An engine `OutOfMemory` is 503, and any other failure is 500.
|
||||
- **Concurrency.** Requests are serialised: two never overlap inside the engine, and the
|
||||
chunks of one request are not interleaved with another's. `/health` answers while a call is
|
||||
blocked.
|
||||
- **Engine against a fake torch.**
|
||||
- An OOM is re-raised unchained, and `empty_cache` runs only after the failed call's tensors
|
||||
are freed.
|
||||
- A RuntimeError saying "out of memory" becomes `OutOfMemory`.
|
||||
- Another failure becomes `ScoringFailed`, unchained.
|
||||
- `ValueError` passes through.
|
||||
- A burst over baseline plus slack is released, and one at or under it is left alone.
|
||||
- **Engine load (fake).**
|
||||
- A wrong `inference.py` hash, or a checkpoint path that is not the pinned snapshot, refuses
|
||||
to start.
|
||||
- The cap is applied before the engine is constructed.
|
||||
- The vision swap refuses to start when the warm-up changes.
|
||||
- The prompt-hash check refuses to start on a token-count mismatch.
|
||||
- **Config.** Every value is validated.
|
||||
|
||||
## Acceptance (fv-ml1, real model; not unit tests)
|
||||
|
||||
1. **Positive control:** through the service, the bench's native numbers on the pooled 259 rows
|
||||
and on Wyrd, single ordering (bench: 240/259, Wyrd 79/84; floor: the pooled set resolves
|
||||
±4 pts, and 0 labels moved across 4 restarts). Also a row-by-row comparison against the
|
||||
bench's own rows.
|
||||
2. **Negative control:** descriptions rotated one place. The top follows the moved description
|
||||
(bench: 122/144 follow, 10/144 same top).
|
||||
3. A-vs-A repeat stability, within the process and across restarts.
|
||||
4. Latency at our shape: 21 binary criteria, and 16 criteria over the ~3,900-token state, both
|
||||
server-side and from nh3-dev.
|
||||
5. Resident and peak VRAM (nvidia-smi and torch), and the cap chosen from them.
|
||||
6. An over-cap request is a 503, memory returns to baseline, and the service keeps answering.
|
||||
7. A `/decide/shared` with more than 16 questions is chunked, and its answers equal the same
|
||||
questions asked chunk by chunk.
|
||||
8. 401 without the token, and 429 past the queue.
|
||||
@@ -0,0 +1,52 @@
|
||||
[project]
|
||||
name = "intern-decision-serve"
|
||||
version = "0.1.0"
|
||||
description = "Intern-Decision-4B (its own inference.py) behind semif-serve's HTTP surface"
|
||||
requires-python = ">=3.12"
|
||||
dependencies = [
|
||||
"fastapi==0.118.0",
|
||||
"uvicorn==0.37.0",
|
||||
]
|
||||
|
||||
[project.optional-dependencies]
|
||||
# The real engine: the 2026-09-30 bench stack (semif-serve 0.1.4's torch/transformers plus what
|
||||
# the checkpoint's own inference.py imports: PIL and the Qwen3.5 processor, which needs torchvision).
|
||||
model = [
|
||||
"torch==2.10.0",
|
||||
"torchvision==0.25.0",
|
||||
"transformers==5.17.0",
|
||||
"pillow==12.3.0",
|
||||
]
|
||||
# Qwen3.5's fast kernels, as in the bench image (without them transformers runs its slower
|
||||
# reference PyTorch paths, and the bench numbers were measured with them).
|
||||
fast = [
|
||||
"flash-linear-attention==0.5.2",
|
||||
"causal-conv1d @ https://github.com/Dao-AILab/causal-conv1d/releases/download/v1.7.0/causal_conv1d-1.7.0+cu12torch2.10cxx11abiTRUE-cp312-cp312-linux_x86_64.whl ; sys_platform == 'linux' and platform_machine == 'x86_64'",
|
||||
]
|
||||
|
||||
[dependency-groups]
|
||||
dev = ["pytest==8.4.2", "httpx==0.28.1"]
|
||||
|
||||
[build-system]
|
||||
requires = ["setuptools>=68"]
|
||||
build-backend = "setuptools.build_meta"
|
||||
|
||||
[tool.setuptools.packages.find]
|
||||
where = ["src"]
|
||||
|
||||
[tool.pytest.ini_options]
|
||||
testpaths = ["tests"]
|
||||
|
||||
[[tool.uv.index]]
|
||||
name = "pytorch-cu128"
|
||||
url = "https://download.pytorch.org/whl/cu128"
|
||||
explicit = true
|
||||
|
||||
[tool.uv.sources]
|
||||
torch = { index = "pytorch-cu128" }
|
||||
torchvision = { index = "pytorch-cu128" }
|
||||
|
||||
[tool.uv]
|
||||
# Hold the transitive pins to the 2026-09-30 bench image (semif-serve:0.1.4's lock), so the
|
||||
# service runs the stack its acceptance numbers are compared against.
|
||||
constraint-dependencies = ["numpy==2.2.6", "huggingface-hub==1.31.0", "regex==2026.9.10", "tokenizers==0.23.2", "safetensors==0.8.0"]
|
||||
@@ -0,0 +1,331 @@
|
||||
"""intern-decision-serve HTTP layer. Contract: intern-decision-serve.contract.md.
|
||||
|
||||
The request surface is semif-serve's. Each semif decision becomes one Jev `choice` question; the
|
||||
questions of a request are asked through the engine (the checkpoint's own DecisionEngine.predict)
|
||||
in calls of at most 16, and the answers are re-keyed into semif's result shape.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import hmac
|
||||
import itertools
|
||||
import json
|
||||
import math
|
||||
import threading
|
||||
import time
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Literal
|
||||
|
||||
from fastapi import FastAPI, Request
|
||||
from fastapi.concurrency import run_in_threadpool
|
||||
from fastapi.responses import JSONResponse
|
||||
from pydantic import BaseModel, ConfigDict, ValidationError
|
||||
|
||||
from .config import INFERENCE_PY_SHA256, MAX_QUESTIONS_PER_CALL, Settings
|
||||
from .errors import OutOfMemory, ScoringFailed
|
||||
|
||||
OPEN_PATHS = frozenset({"/health"})
|
||||
State = str | dict[str, Any] | list[Any]
|
||||
|
||||
MIN_OPTIONS, MAX_OPTIONS = 2, 16 # SemIf's rule (its 16 answer letters), kept for the surface
|
||||
MAX_OPTIONS_FOR_ALL = 4 # "orderings": "all" asks n! orderings
|
||||
LOG_FLOOR = 1e-300 # a probability of exactly 0 before the log (combine)
|
||||
|
||||
PROMPT_VERSION = f"intern-decision-jev/inference.py@{INFERENCE_PY_SHA256[:12]}"
|
||||
READOUT = ("logits at the position before each <decision> marker, softmax over the field's answer "
|
||||
"symbols (the checkpoint's own inference.py, DecisionEngine.predict)")
|
||||
PROBABILITY_STATUS = ("temperature-scaled by the checkpoint's own shipped calibration (see calibration); "
|
||||
"vendor-fitted, not fitted on our workloads")
|
||||
CHUNKING = (f"/decide/shared questions are packed greedily, in request order, into calls of at most "
|
||||
f"{MAX_QUESTIONS_PER_CALL} (1-{MAX_QUESTIONS_PER_CALL}, {MAX_QUESTIONS_PER_CALL + 1}-"
|
||||
f"{2 * MAX_QUESTIONS_PER_CALL}, ...); each call is one prompt, so the questions in a call are "
|
||||
f"asked together. With orderings, ordering k of every decision forms wave k, packed the same way.")
|
||||
|
||||
|
||||
class Option(BaseModel):
|
||||
model_config = ConfigDict(extra="ignore")
|
||||
id: str
|
||||
description: str
|
||||
|
||||
|
||||
class Decision(BaseModel):
|
||||
model_config = ConfigDict(extra="ignore")
|
||||
id: str
|
||||
question: str
|
||||
options: list[Option]
|
||||
orderings: Literal["none", "rotations", "all"] = "none"
|
||||
|
||||
|
||||
class DecideBody(Decision):
|
||||
state: State
|
||||
workload: str | None = None
|
||||
|
||||
|
||||
class SharedBody(BaseModel):
|
||||
model_config = ConfigDict(extra="ignore")
|
||||
state: State
|
||||
decisions: list[Decision]
|
||||
workload: str | None = None
|
||||
|
||||
|
||||
class ApiError(Exception):
|
||||
def __init__(self, status: int, code: str, message: str):
|
||||
super().__init__(message)
|
||||
self.status, self.code, self.message = status, code, message
|
||||
|
||||
|
||||
def error(status: int, code: str, message: str) -> JSONResponse:
|
||||
return JSONResponse(status_code=status, content={"error": {"code": code, "message": message}})
|
||||
|
||||
|
||||
def _first_error(exc: ValidationError) -> str:
|
||||
first = exc.errors()[0]
|
||||
where = ".".join(str(p) for p in first.get("loc", ())) or "body"
|
||||
return f"{where}: {first.get('msg', 'invalid')}"
|
||||
|
||||
|
||||
async def read_limited(stream, declared: str | None, limit: int) -> bytes:
|
||||
"""Read a request body, refusing it once it would exceed `limit` bytes. The check runs BEFORE a
|
||||
chunk is kept, and nothing after the crossing chunk is read. A declared length is trusted only
|
||||
as ASCII digits: `"²".isdigit()` is True but `int("²")` raises."""
|
||||
too_large = ApiError(413, "request_too_large", f"request body exceeds {limit} bytes")
|
||||
if declared is not None and declared.isascii() and declared.isdigit() and int(declared) > limit:
|
||||
raise too_large
|
||||
body = bytearray()
|
||||
async for chunk in stream:
|
||||
if len(body) + len(chunk) > limit:
|
||||
raise too_large
|
||||
body.extend(chunk)
|
||||
return bytes(body)
|
||||
|
||||
|
||||
def check_request(state: State, decisions: list[Decision], workload: str | None) -> None:
|
||||
"""SemIf's row validation, re-stated because SemIf is gone; a violation is a 422."""
|
||||
def bad(message: str) -> ApiError:
|
||||
return ApiError(422, "invalid_request", message)
|
||||
|
||||
if workload is not None:
|
||||
raise bad(f"unknown workload {workload!r}: no per-workload calibration is configured; "
|
||||
"the model's own calibration always applies")
|
||||
if not state:
|
||||
raise bad("state must be a nonempty string, object, or array")
|
||||
try:
|
||||
json.dumps(state, ensure_ascii=False, allow_nan=False)
|
||||
except (TypeError, ValueError):
|
||||
raise bad("state must be finite JSON-compatible data") from None
|
||||
if len({d.id for d in decisions}) != len(decisions):
|
||||
raise bad("decision ids must be unique")
|
||||
for d in decisions:
|
||||
if not d.id or not d.question:
|
||||
raise bad("id and question must be nonempty strings")
|
||||
if not MIN_OPTIONS <= len(d.options) <= MAX_OPTIONS:
|
||||
raise bad(f"decision {d.id!r}: options must contain {MIN_OPTIONS}-{MAX_OPTIONS} entries")
|
||||
if len({o.id for o in d.options}) != len(d.options):
|
||||
raise bad(f"decision {d.id!r}: option ids must be unique")
|
||||
|
||||
|
||||
def ordering_perms(d: Decision) -> list[tuple[int, ...]]:
|
||||
"""Index permutations of the caller's options, the caller's own order first."""
|
||||
n = len(d.options)
|
||||
if d.orderings == "none":
|
||||
return [tuple(range(n))]
|
||||
if d.orderings == "rotations":
|
||||
return [tuple((start + k) % n for k in range(n)) for start in range(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 list(itertools.permutations(range(n)))
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class Slot:
|
||||
"""One question to ask: ordering `k` of decision number `decision`."""
|
||||
decision: int
|
||||
k: int
|
||||
row_id: str
|
||||
option_ids: list[str]
|
||||
question: dict
|
||||
|
||||
|
||||
def plan_waves(decisions: list[Decision]) -> list[list[Slot]]:
|
||||
"""Wave k holds ordering k of every decision that has more than k orderings, in request order.
|
||||
Wave 0 is the request as written; no wave holds two orderings of one decision."""
|
||||
perms = [ordering_perms(d) for d in decisions]
|
||||
waves = []
|
||||
for k in range(max(map(len, perms), default=0)):
|
||||
wave = []
|
||||
for i, (d, ps) in enumerate(zip(decisions, perms)):
|
||||
if k < len(ps):
|
||||
options = [d.options[j] for j in ps[k]]
|
||||
wave.append(Slot(i, k, d.id if d.orderings == "none" else f"{d.id}#o{k}", [o.id for o in options],
|
||||
{"type": "choice", "instructions": d.question,
|
||||
"criteria": {o.id: o.description for o in options}}))
|
||||
waves.append(wave)
|
||||
return waves
|
||||
|
||||
|
||||
def pack(slots: list[Slot]) -> list[list[Slot]]:
|
||||
"""Greedy, in request order: at most MAX_QUESTIONS_PER_CALL questions per call."""
|
||||
return [slots[k:k + MAX_QUESTIONS_PER_CALL] for k in range(0, len(slots), MAX_QUESTIONS_PER_CALL)]
|
||||
|
||||
|
||||
def field_names(count: int) -> list[str]:
|
||||
"""Positional: `q` for a one-question call, `q1`..`qN` otherwise. Decision ids never reach the prompt."""
|
||||
return ["q"] if count == 1 else [f"q{i + 1}" for i in range(count)]
|
||||
|
||||
|
||||
def row_result(slot: Slot, response: dict, sha: str, field: str, questions: int, index: int, model: dict) -> dict:
|
||||
"""INV-1: re-key the model's answer for `field`; never fill in a number it did not return."""
|
||||
answer = response["answers"][field]
|
||||
probs = answer["probabilities"]
|
||||
missing = [i for i in slot.option_ids if i not in probs]
|
||||
if missing:
|
||||
raise ScoringFailed(f"the model's answer for field {field!r} lacks option ids {missing}")
|
||||
return {"id": slot.row_id, "option_ids": slot.option_ids, "probabilities": [probs[i] for i in slot.option_ids],
|
||||
"top": answer["choice"], "confidence": answer["confidence"],
|
||||
"calibration": response["calibration"], "native": answer,
|
||||
"input_tokens": response["usage"]["input_tokens"], "prompt_sha256": sha,
|
||||
"prompt_version": PROMPT_VERSION, "model": model, "readout": READOUT,
|
||||
"probability_status": PROBABILITY_STATUS,
|
||||
"call": {"index": index, "field": field, "questions": questions}}
|
||||
|
||||
|
||||
def combine(d: Decision, results: list[dict]) -> dict:
|
||||
"""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."""
|
||||
option_ids = [o.id for o in d.options]
|
||||
logp: dict[str, list[float]] = {i: [] for i in option_ids}
|
||||
probs: dict[str, list[float]] = {i: [] for i in option_ids}
|
||||
for result in results:
|
||||
for oid, p in zip(result["option_ids"], result["probabilities"]):
|
||||
logp[oid].append(math.log(max(p, LOG_FLOOR)))
|
||||
probs[oid].append(p)
|
||||
means = [sum(logp[i]) / len(logp[i]) for i in option_ids]
|
||||
peak = max(means)
|
||||
weights = [math.exp(m - peak) for m in means]
|
||||
combined = [w / sum(weights) for w in weights]
|
||||
winner = option_ids[combined.index(max(combined))]
|
||||
tops = [r["top"] for r in results]
|
||||
return {"id": d.id, "option_ids": option_ids,
|
||||
"combined": {"method": d.orderings, "orderings": len(results), "probabilities": combined,
|
||||
"top": winner, "agreement": tops.count(winner) / len(tops),
|
||||
"spread": {i: [min(probs[i]), max(probs[i])] for i in option_ids}},
|
||||
"orderings": results}
|
||||
|
||||
|
||||
def create_app(settings: Settings, engine: Any) -> FastAPI:
|
||||
app = FastAPI(title="intern-decision-serve")
|
||||
expected = f"Bearer {settings.api_token}".encode()
|
||||
inference = threading.Lock() # INV-2: one request's calls at a time, off the event loop
|
||||
in_progress = 0 # POSTs admitted and not yet answered
|
||||
|
||||
def admit():
|
||||
nonlocal in_progress
|
||||
if in_progress >= settings.max_queue:
|
||||
raise ApiError(429, "busy", f"{in_progress} requests already in progress (limit {settings.max_queue})")
|
||||
in_progress += 1
|
||||
|
||||
def leave():
|
||||
nonlocal in_progress
|
||||
in_progress -= 1
|
||||
|
||||
@app.middleware("http")
|
||||
async def require_bearer(request: Request, call_next):
|
||||
if request.url.path not in OPEN_PATHS:
|
||||
supplied = request.headers.get("authorization", "").encode()
|
||||
if not hmac.compare_digest(supplied, expected): # INV-6
|
||||
return error(401, "unauthorized", "missing or wrong bearer token")
|
||||
return await call_next(request)
|
||||
|
||||
@app.exception_handler(ApiError)
|
||||
async def api_error(_request: Request, exc: ApiError):
|
||||
return error(exc.status, exc.code, exc.message)
|
||||
|
||||
async def parse(request: Request, model: type[BaseModel]):
|
||||
try:
|
||||
return model.model_validate_json(
|
||||
await read_limited(request.stream(), request.headers.get("content-length"), settings.max_body_bytes))
|
||||
except ValidationError as exc:
|
||||
raise ApiError(422, "invalid_request", _first_error(exc)) from exc
|
||||
|
||||
def planned(state: State, decisions: list[Decision], workload: str | None) -> list[list[Slot]]:
|
||||
check_request(state, decisions, workload)
|
||||
waves = plan_waves(decisions)
|
||||
rows = sum(map(len, waves))
|
||||
if not 1 <= rows <= settings.max_decisions:
|
||||
raise ApiError(422, "invalid_request",
|
||||
f"this request scores {rows} rows; the limit is 1..{settings.max_decisions}")
|
||||
return waves
|
||||
|
||||
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."""
|
||||
with inference:
|
||||
return _run(state, decisions, waves)
|
||||
|
||||
def _run(state: State, decisions: list[Decision], waves: list[list[Slot]]) -> tuple[list[dict], dict]:
|
||||
started = time.perf_counter()
|
||||
model = engine.metadata # static: results never carry live memory numbers
|
||||
by_slot: dict[tuple[int, int], dict] = {}
|
||||
sizes, tokens, inference_ms = [], [], 0.0
|
||||
for chunk in (chunk for wave in waves for chunk in pack(wave)):
|
||||
fields = field_names(len(chunk))
|
||||
response, sha = engine.predict({"state": state,
|
||||
"questions": {f: s.question for f, s in zip(fields, chunk)}})
|
||||
index = len(sizes)
|
||||
sizes.append(len(chunk))
|
||||
tokens.append(response["usage"]["input_tokens"])
|
||||
inference_ms += response["timing"]["inference_ms"]
|
||||
for f, s in zip(fields, chunk):
|
||||
by_slot[(s.decision, s.k)] = row_result(s, response, sha, f, len(chunk), index, model)
|
||||
out = []
|
||||
for i, d in enumerate(decisions):
|
||||
if d.orderings == "none":
|
||||
out.append(by_slot[(i, 0)])
|
||||
else:
|
||||
out.append(combine(d, [by_slot[(i, k)] for k in range(len(ordering_perms(d)))]))
|
||||
timing = {"total_seconds": time.perf_counter() - started, "batch_size": len(by_slot),
|
||||
"calls": len(sizes), "questions_per_call": sizes, "input_tokens": tokens,
|
||||
"inference_seconds": inference_ms / 1000}
|
||||
return out, timing
|
||||
|
||||
async def score(state: State, decisions: list[Decision], waves: list[list[Slot]]):
|
||||
"""Run one request's calls in a worker thread; map their failures to contract codes."""
|
||||
try:
|
||||
return await run_in_threadpool(run, state, decisions, waves)
|
||||
except ValueError as exc: # the model's own validation, token limit, ...
|
||||
raise ApiError(422, "invalid_request", str(exc)) from exc
|
||||
except OutOfMemory as exc:
|
||||
raise ApiError(503, "out_of_memory", str(exc)) from exc
|
||||
except Exception as exc: # noqa: BLE001 — any other failure, building the response included
|
||||
raise ApiError(500, "scoring_failed", f"{type(exc).__name__}: {exc}") from exc
|
||||
|
||||
@app.get("/health")
|
||||
async def health():
|
||||
return {"status": "ok", "model": engine.health(), "vram_cap_gib": settings.vram_cap_gib,
|
||||
"max_tokens": settings.max_tokens, "max_decisions": settings.max_decisions,
|
||||
"max_questions_per_call": MAX_QUESTIONS_PER_CALL, "chunking": CHUNKING, "workloads": []}
|
||||
|
||||
@app.post("/decide")
|
||||
async def decide(request: Request):
|
||||
admit() # before the body is read
|
||||
try:
|
||||
body = await parse(request, DecideBody)
|
||||
results, timing = await score(body.state, [body], planned(body.state, [body], body.workload))
|
||||
if body.orderings != "none":
|
||||
return results[0]
|
||||
return {**results[0], "total_seconds": timing["total_seconds"],
|
||||
"forward_seconds": timing["inference_seconds"]}
|
||||
finally:
|
||||
leave()
|
||||
|
||||
@app.post("/decide/shared")
|
||||
async def decide_shared(request: Request):
|
||||
admit()
|
||||
try:
|
||||
body = await parse(request, SharedBody)
|
||||
results, timing = await score(body.state, body.decisions,
|
||||
planned(body.state, body.decisions, body.workload))
|
||||
return {"results": results, "timing": timing}
|
||||
finally:
|
||||
leave()
|
||||
|
||||
return app
|
||||
@@ -0,0 +1,87 @@
|
||||
"""Settings for intern-decision-serve. Contract: intern-decision-serve.contract.md § Configuration.
|
||||
|
||||
Every value is validated at startup and a bad one is refused with a ValueError naming the
|
||||
variable: a service that starts and then rejects every request (or runs uncapped) is worse than
|
||||
one that does not start.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import math
|
||||
from collections.abc import Mapping
|
||||
from dataclasses import dataclass
|
||||
|
||||
PREFIX = "INTERN_DECISION_"
|
||||
MIN_TOKEN_CHARS = 32
|
||||
MODEL_ID = "internlm/Intern-Decision-4B"
|
||||
REVISION = "0e5e6aa7d6d750e2b1504ba11a8136cb58aeb3cd"
|
||||
# The checkpoint's own inference.py is EXECUTED from the (data) mount, so its content is pinned.
|
||||
INFERENCE_PY_SHA256 = "c904e2c67ca0775621a22375ee373d2ba30b52117cda870c6c9ef74143b29863"
|
||||
DEFAULT_CHECKPOINT = f"/hf/hub/models--internlm--Intern-Decision-4B/snapshots/{REVISION}"
|
||||
MAX_QUESTIONS_PER_CALL = 16 # the model's own limit (inference.validate_request)
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class Settings:
|
||||
api_token: str
|
||||
checkpoint: str = DEFAULT_CHECKPOINT
|
||||
device: str = "cuda"
|
||||
vram_cap_gib: float | None = None
|
||||
max_tokens: int = 8192
|
||||
max_decisions: int = 64
|
||||
max_body_bytes: int = 1024 * 1024
|
||||
max_queue: int = 32
|
||||
release_slack_mib: int = 512
|
||||
keep_vision: bool = False
|
||||
|
||||
@classmethod
|
||||
def from_env(cls, env: Mapping[str, str]) -> "Settings":
|
||||
token = env.get(f"{PREFIX}API_TOKEN", "")
|
||||
# INV-6: visible ASCII only. A CR, LF or NUL can never arrive in a header, so a token
|
||||
# carrying one would lock every caller out while /health still said ok.
|
||||
if len(token) < MIN_TOKEN_CHARS or not all(33 <= ord(c) <= 126 for c in token):
|
||||
raise ValueError(f"{PREFIX}API_TOKEN must be at least {MIN_TOKEN_CHARS} visible ASCII characters")
|
||||
device = env.get(f"{PREFIX}DEVICE", "cuda")
|
||||
if device not in ("cuda", "cpu"):
|
||||
raise ValueError(f"{PREFIX}DEVICE must be cuda or cpu, not {device!r}")
|
||||
keep_vision = env.get(f"{PREFIX}KEEP_VISION", "0")
|
||||
if keep_vision not in ("0", "1"):
|
||||
raise ValueError(f"{PREFIX}KEEP_VISION must be 0 or 1, not {keep_vision!r}")
|
||||
return cls(
|
||||
api_token=token,
|
||||
checkpoint=env.get(f"{PREFIX}CHECKPOINT", DEFAULT_CHECKPOINT),
|
||||
device=device,
|
||||
vram_cap_gib=_positive_float(env, "VRAM_CAP_GIB"),
|
||||
max_tokens=_int(env, "MAX_TOKENS", 8192),
|
||||
max_decisions=_int(env, "MAX_DECISIONS", 64),
|
||||
max_body_bytes=_int(env, "MAX_BODY_BYTES", 1024 * 1024),
|
||||
max_queue=_int(env, "MAX_QUEUE", 32),
|
||||
release_slack_mib=_int(env, "RELEASE_SLACK_MIB", 512, minimum=0),
|
||||
keep_vision=keep_vision == "1",
|
||||
)
|
||||
|
||||
|
||||
def _int(env: Mapping[str, str], name: str, default: int, minimum: int = 1) -> int:
|
||||
raw = env.get(PREFIX + name)
|
||||
if raw is None:
|
||||
return default
|
||||
try:
|
||||
value = int(raw)
|
||||
except ValueError:
|
||||
raise ValueError(f"{PREFIX}{name} must be an integer, got {raw!r}") from None
|
||||
if value < minimum:
|
||||
raise ValueError(f"{PREFIX}{name} must be >= {minimum}, got {value}")
|
||||
return value
|
||||
|
||||
|
||||
def _positive_float(env: Mapping[str, str], name: str) -> float | None:
|
||||
"""Unset or empty means no cap. When set it must be finite and > 0."""
|
||||
raw = env.get(PREFIX + name)
|
||||
if raw is None or raw == "":
|
||||
return None
|
||||
try:
|
||||
value = float(raw)
|
||||
except ValueError:
|
||||
raise ValueError(f"{PREFIX}{name} must be a number, got {raw!r}") from None
|
||||
if not math.isfinite(value) or value <= 0:
|
||||
raise ValueError(f"{PREFIX}{name} must be a finite number > 0, got {raw!r}")
|
||||
return value
|
||||
@@ -0,0 +1,197 @@
|
||||
"""The real engine: Intern-Decision-4B's own DecisionEngine (the checkpoint's inference.py) over one
|
||||
resident model. Needs the `model` extra.
|
||||
|
||||
Contract: intern-decision-serve.contract.md, INV-3 (fail-closed startup), INV-4 (VRAM cap + OOM),
|
||||
INV-5 (offline weights), INV-7 (text-only model), INV-8 (honest prompt hash). load() is exercised on
|
||||
the card at acceptance; its checks and the OOM path are unit-tested against a fake torch and a
|
||||
fake checkpoint (tests/test_engine.py).
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import gc
|
||||
import hashlib
|
||||
import importlib.metadata
|
||||
import importlib.util
|
||||
import logging
|
||||
import traceback
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
from .config import INFERENCE_PY_SHA256, MODEL_ID, REVISION, Settings
|
||||
from .errors import OutOfMemory, ScoringFailed
|
||||
|
||||
log = logging.getLogger("intern_decision_serve.engine")
|
||||
|
||||
MIB = 2**20
|
||||
WARMUP_REQUEST = {
|
||||
"state": "The deployment completed at 14:02 UTC. Health checks passed in all three zones.",
|
||||
"questions": {"q": {"type": "choice", "instructions": "Is there evidence that the deployment succeeded?",
|
||||
"criteria": {"yes": "The deployment succeeded.", "no": "The deployment did not succeed."}}},
|
||||
}
|
||||
|
||||
|
||||
def _first_line(exc: BaseException) -> str:
|
||||
lines = str(exc).splitlines()
|
||||
return lines[0] if lines else ""
|
||||
|
||||
|
||||
def _version(distribution: str) -> str:
|
||||
try:
|
||||
return importlib.metadata.version(distribution)
|
||||
except importlib.metadata.PackageNotFoundError:
|
||||
return "n/a"
|
||||
|
||||
|
||||
def _sha256(path: Path) -> str:
|
||||
return hashlib.sha256(path.read_bytes()).hexdigest()
|
||||
|
||||
|
||||
def _import_inference(path: Path):
|
||||
"""The checkpoint's own inference.py, imported by file path under a private module name."""
|
||||
spec = importlib.util.spec_from_file_location("intern_decision_inference", path)
|
||||
module = importlib.util.module_from_spec(spec)
|
||||
spec.loader.exec_module(module)
|
||||
return module
|
||||
|
||||
|
||||
def _vision_stub(torch: Any):
|
||||
class VisionTowerRemoved(torch.nn.Module):
|
||||
"""INV-7: stands where the vision tower was. The service takes no images; if anything ever
|
||||
routes pixels here, fail loudly rather than answer from a missing tower."""
|
||||
def forward(self, *_args, **_kwargs):
|
||||
raise RuntimeError("the vision tower was removed at startup (text-only service, INV-7)")
|
||||
return VisionTowerRemoved()
|
||||
|
||||
|
||||
class TorchEngine:
|
||||
def __init__(self, torch: Any, engine: Any, inference: Any, tokenizer: Any, metadata: dict,
|
||||
settings: Settings, release_above_bytes: int | None = None):
|
||||
self._torch, self._engine, self._inference, self._tokenizer = torch, engine, inference, tokenizer
|
||||
self.metadata, self._settings = metadata, settings
|
||||
self._release_above = release_above_bytes
|
||||
|
||||
@classmethod
|
||||
def load(cls, settings: Settings, *, torch: Any = None,
|
||||
inference_sha256: str = INFERENCE_PY_SHA256) -> "TorchEngine":
|
||||
checkpoint = Path(settings.checkpoint)
|
||||
# INV-3: the pinned snapshot, and the pinned code. Checked before torch touches the card.
|
||||
if checkpoint.name != REVISION or checkpoint.parent.name != "snapshots":
|
||||
raise RuntimeError(f"INTERN_DECISION_CHECKPOINT must be a snapshots/{REVISION} directory, "
|
||||
f"got {checkpoint}")
|
||||
code = checkpoint / "inference.py"
|
||||
found = _sha256(code)
|
||||
if found != inference_sha256:
|
||||
raise RuntimeError(f"{code} has sha256 {found}, not the pinned {inference_sha256}: "
|
||||
"re-check the checkpoint's inference.py before serving it")
|
||||
if torch is None:
|
||||
import torch
|
||||
|
||||
if settings.device == "cuda":
|
||||
if not torch.cuda.is_available():
|
||||
raise RuntimeError("INTERN_DECISION_DEVICE=cuda but torch sees no CUDA device")
|
||||
major, minor = torch.cuda.get_device_capability(0)
|
||||
arch = f"sm_{major}{minor}"
|
||||
if arch not in torch.cuda.get_arch_list(): # INV-3: no silent PTX/CPU fallback
|
||||
raise RuntimeError(f"torch {torch.__version__} has no kernels for {arch}: {torch.cuda.get_arch_list()}")
|
||||
if settings.vram_cap_gib is not None: # INV-4: cap BEFORE the weights land
|
||||
total = torch.cuda.get_device_properties(0).total_memory
|
||||
fraction = settings.vram_cap_gib * 2**30 / total
|
||||
if not 0 < fraction <= 1:
|
||||
raise ValueError(f"INTERN_DECISION_VRAM_CAP_GIB={settings.vram_cap_gib} does not fit a "
|
||||
f"{total / 2**30:.1f} GiB card")
|
||||
torch.cuda.set_per_process_memory_fraction(fraction, 0)
|
||||
|
||||
inference = _import_inference(code)
|
||||
decision_engine = inference.DecisionEngine(checkpoint=str(checkpoint), max_length=settings.max_tokens,
|
||||
device=settings.device, dtype="bfloat16",
|
||||
attn_implementation="sdpa")
|
||||
backend = decision_engine.backend
|
||||
placed = next(backend.model.parameters()).device.type
|
||||
if placed != settings.device: # INV-3
|
||||
raise RuntimeError(f"model landed on {placed}, expected {settings.device}")
|
||||
metadata = {"name": inference.MODEL_NAME, "source": MODEL_ID, "revision": REVISION,
|
||||
"checkpoint": str(checkpoint), "inference_py_sha256": found,
|
||||
"temperature": decision_engine.temperature, "dtype": "bfloat16", "attn_implementation": "sdpa",
|
||||
"device": settings.device, "max_length": settings.max_tokens,
|
||||
"torch_version": torch.__version__, "transformers_version": _version("transformers"),
|
||||
"vision_tower": "loaded" if settings.keep_vision else "removed"}
|
||||
engine = cls(torch, decision_engine, inference, backend.tokenizer, metadata, settings)
|
||||
|
||||
first, _ = engine.predict(WARMUP_REQUEST) # INV-3: one decision must score
|
||||
engine._prove_prompt_hash(first) # INV-8
|
||||
if not settings.keep_vision: # INV-7
|
||||
engine._remove_vision_tower(first)
|
||||
if settings.device == "cuda": # INV-4: the resting footprint
|
||||
gc.collect()
|
||||
torch.cuda.empty_cache()
|
||||
engine._release_above = torch.cuda.memory_reserved(0) + settings.release_slack_mib * MIB
|
||||
return engine
|
||||
|
||||
def _prompt_text(self, request: dict) -> str:
|
||||
"""INV-8: the chat-template text, rendered with the same arguments HFBackend.encode uses for a
|
||||
text-only row (the tokenizer is its template when there are no images)."""
|
||||
compiled = self._inference.compile_row(self._inference.validate_request(request))
|
||||
return self._tokenizer.apply_chat_template(compiled.messages, tokenize=False, add_generation_prompt=False,
|
||||
enable_thinking=False, add_vision_id=True)
|
||||
|
||||
def _prove_prompt_hash(self, warmup: dict) -> None:
|
||||
tokens = len(self._tokenizer(self._prompt_text(WARMUP_REQUEST), add_special_tokens=False)["input_ids"])
|
||||
if tokens != warmup["usage"]["input_tokens"]:
|
||||
raise RuntimeError(f"INV-8: the rendered prompt tokenises to {tokens} tokens but the model read "
|
||||
f"{warmup['usage']['input_tokens']}: prompt_sha256 would hash a prompt it never saw")
|
||||
|
||||
def _remove_vision_tower(self, before: dict) -> None:
|
||||
inner = getattr(getattr(self._engine.backend, "model", None), "model", None)
|
||||
if inner is None or not hasattr(inner, "visual"):
|
||||
raise RuntimeError("INV-7: the vision tower is not at backend.model.model.visual; "
|
||||
"set INTERN_DECISION_KEEP_VISION=1 or re-check the model class")
|
||||
inner.visual = _vision_stub(self._torch)
|
||||
gc.collect()
|
||||
after, _ = self.predict(WARMUP_REQUEST)
|
||||
if after["answers"] != before["answers"]:
|
||||
raise RuntimeError(f"INV-7: removing the vision tower changed the warm-up answer "
|
||||
f"({before['answers']} -> {after['answers']})")
|
||||
|
||||
def health(self) -> dict:
|
||||
info = dict(self.metadata)
|
||||
if self._settings.device == "cuda":
|
||||
cuda = self._torch.cuda
|
||||
info["device_name"] = cuda.get_device_name(0)
|
||||
info["allocated_gib"] = round(cuda.memory_allocated(0) / 2**30, 3)
|
||||
info["reserved_gib"] = round(cuda.memory_reserved(0) / 2**30, 3)
|
||||
if hasattr(cuda, "max_memory_reserved"):
|
||||
info["max_reserved_gib"] = round(cuda.max_memory_reserved(0) / 2**30, 3)
|
||||
return info
|
||||
|
||||
def _release_burst(self) -> None:
|
||||
"""INV-4: hand a burst back to the driver so GPU 1's shared headroom (scriberr, the vLLM
|
||||
seats) returns after a big request, instead of sitting in torch's cache."""
|
||||
if self._release_above is not None and self._torch.cuda.memory_reserved(0) > self._release_above:
|
||||
self._torch.cuda.empty_cache()
|
||||
|
||||
def predict(self, request: dict) -> tuple[dict, str]:
|
||||
"""One call: the model's own predict(), then the prompt hash (INV-8). Returns (response, sha)."""
|
||||
try:
|
||||
response = self._engine.predict(request)
|
||||
sha = hashlib.sha256(self._prompt_text(request).encode()).hexdigest()
|
||||
except ValueError:
|
||||
raise # the model's validation: raised before any GPU work
|
||||
except self._torch.cuda.OutOfMemoryError as exc:
|
||||
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
|
||||
message = _first_line(exc)
|
||||
if "out of memory" in message.lower(): # cuBLAS/cuDNN allocation failures
|
||||
failure = OutOfMemory
|
||||
else:
|
||||
failure, message = ScoringFailed, f"{type(exc).__name__}: {message}"
|
||||
# Formatted text, not exc_info: a record that keeps the traceback alive pins the tensors.
|
||||
log.error("predict failed:\n%s", traceback.format_exc())
|
||||
else:
|
||||
self._release_burst()
|
||||
return response, sha
|
||||
# INV-4, outside the except block on purpose: the exception's traceback holds the failed
|
||||
# forward's frames and with them its tensors. Raising inside the block, or `from exc`, would
|
||||
# chain to it and keep them allocated after the response (semif-serve, found on the card).
|
||||
gc.collect()
|
||||
self._torch.cuda.empty_cache()
|
||||
raise failure(message)
|
||||
@@ -0,0 +1,10 @@
|
||||
"""Torch-free exceptions shared by the HTTP layer and the engine."""
|
||||
|
||||
|
||||
class OutOfMemory(RuntimeError):
|
||||
"""The engine ran out of GPU memory during a call and has already released its cache (INV-4)."""
|
||||
|
||||
|
||||
class ScoringFailed(RuntimeError):
|
||||
"""A call failed for a reason other than validation or OOM. Raised unchained, after the failed
|
||||
call's memory has been released; the original traceback is logged, not carried (INV-4)."""
|
||||
@@ -0,0 +1,20 @@
|
||||
"""uvicorn entry point: `uvicorn intern_decision_serve.main:app_from_env --factory --workers 1`."""
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
|
||||
from fastapi import FastAPI
|
||||
|
||||
from .app import create_app
|
||||
from .config import Settings
|
||||
|
||||
|
||||
def app_from_env() -> FastAPI:
|
||||
settings = Settings.from_env(os.environ)
|
||||
# INV-5: never download at runtime, inside the image or out of it. Set before torch /
|
||||
# transformers / huggingface_hub are imported, since they read it at import time.
|
||||
os.environ["HF_HUB_OFFLINE"] = "1"
|
||||
os.environ["TRANSFORMERS_OFFLINE"] = "1"
|
||||
from .engine import TorchEngine # torch loads only inside load(), never in the unit tests
|
||||
|
||||
return create_app(settings, TorchEngine.load(settings))
|
||||
@@ -0,0 +1,4 @@
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
sys.path.insert(0, str(Path(__file__).resolve().parent))
|
||||
@@ -0,0 +1,55 @@
|
||||
"""A torch-free stand-in for the real engine: answers Jev requests in the shape
|
||||
DecisionEngine.predict() returns, and records every call it was given."""
|
||||
from __future__ import annotations
|
||||
|
||||
import copy
|
||||
import hashlib
|
||||
import json
|
||||
import math
|
||||
|
||||
CALIBRATION = {"method": "temperature-scaling", "temperature": 1.99241824}
|
||||
|
||||
|
||||
def softmax(xs: list[float]) -> list[float]:
|
||||
top = max(xs)
|
||||
e = [math.exp(x - top) for x in xs]
|
||||
return [v / sum(e) for v in e]
|
||||
|
||||
|
||||
def score_by_description(field: str, question: dict) -> list[float]:
|
||||
"""Default scorer: the option whose description is longest wins; a small first-position bias."""
|
||||
descs = list(question["criteria"].values())
|
||||
return [len(d) + (0.5 if i == 0 else 0.0) for i, d in enumerate(descs)]
|
||||
|
||||
|
||||
class FakeEngine:
|
||||
def __init__(self, scorer=score_by_description, tokens_per_question: int = 100):
|
||||
self.scorer = scorer
|
||||
self.tokens_per_question = tokens_per_question
|
||||
self.calls: list[dict] = []
|
||||
|
||||
metadata = {"name": "Intern-Decision-4B", "revision": "0" * 40}
|
||||
|
||||
def health(self) -> dict:
|
||||
return {**self.metadata, "reserved_gib": 9.1}
|
||||
|
||||
def predict(self, request: dict) -> tuple[dict, str]:
|
||||
self.calls.append(copy.deepcopy(request))
|
||||
answers = {}
|
||||
for field, question in request["questions"].items():
|
||||
ids = list(question["criteria"])
|
||||
probs = dict(zip(ids, softmax(self.scorer(field, question))))
|
||||
best = min(ids, key=lambda i: (-probs[i], i))
|
||||
answers[field] = {"type": "choice", "probabilities": probs, "confidence": probs[best],
|
||||
"choice": best, "source": "local", "decision": best}
|
||||
response = {"answers": answers,
|
||||
"usage": {"input_tokens": self.tokens_per_question * len(answers),
|
||||
"output_tokens": len(answers), "decision_count": len(answers)},
|
||||
"timing": {"inference_ms": 12.5}, "calibration": dict(CALIBRATION),
|
||||
"model": "Intern-Decision-4B", "backend": "hf"}
|
||||
return response, request_sha(request)
|
||||
|
||||
|
||||
def request_sha(request: dict) -> str:
|
||||
"""Stands in for the real prompt hash: the same call always hashes the same."""
|
||||
return hashlib.sha256(json.dumps(request, sort_keys=True).encode()).hexdigest()
|
||||
@@ -0,0 +1,434 @@
|
||||
"""intern-decision-serve HTTP behaviour against a fake engine (no torch, no model).
|
||||
Contract: services/intern-decision-serve/intern-decision-serve.contract.md"""
|
||||
import json
|
||||
|
||||
from fastapi.testclient import TestClient
|
||||
|
||||
from fake_engine import CALIBRATION, FakeEngine, request_sha
|
||||
from intern_decision_serve.app import create_app
|
||||
from intern_decision_serve.config import Settings
|
||||
|
||||
TOKEN = "t" * 40
|
||||
AUTH = {"Authorization": f"Bearer {TOKEN}"}
|
||||
OPTIONS = [{"id": "yes", "description": "It passed."}, {"id": "no", "description": "It did not pass at all."}]
|
||||
ROW = {"id": "r1", "state": "The deploy passed.", "question": "Did it pass?", "options": OPTIONS}
|
||||
|
||||
|
||||
def make_client(engine=None, **overrides):
|
||||
return TestClient(create_app(Settings(api_token=TOKEN, **overrides), engine or FakeEngine()))
|
||||
|
||||
|
||||
def test_decide_asks_one_question_named_q_and_maps_the_answer_back_in_option_order():
|
||||
engine = FakeEngine()
|
||||
response = make_client(engine).post("/decide", json=ROW, headers=AUTH)
|
||||
assert response.status_code == 200
|
||||
assert engine.calls == [{"state": "The deploy passed.", "questions": {"q": {
|
||||
"type": "choice", "instructions": "Did it pass?",
|
||||
"criteria": {"yes": "It passed.", "no": "It did not pass at all."}}}}]
|
||||
out = response.json()
|
||||
native, _ = FakeEngine().predict(engine.calls[0])
|
||||
answer = native["answers"]["q"]
|
||||
assert out["id"] == "r1"
|
||||
assert out["option_ids"] == ["yes", "no"]
|
||||
assert out["probabilities"] == [answer["probabilities"]["yes"], answer["probabilities"]["no"]]
|
||||
assert out["top"] == "no" == answer["choice"]
|
||||
assert out["confidence"] == answer["confidence"]
|
||||
assert out["calibration"] == CALIBRATION
|
||||
assert out["native"] == answer
|
||||
assert out["input_tokens"] == 100
|
||||
assert out["prompt_sha256"] == request_sha(engine.calls[0])
|
||||
assert out["call"] == {"index": 0, "field": "q", "questions": 1}
|
||||
assert out["forward_seconds"] == 0.0125
|
||||
assert out["total_seconds"] >= 0
|
||||
assert out["model"] == FakeEngine.metadata
|
||||
for key in ("prompt_version", "readout", "probability_status"):
|
||||
assert out[key]
|
||||
|
||||
|
||||
import pytest # noqa: E402
|
||||
|
||||
|
||||
@pytest.mark.parametrize("headers", [{}, {"Authorization": "Bearer wrong"}, {"Authorization": TOKEN}])
|
||||
def test_posts_without_the_right_bearer_are_401_and_never_reach_the_engine(headers):
|
||||
engine = FakeEngine()
|
||||
client = make_client(engine)
|
||||
for path, body in (("/decide", ROW), ("/decide/shared", {"state": "s", "decisions": [ROW]})):
|
||||
response = client.post(path, json=body, headers=headers)
|
||||
assert response.status_code == 401
|
||||
assert response.json()["error"]["code"] == "unauthorized"
|
||||
assert engine.calls == []
|
||||
|
||||
|
||||
def test_health_needs_no_auth_and_reports_the_limits_and_the_chunking_rule():
|
||||
response = make_client(vram_cap_gib=11.0).get("/health")
|
||||
assert response.status_code == 200
|
||||
out = response.json()
|
||||
assert out["status"] == "ok"
|
||||
assert out["model"] == FakeEngine().health()
|
||||
assert (out["vram_cap_gib"], out["max_tokens"], out["max_decisions"]) == (11.0, 8192, 64)
|
||||
assert out["max_questions_per_call"] == 16
|
||||
assert "16" in out["chunking"]
|
||||
assert out["workloads"] == []
|
||||
|
||||
|
||||
def decisions(n, options=OPTIONS):
|
||||
return [{"id": f"decision-{i}", "question": f"Question {i}?", "options": options} for i in range(n)]
|
||||
|
||||
|
||||
def test_shared_asks_q1_to_qn_over_the_shared_state_in_one_call():
|
||||
engine = FakeEngine()
|
||||
body = {"state": {"deploy": "passed"}, "decisions": decisions(3)}
|
||||
response = make_client(engine).post("/decide/shared", json=body, headers=AUTH)
|
||||
assert response.status_code == 200
|
||||
assert engine.calls == [{"state": {"deploy": "passed"}, "questions": {
|
||||
f"q{i + 1}": {"type": "choice", "instructions": f"Question {i}?",
|
||||
"criteria": {"yes": "It passed.", "no": "It did not pass at all."}} for i in range(3)}}]
|
||||
out = response.json()
|
||||
assert [r["id"] for r in out["results"]] == ["decision-0", "decision-1", "decision-2"]
|
||||
assert [r["call"] for r in out["results"]] == [{"index": 0, "field": f"q{i}", "questions": 3} for i in (1, 2, 3)]
|
||||
assert out["timing"]["calls"] == 1
|
||||
assert out["timing"]["batch_size"] == 3
|
||||
assert out["timing"]["questions_per_call"] == [3]
|
||||
assert out["timing"]["input_tokens"] == [300]
|
||||
assert out["timing"]["inference_seconds"] == 0.0125
|
||||
assert out["timing"]["total_seconds"] >= 0
|
||||
|
||||
|
||||
def test_one_shared_decision_is_the_same_call_as_decide():
|
||||
engine = FakeEngine()
|
||||
client = make_client(engine)
|
||||
client.post("/decide", json=ROW, headers=AUTH)
|
||||
client.post("/decide/shared", json={"state": ROW["state"], "decisions": [ROW]}, headers=AUTH)
|
||||
assert engine.calls[0] == engine.calls[1]
|
||||
|
||||
|
||||
@pytest.mark.parametrize("n, sizes", [(16, [16]), (17, [16, 1]), (40, [16, 16, 8])])
|
||||
def test_more_than_16_questions_are_split_greedily_in_request_order(n, sizes):
|
||||
engine = FakeEngine()
|
||||
out = make_client(engine).post("/decide/shared", json={"state": "s", "decisions": decisions(n)},
|
||||
headers=AUTH).json()
|
||||
assert [len(c["questions"]) for c in engine.calls] == sizes
|
||||
asked = [q["instructions"] for c in engine.calls for q in c["questions"].values()]
|
||||
assert asked == [f"Question {i}?" for i in range(n)]
|
||||
assert [r["id"] for r in out["results"]] == [f"decision-{i}" for i in range(n)]
|
||||
for i, result in enumerate(out["results"]):
|
||||
call = i // 16
|
||||
size = sizes[call]
|
||||
field = "q" if size == 1 else f"q{i % 16 + 1}"
|
||||
assert result["call"] == {"index": call, "field": field, "questions": size}
|
||||
assert result["prompt_sha256"] == request_sha(engine.calls[call])
|
||||
assert out["timing"]["questions_per_call"] == sizes
|
||||
assert out["timing"]["calls"] == len(sizes)
|
||||
|
||||
|
||||
def test_decision_ids_never_reach_the_model():
|
||||
engine = FakeEngine()
|
||||
make_client(engine).post("/decide/shared", json={"state": "s", "decisions": decisions(20)}, headers=AUTH)
|
||||
make_client(engine).post("/decide", json=ROW, headers=AUTH)
|
||||
assert "decision-" not in repr(engine.calls) and "r1" not in repr(engine.calls)
|
||||
|
||||
|
||||
from intern_decision_serve.errors import OutOfMemory, ScoringFailed # noqa: E402
|
||||
|
||||
|
||||
class RaisingEngine(FakeEngine):
|
||||
def __init__(self, exc, after=0):
|
||||
super().__init__()
|
||||
self.exc, self.after = exc, after
|
||||
|
||||
def predict(self, request):
|
||||
if len(self.calls) >= self.after:
|
||||
self.calls.append(request)
|
||||
raise self.exc
|
||||
return super().predict(request)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("exc, status, code", [
|
||||
(ValueError("Example has 9000 tokens, above 8192; truncation is forbidden"), 422, "invalid_request"),
|
||||
(OutOfMemory("CUDA out of memory. Tried to allocate 2.00 GiB."), 503, "out_of_memory"),
|
||||
(ScoringFailed("RuntimeError: boom"), 500, "scoring_failed"),
|
||||
(ArithmeticError("Floating-point temperature scaling changed argmax"), 500, "scoring_failed"),
|
||||
])
|
||||
def test_engine_failures_map_to_the_contract_codes_on_both_endpoints(exc, status, code):
|
||||
for path, body in (("/decide", ROW), ("/decide/shared", {"state": "s", "decisions": decisions(2)})):
|
||||
response = make_client(RaisingEngine(exc)).post(path, json=body, headers=AUTH)
|
||||
assert response.status_code == status
|
||||
assert response.json()["error"]["code"] == code
|
||||
assert str(exc) in response.json()["error"]["message"]
|
||||
|
||||
|
||||
def test_an_oom_in_a_later_chunk_fails_the_whole_request_with_503():
|
||||
engine = RaisingEngine(OutOfMemory("CUDA out of memory."), after=1)
|
||||
response = make_client(engine).post("/decide/shared", json={"state": "s", "decisions": decisions(20)},
|
||||
headers=AUTH)
|
||||
assert response.status_code == 503
|
||||
assert len(engine.calls) == 2
|
||||
|
||||
|
||||
class DroppingEngine(FakeEngine):
|
||||
"""Returns an answer that lacks one of the asked option ids."""
|
||||
def predict(self, request):
|
||||
response, sha = super().predict(request)
|
||||
for answer in response["answers"].values():
|
||||
answer["probabilities"].pop("no")
|
||||
return response, sha
|
||||
|
||||
|
||||
def test_an_answer_missing_an_option_id_is_a_500_not_a_guess():
|
||||
response = make_client(DroppingEngine()).post("/decide", json=ROW, headers=AUTH)
|
||||
assert response.status_code == 500
|
||||
assert response.json()["error"]["code"] == "scoring_failed"
|
||||
assert "no" in response.json()["error"]["message"]
|
||||
|
||||
|
||||
def opts(n):
|
||||
return [{"id": f"o{i}", "description": f"Option {i}"} for i in range(n)]
|
||||
|
||||
|
||||
BAD_DECIDE = [
|
||||
{**ROW, "options": opts(1)},
|
||||
{**ROW, "options": opts(17)},
|
||||
{**ROW, "options": [{"id": "a", "description": "x"}, {"id": "a", "description": "y"}]},
|
||||
{**ROW, "options": [{"id": "a"}, {"id": "b", "description": "y"}]},
|
||||
{**ROW, "options": [{"id": 1, "description": "x"}, {"id": "b", "description": "y"}]},
|
||||
{**ROW, "id": ""},
|
||||
{**ROW, "question": ""},
|
||||
{**ROW, "state": ""},
|
||||
{**ROW, "state": {}},
|
||||
{**ROW, "state": []},
|
||||
{**ROW, "state": 7},
|
||||
{**ROW, "workload": "triage"},
|
||||
{k: v for k, v in ROW.items() if k != "question"},
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.parametrize("body", BAD_DECIDE)
|
||||
def test_semif_rule_violations_are_422_before_the_engine_runs(body):
|
||||
engine = FakeEngine()
|
||||
client = make_client(engine)
|
||||
response = client.post("/decide", json=body, headers=AUTH)
|
||||
assert response.status_code == 422
|
||||
assert response.json()["error"]["code"] == "invalid_request"
|
||||
decision = {k: v for k, v in body.items() if k not in ("state", "workload")}
|
||||
shared = {"state": body.get("state", "s"), "decisions": [decision],
|
||||
**({"workload": body["workload"]} if "workload" in body else {})}
|
||||
assert client.post("/decide/shared", json=shared, headers=AUTH).status_code == 422
|
||||
assert engine.calls == []
|
||||
|
||||
|
||||
def test_a_non_finite_state_is_422():
|
||||
engine = FakeEngine()
|
||||
raw = '{"id": "r", "state": {"x": NaN}, "question": "q?", "options": [{"id": "a", "description": "A"}, {"id": "b", "description": "B"}]}'
|
||||
response = make_client(engine).post("/decide", content=raw, headers={**AUTH, "Content-Type": "application/json"})
|
||||
assert response.status_code == 422
|
||||
assert engine.calls == []
|
||||
|
||||
|
||||
def test_duplicate_decision_ids_are_422():
|
||||
body = {"state": "s", "decisions": [{**decisions(1)[0]}, {**decisions(1)[0]}]}
|
||||
assert make_client().post("/decide/shared", json=body, headers=AUTH).status_code == 422
|
||||
|
||||
|
||||
@pytest.mark.parametrize("raw", ["{not json", "[]", '{"state": "s"}', '{"state": "s", "decisions": "x"}'])
|
||||
def test_malformed_bodies_are_422(raw):
|
||||
client = make_client()
|
||||
for path in ("/decide", "/decide/shared"):
|
||||
response = client.post(path, content=raw, headers={**AUTH, "Content-Type": "application/json"})
|
||||
assert response.status_code == 422
|
||||
assert response.json()["error"]["code"] == "invalid_request"
|
||||
|
||||
|
||||
@pytest.mark.parametrize("n", [0, 5])
|
||||
def test_a_row_count_outside_1_to_max_decisions_is_422(n):
|
||||
engine = FakeEngine()
|
||||
response = make_client(engine, max_decisions=4).post(
|
||||
"/decide/shared", json={"state": "s", "decisions": decisions(n)}, headers=AUTH)
|
||||
assert response.status_code == 422
|
||||
assert engine.calls == []
|
||||
|
||||
|
||||
import math # noqa: E402
|
||||
|
||||
|
||||
def position_bias(field, question):
|
||||
"""Pure position bias: the first-listed option gets +2, whatever it says."""
|
||||
return [2.0 if i == 0 else 0.0 for i in range(len(question["criteria"]))]
|
||||
|
||||
|
||||
def test_rotations_ask_each_ordering_in_its_own_wave_and_cancel_a_position_bias_exactly():
|
||||
engine = FakeEngine(scorer=position_bias)
|
||||
body = {**ROW, "options": opts(3), "orderings": "rotations"}
|
||||
out = make_client(engine).post("/decide", json=body, headers=AUTH).json()
|
||||
assert [list(c["questions"]) for c in engine.calls] == [["q"], ["q"], ["q"]]
|
||||
assert [list(c["questions"]["q"]["criteria"]) for c in engine.calls] == [
|
||||
["o0", "o1", "o2"], ["o1", "o2", "o0"], ["o2", "o0", "o1"]]
|
||||
assert out["id"] == "r1" and out["option_ids"] == ["o0", "o1", "o2"]
|
||||
c = out["combined"]
|
||||
assert (c["method"], c["orderings"]) == ("rotations", 3)
|
||||
assert c["probabilities"] == pytest.approx([1 / 3] * 3)
|
||||
assert c["top"] == "o0" # first maximum in the caller's order
|
||||
assert c["agreement"] == pytest.approx(1 / 3)
|
||||
assert [r["id"] for r in out["orderings"]] == ["r1#o0", "r1#o1", "r1#o2"]
|
||||
assert [r["option_ids"] for r in out["orderings"]] == [["o0", "o1", "o2"], ["o1", "o2", "o0"], ["o2", "o0", "o1"]]
|
||||
high, low = softmax_pair = (math.exp(2) / (math.exp(2) + 2), 1 / (math.exp(2) + 2))
|
||||
assert c["spread"] == {i: [pytest.approx(low), pytest.approx(high)] for i in ("o0", "o1", "o2")}
|
||||
del softmax_pair
|
||||
|
||||
|
||||
def test_all_asks_every_permutation_and_is_422_above_4_options():
|
||||
engine = FakeEngine()
|
||||
client = make_client(engine)
|
||||
out = client.post("/decide", json={**ROW, "options": opts(3), "orderings": "all"}, headers=AUTH).json()
|
||||
assert len(engine.calls) == 6 and out["combined"]["orderings"] == 6
|
||||
assert len({tuple(c["questions"]["q"]["criteria"]) for c in engine.calls}) == 6
|
||||
assert list(engine.calls[0]["questions"]["q"]["criteria"]) == ["o0", "o1", "o2"]
|
||||
response = client.post("/decide", json={**ROW, "options": opts(5), "orderings": "all"}, headers=AUTH)
|
||||
assert response.status_code == 422 and len(engine.calls) == 6
|
||||
|
||||
|
||||
def test_a_mixed_shared_request_runs_in_waves_and_wave_0_is_the_request_as_written():
|
||||
engine = FakeEngine()
|
||||
plain = {"state": "s", "decisions": [
|
||||
{"id": "a", "question": "A?", "options": OPTIONS},
|
||||
{"id": "b", "question": "B?", "options": opts(3)},
|
||||
{"id": "c", "question": "C?", "options": OPTIONS}]}
|
||||
client = make_client(engine)
|
||||
plain_out = client.post("/decide/shared", json=plain, headers=AUTH).json()
|
||||
plain_call = engine.calls.pop()
|
||||
mixed = {**plain, "decisions": [plain["decisions"][0], {**plain["decisions"][1], "orderings": "rotations"},
|
||||
plain["decisions"][2]]}
|
||||
out = client.post("/decide/shared", json=mixed, headers=AUTH).json()
|
||||
assert engine.calls[0] == plain_call # wave 0
|
||||
assert [list(c["questions"]) for c in engine.calls] == [["q1", "q2", "q3"], ["q"], ["q"]]
|
||||
assert [c["questions"]["q"]["instructions"] for c in engine.calls[1:]] == ["B?", "B?"]
|
||||
assert out["results"][0] == plain_out["results"][0] and out["results"][2] == plain_out["results"][2]
|
||||
assert out["results"][1]["combined"]["orderings"] == 3
|
||||
assert out["timing"]["batch_size"] == 5 and out["timing"]["calls"] == 3
|
||||
|
||||
|
||||
def test_waves_never_put_two_orderings_of_one_decision_in_one_call_and_are_packed_at_16():
|
||||
engine = FakeEngine()
|
||||
body = {"state": "s", "decisions": [{**d, "orderings": "rotations"} for d in decisions(17)]}
|
||||
out = make_client(engine).post("/decide/shared", json=body, headers=AUTH).json()
|
||||
assert [len(c["questions"]) for c in engine.calls] == [16, 1, 16, 1]
|
||||
for call in engine.calls:
|
||||
asked = [q["instructions"] for q in call["questions"].values()]
|
||||
assert len(asked) == len(set(asked))
|
||||
assert [r["id"] for r in out["results"]] == [f"decision-{i}" for i in range(17)]
|
||||
assert out["timing"]["batch_size"] == 34
|
||||
|
||||
|
||||
def test_orderings_count_toward_the_row_cap():
|
||||
engine = FakeEngine()
|
||||
body = {"state": "s", "decisions": [{**d, "orderings": "rotations"} for d in decisions(3)]}
|
||||
assert make_client(engine, max_decisions=5).post("/decide/shared", json=body, headers=AUTH).status_code == 422
|
||||
assert engine.calls == []
|
||||
|
||||
|
||||
def test_a_zero_probability_is_floored_before_the_log():
|
||||
class ZeroEngine(FakeEngine):
|
||||
def predict(self, request):
|
||||
response, sha = super().predict(request)
|
||||
for answer in response["answers"].values():
|
||||
first = next(iter(answer["probabilities"]))
|
||||
answer["probabilities"] = {k: (0.0 if k == first else 1.0 / (len(answer["probabilities"]) - 1))
|
||||
for k in answer["probabilities"]}
|
||||
return response, sha
|
||||
out = make_client(ZeroEngine()).post("/decide", json={**ROW, "orderings": "rotations"}, headers=AUTH)
|
||||
assert out.status_code == 200
|
||||
assert sum(out.json()["combined"]["probabilities"]) == pytest.approx(1.0)
|
||||
|
||||
|
||||
import threading # noqa: E402
|
||||
import time # noqa: E402
|
||||
from concurrent.futures import ThreadPoolExecutor # noqa: E402
|
||||
|
||||
|
||||
@pytest.mark.parametrize("chunked", [False, True])
|
||||
def test_a_body_over_the_limit_is_413_whether_or_not_it_declares_its_length(chunked):
|
||||
engine = FakeEngine()
|
||||
client = make_client(engine, max_body_bytes=200)
|
||||
body = ('{"id": "r1", "state": "' + "x" * 500 + '", "question": "Q?", "options": []}').encode()
|
||||
content = (chunk for chunk in [body[:100], body[100:]]) if chunked else body
|
||||
response = client.post("/decide", content=content, headers={**AUTH, "content-type": "application/json"})
|
||||
assert response.status_code == 413
|
||||
assert response.json()["error"]["code"] == "request_too_large"
|
||||
assert engine.calls == []
|
||||
|
||||
|
||||
def test_a_body_of_exactly_the_limit_is_accepted():
|
||||
body = json.dumps(ROW).encode()
|
||||
client = make_client(max_body_bytes=len(body))
|
||||
assert client.post("/decide", content=body, headers={**AUTH, "content-type": "application/json"}).status_code == 200
|
||||
|
||||
|
||||
def test_read_limited_stops_at_the_crossing_chunk_and_ignores_a_non_ascii_length():
|
||||
import asyncio
|
||||
from intern_decision_serve.app import ApiError, read_limited
|
||||
consumed = []
|
||||
|
||||
async def chunks():
|
||||
for i in range(10):
|
||||
consumed.append(i)
|
||||
yield b"x" * 100
|
||||
|
||||
with pytest.raises(ApiError) as info:
|
||||
asyncio.run(read_limited(chunks(), None, 250))
|
||||
assert info.value.status == 413 and consumed == [0, 1, 2]
|
||||
|
||||
async def small():
|
||||
yield b"{}"
|
||||
|
||||
assert asyncio.run(read_limited(small(), "²", 100)) == b"{}" # int("²") would raise
|
||||
|
||||
|
||||
class SlowEngine(FakeEngine):
|
||||
"""Holds each call until released; records the peak number of calls inside at once and the
|
||||
order in which requests' calls ran."""
|
||||
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
self.inside = self.peak = 0
|
||||
self.guard = threading.Lock()
|
||||
self.release = threading.Event()
|
||||
self.entered = threading.Event()
|
||||
|
||||
def predict(self, request):
|
||||
with self.guard:
|
||||
self.inside += 1
|
||||
self.peak = max(self.peak, self.inside)
|
||||
self.entered.set()
|
||||
self.release.wait(5)
|
||||
with self.guard:
|
||||
self.inside -= 1
|
||||
return super().predict(request)
|
||||
|
||||
|
||||
def test_concurrent_requests_never_overlap_and_a_requests_chunks_are_not_interleaved():
|
||||
engine = SlowEngine()
|
||||
with make_client(engine) as client, ThreadPoolExecutor(4) as pool:
|
||||
futures = [pool.submit(client.post, "/decide/shared",
|
||||
json={"state": f"state {i}", "decisions": decisions(20)}, headers=AUTH)
|
||||
for i in range(3)]
|
||||
assert engine.entered.wait(5)
|
||||
started = time.monotonic()
|
||||
assert client.get("/health").status_code == 200
|
||||
assert time.monotonic() - started < 1.0 # answered while a call is held
|
||||
engine.release.set()
|
||||
assert [f.result().status_code for f in futures] == [200] * 3
|
||||
assert engine.peak == 1
|
||||
states = [c["state"] for c in engine.calls]
|
||||
assert all(states[k] == states[k + 1] for k in range(0, 6, 2)) # each request's 2 chunks back to back
|
||||
|
||||
|
||||
def test_more_than_max_queue_requests_in_progress_get_429_busy_before_the_body_is_read():
|
||||
engine = SlowEngine()
|
||||
with make_client(engine, max_queue=2) as client, ThreadPoolExecutor(3) as pool:
|
||||
held = [pool.submit(client.post, "/decide", json={**ROW, "id": f"r{i}"}, headers=AUTH) for i in range(2)]
|
||||
assert engine.entered.wait(5)
|
||||
time.sleep(0.2) # let the second request reach the queue
|
||||
extra = client.post("/decide", content=b"{not even json", headers={**AUTH, "content-type": "application/json"})
|
||||
assert extra.status_code == 429 and extra.json()["error"]["code"] == "busy"
|
||||
engine.release.set()
|
||||
assert [f.result().status_code for f in held] == [200, 200]
|
||||
assert make_client(FakeEngine(), max_queue=2).post("/decide", json=ROW, headers=AUTH).status_code == 200
|
||||
@@ -0,0 +1,48 @@
|
||||
"""Settings.from_env: every value validated at startup, a bad one refused naming the variable.
|
||||
Contract: intern-decision-serve.contract.md § Configuration."""
|
||||
import pytest
|
||||
|
||||
from intern_decision_serve.config import DEFAULT_CHECKPOINT, Settings
|
||||
|
||||
TOKEN = "t" * 40
|
||||
P = "INTERN_DECISION_"
|
||||
|
||||
|
||||
def env(**kw):
|
||||
return {f"{P}API_TOKEN": TOKEN, **{f"{P}{k}": v for k, v in kw.items()}}
|
||||
|
||||
|
||||
def test_defaults():
|
||||
s = Settings.from_env(env())
|
||||
assert (s.api_token, s.checkpoint, s.device, s.vram_cap_gib) == (TOKEN, DEFAULT_CHECKPOINT, "cuda", None)
|
||||
assert (s.max_tokens, s.max_decisions, s.max_body_bytes, s.max_queue) == (8192, 64, 1024 * 1024, 32)
|
||||
assert (s.release_slack_mib, s.keep_vision) == (512, False)
|
||||
|
||||
|
||||
def test_every_value_is_read():
|
||||
s = Settings.from_env(env(CHECKPOINT="/x", DEVICE="cpu", VRAM_CAP_GIB="10.5", MAX_TOKENS="6000",
|
||||
MAX_DECISIONS="32", MAX_BODY_BYTES="2048", MAX_QUEUE="4", RELEASE_SLACK_MIB="0",
|
||||
KEEP_VISION="1"))
|
||||
assert (s.checkpoint, s.device, s.vram_cap_gib, s.max_tokens, s.max_decisions) == ("/x", "cpu", 10.5, 6000, 32)
|
||||
assert (s.max_body_bytes, s.max_queue, s.release_slack_mib, s.keep_vision) == (2048, 4, 0, True)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("token", ["", "short", "x" * 31, "x" * 40 + " ", "x" * 40 + "\n", "x" * 39 + "é"])
|
||||
def test_a_short_or_non_visible_ascii_token_is_refused(token):
|
||||
with pytest.raises(ValueError, match=f"{P}API_TOKEN"):
|
||||
Settings.from_env({f"{P}API_TOKEN": token})
|
||||
|
||||
|
||||
@pytest.mark.parametrize("name, value", [
|
||||
("VRAM_CAP_GIB", "0"), ("VRAM_CAP_GIB", "-1"), ("VRAM_CAP_GIB", "nan"), ("VRAM_CAP_GIB", "inf"),
|
||||
("VRAM_CAP_GIB", "lots"), ("MAX_TOKENS", "0"), ("MAX_TOKENS", "1.5"), ("MAX_DECISIONS", "0"),
|
||||
("MAX_BODY_BYTES", "x"), ("MAX_QUEUE", "-2"), ("RELEASE_SLACK_MIB", "-1"), ("KEEP_VISION", "yes"),
|
||||
("DEVICE", "mps"),
|
||||
])
|
||||
def test_a_bad_value_is_refused_naming_the_variable(name, value):
|
||||
with pytest.raises(ValueError, match=f"{P}{name}"):
|
||||
Settings.from_env(env(**{name: value}))
|
||||
|
||||
|
||||
def test_an_empty_cap_means_uncapped():
|
||||
assert Settings.from_env(env(VRAM_CAP_GIB="")).vram_cap_gib is None
|
||||
@@ -0,0 +1,321 @@
|
||||
"""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('''
|
||||
import json
|
||||
MODEL_NAME = "Intern-Decision-4B"
|
||||
EVENTS = []
|
||||
WARMUP_SHIFT = {"after_swap": 0.0}
|
||||
TOKEN_SKEW = {"n": 0}
|
||||
|
||||
class Compiled:
|
||||
def __init__(self, messages):
|
||||
self.messages = messages
|
||||
|
||||
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_value_errors_pass_through_untouched():
|
||||
def bad(_request):
|
||||
raise ValueError("Example has 9000 tokens, above 8192; truncation is forbidden")
|
||||
with pytest.raises(ValueError, match="truncation is forbidden"):
|
||||
direct(bad).predict(REQUEST)
|
||||
assert FakeTorch.cuda.empties == []
|
||||
|
||||
|
||||
@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
|
||||
@@ -0,0 +1,23 @@
|
||||
"""Entry point: INV-5, offline before torch/transformers load. Contract: intern-decision-serve.contract.md."""
|
||||
import os
|
||||
|
||||
from fake_engine import FakeEngine
|
||||
|
||||
|
||||
def test_app_from_env_goes_offline_before_the_engine_loads(monkeypatch):
|
||||
from intern_decision_serve import engine as engine_module
|
||||
from intern_decision_serve import main
|
||||
seen = {}
|
||||
|
||||
def fake_load(settings, **_kw):
|
||||
seen["offline"] = (os.environ.get("HF_HUB_OFFLINE"), os.environ.get("TRANSFORMERS_OFFLINE"))
|
||||
seen["token"] = settings.api_token
|
||||
return FakeEngine()
|
||||
|
||||
monkeypatch.delenv("HF_HUB_OFFLINE", raising=False)
|
||||
monkeypatch.delenv("TRANSFORMERS_OFFLINE", raising=False)
|
||||
monkeypatch.setenv("INTERN_DECISION_API_TOKEN", "k" * 40)
|
||||
monkeypatch.setattr(engine_module.TorchEngine, "load", staticmethod(fake_load))
|
||||
app = main.app_from_env()
|
||||
assert seen == {"offline": ("1", "1"), "token": "k" * 40}
|
||||
assert app.title == "intern-decision-serve"
|
||||
Generated
+1282
File diff suppressed because it is too large
Load Diff
Reference in New Issue
Block a user