Team Ai
Modelpublic

Codeseys/composer-replication-framework

sourceHugging Facemitupdated 4mo agoView on Hugging Face
0likes
1"""real_batch.py — build a real, tokenized 3-channel batch from a HF tokenizer.2 3Used by Spike 006's smoke to generate inputs for `compose_loss` from a real4chat-template-formatted conversation, NOT random ints.5"""6from __future__ import annotations7 8from typing import Any9 10import torch11 12 13def build_batch(14    tokenizer: Any,15    *,16    device: torch.device | str = "cpu",17    seed: int = 42,18    variant: str = "factorial",19    align_sdpo_shapes: bool = False,20) -> dict[str, torch.Tensor]:21    """Construct a full 3-channel input batch from a real tokenizer.22 23    Returns a dict with all keys `compose_loss` may consume:24        input_ids, response_mask25        ctx_teacher_input_ids, sdpo_loss_mask26        dpo_chosen_input_ids, dpo_chosen_response_mask27        dpo_rejected_input_ids, dpo_rejected_response_mask28        dpo_chosen_ref_logprobs, dpo_rejected_ref_logprobs29 30    The DPO ref logprobs are dummy tensors (not from a real reference policy31    forward); the smoke is verifying the loss composition wires together,32    not the reference-policy precompute pipeline.33 34    Args:35        tokenizer: real HF tokenizer36        device: torch device for the returned tensors37        seed: reproducibility — fixes torch.manual_seed before any random38            tensor (only the dummy logprobs use random; the chat-template39            text is deterministic)40        variant: "factorial" or "binary_search" — pick which canned41            conversation. Used by Spike 006-strict to alternate batches42            so the loss-decrease isn't memorization of a single sample.43        align_sdpo_shapes: if True, truncate ctx_teacher_input_ids to44            match input_ids length so the SDPO channel actually fires45            (no shape-mismatch fallback). Used by Spike 006-strict to46            exercise the SDPO loss on a real model.47    """48    torch.manual_seed(seed)49 50    # ------------------------------------------------------------------51    # Conversation 1: student rollout (variants for non-tautological tests)52    # ------------------------------------------------------------------53    if variant == "factorial":54        student_msgs = [55            {"role": "system", "content": "You are a careful coding assistant."},56            {"role": "user", "content": "Write a Python function to compute the factorial of n."},57            {"role": "assistant", "content": "def factorial(n):\n    if n <= 1: return 1\n    return n * factorial(n - 1)"},58        ]59        teacher_msgs = [60            {"role": "system", "content": "You are a careful coding assistant."},61            {"role": "user", "content": "Write a Python function to compute the factorial of n."},62            {"role": "user", "content": "[HINT] Recursion overflows for n>1000. Use an iterative loop."},63            {"role": "assistant", "content": "def factorial(n):\n    result = 1\n    for i in range(2, n + 1):\n        result *= i\n    return result"},64        ]65    elif variant == "binary_search":66        student_msgs = [67            {"role": "system", "content": "You are a careful coding assistant."},68            {"role": "user", "content": "Implement binary search in Python."},69            {"role": "assistant", "content": "def bsearch(a, t):\n    l, r = 0, len(a)\n    while l < r:\n        m = (l + r) // 2\n        if a[m] < t: l = m + 1\n        else: r = m\n    return l"},70        ]71        teacher_msgs = [72            {"role": "system", "content": "You are a careful coding assistant."},73            {"role": "user", "content": "Implement binary search in Python."},74            {"role": "user", "content": "[HINT] Use right = len(a) - 1 with inclusive upper bound is more standard."},75            {"role": "assistant", "content": "def bsearch(a, t):\n    l, r = 0, len(a) - 1\n    while l <= r:\n        m = (l + r) // 2\n        if a[m] == t: return m\n        if a[m] < t: l = m + 1\n        else: r = m - 1\n    return -1"},76        ]77    else:78        raise ValueError(f"unknown variant: {variant!r}")79 80    student_text = tokenizer.apply_chat_template(student_msgs, tokenize=False, add_generation_prompt=False)81    student_enc = tokenizer(student_text, return_tensors="pt", add_special_tokens=False)82    input_ids = student_enc["input_ids"].to(device)83 84    T = input_ids.shape[1]85    response_mask = torch.zeros_like(input_ids)86    response_mask[:, int(T * 0.7):] = 187 88    # ------------------------------------------------------------------89    # Conversation 2: hint-conditioned teacher context (SDPO)90    # ------------------------------------------------------------------91    teacher_text = tokenizer.apply_chat_template(teacher_msgs, tokenize=False, add_generation_prompt=False)92    teacher_enc = tokenizer(teacher_text, return_tensors="pt", add_special_tokens=False)93    ctx_teacher_input_ids = teacher_enc["input_ids"].to(device)94 95    if align_sdpo_shapes:96        # Truncate the teacher context to the student length so SDPO actually fires97        # (compose_loss falls back to zero when shapes mismatch). This is a98        # correctness-relaxing test mode — production will pad/align via the99        # real data collator, but for the smoke we just need the SDPO loss100        # to exercise the generalized_jsd_loss code path on a real HF model.101        T_t = ctx_teacher_input_ids.shape[1]102        if T_t > T:103            ctx_teacher_input_ids = ctx_teacher_input_ids[:, :T]104        elif T_t < T:105            pad_id = tokenizer.pad_token_id if tokenizer.pad_token_id is not None else tokenizer.eos_token_id106            pad = torch.full((1, T - T_t), pad_id, dtype=ctx_teacher_input_ids.dtype, device=device)107            ctx_teacher_input_ids = torch.cat([ctx_teacher_input_ids, pad], dim=1)108 109    T_t = ctx_teacher_input_ids.shape[1]110    sdpo_loss_mask = torch.zeros_like(ctx_teacher_input_ids)111    sdpo_loss_mask[:, int(T_t * 0.7):] = 1112 113    # ------------------------------------------------------------------114    # Conversation 3 + 4: DPO chosen / rejected pairs115    # ------------------------------------------------------------------116    dpo_chosen_msgs = [117        {"role": "system", "content": "You are a careful coding assistant."},118        {"role": "user", "content": "What's the time complexity of binary search?"},119        {"role": "assistant", "content": "Binary search is O(log n) because each comparison halves the search space."},120    ]121    dpo_rejected_msgs = [122        {"role": "system", "content": "You are a careful coding assistant."},123        {"role": "user", "content": "What's the time complexity of binary search?"},124        {"role": "assistant", "content": "It's O(n) I think, you have to look at every element."},125    ]126    chosen_text = tokenizer.apply_chat_template(dpo_chosen_msgs, tokenize=False, add_generation_prompt=False)127    rejected_text = tokenizer.apply_chat_template(dpo_rejected_msgs, tokenize=False, add_generation_prompt=False)128 129    # Pad both sequences to the same length so we can stack them130    chosen_enc = tokenizer(chosen_text, return_tensors="pt", add_special_tokens=False, padding=False)131    rejected_enc = tokenizer(rejected_text, return_tensors="pt", add_special_tokens=False, padding=False)132 133    pad_id = tokenizer.pad_token_id if tokenizer.pad_token_id is not None else tokenizer.eos_token_id134 135    chosen_ids = chosen_enc["input_ids"]136    rejected_ids = rejected_enc["input_ids"]137    L = max(chosen_ids.shape[1], rejected_ids.shape[1])138 139    def _pad(ids: torch.Tensor, length: int) -> torch.Tensor:140        cur = ids.shape[1]141        if cur >= length:142            return ids[:, :length]143        return torch.cat([ids, torch.full((1, length - cur), pad_id, dtype=ids.dtype)], dim=1)144 145    dpo_chosen_input_ids = _pad(chosen_ids, L).to(device)146    dpo_rejected_input_ids = _pad(rejected_ids, L).to(device)147 148    chosen_resp_mask = torch.zeros_like(dpo_chosen_input_ids)149    chosen_resp_mask[:, int(L * 0.6):chosen_ids.shape[1]] = 1150    rejected_resp_mask = torch.zeros_like(dpo_rejected_input_ids)151    rejected_resp_mask[:, int(L * 0.6):rejected_ids.shape[1]] = 1152 153    # Dummy reference-policy logprobs (in production: precomputed by data collator)154    dpo_chosen_ref_logprobs = torch.tensor([-30.0], device=device)155    dpo_rejected_ref_logprobs = torch.tensor([-35.0], device=device)156 157    return {158        "input_ids": input_ids,159        "response_mask": response_mask,160        "ctx_teacher_input_ids": ctx_teacher_input_ids,161        "sdpo_loss_mask": sdpo_loss_mask,162        "dpo_chosen_input_ids": dpo_chosen_input_ids,163        "dpo_chosen_response_mask": chosen_resp_mask,164        "dpo_rejected_input_ids": dpo_rejected_input_ids,165        "dpo_rejected_response_mask": rejected_resp_mask,166        "dpo_chosen_ref_logprobs": dpo_chosen_ref_logprobs,167        "dpo_rejected_ref_logprobs": dpo_rejected_ref_logprobs,168    }169 170 171__all__ = ["build_batch"]172