#!/usr/bin/env python3 """Convert PocketTTS ONNX graphs to MNN with INT8 weight quantization. Rewrites unsupported ONNX::IsNaN -> Not(Equal(x,x)) before conversion. """ from __future__ import annotations import argparse import subprocess import sys from pathlib import Path import onnx from onnx import helper ROOT = Path(__file__).resolve().parent STEMS = [ "flow_lm_main", "flow_lm_flow", "mimi_decoder", "mimi_encoder", "text_conditioner", ] def find_converter() -> str: for candidate in ( "/opt/homebrew/bin/MNNConvert", "MNNConvert", str(ROOT / ".venv/bin/MNNConvert"), ): path = Path(candidate) if candidate.startswith("/") else None if path and path.exists(): return str(path) return "MNNConvert" def replace_isnan(model: onnx.ModelProto) -> int: graph = model.graph new_nodes = [] replaced = 0 for node in list(graph.node): if node.op_type != "IsNaN": new_nodes.append(node) continue inp = node.input[0] out = node.output[0] eq_out = f"{out}__eq_self" base = node.name or out new_nodes.append(helper.make_node("Equal", [inp, inp], [eq_out], name=f"{base}__eq")) new_nodes.append(helper.make_node("Not", [eq_out], [out], name=f"{base}__not")) replaced += 1 del graph.node[:] graph.node.extend(new_nodes) return replaced def convert_one( converter: str, src: Path, dst: Path, weight_bits: int, optimize_level: int, ) -> None: dst.parent.mkdir(parents=True, exist_ok=True) cmd = [ converter, "-f", "ONNX", "--modelFile", str(src), "--MNNModel", str(dst), "--bizCode", "PocketTTS", "--optimizeLevel", str(optimize_level), ] if weight_bits > 0: cmd.extend(["--weightQuantBits", str(weight_bits), "--weightQuantAsymmetric"]) print(" ".join(cmd), flush=True) subprocess.run(cmd, check=True) print(f" -> {dst} ({dst.stat().st_size / 1024 / 1024:.1f} MB)", flush=True) def main() -> int: parser = argparse.ArgumentParser() parser.add_argument("--models-dir", default=str(ROOT / "models")) parser.add_argument("--out-dir", default=str(ROOT / "models" / "mnn")) parser.add_argument("--weight-bits", type=int, default=8) parser.add_argument( "--text-fp32", action="store_true", default=True, help="Keep text_conditioner as FP32 MNN (embedding table)", ) parser.add_argument("--optimize-level", type=int, default=1) args = parser.parse_args() models_dir = Path(args.models_dir) out_dir = Path(args.out_dir) onnx_src = models_dir / "onnx_mnn_src" onnx_src.mkdir(exist_ok=True) out_dir.mkdir(exist_ok=True) converter = find_converter() print(f"Using converter: {converter}") for stem in STEMS: src_onnx = models_dir / f"{stem}.onnx" if not src_onnx.exists(): print(f"Missing {src_onnx}", file=sys.stderr) return 1 model = onnx.load(str(src_onnx)) n = replace_isnan(model) rewritten = onnx_src / f"{stem}.onnx" onnx.save(model, str(rewritten)) print(f"{stem}: rewrote IsNaN={n}") bits = 0 if (stem == "text_conditioner" and args.text_fp32) else args.weight_bits suffix = "fp32" if bits == 0 else f"w{bits}" dst = out_dir / f"{stem}_{suffix}.mnn" convert_one(converter, rewritten, dst, bits, args.optimize_level) # also keep alias expected by runtime if stem == "text_conditioner" and bits == 0: alias = out_dir / "text_conditioner_w8.mnn" if not alias.exists(): convert_one(converter, rewritten, alias, 8, args.optimize_level) total = sum(p.stat().st_size for p in out_dir.glob("*.mnn") if "opt" not in p.name and "test" not in p.name and "static" not in p.name) print(f"Done. Core MNN models under {out_dir}") return 0 if __name__ == "__main__": raise SystemExit(main())