suvradeepp's picture
Publish Tiny Hinglish Turn Detector development preview
35d483e verified
Raw
History Blame Contribute Delete
9.15 kB
#!/usr/bin/env python3
"""Train a TinyTCN student or optional Whisper teacher.
Examples
--------
Fast end-to-end validation without corpus access::
python scripts/train.py --config configs/smoke.json --smoke-test
Real split manifests::
python scripts/train.py --config configs/tiny_tcn.yaml
"""
from __future__ import annotations
import argparse
import hashlib
import json
import os
import sys
from dataclasses import asdict, fields
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("--config", default="configs/tiny_tcn.yaml")
parser.add_argument(
"--set",
action="append",
default=[],
metavar="KEY=VALUE",
help="dotted JSON-valued config override; may be repeated",
)
parser.add_argument(
"--smoke-test",
action="store_true",
help="train only on deterministic generated features",
)
parser.add_argument(
"--max-examples",
type=int,
help="debug cap per split (not suitable for reported experiments)",
)
return parser.parse_args()
def _dataclass_kwargs(cls: type, values: dict[str, Any]) -> dict[str, Any]:
allowed = {field.name for field in fields(cls)}
return {key: value for key, value in values.items() if key in allowed}
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 _warm_start_model(model: Any, checkpoint_path: Path, torch: Any) -> dict[str, Any]:
"""Load model weights only, deliberately starting a fresh optimizer/schedule."""
checkpoint = torch.load(checkpoint_path, map_location="cpu", weights_only=False)
if not isinstance(checkpoint, dict):
raise ValueError("initialization checkpoint must be a mapping")
expected_config = model.model_config() if hasattr(model, "model_config") else None
if checkpoint.get("model_config") != expected_config:
raise ValueError("initialization checkpoint architecture does not match this run")
state = checkpoint.get("model_state")
if not isinstance(state, dict):
raise ValueError("initialization checkpoint has no model_state")
model.load_state_dict(state, strict=True)
try:
portable = checkpoint_path.resolve().relative_to(REPOSITORY_ROOT).as_posix()
except ValueError:
portable = checkpoint_path.name
return {
"mode": "weights_only_fresh_optimizer",
"path": portable,
"sha256": _sha256(checkpoint_path),
"selected_epoch": checkpoint.get("epoch"),
"source_run": checkpoint.get("metadata", {}).get("run_name"),
}
def main() -> int:
args = parse_args()
try:
import torch
except ImportError as exc:
raise SystemExit(
"Training requires PyTorch. Install the project's training dependencies first."
) from exc
from turn_detection.models import LogMelConfig, LogMelFrontend, build_model
from turn_detection.training.config import apply_overrides, load_config
from turn_detection.training.datasets import (
build_record_dataloader,
build_smoke_dataloaders,
)
from turn_detection.training.losses import MultiTaskLossConfig
from turn_detection.training.trainer import Trainer, TrainerConfig, seed_everything
config_path = Path(args.config)
if not config_path.is_absolute():
config_path = REPOSITORY_ROOT / config_path
config = apply_overrides(load_config(config_path), args.set)
model_config = dict(config.get("model", {}))
feature_values = dict(config.get("features", {}))
data_config = dict(config.get("data", {}))
training_config = TrainerConfig.from_mapping(config.get("training", {}))
loss_config = MultiTaskLossConfig(
**_dataclass_kwargs(MultiTaskLossConfig, dict(config.get("loss", {})))
)
run_config = dict(config.get("run", {}))
feature_config = LogMelConfig.from_mapping(feature_values)
if int(model_config.get("n_mels", feature_config.n_mels)) != feature_config.n_mels:
raise SystemExit("model.n_mels must equal features.n_mels")
max_seconds = float(feature_values.get("max_seconds", 8.0))
if max_seconds <= 0:
raise SystemExit("features.max_seconds must be positive")
# Model initialization is seeded here; seeding only inside fit() would be too late.
seed_everything(training_config.seed, training_config.deterministic)
model = build_model(model_config)
initialization: dict[str, Any] | None = None
init_checkpoint = run_config.get("init_checkpoint")
if init_checkpoint:
init_path = Path(str(init_checkpoint))
if not init_path.is_absolute():
init_path = REPOSITORY_ROOT / init_path
init_path = init_path.resolve()
try:
init_path.relative_to(REPOSITORY_ROOT)
except ValueError as exc:
raise SystemExit("run.init_checkpoint must stay inside the project") from exc
if not init_path.is_file() or init_path.is_symlink():
raise SystemExit(f"run.init_checkpoint is not a regular file: {init_path}")
try:
initialization = _warm_start_model(model, init_path, torch)
except (OSError, RuntimeError, ValueError) as exc:
raise SystemExit(f"cannot warm-start model: {exc}") from exc
frontend = LogMelFrontend(feature_config)
batch_size = int(data_config.get("batch_size", 32))
if args.smoke_test:
train_loader, validation_loader = build_smoke_dataloaders(
frontend, batch_size=min(batch_size, 16), seed=training_config.seed
)
output_dir = Path(run_config.get("output_dir", "artifacts/smoke"))
else:
train_source = data_config.get("train_source")
validation_source = data_config.get("validation_source", train_source)
if not train_source or not validation_source:
raise SystemExit("data.train_source and data.validation_source are required")
common = {
"frontend": frontend,
"batch_size": batch_size,
"max_seconds": max_seconds,
"num_workers": int(data_config.get("num_workers", 0)),
"seed": training_config.seed,
"revision": data_config.get("revision"),
"token": os.environ.get("HF_TOKEN") or os.environ.get("HUGGING_FACE_HUB_TOKEN"),
"shuffle_buffer": int(data_config.get("shuffle_buffer", 64)),
"max_examples": args.max_examples,
"source_root": data_config.get("source_root"),
}
train_loader = build_record_dataloader(
train_source,
split=str(data_config.get("train_split", "train")),
shuffle=True,
**common,
)
validation_loader = build_record_dataloader(
validation_source,
split=str(data_config.get("validation_split", "validation")),
shuffle=False,
**common,
)
output_dir = Path(run_config.get("output_dir", "artifacts/run"))
if not output_dir.is_absolute():
output_dir = REPOSITORY_ROOT / output_dir
output_dir.mkdir(parents=True, exist_ok=True)
resolved_path = output_dir / "resolved_config.json"
resolved_path.write_text(json.dumps(config, indent=2, sort_keys=True), encoding="utf-8")
parameter_count = sum(parameter.numel() for parameter in model.parameters())
trainable_count = sum(
parameter.numel() for parameter in model.parameters() if parameter.requires_grad
)
print(
json.dumps(
{
"run": run_config.get("name", output_dir.name),
"device": training_config.device,
"parameters": parameter_count,
"trainable_parameters": trainable_count,
"smoke_test": args.smoke_test,
},
indent=2,
)
)
artifact_metadata = {
"feature_config": asdict(feature_config),
"max_seconds": max_seconds,
"run_name": run_config.get("name", output_dir.name),
"run_metadata": run_config,
"data_revision": data_config.get("revision"),
"data_scope": data_config.get("scope"),
"smoke_test": args.smoke_test,
"torch_version": torch.__version__,
"initialization": initialization,
}
trainer = Trainer(
model,
config=training_config,
loss_config=loss_config,
output_dir=output_dir,
artifact_metadata=artifact_metadata,
)
result = trainer.fit(train_loader, validation_loader)
summary = {key: value for key, value in result.items() if key != "history"}
print(json.dumps(summary, indent=2))
return 0
if __name__ == "__main__":
raise SystemExit(main())