Deep CFR self-play league 추가
This commit is contained in:
@@ -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
|
||||||
|
|||||||
Reference in New Issue
Block a user