Deep CFR self-play league 추가
This commit is contained in:
@@ -19,6 +19,10 @@ class DeepCFRConfig:
|
||||
cutoff_rollouts: int = 0
|
||||
cutoff_rollout_policy: str = "random"
|
||||
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
|
||||
store_strategy_on_traverser_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.progress_path = self.run_dir / "runtime_progress.json"
|
||||
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:
|
||||
return {
|
||||
@@ -108,6 +109,7 @@ class DeepCFRTrainer:
|
||||
"action_size": self.action_size,
|
||||
"advantage_networks": [network.state_dict() for network in self.advantage_networks],
|
||||
"strategy_network": self.strategy_network.state_dict(),
|
||||
"self_play_league_snapshots": self.self_play_league_snapshots,
|
||||
"advantage_optimizers": [
|
||||
optimizer.state_dict() for optimizer in self.advantage_optimizers
|
||||
],
|
||||
@@ -132,6 +134,7 @@ class DeepCFRTrainer:
|
||||
optimizer.load_state_dict(state_dict)
|
||||
if "strategy_optimizer" in payload:
|
||||
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:
|
||||
self.iteration = iteration
|
||||
@@ -178,6 +181,9 @@ class DeepCFRTrainer:
|
||||
cutoff_rollouts=self.config.cutoff_rollouts,
|
||||
cutoff_rollout_policy=self.config.cutoff_rollout_policy,
|
||||
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,
|
||||
)
|
||||
for network in self.advantage_networks:
|
||||
@@ -233,12 +239,49 @@ class DeepCFRTrainer:
|
||||
input_dim=self.input_dim,
|
||||
action_size=self.action_size,
|
||||
advantage_networks=network_payloads,
|
||||
league_advantage_networks=self._league_payloads(),
|
||||
worker_seed=self.config.seed + iteration * 1_000_003 + batch_index,
|
||||
)
|
||||
)
|
||||
batch_index += 1
|
||||
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]:
|
||||
self._start_run_logging()
|
||||
metrics: list[IterationMetrics] = []
|
||||
@@ -250,6 +293,7 @@ class DeepCFRTrainer:
|
||||
elapsed = time.perf_counter() - started
|
||||
metrics.append(item)
|
||||
self._append_metrics(item, elapsed)
|
||||
self._maybe_record_self_play_snapshot(iteration)
|
||||
if self.config.save_every_iteration:
|
||||
checkpoint_dir = self.run_dir
|
||||
self.save_checkpoint(checkpoint_dir / f"iteration_{iteration:05d}.pt", item)
|
||||
|
||||
@@ -95,6 +95,9 @@ class DeepCFRTraverser:
|
||||
cutoff_rollouts: int = 0,
|
||||
cutoff_rollout_policy: str = "random",
|
||||
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,
|
||||
) -> None:
|
||||
self.advantage_networks = advantage_networks
|
||||
@@ -125,10 +128,22 @@ class DeepCFRTraverser:
|
||||
if self.cutoff_rollout_policy not in {"random", "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.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._safe_heuristic_rollout_bot = (
|
||||
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(
|
||||
self, state: GameState, traverser: int, iteration: int
|
||||
@@ -163,6 +178,24 @@ class DeepCFRTraverser:
|
||||
return self._cutoff_value(state, traverser, stats)
|
||||
|
||||
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)
|
||||
self._record_strategy(info_state, legal, policy, player, traverser, iteration, depth, stats)
|
||||
|
||||
@@ -224,21 +257,53 @@ class DeepCFRTraverser:
|
||||
return node_value
|
||||
|
||||
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)
|
||||
legal = np.asarray(state.unified_legal_mask(), dtype=bool)
|
||||
with torch.inference_mode():
|
||||
x = torch.as_tensor(info_state, dtype=torch.float32, device=self.device).unsqueeze(0)
|
||||
advantages = (
|
||||
self.advantage_networks[player](x)
|
||||
.squeeze(0)
|
||||
.detach()
|
||||
.cpu()
|
||||
.numpy()
|
||||
.astype(np.float32)
|
||||
)
|
||||
advantages = networks[player](x).squeeze(0).detach().cpu().numpy().astype(np.float32)
|
||||
policy = regret_matching(advantages, legal, self.epsilon).astype(np.float32)
|
||||
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:
|
||||
probs = policy[legal_actions].astype(np.float64)
|
||||
total = float(probs.sum())
|
||||
|
||||
@@ -23,6 +23,7 @@ class TraversalWorkerBatch:
|
||||
input_dim: int
|
||||
action_size: int
|
||||
advantage_networks: list[dict[str, Any]]
|
||||
league_advantage_networks: list[list[dict[str, Any]]]
|
||||
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):
|
||||
network.load_state_dict(state_dict)
|
||||
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()
|
||||
strategy_memory = ReservoirMemory()
|
||||
traverser = DeepCFRTraverser(
|
||||
@@ -65,6 +76,9 @@ def run_traversal_worker_batch(batch: TraversalWorkerBatch) -> TraversalWorkerRe
|
||||
cutoff_rollouts=cfg.cutoff_rollouts,
|
||||
cutoff_rollout_policy=cfg.cutoff_rollout_policy,
|
||||
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),
|
||||
)
|
||||
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["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
|
||||
|
||||
Reference in New Issue
Block a user