Skip to content
Merged
Show file tree
Hide file tree
Changes from 2 commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
3 changes: 3 additions & 0 deletions deep_quoridor/rust/src/bin/selfplay.rs
Original file line number Diff line number Diff line change
Expand Up @@ -363,6 +363,7 @@ fn run_continuous(
agent_p1.reset_game();
agent_p2.reset_game();

let game_start = std::time::Instant::now();
let result = play_game(
agent_p1.as_mut(),
agent_p2.as_mut(),
Expand All @@ -371,6 +372,8 @@ fn run_continuous(
q.max_steps as i32,
false,
)?;
let game_elapsed = game_start.elapsed().as_secs_f64();
println!("{} - selfplay finished in {:.4}", pid, game_elapsed);

// Atomic write: write to tmp dir, then rename to output dir
let ts = std::time::SystemTime::now()
Expand Down
39 changes: 5 additions & 34 deletions deep_quoridor/src/train_v2.py
Original file line number Diff line number Diff line change
Expand Up @@ -10,9 +10,7 @@
if __name__ == "__main__":
parser = argparse.ArgumentParser(description="Train Quoridor agent")
parser.add_argument("config_file", type=str, help="Path to YAML configuration file")
parser.add_argument(
"-r", "--runs-dir", type=str, default=None, help="Directory for runs"
)
parser.add_argument("-r", "--runs-dir", type=str, default=None, help="Directory for runs")
# TODO: implement this
# parser.add_argument("-c", "--continue", dest="continue_run", action="store_true", help="Continue an existing run")
parser.add_argument(
Expand All @@ -24,35 +22,10 @@

args = parser.parse_args()

runs_dir = (
args.runs_dir
if args.runs_dir is not None
else str(Path(__file__).parent.parent)
)
runs_dir = args.runs_dir if args.runs_dir is not None else str(Path(__file__).parent.parent)

config = load_config_and_setup_run(args.config_file, runs_dir, overrides=args.overrides)

use_rust = config.self_play.program == "rust"
if use_rust:
# Apply default Rust binary path if not specified in config
if config.self_play.rust_selfplay_binary is None:
config.self_play.rust_selfplay_binary = str(
Path(__file__).parent.parent
/ "rust"
/ "target"
/ "release"
/ "selfplay"
)
rust_binary = config.self_play.rust_selfplay_binary
if not Path(rust_binary).exists():
print(f"ERROR: Rust self-play binary not found at {rust_binary}")
print(
"Build it with: cd deep_quoridor/rust && cargo build --release --features binary --bin selfplay"
)
Comment on lines -35 to -55

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

I moved this to load_config_and_setup_run for clarity

exit(1)
# Rust self-play requires ONNX model exports
config.training.save_onnx = True

mp.set_start_method("spawn", force=True)

# Make sure we don't have the shutdown signal from a previous run
Expand All @@ -67,12 +40,12 @@
self_play_processes = []
rust_subprocesses = []

if use_rust:
if config.self_play.program == "rust":
# Spawn Rust self-play processes in continuous mode
config_file_path = str(config.paths.config_file)
for i in range(config.self_play.num_workers):
cmd = [
rust_binary,
config.self_play.rust_selfplay_binary,
"--config",
config_file_path,
"--output-dir",
Expand Down Expand Up @@ -102,9 +75,7 @@
sf_count = sum([p.is_alive() for p in self_play_processes])
sf_count += sum([p.poll() is None for p in rust_subprocesses])
if b_count_prev != b_count or sf_count_prev != sf_count:
print(
f"Waiting for {b_count} benchmark processes and {sf_count} self_play processes"
)
print(f"Waiting for {b_count} benchmark processes and {sf_count} self_play processes")
b_count_prev, sf_count_prev = b_count, sf_count

if (b_count + sf_count) == 0:
Expand Down
157 changes: 157 additions & 0 deletions deep_quoridor/src/tune_selfplay.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,157 @@
import argparse
import re
import shutil
import subprocess
import sys
import time
from pathlib import Path

SELFPLAY_RE = re.compile(r"selfplay finished in ([\d.]+)")


def parse_selfplay_line(line):
"""Parse a self-play round completion line.

Returns (elapsed_seconds, num_games) or None if the line doesn't match.
"""
m = SELFPLAY_RE.search(line)
if not m:
return None

elapsed = float(m.group(1))

return elapsed


def run_single_benchmark(config_file, num_workers, parallel_games, duration, runs_dir, extra_overrides):
run_id = f"bench-w{num_workers}-g{parallel_games}-{int(time.time())}"

overrides = [
f"run_id={run_id}",
f"self_play.num_workers={num_workers}",
f"self_play.parallel_games={parallel_games}",
f"training.finish_after={duration}",
"benchmarks=[]",
"wandb=None",
] + extra_overrides

cmd = [sys.executable, "train_v2.py", config_file, "-o"] + overrides
if runs_dir:
cmd.insert(3, "--runs-dir")
cmd.insert(4, runs_dir)

src_dir = Path(__file__).parent

wall_start = time.time()
proc = subprocess.Popen(
cmd,
stdout=subprocess.PIPE,
stderr=subprocess.STDOUT,
text=True,
cwd=str(src_dir),
)
durations = []
for line in proc.stdout:
line = line.rstrip()
parsed = parse_selfplay_line(line)
if parsed:
print(f" {line}")
durations.append(parsed)

proc.wait()
wall_elapsed = time.time() - wall_start

# Cleanup run directory
effective_runs_dir = runs_dir or str(src_dir.parent)
run_dir = Path(effective_runs_dir) / "runs" / run_id
if run_dir.exists():
shutil.rmtree(run_dir)

return compute_metrics(num_workers, parallel_games, durations, wall_elapsed)


def compute_metrics(num_workers, parallel_games, durations, wall_elapsed):
if not durations:
return {
"num_workers": num_workers,
"parallel_games": parallel_games,
"total_rounds": 0,
"total_games": 0,
"avg_round_time": float("nan"),
"avg_throughput": 0.0,
}

total_rounds = len(durations)
total_worker_time = sum(durations)
total_games = parallel_games * total_rounds
avg_throughput = num_workers * total_games / total_worker_time

return {
"num_workers": num_workers,
"parallel_games": parallel_games,
"total_rounds": total_rounds,
"total_games": total_games,
"avg_round_time": total_worker_time / total_rounds,
"avg_throughput": avg_throughput,
}


def print_results_table(results):
print(f"\n{'=' * 80}")
print("BENCHMARK RESULTS")
print(f"{'=' * 80}")

header = f"{'nw':>10} {'pg':>10} {'rounds':>10} {'tot games':>10} {'avg round t':>14} {'games/s':>14}"
print(header)
print("-" * len(header))

for r in sorted(results, key=lambda x: x["avg_throughput"], reverse=True):
print(
f"{r['num_workers']:>10} "
f"{r['parallel_games']:>10} "
f"{r['total_rounds']:>10} "
f"{r['total_games']:>10} "
f"{r['avg_round_time']:>14.2f} "
f"{r['avg_throughput']:>14.3f} "
)


def main():
parser = argparse.ArgumentParser(description="Benchmark self-play throughput across configurations")
parser.add_argument("config_file", type=str, help="Base config YAML file")
parser.add_argument("--workers", type=str, help="Comma-separated num_workers values")
parser.add_argument("--games", type=str, help="Comma-separated parallel_games values")
parser.add_argument("--duration", type=str, default="2 minutes", help="Duration per combo (default: '2 minutes')")
parser.add_argument("--runs-dir", type=str, default=None, help="Directory for runs")
parser.add_argument("--extra-overrides", nargs="*", default=[], help="Additional config overrides for train_v2.py")

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

isn't this the same as -o? Shouldn't it also be -o for consistency?

args = parser.parse_args()

workers_list = [int(x) for x in args.workers.split(",")]
games_list = [int(x) for x in args.games.split(",")]
combos = [(w, g) for w in workers_list for g in games_list]

print(f"Benchmarking {len(combos)} configurations, {args.duration} each")
print(f" Workers: {workers_list}")
print(f" Parallel games: {games_list}")

results = []
for i, (num_workers, parallel_games) in enumerate(combos):
print(f"\n{'=' * 60}")
print(f"[{i + 1}/{len(combos)}] num_workers={num_workers}, parallel_games={parallel_games}")
print(f"{'=' * 60}")

result = run_single_benchmark(
config_file=args.config_file,
num_workers=num_workers,
parallel_games=parallel_games,
duration=args.duration,
runs_dir=args.runs_dir,
extra_overrides=args.extra_overrides,
)
results.append(result)

print_results_table(results)


if __name__ == "__main__":
main()
1 change: 0 additions & 1 deletion deep_quoridor/src/v2/TODO.md
Original file line number Diff line number Diff line change
Expand Up @@ -22,4 +22,3 @@
- Allow to dynamically change the number of workers and parallel games, to experiment with performance
- Mount the run directory and make other processes play from another computer
- The processes could write status files and we could have a script to watch the status (e.g. elapsed time.)
- Tuning: a script that would take a list of num-workers and a list of parallel games, as well as time to run, and will execute the different combinations for that time and output how many games were played in each combination, to find the fastest one.
27 changes: 18 additions & 9 deletions deep_quoridor/src/v2/config.py
Original file line number Diff line number Diff line change
Expand Up @@ -151,9 +151,7 @@ class PathsConfig(StrictBaseModel):
config_file: Path

@classmethod
def create(
cls, base_dir: str, run_id: str, create_dirs: bool = True
) -> "PathsConfig":
def create(cls, base_dir: str, run_id: str, create_dirs: bool = True) -> "PathsConfig":
run_root = Path(base_dir) / "runs"
run_dir = run_root / run_id
config_file = run_dir / "config.yaml"
Expand Down Expand Up @@ -190,9 +188,7 @@ class Config(UserConfig):
paths: PathsConfig

@classmethod
def from_user(
cls, user: UserConfig, base_dir: str, create_dirs: bool = True
) -> "Config":
def from_user(cls, user: UserConfig, base_dir: str, create_dirs: bool = True) -> "Config":
paths = PathsConfig.create(base_dir, user.run_id, create_dirs=create_dirs)
return cls(**user.model_dump(), paths=paths)

Expand Down Expand Up @@ -282,9 +278,7 @@ def _apply_overrides(data: dict, overrides: list[str]) -> dict:
"""
for override in overrides:
if "=" not in override:
raise ValueError(
f"Invalid override format '{override}', expected 'key=value'"
)
raise ValueError(f"Invalid override format '{override}', expected 'key=value'")
key, value = override.split("=", 1)
parts = key.split(".")
parsed_value = _parse_override_value(value)
Expand Down Expand Up @@ -317,4 +311,19 @@ def load_config_and_setup_run(
with config_filename.open(mode="w") as f:
f.write(to_yaml_str_ordered(user_config))

use_rust = config.self_play.program == "rust"
if use_rust:
# Apply default Rust binary path if not specified in config
if config.self_play.rust_selfplay_binary is None:
config.self_play.rust_selfplay_binary = str(
Path(__file__).parent.parent.parent / "rust" / "target" / "release" / "selfplay"
)
rust_binary = config.self_play.rust_selfplay_binary
if not Path(rust_binary).exists():
print(f"ERROR: Rust self-play binary not found at {rust_binary}")
print("Build it with: cd deep_quoridor/rust && cargo build --release --features binary --bin selfplay")
exit(1)
# Rust self-play requires ONNX model exports
config.training.save_onnx = True

return config
2 changes: 1 addition & 1 deletion deep_quoridor/src/v2/self_play.py
Original file line number Diff line number Diff line change
Expand Up @@ -71,4 +71,4 @@ def self_play(config: Config):
)
num_truncated = n - len(finished_in)
elapsed = Timer.finish("self-play")
print(f"{os.getpid()} - finsihed in {elapsed} {sorted(finished_in)}, {num_truncated}")
print(f"{os.getpid()} - selfplay finished in {elapsed} {sorted(finished_in)}, {num_truncated}")