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