| |
| """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") |
|
|
| |
| 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()) |
|
|