"""Packen für M1/M2 mit dem SFT3-Prozessor (identisch zu sft_train2.build_examples, ohne Prompt-Zufall).""" import sys, torch sys.path.insert(0, '/e/data1/datasets/playground/mmlaion/schuhmann1/dramabox/train') from moss_pack import collate # Right-Padding, Labels maskiert from prompt_fmt import user_content, AUDIO_PLACEHOLDER class MossPacker: def __init__(self, proc, cfg): self.proc = proc self.assist = int(cfg.audio_assistant_slot_token_id); self.aend = int(cfg.audio_end_token_id) self.audio_pad = int(cfg.audio_pad_token_id); self.text_pad = int(cfg.pad_token_id) def pack(self, meta, codes, ref_codes=None): use_ref = ref_codes is not None content = user_content(meta, use_ref) um = {"role": "user", "content": content, "audio_codes_list": [torch.as_tensor(ref_codes)] if use_ref else []} am = self.proc.build_assistant_message(audio_codes_list=[torch.as_tensor(codes)]) b = self.proc([[um, am]], mode="computing_loss") ids = b["input_ids"][0].long() lab = torch.full_like(ids, -100) lab[:-1] = ids[1:] tt = lab[:, 0] sup = (tt == self.assist) | (tt == self.aend) fa = (ids[:, 0] == self.assist).nonzero() if fa.numel() > 0: sup[: max(0, int(fa[0]) - 1)] = False # Referenz-<|audio_end|> nicht als Stop lernen lab[~sup] = -100 au = lab[:, 1:] au[(au >= self.audio_pad) | (au < 0)] = -100 lab[:, 1:] = au return {"input_ids": ids, "labels": lab, "n_sup": int(sup.sum())} def prompt_ids(self, meta, ref_codes=None, frames=None): """Generierungs-Eingabe [T,13] (endet mit <|audio_start|>).""" use_ref = ref_codes is not None content = user_content(meta, use_ref, frames=frames) um = {"role": "user", "content": content, "audio_codes_list": [torch.as_tensor(ref_codes)] if use_ref else []} b = self.proc([[um]], mode="generation") return b["input_ids"][0].long()