#!/usr/bin/env python3 """Train Go1 go2style flat-terrain locomotion. Usage: uv run scripts/train_go2style.py uv run scripts/train_go2style.py --rllib rslrl """ import logging # IMPORTANT: trigger registration of go2style env + rl config import motrix_rl.tasks.go1_go2style # noqa: F401 # Now run the standard training pipeline from absl import app, flags from skrl import config from motrix_rl import utils logger = logging.getLogger(__name__) _ENV = flags.DEFINE_string("env", "go1-flat-terrain-walk-go2style", "The env to train") _SIM_BACKEND = flags.DEFINE_string("sim-backend", None, "Simulation backend") _NUM_ENVS = flags.DEFINE_integer("num-envs", 2048, "Number of envs") _RENDER = flags.DEFINE_bool("render", False, "Render the env") _TRAIN_BACKEND = flags.DEFINE_string("train-backend", None, "learning backend (jax/torch)") _SEED = flags.DEFINE_integer("seed", None, "Random seed") _RAND_SEED = flags.DEFINE_bool("rand-seed", False, "Generate random seed") _RLLIB = flags.DEFINE_string("rllib", "skrl", "RL framework (skrl/rslrl)") def get_train_backend(supports, train_backend_arg, rllib): if rllib == "rslrl": if train_backend_arg is not None and train_backend_arg != "torch": raise Exception("RSLRL only supports PyTorch backend.") if not supports.torch: raise Exception("RSLRL requires PyTorch.") return "torch" if train_backend_arg is not None: backend = train_backend_arg if backend == "jax" and not supports.jax: raise Exception("JAX not available.") if backend == "torch" and not supports.torch: raise Exception("PyTorch not available.") return backend if supports.jax and supports.jax_gpu: return "jax" elif supports.torch and supports.torch_gpu: return "torch" elif supports.jax: return "jax" elif supports.torch: return "torch" else: raise Exception("Neither JAX nor PyTorch available.") def main(argv): device_supports = utils.get_device_supports() logger.info(device_supports) env_name = _ENV.value enable_render = _RENDER.value rl_override = {} if _NUM_ENVS.present: rl_override["num_envs"] = _NUM_ENVS.value if _RAND_SEED.value: rl_override["runner.seed"] = None elif _SEED.present: rl_override["runner.seed"] = _SEED.value sim_backend = _SIM_BACKEND.value rllib = _RLLIB.value train_backend = get_train_backend(device_supports, _TRAIN_BACKEND.value, rllib) if rllib == "rslrl": assert device_supports.torch from motrix_rl.rslrl.torch.train import ppo trainer = ppo.Trainer(env_name, sim_backend, cfg_override=rl_override, enable_render=enable_render) elif train_backend == "jax": from motrix_rl.skrl.jax.train import ppo config.jax.backend = "jax" trainer = ppo.Trainer(env_name, sim_backend, cfg_override=rl_override, enable_render=enable_render) elif train_backend == "torch": from motrix_rl.skrl.torch.train import ppo config.torch.backend = "torch" trainer = ppo.Trainer(env_name, sim_backend, cfg_override=rl_override, enable_render=enable_render) else: raise Exception(f"Unknown train backend: {train_backend}") trainer.train() if __name__ == "__main__": app.run(main)