72 lines
3.1 KiB
Python
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()
|