Download scripts/train.py from suvradeepp/tiny-hinglish-turn-detector: direct link, hf CLI and curl.
- Browser
- Download file 9.15 kB
-
https://huggingface.co/suvradeepp/tiny-hinglish-turn-detector/resolve/main/scripts/train.py
- Command line
-
hf download hf://suvradeepp/tiny-hinglish-turn-detector/scripts/train.py
-
curl -L -o train.py https://huggingface.co/suvradeepp/tiny-hinglish-turn-detector/resolve/main/scripts/train.py
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()) | |