Kazeia-engine/dist/distill/train_reflow.py

225 lines
9.3 KiB
Python

# 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()