43 lines
1.2 KiB
Python
43 lines
1.2 KiB
Python
#!/usr/bin/env python3
|
|
"""DreamWaQ training — Manaro-Alpha aligned.
|
|
|
|
Usage:
|
|
uv run scripts/train_dreamwaq.py # default (2048 envs, 100M steps)
|
|
uv run scripts/train_dreamwaq.py --num-envs 4096 --timesteps 150M
|
|
"""
|
|
import argparse
|
|
|
|
# Register env + config
|
|
import motrix_envs.locomotion.go1.dreamwaq # noqa
|
|
import motrix_rl.tasks.go1_dreamwaq # noqa
|
|
from motrix_rl.skrl.jax.train.dreamwaq_ppo import DreamWaQTrainer
|
|
|
|
|
|
def main():
|
|
p = argparse.ArgumentParser()
|
|
p.add_argument("--num-envs", type=int, default=2048)
|
|
p.add_argument("--timesteps", type=str, default="100M")
|
|
p.add_argument("--seed", type=int, default=42)
|
|
args = p.parse_args()
|
|
|
|
ts = args.timesteps
|
|
if ts.endswith("M"): ts = int(float(ts[:-1]) * 1_000_000)
|
|
elif ts.endswith("K"): ts = int(float(ts[:-1]) * 1_000)
|
|
else: ts = int(ts)
|
|
|
|
# SKRL timesteps = env.step() calls, NOT individual env steps
|
|
skrl_ts = ts // args.num_envs
|
|
|
|
override = {
|
|
"num_envs": args.num_envs,
|
|
"runner.seed": args.seed,
|
|
"runner.trainer.timesteps": skrl_ts,
|
|
}
|
|
|
|
trainer = DreamWaQTrainer(env_name="go1-dreamwaq-walk", cfg_override=override)
|
|
trainer.train()
|
|
|
|
|
|
if __name__ == "__main__":
|
|
main()
|