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