Skip to content

Commit c97673e

Browse files
committed
[fix][generators] VLM generator: honor generator.chat_template in renders + image-token/feature integrity guard
Fixes NovaSky-AI#2075. The VLM generator's deferred obs-token extraction requires each re-render to be a token-prefix extension of the previous one (the NOTE in agent_loop). Thinking-model vendor templates strip reasoning from history and violate that, silently corrupting trajectories -- loud only when images make the token/feature pairing fail inside the model forward. * _render_conversation forwards chat_template / chat_template_kwargs (both already exist on the base generator via generator.chat_template, including source=file; the VLM path never sent them). A thinking- preserving template restores the prefix property and makes training on-policy. Servers need trust_request_chat_template for per-request templates. * Trajectory-end integrity guard via the render response's own mm_placeholders: actionable per-trajectory error instead of the opaque torch._check deep into training. Validated end to end: Qwen3.5-9B multi-turn computer-use GRPO (64-turn episodes with screenshots), 6 GRPO steps, one epoch, completed cleanly with a thinking-preserving template supplied via source=file; the same setup crashed at the first policy update twice without these changes (traces in NovaSky-AI#2075).
1 parent cb5334c commit c97673e

2 files changed

Lines changed: 79 additions & 3 deletions

File tree

‎skyrl/train/generators/skyrl_vlm_generator.py‎

Lines changed: 38 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -59,9 +59,20 @@ def _validate_cfg(self, generator_cfg: GeneratorConfig):
5959
)
6060

6161
async def _render_conversation(self, conversation: ConversationType) -> RenderedConversation:
62-
rendered = await self.inference_engine_client.render_chat_completion(
63-
{"json": {"model": self.inference_engine_client.model_name, "messages": conversation}}
64-
)
62+
body: Dict[str, Any] = {"model": self.inference_engine_client.model_name, "messages": conversation}
63+
# Honor generator.chat_template / chat_template_kwargs (already supported
64+
# by the base generator) on the render path. The deferred obs-token
65+
# extraction below requires each re-render to be a token-prefix
66+
# extension of the previous one (see NOTE in agent_loop); thinking-model
67+
# vendor templates strip reasoning from history and violate that, so a
68+
# thinking-preserving custom template is the supported way to train
69+
# them (issue #2075). Servers need trust_request_chat_template for
70+
# per-request templates.
71+
if self.custom_chat_template:
72+
body["chat_template"] = self.custom_chat_template
73+
if self.generator_cfg.chat_template_kwargs:
74+
body["chat_template_kwargs"] = dict(self.generator_cfg.chat_template_kwargs)
75+
rendered = await self.inference_engine_client.render_chat_completion({"json": body})
6576
return RenderedConversation(prompt_ids=rendered["token_ids"], features=rendered.get("features", None))
6677

6778
async def agent_loop(
@@ -129,12 +140,14 @@ async def agent_loop(
129140
# sequences which are a prefix of subsequent turns. In general, this
130141
# will not hold for thinking models, e.g., Qwen3-Thinking.
131142
pending_obs_offset: Optional[int] = None
143+
last_render_ids: Optional[List[int]] = None
132144

133145
while not done:
134146
# 1. Render full conversation for this turn's generation input
135147
rendered_conversation = await self._render_conversation(conversation)
136148
input_ids = rendered_conversation["prompt_ids"]
137149
latest_features = rendered_conversation["features"]
150+
last_render_ids = input_ids
138151

139152
# 1b. Flush pending obs tokens from the previous turn
140153
if pending_obs_offset is not None:
@@ -209,6 +222,28 @@ async def agent_loop(
209222
pixel_values = mm_kwargs["pixel_values"]
210223
image_grid_thw = mm_kwargs["image_grid_thw"]
211224

225+
# ── Integrity guard: image tokens vs features ─────────────────
226+
# The model forward pairs each image token with exactly one
227+
# feature row; a mismatch fails there as an opaque torch._check
228+
# deep into training. Detect it here, per trajectory, with an
229+
# actionable message. Uses the render response's mm_placeholders,
230+
# so no model constants are needed.
231+
placeholders = (latest_features or {}).get("mm_placeholders", {}).get("image", [])
232+
if placeholders:
233+
render_ids = last_render_ids if last_render_ids is not None else prompt_ids
234+
ph0 = placeholders[0]
235+
image_token = render_ids[ph0["offset"] + ph0["length"] // 2]
236+
expected = sum(ph["length"] for ph in placeholders)
237+
actual = (list(prompt_ids) + list(response_ids)).count(image_token)
238+
if actual != expected:
239+
raise RuntimeError(
240+
f"trajectory image-token mismatch: sequence has {actual} image tokens, "
241+
f"render declares {expected} across {len(placeholders)} images "
242+
f"(stop_reason={stop_reason}). The trajectory token stream diverged from "
243+
f"the render -- typically a history-editing chat template breaking the "
244+
f"prefix-extension assumption (issue #2075)."
245+
)
246+
212247
# ── Cleanup ───────────────────────────────────────────────────
213248
env_metrics = env.get_metrics()
214249
await self._run_in_executor_if_available(env.close)
Lines changed: 41 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,41 @@
1+
"""The VLM render path honors generator.chat_template (issue #2075).
2+
3+
Unit seam: _render_conversation built without __init__ so the test needs no
4+
engines. Asserts the request body carries the custom template exactly when
5+
configured, and stays untouched otherwise.
6+
"""
7+
8+
from types import SimpleNamespace
9+
from unittest.mock import AsyncMock
10+
11+
import pytest
12+
13+
from skyrl.train.generators.skyrl_vlm_generator import SkyRLVLMGymGenerator
14+
15+
16+
def _bare_generator(custom_template, kwargs=None):
17+
gen = SkyRLVLMGymGenerator.__new__(SkyRLVLMGymGenerator)
18+
gen.custom_chat_template = custom_template
19+
gen.generator_cfg = SimpleNamespace(chat_template_kwargs=kwargs or {})
20+
gen.inference_engine_client = SimpleNamespace(
21+
model_name="test-model",
22+
render_chat_completion=AsyncMock(return_value={"token_ids": [1, 2], "features": None}),
23+
)
24+
return gen
25+
26+
27+
@pytest.mark.asyncio
28+
async def test_custom_template_reaches_render_body():
29+
gen = _bare_generator("{{ messages }}", {"enable_thinking": True})
30+
await gen._render_conversation([{"role": "user", "content": "hi"}])
31+
body = gen.inference_engine_client.render_chat_completion.call_args.args[0]["json"]
32+
assert body["chat_template"] == "{{ messages }}"
33+
assert body["chat_template_kwargs"] == {"enable_thinking": True}
34+
35+
36+
@pytest.mark.asyncio
37+
async def test_no_template_leaves_body_untouched():
38+
gen = _bare_generator(None)
39+
await gen._render_conversation([{"role": "user", "content": "hi"}])
40+
body = gen.inference_engine_client.render_chat_completion.call_args.args[0]["json"]
41+
assert "chat_template" not in body and "chat_template_kwargs" not in body

0 commit comments

Comments
 (0)