pocket-tts-mnn / convert_to_mnn.py
developerabu's picture
Add Pocket TTS MNN INT8 conversion and hybrid runtime
4e75a38 verified
Raw
History Blame Contribute Delete
4.11 kB
#!/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())