Kazeia-engine/dist/distill/export_onnx.py

54 lines
2.2 KiB
Python

# Export du student distille (DiT) en ONNX batch=1 pour bench tablette (ORT).
# Inference few-step = appeler ce graphe N fois (chemin droit, pas de CFG).
# Usage: python export_onnx.py --ckpt runs/rf1/ckpt_010000.pt --out student.onnx
import argparse, os, sys
import torch
sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))
from student import CV3_DIT_KWARGS
from cosyvoice.flow.DiT.dit import DiT
class StudentWrap(torch.nn.Module):
# ordonne les entrees comme le flow estimator CV3 (x,mask,mu,t,spks,cond), batch1
def __init__(self, dit):
super().__init__()
self.dit = dit
def forward(self, x, mask, mu, t, spks, cond):
return self.dit(x, mask, mu, t, spks, cond, streaming=False)
def main():
ap = argparse.ArgumentParser()
ap.add_argument("--ckpt", required=True)
ap.add_argument("--out", required=True)
ap.add_argument("--T", type=int, default=650)
args = ap.parse_args()
dit = DiT(**CV3_DIT_KWARGS).eval()
sd = torch.load(args.ckpt, map_location="cpu")
dit.load_state_dict(sd["ema"] if "ema" in sd else sd["student"], strict=True)
T = args.T
x = torch.randn(1, 80, T); mask = torch.ones(1, 1, T); mu = torch.randn(1, 80, T)
t = torch.tensor([0.5]); spks = torch.randn(1, 80); cond = torch.randn(1, 80, T)
w = StudentWrap(dit)
with torch.no_grad():
ref = w(x, mask, mu, t, spks, cond)
torch.onnx.export(
w, (x, mask, mu, t, spks, cond), args.out,
input_names=["x", "mask", "mu", "t", "spks", "cond"], output_names=["v"],
opset_version=17, dynamo=False,
dynamic_axes={"x": {2: "T"}, "mask": {2: "T"}, "mu": {2: "T"}, "cond": {2: "T"}, "v": {2: "T"}})
print(f"export OK -> {args.out} ({os.path.getsize(args.out)/1e6:.1f} MB), out {tuple(ref.shape)} step {sd.get('step','?')}", flush=True)
import onnxruntime as ort, numpy as np
s = ort.InferenceSession(args.out, providers=["CPUExecutionProvider"])
o = s.run(None, {"x": x.numpy(), "mask": mask.numpy(), "mu": mu.numpy(),
"t": t.numpy(), "spks": spks.numpy(), "cond": cond.numpy()})[0]
print("parite max|diff|:", float(np.abs(o - ref.numpy()).max()), flush=True)
if __name__ == "__main__":
main()