Text-to-Speech
English
German
voice-acting
qwen3
moss-audio-tokenizer-v2
audio-generation
File size: 4,533 Bytes
d911efa
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
"""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