Skip to content

Commit 67dc0ad

Browse files
authored
[llm_server] Fix transcript replay and reasoning response handling (pytorch#22852)
Fall back to rendered text when the assistant boundary cannot be verified. Compare supplied reasoning against the finalized client-visible response before replaying stored tokens, preserving omitted reasoning while invalidating explicit edits. Discard whitespace-only Muse Glimmer blocks and reject non-boolean return_reasoning values before generation. Document the response contract and verify complete and streaming turns with full prompt assertions. Run serving tests for relevant pull requests, including the Muse adapter tests. Make tokenizer paths explicit and cover exact BPE prompt assembly with a portable fixture.
1 parent 338ac76 commit 67dc0ad

20 files changed

Lines changed: 842 additions & 203 deletions

‎.github/workflows/_llm_server.yml‎

Lines changed: 5 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -36,12 +36,15 @@ jobs:
3636
python -m pip install --progress-bar off \
3737
-r examples/llm_server/python/requirements.txt \
3838
httpx \
39-
pytest
39+
pytest \
40+
tokenizers
4041
4142
export PYTHONPATH="$(dirname "${PWD}"):${PYTHONPATH:-}"
4243
export PYTHONDONTWRITEBYTECODE=1
4344
44-
python -m pytest -q examples/llm_server/python/tests
45+
python -m pytest -q \
46+
examples/llm_server/python/tests \
47+
examples/models/muse-glimmer/tests/test_serve.py
4548
4649
cmake -S . -B cmake-out \
4750
-DCMAKE_BUILD_TYPE=Release \

‎.github/workflows/pull.yml‎

Lines changed: 22 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -752,6 +752,28 @@ jobs:
752752
id-token: write
753753
contents: read
754754

755+
llm-server:
756+
needs: [changed-files, run-decision]
757+
if: |
758+
github.event_name == 'pull_request' && (
759+
contains(needs.changed-files.outputs.changed-files, 'examples/llm_server') ||
760+
contains(needs.changed-files.outputs.changed-files, 'examples/models/muse-glimmer/serving') ||
761+
contains(needs.changed-files.outputs.changed-files, 'examples/models/muse-glimmer/tests') ||
762+
contains(needs.changed-files.outputs.changed-files, 'examples/models/muse_glimmer') ||
763+
contains(needs.changed-files.outputs.changed-files, 'extension/llm/runner') ||
764+
contains(needs.changed-files.outputs.changed-files, 'extension/llm/tokenizers') ||
765+
contains(needs.changed-files.outputs.changed-files, '.github/workflows/_get-changed-files.yml') ||
766+
contains(needs.changed-files.outputs.changed-files, '.github/workflows/_ci-run-decision.yml') ||
767+
contains(needs.changed-files.outputs.changed-files, '.github/workflows/_llm_server.yml') ||
768+
contains(needs.changed-files.outputs.changed-files, '.github/workflows/pull.yml') ||
769+
contains(needs.changed-files.outputs.changed-files, '.github/workflows/trunk.yml') ||
770+
needs.run-decision.outputs.is-full-run == 'true'
771+
)
772+
uses: ./.github/workflows/_llm_server.yml
773+
permissions:
774+
id-token: write
775+
contents: read
776+
755777
unittest:
756778
uses: ./.github/workflows/_unittest.yml
757779
permissions:

‎.github/workflows/trunk.yml‎

Lines changed: 14 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -994,12 +994,20 @@ jobs:
994994
llm-server:
995995
needs: [changed-files, run-decision]
996996
if: |
997-
contains(needs.changed-files.outputs.changed-files, 'examples/llm_server') ||
998-
contains(needs.changed-files.outputs.changed-files, 'extension/llm/runner') ||
999-
contains(needs.changed-files.outputs.changed-files, 'extension/llm/tokenizers') ||
1000-
contains(needs.changed-files.outputs.changed-files, '.github/workflows/_llm_server.yml') ||
1001-
contains(needs.changed-files.outputs.changed-files, '.github/workflows/trunk.yml') ||
1002-
needs.run-decision.outputs.is-full-run == 'true'
997+
github.event_name != 'pull_request' && (
998+
contains(needs.changed-files.outputs.changed-files, 'examples/llm_server') ||
999+
contains(needs.changed-files.outputs.changed-files, 'examples/models/muse-glimmer/serving') ||
1000+
contains(needs.changed-files.outputs.changed-files, 'examples/models/muse-glimmer/tests') ||
1001+
contains(needs.changed-files.outputs.changed-files, 'examples/models/muse_glimmer') ||
1002+
contains(needs.changed-files.outputs.changed-files, 'extension/llm/runner') ||
1003+
contains(needs.changed-files.outputs.changed-files, 'extension/llm/tokenizers') ||
1004+
contains(needs.changed-files.outputs.changed-files, '.github/workflows/_get-changed-files.yml') ||
1005+
contains(needs.changed-files.outputs.changed-files, '.github/workflows/_ci-run-decision.yml') ||
1006+
contains(needs.changed-files.outputs.changed-files, '.github/workflows/_llm_server.yml') ||
1007+
contains(needs.changed-files.outputs.changed-files, '.github/workflows/pull.yml') ||
1008+
contains(needs.changed-files.outputs.changed-files, '.github/workflows/trunk.yml') ||
1009+
needs.run-decision.outputs.is-full-run == 'true'
1010+
)
10031011
uses: ./.github/workflows/_llm_server.yml
10041012
permissions:
10051013
id-token: write

‎examples/llm_server/python/README.md‎

Lines changed: 12 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -60,12 +60,20 @@ Key flags:
6060
| Flag | Effect |
6161
|------|--------|
6262
| `--hf-tokenizer` | model's HF chat template (required unless fallback) |
63+
| `--assistant-header` | exact assistant generation header, including trailing whitespace (default: ChatML) |
6364
| `--allow-chatml-fallback` | opt into approximate ChatML when no HF tokenizer |
6465
| `--no-think` | default `enable_thinking=False` (e.g. Qwen3) |
6566
| `--max-context N` | reject over-long prompts with 400 instead of failing mid-gen |
6667
| `--num-runners N` | Worker processes — **1 only** (one worker hosts many isolated sessions on one weight load; more would duplicate weights) |
6768
| `--worker-bin PATH` | path to a model worker binary that speaks the llm_server JSONL protocol |
6869

70+
Set `--assistant-header` to the model template's exact generation boundary when
71+
it differs from ChatML. For Llama 3 templates, add
72+
`--assistant-header $'<|start_header_id|>assistant<|end_header_id|>\n\n'` in Bash;
73+
the `$'...'` quoting supplies literal newlines. The launcher warns once at startup
74+
if the configured header is absent from a rendered probe. Unverified boundaries
75+
use the rendered text, which can reduce KV reuse.
76+
6977
## Smoke test
7078

7179
```bash
@@ -114,7 +122,7 @@ Two layers, both contract-focused (assert on the wire, not internals):
114122

115123
```bash
116124
# 1. Model-free tests — unit coverage plus loopback disconnect integration.
117-
pip install pytest httpx
125+
pip install pytest httpx tokenizers
118126
pytest tests/
119127

120128
# 2. Conformance — black-box, against a LIVE server (real model, or llama.cpp/mlx-lm).
@@ -126,6 +134,9 @@ real server/protocol/streaming code is tested over HTTP without a `.pte`. The
126134
worker JSONL protocol is covered separately by `tests/test_worker_client.py`,
127135
and `tests/test_stream_disconnect.py` uses real loopback Uvicorn/TCP plus a
128136
model-free subprocess to verify disconnect cancellation end to end.
137+
The BPE splice tests use an in-memory tokenizer with no model downloads. Optional
138+
integration tests use local tokenizer directories set with `QWEN_HF_DIR`,
139+
`GEMMA_HF_DIR`, or `MUSE_GLIMMER_HF_DIR` and require `transformers`.
129140

130141
## Architecture
131142

‎examples/llm_server/python/chat_template.py‎

Lines changed: 10 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -106,6 +106,7 @@ def __init__(
106106
# chat_template_kwargs override these.
107107
self._defaults = default_template_kwargs or {}
108108
self._assistant_header = assistant_header
109+
self._warned_assistant_header = False
109110
self._strip_rendered_prefix = strip_rendered_prefix
110111
self._append_generation_prompt_after_tool_response = (
111112
append_generation_prompt_after_tool_response
@@ -229,8 +230,6 @@ def generation_preamble(
229230
the resident one (for Qwen3 the scaffold is tool-independent -> same key).
230231
Returns ``""`` for the fallback / no-scaffold templates (fix is a no-op).
231232
"""
232-
if self._hf is None:
233-
return ""
234233
merged = {**self._defaults, **(template_kwargs or {})}
235234
if tools:
236235
try:
@@ -251,7 +250,15 @@ def generation_preamble(
251250
template_kwargs=template_kwargs,
252251
)
253252
marker = self._assistant_header
254-
idx = rendered.rfind(marker)
253+
idx = rendered.rfind(marker) if marker else -1
254+
if idx == -1 and not self._warned_assistant_header:
255+
logger.warning(
256+
"Assistant header %r was not found in the chat template's generation "
257+
"prompt. Stored-token replay may fall back to rendered text; configure "
258+
"assistant_header (--assistant-header for the generic server).",
259+
marker,
260+
)
261+
self._warned_assistant_header = True
255262
preamble = rendered[idx + len(marker) :] if idx != -1 else ""
256263
self._preamble_cache[key] = preamble
257264
return preamble

‎examples/llm_server/python/openai_transcript.py‎

Lines changed: 49 additions & 45 deletions
Original file line numberDiff line numberDiff line change
@@ -16,10 +16,10 @@
1616
token ids and a fingerprint of the response. On the next request each prior
1717
assistant turn is replaced with a sentinel, the conversation is rendered once,
1818
and the rendered text is split on the sentinels with the stored ids spliced back
19-
in -- only for turns whose fingerprint matches the incoming message (an edited,
20-
branched, or reused history is never substituted with stale ids) and whose ids
21-
are present (a stop-trimmed turn is left as text). The worker's exact-token
22-
prefix check is the final backstop.
19+
in -- only for turns whose content/tool fingerprint and any supplied reasoning string
20+
match the recorded response, and whose ids are present (a stop-trimmed turn is
21+
left as text). An edit invalidates that turn and all later records. The worker's
22+
exact-token prefix check separately protects KV reuse.
2323
"""
2424

2525
import hashlib
@@ -111,9 +111,8 @@ def __init__(self, template: ChatTemplate):
111111
# boundary verification.
112112
_header_fn = getattr(template, "assistant_header", None)
113113
self._assist_hdr = _header_fn() if _header_fn else _ASSIST_HDR
114-
# session_id -> [{"fp": str, "ids": list[int] | None}, ...] (one per
115-
# assistant turn we produced, in order). Cleared on reset/close.
116-
self._turns: dict[str, list[dict]] = {}
114+
# Keyed by assistant-turn index, including gaps after invalidation.
115+
self._turns: dict[str, dict[int, dict]] = {}
117116

118117
@staticmethod
119118
def _assistant_fingerprint(content, tool_calls) -> str:
@@ -135,19 +134,24 @@ def _assistant_fingerprint(content, tool_calls) -> str:
135134
blob = json.dumps([content or "", norm], sort_keys=True, ensure_ascii=False)
136135
return hashlib.sha1(blob.encode("utf-8")).hexdigest()
137136

137+
@staticmethod
138+
def _reasoning_fingerprint(reasoning_content: Optional[str]) -> Optional[bytes]:
139+
if reasoning_content is None:
140+
return None
141+
return hashlib.sha256(reasoning_content.encode("utf-8")).digest()
142+
138143
def _normalize_scaffold(self, text_chunk: str, preamble: str) -> Optional[str]:
139144
"""Force the scaffold region (between the last assistant header in
140145
`text_chunk` and its end) to equal `preamble`, so the worker re-tokenizes
141146
the exact resident scaffold. The region is empty (history stripped it ->
142147
insert) or a think scaffold (history preserved it -> replace). Returns the
143148
adjusted text, or None if it isn't a recognized scaffold (-> text fallback)."""
149+
if not self._assist_hdr:
150+
return None
144151
h = text_chunk.rfind(self._assist_hdr)
145152
if h == -1:
146-
# No assistant header: with a scaffold to reproduce this is
147-
# ambiguous (-> text fallback); without one there is nothing to
148-
# normalize, so splicing still works for templates with a different
149-
# assistant header.
150-
return None if preamble else text_chunk
153+
# Without a verified boundary, splicing can duplicate template framing.
154+
return None
151155
base = h + len(self._assist_hdr)
152156
if not preamble:
153157
# No generation scaffold: the worker prefills nothing ahead of the
@@ -266,40 +270,37 @@ def build_prompt_input(
266270
"""Return a PromptInput: token-ID segments when this session has faithful
267271
stored ids for matching prior assistant turns, else the plain rendered
268272
text. Each incoming assistant turn is matched IN ORDER against the stored
269-
records and only spliced when (a) its fingerprint matches what we returned
270-
(else the history diverged -> stop, splice nothing further) and (b) we
271-
kept faithful ids for it (a stop-trimmed turn's None -> rendered as text).
273+
records and only spliced when its content/tool calls and any supplied
274+
reasoning string match what we returned, and we kept faithful ids for it.
275+
Omitted or null reasoning permits reuse; a string edit invalidates the tail.
272276
Falls back to text on a sentinel collision or a render that
273277
dropped/duplicated a sentinel."""
274278
stored = self._turns.get(session_id or "")
275279
if not stored:
276280
return PromptInput(text=rendered_prompt)
277-
# Positional: stored[k] is the k-th assistant turn WE generated, matched
278-
# against the k-th assistant message in the request. A client-injected
279-
# turn (few-shot exemplar, pre-seeded turn, reused session) shifts that
280-
# alignment -> fingerprint mismatch at k -> stop splicing. Always safe
281-
# (text fallback + worker prefix backstop); just a lower hit rate.
281+
# Missing records render as text without shifting later turn indices.
282282
positions = [i for i, m in enumerate(messages) if m.role == "assistant"]
283283
splice: dict[int, dict] = {} # message index -> {"ids", "preamble"}
284-
diverged_at = None
285284
for k, pos in enumerate(positions):
286-
if k >= len(stored):
287-
break
285+
record = stored.get(k)
286+
if record is None:
287+
continue
288288
m = messages[pos]
289-
if self._assistant_fingerprint(m.content, m.tool_calls) != stored[k]["fp"]:
290-
diverged_at = k # this stored turn and every later one are stale
289+
if self._assistant_fingerprint(m.content, m.tool_calls) != record["fp"] or (
290+
m.reasoning_content is not None
291+
and self._reasoning_fingerprint(m.reasoning_content)
292+
!= record["reasoning_fp"]
293+
):
294+
# Discard the stale tail without shifting subsequent turn indices.
295+
self._turns[session_id or ""] = {
296+
index: record for index, record in stored.items() if index < k
297+
}
291298
break
292-
if stored[k]["ids"] is not None:
299+
if record["ids"] is not None:
293300
splice[pos] = {
294-
"ids": stored[k]["ids"],
295-
"preamble": stored[k].get("preamble", ""),
301+
"ids": record["ids"],
302+
"preamble": record.get("preamble", ""),
296303
}
297-
if diverged_at is not None:
298-
# Drop the stale tail from the first mismatch so an edited/branched
299-
# earlier turn can't shadow future requests; the matched prefix still
300-
# splices, the rest stays text until reset/close. Safe either way:
301-
# stale ids are never spliced and the worker's prefix check backstops.
302-
del stored[diverged_at:]
303304
if not splice:
304305
return PromptInput(text=rendered_prompt)
305306
tool_splice = {
@@ -351,25 +352,28 @@ def record_assistant_turn(
351352
generated_token_ids: list,
352353
prior_turns: int,
353354
preamble: str = "",
355+
reasoning_content: Optional[str] = None,
354356
) -> None:
355357
"""Record this turn's {fingerprint, generated ids, generation preamble} at
356358
`prior_turns` (the assistant-turn count of the request it answers).
357359
Records at/after that index are dropped first, so a regenerated/branched
358360
turn replaces stale records rather than shadowing later hits. ids is None
359361
when the worker omitted them (stop-trimmed -> non-resumable), kept for
360-
positional alignment. `preamble` is the generation scaffold (e.g. the
361-
Qwen3 `<think>` block) reproduced ahead of the spliced ids next request."""
362+
positional alignment. `reasoning_content` is the client-visible value,
363+
including None when the client opted out. `preamble` is the generation
364+
scaffold (e.g. the Qwen3 `<think>` block) reproduced ahead of the spliced
365+
ids next request."""
362366
if not session_id:
363367
return
364-
turns = self._turns.setdefault(session_id, [])
365-
del turns[prior_turns:]
366-
turns.append(
367-
{
368-
"fp": self._assistant_fingerprint(content, tool_calls),
369-
"ids": list(generated_token_ids) if generated_token_ids else None,
370-
"preamble": preamble,
371-
}
372-
)
368+
turns = self._turns.setdefault(session_id, {})
369+
for index in [index for index in turns if index >= prior_turns]:
370+
del turns[index]
371+
turns[prior_turns] = {
372+
"fp": self._assistant_fingerprint(content, tool_calls),
373+
"reasoning_fp": self._reasoning_fingerprint(reasoning_content),
374+
"ids": list(generated_token_ids) if generated_token_ids else None,
375+
"preamble": preamble,
376+
}
373377

374378
def reset(self, session_id: str) -> None:
375379
self._turns.pop(session_id, None)

‎examples/llm_server/python/protocol.py‎

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -42,6 +42,7 @@ class ChatMessage(BaseModel):
4242
# ResponseMessage.reasoning_content so a multi-turn client can echo an assistant
4343
# turn's reasoning back into the request; a chat template that renders prior
4444
# reasoning needs it here, and without the field it is dropped at parse.
45+
# Omission/null allows stored-token replay; a string edit invalidates that turn.
4546
reasoning_content: Optional[str] = None
4647
tool_calls: Optional[list[ToolCall]] = None
4748
tool_call_id: Optional[str] = None

‎examples/llm_server/python/server.py‎

Lines changed: 8 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -171,6 +171,12 @@ def main() -> None:
171171
help="Allow approximate generic ChatML templating when --hf-tokenizer is absent. "
172172
"Off by default: the fallback can't reproduce model-specific controls.",
173173
)
174+
p.add_argument(
175+
"--assistant-header",
176+
default="<|im_start|>assistant\n",
177+
help="Exact template text ending the assistant generation header, including "
178+
"any trailing whitespace. Defaults to the ChatML header.",
179+
)
174180
p.add_argument(
175181
"--model-id", default="executorch", help="Model id reported on /v1/models"
176182
)
@@ -218,7 +224,9 @@ def main() -> None:
218224
args.hf_tokenizer,
219225
default_template_kwargs=default_template_kwargs,
220226
allow_fallback=args.allow_chatml_fallback,
227+
assistant_header=args.assistant_header,
221228
)
229+
template.generation_preamble()
222230
worker = _spawn(args) # one worker hosting many isolated sessions
223231
runtime = SessionRuntime(worker)
224232
serving = ServingChat(

0 commit comments

Comments
 (0)