fixed env eval setup
This commit is contained in:
+79
-35
@@ -1,10 +1,43 @@
|
||||
import argparse
|
||||
import os
|
||||
import argparse
|
||||
from pathlib import Path
|
||||
|
||||
from stable_baselines3.common.callbacks import BaseCallback
|
||||
from stable_baselines3.common.vec_env import SubprocVecEnv, DummyVecEnv
|
||||
from .env import JackBotEnv
|
||||
|
||||
|
||||
class MilestoneCheckpointCallback(BaseCallback):
|
||||
"""
|
||||
Saves a model checkpoint the FIRST time total_timesteps
|
||||
crosses every multiple of step_interval (e.g., 100,000).
|
||||
"""
|
||||
def __init__(self, save_path: str, name_prefix: str = "ppo_jackbot", step_interval: int = 100_000, verbose: int = 1):
|
||||
super().__init__(verbose)
|
||||
self.save_path = save_path
|
||||
self.name_prefix = name_prefix
|
||||
self.step_interval = step_interval
|
||||
self.last_milestone = 0
|
||||
os.makedirs(self.save_path, exist_ok=True)
|
||||
|
||||
def _on_step(self) -> bool:
|
||||
current_milestone = self.num_timesteps // self.step_interval
|
||||
|
||||
if current_milestone > self.last_milestone:
|
||||
self.last_milestone = current_milestone
|
||||
milestone_step = current_milestone * self.step_interval
|
||||
|
||||
save_file = os.path.join(
|
||||
self.save_path,
|
||||
f"{self.name_prefix}_{milestone_step}_steps.zip"
|
||||
)
|
||||
self.model.save(save_file)
|
||||
|
||||
if self.verbose > 0:
|
||||
print(f"\n[Checkpoint] Saved milestone model at {self.num_timesteps} steps -> {save_file}\n")
|
||||
|
||||
return True
|
||||
|
||||
|
||||
def parse_args():
|
||||
parser = argparse.ArgumentParser(description="Train a joint-command policy for JackBot.")
|
||||
parser.add_argument("--timesteps", type=int, default=500_000, help="Total training timesteps")
|
||||
@@ -12,12 +45,26 @@ def parse_args():
|
||||
parser.add_argument("--seed", type=int, default=0, help="Random seed")
|
||||
parser.add_argument("--device", type=str, default="auto", help="Device to use: 'cpu', 'cuda', or 'auto' to autodetect")
|
||||
parser.add_argument("--use-gui", action="store_true", help="Enable PyBullet GUI during training")
|
||||
parser.add_argument("--num-robots", type=int, default=1, help="Number of robots in the training environment")
|
||||
parser.add_argument("--num-workers", type=int, default=8, help="Number of parallel CPU worker processes")
|
||||
parser.add_argument("--robot-spacing", type=float, default=3.0, help="Spacing between robots in meters")
|
||||
parser.add_argument("--start-pose", type=str, choices=["init_deg", "init90_deg"], default="init_deg", help="Initial robot pose at reset")
|
||||
return parser.parse_args()
|
||||
|
||||
|
||||
def make_env(num_robots, robot_spacing, start_pose, use_gui, rank, seed=0):
|
||||
def _init():
|
||||
env = JackBotEnv(
|
||||
use_gui=use_gui if rank == 0 else False, # Only rank 0 gets GUI if requested
|
||||
random_command=True,
|
||||
num_robots=num_robots,
|
||||
robot_spacing=robot_spacing,
|
||||
start_pose=start_pose,
|
||||
)
|
||||
env.reset(seed=seed + rank)
|
||||
return env
|
||||
return _init
|
||||
|
||||
|
||||
def train(
|
||||
total_timesteps: int,
|
||||
model_path: str,
|
||||
@@ -25,37 +72,33 @@ def train(
|
||||
device: str = "auto",
|
||||
use_gui: bool = False,
|
||||
num_robots: int = 1,
|
||||
num_workers: int = 8,
|
||||
robot_spacing: float = 0.5,
|
||||
start_pose: str = "init_deg",
|
||||
):
|
||||
try:
|
||||
from stable_baselines3 import PPO
|
||||
from stable_baselines3.common.vec_env import DummyVecEnv
|
||||
except ImportError as exc:
|
||||
raise ImportError(
|
||||
"stable-baselines3 and gym are required for training. "
|
||||
"Install them with: pip install stable-baselines3 gym"
|
||||
) from exc
|
||||
raise ImportError("stable-baselines3 is required. Install with: pip install stable-baselines3") from exc
|
||||
|
||||
env = DummyVecEnv([
|
||||
lambda: JackBotEnv(
|
||||
use_gui=use_gui,
|
||||
random_command=True,
|
||||
num_robots=num_robots,
|
||||
robot_spacing=robot_spacing,
|
||||
start_pose=start_pose,
|
||||
)
|
||||
])
|
||||
# Create multi-process vector environment
|
||||
if num_workers > 1:
|
||||
env_fns = [
|
||||
make_env(num_robots, robot_spacing, start_pose, use_gui, rank=i, seed=seed)
|
||||
for i in range(num_workers)
|
||||
]
|
||||
env = SubprocVecEnv(env_fns)
|
||||
else:
|
||||
env = DummyVecEnv([
|
||||
make_env(num_robots, robot_spacing, start_pose, use_gui, rank=0, seed=seed)
|
||||
])
|
||||
|
||||
def resolve_device(requested_device: str) -> str:
|
||||
try:
|
||||
import torch
|
||||
except ImportError:
|
||||
if requested_device != "cpu":
|
||||
raise RuntimeError(
|
||||
"PyTorch is not installed in the active environment. "
|
||||
"Install torch with a GPU-enabled build before using --device cuda."
|
||||
)
|
||||
raise RuntimeError("PyTorch is not installed.")
|
||||
return "cpu"
|
||||
|
||||
hip_supported = getattr(torch.version, "hip", None) is not None
|
||||
@@ -63,24 +106,14 @@ def train(
|
||||
hip_available = hip_supported and getattr(torch.backends, "hip", None) is not None and torch.backends.hip.is_available()
|
||||
|
||||
if requested_device == "auto":
|
||||
if hip_available or cuda_available:
|
||||
return "cuda"
|
||||
return "cpu"
|
||||
return "cuda" if (hip_available or cuda_available) else "cpu"
|
||||
|
||||
if requested_device in {"cuda", "gpu", "hip"}:
|
||||
if hip_available or cuda_available:
|
||||
return "cuda"
|
||||
raise RuntimeError(
|
||||
f"GPU device requested ({requested_device}) but no CUDA/ROCm-capable PyTorch is available. "
|
||||
f"Installed torch build: {torch.__version__} (hip={getattr(torch.version, 'hip', None)}, cuda={cuda_available})"
|
||||
)
|
||||
raise RuntimeError(f"GPU requested ({requested_device}) but not available.")
|
||||
|
||||
if requested_device == "cpu":
|
||||
return "cpu"
|
||||
|
||||
raise ValueError(
|
||||
f"Unsupported device '{requested_device}'. Use 'cpu', 'cuda', or 'auto'."
|
||||
)
|
||||
return "cpu"
|
||||
|
||||
device = resolve_device(device)
|
||||
|
||||
@@ -92,7 +125,17 @@ def train(
|
||||
device=device,
|
||||
tensorboard_log=str(Path(__file__).resolve().parent / "tensorboard"),
|
||||
)
|
||||
model.learn(total_timesteps=total_timesteps)
|
||||
|
||||
save_dir = str(Path(model_path).parent)
|
||||
model_prefix = Path(model_path).stem
|
||||
|
||||
milestone_cb = MilestoneCheckpointCallback(
|
||||
save_path=save_dir,
|
||||
name_prefix=model_prefix,
|
||||
step_interval=100_000
|
||||
)
|
||||
|
||||
model.learn(total_timesteps=total_timesteps, callback=milestone_cb)
|
||||
|
||||
Path(model_path).parent.mkdir(parents=True, exist_ok=True)
|
||||
model.save(model_path)
|
||||
@@ -108,6 +151,7 @@ if __name__ == "__main__":
|
||||
device=args.device,
|
||||
use_gui=args.use_gui,
|
||||
num_robots=args.num_robots,
|
||||
num_workers=args.num_workers,
|
||||
robot_spacing=args.robot_spacing,
|
||||
start_pose=args.start_pose,
|
||||
)
|
||||
Reference in New Issue
Block a user