Text-to-Speech
English
German
voice-acting
qwen3
moss-audio-tokenizer-v2
audio-generation
ChristophSchuhmann's picture
Document architecture, prompts, code, and full run statistics
d911efa verified
Raw History Blame Contribute Delete
5.53 kB
"""Channel-wise supervised loss for MOSS-TTS-Local-Transformer.
ADOPTED VERBATIM from the reference implementation
`ref/code/lora-self-distillation/training/va_train.py` (which in turn copies
`moss_tts_local_v1.5/finetuning/sft.py`). The previous agent had derived the same structure
independently; per the brief, where the two differ this one wins. Only the module docstring and
the import of `torch.utils.checkpoint` are new.
Structure (matches PROTOCOL 0.a): the global Qwen3 emits h_t; frame t+1 is decoded by the 1-layer
local GPT-2 over a 12-step sequence whose inputs are [proj(h_t), emb_0(c_0), ..., emb_10(c_10)],
with head k applied to local output k. Fully teacher-forceable in a single causal pass.
"""
import torch
import torch.nn.functional as F
import torch.utils.checkpoint
def unwrap_training_model(model):
unwrapped = model
while hasattr(unwrapped, "module"):
unwrapped = unwrapped.module
return unwrapped
def _local_chunk_sums(base_model, h, lab, n_vq, use_binary, local_dtype):
"""Per-channel SUM of cross-entropy over a chunk of positions (text at index 0)."""
M = h.shape[0]
local_prefix = base_model._global_hidden_to_local(h).to(dtype=local_dtype)
local_inputs = torch.zeros((M, n_vq, int(local_prefix.shape[-1])), dtype=local_dtype, device=h.device)
local_inputs[:, 0, :] = local_prefix
audio_t = lab[:, 1:]
for ci in range(n_vq - 1):
tid = audio_t[:, ci]; emb = base_model.audio_embeddings[ci]
vm = (tid >= 0) & (tid < emb.num_embeddings)
local_inputs[:, ci + 1, :] = emb(tid.masked_fill(~vm, 0)).to(dtype=local_dtype) * vm.unsqueeze(-1)
lh = base_model.local_transformer(
input_ids=None, attention_mask=None, position_ids=None, inputs_embeds=local_inputs,
use_cache=False, output_attentions=False, output_hidden_states=False, return_dict=True,
cu_seqlens=None, num_sequences=None).last_hidden_state
sums = []
tt = lab[:, 0]
if use_binary:
bt = torch.full_like(tt, -100)
bt = bt.masked_fill(tt.eq(int(base_model.config.audio_assistant_slot_token_id)), 0)
bt = bt.masked_fill(tt.eq(int(base_model.config.audio_end_token_id)), 1)
sums.append(F.cross_entropy(base_model.local_text_lm_head(lh[:, 0, :]).float(), bt,
ignore_index=-100, reduction="sum"))
else:
sums.append(F.cross_entropy(base_model.text_lm_head(lh[:, 0, :]).float(), tt,
ignore_index=-100, reduction="sum"))
for ci in range(n_vq):
sums.append(F.cross_entropy(base_model.audio_lm_heads[ci](lh[:, ci, :]).float(), audio_t[:, ci],
ignore_index=-100, reduction="sum"))
return torch.stack(sums) # (n_vq+1,)
def compute_supervised_loss_from_hidden(base_model, *, global_hidden_states, labels,
channelwise_loss_weight, chunk_size=1024,
use_checkpoint=True, return_per_channel=False):
"""Chunked over positions so peak memory is independent of sequence length."""
batch_size, seq_len, hidden_size = global_hidden_states.shape
n_vq = int(base_model.config.n_vq)
if labels.shape[-1] != n_vq + 1:
raise ValueError(f"Expected labels with {n_vq + 1} channels, got {labels.shape[-1]}.")
weights = channelwise_loss_weight or [1.0] * (n_vq + 1)
if len(weights) != n_vq + 1:
raise ValueError(f"`channelwise_loss_weight` length {len(weights)} != {n_vq + 1}.")
dev = global_hidden_states.device
flat_hidden = global_hidden_states.reshape(batch_size * seq_len, hidden_size)
flat_labels = labels.reshape(batch_size * seq_len, n_vq + 1)
local_dtype = base_model.local_transformer.ln_f.weight.dtype
use_binary = (hasattr(base_model, "_use_binary_local_text_head")
and base_model._use_binary_local_text_head()
and getattr(base_model, "local_text_lm_head", None) is not None)
N = batch_size * seq_len
sums_total = torch.zeros(n_vq + 1, device=dev, dtype=torch.float32)
cnts = torch.zeros(n_vq + 1, device=dev)
for s in range(0, N, chunk_size):
h = flat_hidden[s:s + chunk_size]; lab = flat_labels[s:s + chunk_size]
with torch.no_grad():
tt = lab[:, 0]; at = lab[:, 1:]
cnts[0] += ((tt.eq(int(base_model.config.audio_assistant_slot_token_id)) |
tt.eq(int(base_model.config.audio_end_token_id))).sum() if use_binary
else (tt != -100).sum())
for ci in range(n_vq):
cnts[ci + 1] += (at[:, ci] != -100).sum()
cs = (torch.utils.checkpoint.checkpoint(_local_chunk_sums, base_model, h, lab, n_vq,
use_binary, local_dtype, use_reentrant=False)
if use_checkpoint else _local_chunk_sums(base_model, h, lab, n_vq, use_binary, local_dtype))
sums_total = sums_total + cs
total_loss = torch.zeros((), device=dev, dtype=torch.float32); total_weight = 0.0
for k in range(n_vq + 1):
if cnts[k] > 0:
total_loss = total_loss + float(weights[k]) * (sums_total[k] / cnts[k])
total_weight += float(weights[k])
if total_weight <= 0:
raise RuntimeError("All labels are ignored; check dataset packing.")
loss = total_loss / total_weight
if return_per_channel:
per = (sums_total / cnts.clamp(min=1)).detach()
return loss, per
return loss