Deep CFR self-play league 추가

This commit is contained in:
2026-05-07 00:01:19 +09:00
parent a5bd5eeb7e
commit 46d841cf11
5 changed files with 161 additions and 8 deletions
@@ -19,6 +19,10 @@ class DeepCFRConfig:
cutoff_rollouts: int = 0 cutoff_rollouts: int = 0
cutoff_rollout_policy: str = "random" cutoff_rollout_policy: str = "random"
cutoff_rollout_max_steps: int = 10_000 cutoff_rollout_max_steps: int = 10_000
opponent_policy: str = "network"
self_play_snapshot_every: int = 1
self_play_max_snapshots: int = 20
self_play_anchor_probability: float = 0.0
strategy_sample_interval: int = 1 strategy_sample_interval: int = 1
store_strategy_on_traverser_nodes: bool = True store_strategy_on_traverser_nodes: bool = True
store_strategy_on_opponent_nodes: bool = True store_strategy_on_opponent_nodes: bool = True
@@ -98,6 +98,7 @@ class DeepCFRTrainer:
self.metrics_path = self.run_dir / "metrics.jsonl" self.metrics_path = self.run_dir / "metrics.jsonl"
self.progress_path = self.run_dir / "runtime_progress.json" self.progress_path = self.run_dir / "runtime_progress.json"
self.log_path = self.run_dir / "train.log" self.log_path = self.run_dir / "train.log"
self.self_play_league_snapshots: list[list[dict]] = []
def checkpoint_payload(self, metrics: IterationMetrics | None = None) -> dict: def checkpoint_payload(self, metrics: IterationMetrics | None = None) -> dict:
return { return {
@@ -108,6 +109,7 @@ class DeepCFRTrainer:
"action_size": self.action_size, "action_size": self.action_size,
"advantage_networks": [network.state_dict() for network in self.advantage_networks], "advantage_networks": [network.state_dict() for network in self.advantage_networks],
"strategy_network": self.strategy_network.state_dict(), "strategy_network": self.strategy_network.state_dict(),
"self_play_league_snapshots": self.self_play_league_snapshots,
"advantage_optimizers": [ "advantage_optimizers": [
optimizer.state_dict() for optimizer in self.advantage_optimizers optimizer.state_dict() for optimizer in self.advantage_optimizers
], ],
@@ -132,6 +134,7 @@ class DeepCFRTrainer:
optimizer.load_state_dict(state_dict) optimizer.load_state_dict(state_dict)
if "strategy_optimizer" in payload: if "strategy_optimizer" in payload:
self.strategy_optimizer.load_state_dict(payload["strategy_optimizer"]) self.strategy_optimizer.load_state_dict(payload["strategy_optimizer"])
self.self_play_league_snapshots = payload.get("self_play_league_snapshots", [])
def run_iteration(self, iteration: int) -> IterationMetrics: def run_iteration(self, iteration: int) -> IterationMetrics:
self.iteration = iteration self.iteration = iteration
@@ -178,6 +181,9 @@ class DeepCFRTrainer:
cutoff_rollouts=self.config.cutoff_rollouts, cutoff_rollouts=self.config.cutoff_rollouts,
cutoff_rollout_policy=self.config.cutoff_rollout_policy, cutoff_rollout_policy=self.config.cutoff_rollout_policy,
cutoff_rollout_max_steps=self.config.cutoff_rollout_max_steps, cutoff_rollout_max_steps=self.config.cutoff_rollout_max_steps,
opponent_policy=self.config.opponent_policy,
league_advantage_networks=self._materialize_league_networks(),
self_play_anchor_probability=self.config.self_play_anchor_probability,
rng=self.rng, rng=self.rng,
) )
for network in self.advantage_networks: for network in self.advantage_networks:
@@ -233,12 +239,49 @@ class DeepCFRTrainer:
input_dim=self.input_dim, input_dim=self.input_dim,
action_size=self.action_size, action_size=self.action_size,
advantage_networks=network_payloads, advantage_networks=network_payloads,
league_advantage_networks=self._league_payloads(),
worker_seed=self.config.seed + iteration * 1_000_003 + batch_index, worker_seed=self.config.seed + iteration * 1_000_003 + batch_index,
) )
) )
batch_index += 1 batch_index += 1
return batches return batches
def _frozen_advantage_state_dicts(self) -> list[dict]:
return [
{name: value.detach().cpu().clone() for name, value in network.state_dict().items()}
for network in self.advantage_networks
]
def _maybe_record_self_play_snapshot(self, iteration: int) -> None:
if self.config.opponent_policy != "self_play_league":
return
if self.config.self_play_max_snapshots <= 0:
return
if iteration % max(1, self.config.self_play_snapshot_every) != 0:
return
self.self_play_league_snapshots.append(self._frozen_advantage_state_dicts())
overflow = len(self.self_play_league_snapshots) - self.config.self_play_max_snapshots
if overflow > 0:
del self.self_play_league_snapshots[:overflow]
def _materialize_league_networks(self) -> list[list[nn.Module]]:
league: list[list[nn.Module]] = []
for snapshot in self.self_play_league_snapshots:
networks = [
DeepCFRMLP(self.input_dim, self.action_size, self.config.hidden_size).to(
self.device
)
for _ in range(2)
]
for network, state_dict in zip(networks, snapshot, strict=True):
network.load_state_dict(state_dict)
network.eval()
league.append(networks)
return league
def _league_payloads(self) -> list[list[dict]]:
return self.self_play_league_snapshots
def train(self) -> list[IterationMetrics]: def train(self) -> list[IterationMetrics]:
self._start_run_logging() self._start_run_logging()
metrics: list[IterationMetrics] = [] metrics: list[IterationMetrics] = []
@@ -250,6 +293,7 @@ class DeepCFRTrainer:
elapsed = time.perf_counter() - started elapsed = time.perf_counter() - started
metrics.append(item) metrics.append(item)
self._append_metrics(item, elapsed) self._append_metrics(item, elapsed)
self._maybe_record_self_play_snapshot(iteration)
if self.config.save_every_iteration: if self.config.save_every_iteration:
checkpoint_dir = self.run_dir checkpoint_dir = self.run_dir
self.save_checkpoint(checkpoint_dir / f"iteration_{iteration:05d}.pt", item) self.save_checkpoint(checkpoint_dir / f"iteration_{iteration:05d}.pt", item)
@@ -95,6 +95,9 @@ class DeepCFRTraverser:
cutoff_rollouts: int = 0, cutoff_rollouts: int = 0,
cutoff_rollout_policy: str = "random", cutoff_rollout_policy: str = "random",
cutoff_rollout_max_steps: int = 10_000, cutoff_rollout_max_steps: int = 10_000,
opponent_policy: str = "network",
league_advantage_networks: list[list[torch.nn.Module]] | None = None,
self_play_anchor_probability: float = 0.0,
rng: np.random.Generator | None = None, rng: np.random.Generator | None = None,
) -> None: ) -> None:
self.advantage_networks = advantage_networks self.advantage_networks = advantage_networks
@@ -125,10 +128,22 @@ class DeepCFRTraverser:
if self.cutoff_rollout_policy not in {"random", "safe_heuristic"}: if self.cutoff_rollout_policy not in {"random", "safe_heuristic"}:
raise ValueError("cutoff_rollout_policy must be 'random' or 'safe_heuristic'") raise ValueError("cutoff_rollout_policy must be 'random' or 'safe_heuristic'")
self.cutoff_rollout_max_steps = max(1, int(cutoff_rollout_max_steps)) self.cutoff_rollout_max_steps = max(1, int(cutoff_rollout_max_steps))
self.opponent_policy = opponent_policy
if self.opponent_policy not in {"network", "safe_heuristic", "self_play_league"}:
raise ValueError(
"opponent_policy must be 'network', 'safe_heuristic', or 'self_play_league'"
)
self.league_advantage_networks = league_advantage_networks or []
self.self_play_anchor_probability = min(1.0, max(0.0, float(self_play_anchor_probability)))
self.rng = rng or np.random.default_rng() self.rng = rng or np.random.default_rng()
self._safe_heuristic_rollout_bot = ( self._safe_heuristic_rollout_bot = (
SafeHeuristicBot() if self.cutoff_rollout_policy == "safe_heuristic" else None SafeHeuristicBot() if self.cutoff_rollout_policy == "safe_heuristic" else None
) )
self._safe_heuristic_opponent_bot = (
SafeHeuristicBot()
if self.opponent_policy == "safe_heuristic" or self.self_play_anchor_probability > 0.0
else None
)
def traverse( def traverse(
self, state: GameState, traverser: int, iteration: int self, state: GameState, traverser: int, iteration: int
@@ -163,6 +178,24 @@ class DeepCFRTraverser:
return self._cutoff_value(state, traverser, stats) return self._cutoff_value(state, traverser, stats)
player = state.current_player player = state.current_player
fixed_action = self._fixed_opponent_action(state, player, traverser)
if fixed_action is not None:
unified_action = state.to_unified_action(fixed_action)
swapped_deck_index = self._sample_deck_draw_chance(state, unified_action)
state.push_action(fixed_action)
try:
return self._traverse(
state,
traverser,
iteration,
depth=depth + 1,
stats=stats,
)
finally:
state.pop_action()
if swapped_deck_index is not None:
state.swap_deck_cards(swapped_deck_index, len(state.deck) - 1)
info_state, legal, policy = self._policy(state, player) info_state, legal, policy = self._policy(state, player)
self._record_strategy(info_state, legal, policy, player, traverser, iteration, depth, stats) self._record_strategy(info_state, legal, policy, player, traverser, iteration, depth, stats)
@@ -224,21 +257,53 @@ class DeepCFRTraverser:
return node_value return node_value
def _policy(self, state: GameState, player: int) -> tuple[np.ndarray, np.ndarray, np.ndarray]: def _policy(self, state: GameState, player: int) -> tuple[np.ndarray, np.ndarray, np.ndarray]:
return self._policy_from_networks(self.advantage_networks, state, player)
def _policy_from_networks(
self,
networks: list[torch.nn.Module],
state: GameState,
player: int,
) -> tuple[np.ndarray, np.ndarray, np.ndarray]:
info_state = encode_info_state(state, player) info_state = encode_info_state(state, player)
legal = np.asarray(state.unified_legal_mask(), dtype=bool) legal = np.asarray(state.unified_legal_mask(), dtype=bool)
with torch.inference_mode(): with torch.inference_mode():
x = torch.as_tensor(info_state, dtype=torch.float32, device=self.device).unsqueeze(0) x = torch.as_tensor(info_state, dtype=torch.float32, device=self.device).unsqueeze(0)
advantages = ( advantages = networks[player](x).squeeze(0).detach().cpu().numpy().astype(np.float32)
self.advantage_networks[player](x)
.squeeze(0)
.detach()
.cpu()
.numpy()
.astype(np.float32)
)
policy = regret_matching(advantages, legal, self.epsilon).astype(np.float32) policy = regret_matching(advantages, legal, self.epsilon).astype(np.float32)
return info_state, legal, policy return info_state, legal, policy
def _fixed_opponent_action(
self,
state: GameState,
player: int,
traverser: int,
) -> int | None:
if player == traverser or self.opponent_policy == "network":
return None
if self.opponent_policy == "safe_heuristic":
if self._safe_heuristic_opponent_bot is None:
self._safe_heuristic_opponent_bot = SafeHeuristicBot()
return self._safe_heuristic_opponent_bot.act(state)
if (
self.self_play_anchor_probability > 0.0
and self.rng.random() < self.self_play_anchor_probability
):
if self._safe_heuristic_opponent_bot is None:
self._safe_heuristic_opponent_bot = SafeHeuristicBot()
return self._safe_heuristic_opponent_bot.act(state)
if not self.league_advantage_networks:
return None
networks = self.league_advantage_networks[
int(self.rng.integers(0, len(self.league_advantage_networks)))
]
_, legal, policy = self._policy_from_networks(networks, state, player)
legal_actions = np.flatnonzero(legal)
if len(legal_actions) == 0:
return None
unified_action = self._sample_action(policy, legal_actions)
return state.from_unified_action(unified_action)
def _sample_action(self, policy: np.ndarray, legal_actions: np.ndarray) -> int: def _sample_action(self, policy: np.ndarray, legal_actions: np.ndarray) -> int:
probs = policy[legal_actions].astype(np.float64) probs = policy[legal_actions].astype(np.float64)
total = float(probs.sum()) total = float(probs.sum())
@@ -23,6 +23,7 @@ class TraversalWorkerBatch:
input_dim: int input_dim: int
action_size: int action_size: int
advantage_networks: list[dict[str, Any]] advantage_networks: list[dict[str, Any]]
league_advantage_networks: list[list[dict[str, Any]]]
worker_seed: int worker_seed: int
@@ -44,6 +45,16 @@ def run_traversal_worker_batch(batch: TraversalWorkerBatch) -> TraversalWorkerRe
for network, state_dict in zip(networks, batch.advantage_networks, strict=True): for network, state_dict in zip(networks, batch.advantage_networks, strict=True):
network.load_state_dict(state_dict) network.load_state_dict(state_dict)
network.eval() network.eval()
league_networks: list[list[torch.nn.Module]] = []
for snapshot in batch.league_advantage_networks:
snapshot_networks = [
DeepCFRMLP(batch.input_dim, batch.action_size, cfg.hidden_size).to(device)
for _ in range(2)
]
for network, state_dict in zip(snapshot_networks, snapshot, strict=True):
network.load_state_dict(state_dict)
network.eval()
league_networks.append(snapshot_networks)
advantage_memory = ReservoirMemory() advantage_memory = ReservoirMemory()
strategy_memory = ReservoirMemory() strategy_memory = ReservoirMemory()
traverser = DeepCFRTraverser( traverser = DeepCFRTraverser(
@@ -65,6 +76,9 @@ def run_traversal_worker_batch(batch: TraversalWorkerBatch) -> TraversalWorkerRe
cutoff_rollouts=cfg.cutoff_rollouts, cutoff_rollouts=cfg.cutoff_rollouts,
cutoff_rollout_policy=cfg.cutoff_rollout_policy, cutoff_rollout_policy=cfg.cutoff_rollout_policy,
cutoff_rollout_max_steps=cfg.cutoff_rollout_max_steps, cutoff_rollout_max_steps=cfg.cutoff_rollout_max_steps,
opponent_policy=cfg.opponent_policy,
league_advantage_networks=league_networks,
self_play_anchor_probability=cfg.self_play_anchor_probability,
rng=np.random.default_rng(batch.worker_seed), rng=np.random.default_rng(batch.worker_seed),
) )
game_config = LostCitiesConfig(**batch.game_config) game_config = LostCitiesConfig(**batch.game_config)
@@ -234,3 +234,29 @@ def test_deep_cfr_traversal_benchmark_smoke() -> None:
assert result["traversal_nodes"] > 0 assert result["traversal_nodes"] > 0
assert result["nodes_per_second"] > 0.0 assert result["nodes_per_second"] > 0.0
def test_deep_cfr_self_play_league_records_snapshots(tmp_path) -> None:
trainer = DeepCFRTrainer(
DeepCFRConfig(
iterations=2,
traversals_per_iteration=1,
max_traversal_depth=2,
max_nodes_per_traversal=32,
batch_size=2,
hidden_size=16,
seed=53,
checkpoint_dir=str(tmp_path / "league"),
save_every_iteration=False,
opponent_policy="self_play_league",
self_play_snapshot_every=1,
self_play_max_snapshots=1,
self_play_anchor_probability=1.0,
),
LostCitiesConfig(seed=53),
)
metrics = trainer.train()
assert len(metrics) == 2
assert len(trainer.self_play_league_snapshots) == 1