100 lines
3.5 KiB
Python
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)")
|