96 lines
3.3 KiB
Python
96 lines
3.3 KiB
Python
#!/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)
|