fugu-lite / src /fugu_lite /train_es.py
tahsinsoyak's picture
Upload private Fugu-Lite V1 snapshot
88e15cd verified
Raw
History Blame
5.13 kB
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