Download code/va_loss.py from laion/Humaneness-Voice-Small: direct link, hf CLI and curl.
- Browser
- Download file 5.53 kB
-
https://huggingface.co/laion/Humaneness-Voice-Small/resolve/main/code/va_loss.py
- Command line
-
hf download hf://laion/Humaneness-Voice-Small/code/va_loss.py
-
curl -L -o va_loss.py https://huggingface.co/laion/Humaneness-Voice-Small/resolve/main/code/va_loss.py
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 | |