fix: clamp std before distribution, lower init_noise to 0.5, NaN guard
This commit is contained in:
42
scripts/train_cts.py
Normal file
42
scripts/train_cts.py
Normal file
@@ -0,0 +1,42 @@
|
||||
#!/usr/bin/env python3
|
||||
"""Train CTS (Concurrent Teacher-Student) Go1 locomotion.
|
||||
|
||||
Usage:
|
||||
uv run scripts/train_cts.py
|
||||
uv run scripts/train_cts.py --num-envs 512
|
||||
"""
|
||||
import logging
|
||||
import motrix_rl.tasks.go1_go2style # noqa: triggers env + rlcfg registration
|
||||
|
||||
from absl import app, flags
|
||||
from skrl import config as skrl_config
|
||||
|
||||
from motrix_rl import utils
|
||||
from motrix_rl.skrl.jax.train.cts_ppo import CTSTrainer
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
_ENV = flags.DEFINE_string("env", "go1-cts-flat-walk-go2style", "CTS env to train")
|
||||
_NUM_ENVS = flags.DEFINE_integer("num-envs", 1024, "Number of environments")
|
||||
_SEED = flags.DEFINE_integer("seed", None, "Random seed")
|
||||
|
||||
|
||||
def main(argv):
|
||||
supports = utils.get_device_supports()
|
||||
logger.info(supports)
|
||||
|
||||
env_name = _ENV.value
|
||||
override = {}
|
||||
if _NUM_ENVS.present:
|
||||
override["num_envs"] = _NUM_ENVS.value
|
||||
if _SEED.present:
|
||||
override["runner.seed"] = _SEED.value
|
||||
|
||||
skrl_config.jax.backend = "jax"
|
||||
|
||||
trainer = CTSTrainer(env_name=env_name, cfg_override=override)
|
||||
trainer.train()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
app.run(main)
|
||||
Reference in New Issue
Block a user