Deep CFR 재현 config 실행 보강
This commit is contained in:
@@ -2,8 +2,8 @@ from __future__ import annotations
|
|||||||
|
|
||||||
from collections.abc import Callable
|
from collections.abc import Callable
|
||||||
|
|
||||||
from ..policy import LostCitiesPolicy
|
from ..policy import LostCitiesPolicy, PolicyInput
|
||||||
from .heuristic import SafeHeuristicBot
|
from .heuristic import SafeHeuristicBot, SafeHeuristicParams
|
||||||
from .passive import PassiveDiscardBot
|
from .passive import PassiveDiscardBot
|
||||||
from .random import RandomBot
|
from .random import RandomBot
|
||||||
|
|
||||||
@@ -11,20 +11,67 @@ BotName = str
|
|||||||
DEFAULT_BOT: BotName = "random"
|
DEFAULT_BOT: BotName = "random"
|
||||||
PolicyFactory = Callable[[int | None], LostCitiesPolicy]
|
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] = {
|
BOT_REGISTRY: dict[BotName, PolicyFactory] = {
|
||||||
DEFAULT_BOT: RandomBot,
|
DEFAULT_BOT: RandomBot,
|
||||||
"passive-discard": lambda seed: PassiveDiscardBot(),
|
"passive-discard": lambda seed: PassiveDiscardBot(),
|
||||||
"safe-heuristic": lambda seed: SafeHeuristicBot(),
|
"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]:
|
def available_bot_names() -> list[BotName]:
|
||||||
return sorted(BOT_REGISTRY)
|
return sorted(BOT_REGISTRY)
|
||||||
|
|
||||||
|
|
||||||
def build_bot(name: BotName, *, seed: int | None = None) -> LostCitiesPolicy:
|
def build_bot(name: BotName, *, seed: int | None = None) -> LostCitiesPolicy:
|
||||||
|
canonical = canonical_bot_name(name)
|
||||||
try:
|
try:
|
||||||
policy_factory = BOT_REGISTRY[name]
|
policy_factory = BOT_REGISTRY[canonical]
|
||||||
except KeyError as exc:
|
except KeyError as exc:
|
||||||
raise ValueError(f"unknown Lost Cities bot: {name}") from exc
|
raise ValueError(f"unknown Lost Cities bot: {name}") from exc
|
||||||
return policy_factory(seed)
|
return policy_factory(seed)
|
||||||
|
|||||||
@@ -45,17 +45,26 @@ def _with_overrides(config: DeepCFRConfig, overrides: dict[str, Any]) -> DeepCFR
|
|||||||
return DeepCFRConfig.model_validate(data)
|
return DeepCFRConfig.model_validate(data)
|
||||||
|
|
||||||
|
|
||||||
def train_command(args: argparse.Namespace) -> None:
|
def _train_overrides_from_args(args: argparse.Namespace) -> dict[str, Any]:
|
||||||
config = _load_config(args.config)
|
|
||||||
overrides: dict[str, Any] = {}
|
overrides: dict[str, Any] = {}
|
||||||
|
run_overrides = overrides.setdefault("run", {})
|
||||||
if args.iterations is not None:
|
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:
|
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:
|
if args.traversals_per_iteration is not None:
|
||||||
overrides.setdefault("traversal", {})["traversals_per_iteration"] = (
|
traversal_overrides = overrides.setdefault("traversal", {})
|
||||||
args.traversals_per_iteration
|
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:
|
if args.checkpoint_dir is not None:
|
||||||
overrides.setdefault("checkpoint", {})["directory"] = args.checkpoint_dir
|
overrides.setdefault("checkpoint", {})["directory"] = args.checkpoint_dir
|
||||||
if args.eval_every is not None:
|
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
|
overrides.setdefault("evaluation", {})["games"] = args.eval_games
|
||||||
if args.no_save:
|
if args.no_save:
|
||||||
overrides.setdefault("checkpoint", {})["save_every_iteration"] = False
|
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)
|
config = _with_overrides(config, overrides)
|
||||||
trainer = DeepCFRTrainer(
|
trainer = DeepCFRTrainer(
|
||||||
config,
|
config,
|
||||||
@@ -165,7 +180,10 @@ def main(argv: list[str] | None = None) -> None:
|
|||||||
train = subparsers.add_parser("train")
|
train = subparsers.add_parser("train")
|
||||||
train.add_argument("--config")
|
train.add_argument("--config")
|
||||||
train.add_argument("--iterations", type=int)
|
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("--traversals-per-iteration", type=int)
|
||||||
|
train.add_argument("--num-workers")
|
||||||
train.add_argument("--checkpoint-dir")
|
train.add_argument("--checkpoint-dir")
|
||||||
train.add_argument("--resume")
|
train.add_argument("--resume")
|
||||||
train.add_argument("--device")
|
train.add_argument("--device")
|
||||||
|
|||||||
@@ -8,6 +8,10 @@ from coolrl_lost_cities.games.classic.deep_cfr.benchmark import (
|
|||||||
benchmark_traversal,
|
benchmark_traversal,
|
||||||
benchmark_traversal_modes,
|
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.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.memory import ReservoirMemory, TrainingSample
|
||||||
from coolrl_lost_cities.games.classic.deep_cfr.trainer import DeepCFRTrainer
|
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:
|
def test_deep_cfr_playability_encoding_extends_input_shape() -> None:
|
||||||
state = GameState.new_game(LostCitiesConfig(seed=61), seed=61)
|
state = GameState.new_game(LostCitiesConfig(seed=61), seed=61)
|
||||||
base_dim = input_dim(state)
|
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:
|
def test_classic_package_exports_bot_registry_helpers() -> None:
|
||||||
assert "random" in classic.available_bot_names()
|
assert "random" in classic.available_bot_names()
|
||||||
assert isinstance(classic.build_bot("random", seed=1), classic.LostCitiesPolicy)
|
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