llm / dipaugnet /training /engine.py
abersbail's picture
Replace llm Space with DIPAug project hub
9c2e807 verified
Raw
History Blame Contribute Delete
1.91 kB
"""Phase 1 training engine skeleton."""
from __future__ import annotations
from dataclasses import dataclass
from typing import Any
import torch
from torch.cuda.amp import GradScaler, autocast
from dipauglib.schedulers.adaptive import AdaptiveAugmentationScheduler
@dataclass
class TrainState:
"""Training state bundle."""
epoch: int = 0
best_val_f1: float = 0.0
patience_counter: int = 0
def train_one_epoch(
model: torch.nn.Module,
loader: Any,
optimizer: torch.optim.Optimizer,
criterion: torch.nn.Module,
device: torch.device,
use_amp: bool = True,
) -> float:
"""Train one epoch."""
model.train()
scaler = GradScaler(enabled=use_amp and device.type == "cuda")
losses: list[float] = []
for batch in loader:
images = batch["image"].to(device)
targets = batch["target"].to(device)
optimizer.zero_grad(set_to_none=True)
with autocast(enabled=use_amp and device.type == "cuda"):
logits = model(images)["logits"]
loss = criterion(logits, targets)
scaler.scale(loss).backward()
scaler.unscale_(optimizer)
torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)
scaler.step(optimizer)
scaler.update()
losses.append(float(loss.detach().cpu()))
return float(sum(losses) / max(1, len(losses)))
def fit_phase1(model: torch.nn.Module, optimizer: torch.optim.Optimizer, scheduler: Any | None = None) -> dict[str, Any]:
"""Placeholder fit entry for config-driven runners."""
aug_scheduler = AdaptiveAugmentationScheduler()
return {
"status": "scaffold_only",
"message": "Phase 1 training loop skeleton created.",
"initial_aug_intensity": aug_scheduler.intensity_at(0),
"mid_aug_intensity": aug_scheduler.intensity_at(50),
"scheduler_present": scheduler is not None,
}