Deep CFR 재현 config 실행 보강

This commit is contained in:
2026-05-07 01:12:09 +09:00
parent bbf8950c3a
commit df8f979ce4
4 changed files with 123 additions and 10 deletions
@@ -2,8 +2,8 @@ from __future__ import annotations
from collections.abc import Callable
from ..policy import LostCitiesPolicy
from .heuristic import SafeHeuristicBot
from ..policy import LostCitiesPolicy, PolicyInput
from .heuristic import SafeHeuristicBot, SafeHeuristicParams
from .passive import PassiveDiscardBot
from .random import RandomBot
@@ -11,20 +11,67 @@ BotName = str
DEFAULT_BOT: BotName = "random"
PolicyFactory = Callable[[int | None], LostCitiesPolicy]
class NoisyPolicy(LostCitiesPolicy):
def __init__(
self,
base: LostCitiesPolicy,
random_policy: RandomBot,
*,
epsilon: float = 0.15,
):
self.base = base
self.random_policy = random_policy
self.epsilon = epsilon
def act(self, obs_or_state: PolicyInput) -> int:
if self.random_policy.rng.random() < self.epsilon:
return self.random_policy.act(obs_or_state)
return self.base.act(obs_or_state)
LOOSE_SAFE_HEURISTIC_PARAMS = SafeHeuristicParams(
open_target_ratio=0.42,
open_min_card_ratio=0.30,
handshake_target_multiplier=1.00,
handshake_min_card_ratio=0.25,
late_open_block_ratio=0.12,
)
STRICT_SAFE_HEURISTIC_PARAMS = SafeHeuristicParams(
open_target_ratio=0.62,
open_min_card_ratio=0.50,
handshake_target_multiplier=1.35,
handshake_min_card_ratio=0.45,
late_open_block_ratio=0.30,
)
BOT_REGISTRY: dict[BotName, PolicyFactory] = {
DEFAULT_BOT: RandomBot,
"passive-discard": lambda seed: PassiveDiscardBot(),
"safe-heuristic": lambda seed: SafeHeuristicBot(),
"safe-heuristic-loose": lambda seed: SafeHeuristicBot(LOOSE_SAFE_HEURISTIC_PARAMS),
"safe-heuristic-strict": lambda seed: SafeHeuristicBot(STRICT_SAFE_HEURISTIC_PARAMS),
"noisy-safe": lambda seed: NoisyPolicy(
SafeHeuristicBot(),
RandomBot(seed),
),
}
def canonical_bot_name(name: BotName) -> BotName:
return name.strip().lower().replace("_", "-")
def available_bot_names() -> list[BotName]:
return sorted(BOT_REGISTRY)
def build_bot(name: BotName, *, seed: int | None = None) -> LostCitiesPolicy:
canonical = canonical_bot_name(name)
try:
policy_factory = BOT_REGISTRY[name]
policy_factory = BOT_REGISTRY[canonical]
except KeyError as exc:
raise ValueError(f"unknown Lost Cities bot: {name}") from exc
return policy_factory(seed)
@@ -45,17 +45,26 @@ def _with_overrides(config: DeepCFRConfig, overrides: dict[str, Any]) -> DeepCFR
return DeepCFRConfig.model_validate(data)
def train_command(args: argparse.Namespace) -> None:
config = _load_config(args.config)
def _train_overrides_from_args(args: argparse.Namespace) -> dict[str, Any]:
overrides: dict[str, Any] = {}
run_overrides = overrides.setdefault("run", {})
if args.iterations is not None:
overrides.setdefault("run", {})["iterations"] = args.iterations
run_overrides["iterations"] = args.iterations
if args.max_hours is None and args.max_iterations is None:
run_overrides["max_hours"] = None
run_overrides["max_iterations"] = None
if args.max_hours is not None:
run_overrides["max_hours"] = args.max_hours
if args.max_iterations is not None:
run_overrides["max_iterations"] = args.max_iterations
if args.seed is not None:
overrides.setdefault("run", {})["seed"] = args.seed
run_overrides["seed"] = args.seed
if args.traversals_per_iteration is not None:
overrides.setdefault("traversal", {})["traversals_per_iteration"] = (
args.traversals_per_iteration
)
traversal_overrides = overrides.setdefault("traversal", {})
traversal_overrides["traversals_per_iteration"] = args.traversals_per_iteration
traversal_overrides["traversals_per_player"] = None
if args.num_workers is not None:
overrides.setdefault("traversal", {})["num_workers"] = args.num_workers
if args.checkpoint_dir is not None:
overrides.setdefault("checkpoint", {})["directory"] = args.checkpoint_dir
if args.eval_every is not None:
@@ -64,6 +73,12 @@ def train_command(args: argparse.Namespace) -> None:
overrides.setdefault("evaluation", {})["games"] = args.eval_games
if args.no_save:
overrides.setdefault("checkpoint", {})["save_every_iteration"] = False
return overrides
def train_command(args: argparse.Namespace) -> None:
config = _load_config(args.config)
overrides = _train_overrides_from_args(args)
config = _with_overrides(config, overrides)
trainer = DeepCFRTrainer(
config,
@@ -165,7 +180,10 @@ def main(argv: list[str] | None = None) -> None:
train = subparsers.add_parser("train")
train.add_argument("--config")
train.add_argument("--iterations", type=int)
train.add_argument("--max-hours", type=float)
train.add_argument("--max-iterations", type=int)
train.add_argument("--traversals-per-iteration", type=int)
train.add_argument("--num-workers")
train.add_argument("--checkpoint-dir")
train.add_argument("--resume")
train.add_argument("--device")
@@ -8,6 +8,10 @@ from coolrl_lost_cities.games.classic.deep_cfr.benchmark import (
benchmark_traversal,
benchmark_traversal_modes,
)
from coolrl_lost_cities.games.classic.deep_cfr.cli import (
_train_overrides_from_args,
_with_overrides,
)
from coolrl_lost_cities.games.classic.deep_cfr.config import DeepCFRConfig, load_config
from coolrl_lost_cities.games.classic.deep_cfr.memory import ReservoirMemory, TrainingSample
from coolrl_lost_cities.games.classic.deep_cfr.trainer import DeepCFRTrainer
@@ -58,6 +62,38 @@ def test_deep_cfr_loads_mapped_legacy_reproduction_config() -> None:
)
def test_deep_cfr_train_cli_count_overrides_disable_duration_limits() -> None:
args = type(
"Args",
(),
{
"iterations": 1,
"max_hours": None,
"max_iterations": None,
"seed": None,
"traversals_per_iteration": 1,
"num_workers": "0",
"checkpoint_dir": None,
"eval_every": None,
"eval_games": None,
"no_save": True,
},
)()
config = load_config(
"configs/deep_cfr/pure_self_play_zero_pit_poc_full_depth_slot_aware_playability.yaml"
)
overridden = _with_overrides(config, _train_overrides_from_args(args))
assert overridden.run.iterations == 1
assert overridden.run.max_hours is None
assert overridden.run.max_iterations is 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.checkpoint.save_every_iteration is False
def test_deep_cfr_playability_encoding_extends_input_shape() -> None:
state = GameState.new_game(LostCitiesConfig(seed=61), seed=61)
base_dim = input_dim(state)
+12
View File
@@ -38,3 +38,15 @@ def test_classic_package_exports_snapshot_alias() -> None:
def test_classic_package_exports_bot_registry_helpers() -> None:
assert "random" in classic.available_bot_names()
assert isinstance(classic.build_bot("random", seed=1), classic.LostCitiesPolicy)
def test_classic_bot_registry_accepts_reproduction_opponent_names() -> None:
for name in [
"random",
"passive_discard",
"safe_heuristic",
"safe_heuristic_loose",
"safe_heuristic_strict",
"noisy_safe",
]:
assert isinstance(classic.build_bot(name, seed=1), classic.LostCitiesPolicy)