File size: 3,814 Bytes
88e15cd | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 | 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
|