fugu-lite / src /fugu_lite /training_common.py
tahsinsoyak's picture
Upload private Fugu-Lite V1 snapshot
88e15cd verified
Raw
History Blame
3.81 kB
from __future__ import annotations
import random
from collections.abc import Iterable
import numpy as np
import torch
from torch.utils.data import DataLoader
from .config import ModelConfig
from .dataset import RouterCollator, RouterDataset
from .model import RouterModel
from .schemas import RewardRecord
def seed_everything(seed: int) -> None:
random.seed(seed)
np.random.seed(seed)
torch.manual_seed(seed)
if torch.cuda.is_available():
torch.cuda.manual_seed_all(seed)
def resolve_device() -> torch.device:
return torch.device("cuda" if torch.cuda.is_available() else "cpu")
def make_model(
model_config: ModelConfig,
worker_ids: list[str],
checkpoint: str | None = None,
) -> RouterModel:
device = resolve_device()
if checkpoint:
model = RouterModel.from_checkpoint(checkpoint, device=device)
if model.worker_ids != worker_ids:
raise ValueError(
f"Checkpoint workers {model.worker_ids} do not match dataset workers {worker_ids}"
)
return model
return RouterModel(model_config, worker_ids, device)
def make_loader(
records: list[RewardRecord],
model: RouterModel,
split: str,
batch_size: int,
shuffle: bool,
) -> DataLoader:
dataset = RouterDataset(records, split=split)
if not dataset:
raise ValueError(f"No {split} records found")
generator = torch.Generator().manual_seed(0)
return DataLoader(
dataset,
batch_size=batch_size,
shuffle=shuffle,
generator=generator,
collate_fn=RouterCollator(model.tokenizer, model.router_config.max_length),
num_workers=0,
pin_memory=torch.cuda.is_available(),
)
def move_batch(batch: dict, device: torch.device) -> tuple[dict, torch.Tensor]:
inputs = {
"input_ids": batch["input_ids"].to(device, non_blocking=True),
"attention_mask": batch["attention_mask"].to(device, non_blocking=True),
}
return inputs, batch["rewards"].to(device, non_blocking=True)
def optimizer_parameters(model: RouterModel) -> Iterable[torch.nn.Parameter]:
return (parameter for parameter in model.parameters() if parameter.requires_grad)
@torch.no_grad()
def evaluate_loader(model: RouterModel, loader: DataLoader) -> dict:
model.eval()
chosen_rewards = []
oracle_rewards = []
oracle_hits = []
all_rewards = []
route_counts = torch.zeros(len(model.worker_ids), dtype=torch.long)
for batch in loader:
inputs, rewards = move_batch(batch, model.device_ref)
logits = model(**inputs)
actions = logits.argmax(dim=-1)
chosen = rewards.gather(1, actions.unsqueeze(1)).squeeze(1)
oracle, oracle_actions = rewards.max(dim=1)
chosen_rewards.append(chosen.cpu())
oracle_rewards.append(oracle.cpu())
oracle_hits.append(actions.eq(oracle_actions).float().cpu())
all_rewards.append(rewards.cpu())
route_counts += torch.bincount(actions.cpu(), minlength=len(model.worker_ids))
chosen = torch.cat(chosen_rewards)
oracle = torch.cat(oracle_rewards)
rewards = torch.cat(all_rewards)
fixed_means = rewards.mean(dim=0)
report = {
"examples": int(chosen.numel()),
"router_utility": float(chosen.mean()),
"oracle_utility": float(oracle.mean()),
"regret": float((oracle - chosen).mean()),
"oracle_route_accuracy": float(torch.cat(oracle_hits).mean()),
"best_fixed_worker": model.worker_ids[int(fixed_means.argmax())],
"best_fixed_utility": float(fixed_means.max()),
"random_utility": float(rewards.mean()),
"route_counts": {
worker_id: int(count) for worker_id, count in zip(model.worker_ids, route_counts)
},
}
return report