267 lines
12 KiB
Python
267 lines
12 KiB
Python
# 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()
|