"""Training pipeline for specialist extractor fine-tuning. Manages training runs on the Stonks Oracle schema, tracks artifacts, and produces evaluation-ready models for holdout testing. """ from __future__ import annotations import enum from dataclasses import dataclass, field from datetime import datetime, timezone from typing import Any from uuid import UUID, uuid4 class TrainingStatus(str, enum.Enum): """Status of a training run.""" PENDING = "pending" PREPARING_DATA = "preparing_data" TRAINING = "training" EVALUATING = "evaluating" COMPLETED = "completed" FAILED = "failed" @dataclass class TrainingConfig: """Configuration for specialist model fine-tuning.""" base_model: str = "GLiNER2-large" schema_version: str = "1.0" dataset_version: str = "" training_range: str = "" # e.g., "2024-01 to 2024-06" # Training parameters learning_rate: float = 2e-5 batch_size: int = 16 max_epochs: int = 10 warmup_steps: int = 100 weight_decay: float = 0.01 # Data split train_ratio: float = 0.8 validation_ratio: float = 0.1 holdout_ratio: float = 0.1 # Frozen holdout — never used in training # Entity types to fine-tune entity_types: list[str] = field(default_factory=lambda: [ "company", "person", "event", "financial_metric", "date", "money", "percentage", "ticker", ]) @dataclass class TrainingRun: """A single training run for the specialist extractor.""" run_id: UUID config: TrainingConfig status: TrainingStatus = TrainingStatus.PENDING started_at: datetime | None = None completed_at: datetime | None = None # Training metrics train_loss: float = 0.0 validation_loss: float = 0.0 best_epoch: int = 0 total_examples: int = 0 # Artifact tracking artifact_path: str = "" model_version: str = "" parent_model_version: str = "" # Metadata notes: str = "" errors: list[str] = field(default_factory=list) @classmethod def create(cls, config: TrainingConfig) -> TrainingRun: return cls( run_id=uuid4(), config=config, ) def start(self) -> None: """Begin training.""" self.status = TrainingStatus.PREPARING_DATA self.started_at = datetime.now(timezone.utc) def begin_training(self) -> None: """Transition to active training.""" self.status = TrainingStatus.TRAINING def begin_evaluation(self) -> None: """Transition to evaluation phase.""" self.status = TrainingStatus.EVALUATING def complete( self, artifact_path: str, model_version: str, train_loss: float = 0.0, validation_loss: float = 0.0, best_epoch: int = 0, ) -> None: """Mark training as complete with artifact metadata.""" self.status = TrainingStatus.COMPLETED self.completed_at = datetime.now(timezone.utc) self.artifact_path = artifact_path self.model_version = model_version self.train_loss = train_loss self.validation_loss = validation_loss self.best_epoch = best_epoch def fail(self, error: str) -> None: """Mark training as failed.""" self.status = TrainingStatus.FAILED self.completed_at = datetime.now(timezone.utc) self.errors.append(error) @property def duration_seconds(self) -> float | None: if self.started_at and self.completed_at: return (self.completed_at - self.started_at).total_seconds() return None def to_dict(self) -> dict[str, Any]: return { "run_id": str(self.run_id), "status": self.status.value, "base_model": self.config.base_model, "schema_version": self.config.schema_version, "dataset_version": self.config.dataset_version, "model_version": self.model_version, "artifact_path": self.artifact_path, "train_loss": self.train_loss, "validation_loss": self.validation_loss, "best_epoch": self.best_epoch, "duration_seconds": self.duration_seconds, }