"""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