Deep CFR weighted self-play league 추가
This commit is contained in:
@@ -23,6 +23,11 @@ class DeepCFRConfig:
|
|||||||
self_play_snapshot_every: int = 1
|
self_play_snapshot_every: int = 1
|
||||||
self_play_max_snapshots: int = 20
|
self_play_max_snapshots: int = 20
|
||||||
self_play_anchor_probability: float = 0.0
|
self_play_anchor_probability: float = 0.0
|
||||||
|
self_play_current_weight: float = 0.5
|
||||||
|
self_play_recent_weight: float = 0.3
|
||||||
|
self_play_older_weight: float = 0.2
|
||||||
|
self_play_anchor_weight: float = 0.0
|
||||||
|
self_play_recent_window: int = 5
|
||||||
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
|
||||||
|
|||||||
@@ -184,6 +184,11 @@ class DeepCFRTrainer:
|
|||||||
opponent_policy=self.config.opponent_policy,
|
opponent_policy=self.config.opponent_policy,
|
||||||
league_advantage_networks=self._materialize_league_networks(),
|
league_advantage_networks=self._materialize_league_networks(),
|
||||||
self_play_anchor_probability=self.config.self_play_anchor_probability,
|
self_play_anchor_probability=self.config.self_play_anchor_probability,
|
||||||
|
self_play_current_weight=self.config.self_play_current_weight,
|
||||||
|
self_play_recent_weight=self.config.self_play_recent_weight,
|
||||||
|
self_play_older_weight=self.config.self_play_older_weight,
|
||||||
|
self_play_anchor_weight=self.config.self_play_anchor_weight,
|
||||||
|
self_play_recent_window=self.config.self_play_recent_window,
|
||||||
rng=self.rng,
|
rng=self.rng,
|
||||||
)
|
)
|
||||||
for network in self.advantage_networks:
|
for network in self.advantage_networks:
|
||||||
|
|||||||
@@ -98,6 +98,11 @@ class DeepCFRTraverser:
|
|||||||
opponent_policy: str = "network",
|
opponent_policy: str = "network",
|
||||||
league_advantage_networks: list[list[torch.nn.Module]] | None = None,
|
league_advantage_networks: list[list[torch.nn.Module]] | None = None,
|
||||||
self_play_anchor_probability: float = 0.0,
|
self_play_anchor_probability: float = 0.0,
|
||||||
|
self_play_current_weight: float = 0.5,
|
||||||
|
self_play_recent_weight: float = 0.3,
|
||||||
|
self_play_older_weight: float = 0.2,
|
||||||
|
self_play_anchor_weight: float = 0.0,
|
||||||
|
self_play_recent_window: int = 5,
|
||||||
rng: np.random.Generator | None = None,
|
rng: np.random.Generator | None = None,
|
||||||
) -> None:
|
) -> None:
|
||||||
self.advantage_networks = advantage_networks
|
self.advantage_networks = advantage_networks
|
||||||
@@ -135,6 +140,11 @@ class DeepCFRTraverser:
|
|||||||
)
|
)
|
||||||
self.league_advantage_networks = league_advantage_networks or []
|
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.self_play_anchor_probability = min(1.0, max(0.0, float(self_play_anchor_probability)))
|
||||||
|
self.self_play_current_weight = max(0.0, float(self_play_current_weight))
|
||||||
|
self.self_play_recent_weight = max(0.0, float(self_play_recent_weight))
|
||||||
|
self.self_play_older_weight = max(0.0, float(self_play_older_weight))
|
||||||
|
self.self_play_anchor_weight = max(0.0, float(self_play_anchor_weight))
|
||||||
|
self.self_play_recent_window = max(0, int(self_play_recent_window))
|
||||||
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
|
||||||
@@ -285,18 +295,16 @@ class DeepCFRTraverser:
|
|||||||
if self._safe_heuristic_opponent_bot is None:
|
if self._safe_heuristic_opponent_bot is None:
|
||||||
self._safe_heuristic_opponent_bot = SafeHeuristicBot()
|
self._safe_heuristic_opponent_bot = SafeHeuristicBot()
|
||||||
return self._safe_heuristic_opponent_bot.act(state)
|
return self._safe_heuristic_opponent_bot.act(state)
|
||||||
if (
|
bucket = self._self_play_bucket()
|
||||||
self.self_play_anchor_probability > 0.0
|
if bucket == "current":
|
||||||
and self.rng.random() < self.self_play_anchor_probability
|
return None
|
||||||
):
|
if bucket == "anchor":
|
||||||
if self._safe_heuristic_opponent_bot is None:
|
if self._safe_heuristic_opponent_bot is None:
|
||||||
self._safe_heuristic_opponent_bot = SafeHeuristicBot()
|
self._safe_heuristic_opponent_bot = SafeHeuristicBot()
|
||||||
return self._safe_heuristic_opponent_bot.act(state)
|
return self._safe_heuristic_opponent_bot.act(state)
|
||||||
if not self.league_advantage_networks:
|
networks = self._self_play_snapshot_networks(bucket)
|
||||||
|
if networks is None:
|
||||||
return None
|
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, policy = self._policy_from_networks(networks, state, player)
|
||||||
legal_actions = np.flatnonzero(legal)
|
legal_actions = np.flatnonzero(legal)
|
||||||
if len(legal_actions) == 0:
|
if len(legal_actions) == 0:
|
||||||
@@ -304,6 +312,46 @@ class DeepCFRTraverser:
|
|||||||
unified_action = self._sample_action(policy, legal_actions)
|
unified_action = self._sample_action(policy, legal_actions)
|
||||||
return state.from_unified_action(unified_action)
|
return state.from_unified_action(unified_action)
|
||||||
|
|
||||||
|
def _self_play_bucket(self) -> str:
|
||||||
|
if (
|
||||||
|
self.self_play_anchor_probability > 0.0
|
||||||
|
and self.rng.random() < self.self_play_anchor_probability
|
||||||
|
):
|
||||||
|
return "anchor"
|
||||||
|
recent_count = min(len(self.league_advantage_networks), self.self_play_recent_window)
|
||||||
|
older_count = max(0, len(self.league_advantage_networks) - recent_count)
|
||||||
|
labels = ["current", "recent", "older", "anchor"]
|
||||||
|
weights = np.asarray(
|
||||||
|
[
|
||||||
|
self.self_play_current_weight,
|
||||||
|
self.self_play_recent_weight if recent_count > 0 else 0.0,
|
||||||
|
self.self_play_older_weight if older_count > 0 else 0.0,
|
||||||
|
self.self_play_anchor_weight,
|
||||||
|
],
|
||||||
|
dtype=np.float64,
|
||||||
|
)
|
||||||
|
total = float(weights.sum())
|
||||||
|
if total <= 0.0:
|
||||||
|
return "current"
|
||||||
|
weights /= total
|
||||||
|
return str(self.rng.choice(labels, p=weights))
|
||||||
|
|
||||||
|
def _self_play_snapshot_networks(self, bucket: str) -> list[torch.nn.Module] | None:
|
||||||
|
if not self.league_advantage_networks:
|
||||||
|
return None
|
||||||
|
recent_count = min(len(self.league_advantage_networks), self.self_play_recent_window)
|
||||||
|
if bucket == "recent" and recent_count > 0:
|
||||||
|
candidates = self.league_advantage_networks[-recent_count:]
|
||||||
|
elif bucket == "older":
|
||||||
|
candidates = self.league_advantage_networks[
|
||||||
|
: max(0, len(self.league_advantage_networks) - recent_count)
|
||||||
|
]
|
||||||
|
else:
|
||||||
|
candidates = self.league_advantage_networks
|
||||||
|
if not candidates:
|
||||||
|
return None
|
||||||
|
return candidates[int(self.rng.integers(0, len(candidates)))]
|
||||||
|
|
||||||
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())
|
||||||
|
|||||||
@@ -79,6 +79,11 @@ def run_traversal_worker_batch(batch: TraversalWorkerBatch) -> TraversalWorkerRe
|
|||||||
opponent_policy=cfg.opponent_policy,
|
opponent_policy=cfg.opponent_policy,
|
||||||
league_advantage_networks=league_networks,
|
league_advantage_networks=league_networks,
|
||||||
self_play_anchor_probability=cfg.self_play_anchor_probability,
|
self_play_anchor_probability=cfg.self_play_anchor_probability,
|
||||||
|
self_play_current_weight=cfg.self_play_current_weight,
|
||||||
|
self_play_recent_weight=cfg.self_play_recent_weight,
|
||||||
|
self_play_older_weight=cfg.self_play_older_weight,
|
||||||
|
self_play_anchor_weight=cfg.self_play_anchor_weight,
|
||||||
|
self_play_recent_window=cfg.self_play_recent_window,
|
||||||
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)
|
||||||
|
|||||||
@@ -260,3 +260,34 @@ def test_deep_cfr_self_play_league_records_snapshots(tmp_path) -> None:
|
|||||||
|
|
||||||
assert len(metrics) == 2
|
assert len(metrics) == 2
|
||||||
assert len(trainer.self_play_league_snapshots) == 1
|
assert len(trainer.self_play_league_snapshots) == 1
|
||||||
|
|
||||||
|
|
||||||
|
def test_deep_cfr_weighted_self_play_league_uses_snapshot_bucket(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=59,
|
||||||
|
checkpoint_dir=str(tmp_path / "weighted-league"),
|
||||||
|
save_every_iteration=False,
|
||||||
|
opponent_policy="self_play_league",
|
||||||
|
self_play_snapshot_every=1,
|
||||||
|
self_play_max_snapshots=2,
|
||||||
|
self_play_current_weight=0.0,
|
||||||
|
self_play_recent_weight=1.0,
|
||||||
|
self_play_older_weight=0.0,
|
||||||
|
self_play_anchor_weight=0.0,
|
||||||
|
self_play_recent_window=1,
|
||||||
|
),
|
||||||
|
LostCitiesConfig(seed=59),
|
||||||
|
)
|
||||||
|
|
||||||
|
metrics = trainer.train()
|
||||||
|
|
||||||
|
assert len(metrics) == 2
|
||||||
|
assert len(trainer.self_play_league_snapshots) == 2
|
||||||
|
assert metrics[1].traversal_nodes > 0
|
||||||
|
|||||||
Reference in New Issue
Block a user