43 lines
1.1 KiB
Python
43 lines
1.1 KiB
Python
#!/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)
|