# Distillation vitesse-moyenne (IntMeanFlow-style, sans JVP) du flow CosyVoice3. # Teacher gele (CFG 0.7 bake dans la cible), student init=teacher, intervalle (t,r) cosine-warpe. # Usage: python train.py --model_dir models/cosyvoice3-0.5b --shards "shards/*.pt" --out runs/mv1 import argparse, glob, math, os, random, time import torch import torch.nn.functional as F from hyperpyyaml import load_hyperpyyaml import sys sys.path.insert(0, os.path.dirname(os.path.abspath(__file__))) from student import MeanFlowDiT, CV3_DIT_KWARGS, student_infer SIGMA = 1e-6 CFG = 0.7 def cosine_warp(s): return 1 - torch.cos(s * 0.5 * math.pi) def load_teacher(model_dir, device): with open(f"{model_dir}/cosyvoice3.yaml") as f: # llm/hift non requis pour la distillation du flow (evite transformers & co) cfg = load_hyperpyyaml(f, overrides={"llm": None, "hift": None}) flow = cfg["flow"] sd = torch.load(f"{model_dir}/flow.pt", map_location="cpu") flow.load_state_dict(sd, strict=True) flow.to(device).eval() for p in flow.parameters(): p.requires_grad_(False) return flow class ShardDataset: def __init__(self, pattern, max_mel=1200, prompt_s=(3.0, 6.0), holdout=8): self.files = sorted(glob.glob(pattern)) assert self.files, f"aucun shard: {pattern}" self.max_mel = max_mel self.prompt_s = prompt_s first = torch.load(self.files[0], map_location="cpu") self.val_items = first[:holdout] self._cache_idx, self._cache = -1, None def sample(self, rng): si = rng.randrange(len(self.files)) if si != self._cache_idx: self._cache, self._cache_idx = torch.load(self.files[si], map_location="cpu"), si it = self._cache[rng.randrange(len(self._cache))] return it def make_batch(self, bs, rng, device): items = [self.sample(rng) for _ in range(bs)] return collate(items, self.max_mel, self.prompt_s, rng, device) def collate(items, max_mel, prompt_s, rng, device): B = len(items) toks, mels, xvs, p_mel_lens, mel_lens = [], [], [], [], [] for it in items: mel = it["mel"].float() tok = it["tokens"].long() tl = min(mel.shape[0] // 2, tok.shape[0], max_mel // 2) mel, tok = mel[: 2 * tl], tok[:tl] p_tok = int(rng.uniform(*prompt_s) * 25) p_tok = min(p_tok, tl - 25) # garder >=1s de cible toks.append(tok); mels.append(mel); xvs.append(it["xvec"].float()) p_mel_lens.append(2 * p_tok); mel_lens.append(2 * tl) T = max(mel_lens) tok_pad = torch.zeros(B, T // 2, dtype=torch.long) mel_pad = torch.zeros(B, T, 80) for i in range(B): tok_pad[i, : len(toks[i])] = toks[i] mel_pad[i, : mel_lens[i]] = mels[i] return (tok_pad.to(device), mel_pad.to(device), torch.stack(xvs).to(device), torch.tensor(p_mel_lens, device=device), torch.tensor(mel_lens, device=device)) @torch.no_grad() def build_conditions(flow, tok, mel, xv, p_mel_len, mel_len): """Replique flow.inference (finalize) en batch: mu, conds, mask, spks.""" B, T = mel.shape[0], mel.shape[1] emb = F.normalize(xv, dim=1) spks = flow.spk_embed_affine_layer(emb) tok_mask = (torch.arange(T // 2, device=tok.device)[None, :] < (mel_len[:, None] // 2)).unsqueeze(-1).float() h = flow.input_embedding(torch.clamp(tok, min=0)) * tok_mask h = flow.pre_lookahead_layer(h) h = h.repeat_interleave(flow.token_mel_ratio, dim=1) # [B,T,80] conds = torch.zeros_like(mel) for i in range(B): conds[i, : p_mel_len[i]] = mel[i, : p_mel_len[i]] mask = (torch.arange(T, device=mel.device)[None, :] < mel_len[:, None]).float() return h.transpose(1, 2), conds.transpose(1, 2), mask.unsqueeze(1), spks # mu/conds [B,80,T] @torch.no_grad() def teacher_cfg_v(est, x, mask, mu, t, spks, cond): B = x.size(0) x2 = torch.cat([x, x], 0); m2 = torch.cat([mask, mask], 0) mu2 = torch.cat([mu, torch.zeros_like(mu)], 0) t2 = torch.cat([t, t], 0) s2 = torch.cat([spks, torch.zeros_like(spks)], 0) c2 = torch.cat([cond, torch.zeros_like(cond)], 0) v = est(x2, m2, mu2, t2, s2, c2) vc, vu = v[:B], v[B:] return (1.0 + CFG) * vc - CFG * vu def main(): ap = argparse.ArgumentParser() ap.add_argument("--model_dir", required=True) ap.add_argument("--shards", required=True) ap.add_argument("--out", required=True) ap.add_argument("--bs", type=int, default=8) ap.add_argument("--steps", type=int, default=100000) ap.add_argument("--lr", type=float, default=2e-5) ap.add_argument("--warmup", type=int, default=500) ap.add_argument("--substeps", type=int, default=4) # pas Euler teacher par intervalle ap.add_argument("--grid_prob", type=float, default=0.5) # prob d'ancrer (t,r) sur la grille 2-NFE ap.add_argument("--ema", type=float, default=0.999) ap.add_argument("--save_every", type=int, default=2000) ap.add_argument("--val_every", type=int, default=2000) ap.add_argument("--seed", type=int, default=42) args = ap.parse_args() os.makedirs(args.out, exist_ok=True) device = "cuda" rng = random.Random(args.seed) torch.manual_seed(args.seed) flow = load_teacher(args.model_dir, device) teacher_est = flow.decoder.estimator student = MeanFlowDiT.from_teacher(teacher_est, **CV3_DIT_KWARGS).to(device).train() ema = MeanFlowDiT.from_teacher(teacher_est, **CV3_DIT_KWARGS).to(device).eval() ema.load_state_dict(student.state_dict()) for p in ema.parameters(): p.requires_grad_(False) ds = ShardDataset(args.shards) opt = torch.optim.AdamW(student.parameters(), lr=args.lr, weight_decay=0.01, betas=(0.9, 0.95)) sched = torch.optim.lr_scheduler.LambdaLR( opt, lambda s: min(1.0, s / max(1, args.warmup)) * (0.5 * (1 + math.cos(math.pi * min(1.0, s / args.steps))))) # grille 2-NFE de reference (cosine warp de {0, .5, 1}) grid2 = cosine_warp(torch.tensor([0.0, 0.5, 1.0])) t0 = time.time() for step in range(1, args.steps + 1): tok, mel, xv, p_len, m_len = ds.make_batch(args.bs, rng, device) with torch.autocast("cuda", dtype=torch.bfloat16): mu, conds, mask, spks = build_conditions(flow, tok, mel, xv, p_len, m_len) x1 = mel.transpose(1, 2) # [B,80,T] B = x1.size(0) # intervalle (t,r) z = torch.randn_like(x1) grid_anchor = rng.random() < args.grid_prob second_leg = False if grid_anchor: i = rng.randrange(2) second_leg = (i == 1) t_s = grid2[i].expand(B).to(device) r_s = grid2[i + 1].expand(B).to(device) else: u = torch.rand(B, 2, device=device) lo, hi = u.min(1).values, u.max(1).values hi = torch.clamp(hi, min=lo + 0.05) t_s, r_s = cosine_warp(lo), cosine_warp(torch.clamp(hi, max=1.0)) tt = t_s.view(B, 1, 1) if second_leg: # entree du 2e pas = SORTIE REELLE du 1er pas du student (composition d'inference), # pas l'interpolation donnees -> corrige le mismatch on-policy with torch.no_grad(): ta = grid2[0].expand(B).to(device) ra = grid2[1].expand(B).to(device) u0 = student(z, mask, mu, ta, r=ra, spks=spks, cond=conds) x_t = z + (grid2[1] - grid2[0]).to(device) * u0 else: x_t = (1 - (1 - SIGMA) * tt) * z + tt * x1 # cible: vitesse moyenne du teacher CFG sur [t,r] (substeps Euler) with torch.no_grad(): x = x_t cur = t_s.clone() dt = (r_s - t_s) / args.substeps for _ in range(args.substeps): v = teacher_cfg_v(teacher_est, x, mask, mu, cur, spks, conds) x = x + dt.view(B, 1, 1) * v cur = cur + dt u_tgt = (x - x_t) / (r_s - t_s).view(B, 1, 1) u_pred = student(x_t, mask, mu, t_s, r=r_s, spks=spks, cond=conds) # perte sur la region cible uniquement (apres le prompt), masquee tgt_mask = mask.clone() for i in range(B): tgt_mask[i, :, : p_len[i]] = 0 loss = ((u_pred - u_tgt) ** 2 * tgt_mask).sum() / (tgt_mask.sum() * 80) opt.zero_grad(set_to_none=True) loss.backward() gn = torch.nn.utils.clip_grad_norm_(student.parameters(), 1.0) opt.step() sched.step() with torch.no_grad(): for pe, ps in zip(ema.parameters(), student.parameters()): pe.mul_(args.ema).add_(ps, alpha=1 - args.ema) if step % 50 == 0: print(f"step {step} loss {loss.item():.4f} gnorm {gn:.2f} lr {sched.get_last_lr()[0]:.2e} " f"{(time.time()-t0)/step:.2f}s/step", flush=True) if step % args.val_every == 0: val(flow, teacher_est, ema, ds, device, args.out, step) val(flow, teacher_est, student, ds, device, args.out, step, tag="raw") student.train() if step % args.save_every == 0: torch.save({"student": student.state_dict(), "ema": ema.state_dict(), "step": step, "args": vars(args)}, f"{args.out}/ckpt_{step:06d}.pt") torch.save({"student": student.state_dict(), "ema": ema.state_dict(), "step": args.steps, "args": vars(args)}, f"{args.out}/ckpt_final.pt") @torch.no_grad() def val(flow, teacher_est, ema, ds, device, out, step, tag="ema"): ema.eval() rng = random.Random(0) items = ds.val_items[:4] tok, mel, xv, p_len, m_len = collate(items, 1200, (4.0, 4.0), rng, device) with torch.autocast("cuda", dtype=torch.bfloat16): mu, conds, mask, spks = build_conditions(flow, tok, mel, xv, p_len, m_len) g = torch.Generator(device=device).manual_seed(123) z = torch.randn(mel.shape[0], 80, mel.shape[1], device=device, generator=g) # teacher 10 steps CFG (reference) ts10 = cosine_warp(torch.linspace(0, 1, 11)).to(device) x = z.clone() for i in range(10): t = ts10[i].expand(z.size(0)) v = teacher_cfg_v(teacher_est, x, mask, mu, t, spks, conds) x = x + (ts10[i + 1] - ts10[i]) * v mel_teacher = x # student EMA 2-NFE et 1-NFE ts2 = cosine_warp(torch.linspace(0, 1, 3)).to(device) mel_s2 = student_infer(ema, z.clone(), mu, mask, spks, conds, ts2) ts1 = cosine_warp(torch.linspace(0, 1, 2)).to(device) mel_s1 = student_infer(ema, z.clone(), mu, mask, spks, conds, ts1) # baseline = teacher 2 steps (la ligne de depart du student a l'init) x = z.clone() for i in range(2): t = ts2[i].expand(z.size(0)) v = teacher_cfg_v(teacher_est, x, mask, mu, t, spks, conds) x = x + (ts2[i + 1] - ts2[i]) * v mel_t2 = x tm = mask.clone() for i in range(len(items)): tm[i, :, : p_len[i]] = 0 l2 = lambda a, b: (((a - b) ** 2 * tm).sum() / (tm.sum() * 80)).item() print(f"VAL[{tag}] step {step}: student2={l2(mel_s2.float(), mel_teacher.float()):.4f} " f"student1={l2(mel_s1.float(), mel_teacher.float()):.4f} " f"BASELINE teacher2={l2(mel_t2.float(), mel_teacher.float()):.4f} (vs teacher10)", flush=True) torch.save({"teacher10": mel_teacher.float().cpu(), "student2": mel_s2.float().cpu(), "student1": mel_s1.float().cpu(), "p_len": p_len.cpu(), "m_len": m_len.cpu()}, f"{out}/val_{tag}_{step:06d}.pt") if __name__ == "__main__": main()