fix: clamp std before distribution, lower init_noise to 0.5, NaN guard

This commit is contained in:
8x54zj-m
2026-06-30 15:35:58 +08:00
parent f2a8e0e2ff
commit b0f1da4596
61 changed files with 6522 additions and 45 deletions

95
scripts/train_go2style.py Normal file
View File

@@ -0,0 +1,95 @@
#!/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)