Files
Agentic-AI/scripts/export-tamil-tts.py
T

100 lines
3.5 KiB
Python

"""One-time export of facebook/mms-tts-tam to ONNX.
No public ONNX build of Tamil MMS-TTS exists, so we make one. This runs ONCE on
a workstation; the committed artefact is what ships. The service itself is pure
JavaScript and never needs Python or this script.
The output must match the contract Transformers.js expects for VITS, taken from
the working English model (Xenova/mms-tts-eng):
inputs : input_ids, attention_mask
outputs: waveform, spectrogram
Layout produced (mirrors the HF repo so Transformers.js can load the folder):
assets/tts/mms-tts-tam/
config.json, tokenizer.json, vocab.json, …
onnx/model.onnx fp32
onnx/model_quantized.onnx int8 ← what we ship
"""
from __future__ import annotations
import json
import shutil
import sys
from pathlib import Path
import torch
from transformers import AutoTokenizer, VitsModel
MODEL = "facebook/mms-tts-tam"
OUT = Path("assets/tts/mms-tts-tam")
ONNX_DIR = OUT / "onnx"
ONNX_DIR.mkdir(parents=True, exist_ok=True)
print(f"loading {MODEL} …")
model = VitsModel.from_pretrained(MODEL).eval()
tok = AutoTokenizer.from_pretrained(MODEL)
# VITS has a stochastic duration predictor. Exporting with noise left on bakes
# RandomNormalLike nodes into the graph, which is fine and keeps prosody
# natural — but seed it so this export is reproducible.
torch.manual_seed(0)
class Exportable(torch.nn.Module):
"""Return only (waveform, spectrogram), in that order — the JS side indexes
outputs by name, but a plain tuple keeps the exported graph simple."""
def __init__(self, m: VitsModel) -> None:
super().__init__()
self.m = m
def forward(self, input_ids: torch.Tensor, attention_mask: torch.Tensor):
out = self.m(input_ids=input_ids, attention_mask=attention_mask)
return out.waveform, out.spectrogram
sample = tok("வணக்கம், இது ஒரு சோதனை.", return_tensors="pt")
fp32 = ONNX_DIR / "model.onnx"
print("exporting to ONNX …")
torch.onnx.export(
Exportable(model),
(sample["input_ids"], sample["attention_mask"]),
str(fp32),
input_names=["input_ids", "attention_mask"],
output_names=["waveform", "spectrogram"],
dynamic_axes={
"input_ids": {0: "batch", 1: "sequence"},
"attention_mask": {0: "batch", 1: "sequence"},
"waveform": {0: "batch", 1: "samples"},
"spectrogram": {0: "batch", 2: "frames"},
},
opset_version=17,
do_constant_folding=True,
)
print(f" fp32: {fp32.stat().st_size / 1e6:.1f} MB")
# ── int8 ────────────────────────────────────────────────────────────────────
try:
from onnxruntime.quantization import QuantType, quantize_dynamic
q = ONNX_DIR / "model_quantized.onnx"
quantize_dynamic(str(fp32), str(q), weight_type=QuantType.QUInt8)
print(f" int8: {q.stat().st_size / 1e6:.1f} MB")
except Exception as e: # noqa: BLE001
print(f" quantisation skipped: {e}")
# ── tokenizer + config, so the folder loads standalone ──────────────────────
tok.save_pretrained(OUT)
model.config.to_json_file(OUT / "config.json")
# Transformers.js reads this to pick a default dtype.
(OUT / "quantize_config.json").write_text(json.dumps({"per_channel": False, "reduce_range": False}, indent=2))
print("\nwrote:")
for p in sorted(OUT.rglob("*")):
if p.is_file():
print(f" {p.relative_to(OUT)} ({p.stat().st_size / 1e6:.2f} MB)")