diff --git a/src/coolrl_lost_cities/games/classic/deep_cfr/cli.py b/src/coolrl_lost_cities/games/classic/deep_cfr/cli.py index 6ddaf69..87963fa 100644 --- a/src/coolrl_lost_cities/games/classic/deep_cfr/cli.py +++ b/src/coolrl_lost_cities/games/classic/deep_cfr/cli.py @@ -17,6 +17,9 @@ from coolrl_lost_cities.games.classic.deep_cfr.evaluate import ( from coolrl_lost_cities.games.classic.deep_cfr.imitation import ( new_pretrained_strategy_network, ) +from coolrl_lost_cities.games.classic.deep_cfr.policy_gradient import ( + fine_tune_strategy_policy_gradient, +) from coolrl_lost_cities.games.classic.deep_cfr.trainer import DeepCFRTrainer from coolrl_lost_cities.games.classic.game import classic_config @@ -100,6 +103,28 @@ def pretrain_command(args: argparse.Namespace) -> None: print(json.dumps(metrics.__dict__, sort_keys=True)) +def policy_gradient_command(args: argparse.Namespace) -> None: + policy, game_config = load_strategy_policy_from_checkpoint(args.checkpoint, device=args.device) + metrics = fine_tune_strategy_policy_gradient( + policy.strategy_network, + game_config, + episodes=args.episodes, + seed=args.seed, + opponent=args.opponent, + learning_rate=args.learning_rate, + max_steps=args.max_steps, + device=args.device, + ) + if args.output: + import torch + + torch.save( + {"strategy_network": policy.strategy_network.state_dict(), "metrics": metrics.__dict__}, + args.output, + ) + print(json.dumps(metrics.__dict__, sort_keys=True)) + + def main(argv: list[str] | None = None) -> None: parser = argparse.ArgumentParser(description="Lost Cities classic Deep CFR tools.") subparsers = parser.add_subparsers(dest="command", required=True) @@ -142,6 +167,17 @@ def main(argv: list[str] | None = None) -> None: pretrain.add_argument("--output") pretrain.set_defaults(func=pretrain_command) + pg = subparsers.add_parser("policy-gradient") + pg.add_argument("--checkpoint", required=True) + pg.add_argument("--episodes", type=int, default=2) + pg.add_argument("--opponent", default="random") + pg.add_argument("--learning-rate", type=float, default=1.0e-4) + pg.add_argument("--max-steps", type=int, default=10_000) + pg.add_argument("--seed", type=int, default=1) + pg.add_argument("--device", default="cpu") + pg.add_argument("--output") + pg.set_defaults(func=policy_gradient_command) + args = parser.parse_args(argv) args.func(args) diff --git a/src/coolrl_lost_cities/games/classic/deep_cfr/policy_gradient.py b/src/coolrl_lost_cities/games/classic/deep_cfr/policy_gradient.py new file mode 100644 index 0000000..c2a9893 --- /dev/null +++ b/src/coolrl_lost_cities/games/classic/deep_cfr/policy_gradient.py @@ -0,0 +1,79 @@ +from __future__ import annotations + +from dataclasses import dataclass + +import numpy as np +import torch +from torch import nn + +from coolrl_lost_cities.games.classic.bots import build_bot +from coolrl_lost_cities.games.classic.deep_cfr.encoding import encode_info_state +from coolrl_lost_cities.games.classic.game import GameState, LostCitiesConfig, classic_config + + +@dataclass(frozen=True) +class PolicyGradientMetrics: + episodes: int + avg_reward: float + loss: float + + +def fine_tune_strategy_policy_gradient( + strategy_network: nn.Module, + config: LostCitiesConfig | None = None, + *, + episodes: int = 2, + seed: int = 1, + opponent: str = "random", + learning_rate: float = 1.0e-4, + max_steps: int = 10_000, + device: torch.device | str = "cpu", +) -> PolicyGradientMetrics: + game_config = config or classic_config(seed=seed) + device = torch.device(device) + strategy_network.to(device) + strategy_network.train() + optimizer = torch.optim.Adam(strategy_network.parameters(), lr=learning_rate) + rng = np.random.default_rng(seed + 31) + rewards: list[float] = [] + losses: list[torch.Tensor] = [] + for episode in range(episodes): + state = GameState.new_game(game_config, seed=seed + episode) + opponent_policy = build_bot(opponent, seed=seed + episode + 1000) + log_probs: list[torch.Tensor] = [] + for _ in range(max_steps): + if state.terminal: + break + if state.current_player == 0: + info = encode_info_state(state, 0) + legal = torch.as_tensor(state.unified_legal_mask(), dtype=torch.bool, device=device) + x = torch.as_tensor(info, dtype=torch.float32, device=device).unsqueeze(0) + logits = ( + strategy_network(x) + .squeeze(0) + .masked_fill(~legal, torch.finfo(torch.float32).min) + ) + probs = torch.softmax(logits, dim=-1) + action_index = int(rng.choice(len(probs), p=probs.detach().cpu().numpy())) + log_probs.append(torch.log(probs[action_index].clamp_min(1.0e-12))) + action = state.from_unified_action(action_index) + else: + action = opponent_policy.act(state) + state.apply_action(action) + reward = float(state.score_diff(0)) + rewards.append(reward) + if log_probs: + losses.append(-torch.stack(log_probs).sum() * reward) + if losses: + loss = torch.stack(losses).mean() + optimizer.zero_grad(set_to_none=True) + loss.backward() + optimizer.step() + loss_value = float(loss.detach().cpu()) + else: + loss_value = 0.0 + return PolicyGradientMetrics( + episodes=episodes, + avg_reward=float(np.mean(rewards)) if rewards else 0.0, + loss=loss_value, + ) diff --git a/tests/games/classic/test_deep_cfr_policy_gradient.py b/tests/games/classic/test_deep_cfr_policy_gradient.py new file mode 100644 index 0000000..9029342 --- /dev/null +++ b/tests/games/classic/test_deep_cfr_policy_gradient.py @@ -0,0 +1,27 @@ +from __future__ import annotations + +from coolrl_lost_cities.games.classic.deep_cfr.encoding import input_dim +from coolrl_lost_cities.games.classic.game import GameState, LostCitiesConfig + +from coolrl_lost_cities.games.classic.deep_cfr.networks import DeepCFRMLP +from coolrl_lost_cities.games.classic.deep_cfr.policy_gradient import ( + fine_tune_strategy_policy_gradient, +) + + +def test_policy_gradient_fine_tune_smoke() -> None: + config = LostCitiesConfig(seed=71) + state = GameState.new_game(config, seed=71) + network = DeepCFRMLP(input_dim(state), 2 * config.hand_size + 1 + config.n_colors, 16) + + metrics = fine_tune_strategy_policy_gradient( + network, + config, + episodes=1, + seed=71, + max_steps=64, + ) + + assert metrics.episodes == 1 + assert isinstance(metrics.avg_reward, float) + assert isinstance(metrics.loss, float)