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 + 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: .{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 # exact for every start index (441/128 ratio) TOPK = 50 SEM_TOKEN_WRAP = 151_675 % 65_536 # = 20603; see Corpus8k.__getitem__ def load_adapter_module(adapter_dir: str): sys.path.insert(0, adapter_dir) import minimax_music3_reference_adapter as ref return ref # ---------------------------------------------------------------- corpus ---- 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"))) # [T_total, 8] int32 vae = np.load(io.BytesIO(read(".vae.npy"))) # [F_dav, 128] fp16 frame-major probs_z = np.load(io.BytesIO(read(".probs.npz"))) p_idx, p_val = probs_z["idx"], probs_z["prob"] # [T_total, 8, 50] emitted = codes.shape[0] - warmup # window start s needs stitched-timeline latents [stitched_lat(s), # stitched_lat(s)+441) and frames [s, s+128) emitted. stitched_lat runs # slightly FASTER than uniform (345 per 100 frames), so walk max_start # down until its window fits (a handful of iterations at most). 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: # belt-and-braces: manifest filtering should prevent this; fall # back to a neighbour rather than killing the DataLoader worker. 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) # [128, 441] fp32 rows = slice(start + warmup, start + warmup + WINDOW) target_codes = codes[rows].astype(np.int64) # [128, 8] t_idx = p_idx[rows].astype(np.int64) # [128, 8, 50] t_val = p_val[rows].astype(np.float32) # Head 0's ids were stored as raw LM token ids (code + 151675) and the # uint16 dtype wrapped them mod 65536 -> code + 20603 (verified: own # code lands in the unwrapped top-50 for ~100% of frames). Unwrap, and # zero out anything outside the semantic vocab (e.g. EOS 151670 -> -5). 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), # [441, 128] pool, # [128, 441] 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]: # min_duration_s: a 128-frame window needs 5.12 s of emitted codes plus # stitched-latent margin; 8 s is conservative and drops only 24/8120 # early-EOS outliers (min in corpus: 3.7 s). 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 # ----------------------------------------------------------------- model ---- 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)) # [B*F, 8, D] 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) # [B, F, 16384] ac_logits = depth_teacher_forced(model.depth_decoder, hidden, codes) if log_taus is not None: # Released checkpoints carry uncalibrated (muP-sharp) readouts — # harmless for argmax/top-K, fatal for CE. Learnable per-head # temperature; folded into head weights at save time. 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()} # ------------------------------------------------------------------ main ---- 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()} # Fold the learned temperatures into the readout heads so the saved # file is a plain calibrated encoder (argmax/top-K ordering unchanged, # CE now meaningful). heads.0 = semantic; depth_decoder.heads.i = c(i+1). 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 \\rvq_encoder.safetensors " "--config-file \\rvq_encoder_config.json, then the 1-step lm_sft_train2 [Base] CE.") if __name__ == "__main__": main()