Codeseys/composer-replication-framework
0
1"""Wave 20 — chat-template alignment regression guard for the PACKAGE collator.2 3`composer_replication.trainer.data_collator.ComposerDataCollator` builds the4SDPO `sdpo_loss_mask` (and the aligned-student `response_mask`) so that in-loss5positions sit exactly on content tokens. The hard part is that6`apply_chat_template` inserts role/BOS/EOS scaffolding around each message; the7old `_build_segment_mask` tokenized each content string in isolation and8concatenated, so the mask drifted left of the real content tokens. The Wave 199production audit measured this drift at ~67% aligned. Wave 20's10`_build_chat_aligned_mask` derives the mask from per-message11`apply_chat_template` prefix deltas instead, restoring ~100% alignment.12 13These tests use a REAL chat-template tokenizer (the stub used by14spikes/005 cannot expose the drift — its `apply_chat_template` adds no15scaffolding). They skip cleanly when transformers / the model cache is absent.16"""17from __future__ import annotations18 19import pytest20 21from composer_replication.trainer.data_collator import (22 CollatorConfig,23 ComposerDataCollator,24)25 26 27def _load_real_chat_tokenizer():28 """Return a real tokenizer with a chat template, or None to skip."""29 try:30 import os31 32 os.environ.setdefault("HF_HUB_OFFLINE", "1")33 os.environ.setdefault("TRANSFORMERS_OFFLINE", "1")34 from transformers import AutoTokenizer35 except Exception:36 return None37 for model in ("Qwen/Qwen2.5-0.5B-Instruct", "Qwen/Qwen2.5-1.5B-Instruct"):38 try:39 t = AutoTokenizer.from_pretrained(model)40 if getattr(t, "chat_template", None):41 return t42 except Exception:43 continue44 return None45 46 47_REAL_TOK = _load_real_chat_tokenizer()48_SKIP_REASON = "real chat-template tokenizer not available (offline / not cached)"49 50 51@pytest.fixture52def real_chat_tok():53 if _REAL_TOK is None:54 pytest.skip(_SKIP_REASON)55 return _REAL_TOK56 57 58@pytest.fixture59def multiturn_error_trace():60 """Multi-turn trace with an error site after several turns, so the61 chat-template scaffolding drift compounds (what exposed the old 33%)."""62 return {63 "trace_id": "real-align-1",64 "turns": [65 {"role": "user", "content": "Read /etc/app/config.yaml and summarize it."},66 {"role": "assistant", "content": '[TOOL_USE] name=Read input={"path":"/etc/app/config.yaml"}'},67 {"role": "user", "content": "[TOOL_RESULT (ERROR)] (id=t1)\nError: no such file or directory"},68 {69 "role": "assistant",70 "content": "The file does not exist there. Let me search for it instead.",71 "tool_error": "file_not_found",72 "error_meta": {"source_role": "user"},73 },74 {"role": "user", "content": "[TOOL_RESULT] (id=t2)\nFound /opt/app/config.yaml"},75 {"role": "assistant", "content": "Found it at /opt/app/config.yaml. Reading now."},76 ],77 "final_reward": 0.0,78 }79 80 81def _hint_gen(kind, _meta):82 return f"The path was wrong (kind: {kind}). Search with Glob before reading."83 84 85def test_real_chat_template_sdpo_mask_fully_aligned(real_chat_tok, multiturn_error_trace):86 """THE Wave 20 guarantee: with a REAL chat template, every in-loss87 sdpo_loss_mask position must have student==teacher token id. Before the88 fix this drifted to ~67% because the mask was built from per-segment89 tokenization that ignored apply_chat_template scaffolding."""90 cfg = CollatorConfig(hint_generator=_hint_gen, enable_replay_dpo=False)91 collator = ComposerDataCollator(tokenizer=real_chat_tok, config=cfg)92 batch = collator([multiturn_error_trace])93 94 assert "sdpo_loss_mask" in batch, "SDPO channel did not fire on the error trace"95 s_in = batch["input_ids"]96 t_in = batch["ctx_teacher_input_ids"]97 m_in = batch["sdpo_loss_mask"]98 assert s_in.shape == t_in.shape == m_in.shape99 100 n_aligned = n_total = 0101 for row in range(s_in.shape[0]):102 in_loss = m_in[row] == 1103 if int(in_loss.sum()) == 0:104 continue105 s_at = s_in[row][in_loss]106 t_at = t_in[row][in_loss]107 n_aligned += int((s_at == t_at).sum().item())108 n_total += int(in_loss.sum().item())109 110 assert n_total > 0, "No in-loss positions — SDPO mask is empty"111 ratio = n_aligned / n_total112 assert ratio >= 0.95, (113 f"SDPO mask alignment is only {100 * ratio:.1f}% ({n_aligned}/{n_total}); "114 f"the chat-template drift fix has regressed. Expected ~100%."115 )116 117 118def test_real_chat_template_in_loss_tokens_are_content_not_scaffolding(119 real_chat_tok, multiturn_error_trace120):121 """The in-loss teacher tokens must decode to the recovery turn's CONTENT,122 not chat-template markers (<|im_start|>, role strings, etc.)."""123 cfg = CollatorConfig(hint_generator=_hint_gen, enable_replay_dpo=False)124 collator = ComposerDataCollator(tokenizer=real_chat_tok, config=cfg)125 batch = collator([multiturn_error_trace])126 127 t_in = batch["ctx_teacher_input_ids"][0]128 m_in = batch["sdpo_loss_mask"][0]129 in_loss = m_in == 1130 decoded = real_chat_tok.decode(t_in[in_loss].tolist())131 assert "does not exist" in decoded, (132 f"In-loss tokens don't contain the recovery content; got: {decoded!r}"133 )134 for marker in ("<|im_start|>", "<|im_end|>", "<|endoftext|>"):135 assert marker not in decoded, (136 f"Chat-template marker {marker!r} leaked into the in-loss span: {decoded!r}"137 )138 139 140def test_real_chat_template_student_teacher_shapes_match(real_chat_tok, multiturn_error_trace):141 """The SDPO gate requires student_logits.shape == teacher_logits.shape;142 verify the aligned-student path produces matching sequence lengths."""143 cfg = CollatorConfig(hint_generator=_hint_gen, enable_replay_dpo=False)144 collator = ComposerDataCollator(tokenizer=real_chat_tok, config=cfg)145 batch = collator([multiturn_error_trace])146 assert batch["input_ids"].shape == batch["ctx_teacher_input_ids"].shape147 148 149# ----------------------------------------------------------------------------150# Empty-recovery guard (Wave 21 — discovered on real Claude Code traces)151# ----------------------------------------------------------------------------152#153# ~67% of real error sites have EMPTY recovery content: when strip_thinking=True154# the recovery turn (which was pure [THINKING] reasoning) becomes empty. Injecting155# an SDPO hint with no recovery content yields an all-ignore_index mask — a156# zero-signal row that wastes a forward pass and dilutes the channel. The collator157# must treat empty-recovery error turns as non-error sites. These use a stub158# tokenizer (pure logic, no model needed) so they always run.159 160 161class _StubTok:162 """Word-level deterministic tokenizer; apply_chat_template space-joins."""163 164 pad_token_id = 0165 166 def __init__(self) -> None:167 self._v: dict[str, int] = {"<pad>": 0, "<bos>": 1, "<eos>": 2}168 169 def _id(self, w: str) -> int:170 if w not in self._v:171 self._v[w] = len(self._v)172 return self._v[w]173 174 def __call__(self, text, **_k):175 return {"input_ids": [self._id(w) for w in text.split()] if text else []}176 177 def apply_chat_template(self, messages, tokenize=True, **_k): # noqa: ARG002178 return [self._id(w) for w in " ".join(m.get("content", "") for m in messages).split()]179 180 181def _hint_for_tnf(kind, _meta):182 return "HINT use a real tool" if kind == "tool_not_found" else None183 184 185def test_empty_recovery_does_not_fire_sdpo():186 """An error turn with empty recovery content must NOT emit an SDPO mask."""187 tok = _StubTok()188 trace = {189 "trace_id": "empty-recovery",190 "turns": [191 {"role": "user", "content": "do the thing"},192 {"role": "assistant", "content": "", "tool_error": "tool_not_found", "error_meta": {}},193 {"role": "user", "content": "tool not found"},194 ],195 "final_reward": 0.0,196 }197 cfg = CollatorConfig(hint_generator=_hint_for_tnf)198 collator = ComposerDataCollator(tokenizer=tok, config=cfg)199 batch = collator([trace])200 assert "sdpo_loss_mask" not in batch, (201 "Empty-recovery error turn fired a zero-signal SDPO mask; it must be skipped."202 )203 204 205def test_mixed_recovery_fires_on_nonempty_only():206 """A trace mixing empty + non-empty recovery turns fires SDPO from the207 non-empty one and has loss-active positions."""208 tok = _StubTok()209 trace = {210 "trace_id": "mixed-recovery",211 "turns": [212 {"role": "user", "content": "first task"},213 {"role": "assistant", "content": "", "tool_error": "tool_not_found", "error_meta": {}},214 {"role": "user", "content": "tool not found"},215 {"role": "assistant", "content": "let me use a real tool instead",216 "tool_error": "tool_not_found", "error_meta": {}},217 ],218 "final_reward": 0.0,219 }220 cfg = CollatorConfig(hint_generator=_hint_for_tnf)221 collator = ComposerDataCollator(tokenizer=tok, config=cfg)222 batch = collator([trace])223 assert "sdpo_loss_mask" in batch224 assert int((batch["sdpo_loss_mask"] == 1).sum()) > 0225 226 227def test_empty_recovery_keeps_student_teacher_shapes_matched():228 """Even with a skipped empty-recovery turn, when SDPO DOES fire elsewhere229 the student/teacher shapes must still match (lockstep skip on both sides)."""230 tok = _StubTok()231 trace = {232 "trace_id": "mixed-shape",233 "turns": [234 {"role": "user", "content": "task"},235 {"role": "assistant", "content": "", "tool_error": "tool_not_found", "error_meta": {}},236 {"role": "user", "content": "tool not found"},237 {"role": "assistant", "content": "recover now with a real tool",238 "tool_error": "tool_not_found", "error_meta": {}},239 ],240 "final_reward": 0.0,241 }242 cfg = CollatorConfig(hint_generator=_hint_for_tnf)243 collator = ComposerDataCollator(tokenizer=tok, config=cfg)244 batch = collator([trace])245 assert batch["input_ids"].shape == batch["ctx_teacher_input_ids"].shape246 