File size: 3,829 Bytes
9726ddb
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
#!/usr/bin/env python3
"""
Prepare Supertone/supertonic assets for TranslateBlue on-device TTS.

Native Core ML ``.mlpackage`` export from these ONNX graphs is not reliably supported
by current ``coremltools`` (ONNX path removed; legacy converters fail on modern ops).

This script downloads ``Supertone/supertonic`` ONNX + configs + voice styles from the Hub.
The app loads them with ONNX Runtime and the Core ML execution provider when available.

Usage::

    python3 -m venv .venv && . .venv/bin/activate
    pip install -r requirements-export.txt
    python convert_supertonic_coreml.py --output ./export

See README.md for layout and Hub upload paths.
"""

from __future__ import annotations

import argparse
import json
import shutil
import sys
from pathlib import Path
from typing import Iterable, List

from huggingface_hub import HfApi, hf_hub_download

REPO = "Supertone/supertonic"
ONNX_FILES = [
    "duration_predictor.onnx",
    "text_encoder.onnx",
    "vector_estimator.onnx",
    "vocoder.onnx",
]
JSON_FILES = ["tts.json", "unicode_indexer.json"]


def _download(repo_id: str, rel: str, dest: Path, revision: str) -> None:
    dest.parent.mkdir(parents=True, exist_ok=True)
    p = hf_hub_download(repo_id=repo_id, filename=rel, revision=revision)
    shutil.copy2(p, dest)


def _iter_voice_files(repo_id: str, revision: str) -> List[str]:
    api = HfApi()
    out: List[str] = []
    for ent in api.list_repo_tree(repo_id, revision=revision, recursive=True):
        p = getattr(ent, "path", None)
        if p and p.startswith("voice_styles/") and p.endswith(".json"):
            out.append(p)
    return sorted(out)


def copy_json_and_voices(repo_id: str, out_onnx: Path, out_voice: Path, revision: str) -> None:
    out_onnx.mkdir(parents=True, exist_ok=True)
    for j in JSON_FILES:
        _download(repo_id, f"onnx/{j}", out_onnx / j, revision)
    out_voice.mkdir(parents=True, exist_ok=True)
    for vf in _iter_voice_files(repo_id, revision):
        _download(repo_id, vf, out_voice / Path(vf).name, revision)


def write_manifest(path: Path, repo_id: str, revision: str, max_text_len: int) -> None:
    meta = {
        "repo_id": repo_id,
        "revision": revision,
        "max_text_len": max_text_len,
        "onnx_models": ONNX_FILES,
        "note": "fp32 ONNX from Hub under onnx/; runtime uses ORT CoreML EP on device when enabled.",
    }
    path.write_text(json.dumps(meta, indent=2), encoding="utf-8")


def validate_ort(onnx_dir: Path) -> None:
    import onnxruntime as ort

    dp = str(onnx_dir / "duration_predictor.onnx")
    sess = ort.InferenceSession(dp, providers=["CPUExecutionProvider"])
    names = [i.name for i in sess.get_inputs()]
    assert "text_ids" in names and "text_mask" in names, names


def main(argv: Iterable[str] | None = None) -> int:
    ap = argparse.ArgumentParser(description="Download Supertone/supertonic ONNX for TranslateBlue")
    ap.add_argument("--output", type=Path, default=Path("./export"))
    ap.add_argument("--repo-id", default=REPO)
    ap.add_argument("--revision", default="main")
    ap.add_argument("--max-text-len", type=int, default=300)
    ap.add_argument("--skip-validate", action="store_true")
    args = ap.parse_args(list(argv) if argv is not None else None)

    root: Path = args.output
    onnx_fp = root / "onnx"
    voices = root / "voice_styles"

    for name in ONNX_FILES:
        _download(args.repo_id, f"onnx/{name}", onnx_fp / name, args.revision)
    copy_json_and_voices(args.repo_id, onnx_fp, voices, args.revision)

    write_manifest(root / "manifest.json", args.repo_id, args.revision, args.max_text_len)

    if not args.skip_validate:
        validate_ort(onnx_fp)

    print(f"OK: wrote {root.resolve()} (fp32 onnx + voice_styles)")
    return 0


if __name__ == "__main__":
    sys.exit(main())