| r"""Further-train a SimpleTuner-class RVQ encoder on the published distillation corpora. |
| |
| Warm-starts from a released checkpoint (default: Mothersuperior pooled-v4, 169M) |
| and trains with SOFT DISTILLATION against the teacher top-50 distributions the |
| corpus ships — the axis the published runs (hard-label CE, 12 epochs) never used. |
| Checkpoints regularly; final selection happens OFFLINE on our real-audio gate |
| (export_codes_v4.py --encoder-file <ckpt> + 1-step lm_sft_train2 [Base] CE), |
| which nobody else measures. |
| |
| Corpus support: |
| * Mothersuperior/minimax-music3-rvq-distill-corpus-8k (PRIMARY — ships |
| precomputed fp16 DAV latents frame-major [F,128] @ 86.1328 Hz, plus |
| top-50 PROBS [T,8,50]). Records: <id>.{flac,codes.npy,probs.npz,vae.npy,json} |
| inside data/shard-NNN.zip; manifest.jsonl has align_ok/probs_ok flags. |
| Warm-up convention: codes/probs row 0 is un-emitted; frame i <-> row i+1. |
| * bghira/minimax-music3-rvq-reverse-distillation: NOT yet wired — ships |
| top-50 LOGITS but NO DAV latents, so it needs a ChunkedDAV precache pass |
| first (see docs/plans/2026-08-18-encoder-training-plan.md). Respect its |
| deterministic dataset_split when added. |
| |
| Key geometry: every 128-frame window spans EXACTLY 441 DAV latents |
| (bounds[i] = floor(i*441/128); the +128 difference is always 441), so batches |
| need no padding — which matters because GroupNorm(1, C) in the conv stack |
| normalises over length and padding would shift its statistics. |
| |
| The depth decoder is TEACHER-FORCED here (ground-truth c0..c6 as priors, causal |
| mask, head i read from position i+1) — the published forward() feeds its own |
| argmax chain and is inference-only. |
| |
| Usage (single GPU, 5090): |
| python rvq_distill_train.py ^ |
| --corpus M:\HOT-Step-CPP\_corpora\mm3-rvq-distill-8k ^ |
| --encoder-dir M:\HOT-Step-CPP\_experiments\open-rvq-pooled-v4 ^ |
| --out M:\HOT-Step-CPP\_experiments\rvq-train\pv4-softdistill-r1 ^ |
| --steps 20000 --batch 16 --grad-accum 4 --lr 1e-4 --vram-frac 0.85 |
| """ |
|
|
| from __future__ import annotations |
|
|
| import argparse |
| import hashlib |
| import io |
| import json |
| import math |
| import os |
| import shutil |
| import sys |
| import time |
| import zipfile |
| from pathlib import Path |
|
|
| import numpy as np |
| import torch |
| import torch.nn.functional as F |
| from torch.utils.data import DataLoader, Dataset |
|
|
| ADAPTER_DIR_DEFAULT = r"M:\HOT-Step-CPP\_experiments\open-rvq-v4" |
| WINDOW = 128 |
| LATENTS_PER_WINDOW = 441 |
| TOPK = 50 |
| SEM_TOKEN_WRAP = 151_675 % 65_536 |
|
|
|
|
| def load_adapter_module(adapter_dir: str): |
| sys.path.insert(0, adapter_dir) |
| import minimax_music3_reference_adapter as ref |
|
|
| return ref |
|
|
|
|
| |
|
|
| def holdout_of(record_id: str, frac: float, salt: str = "hotstep-holdout-v1") -> bool: |
| digest = hashlib.sha1(f"{salt}:{record_id}".encode()).digest() |
| return (int.from_bytes(digest[:8], "big") / 2**64) < frac |
|
|
|
|
| def stitched_lat(i: int) -> int: |
| """Frame index -> latent index in the corpus's RENDERED audio. |
| |
| The 8k corpus audio is stitched from 200-frame DiT windows at a 100-frame / |
| 345-LATENT hop (README: 'constant stitched timeline'; non-first-chunk |
| ownership from local frame 25). 100 frames nominally map to 344.53 latents, |
| so the uniform i*441//128 mapping drifts ~1 latent per 800 and top-1 decays |
| to ~0 by frame 2400 — measured, then eliminated by this rule. The +1 is a |
| constant global offset, verified best at every position on every record |
| probed (t1 0.4-0.6 flat vs position). Real audio has NO such drift — the |
| uniform rule stays correct at inference/export time. |
| """ |
| k = 0 if i < 125 else (i - 25) // 100 |
| return 345 * k + ((i - 100 * k) * 441) // 128 + 1 |
|
|
|
|
| def pool_matrix(bounds: list[int]) -> torch.Tensor: |
| """Local copy of the adapter's build_pool_matrix so DataLoader workers can |
| unpickle this Dataset without the runtime-sys.path adapter import.""" |
| origin = bounds[0] |
| local = [b - origin for b in bounds] |
| pool = torch.zeros((len(local) - 1, local[-1]), dtype=torch.float32) |
| for i, (a, b) in enumerate(zip(local[:-1], local[1:])): |
| pool[i, a:b] = 1.0 / (b - a) |
| return pool |
|
|
|
|
| class Corpus8k(Dataset): |
| """Random 128-frame windows from the Mothersuperior 8k distill corpus.""" |
|
|
| def __init__(self, root: str, records: list[dict], seed: int, fixed_windows: bool = False): |
| self.root = Path(root) |
| self.records = records |
| self.seed = seed |
| self.fixed_windows = fixed_windows |
| self._zips: dict[str, zipfile.ZipFile] = {} |
|
|
| def __len__(self): |
| return len(self.records) |
|
|
| def _zip(self, shard: str) -> zipfile.ZipFile: |
| z = self._zips.get(shard) |
| if z is None: |
| z = zipfile.ZipFile(self.root / "data" / shard) |
| self._zips[shard] = z |
| return z |
|
|
| def __getitem__(self, index: int): |
| rec = self.records[index] |
| z = self._zip(rec["shard"]) |
| rid = rec["id"] |
|
|
| def read(ext): |
| return z.read(f"{rid}/{rid}{ext}") |
|
|
| meta = json.loads(read(".json")) |
| warmup = int(meta.get("codes_warmup_frames", 1)) |
| codes = np.load(io.BytesIO(read(".codes.npy"))) |
| vae = np.load(io.BytesIO(read(".vae.npy"))) |
| probs_z = np.load(io.BytesIO(read(".probs.npz"))) |
| p_idx, p_val = probs_z["idx"], probs_z["prob"] |
|
|
| emitted = codes.shape[0] - warmup |
| |
| |
| |
| |
| max_start = emitted - WINDOW |
| while max_start >= 0 and stitched_lat(max_start) + LATENTS_PER_WINDOW > vae.shape[0]: |
| max_start -= 1 |
| if max_start < 0: |
| |
| |
| return self[(index + 1) % len(self.records)] |
| if self.fixed_windows: |
| start = (max_start // 2) if index % 2 else 0 |
| else: |
| g = np.random.default_rng( |
| (self.seed * 1_000_003 + index) ^ int.from_bytes(os.urandom(4), "big") |
| ) |
| start = int(g.integers(0, max_start + 1)) |
|
|
| lat_start = stitched_lat(start) |
| latents = vae[lat_start : lat_start + LATENTS_PER_WINDOW].astype(np.float32) |
|
|
| bounds = [(start + i) * 441 // 128 for i in range(WINDOW + 1)] |
| pool = pool_matrix(bounds) |
|
|
| rows = slice(start + warmup, start + warmup + WINDOW) |
| target_codes = codes[rows].astype(np.int64) |
| t_idx = p_idx[rows].astype(np.int64) |
| t_val = p_val[rows].astype(np.float32) |
| |
| |
| |
| |
| sem = t_idx[:, 0, :] - SEM_TOKEN_WRAP |
| bad = (sem < 0) | (sem >= 16384) |
| sem[bad] = 0 |
| t_val[:, 0, :][bad] = 0.0 |
| t_idx[:, 0, :] = sem |
| t_val = t_val / np.clip(t_val.sum(-1, keepdims=True), 1e-8, None) |
|
|
| return ( |
| torch.from_numpy(latents), |
| pool, |
| torch.from_numpy(target_codes), |
| torch.from_numpy(t_idx), |
| torch.from_numpy(t_val), |
| ) |
|
|
|
|
| def load_manifest(root: str, min_duration_s: float = 8.0) -> list[dict]: |
| |
| |
| |
| records = [] |
| with open(Path(root) / "manifest.jsonl", encoding="utf-8") as fh: |
| for line in fh: |
| rec = json.loads(line) |
| if rec.get("align_ok") and rec.get("probs_ok") and rec.get("duration_s", 0) >= min_duration_s: |
| records.append(rec) |
| return records |
|
|
|
|
| |
|
|
| def encoder_trunk(model, latents: torch.Tensor, pool: torch.Tensor) -> torch.Tensor: |
| """conv stack + pooled transformer -> per-frame hidden [B, 128, d_model].""" |
| h = model.conv_in(latents.transpose(1, 2)) |
| for block in model.blocks: |
| h = block(h) |
| h = torch.bmm(pool.to(h.dtype), h.transpose(1, 2)) |
| h = h + model.position[:, : pool.shape[1]].to(h.dtype) |
| layers = model.transformer if isinstance(model.transformer, torch.nn.ModuleList) else model.transformer.layers |
| for layer in layers: |
| h = layer(h) |
| return model.norm_out(h) |
|
|
|
|
| def depth_teacher_forced(dd, frame_context: torch.Tensor, codes: torch.Tensor) -> list[torch.Tensor]: |
| """Parallel teacher-forced depth pass. codes [B, F, 8] ground truth. |
| Sequence = [ctx, e0(c0), .., e6(c6)]; causal mask; head i reads position i+1.""" |
| batch, frames, _ = frame_context.shape |
| parts = [dd.context_projection(frame_context).flatten(0, 1).unsqueeze(1)] |
| for i in range(7): |
| parts.append(dd.prior_embeddings[i](codes[..., i]).flatten(0, 1).unsqueeze(1)) |
| hidden = dd._decode(torch.cat(parts, dim=1)) |
| return [dd.heads[i](hidden[:, i + 1]).view(batch, frames, -1) for i in range(7)] |
|
|
|
|
| def soft_ce(logits: torch.Tensor, t_idx: torch.Tensor, t_val: torch.Tensor) -> torch.Tensor: |
| """-sum p_teacher * log q_student over the teacher's top-50 support.""" |
| logq = F.log_softmax(logits.float(), dim=-1) |
| return -(t_val * logq.gather(-1, t_idx)).sum(-1).mean() |
|
|
|
|
| def compute_losses(model, batch, device, acoustic_weight: float, hard_weight: float, log_taus=None): |
| latents, pool, codes, t_idx, t_val = (x.to(device, non_blocking=True) for x in batch) |
| hidden = encoder_trunk(model, latents, pool) |
| sem_logits = model.heads[0](hidden) |
| ac_logits = depth_teacher_forced(model.depth_decoder, hidden, codes) |
| if log_taus is not None: |
| |
| |
| |
| taus = log_taus.exp() |
| sem_logits = sem_logits / taus[0] |
| ac_logits = [ac_logits[i] / taus[i + 1] for i in range(7)] |
|
|
| losses = {"sem_soft": soft_ce(sem_logits, t_idx[..., 0, :], t_val[..., 0, :])} |
| ac = [soft_ce(ac_logits[i], t_idx[..., i + 1, :], t_val[..., i + 1, :]) for i in range(7)] |
| losses["ac_soft"] = torch.stack(ac).mean() |
|
|
| if hard_weight > 0: |
| losses["sem_hard"] = F.cross_entropy(sem_logits.flatten(0, 1).float(), codes[..., 0].flatten()) |
| ach = [ |
| F.cross_entropy(ac_logits[i].flatten(0, 1).float(), codes[..., i + 1].flatten()) |
| for i in range(7) |
| ] |
| losses["ac_hard"] = torch.stack(ach).mean() |
|
|
| total = losses["sem_soft"] + acoustic_weight * losses["ac_soft"] |
| if hard_weight > 0: |
| total = total + hard_weight * (losses["sem_hard"] + acoustic_weight * losses["ac_hard"]) |
| losses["total"] = total |
|
|
| with torch.no_grad(): |
| losses["sem_t1"] = (sem_logits.argmax(-1) == codes[..., 0]).float().mean() |
| losses["ac_t1"] = torch.stack( |
| [(ac_logits[i].argmax(-1) == codes[..., i + 1]).float().mean() for i in range(7)] |
| ).mean() |
| return losses |
|
|
|
|
| @torch.no_grad() |
| def evaluate(model, loader, device, acoustic_weight, autocast_ctx, log_taus=None): |
| model.eval() |
| sums: dict[str, float] = {} |
| n = 0 |
| for batch in loader: |
| with autocast_ctx(): |
| losses = compute_losses(model, batch, device, acoustic_weight, hard_weight=1.0, log_taus=log_taus) |
| for k, v in losses.items(): |
| sums[k] = sums.get(k, 0.0) + float(v) |
| n += 1 |
| model.train() |
| return {k: v / max(n, 1) for k, v in sums.items()} |
|
|
|
|
| |
|
|
| def main(): |
| ap = argparse.ArgumentParser() |
| ap.add_argument("--corpus", default=r"M:\HOT-Step-CPP\_corpora\mm3-rvq-distill-8k") |
| ap.add_argument("--encoder-dir", default=r"M:\HOT-Step-CPP\_experiments\open-rvq-pooled-v4") |
| ap.add_argument("--encoder-file", default="rvq_encoder.safetensors") |
| ap.add_argument("--config-file", default="rvq_encoder_config.json") |
| ap.add_argument("--adapter-dir", default=ADAPTER_DIR_DEFAULT) |
| ap.add_argument("--out", required=True) |
| ap.add_argument("--steps", type=int, default=20000) |
| ap.add_argument("--batch", type=int, default=16) |
| ap.add_argument("--grad-accum", type=int, default=4) |
| ap.add_argument("--lr", type=float, default=1e-4) |
| ap.add_argument("--warmup", type=int, default=200) |
| ap.add_argument("--weight-decay", type=float, default=0.01) |
| ap.add_argument("--acoustic-weight", type=float, default=1.0) |
| ap.add_argument("--hard-weight", type=float, default=0.0, |
| help="mix in hard-label CE alongside soft distillation") |
| ap.add_argument("--sem-tau", type=float, default=16.0, |
| help="initial semantic logit temperature (line-searched on pooled-v4)") |
| ap.add_argument("--ac-tau", type=float, default=2.0, |
| help="initial acoustic logit temperature") |
| ap.add_argument("--holdout-frac", type=float, default=0.05) |
| ap.add_argument("--eval-every", type=int, default=500) |
| ap.add_argument("--save-every", type=int, default=1000) |
| ap.add_argument("--workers", type=int, default=2) |
| ap.add_argument("--seed", type=int, default=1) |
| ap.add_argument("--vram-frac", type=float, default=0.85) |
| ap.add_argument("--device", default="cuda" if torch.cuda.is_available() else "cpu") |
| ap.add_argument("--smoke", action="store_true", help="2 optimizer steps on CPU-sized batch, then exit") |
| args = ap.parse_args() |
|
|
| torch.manual_seed(args.seed) |
| device = torch.device(args.device) |
| if device.type == "cuda": |
| torch.cuda.set_per_process_memory_fraction(args.vram_frac) |
|
|
| ref = load_adapter_module(args.adapter_dir) |
| from safetensors.torch import load_file, save_file |
|
|
| enc_dir = Path(args.encoder_dir) |
| config = ref.RVQEncoderConfig.from_dict( |
| json.loads((enc_dir / args.config_file).read_text(encoding="utf-8")) |
| ) |
| model = ref.MiniMaxMusicRVQEncoder(config) |
| model.load_state_dict(load_file(str(enc_dir / args.encoder_file)), strict=True) |
| model.to(device).train() |
| n_params = sum(p.numel() for p in model.parameters()) |
| print(f"[model] {n_params/1e6:.1f}M params, warm start from {enc_dir / args.encoder_file}") |
|
|
| records = load_manifest(args.corpus) |
| train_recs = [r for r in records if not holdout_of(r["id"], args.holdout_frac)] |
| hold_recs = [r for r in records if holdout_of(r["id"], args.holdout_frac)] |
| print(f"[data] {len(records)} usable records -> {len(train_recs)} train / {len(hold_recs)} holdout") |
|
|
| train_ds = Corpus8k(args.corpus, train_recs, args.seed) |
| hold_ds = Corpus8k(args.corpus, hold_recs, args.seed, fixed_windows=True) |
| loader_kw = dict( |
| batch_size=args.batch, |
| num_workers=args.workers, |
| pin_memory=device.type == "cuda", |
| persistent_workers=args.workers > 0, |
| ) |
| train_loader = DataLoader(train_ds, shuffle=True, drop_last=True, **loader_kw) |
| hold_loader = DataLoader(hold_ds, shuffle=False, **loader_kw) |
|
|
| log_taus = torch.nn.Parameter( |
| torch.log(torch.tensor([args.sem_tau] + [args.ac_tau] * 7, dtype=torch.float32, device=device)) |
| ) |
| decay, no_decay = [], [] |
| for name, p in model.named_parameters(): |
| (no_decay if p.ndim <= 1 or "position" in name else decay).append(p) |
| opt = torch.optim.AdamW( |
| [ |
| {"params": decay, "weight_decay": args.weight_decay}, |
| {"params": no_decay, "weight_decay": 0.0}, |
| {"params": [log_taus], "weight_decay": 0.0}, |
| ], |
| lr=args.lr, betas=(0.9, 0.95), |
| ) |
|
|
| def lr_at(step): |
| if step < args.warmup: |
| return args.lr * (step + 1) / args.warmup |
| t = (step - args.warmup) / max(args.steps - args.warmup, 1) |
| return args.lr * 0.5 * (1 + math.cos(math.pi * min(t, 1.0))) |
|
|
| use_bf16 = device.type == "cuda" |
|
|
| def autocast_ctx(): |
| return torch.autocast("cuda", dtype=torch.bfloat16) if use_bf16 else torch.autocast("cpu", enabled=False) |
|
|
| out = Path(args.out) |
| out.mkdir(parents=True, exist_ok=True) |
| shutil.copyfile(enc_dir / args.config_file, out / "rvq_encoder_config.json") |
| (out / "run_args.json").write_text(json.dumps(vars(args), indent=1), encoding="utf-8") |
| log_path = out / "train_log.jsonl" |
|
|
| def save_ckpt(tag: str): |
| ck = out / tag |
| ck.mkdir(exist_ok=True) |
| state = {k: v.detach().to(torch.float32).cpu() for k, v in model.state_dict().items()} |
| |
| |
| |
| taus = log_taus.detach().exp().cpu() |
| for key in ("heads.0.weight", "heads.0.bias"): |
| state[key] = state[key] / taus[0] |
| for i in range(7): |
| for suffix in ("weight", "bias"): |
| key = f"depth_decoder.heads.{i}.{suffix}" |
| state[key] = state[key] / taus[i + 1] |
| save_file(state, str(ck / "rvq_encoder.safetensors")) |
| shutil.copyfile(out / "rvq_encoder_config.json", ck / "rvq_encoder_config.json") |
| (ck / "calibration.json").write_text( |
| json.dumps({"folded_taus": [round(float(t), 4) for t in taus]}), encoding="utf-8" |
| ) |
| print(f"[ckpt] {ck} (taus folded: {[round(float(t), 2) for t in taus]})") |
|
|
| best_hold = float("inf") |
| step = 0 |
| t0 = time.time() |
| opt.zero_grad(set_to_none=True) |
| data_iter = iter(train_loader) |
| while step < args.steps: |
| for micro in range(args.grad_accum): |
| try: |
| batch = next(data_iter) |
| except StopIteration: |
| data_iter = iter(train_loader) |
| batch = next(data_iter) |
| with autocast_ctx(): |
| losses = compute_losses( |
| model, batch, device, args.acoustic_weight, args.hard_weight, log_taus=log_taus |
| ) |
| (losses["total"] / args.grad_accum).backward() |
| if not torch.isfinite(losses["total"]): |
| raise FloatingPointError(f"non-finite loss at step {step}: {losses}") |
| torch.nn.utils.clip_grad_norm_(list(model.parameters()) + [log_taus], 1.0) |
| for group in opt.param_groups: |
| group["lr"] = lr_at(step) |
| opt.step() |
| opt.zero_grad(set_to_none=True) |
| step += 1 |
|
|
| if step % 25 == 0 or args.smoke: |
| row = {k: round(float(v.detach()), 4) for k, v in losses.items()} |
| row.update( |
| step=step, lr=round(lr_at(step), 8), elapsed_s=round(time.time() - t0, 1), |
| taus=[round(float(t), 2) for t in log_taus.detach().exp()], |
| ) |
| print(f"[{step}/{args.steps}] " + " ".join(f"{k}={v}" for k, v in row.items() if k != "step")) |
| with open(log_path, "a", encoding="utf-8") as fh: |
| fh.write(json.dumps(row) + "\n") |
|
|
| if args.smoke and step >= 2: |
| print("[smoke] OK — forward/backward/step ran clean") |
| return |
|
|
| if step % args.eval_every == 0: |
| ev = evaluate(model, hold_loader, device, args.acoustic_weight, autocast_ctx, log_taus=log_taus) |
| row = {("hold_" + k): round(float(v), 4) for k, v in ev.items()} |
| row["step"] = step |
| print(f"[eval @{step}] " + " ".join(f"{k}={v}" for k, v in row.items() if k != "step")) |
| with open(log_path, "a", encoding="utf-8") as fh: |
| fh.write(json.dumps(row) + "\n") |
| score = ev["sem_hard"] + args.acoustic_weight * ev["ac_hard"] |
| if score < best_hold: |
| best_hold = score |
| save_ckpt("best_holdout") |
|
|
| if step % args.save_every == 0: |
| save_ckpt(f"step{step:06d}") |
|
|
| save_ckpt("final") |
| print(f"[done] {args.steps} steps in {(time.time() - t0)/3600:.2f} h; best holdout hard-CE sum {best_hold:.4f}") |
| print("[next] run each kept checkpoint through the REAL-AUDIO gate: " |
| "export_codes_v4.py --encoder-file <ckpt>\\rvq_encoder.safetensors " |
| "--config-file <ckpt>\\rvq_encoder_config.json, then the 1-step lm_sft_train2 [Base] CE.") |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|