Deep CFR 재현 config 실행 보강
This commit is contained in:
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user