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