From 647baa7d6d8f6a06fd2c2460b4742421e5eba51c Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=EC=A0=95=EC=8B=9C=EC=9B=90?= Date: Thu, 7 May 2026 00:12:41 +0900 Subject: [PATCH] =?UTF-8?q?Deep=20CFR=20imitation=20pretraining=20?= =?UTF-8?q?=EC=B6=94=EA=B0=80?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- .../games/classic/deep_cfr/cli.py | 28 +++++ .../games/classic/deep_cfr/imitation.py | 108 ++++++++++++++++++ .../games/classic/test_deep_cfr_imitation.py | 31 +++++ 3 files changed, 167 insertions(+) create mode 100644 src/coolrl_lost_cities/games/classic/deep_cfr/imitation.py create mode 100644 tests/games/classic/test_deep_cfr_imitation.py 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 6597000..6ddaf69 100644 --- a/src/coolrl_lost_cities/games/classic/deep_cfr/cli.py +++ b/src/coolrl_lost_cities/games/classic/deep_cfr/cli.py @@ -14,6 +14,9 @@ from coolrl_lost_cities.games.classic.deep_cfr.evaluate import ( evaluate_strategy_network, load_strategy_policy_from_checkpoint, ) +from coolrl_lost_cities.games.classic.deep_cfr.imitation import ( + new_pretrained_strategy_network, +) from coolrl_lost_cities.games.classic.deep_cfr.trainer import DeepCFRTrainer from coolrl_lost_cities.games.classic.game import classic_config @@ -80,6 +83,23 @@ def benchmark_command(args: argparse.Namespace) -> None: print(json.dumps(result, indent=2, sort_keys=True)) +def pretrain_command(args: argparse.Namespace) -> None: + network, metrics = new_pretrained_strategy_network( + classic_config(seed=args.seed), + hidden_size=args.hidden_size, + games=args.games, + seed=args.seed, + steps=args.steps, + ) + if args.output: + import torch + + torch.save( + {"strategy_network": 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) @@ -114,6 +134,14 @@ def main(argv: list[str] | None = None) -> None: benchmark.add_argument("--compare", action="store_true") benchmark.set_defaults(func=benchmark_command) + pretrain = subparsers.add_parser("pretrain") + pretrain.add_argument("--games", type=int, default=4) + pretrain.add_argument("--steps", type=int, default=32) + pretrain.add_argument("--hidden-size", type=int, default=64) + pretrain.add_argument("--seed", type=int, default=1) + pretrain.add_argument("--output") + pretrain.set_defaults(func=pretrain_command) + args = parser.parse_args(argv) args.func(args) diff --git a/src/coolrl_lost_cities/games/classic/deep_cfr/imitation.py b/src/coolrl_lost_cities/games/classic/deep_cfr/imitation.py new file mode 100644 index 0000000..2325ef9 --- /dev/null +++ b/src/coolrl_lost_cities/games/classic/deep_cfr/imitation.py @@ -0,0 +1,108 @@ +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 SafeHeuristicBot +from coolrl_lost_cities.games.classic.deep_cfr.encoding import encode_info_state, input_dim +from coolrl_lost_cities.games.classic.deep_cfr.networks import DeepCFRMLP +from coolrl_lost_cities.games.classic.game import GameState, LostCitiesConfig, classic_config + + +@dataclass(frozen=True) +class ImitationMetrics: + samples: int + loss: float + + +def collect_safe_heuristic_samples( + config: LostCitiesConfig | None = None, + *, + games: int = 4, + seed: int = 1, + max_steps: int = 10_000, +) -> tuple[np.ndarray, np.ndarray, np.ndarray]: + game_config = config or classic_config(seed=seed) + probe = GameState.new_game(game_config, seed=seed) + action_size = 2 * probe.config.hand_size + 1 + probe.config.n_colors + bot = SafeHeuristicBot() + infos: list[np.ndarray] = [] + targets: list[np.ndarray] = [] + masks: list[np.ndarray] = [] + for game_index in range(games): + state = GameState.new_game(game_config, seed=seed + game_index) + for _ in range(max_steps): + if state.terminal: + break + info = encode_info_state(state, state.current_player) + legal = np.asarray(state.unified_legal_mask(), dtype=bool) + action = bot.act(state) + unified = state.to_unified_action(action) + target = np.zeros(action_size, dtype=np.float32) + target[unified] = 1.0 + infos.append(info) + targets.append(target) + masks.append(legal) + state.apply_action(action) + if not infos: + raise RuntimeError("no imitation samples collected") + return ( + np.stack(infos).astype(np.float32), + np.stack(targets).astype(np.float32), + np.stack(masks).astype(bool), + ) + + +def pretrain_strategy_network( + strategy_network: nn.Module, + config: LostCitiesConfig | None = None, + *, + games: int = 4, + seed: int = 1, + steps: int = 32, + batch_size: int = 64, + learning_rate: float = 1.0e-3, + device: torch.device | str = "cpu", +) -> ImitationMetrics: + x_np, y_np, legal_np = collect_safe_heuristic_samples(config, games=games, 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 + 17) + last_loss = 0.0 + for _ in range(max(0, steps)): + indices = rng.choice( + len(x_np), size=min(batch_size, len(x_np)), replace=len(x_np) < batch_size + ) + x = torch.as_tensor(x_np[indices], dtype=torch.float32, device=device) + y = torch.as_tensor(y_np[indices], dtype=torch.float32, device=device) + legal = torch.as_tensor(legal_np[indices], dtype=torch.bool, device=device) + logits = strategy_network(x).masked_fill(~legal, torch.finfo(torch.float32).min) + log_probs = nn.functional.log_softmax(logits, dim=-1).masked_fill(~legal, 0.0) + loss = -(y * log_probs).sum(dim=-1).mean() + optimizer.zero_grad(set_to_none=True) + loss.backward() + optimizer.step() + last_loss = float(loss.detach().cpu()) + return ImitationMetrics(samples=len(x_np), loss=last_loss) + + +def new_pretrained_strategy_network( + config: LostCitiesConfig | None = None, + *, + hidden_size: int = 64, + games: int = 4, + seed: int = 1, + steps: int = 32, +) -> tuple[DeepCFRMLP, ImitationMetrics]: + game_config = config or classic_config(seed=seed) + probe = GameState.new_game(game_config, seed=seed) + network = DeepCFRMLP( + input_dim(probe), 2 * probe.config.hand_size + 1 + probe.config.n_colors, hidden_size + ) + metrics = pretrain_strategy_network(network, game_config, games=games, seed=seed, steps=steps) + return network, metrics diff --git a/tests/games/classic/test_deep_cfr_imitation.py b/tests/games/classic/test_deep_cfr_imitation.py new file mode 100644 index 0000000..981d607 --- /dev/null +++ b/tests/games/classic/test_deep_cfr_imitation.py @@ -0,0 +1,31 @@ +from __future__ import annotations + +from coolrl_lost_cities.games.classic.game import LostCitiesConfig + +from coolrl_lost_cities.games.classic.deep_cfr.imitation import ( + collect_safe_heuristic_samples, + new_pretrained_strategy_network, +) + + +def test_collect_safe_heuristic_samples_shapes() -> None: + x, y, legal = collect_safe_heuristic_samples(LostCitiesConfig(seed=61), games=1, seed=61) + + assert len(x) == len(y) == len(legal) + assert x.ndim == 2 + assert y.shape == legal.shape + assert y.sum(axis=1).min() == 1.0 + + +def test_pretrain_strategy_network_smoke() -> None: + network, metrics = new_pretrained_strategy_network( + LostCitiesConfig(seed=67), + hidden_size=16, + games=1, + seed=67, + steps=1, + ) + + assert metrics.samples > 0 + assert metrics.loss >= 0.0 + assert network is not None