Kazeia-engine/dist/distill/eval_audio.py

72 lines
3.1 KiB
Python

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