# Eval audio d'un checkpoint Reflow : greffe le student dans le pipeline CosyVoice3 complet, # genere le clone (student 1-NFE et 2-NFE) + la reference teacher 10-step, sauve les WAV. # Usage: python eval_audio.py --model_dir models/cosyvoice3-0.5b --ckpt runs/rf1/ckpt_010000.pt \ # --ref data/damien_15s.wav --ref_text data/damien_15s.txt --text "..." --out_dir evals/rf1_10k import argparse, os, sys, types import torch, torchaudio sys.path.insert(0, os.path.dirname(os.path.abspath(__file__))) from student import CV3_DIT_KWARGS from cosyvoice.flow.DiT.dit import DiT from cosyvoice.cli.cosyvoice import CosyVoice3 def make_reflow_forward(student, nfe): # remplace CausalConditionalCFM.forward : solveur chemin droit, student, sans CFG (absorbe) @torch.inference_mode() def forward(self, mu, mask, n_timesteps, temperature=1.0, spks=None, cond=None, prompt_len=0, cache=torch.zeros(1, 80, 0, 2), streaming=False, **kw): z = torch.randn_like(mu) * temperature ts = torch.linspace(0, 1, nfe + 1, device=mu.device, dtype=mu.dtype) x = z for i in range(nfe): v = student(x, mask, mu, ts[i].expand(x.size(0)), spks, cond) x = x + (ts[i + 1] - ts[i]) * v return x.float(), cache return forward def main(): ap = argparse.ArgumentParser() ap.add_argument("--model_dir", required=True) ap.add_argument("--ckpt", required=True) ap.add_argument("--ref", required=True) ap.add_argument("--ref_text", required=True) ap.add_argument("--text", required=True) ap.add_argument("--out_dir", required=True) ap.add_argument("--nfe", type=int, nargs="+", default=[1, 2]) args = ap.parse_args() os.makedirs(args.out_dir, exist_ok=True) rt = open(args.ref_text).read().strip() if os.path.exists(args.ref_text) else args.ref_text ref_text = "You are a helpful assistant.<|endofprompt|>" + rt cv = CosyVoice3(args.model_dir, load_trt=False, fp16=False) decoder = cv.model.flow.decoder orig_forward = decoder.forward # teacher 10-step CFG (bound method) student = DiT(**CV3_DIT_KWARGS).to("cuda").eval() sd = torch.load(args.ckpt, map_location="cpu") student.load_state_dict(sd["ema"] if "ema" in sd else sd["student"], strict=True) print(f"ckpt {args.ckpt} step {sd.get('step','?')} charge", flush=True) def synth(out_path): chunks = [o["tts_speech"] for o in cv.inference_zero_shot(args.text, ref_text, args.ref, stream=False)] wav = torch.cat(chunks, dim=1) torchaudio.save(out_path, wav, cv.sample_rate) return wav.shape[1] / cv.sample_rate # teacher reference (pipeline intact) d = synth(f"{args.out_dir}/teacher10.wav") print(f"teacher10.wav {d:.2f}s", flush=True) # student few-step for nfe in args.nfe: decoder.forward = types.MethodType(make_reflow_forward(student, nfe), decoder) d = synth(f"{args.out_dir}/student_{nfe}nfe.wav") print(f"student_{nfe}nfe.wav {d:.2f}s", flush=True) decoder.forward = orig_forward print("EVAL-DONE", flush=True) if __name__ == "__main__": main()