Add first-open counterfactual audit

This commit is contained in:
2026-05-09 17:06:06 +09:00
parent ba54032703
commit 7c3a3499dc
2 changed files with 417 additions and 0 deletions
+34
View File
@@ -160,6 +160,40 @@ higher policy probability and sampled rate than good-open candidates. The
sampled target mean is also better for bad opens than good opens. This points
to a target or metric-alignment problem before model capacity or LCFR tuning.
## First-open counterfactual audit
Script:
- `scripts/analyze_first_open_counterfactual.py`
Output:
- `runs/tmp/first_open_counterfactual_confirm_eps005_200_vs_500.jsonl`
Method: collect first-open candidate states from existing checkpoints, force
each first-open candidate once, and compare the resulting continuation value
against the current policy's best non-open action from the same state.
`delta_open = value(force open) - value(best non-open)`.
Counterfactual summary against `safe_heuristic_strict`:
| checkpoint | bucket | candidates | delta mean | delta median | delta positive | policy prob | selected rate |
| --- | --- | ---: | ---: | ---: | ---: | ---: | ---: |
| iter 200 | good open | 40 | -26.57 | -26.5 | 0.050 | 0.000 | 0.000 |
| iter 200 | bad open | 460 | -23.15 | -21.0 | 0.130 | 0.036 | 0.037 |
| iter 500 | good open | 30 | -23.43 | -17.5 | 0.167 | 0.067 | 0.067 |
| iter 500 | bad open | 470 | -11.18 | -8.0 | 0.226 | 0.020 | 0.019 |
Interpretation: the heuristic `open_bad` label is not entirely misaligned with
continuation value. Forced bad opens are usually worse than the best non-open
alternative. However, heuristic `open_good` also often loses to best non-open in
these sampled states, so "recoverable eventually" is not the same as "open now."
Combined with the target audit, this points toward target/objective alignment:
the traversal target is not making the bad-open-vs-non-open mistake clearly
negative, even when the counterfactual continuation usually is negative.
## Open questions
1. Does the traversal target itself provide separable labels for good first
@@ -0,0 +1,383 @@
#!/usr/bin/env python
"""Compare first-open heuristic labels against forced-action continuation value."""
from __future__ import annotations
import argparse
import json
import time
from collections import defaultdict
from dataclasses import dataclass, field
from pathlib import Path
from typing import Any
import numpy as np
import torch
from coolrl_lost_cities.games.classic.deep_cfr.encoding import encode_info_state
from coolrl_lost_cities.games.classic.game import GameState, LostCitiesConfig
from coolrl_lost_cities.games.classic.bots import build_bot
from coolrl_lost_cities.games.classic.deep_cfr.config import config_from_dict
from coolrl_lost_cities.games.classic.deep_cfr.networks import DeepCFRMLP
@dataclass
class Bucket:
deltas: list[float] = field(default_factory=list)
open_values: list[float] = field(default_factory=list)
baseline_values: list[float] = field(default_factory=list)
policy_probs: list[float] = field(default_factory=list)
selected: int = 0
def add(
self,
*,
delta: float,
open_value: float,
baseline_value: float,
policy_prob: float,
selected: bool,
) -> None:
self.deltas.append(float(delta))
self.open_values.append(float(open_value))
self.baseline_values.append(float(baseline_value))
self.policy_probs.append(float(policy_prob))
self.selected += int(selected)
def to_dict(self) -> dict[str, float | int]:
return {
"count": len(self.deltas),
"delta_mean": _mean(self.deltas),
"delta_p10": _percentile(self.deltas, 10),
"delta_p25": _percentile(self.deltas, 25),
"delta_p50": _percentile(self.deltas, 50),
"delta_p75": _percentile(self.deltas, 75),
"delta_p90": _percentile(self.deltas, 90),
"delta_positive_rate": _positive_rate(self.deltas),
"open_value_mean": _mean(self.open_values),
"baseline_value_mean": _mean(self.baseline_values),
"policy_prob_mean": _mean(self.policy_probs),
"selected": self.selected,
"selected_rate": self.selected / max(1, len(self.deltas)),
}
def _mean(values: list[float]) -> float:
return float(np.mean(values)) if values else 0.0
def _percentile(values: list[float], percentile: float) -> float:
return float(np.percentile(values, percentile)) if values else 0.0
def _positive_rate(values: list[float]) -> float:
if not values:
return 0.0
return sum(1 for value in values if value > 0.0) / len(values)
def _numeric_value(card: Any, min_rank: int) -> int:
if card.rank == 0:
return 0
return min_rank + card.rank - 1
def _open_quality(state: GameState, player: int, color: int) -> str:
expedition = state.expeditions[player][color]
hand_cards = [
card for card in state.hand_slots(player) if card is not None and card.color == color
]
last_numeric = state.last_numeric_rank(player, color)
current_sum = sum(_numeric_value(card, state.config.min_rank) for card in expedition)
current_wagers = sum(1 for card in expedition if card.rank == 0)
playable_numeric = [card for card in hand_cards if card.rank > 0 and card.rank > last_numeric]
playable_wagers = [card for card in hand_cards if card.rank == 0 and last_numeric == 0]
projected_sum = current_sum + sum(
_numeric_value(card, state.config.min_rank) for card in playable_numeric
)
projected_wagers = current_wagers + len(playable_wagers)
projected_len = len(expedition) + len(playable_numeric) + len(playable_wagers)
recoverable_score = (projected_sum + state.config.expedition_penalty) * (projected_wagers + 1)
if recoverable_score >= 0:
return "open_good"
if projected_len >= state.config.bonus_threshold:
return "open_weak"
return "open_bad"
def _classify_action(state: GameState, unified_action: int, player: int) -> str:
card_action_size = state.config.hand_size * 2
if unified_action >= card_action_size:
return "draw_deck" if unified_action == card_action_size else "draw_pile"
if unified_action % 2 == 1:
return "discard"
card = state.hand_slots(player)[unified_action // 2]
if card is None:
return "invalid_play"
color = int(card.color)
if state.expeditions[player][color]:
return "play_existing"
return _open_quality(state, player, color)
def _regret_matching(
advantages: np.ndarray,
legal: np.ndarray,
*,
epsilon: float,
fallback: str,
) -> np.ndarray:
legal_actions = np.flatnonzero(legal)
policy = np.zeros_like(advantages, dtype=np.float32)
if len(legal_actions) == 0:
return policy
positive = np.where(legal, np.maximum(advantages, 0.0), 0.0).astype(np.float32)
total = float(positive.sum())
if total > epsilon:
return positive / total
if fallback == "uniform":
policy[legal_actions] = 1.0 / float(len(legal_actions))
return policy
best = float(np.max(advantages[legal_actions]))
best_actions = legal_actions[advantages[legal_actions] == best]
policy[int(best_actions[0])] = 1.0
return policy
class AdvantagePolicy:
def __init__(
self,
networks: list[torch.nn.Module],
*,
device: torch.device,
encoding: Any,
epsilon: float,
fallback: str,
) -> None:
self.networks = networks
self.device = device
self.encoding = encoding
self.epsilon = epsilon
self.fallback = fallback
def distribution(self, state: GameState) -> tuple[np.ndarray, np.ndarray, np.ndarray]:
player = int(state.current_player)
legal = np.asarray(state.unified_legal_mask(), dtype=bool)
info = encode_info_state(state, player, self.encoding)
with torch.inference_mode():
x = torch.as_tensor(info, dtype=torch.float32, device=self.device).unsqueeze(0)
advantages = self.networks[player](x).squeeze(0).detach().cpu().numpy()
policy = _regret_matching(
advantages,
legal,
epsilon=self.epsilon,
fallback=self.fallback,
)
return np.flatnonzero(legal), policy, advantages
def select_unified(self, state: GameState) -> int:
legal_actions, policy, advantages = self.distribution(state)
if len(legal_actions) == 0:
raise RuntimeError("no legal action available")
probs = policy[legal_actions]
if float(probs.sum()) > 0.0:
return int(legal_actions[int(np.argmax(probs))])
return int(legal_actions[int(np.argmax(advantages[legal_actions]))])
def _load_checkpoint(
checkpoint: Path,
device: torch.device,
) -> tuple[Any, LostCitiesConfig, AdvantagePolicy, int]:
payload = torch.load(checkpoint, map_location="cpu")
cfg = config_from_dict(payload["config"])
game_config = LostCitiesConfig(**payload["game_config"])
action_size = int(payload["action_size"])
networks = [
DeepCFRMLP.from_config(int(payload["input_dim"]), action_size, cfg.network).to(device)
for _ in range(2)
]
for network, state_dict in zip(networks, payload["advantage_networks"], strict=True):
network.load_state_dict(state_dict)
network.eval()
policy = AdvantagePolicy(
networks,
device=device,
encoding=cfg.encoding,
epsilon=cfg.traversal.regret_matching_epsilon,
fallback=cfg.regret_matching.all_negative_fallback,
)
return cfg, game_config, policy, int(payload.get("iteration", -1))
def _rollout_value(
state: GameState,
*,
policy_player: int,
policy: AdvantagePolicy,
opponent: str,
seed: int,
max_steps: int,
) -> float:
rollout = state.clone()
opponent_policy = build_bot(opponent, seed=seed)
steps = 0
while not rollout.terminal and steps < max_steps:
current = int(rollout.current_player)
if current == policy_player:
unified = policy.select_unified(rollout)
action = rollout.from_unified_action(unified)
else:
action = opponent_policy.act(rollout)
rollout.apply_action(action)
steps += 1
return float(rollout.score_diff(policy_player))
def _best_non_open_action(
labels: dict[int, str],
legal_actions: np.ndarray,
policy_probs: np.ndarray,
advantages: np.ndarray,
) -> int | None:
candidates = [
int(action) for action in legal_actions if not labels[int(action)].startswith("open_")
]
if not candidates:
return None
return max(
candidates,
key=lambda action: (float(policy_probs[action]), float(advantages[action]), -action),
)
def analyze_checkpoint(
checkpoint: Path,
*,
games: int,
seed: int,
opponent: str,
device: torch.device,
max_steps: int,
max_candidates: int,
) -> dict[str, Any]:
_cfg, game_config, policy, iteration = _load_checkpoint(checkpoint, device)
buckets: dict[str, Bucket] = defaultdict(Bucket)
candidate_states = 0
first_open_candidates = 0
evaluated_open_candidates = 0
policy_turns = 0
started = time.perf_counter()
for game_index in range(games):
if evaluated_open_candidates >= max_candidates:
break
game_seed = seed + game_index
swap = game_index % 2 == 1
policy_player = 1 if swap else 0
opponent_policy = build_bot(opponent, seed=game_seed * 2 + (1 - policy_player))
state = GameState.new_game(game_config, seed=game_seed)
for _step in range(max_steps):
if state.terminal or evaluated_open_candidates >= max_candidates:
break
current = int(state.current_player)
if current != policy_player:
state.apply_action(opponent_policy.act(state))
continue
policy_turns += 1
legal_actions, policy_probs, advantages = policy.distribution(state)
labels = {
int(action): _classify_action(state, int(action), current)
for action in legal_actions
}
open_actions = [
int(action) for action, label in labels.items() if label.startswith("open_")
]
best_non_open = _best_non_open_action(labels, legal_actions, policy_probs, advantages)
selected = policy.select_unified(state)
if open_actions and best_non_open is not None:
candidate_states += 1
baseline_state = state.clone()
baseline_state.apply_action(baseline_state.from_unified_action(best_non_open))
baseline_value = _rollout_value(
baseline_state,
policy_player=policy_player,
policy=policy,
opponent=opponent,
seed=game_seed * 10_000 + candidate_states * 101 + 1,
max_steps=max_steps,
)
for open_action in open_actions:
if evaluated_open_candidates >= max_candidates:
break
open_state = state.clone()
open_state.apply_action(open_state.from_unified_action(open_action))
open_value = _rollout_value(
open_state,
policy_player=policy_player,
policy=policy,
opponent=opponent,
seed=game_seed * 10_000 + candidate_states * 101 + 2,
max_steps=max_steps,
)
label = labels[open_action]
buckets[label].add(
delta=open_value - baseline_value,
open_value=open_value,
baseline_value=baseline_value,
policy_prob=float(policy_probs[open_action]),
selected=open_action == selected,
)
first_open_candidates += 1
evaluated_open_candidates += 1
state.apply_action(state.from_unified_action(selected))
return {
"checkpoint": str(checkpoint),
"iteration": iteration,
"opponent": opponent,
"games": games,
"seed": seed,
"device": str(device),
"max_candidates": max_candidates,
"elapsed_seconds": time.perf_counter() - started,
"policy_turns": policy_turns,
"candidate_states": candidate_states,
"first_open_candidates": first_open_candidates,
"buckets": {key: bucket.to_dict() for key, bucket in sorted(buckets.items())},
}
def main() -> None:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("checkpoints", nargs="+", type=Path)
parser.add_argument("--opponent", default="safe_heuristic_strict")
parser.add_argument("--games", type=int, default=100)
parser.add_argument("--seed", type=int, default=231_000)
parser.add_argument("--device", default="cuda")
parser.add_argument("--max-steps", type=int, default=10_000)
parser.add_argument("--max-candidates", type=int, default=500)
parser.add_argument("--output", type=Path, required=True)
args = parser.parse_args()
device = torch.device(args.device)
args.output.parent.mkdir(parents=True, exist_ok=True)
rows = []
for checkpoint in args.checkpoints:
row = analyze_checkpoint(
checkpoint,
games=args.games,
seed=args.seed,
opponent=args.opponent,
device=device,
max_steps=args.max_steps,
max_candidates=args.max_candidates,
)
rows.append(row)
print(json.dumps(row, sort_keys=True))
args.output.write_text("\n".join(json.dumps(row, sort_keys=True) for row in rows) + "\n")
print(f"wrote {args.output}")
if __name__ == "__main__":
main()