Deep CFR regret fallback audit metrics 추가
This commit is contained in:
@@ -0,0 +1,103 @@
|
||||
# Deep CFR Regret Fallback Audit, 2026-05-07
|
||||
|
||||
Goal: test whether all-negative regret matching fallback is a plausible source of
|
||||
early over-opening in Lost Cities Deep CFR.
|
||||
|
||||
## Code Changes
|
||||
|
||||
- Added traversal audit metrics for regret matching fallback decisions.
|
||||
- Default behavior remains unchanged: `regret_matching.all_negative_fallback: uniform`.
|
||||
- Added optional fallback mode: `argmax_tiebreak`.
|
||||
- Added CLI override: `--regret-fallback uniform|argmax_tiebreak`.
|
||||
|
||||
Key metrics:
|
||||
|
||||
- `traversal_regret_matching_decisions`
|
||||
- `traversal_regret_fallback_count`
|
||||
- `traversal_regret_fallback_rate`
|
||||
- `traversal_regret_fallback_avg_depth`
|
||||
- `traversal_regret_fallback_depth_bucket_<range>`
|
||||
- `traversal_regret_fallback_opened_colors_count_<n>`
|
||||
- `traversal_regret_fallback_action_open_new`
|
||||
- `traversal_regret_fallback_open_new_selected`
|
||||
- `traversal_regret_fallback_open_new_selected_rate`
|
||||
- `traversal_regret_fallback_legal_actions_mean`
|
||||
- `traversal_regret_fallback_legal_open_new_mean`
|
||||
- `traversal_regret_fallback_legal_discard_mean`
|
||||
- `traversal_regret_fallback_legal_draw_deck_mean`
|
||||
- `traversal_regret_fallback_legal_draw_pile_mean`
|
||||
- `traversal_regret_fallback_open_new_available_rate`
|
||||
- `traversal_regret_fallback_open_new_selection_over_availability`
|
||||
- `traversal_regret_fallback_avg_opened_colors_before_action`
|
||||
- `traversal_regret_fallback_argmax_tie_rate`
|
||||
- `traversal_regret_fallback_argmax_tie_size_mean`
|
||||
- `traversal_regret_fallback_argmax_full_tie_rate`
|
||||
- `traversal_regret_fallback_open_new_available_color_<color>`
|
||||
- `traversal_regret_fallback_open_new_selected_color_<color>`
|
||||
|
||||
Implementation note: fallback policy state is captured immediately after the
|
||||
network policy is computed. This avoids child recursion overwriting the
|
||||
traversal-level fallback flag before the decision is recorded.
|
||||
|
||||
## Runs
|
||||
|
||||
Baseline long run, analyzed after it had reached iteration 210:
|
||||
|
||||
- `runs/deep_cfr/2026-05-07_legacy_align_full_depth_slot_playability`
|
||||
|
||||
Short comparison runs:
|
||||
|
||||
- `runs/deep_cfr/2026-05-07_regret_fallback_uniform_20iter`
|
||||
- `runs/deep_cfr/2026-05-07_regret_fallback_argmax_tiebreak_20iter`
|
||||
|
||||
Both short runs used the same base config and seed:
|
||||
|
||||
- `configs/deep_cfr/deep_cfr_selfplay_full_depth_slot_playability.yaml`
|
||||
- `seed: 79`
|
||||
- `iterations: 20`
|
||||
- `save_latest_only`
|
||||
|
||||
The 20-iteration runs below were collected before the expanded fallback timing,
|
||||
legal-action composition, and tie diagnostics were added. They should be treated
|
||||
as the first historical audit snapshot. New paired runs are needed to compare
|
||||
the expanded metrics.
|
||||
|
||||
Instrumentation smoke run:
|
||||
|
||||
- `runs/deep_cfr/2026-05-07_regret_fallback_metrics_smoke_1iter_v2`
|
||||
|
||||
This run confirms the expanded metrics are emitted to `metrics.jsonl`.
|
||||
|
||||
## Iteration 20 Snapshot
|
||||
|
||||
| metric | uniform | argmax_tiebreak |
|
||||
|---|---:|---:|
|
||||
| traversal_regret_matching_decisions | 39,358 | 49,296 |
|
||||
| traversal_regret_fallback_count | 18,362 | 7,588 |
|
||||
| traversal_regret_fallback_rate | 0.4665 | 0.1539 |
|
||||
| traversal_regret_fallback_open_new_selected | 641 | 164 |
|
||||
| traversal_regret_fallback_open_new_selected_rate | 0.0349 | 0.0216 |
|
||||
| traversal_regret_fallback_avg_opened_colors_before_action | 4.4325 | 4.6118 |
|
||||
| eval_random_avg_opened_colors | 2.48 | 2.16 |
|
||||
| eval_random_5_color_open_count | 49 | 35 |
|
||||
| eval_safe_heuristic_avg_opened_colors | 2.50 | 2.44 |
|
||||
| eval_safe_heuristic_5_color_open_count | 50 | 44 |
|
||||
| eval_passive_discard_avg_opened_colors | 2.33 | 2.32 |
|
||||
| eval_passive_discard_5_color_open_count | 39 | 41 |
|
||||
| eval_random_avg_score_diff0 | 42.46 | 33.58 |
|
||||
| eval_safe_heuristic_avg_score_diff0 | -52.65 | -58.25 |
|
||||
|
||||
## Read
|
||||
|
||||
The audit confirms that uniform fallback fires frequently in the early run.
|
||||
At iteration 20, almost half of traversal regret-matching decisions use fallback
|
||||
under `uniform`.
|
||||
|
||||
`argmax_tiebreak` sharply reduces fallback frequency and absolute fallback
|
||||
open-new selections in this 20-iteration comparison. It also lowers 5-color
|
||||
counts against random and safe heuristic opponents at iteration 20. The effect is
|
||||
not uniform across every opponent in this very short run.
|
||||
|
||||
This is diagnostic evidence, not enough to promote `argmax_tiebreak` as the
|
||||
default. A longer 50-100 iteration paired run is still needed before deciding
|
||||
whether this fixes the plateau without hurting policy quality.
|
||||
@@ -86,6 +86,8 @@ def _train_overrides_from_args(args: argparse.Namespace) -> dict[str, Any]:
|
||||
overrides.setdefault("evaluation", {})["eval_every"] = args.eval_every
|
||||
if args.eval_games is not None:
|
||||
overrides.setdefault("evaluation", {})["games"] = args.eval_games
|
||||
if args.regret_fallback is not None:
|
||||
overrides.setdefault("regret_matching", {})["all_negative_fallback"] = args.regret_fallback
|
||||
if args.no_save:
|
||||
checkpoint_overrides = overrides.setdefault("checkpoint", {})
|
||||
checkpoint_overrides["save_latest"] = False
|
||||
@@ -218,6 +220,11 @@ def main(argv: list[str] | None = None) -> None:
|
||||
train.add_argument("--device")
|
||||
train.add_argument("--eval-every", type=int)
|
||||
train.add_argument("--eval-games", type=int)
|
||||
train.add_argument(
|
||||
"--regret-fallback",
|
||||
choices=("uniform", "argmax_tiebreak"),
|
||||
help="Override regret_matching.all_negative_fallback.",
|
||||
)
|
||||
train.add_argument("--seed", type=int)
|
||||
train.add_argument("--no-save", action="store_true")
|
||||
train.add_argument("--save-latest-only", action="store_true")
|
||||
|
||||
@@ -163,6 +163,18 @@ class TraversalConfig(StrictModel):
|
||||
return max(1, int(self.worker_chunk_size))
|
||||
|
||||
|
||||
class RegretMatchingConfig(StrictModel):
|
||||
all_negative_fallback: str = "uniform"
|
||||
|
||||
@field_validator("all_negative_fallback")
|
||||
@classmethod
|
||||
def _validate_all_negative_fallback(cls, value: str) -> str:
|
||||
token = value.strip().lower()
|
||||
if token not in {"uniform", "argmax_tiebreak"}:
|
||||
raise ValueError("must be 'uniform' or 'argmax_tiebreak'")
|
||||
return token
|
||||
|
||||
|
||||
class SelfPlayLeagueConfig(StrictModel):
|
||||
snapshot_every: int = 1
|
||||
max_snapshots: int = 20
|
||||
@@ -268,6 +280,7 @@ class DeepCFRConfig(StrictModel):
|
||||
encoding: EncodingConfig = Field(default_factory=EncodingConfig)
|
||||
network: NetworkConfig = Field(default_factory=NetworkConfig)
|
||||
traversal: TraversalConfig = Field(default_factory=TraversalConfig)
|
||||
regret_matching: RegretMatchingConfig = Field(default_factory=RegretMatchingConfig)
|
||||
self_play: SelfPlayLeagueConfig = Field(default_factory=SelfPlayLeagueConfig)
|
||||
optimization: OptimizationConfig = Field(default_factory=OptimizationConfig)
|
||||
memory: MemoryConfig = Field(default_factory=MemoryConfig)
|
||||
|
||||
@@ -294,6 +294,9 @@ class DeepCFRTrainer:
|
||||
for player, memory in enumerate(self.advantage_memories):
|
||||
self._runtime_metrics[f"advantage_player_{player}_memory_size"] = len(memory)
|
||||
self._runtime_metrics["strategy_memory_size"] = len(self.strategy_memory)
|
||||
for key, value in total_stats.to_dict().items():
|
||||
if key.startswith("traversal_regret_") or key == "traversal_sampled_actions":
|
||||
self._runtime_metrics[key] = value
|
||||
return IterationMetrics(
|
||||
iteration=iteration,
|
||||
advantage_samples=self._advantage_memory_size(),
|
||||
@@ -354,6 +357,7 @@ class DeepCFRTrainer:
|
||||
cutoff_rollout_policy=self.config.traversal.cutoff_rollout_policy,
|
||||
cutoff_rollout_max_steps=self.config.traversal.cutoff_rollout_max_steps,
|
||||
opponent_policy=self.config.traversal.opponent_policy,
|
||||
all_negative_fallback=self.config.regret_matching.all_negative_fallback,
|
||||
league_advantage_networks=league_networks,
|
||||
self_play_anchor_probability=self.config.self_play.anchor_probability,
|
||||
self_play_current_weight=self.config.self_play.current_weight,
|
||||
|
||||
@@ -88,6 +88,10 @@ cdef class CythonDeepCFRTraverser:
|
||||
cdef int cutoff_rollouts
|
||||
cdef int cutoff_rollout_max_steps
|
||||
cdef int opponent_policy_id
|
||||
cdef int all_negative_fallback_id
|
||||
cdef bint last_policy_regret_fallback
|
||||
cdef int last_policy_argmax_tie_size
|
||||
cdef bint last_policy_argmax_full_tie
|
||||
cdef float self_play_anchor_probability
|
||||
cdef float self_play_current_weight
|
||||
cdef float self_play_recent_weight
|
||||
@@ -123,6 +127,7 @@ cdef class CythonDeepCFRTraverser:
|
||||
str cutoff_rollout_policy="random",
|
||||
int cutoff_rollout_max_steps=10000,
|
||||
str opponent_policy="network",
|
||||
str all_negative_fallback="uniform",
|
||||
object league_advantage_networks=None,
|
||||
float self_play_anchor_probability=0.0,
|
||||
float self_play_current_weight=0.5,
|
||||
@@ -182,6 +187,15 @@ cdef class CythonDeepCFRTraverser:
|
||||
self.opponent_policy_id = 2
|
||||
else:
|
||||
raise ValueError("opponent_policy must be 'network', 'safe_heuristic', or 'self_play_league'")
|
||||
if all_negative_fallback == "uniform":
|
||||
self.all_negative_fallback_id = 0
|
||||
elif all_negative_fallback == "argmax_tiebreak":
|
||||
self.all_negative_fallback_id = 1
|
||||
else:
|
||||
raise ValueError("all_negative_fallback must be 'uniform' or 'argmax_tiebreak'")
|
||||
self.last_policy_regret_fallback = False
|
||||
self.last_policy_argmax_tie_size = 0
|
||||
self.last_policy_argmax_full_tie = False
|
||||
self.league_advantage_networks = [] if league_advantage_networks is None else league_advantage_networks
|
||||
self.self_play_anchor_probability = min(1.0, max(0.0, self_play_anchor_probability))
|
||||
self.self_play_current_weight = max(0.0, self_play_current_weight)
|
||||
@@ -247,6 +261,9 @@ cdef class CythonDeepCFRTraverser:
|
||||
cdef float sampling_policy[MAX_ACTIONS]
|
||||
cdef unsigned char legal[MAX_ACTIONS]
|
||||
cdef object info_state
|
||||
cdef bint policy_regret_fallback
|
||||
cdef int policy_argmax_tie_size
|
||||
cdef bint policy_argmax_full_tie
|
||||
|
||||
stats.nodes += 1
|
||||
if depth > stats.max_depth_reached:
|
||||
@@ -266,7 +283,7 @@ cdef class CythonDeepCFRTraverser:
|
||||
return self._cutoff_value(state, traverser, stats)
|
||||
|
||||
player = state.current_player
|
||||
fixed_action = self._fixed_opponent_action(state, player, traverser)
|
||||
fixed_action = self._fixed_opponent_action(state, player, traverser, depth, stats)
|
||||
if fixed_action >= 0:
|
||||
fixed_unified_action = self._to_unified_action_c(state, fixed_action)
|
||||
swapped_deck_index = self._sample_deck_draw_chance(state, fixed_unified_action)
|
||||
@@ -279,6 +296,9 @@ cdef class CythonDeepCFRTraverser:
|
||||
state._swap_deck_cards_c(swapped_deck_index, state.deck_len - 1)
|
||||
|
||||
info_state = self._policy(state, player, legal, policy)
|
||||
policy_regret_fallback = self.last_policy_regret_fallback
|
||||
policy_argmax_tie_size = self.last_policy_argmax_tie_size
|
||||
policy_argmax_full_tie = self.last_policy_argmax_full_tie
|
||||
self._record_strategy(info_state, legal, policy, player, traverser, iteration, depth, stats)
|
||||
|
||||
legal_count = 0
|
||||
@@ -304,6 +324,16 @@ cdef class CythonDeepCFRTraverser:
|
||||
state._swap_deck_cards_c(swapped_deck_index, state.deck_len - 1)
|
||||
|
||||
stats.sampled_actions += 1
|
||||
self._record_regret_matching_decision(
|
||||
stats,
|
||||
state,
|
||||
player,
|
||||
action,
|
||||
depth,
|
||||
policy_regret_fallback,
|
||||
policy_argmax_tie_size,
|
||||
policy_argmax_full_tie,
|
||||
)
|
||||
action_prob = sampling_policy[action]
|
||||
if action_prob < self.epsilon:
|
||||
action_prob = self.epsilon
|
||||
@@ -350,6 +380,12 @@ cdef class CythonDeepCFRTraverser:
|
||||
cdef int actions[MAX_ACTIONS]
|
||||
cdef int action_count
|
||||
cdef int i
|
||||
cdef int selected
|
||||
cdef int tie_count
|
||||
cdef int max_tie_count
|
||||
cdef float positive
|
||||
cdef float positive_sum = 0.0
|
||||
cdef float best = 0.0
|
||||
|
||||
info_state = np.empty(self.input_dim, dtype=np.float32)
|
||||
info_view = info_state
|
||||
@@ -369,7 +405,46 @@ cdef class CythonDeepCFRTraverser:
|
||||
x = torch.as_tensor(info_state, dtype=torch.float32, device=self.device).unsqueeze(0)
|
||||
advantages = networks[player](x).squeeze(0).detach().cpu().numpy().astype(np.float32)
|
||||
adv_view = advantages
|
||||
regret_matching_c(&adv_view[0], legal, self.action_size, self.epsilon, policy)
|
||||
for i in range(self.action_size):
|
||||
if legal[i] != 0:
|
||||
positive = adv_view[i] if adv_view[i] > 0.0 else 0.0
|
||||
positive_sum += positive
|
||||
self.last_policy_regret_fallback = positive_sum <= self.epsilon
|
||||
self.last_policy_argmax_tie_size = 0
|
||||
self.last_policy_argmax_full_tie = False
|
||||
if not self.last_policy_regret_fallback or self.all_negative_fallback_id == 0:
|
||||
if self.last_policy_regret_fallback:
|
||||
max_tie_count = 0
|
||||
for i in range(self.action_size):
|
||||
if legal[i] == 0:
|
||||
continue
|
||||
if max_tie_count == 0 or adv_view[i] > best:
|
||||
best = adv_view[i]
|
||||
max_tie_count = 1
|
||||
elif adv_view[i] == best:
|
||||
max_tie_count += 1
|
||||
self.last_policy_argmax_tie_size = max_tie_count
|
||||
self.last_policy_argmax_full_tie = max_tie_count > 1 and max_tie_count == action_count
|
||||
regret_matching_c(&adv_view[0], legal, self.action_size, self.epsilon, policy)
|
||||
else:
|
||||
selected = -1
|
||||
tie_count = 0
|
||||
for i in range(self.action_size):
|
||||
policy[i] = 0.0
|
||||
if legal[i] == 0:
|
||||
continue
|
||||
if selected < 0 or adv_view[i] > best:
|
||||
selected = i
|
||||
best = adv_view[i]
|
||||
tie_count = 1
|
||||
elif adv_view[i] == best:
|
||||
tie_count += 1
|
||||
if _next_u32(&self.rng) % <unsigned int>tie_count == 0:
|
||||
selected = i
|
||||
self.last_policy_argmax_tie_size = tie_count
|
||||
self.last_policy_argmax_full_tie = tie_count > 1 and tie_count == action_count
|
||||
if selected >= 0:
|
||||
policy[selected] = 1.0
|
||||
return info_state
|
||||
|
||||
cdef void _sampling_policy(
|
||||
@@ -402,7 +477,14 @@ cdef class CythonDeepCFRTraverser:
|
||||
else:
|
||||
out_policy[i] = 0.0
|
||||
|
||||
cdef int _fixed_opponent_action(self, GameState state, int player, int traverser) except *:
|
||||
cdef int _fixed_opponent_action(
|
||||
self,
|
||||
GameState state,
|
||||
int player,
|
||||
int traverser,
|
||||
int depth,
|
||||
object stats,
|
||||
) except *:
|
||||
cdef int bucket
|
||||
cdef object networks
|
||||
cdef unsigned char legal[MAX_ACTIONS]
|
||||
@@ -435,8 +517,116 @@ cdef class CythonDeepCFRTraverser:
|
||||
if count <= 0:
|
||||
return -1
|
||||
unified_action = _sample_policy_from_actions_c(policy, actions, count, _next_double(&self.rng))
|
||||
self._record_regret_matching_decision(
|
||||
stats,
|
||||
state,
|
||||
player,
|
||||
unified_action,
|
||||
depth,
|
||||
self.last_policy_regret_fallback,
|
||||
self.last_policy_argmax_tie_size,
|
||||
self.last_policy_argmax_full_tie,
|
||||
)
|
||||
return self._from_unified_action_c(state, unified_action)
|
||||
|
||||
cdef void _record_regret_matching_decision(
|
||||
self,
|
||||
object stats,
|
||||
GameState state,
|
||||
int player,
|
||||
int unified_action,
|
||||
int depth,
|
||||
bint fallback,
|
||||
int argmax_tie_size,
|
||||
bint argmax_full_tie,
|
||||
):
|
||||
cdef int card_action_size = 2 * state.hand_size
|
||||
cdef int actions[MAX_ACTIONS]
|
||||
cdef int legal_count
|
||||
cdef int legal_action
|
||||
cdef int opened_colors
|
||||
cdef int slot
|
||||
cdef int card
|
||||
cdef int color
|
||||
stats.regret_matching_decisions += 1
|
||||
if not fallback:
|
||||
return
|
||||
stats.regret_fallback_count += 1
|
||||
stats.regret_fallback_depth_sum += depth
|
||||
self._record_fallback_depth_bucket(stats, depth)
|
||||
opened_colors = self._opened_color_count(state, player)
|
||||
stats.regret_fallback_opened_colors_sum += opened_colors
|
||||
stats.regret_fallback_opened_colors_buckets[opened_colors] = (
|
||||
stats.regret_fallback_opened_colors_buckets.get(opened_colors, 0) + 1
|
||||
)
|
||||
if argmax_tie_size > 1:
|
||||
stats.regret_fallback_argmax_tie_count += 1
|
||||
stats.regret_fallback_argmax_tie_size_sum += argmax_tie_size
|
||||
if argmax_full_tie:
|
||||
stats.regret_fallback_argmax_full_tie_count += 1
|
||||
legal_count = state._unified_legal_actions_c(actions)
|
||||
stats.regret_fallback_legal_actions_sum += legal_count
|
||||
for slot in range(legal_count):
|
||||
legal_action = actions[slot]
|
||||
if legal_action < card_action_size:
|
||||
if legal_action % 2 == 1:
|
||||
stats.regret_fallback_legal_discard_sum += 1
|
||||
continue
|
||||
card = state.hand_cards[state._hand_index(player, legal_action // 2)]
|
||||
color = state._card_color(card)
|
||||
if state.expedition_lens[state._expedition_len_index(player, color)] == 0:
|
||||
stats.regret_fallback_legal_open_new_sum += 1
|
||||
stats.regret_fallback_open_new_available_by_color[color] = (
|
||||
stats.regret_fallback_open_new_available_by_color.get(color, 0) + 1
|
||||
)
|
||||
else:
|
||||
stats.regret_fallback_legal_play_existing_sum += 1
|
||||
continue
|
||||
if legal_action == card_action_size:
|
||||
stats.regret_fallback_legal_draw_deck_sum += 1
|
||||
else:
|
||||
stats.regret_fallback_legal_draw_pile_sum += 1
|
||||
if unified_action < card_action_size:
|
||||
if unified_action % 2 == 1:
|
||||
stats.regret_fallback_action_discard += 1
|
||||
return
|
||||
slot = unified_action // 2
|
||||
card = state.hand_cards[state._hand_index(player, slot)]
|
||||
color = state._card_color(card)
|
||||
if state.expedition_lens[state._expedition_len_index(player, color)] == 0:
|
||||
stats.regret_fallback_action_open_new += 1
|
||||
stats.regret_fallback_open_new_selected_by_color[color] = (
|
||||
stats.regret_fallback_open_new_selected_by_color.get(color, 0) + 1
|
||||
)
|
||||
else:
|
||||
stats.regret_fallback_action_play_existing += 1
|
||||
return
|
||||
if unified_action == card_action_size:
|
||||
stats.regret_fallback_action_draw_deck += 1
|
||||
else:
|
||||
stats.regret_fallback_action_draw_pile += 1
|
||||
|
||||
cdef void _record_fallback_depth_bucket(self, object stats, int depth):
|
||||
cdef int width = 50
|
||||
cdef int max_depth = 400
|
||||
cdef int start = (depth // width) * width
|
||||
cdef str key
|
||||
if start >= max_depth:
|
||||
key = f"{max_depth}_plus"
|
||||
else:
|
||||
key = f"{start}_{start + width - 1}"
|
||||
stats.regret_fallback_depth_buckets[key] = (
|
||||
stats.regret_fallback_depth_buckets.get(key, 0) + 1
|
||||
)
|
||||
|
||||
cdef int _opened_color_count(self, GameState state, int player) noexcept:
|
||||
cdef int color
|
||||
cdef int count = 0
|
||||
for color in range(state.n_colors):
|
||||
if state.expedition_lens[state._expedition_len_index(player, color)] > 0:
|
||||
count += 1
|
||||
return count
|
||||
|
||||
cdef int _self_play_bucket(self) noexcept:
|
||||
cdef int recent_count
|
||||
cdef int older_count
|
||||
@@ -665,6 +855,7 @@ def run_cython_traversal_batch(
|
||||
str cutoff_rollout_policy="random",
|
||||
int cutoff_rollout_max_steps=10000,
|
||||
str opponent_policy="network",
|
||||
str all_negative_fallback="uniform",
|
||||
object league_advantage_networks=None,
|
||||
float self_play_anchor_probability=0.0,
|
||||
float self_play_current_weight=0.5,
|
||||
@@ -700,6 +891,7 @@ def run_cython_traversal_batch(
|
||||
cutoff_rollout_policy=cutoff_rollout_policy,
|
||||
cutoff_rollout_max_steps=cutoff_rollout_max_steps,
|
||||
opponent_policy=opponent_policy,
|
||||
all_negative_fallback=all_negative_fallback,
|
||||
league_advantage_networks=league_advantage_networks,
|
||||
self_play_anchor_probability=self_play_anchor_probability,
|
||||
self_play_current_weight=self_play_current_weight,
|
||||
|
||||
@@ -13,6 +13,28 @@ class TraversalStats:
|
||||
advantage_samples: int = 0
|
||||
strategy_samples: int = 0
|
||||
sampled_actions: int = 0
|
||||
regret_matching_decisions: int = 0
|
||||
regret_fallback_count: int = 0
|
||||
regret_fallback_depth_sum: int = 0
|
||||
regret_fallback_opened_colors_sum: int = 0
|
||||
regret_fallback_legal_actions_sum: int = 0
|
||||
regret_fallback_legal_play_existing_sum: int = 0
|
||||
regret_fallback_legal_open_new_sum: int = 0
|
||||
regret_fallback_legal_discard_sum: int = 0
|
||||
regret_fallback_legal_draw_deck_sum: int = 0
|
||||
regret_fallback_legal_draw_pile_sum: int = 0
|
||||
regret_fallback_action_play_existing: int = 0
|
||||
regret_fallback_action_open_new: int = 0
|
||||
regret_fallback_action_discard: int = 0
|
||||
regret_fallback_action_draw_deck: int = 0
|
||||
regret_fallback_action_draw_pile: int = 0
|
||||
regret_fallback_argmax_tie_count: int = 0
|
||||
regret_fallback_argmax_tie_size_sum: int = 0
|
||||
regret_fallback_argmax_full_tie_count: int = 0
|
||||
regret_fallback_depth_buckets: dict[str, int] = field(default_factory=dict)
|
||||
regret_fallback_opened_colors_buckets: dict[int, int] = field(default_factory=dict)
|
||||
regret_fallback_open_new_available_by_color: dict[int, int] = field(default_factory=dict)
|
||||
regret_fallback_open_new_selected_by_color: dict[int, int] = field(default_factory=dict)
|
||||
cutoff_rollouts: int = 0
|
||||
cutoff_rollout_steps: int = 0
|
||||
cutoff_rollout_timeouts: int = 0
|
||||
@@ -28,6 +50,42 @@ class TraversalStats:
|
||||
self.advantage_samples += other.advantage_samples
|
||||
self.strategy_samples += other.strategy_samples
|
||||
self.sampled_actions += other.sampled_actions
|
||||
self.regret_matching_decisions += other.regret_matching_decisions
|
||||
self.regret_fallback_count += other.regret_fallback_count
|
||||
self.regret_fallback_depth_sum += other.regret_fallback_depth_sum
|
||||
self.regret_fallback_opened_colors_sum += other.regret_fallback_opened_colors_sum
|
||||
self.regret_fallback_legal_actions_sum += other.regret_fallback_legal_actions_sum
|
||||
self.regret_fallback_legal_play_existing_sum += (
|
||||
other.regret_fallback_legal_play_existing_sum
|
||||
)
|
||||
self.regret_fallback_legal_open_new_sum += other.regret_fallback_legal_open_new_sum
|
||||
self.regret_fallback_legal_discard_sum += other.regret_fallback_legal_discard_sum
|
||||
self.regret_fallback_legal_draw_deck_sum += other.regret_fallback_legal_draw_deck_sum
|
||||
self.regret_fallback_legal_draw_pile_sum += other.regret_fallback_legal_draw_pile_sum
|
||||
self.regret_fallback_action_play_existing += other.regret_fallback_action_play_existing
|
||||
self.regret_fallback_action_open_new += other.regret_fallback_action_open_new
|
||||
self.regret_fallback_action_discard += other.regret_fallback_action_discard
|
||||
self.regret_fallback_action_draw_deck += other.regret_fallback_action_draw_deck
|
||||
self.regret_fallback_action_draw_pile += other.regret_fallback_action_draw_pile
|
||||
self.regret_fallback_argmax_tie_count += other.regret_fallback_argmax_tie_count
|
||||
self.regret_fallback_argmax_tie_size_sum += other.regret_fallback_argmax_tie_size_sum
|
||||
self.regret_fallback_argmax_full_tie_count += other.regret_fallback_argmax_full_tie_count
|
||||
for key, value in other.regret_fallback_depth_buckets.items():
|
||||
self.regret_fallback_depth_buckets[key] = (
|
||||
self.regret_fallback_depth_buckets.get(key, 0) + value
|
||||
)
|
||||
for key, value in other.regret_fallback_opened_colors_buckets.items():
|
||||
self.regret_fallback_opened_colors_buckets[key] = (
|
||||
self.regret_fallback_opened_colors_buckets.get(key, 0) + value
|
||||
)
|
||||
for key, value in other.regret_fallback_open_new_available_by_color.items():
|
||||
self.regret_fallback_open_new_available_by_color[key] = (
|
||||
self.regret_fallback_open_new_available_by_color.get(key, 0) + value
|
||||
)
|
||||
for key, value in other.regret_fallback_open_new_selected_by_color.items():
|
||||
self.regret_fallback_open_new_selected_by_color[key] = (
|
||||
self.regret_fallback_open_new_selected_by_color.get(key, 0) + value
|
||||
)
|
||||
self.cutoff_rollouts += other.cutoff_rollouts
|
||||
self.cutoff_rollout_steps += other.cutoff_rollout_steps
|
||||
self.cutoff_rollout_timeouts += other.cutoff_rollout_timeouts
|
||||
@@ -43,6 +101,73 @@ class TraversalStats:
|
||||
def avg_endpoint_depth(self) -> float:
|
||||
return self.endpoint_depth_sum / max(1, self.endpoints)
|
||||
|
||||
@property
|
||||
def regret_fallback_rate(self) -> float:
|
||||
return self.regret_fallback_count / max(1, self.regret_matching_decisions)
|
||||
|
||||
@property
|
||||
def regret_fallback_open_new_selected_rate(self) -> float:
|
||||
return self.regret_fallback_action_open_new / max(1, self.regret_fallback_count)
|
||||
|
||||
@property
|
||||
def regret_fallback_avg_depth(self) -> float:
|
||||
return self.regret_fallback_depth_sum / max(1, self.regret_fallback_count)
|
||||
|
||||
@property
|
||||
def regret_fallback_avg_opened_colors_before_action(self) -> float:
|
||||
return self.regret_fallback_opened_colors_sum / max(1, self.regret_fallback_count)
|
||||
|
||||
@property
|
||||
def regret_fallback_avg_legal_actions(self) -> float:
|
||||
return self.regret_fallback_legal_actions_sum / max(1, self.regret_fallback_count)
|
||||
|
||||
@property
|
||||
def regret_fallback_avg_legal_play_existing(self) -> float:
|
||||
return self.regret_fallback_legal_play_existing_sum / max(1, self.regret_fallback_count)
|
||||
|
||||
@property
|
||||
def regret_fallback_avg_legal_open_new(self) -> float:
|
||||
return self.regret_fallback_legal_open_new_sum / max(1, self.regret_fallback_count)
|
||||
|
||||
@property
|
||||
def regret_fallback_avg_legal_discard(self) -> float:
|
||||
return self.regret_fallback_legal_discard_sum / max(1, self.regret_fallback_count)
|
||||
|
||||
@property
|
||||
def regret_fallback_avg_legal_draw_deck(self) -> float:
|
||||
return self.regret_fallback_legal_draw_deck_sum / max(1, self.regret_fallback_count)
|
||||
|
||||
@property
|
||||
def regret_fallback_avg_legal_draw_pile(self) -> float:
|
||||
return self.regret_fallback_legal_draw_pile_sum / max(1, self.regret_fallback_count)
|
||||
|
||||
@property
|
||||
def regret_fallback_open_new_available_rate(self) -> float:
|
||||
return self.regret_fallback_legal_open_new_sum / max(
|
||||
1, self.regret_fallback_legal_actions_sum
|
||||
)
|
||||
|
||||
@property
|
||||
def regret_fallback_open_new_selection_over_availability(self) -> float:
|
||||
available_rate = self.regret_fallback_open_new_available_rate
|
||||
if available_rate <= 0.0:
|
||||
return 0.0
|
||||
return self.regret_fallback_open_new_selected_rate / available_rate
|
||||
|
||||
@property
|
||||
def regret_fallback_argmax_tie_rate(self) -> float:
|
||||
return self.regret_fallback_argmax_tie_count / max(1, self.regret_fallback_count)
|
||||
|
||||
@property
|
||||
def regret_fallback_argmax_avg_tie_size(self) -> float:
|
||||
return self.regret_fallback_argmax_tie_size_sum / max(
|
||||
1, self.regret_fallback_argmax_tie_count
|
||||
)
|
||||
|
||||
@property
|
||||
def regret_fallback_argmax_full_tie_rate(self) -> float:
|
||||
return self.regret_fallback_argmax_full_tie_count / max(1, self.regret_fallback_count)
|
||||
|
||||
def to_dict(self) -> dict[str, float | int]:
|
||||
return {
|
||||
"traversal_nodes": self.nodes,
|
||||
@@ -53,6 +178,57 @@ class TraversalStats:
|
||||
"traversal_advantage_samples": self.advantage_samples,
|
||||
"traversal_strategy_samples": self.strategy_samples,
|
||||
"traversal_sampled_actions": self.sampled_actions,
|
||||
"traversal_regret_matching_decisions": self.regret_matching_decisions,
|
||||
"traversal_regret_fallback_count": self.regret_fallback_count,
|
||||
"traversal_regret_fallback_rate": self.regret_fallback_rate,
|
||||
"traversal_regret_fallback_avg_depth": self.regret_fallback_avg_depth,
|
||||
"traversal_regret_fallback_action_play_existing": self.regret_fallback_action_play_existing,
|
||||
"traversal_regret_fallback_action_open_new": self.regret_fallback_action_open_new,
|
||||
"traversal_regret_fallback_action_discard": self.regret_fallback_action_discard,
|
||||
"traversal_regret_fallback_action_draw_deck": self.regret_fallback_action_draw_deck,
|
||||
"traversal_regret_fallback_action_draw_pile": self.regret_fallback_action_draw_pile,
|
||||
"traversal_regret_fallback_legal_actions_mean": (
|
||||
self.regret_fallback_avg_legal_actions
|
||||
),
|
||||
"traversal_regret_fallback_legal_play_existing_mean": (
|
||||
self.regret_fallback_avg_legal_play_existing
|
||||
),
|
||||
"traversal_regret_fallback_legal_open_new_mean": (
|
||||
self.regret_fallback_avg_legal_open_new
|
||||
),
|
||||
"traversal_regret_fallback_legal_discard_mean": (
|
||||
self.regret_fallback_avg_legal_discard
|
||||
),
|
||||
"traversal_regret_fallback_legal_draw_deck_mean": (
|
||||
self.regret_fallback_avg_legal_draw_deck
|
||||
),
|
||||
"traversal_regret_fallback_legal_draw_pile_mean": (
|
||||
self.regret_fallback_avg_legal_draw_pile
|
||||
),
|
||||
"traversal_regret_fallback_open_new_available_rate": (
|
||||
self.regret_fallback_open_new_available_rate
|
||||
),
|
||||
"traversal_regret_fallback_open_new_selected": self.regret_fallback_action_open_new,
|
||||
"traversal_regret_fallback_open_new_selected_rate": (
|
||||
self.regret_fallback_open_new_selected_rate
|
||||
),
|
||||
"traversal_regret_fallback_open_new_selection_over_availability": (
|
||||
self.regret_fallback_open_new_selection_over_availability
|
||||
),
|
||||
"traversal_regret_fallback_avg_opened_colors_before_action": (
|
||||
self.regret_fallback_avg_opened_colors_before_action
|
||||
),
|
||||
"traversal_regret_fallback_argmax_tie_count": (self.regret_fallback_argmax_tie_count),
|
||||
"traversal_regret_fallback_argmax_tie_rate": (self.regret_fallback_argmax_tie_rate),
|
||||
"traversal_regret_fallback_argmax_tie_size_mean": (
|
||||
self.regret_fallback_argmax_avg_tie_size
|
||||
),
|
||||
"traversal_regret_fallback_argmax_full_tie_count": (
|
||||
self.regret_fallback_argmax_full_tie_count
|
||||
),
|
||||
"traversal_regret_fallback_argmax_full_tie_rate": (
|
||||
self.regret_fallback_argmax_full_tie_rate
|
||||
),
|
||||
"traversal_cutoff_rollouts": self.cutoff_rollouts,
|
||||
"traversal_cutoff_rollout_steps": self.cutoff_rollout_steps,
|
||||
"traversal_cutoff_rollout_timeouts": self.cutoff_rollout_timeouts,
|
||||
@@ -63,4 +239,20 @@ class TraversalStats:
|
||||
f"traversal_endpoint_depth_bucket_{key}": value
|
||||
for key, value in self.endpoint_depth_buckets.items()
|
||||
},
|
||||
**{
|
||||
f"traversal_regret_fallback_depth_bucket_{key}": value
|
||||
for key, value in self.regret_fallback_depth_buckets.items()
|
||||
},
|
||||
**{
|
||||
f"traversal_regret_fallback_opened_colors_count_{key}": value
|
||||
for key, value in self.regret_fallback_opened_colors_buckets.items()
|
||||
},
|
||||
**{
|
||||
f"traversal_regret_fallback_open_new_available_color_{key}": value
|
||||
for key, value in self.regret_fallback_open_new_available_by_color.items()
|
||||
},
|
||||
**{
|
||||
f"traversal_regret_fallback_open_new_selected_color_{key}": value
|
||||
for key, value in self.regret_fallback_open_new_selected_by_color.items()
|
||||
},
|
||||
}
|
||||
|
||||
@@ -100,6 +100,7 @@ def run_traversal_worker_batch(batch: TraversalWorkerBatch) -> TraversalWorkerRe
|
||||
cutoff_rollout_policy=cfg.traversal.cutoff_rollout_policy,
|
||||
cutoff_rollout_max_steps=cfg.traversal.cutoff_rollout_max_steps,
|
||||
opponent_policy=cfg.traversal.opponent_policy,
|
||||
all_negative_fallback=cfg.regret_matching.all_negative_fallback,
|
||||
league_advantage_networks=league_networks,
|
||||
self_play_anchor_probability=cfg.self_play.anchor_probability,
|
||||
self_play_current_weight=cfg.self_play.current_weight,
|
||||
|
||||
@@ -64,6 +64,7 @@ def test_deep_cfr_loads_mapped_legacy_reproduction_config() -> None:
|
||||
assert config.evaluation.resolved_batch_size() == 64
|
||||
assert config.evaluation.device == "trainer"
|
||||
assert config.evaluation.resolved_num_workers() == 4
|
||||
assert config.regret_matching.all_negative_fallback == "uniform"
|
||||
assert config.checkpoint.save_iteration_interval == 10
|
||||
assert (
|
||||
config.checkpoint.directory == "runs/deep_cfr/deep_cfr_selfplay_full_depth_slot_playability"
|
||||
@@ -84,6 +85,7 @@ def test_deep_cfr_train_cli_count_overrides_disable_duration_limits() -> None:
|
||||
"checkpoint_dir": None,
|
||||
"eval_every": None,
|
||||
"eval_games": None,
|
||||
"regret_fallback": "argmax_tiebreak",
|
||||
"no_save": True,
|
||||
"save_latest_only": False,
|
||||
"save_iteration_interval": None,
|
||||
@@ -100,6 +102,7 @@ def test_deep_cfr_train_cli_count_overrides_disable_duration_limits() -> None:
|
||||
assert overridden.traversal.traversals_per_player is None
|
||||
assert overridden.traversal.resolved_traversals_per_player() == 1
|
||||
assert overridden.traversal.resolved_num_workers() == 0
|
||||
assert overridden.regret_matching.all_negative_fallback == "argmax_tiebreak"
|
||||
assert overridden.checkpoint.save_every_iteration is False
|
||||
assert overridden.checkpoint.save_latest is False
|
||||
|
||||
@@ -118,6 +121,7 @@ def test_deep_cfr_train_cli_checkpoint_save_overrides() -> None:
|
||||
"checkpoint_dir": None,
|
||||
"eval_every": None,
|
||||
"eval_games": None,
|
||||
"regret_fallback": None,
|
||||
"no_save": False,
|
||||
"save_latest_only": True,
|
||||
"save_iteration_interval": 1,
|
||||
@@ -388,6 +392,57 @@ def test_deep_cfr_cython_traverser_supports_outcome_sampling_and_rollout_cutoffs
|
||||
assert np.all(sample.target[unsampled_legal] == 0.0)
|
||||
|
||||
|
||||
def test_deep_cfr_cython_traverser_records_regret_fallback_metrics() -> None:
|
||||
trainer = DeepCFRTrainer(
|
||||
_deep_cfr_config(
|
||||
{
|
||||
"run": {"iterations": 1, "seed": 37},
|
||||
"network": {"hidden_size": 16},
|
||||
"traversal": {"traversals_per_iteration": 1, "max_depth": 1},
|
||||
"optimization": {"batch_size": 2},
|
||||
"checkpoint": {"save_every_iteration": False},
|
||||
"regret_matching": {"all_negative_fallback": "argmax_tiebreak"},
|
||||
}
|
||||
),
|
||||
LostCitiesConfig(seed=37),
|
||||
)
|
||||
for network in trainer.advantage_networks:
|
||||
for parameter in network.parameters():
|
||||
parameter.data.zero_()
|
||||
state = GameState.new_game(LostCitiesConfig(seed=37), seed=37)
|
||||
traverser = CythonDeepCFRTraverser(
|
||||
trainer.advantage_networks,
|
||||
device=trainer.device,
|
||||
action_size=trainer.action_size,
|
||||
max_depth=1,
|
||||
all_negative_fallback="argmax_tiebreak",
|
||||
seed=37,
|
||||
)
|
||||
|
||||
_, stats = traverser.traverse(state, traverser=0, iteration=1)
|
||||
metrics = stats.to_dict()
|
||||
action_count_sum = (
|
||||
stats.regret_fallback_action_play_existing
|
||||
+ stats.regret_fallback_action_open_new
|
||||
+ stats.regret_fallback_action_discard
|
||||
+ stats.regret_fallback_action_draw_deck
|
||||
+ stats.regret_fallback_action_draw_pile
|
||||
)
|
||||
|
||||
assert stats.regret_matching_decisions > 0
|
||||
assert stats.regret_fallback_count > 0
|
||||
assert action_count_sum == stats.regret_fallback_count
|
||||
assert metrics["traversal_regret_fallback_rate"] > 0.0
|
||||
assert "traversal_regret_fallback_open_new_selected_rate" in metrics
|
||||
assert metrics["traversal_regret_fallback_legal_actions_mean"] > 0.0
|
||||
assert "traversal_regret_fallback_open_new_available_rate" in metrics
|
||||
assert "traversal_regret_fallback_open_new_selection_over_availability" in metrics
|
||||
assert "traversal_regret_fallback_depth_bucket_0_49" in metrics
|
||||
assert "traversal_regret_fallback_opened_colors_count_0" in metrics
|
||||
assert metrics["traversal_regret_fallback_argmax_tie_rate"] > 0.0
|
||||
assert metrics["traversal_regret_fallback_argmax_tie_size_mean"] > 0.0
|
||||
|
||||
|
||||
def test_reservoir_memory_caps_samples_and_filters_player_batches() -> None:
|
||||
memory = ReservoirMemory(capacity=3)
|
||||
rng = np.random.default_rng(37)
|
||||
|
||||
Reference in New Issue
Block a user