#!/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())