scragnog's picture
hotstep-v1: soft-distilled + calibrated pooled-v4
45ed331 verified
Raw
History Blame Contribute Delete
21.2 kB
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 # 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 <ckpt>\\rvq_encoder.safetensors "
"--config-file <ckpt>\\rvq_encoder_config.json, then the 1-step lm_sft_train2 [Base] CE.")
if __name__ == "__main__":
main()