-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathmain.py
More file actions
124 lines (107 loc) · 3.36 KB
/
Copy pathmain.py
File metadata and controls
124 lines (107 loc) · 3.36 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
#!/usr/bin/env python3
"""
Main entry point for PongAI project.
Provides a CLI interface to train models, run demo, or start API server.
"""
import argparse
import sys
from pathlib import Path
# Add project root to path
sys.path.insert(0, str(Path(__file__).parent))
def main():
"""Main CLI entry point."""
parser = argparse.ArgumentParser(
description="PongAI: Human vs AI Reinforcement Learning Demo",
formatter_class=argparse.RawDescriptionHelpFormatter,
epilog="""
Examples:
python main.py train # Train a new model
python main.py demo # Run interactive demo
python main.py api # Start FastAPI server
python main.py demo --model models/rl_model_final.zip # Demo with specific model
"""
)
subparsers = parser.add_subparsers(dest='command', help='Command to run')
# Train command
train_parser = subparsers.add_parser('train', help='Train a PPO model')
train_parser.add_argument(
'--timesteps',
type=int,
default=1_000_000,
help='Total timesteps to train (default: 1,000,000)'
)
train_parser.add_argument(
'--envs',
type=int,
default=4,
help='Number of parallel environments (default: 4)'
)
train_parser.add_argument(
'--no-multiprocessing',
action='store_true',
help='Disable multiprocessing (use DummyVecEnv)'
)
train_parser.add_argument(
'--device',
type=str,
choices=['cuda', 'cpu'],
default=None,
help='Device to use (cuda or cpu, default: auto-detect)'
)
# Demo command
demo_parser = subparsers.add_parser('demo', help='Run interactive demo')
demo_parser.add_argument(
'--model',
type=str,
default='models/rl_model_final.zip',
help='Path to model to load (default: models/rl_model_final.zip)'
)
# API command
api_parser = subparsers.add_parser('api', help='Start FastAPI server')
api_parser.add_argument(
'--host',
type=str,
default='127.0.0.1',
help='Host to bind to (default: 127.0.0.1)'
)
api_parser.add_argument(
'--port',
type=int,
default=8000,
help='Port to bind to (default: 8000)'
)
api_parser.add_argument(
'--reload',
action='store_true',
default=True,
help='Enable auto-reload on file changes'
)
args = parser.parse_args()
if not args.command:
parser.print_help()
return 0
# Route to appropriate module
if args.command == 'train':
from train.ppo import train_ppo
print("[Main] Starting training...")
train_ppo(
total_timesteps=args.timesteps,
num_envs=args.envs,
use_multiprocessing=not args.no_multiprocessing,
device=args.device,
)
elif args.command == 'demo':
from demo.play import main as demo_main
print(f"[Main] Starting demo with model: {args.model}...")
demo_main(args.model)
elif args.command == 'api':
from api.app import run_server
print("[Main] Starting API server...")
run_server(
host=args.host,
port=args.port,
reload=args.reload,
)
return 0
if __name__ == "__main__":
sys.exit(main())