Support Multi-GPU Training (#22)
- Change in isaaclab_env wrapper to explicitly state GPU for each simulation - Removing jax cache to support multi-gpu environment launch in MuJoCo Playground - Removing .train() and .eval() in evaluation and rendering to avoid deadlock in multi-gpu training - Supporting synchronous normalization for multi-gpu training
This commit is contained in:
@@ -2,13 +2,6 @@ from typing import Optional
|
||||
|
||||
import gymnasium as gym
|
||||
import torch
|
||||
from isaaclab.app import AppLauncher
|
||||
|
||||
app_launcher = AppLauncher(headless=True)
|
||||
simulation_app = app_launcher.app
|
||||
|
||||
import isaaclab_tasks
|
||||
from isaaclab_tasks.utils.parse_cfg import parse_env_cfg
|
||||
|
||||
|
||||
class IsaacLabEnv:
|
||||
@@ -22,6 +15,14 @@ class IsaacLabEnv:
|
||||
seed: int,
|
||||
action_bounds: Optional[float] = None,
|
||||
):
|
||||
from isaaclab.app import AppLauncher
|
||||
|
||||
app_launcher = AppLauncher(headless=True, device=device)
|
||||
simulation_app = app_launcher.app
|
||||
|
||||
import isaaclab_tasks
|
||||
from isaaclab_tasks.utils.parse_cfg import parse_env_cfg
|
||||
|
||||
env_cfg = parse_env_cfg(
|
||||
task_name,
|
||||
device=device,
|
||||
|
||||
@@ -6,7 +6,15 @@ import mujoco
|
||||
|
||||
|
||||
class PlaygroundEvalEnvWrapper:
|
||||
def __init__(self, eval_env, max_episode_steps, env_name, num_eval_envs, seed):
|
||||
def __init__(
|
||||
self,
|
||||
eval_env,
|
||||
max_episode_steps,
|
||||
env_name,
|
||||
num_eval_envs,
|
||||
seed,
|
||||
device_rank=None,
|
||||
):
|
||||
"""
|
||||
Wrapper used for evaluation / rendering environments.
|
||||
Note that this is different from training environments that are
|
||||
@@ -24,6 +32,11 @@ class PlaygroundEvalEnvWrapper:
|
||||
self.asymmetric_obs = False
|
||||
|
||||
self.key = jax.random.PRNGKey(seed)
|
||||
|
||||
if device_rank is not None:
|
||||
gpu_devices = jax.devices("gpu")
|
||||
self.key = jax.device_put(self.key, gpu_devices[device_rank])
|
||||
|
||||
self.key_reset = jax.random.split(self.key, num_eval_envs)
|
||||
self.max_episode_steps = max_episode_steps
|
||||
|
||||
@@ -118,7 +131,12 @@ def make_env(
|
||||
eval_env_cfg.push_config.magnitude_range = [0.0, 0.0]
|
||||
eval_env = registry.load(env_name, config=eval_env_cfg)
|
||||
eval_env = PlaygroundEvalEnvWrapper(
|
||||
eval_env, eval_env_cfg.episode_length, env_name, num_eval_envs, seed
|
||||
eval_env,
|
||||
eval_env_cfg.episode_length,
|
||||
env_name,
|
||||
num_eval_envs,
|
||||
seed,
|
||||
device_rank=device_rank,
|
||||
)
|
||||
|
||||
render_env_cfg = registry.get_default_config(env_name)
|
||||
@@ -127,7 +145,12 @@ def make_env(
|
||||
render_env_cfg.push_config.magnitude_range = [0.0, 0.0]
|
||||
render_env = registry.load(env_name, config=render_env_cfg)
|
||||
render_env = PlaygroundEvalEnvWrapper(
|
||||
render_env, render_env_cfg.episode_length, env_name, 1, seed
|
||||
render_env,
|
||||
render_env_cfg.episode_length,
|
||||
env_name,
|
||||
1,
|
||||
seed,
|
||||
device_rank=device_rank,
|
||||
)
|
||||
|
||||
return train_env, eval_env, render_env
|
||||
|
||||
Reference in New Issue
Block a user