MJLabWrapper#
- torchrl.envs.MJLabWrapper(*args, **kwargs)[source]#
TorchRL wrapper for a pre-built
mjlab.envs.ManagerBasedRlEnv.- Parameters:
env – The mjlab manager-based RL environment to wrap.
device – Torch device for actions and returned tensors. When
None, it is inferred fromenv.device. The TorchRL device must match the mjlab simulation device; a mismatch raises an error instead of silently copying tensors across devices in the hot path.batch_size – Batch size of the wrapper. When
None, it is inferred as[env.num_envs]. mjlab exposes a flat vectorized batch, so only one-dimensional batch sizes are accepted.native_autoreset – If
False(default), TorchRL drives resets through the standard"_reset"mask by settingenv.cfg.auto_resettoFalseand callingenv.reset(env_ids=...)on done rows. IfTrue, mjlab’s native auto-reset is kept and TorchRL marks the terminal"next"observation as invalid while carrying mjlab’s post-reset observation into the next root tensordict, matching the native-auto-reset contract used by Isaac Lab.from_pixels – If
True, append RGB pixels underpixels_keyat reset and step time. If the mjlab scene has an RGBCameraSensor, TorchRL uses the sensor’s batched output with shape[num_envs, H, W, 3]. Otherwise it falls back torender(), which requiresnum_envs == 1andrender_mode="rgb_array"because mjlab viewer rendering returns one frame for the whole scene rather than one image per environment row.pixels_only – If
True, return only the pixel observation. Requiresfrom_pixels=True.pixels_key – TensorDict key for pixel observations. Defaults to
"pixels".pixels_sensor – Name of the mjlab RGB
CameraSensorto use forfrom_pixels=True. IfNoneand exactly one RGB camera sensor is present, it is selected automatically.allow_done_after_reset – Passed to
EnvBase. Defaults toFalse.log_extras – If
True, accumulate the numeric entries of theextras["log"]mapping that mjlab returns fromresetandstep(episode reward terms, curriculum levels, termination counts, …). Callpop_logged_extras()to retrieve their running means and clear the accumulator, typically once per training iteration. Entries are averaged over the calls in which they appeared, so metrics that mjlab emits only when an episode ends are not diluted by silent steps. Defaults toFalse.
mjlab reference: Zakka et al., “mjlab: A Lightweight Framework for GPU-Accelerated Robot Learning”, arXiv:2601.22074.
Examples
>>> from mjlab.envs import ManagerBasedRlEnv >>> from mjlab.tasks.registry import load_env_cfg >>> from torchrl.envs import MJLabWrapper >>> cfg = load_env_cfg("Mjlab-Cartpole-Balance") >>> cfg.scene.num_envs = 4 >>> base_env = ManagerBasedRlEnv(cfg, device="cuda:0") >>> env = MJLabWrapper(base_env) >>> td = env.rollout(10)