"""Packing of (instruction, text, tokens, reference codes, target codes) into MOSS-TTS tensors. Derived from $NB/vblora/code/vbpack.py (which itself replaced the private MOSS-TTS repo's `MossTTSLocalV15SFTDataset._pack_record`, that repo not being on this filesystem). Extended here with **reference-audio conditioning**: the processor accepts reference audio as a LongTensor [T, n_vq] of codec codes directly (see `_resolve_audio_items` in processing_moss_tts.py), so no audio tokenizer is needed at train time. Label derivation is pinned to the loss consumer (`va_loss.compute_supervised_loss_from_hidden`): * labels[t] = input_ids[t+1] (hidden at t decodes frame t+1) * supervised iff the TARGET text channel is audio_assistant_slot (continue) or audio_end (stop) * audio channels >= audio_pad_code -> -100 Crucially the reference audio rows sit inside the USER turn, whose target text channel is neither assistant_slot nor audio_end, so they are automatically unsupervised. Verified in verify_pack.py. """ import numpy as np import torch def codes_from_blob(blob, frames, n_vq=12): """target_codes/ref_codes/chosen_codes/... are int16 little-endian [frames, n_vq].""" if blob is None or not frames: return None a = np.frombuffer(blob, dtype=np.int16) n = int(frames) * int(n_vq) if a.size < n: return None return torch.from_numpy(a[:n].reshape(int(frames), int(n_vq)).astype(np.int64)) class Packer: def __init__(self, processor, model_config): self.proc = processor c = model_config self.n_vq = int(c.n_vq) self.assist = int(c.audio_assistant_slot_token_id) self.aend = int(c.audio_end_token_id) self.audio_pad = int(getattr(c, "audio_pad_token_id", 1024)) self.text_pad = int(c.pad_token_id) def pack(self, instruction, text, tokens, codes, ref_codes=None, language="English"): """codes: [T, n_vq] long. ref_codes: [R, n_vq] long or None. -> dict(input_ids [T,13], labels [T,13], n_sup)""" codes = torch.as_tensor(np.asarray(codes), dtype=torch.long) ref = None if ref_codes is not None: ref = [torch.as_tensor(np.asarray(ref_codes), dtype=torch.long)] um = self.proc.build_user_message(text=text, instruction=instruction, language=language, tokens=tokens, reference=ref) am = self.proc.build_assistant_message(audio_codes_list=[codes]) b = self.proc([[um, am]], mode="computing_loss") ids = b["input_ids"][0].long() # [T, 13] lab = torch.full_like(ids, -100) lab[:-1] = ids[1:] tt = lab[:, 0] sup = (tt == self.assist) | (tt == self.aend) # --- reference-audio contamination guard --------------------------------------- # The reference clip inside the USER turn is terminated by <|audio_end|>, which is the # SAME token id the model uses for the assistant's STOP decision. Without this guard the # position preceding the reference's terminator is supervised as "stop", teaching the # model to stop at the end of the user's reference audio. Measured: exactly 1 spurious # supervised position per referenced sample (n_sup 154 vs 153 at target_frames=152). # Supervise only from one position before the first assistant audio row. fa = (ids[:, 0] == self.assist).nonzero() if fa.numel() > 0: k = int(fa[0]) sup[: max(0, k - 1)] = False 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 collate(items, text_pad, audio_pad): """Right-pad a list of pack() outputs. Padding is fully masked in `labels`.""" T = max(it["input_ids"].shape[0] for it in items) C = items[0]["input_ids"].shape[1] B = len(items) ids = torch.empty((B, T, C), dtype=torch.long) ids[:, :, 0] = text_pad ids[:, :, 1:] = audio_pad lab = torch.full((B, T, C), -100, dtype=torch.long) am = torch.zeros((B, T), dtype=torch.long) for i, it in enumerate(items): t = it["input_ids"].shape[0] ids[i, :t] = it["input_ids"] lab[i, :t] = it["labels"] am[i, :t] = 1 out = {"input_ids": ids, "labels": lab, "attention_mask": am} for k in ("meta",): if k in items[0]: out[k] = [it[k] for it in items] return out