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