| 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 |
|
|
|
|