Deep CFR policy gradient fine-tuning 추가

This commit is contained in:
2026-05-07 00:13:53 +09:00
parent 647baa7d6d
commit b71f4b95be
3 changed files with 142 additions and 0 deletions
@@ -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)
@@ -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,
)
@@ -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)