225 lines
9.3 KiB
Python
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()
|