Instructions to use capser54/gomoku-maskable-ppo-stage3-h6 with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- stable-baselines3
How to use capser54/gomoku-maskable-ppo-stage3-h6 with stable-baselines3:
from huggingface_sb3 import load_from_hub checkpoint = load_from_hub( repo_id="capser54/gomoku-maskable-ppo-stage3-h6", filename="{MODEL FILENAME}.zip", ) - Notebooks
- Google Colab
- Kaggle
Download evaluate.py from capser54/gomoku-maskable-ppo-stage3-h6: direct link, hf CLI and curl.
- Browser
- Download file 7.38 kB
-
https://huggingface.co/capser54/gomoku-maskable-ppo-stage3-h6/resolve/main/evaluate.py
- Command line
-
hf download hf://capser54/gomoku-maskable-ppo-stage3-h6/evaluate.py
-
curl -L -o evaluate.py https://huggingface.co/capser54/gomoku-maskable-ppo-stage3-h6/resolve/main/evaluate.py
7.38 kB
| from __future__ import annotations | |
| import argparse | |
| import os | |
| from pathlib import Path | |
| # Disable tqdm's rich integration BEFORE any imports to avoid shutdown errors | |
| # This must be set before tqdm is imported by any module | |
| os.environ["TQDM_DISABLE"] = "1" | |
| import numpy as np | |
| from sb3_contrib import MaskablePPO | |
| from sb3_contrib.common.wrappers import ActionMasker | |
| from gomoku_rl import GomokuEnv, HeuristicOpponent, PolicyGuidedOpponent, RandomOpponent | |
| from gomoku_rl.opponents import predict_masked_action | |
| def mask_fn(env: GomokuEnv): | |
| return env.action_masks() | |
| def build_opponent( | |
| name: str, | |
| seed: int, | |
| heuristic_search_depth: int = 2, | |
| heuristic_candidate_radius: int = 2, | |
| heuristic_max_candidates: int = 10, | |
| heuristic_early_max_candidates: int = 14, | |
| ): | |
| if name == "random": | |
| return RandomOpponent(seed=seed) | |
| if name == "heuristic": | |
| return HeuristicOpponent( | |
| seed=seed, | |
| search_depth=heuristic_search_depth, | |
| candidate_radius=heuristic_candidate_radius, | |
| max_candidates=heuristic_max_candidates, | |
| early_max_candidates=heuristic_early_max_candidates, | |
| ) | |
| raise ValueError(f"Unsupported opponent: {name}") | |
| def build_agent( | |
| model: MaskablePPO, | |
| agent_mode: str, | |
| seed: int, | |
| search_depth: int, | |
| candidate_radius: int, | |
| max_candidates: int, | |
| policy_weight: float, | |
| ): | |
| if agent_mode == "model": | |
| return model | |
| if agent_mode == "hybrid": | |
| return PolicyGuidedOpponent( | |
| model=model, | |
| seed=seed, | |
| search_depth=search_depth, | |
| candidate_radius=candidate_radius, | |
| max_candidates=max_candidates, | |
| early_max_candidates=max(max_candidates + 4, max_candidates), | |
| policy_weight=policy_weight, | |
| ) | |
| raise ValueError(f"Unsupported agent mode: {agent_mode}") | |
| def play_one_game(model: MaskablePPO, agent, env: ActionMasker, seed: int) -> str: | |
| env.reset(seed=seed) | |
| done = False | |
| info: dict[str, bool] = {} | |
| while not done: | |
| board = env.env.board | |
| if isinstance(agent, MaskablePPO): | |
| action = predict_masked_action(agent, board, player=1, deterministic=True) | |
| else: | |
| action = agent.choose_action(board.copy(), player=1, win_length=env.env.win_length) | |
| _, _, terminated, truncated, info = env.step(action) | |
| done = terminated or truncated | |
| if info.get("agent_won"): | |
| return "win" | |
| if info.get("opponent_won"): | |
| return "loss" | |
| if info.get("draw"): | |
| return "draw" | |
| raise RuntimeError(f"Unexpected terminal info: {info}") | |
| def evaluate( | |
| model_path: Path, | |
| opponent_name: str, | |
| games: int, | |
| board_size: int, | |
| win_length: int, | |
| seed: int, | |
| agent_mode: str = "model", | |
| search_depth: int = 3, | |
| candidate_radius: int = 2, | |
| max_candidates: int = 12, | |
| policy_weight: float = 220_000.0, | |
| opponent_search_depth: int = 2, | |
| opponent_candidate_radius: int = 2, | |
| opponent_max_candidates: int = 10, | |
| opponent_early_max_candidates: int = 14, | |
| ) -> dict[str, str | int | float]: | |
| opponent = build_opponent( | |
| opponent_name, | |
| seed, | |
| heuristic_search_depth=opponent_search_depth, | |
| heuristic_candidate_radius=opponent_candidate_radius, | |
| heuristic_max_candidates=opponent_max_candidates, | |
| heuristic_early_max_candidates=opponent_early_max_candidates, | |
| ) | |
| env = GomokuEnv( | |
| board_size=board_size, | |
| win_length=win_length, | |
| opponent=opponent, | |
| opponent_starts_prob=0.5, | |
| seed=seed, | |
| ) | |
| wrapped_env = ActionMasker(env, mask_fn) | |
| model = MaskablePPO.load(model_path) | |
| agent = build_agent( | |
| model=model, | |
| agent_mode=agent_mode, | |
| seed=seed, | |
| search_depth=search_depth, | |
| candidate_radius=candidate_radius, | |
| max_candidates=max_candidates, | |
| policy_weight=policy_weight, | |
| ) | |
| wins = 0 | |
| losses = 0 | |
| draws = 0 | |
| for game_idx in range(games): | |
| result = play_one_game(model, agent, wrapped_env, seed + game_idx) | |
| if result == "win": | |
| wins += 1 | |
| elif result == "loss": | |
| losses += 1 | |
| else: | |
| draws += 1 | |
| total = wins + losses + draws | |
| stats: dict[str, str | int | float] = { | |
| "model_path": str(model_path), | |
| "opponent": opponent_name, | |
| "agent_mode": agent_mode, | |
| "games": total, | |
| "wins": wins, | |
| "losses": losses, | |
| "draws": draws, | |
| "win_rate": wins / total, | |
| "loss_rate": losses / total, | |
| "draw_rate": draws / total, | |
| } | |
| print(f"Model: {model_path}") | |
| print(f"Agent mode: {agent_mode}") | |
| print(f"Opponent: {opponent_name}") | |
| if opponent_name == "heuristic": | |
| print( | |
| "Opponent heuristic config: " | |
| f"depth={opponent_search_depth}, " | |
| f"radius={opponent_candidate_radius}, " | |
| f"max_candidates={opponent_max_candidates}, " | |
| f"early_max_candidates={opponent_early_max_candidates}" | |
| ) | |
| print(f"Games: {total}") | |
| print(f"Wins: {wins}") | |
| print(f"Losses: {losses}") | |
| print(f"Draws: {draws}") | |
| print(f"Win rate: {stats['win_rate']:.3%}") | |
| print(f"Loss rate: {stats['loss_rate']:.3%}") | |
| print(f"Draw rate: {stats['draw_rate']:.3%}") | |
| return stats | |
| def parse_args() -> argparse.Namespace: | |
| parser = argparse.ArgumentParser(description="Evaluate a trained Gomoku MaskablePPO model.") | |
| parser.add_argument("--model-path", type=Path, required=True) | |
| parser.add_argument("--opponent", choices=["random", "heuristic"], required=True) | |
| parser.add_argument("--agent-mode", choices=["model", "hybrid"], default="model") | |
| parser.add_argument("--games", type=int, default=200) | |
| parser.add_argument("--board-size", type=int, default=9) | |
| parser.add_argument("--win-length", type=int, default=5) | |
| parser.add_argument("--seed", type=int, default=42) | |
| parser.add_argument("--search-depth", type=int, default=3) | |
| parser.add_argument("--candidate-radius", type=int, default=2) | |
| parser.add_argument("--max-candidates", type=int, default=12) | |
| parser.add_argument("--policy-weight", type=float, default=220_000.0) | |
| parser.add_argument("--opponent-search-depth", type=int, default=2) | |
| parser.add_argument("--opponent-candidate-radius", type=int, default=2) | |
| parser.add_argument("--opponent-max-candidates", type=int, default=10) | |
| parser.add_argument("--opponent-early-max-candidates", type=int, default=14) | |
| return parser.parse_args() | |
| def main() -> None: | |
| args = parse_args() | |
| evaluate( | |
| model_path=args.model_path, | |
| opponent_name=args.opponent, | |
| games=args.games, | |
| board_size=args.board_size, | |
| win_length=args.win_length, | |
| seed=args.seed, | |
| agent_mode=args.agent_mode, | |
| search_depth=args.search_depth, | |
| candidate_radius=args.candidate_radius, | |
| max_candidates=args.max_candidates, | |
| policy_weight=args.policy_weight, | |
| opponent_search_depth=args.opponent_search_depth, | |
| opponent_candidate_radius=args.opponent_candidate_radius, | |
| opponent_max_candidates=args.opponent_max_candidates, | |
| opponent_early_max_candidates=args.opponent_early_max_candidates, | |
| ) | |
| if __name__ == "__main__": | |
| main() | |