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")
|
||||
|
||||
Reference in New Issue
Block a user