| from __future__ import annotations |
|
|
| from pathlib import Path |
|
|
| import numpy as np |
| import torch |
| from cmaes import SepCMA |
| from rich.console import Console |
|
|
| from .config import ESConfig |
| from .io import write_json |
| from .schemas import RewardRecord |
| from .training_common import make_loader, make_model, move_batch, seed_everything |
|
|
| console = Console() |
|
|
|
|
| @torch.no_grad() |
| def _cached_features(model, loader) -> tuple[np.ndarray, np.ndarray]: |
| model.eval() |
| features = [] |
| rewards = [] |
| for batch in loader: |
| inputs, batch_rewards = move_batch(batch, model.device_ref) |
| features.append(model.encode_features(**inputs).cpu().numpy()) |
| rewards.append(batch_rewards.cpu().numpy()) |
| return np.concatenate(features), np.concatenate(rewards) |
|
|
|
|
| def _unpack(candidate: np.ndarray, worker_count: int, hidden_size: int): |
| weight_size = worker_count * hidden_size |
| weight = candidate[:weight_size].reshape(worker_count, hidden_size) |
| bias = candidate[weight_size:] |
| return weight, bias |
|
|
|
|
| def _candidate_utility( |
| candidate: np.ndarray, |
| features: np.ndarray, |
| rewards: np.ndarray, |
| worker_count: int, |
| objective: str, |
| ) -> float: |
| weight, bias = _unpack(candidate, worker_count, features.shape[1]) |
| logits = features @ weight.T + bias |
| if objective == "expected_utility": |
| shifted = logits - logits.max(axis=1, keepdims=True) |
| probabilities = np.exp(shifted) |
| probabilities /= probabilities.sum(axis=1, keepdims=True) |
| return float((probabilities * rewards).sum(axis=1).mean()) |
| actions = logits.argmax(axis=1) |
| return float(rewards[np.arange(len(rewards)), actions].mean()) |
|
|
|
|
| def train_es( |
| records: list[RewardRecord], |
| config: ESConfig, |
| output_dir: str | Path, |
| checkpoint: str | None = None, |
| ) -> dict: |
| seed_everything(config.training.seed) |
| worker_ids = records[0].worker_ids |
| model = make_model(config.model, worker_ids, checkpoint) |
| train_loader = make_loader(records, model, "train", batch_size=32, shuffle=False) |
| train_features, train_rewards = _cached_features(model, train_loader) |
| try: |
| validation_loader = make_loader( |
| records, model, "validation", batch_size=32, shuffle=False |
| ) |
| validation_features, validation_rewards = _cached_features(model, validation_loader) |
| except ValueError: |
| validation_features = validation_rewards = None |
|
|
| initial = np.concatenate( |
| [ |
| model.classifier.weight.detach().cpu().numpy().ravel(), |
| model.classifier.bias.detach().cpu().numpy().ravel(), |
| ] |
| ).astype(np.float64) |
| kwargs = { |
| "mean": initial, |
| "sigma": config.training.sigma, |
| "seed": config.training.seed, |
| } |
| if config.training.population_size is not None: |
| kwargs["population_size"] = config.training.population_size |
| optimizer = SepCMA(**kwargs) |
| best_candidate = initial.copy() |
| best_train_utility = _candidate_utility( |
| initial, |
| train_features, |
| train_rewards, |
| len(worker_ids), |
| config.training.objective, |
| ) |
| history = [] |
| for generation in range(config.training.generations): |
| solutions = [] |
| generation_best = -float("inf") |
| for _ in range(optimizer.population_size): |
| candidate = optimizer.ask() |
| utility = _candidate_utility( |
| candidate, |
| train_features, |
| train_rewards, |
| len(worker_ids), |
| config.training.objective, |
| ) |
| solutions.append((candidate, -utility)) |
| generation_best = max(generation_best, utility) |
| if utility > best_train_utility: |
| best_train_utility = utility |
| best_candidate = candidate.copy() |
| optimizer.tell(solutions) |
| row = { |
| "generation": generation + 1, |
| "generation_best_train_utility": generation_best, |
| "best_train_utility": best_train_utility, |
| } |
| if validation_features is not None and validation_rewards is not None: |
| row["validation_utility"] = _candidate_utility( |
| best_candidate, |
| validation_features, |
| validation_rewards, |
| len(worker_ids), |
| config.training.objective, |
| ) |
| history.append(row) |
| console.print(row) |
|
|
| weight, bias = _unpack(best_candidate, len(worker_ids), train_features.shape[1]) |
| with torch.no_grad(): |
| model.classifier.weight.copy_(torch.from_numpy(weight).to(model.classifier.weight)) |
| model.classifier.bias.copy_(torch.from_numpy(bias).to(model.classifier.bias)) |
| output = Path(output_dir) |
| model.save_checkpoint( |
| output, |
| metadata={"stage": "sep-cma-es", "history": history}, |
| ) |
| report = { |
| "stage": "sep-cma-es", |
| "checkpoint": str(output), |
| "worker_ids": worker_ids, |
| "best_train_utility": best_train_utility, |
| "history": history, |
| } |
| write_json(output / "training_report.json", report) |
| return report |
|
|
|
|