suvradeepp's picture
Publish Tiny Hinglish Turn Detector development preview
0348402 verified
Raw
History Blame Contribute Delete
17.9 kB
#!/usr/bin/env python3
"""Export a self-describing endpoint model to ONNX and optionally INT8."""
from __future__ import annotations
import argparse
import hashlib
import inspect
import json
import math
import sys
from pathlib import Path
from typing import Any
REPOSITORY_ROOT = Path(__file__).resolve().parents[1]
SOURCE_ROOT = REPOSITORY_ROOT / "src"
if str(SOURCE_ROOT) not in sys.path:
sys.path.insert(0, str(SOURCE_ROOT))
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--checkpoint", required=True)
parser.add_argument("--output", required=True, help="FP32 .onnx output path")
parser.add_argument("--opset", type=int, default=17)
parser.add_argument(
"--quantize",
choices=("none", "dynamic", "static"),
default="none",
help="dynamic suits transformers; static suits the TinyTCN CNN",
)
parser.add_argument(
"--calibration-npz",
help="static INT8 arrays: log_mel [N,M,T], frame_mask [N,T]",
)
parser.add_argument("--skip-parity", action="store_true")
return parser.parse_args()
def _sha256(path: Path) -> str:
digest = hashlib.sha256()
with path.open("rb") as handle:
for block in iter(lambda: handle.read(1024 * 1024), b""):
digest.update(block)
return digest.hexdigest()
def _file_evidence(path: Path) -> dict[str, Any]:
if path.is_symlink() or not path.is_file():
raise SystemExit(f"cannot bind non-regular source file: {path}")
resolved = path.resolve()
try:
portable = resolved.relative_to(REPOSITORY_ROOT).as_posix()
except ValueError:
portable = resolved.name
return {
"path": portable,
"bytes": resolved.stat().st_size,
"sha256": _sha256(resolved),
}
def _deployment_source_paths() -> list[Path]:
"""Return the exact executable source surface shipped with an export."""
paths = [
*sorted((REPOSITORY_ROOT / "src" / "turn_detection").rglob("*.py")),
*sorted((REPOSITORY_ROOT / "scripts").glob("*.py")),
*sorted((REPOSITORY_ROOT / "scripts").glob("*.sh")),
*(
path
for path in sorted((REPOSITORY_ROOT / "deployment").rglob("*"))
if path.is_file() and "__pycache__" not in path.parts and path.suffix != ".pyc"
),
REPOSITORY_ROOT / "app.py",
REPOSITORY_ROOT / "pyproject.toml",
REPOSITORY_ROOT / "space" / "requirements.txt",
*sorted(REPOSITORY_ROOT.glob("requirements-*.txt")),
]
return sorted(set(paths), key=lambda path: path.relative_to(REPOSITORY_ROOT).as_posix())
def _legacy_export_without_onnx_package(
torch: Any,
model: Any,
model_args: tuple[Any, ...],
output_path: Path,
*,
input_names: list[str],
output_names: list[str],
dynamic_axes: dict[str, dict[int, str]],
opset: int,
) -> None:
"""Serialize via Torch's private legacy graph only when ``onnx`` is absent.
This narrow fallback is useful in network-restricted build environments.
It is intentionally not used for arbitrary exporter failures, and the
resulting graph is still required to pass ONNX Runtime parity below.
"""
graph, params, _ = torch.onnx.utils._model_to_graph(
model,
model_args,
input_names=input_names,
output_names=output_names,
operator_export_type=torch.onnx.OperatorExportTypes.ONNX,
do_constant_folding=True,
training=torch.onnx.TrainingMode.EVAL,
dynamic_axes=dynamic_axes,
)
serialized, *_ = graph._export_onnx(
params,
opset,
dynamic_axes,
False,
torch.onnx.OperatorExportTypes.ONNX,
True,
False,
{},
True,
"",
{},
)
output_path.write_bytes(serialized)
def _quantize_dynamic(source: Path, destination: Path) -> None:
try:
from onnxruntime.quantization import QuantType, quantize_dynamic
except ImportError as exc:
raise SystemExit("INT8 export requires onnxruntime") from exc
quantize_dynamic(
str(source),
str(destination),
weight_type=QuantType.QInt8,
per_channel=True,
)
def _quantize_static(source: Path, destination: Path, calibration_path: Path) -> None:
try:
import numpy as np
from onnxruntime.quantization import (
CalibrationDataReader,
CalibrationMethod,
QuantFormat,
QuantType,
quantize_static,
)
except ImportError as exc:
raise SystemExit("static INT8 export requires numpy and onnxruntime") from exc
loaded = np.load(calibration_path)
if "log_mel" not in loaded or "frame_mask" not in loaded:
raise SystemExit("calibration NPZ needs log_mel and frame_mask arrays")
features = loaded["log_mel"].astype("float32")
masks = loaded["frame_mask"].astype("float32")
if features.ndim != 3 or masks.shape != (features.shape[0], features.shape[2]):
raise SystemExit("invalid calibration shapes")
class Reader(CalibrationDataReader):
def __init__(self) -> None:
self.index = 0
def get_next(self) -> dict[str, Any] | None:
if self.index >= features.shape[0]:
return None
item = {
"log_mel": features[self.index : self.index + 1],
"frame_mask": masks[self.index : self.index + 1],
}
self.index += 1
return item
quantize_static(
str(source),
str(destination),
Reader(),
quant_format=QuantFormat.QDQ,
activation_type=QuantType.QInt8,
weight_type=QuantType.QInt8,
per_channel=True,
calibrate_method=CalibrationMethod.MinMax,
)
def _parity_check(model_path: Path, features: Any, mask: Any, expected: Any) -> float:
try:
import numpy as np
import onnxruntime as ort
except ImportError as exc:
raise SystemExit("ONNX parity checking requires numpy and onnxruntime") from exc
session = ort.InferenceSession(str(model_path), providers=["CPUExecutionProvider"])
actual = session.run(
["endpoint_probability"],
{
"log_mel": features.detach().cpu().numpy().astype("float32"),
"frame_mask": mask.detach().cpu().numpy().astype("float32"),
},
)[0]
return float(np.max(np.abs(actual - expected.detach().cpu().numpy())))
def main() -> int:
args = parse_args()
try:
import torch
from torch import nn
except ImportError as exc:
raise SystemExit("ONNX export requires PyTorch") from exc
from turn_detection.models import (
LogMelConfig,
build_runtime_metadata,
load_model_checkpoint,
)
checkpoint_path = Path(args.checkpoint)
if not checkpoint_path.is_absolute():
checkpoint_path = REPOSITORY_ROOT / checkpoint_path
output_path = Path(args.output)
if not output_path.is_absolute():
output_path = REPOSITORY_ROOT / output_path
if output_path.suffix.lower() != ".onnx":
raise SystemExit("--output must end in .onnx")
output_path.parent.mkdir(parents=True, exist_ok=True)
model, checkpoint = load_model_checkpoint(checkpoint_path, map_location="cpu")
model.eval()
checkpoint_metadata = dict(checkpoint.get("metadata", {}))
feature_config = LogMelConfig.from_mapping(checkpoint_metadata.get("feature_config", {}))
max_seconds = float(checkpoint_metadata.get("max_seconds", 8.0))
frames = max(
2, int(round(max_seconds * feature_config.sample_rate / feature_config.hop_length))
)
model_type = str(checkpoint["model_config"].get("type", "tiny_tcn"))
fixed_frames = model_type in {"whisper", "whisper_teacher", "teacher"}
threshold = float(checkpoint.get("threshold", 0.5))
if not math.isfinite(threshold):
raise SystemExit("checkpoint threshold must be finite")
run_metadata = checkpoint_metadata.get("run_metadata", {})
if not isinstance(run_metadata, dict):
run_metadata = {}
smoke_test = bool(checkpoint_metadata.get("smoke_test", False))
training_status = str(run_metadata.get("status", "smoke" if smoke_test else "development"))
# Only the exact, explicit status "final" opens the final-release path.
# Candidate/production/release-like free text remains development-only.
development_only = smoke_test or training_status.lower() != "final"
data_scope = checkpoint_metadata.get("data_scope")
try:
metadata = build_runtime_metadata(
feature_config,
max_seconds=max_seconds,
threshold=threshold,
model_name=str(checkpoint_metadata.get("run_name", output_path.stem)),
architecture=model_type,
model_version=str(checkpoint.get("format_version", 1)),
development_only=development_only,
training_status=training_status,
data_scope=None if data_scope is None else str(data_scope),
data_revision=(
None
if checkpoint_metadata.get("data_revision") is None
else str(checkpoint_metadata["data_revision"])
),
parameter_count=sum(parameter.numel() for parameter in model.parameters()),
)
except ValueError as exc:
raise SystemExit(
f"checkpoint preprocessing cannot be represented by the current runtime: {exc}. "
"Export a deployment-compatible distilled TinyTCN student."
) from exc
class EndpointWrapper(nn.Module):
def __init__(self, wrapped: nn.Module) -> None:
super().__init__()
self.wrapped = wrapped
def forward(self, log_mel: Any, frame_mask: Any) -> Any:
return torch.sigmoid(self.wrapped(log_mel, frame_mask > 0.5).endpoint_logits)
wrapper = EndpointWrapper(model).eval()
generator = torch.Generator().manual_seed(17)
dummy_features = torch.randn(
(1, feature_config.n_mels, frames), generator=generator, dtype=torch.float32
)
dummy_mask = torch.ones((1, frames), dtype=torch.float32)
with torch.inference_mode():
expected = wrapper(dummy_features, dummy_mask)
dynamic_axes = {
"log_mel": {0: "batch"},
"frame_mask": {0: "batch"},
"endpoint_probability": {0: "batch"},
}
if not fixed_frames:
dynamic_axes["log_mel"][2] = "frames"
dynamic_axes["frame_mask"][1] = "frames"
try:
exporter_options: dict[str, Any] = {}
if "dynamo" in inspect.signature(torch.onnx.export).parameters:
exporter_options["dynamo"] = False
torch.onnx.export(
wrapper,
(dummy_features, dummy_mask),
str(output_path),
input_names=["log_mel", "frame_mask"],
output_names=["endpoint_probability"],
dynamic_axes=dynamic_axes,
opset_version=args.opset,
do_constant_folding=True,
**exporter_options,
)
except Exception as exc:
missing_module = isinstance(exc, ModuleNotFoundError) and getattr(exc, "name", None) in {
"onnx",
"onnxscript",
}
missing_message = str(exc) in {
"Module onnx is not installed!",
"No module named 'onnx'",
"No module named 'onnxscript'",
}
if not (missing_module or missing_message):
raise
print(
"warning: onnx package unavailable; using Torch's private legacy serializer",
file=sys.stderr,
)
try:
_legacy_export_without_onnx_package(
torch,
wrapper,
(dummy_features, dummy_mask),
output_path,
input_names=["log_mel", "frame_mask"],
output_names=["endpoint_probability"],
dynamic_axes=dynamic_axes,
opset=args.opset,
)
except Exception as fallback_exc:
raise SystemExit(
"ONNX package is unavailable and Torch's private fallback was incompatible"
) from fallback_exc
parity: dict[str, float | None] = {
"fp32_max_abs_error": None,
"int8_max_abs_error": None,
}
if not args.skip_parity:
parity["fp32_max_abs_error"] = _parity_check(
output_path, dummy_features, dummy_mask, expected
)
if parity["fp32_max_abs_error"] > 1e-4:
raise SystemExit(f"FP32 ONNX parity failed: {parity['fp32_max_abs_error']:.6g}")
quantized_path: Path | None = None
if args.quantize != "none":
quantized_path = output_path.with_name(output_path.stem + ".int8.onnx")
if args.quantize == "dynamic":
_quantize_dynamic(output_path, quantized_path)
else:
if not args.calibration_npz:
raise SystemExit("--quantize static requires --calibration-npz")
_quantize_static(output_path, quantized_path, Path(args.calibration_npz))
if not args.skip_parity:
parity["int8_max_abs_error"] = _parity_check(
quantized_path, dummy_features, dummy_mask, expected
)
files: dict[str, dict[str, Any]] = {
"fp32": {
"filename": output_path.name,
"bytes": output_path.stat().st_size,
"sha256": _sha256(output_path),
}
}
if quantized_path is not None:
files["int8"] = {
"filename": quantized_path.name,
"bytes": quantized_path.stat().st_size,
"sha256": _sha256(quantized_path),
"quantization": args.quantize,
}
resolved_config_path = checkpoint_path.parent / "resolved_config.json"
resolved_config_evidence: dict[str, Any] | None = None
training_data: dict[str, Any] | None = None
if resolved_config_path.is_file():
try:
resolved_config = json.loads(resolved_config_path.read_text(encoding="utf-8"))
except json.JSONDecodeError as exc:
raise SystemExit("resolved_config.json is invalid") from exc
resolved_config_evidence = _file_evidence(resolved_config_path)
data_config = resolved_config.get("data", {})
if isinstance(data_config, dict):
sources: dict[str, Any] = {}
for key in ("train_source", "validation_source"):
value = data_config.get(key)
if not isinstance(value, str):
continue
candidate = Path(value)
if not candidate.is_absolute():
candidate = REPOSITORY_ROOT / candidate
sources[key] = (
_file_evidence(candidate) if candidate.is_file() else {"identifier": value}
)
training_data = {
"revision": data_config.get("revision"),
"scope": data_config.get("scope"),
"sources": sources,
}
source_files = [_file_evidence(path) for path in _deployment_source_paths()]
source_inventory_sha256 = hashlib.sha256(
json.dumps(source_files, sort_keys=True, separators=(",", ":")).encode("utf-8")
).hexdigest()
export_manifest = {
"format_version": 2,
"task": "audio-turn-end-detection",
"model_type": model_type,
"parameter_count": sum(parameter.numel() for parameter in model.parameters()),
"checkpoint": {
"filename": checkpoint_path.name,
"bytes": checkpoint_path.stat().st_size,
"sha256": _sha256(checkpoint_path),
"selected_epoch": checkpoint.get("epoch"),
},
"model_config": checkpoint.get("model_config"),
"threshold": threshold,
"controller": metadata["controller"],
"resolved_config": resolved_config_evidence,
"training_data": training_data,
"source_files": source_files,
"source_inventory_sha256": source_inventory_sha256,
"input_names": ["log_mel", "frame_mask"],
"output_names": ["endpoint_probability"],
"input_dtypes": {"log_mel": "float32", "frame_mask": "float32"},
"input_shapes": {
"log_mel": ["batch", feature_config.n_mels, frames if fixed_frames else "frames"],
"frame_mask": ["batch", frames if fixed_frames else "frames"],
},
"dynamic_frames": not fixed_frames,
"files": files,
"parity": parity,
"quantized_threshold_recalibration_required": quantized_path is not None,
"development_only": development_only,
"training_status": training_status,
"data_scope": data_scope,
"data_revision": checkpoint_metadata.get("data_revision"),
"notes": (
"Whisper export uses a fixed time axis dictated by encoder positional embeddings."
if fixed_frames
else "TinyTCN accepts a dynamic number of log-mel frames."
),
}
metadata_path = output_path.parent / "model_metadata.json"
metadata_path.write_text(
json.dumps(metadata, indent=2, sort_keys=True, allow_nan=False), encoding="utf-8"
)
export_manifest_path = output_path.parent / "export_manifest.json"
export_manifest_path.write_text(
json.dumps(export_manifest, indent=2, sort_keys=True, allow_nan=False),
encoding="utf-8",
)
print(
json.dumps(
{
"model": str(output_path),
"metadata": str(metadata_path),
"export_manifest": str(export_manifest_path),
**parity,
},
indent=2,
)
)
return 0
if __name__ == "__main__":
raise SystemExit(main())