kokoro-kmp-models / scripts /export_onnx.py
Shusek00's picture
Publish complete Kokoro multilingual catalog v2.0.0
0dbae7f verified
Raw
History Blame Contribute Delete
7.41 kB
#!/usr/bin/env python3
"""Export a Kokoro checkpoint to a self-contained ONNX model and verify it."""
from __future__ import annotations
import argparse
import hashlib
import json
from pathlib import Path
import numpy as np
import onnx
import onnxruntime as ort
import torch
from kokoro import KModel
from kokoro.model import KModelForONNX
def sha256(path: Path) -> str:
digest = hashlib.sha256()
with path.open("rb") as source:
for block in iter(lambda: source.read(1024 * 1024), b""):
digest.update(block)
return digest.hexdigest()
def make_reference_inputs(
model: KModel,
voice: Path,
phonemes: str,
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
missing = sorted({phoneme for phoneme in phonemes if model.vocab.get(phoneme) is None})
if missing:
raise RuntimeError(f"Reference phonemes missing from the Kokoro vocabulary: {missing}")
if not 1 <= len(phonemes) <= 510:
raise RuntimeError(f"Reference phoneme length must be in 1..510, got {len(phonemes)}")
voice_pack = np.fromfile(voice, dtype=np.float32)
if voice_pack.size != 510 * 256:
raise RuntimeError(
f"Expected a raw 510x256 float32 voice pack, got {voice_pack.size} values"
)
voice_pack = voice_pack.reshape(510, 256)
token_ids = [int(model.vocab[phoneme]) for phoneme in phonemes]
input_ids = torch.tensor([[0, *token_ids, 0]], dtype=torch.long)
style = torch.from_numpy(voice_pack[len(phonemes) - 1]).unsqueeze(0)
speed = torch.tensor([1.0], dtype=torch.float32)
return input_ids, style, speed
def add_metadata(
output: Path,
*,
language: str,
checkpoint: Path,
source_repository: str,
source_checkpoint_name: str,
source_revision: str,
base_revision: str,
) -> None:
model = onnx.load(output)
onnx.helper.set_model_props(
model,
{
"architecture": "Kokoro-82M",
"language": language,
"license": "Apache-2.0",
"sample_rate_hz": "24000",
"source_repository": source_repository,
"source_revision": source_revision,
"source_checkpoint": source_checkpoint_name,
"source_checkpoint_sha256": sha256(checkpoint),
"base_repository": "hexgrad/Kokoro-82M",
"base_revision": base_revision,
},
)
onnx.save_model(model, output)
def export(args: argparse.Namespace) -> None:
checkpoint = args.checkpoint.resolve()
config = args.config.resolve()
voice = args.voice.resolve()
output = args.output.resolve()
output.parent.mkdir(parents=True, exist_ok=True)
torch.manual_seed(0)
model = KModel(
repo_id="hexgrad/Kokoro-82M",
config=str(config),
model=str(checkpoint),
disable_complex=True,
).cpu().eval()
export_model = KModelForONNX(model).cpu().eval()
input_ids, style, speed = make_reference_inputs(
model,
voice,
args.reference_phonemes,
)
torch.onnx.export(
export_model,
args=(input_ids, style, speed),
f=str(output),
input_names=["input_ids", "style", "speed"],
output_names=["waveform", "duration"],
opset_version=17,
dynamic_axes={
"input_ids": {1: "sequence_length"},
"waveform": {0: "num_samples"},
"duration": {0: "sequence_length"},
},
do_constant_folding=True,
dynamo=False,
)
add_metadata(
output,
language=args.language,
checkpoint=checkpoint,
source_repository=args.source_repository,
source_checkpoint_name=args.source_checkpoint_name or checkpoint.name,
source_revision=args.source_revision,
base_revision=args.base_revision,
)
onnx.checker.check_model(str(output), full_check=True)
options = ort.SessionOptions()
options.graph_optimization_level = ort.GraphOptimizationLevel.ORT_ENABLE_ALL
session = ort.InferenceSession(str(output), sess_options=options, providers=["CPUExecutionProvider"])
# Kokoro deliberately injects a small amount of random excitation noise. As a
# result, PyTorch and ONNX waveforms should be highly correlated, but are not
# expected to be bit-identical or pass a point-wise allclose assertion.
torch.manual_seed(0)
with torch.inference_mode():
torch_waveform, torch_duration = export_model(input_ids, style, speed)
ort_waveform, ort_duration = session.run(
None,
{
"input_ids": input_ids.numpy(),
"style": style.numpy(),
"speed": speed.numpy(),
},
)
expected_waveform = torch_waveform.detach().cpu().numpy()
expected_duration = torch_duration.detach().cpu().numpy()
np.testing.assert_array_equal(ort_duration, expected_duration)
if not np.isfinite(ort_waveform).all():
raise RuntimeError("ONNX Runtime produced non-finite waveform samples")
delta = np.abs(ort_waveform - expected_waveform)
correlation = float(np.corrcoef(ort_waveform, expected_waveform)[0, 1])
mean_absolute_error = float(delta.mean())
if correlation < 0.99 or mean_absolute_error > 0.005:
raise RuntimeError(
"ONNX/PyTorch validation failed: "
f"correlation={correlation:.6f}, meanAbsoluteError={mean_absolute_error:.6f}"
)
validation = {
"language": args.language,
"model": output.name,
"modelBytes": output.stat().st_size,
"modelSha256": sha256(output),
"sourceCheckpoint": args.source_checkpoint_name or checkpoint.name,
"sourceCheckpointSha256": sha256(checkpoint),
"sourceRevision": args.source_revision,
"baseRevision": args.base_revision,
"opset": 17,
"sampleRateHz": 24000,
"voice": voice.name,
"voiceSha256": sha256(voice),
"referencePhonemes": args.reference_phonemes,
"referenceInputTokens": int(input_ids.shape[1]),
"referenceOutputSamples": int(ort_waveform.size),
"maxAbsoluteError": float(delta.max(initial=0.0)),
"meanAbsoluteError": mean_absolute_error,
"waveformCorrelation": correlation,
"durationExactMatch": True,
"onnxRuntimeVersion": ort.__version__,
"torchVersion": torch.__version__,
}
validation_path = output.with_suffix(".validation.json")
validation_path.write_text(json.dumps(validation, indent=2) + "\n", encoding="utf-8")
print(json.dumps(validation, indent=2))
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser()
parser.add_argument("--checkpoint", type=Path, required=True)
parser.add_argument("--config", type=Path, required=True)
parser.add_argument("--voice", type=Path, required=True)
parser.add_argument("--output", type=Path, required=True)
parser.add_argument("--language", required=True)
parser.add_argument("--reference-phonemes", required=True)
parser.add_argument(
"--source-repository",
default="software-mansion/react-native-executorch-kokoro",
)
parser.add_argument("--source-checkpoint-name")
parser.add_argument("--source-revision", required=True)
parser.add_argument("--base-revision", required=True)
return parser.parse_args()
if __name__ == "__main__":
export(parse_args())