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")