Team Ai
Modelpublic

Codeseys/composer-replication-framework

sourceHugging Facemitupdated 4mo agoView on Hugging Face
0likes
test_chat_template_alignment.py246 linesDownload Raw Back to tests
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