# Distillation par rectified-flow (Reflow/InstaFlow) du flow CosyVoice3. # Cible FIXE et detachee: x_teacher = sortie teacher 10-step+CFG depuis z (CFG absorbe). # Student = DiT pur init=teacher; regresse la vitesse du chemin droit z->x_teacher. # Robuste: pas de cible dynamique, pas de bouclage on-policy -> pas de divergence. # Usage: python train_reflow.py --model_dir models/cosyvoice3-0.5b --shards "shards_cml/*.pt" --out runs/rf1 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 CV3_DIT_KWARGS from cosyvoice.flow.DiT.dit import DiT 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: cfg = load_hyperpyyaml(f, overrides={"llm": None, "hift": None}) flow = cfg["flow"] flow.load_state_dict(torch.load(f"{model_dir}/flow.pt", map_location="cpu"), strict=True) flow.to(device).eval() for p in flow.parameters(): p.requires_grad_(False) return flow def make_student(teacher_est, device): s = DiT(**CV3_DIT_KWARGS) s.load_state_dict(teacher_est.state_dict(), strict=True) return s.to(device) class ShardDataset: def __init__(self, pattern, max_mel=1000, 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 self.val_items = torch.load(self.files[0], map_location="cpu")[:holdout] self._ci, self._c = -1, None def _shard(self, i): if i != self._ci: self._c, self._ci = torch.load(self.files[i], map_location="cpu"), i return self._c def make_batch(self, bs, rng, device): items = [] for _ in range(bs): sh = self._shard(rng.randrange(len(self.files))) items.append(sh[rng.randrange(len(sh))]) 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 = min(int(rng.uniform(*prompt_s) * 25), tl - 25) 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): B, T = mel.shape[0], mel.shape[1] spks = flow.spk_embed_affine_layer(F.normalize(xv, dim=1)) 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) 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 @torch.no_grad() def teacher_cfg_v(est, x, mask, mu, t, spks, cond): B = x.size(0) v = est(torch.cat([x, x], 0), torch.cat([mask, mask], 0), torch.cat([mu, torch.zeros_like(mu)], 0), torch.cat([t, t], 0), torch.cat([spks, torch.zeros_like(spks)], 0), torch.cat([cond, torch.zeros_like(cond)], 0)) vc, vu = v[:B], v[B:] return (1.0 + CFG) * vc - CFG * vu @torch.no_grad() def teacher_solve(est, z, mask, mu, spks, cond, nfe=10): ts = cosine_warp(torch.linspace(0, 1, nfe + 1)).to(z.device) x = z for i in range(nfe): v = teacher_cfg_v(est, x, mask, mu, ts[i].expand(z.size(0)), spks, cond) x = x + (ts[i + 1] - ts[i]) * v return x @torch.no_grad() def student_solve(student, z, mask, mu, spks, cond, nfe): ts = torch.linspace(0, 1, nfe + 1).to(z.device) # reflow = chemin droit, pas de warp x = z for i in range(nfe): v = student(x, mask, mu, ts[i].expand(z.size(0)), spks, cond) x = x + (ts[i + 1] - ts[i]) * v return x 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=60000) ap.add_argument("--lr", type=float, default=1e-5) ap.add_argument("--warmup", type=int, default=1000) ap.add_argument("--ema", type=float, default=0.9995) ap.add_argument("--val_every", type=int, default=1000) ap.add_argument("--save_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) test = flow.decoder.estimator student = make_student(test, device).train() ema = make_student(test, device).eval() 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))))) 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 = x1.size(0) z = torch.randn_like(x1) # cible FIXE detachee : sortie teacher 10-step CFG depuis z (CFG absorbe) x_tea = teacher_solve(test, z, mask, mu, spks, conds, nfe=10) # chemin droit reflow : x_t = (1-t) z + t x_tea ; vitesse cible = x_tea - z t = torch.rand(B, 1, 1, device=device) x_t = (1 - t) * z + t * x_tea v_tgt = x_tea - z v_pred = student(x_t, mask, mu, t.view(B), spks, conds) tgt_mask = mask.clone() for i in range(B): tgt_mask[i, :, : p_len[i]] = 0 loss = ((v_pred - v_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 {float(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, test, ema, ds, device, args.out, step) student.train() if step % args.save_every == 0: torch.save({"student": student.state_dict(), "ema": ema.state_dict(), "step": step}, f"{args.out}/ckpt_{step:06d}.pt") torch.save({"student": student.state_dict(), "ema": ema.state_dict(), "step": args.steps}, f"{args.out}/ckpt_final.pt") @torch.no_grad() def val(flow, test, ema, ds, device, out, step): rng = random.Random(0) tok, mel, xv, p_len, m_len = collate(ds.val_items[:4], 1000, (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) mel_tea = teacher_solve(test, z, mask, mu, spks, conds, nfe=10) s1 = student_solve(ema, z, mask, mu, spks, conds, 1) s2 = student_solve(ema, z, mask, mu, spks, conds, 2) s4 = student_solve(ema, z, mask, mu, spks, conds, 4) tm = mask.clone() for i in range(len(p_len)): tm[i, :, : p_len[i]] = 0 l2 = lambda a, b: (((a - b) ** 2 * tm).sum() / (tm.sum() * 80)).item() print(f"VAL step {step}: s1={l2(s1.float(), mel_tea.float()):.4f} " f"s2={l2(s2.float(), mel_tea.float()):.4f} s4={l2(s4.float(), mel_tea.float()):.4f} (vs teacher10)", flush=True) torch.save({"teacher10": mel_tea.float().cpu(), "s1": s1.float().cpu(), "s2": s2.float().cpu(), "s4": s4.float().cpu(), "p_len": p_len.cpu(), "m_len": m_len.cpu()}, f"{out}/val_{step:06d}.pt") if __name__ == "__main__": main()