Text-to-Speech
English
German
voice-acting
qwen3
moss-audio-tokenizer-v2
audio-generation
Humaneness-Voice-Small / code /moss_pack.py
ChristophSchuhmann's picture
Document architecture, prompts, code, and full run statistics
d911efa verified
Raw History Blame Contribute Delete
4.53 kB
"""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