Multi-stage evidence-grounded inference architecture replacing the monolithic 9B model extraction pipeline. CPU-first specialist services handle routine extraction while the 9B vLLM model is preserved for semantic adjudication of ambiguous cases. Key components: - Capability-aware inference gateway (OpenAI-compatible + Ollama) - Endpoint registry with DB migrations and REST API - Sentence-aware document segmenter (property tests) - Deterministic financial parsing with offset integrity - Symbol resolution with ambiguity detection - Specialist service (GLiNER2, dynamic batching, K8s deployment) - Company-specific sentiment (FinBERT, calibration) - Retrieval-based novelty and duplicate detection - Confidence calibration pipeline - Deterministic routing engine (property tests) - 9B adjudication layer with VRAM gating - Stock-specific impact model (features, labels, baseline, trained) - Pipeline orchestrator (state machine, queues, leases, feature flags) - Bounded parallelism (async workers, semaphore, load shedding) - Observability (tracing, metrics, alerts) - Compatibility adapter (v3→v2 golden mapping tests) - Shadow/canary promotion framework - Active learning and fine-tuning pipeline Test results: 1,161 tests pass, ruff lint clean. All 282 spec tasks completed.
143 lines
4.1 KiB
Python
143 lines
4.1 KiB
Python
"""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,
|
|
}
|