fix: clamp std before distribution, lower init_noise to 0.5, NaN guard
This commit is contained in:
@@ -35,6 +35,9 @@ _TRAIN_BACKEND = flags.DEFINE_string("train-backend", None, "The learning backen
|
||||
_SEED = flags.DEFINE_integer("seed", None, "Random seed for reproducibility")
|
||||
_RAND_SEED = flags.DEFINE_bool("rand-seed", False, "Generate random seed")
|
||||
_RLLIB = flags.DEFINE_string("rllib", "skrl", "The RL framework (skrl/rslrl)")
|
||||
_CHECKPOINT = flags.DEFINE_string("checkpoint", None, "Resume training from a checkpoint (.pickle/.pt)")
|
||||
_FORCE_PHASE = flags.DEFINE_integer("force-phase", None, "Lock curriculum to a specific phase (0=flat,1=rough,2=stairs,3=mixed)")
|
||||
_TRACKING_LINVEL_SCALE = flags.DEFINE_float("tracking-linvel-scale", None, "Override tracking_lin_vel reward scale")
|
||||
|
||||
|
||||
def get_train_backend(supports: utils.DeviceSupports, train_backend_arg: str | None, rllib: str):
|
||||
@@ -104,6 +107,15 @@ def main(argv):
|
||||
# Determine the training backend
|
||||
train_backend = get_train_backend(device_supports, _TRAIN_BACKEND.value, rllib)
|
||||
|
||||
# Build env config overrides from command-line flags
|
||||
env_cfg_override = {}
|
||||
if _FORCE_PHASE.present:
|
||||
env_cfg_override["force_phase"] = _FORCE_PHASE.value
|
||||
if _TRACKING_LINVEL_SCALE.present:
|
||||
env_cfg_override["tracking_lin_vel_scale"] = _TRACKING_LINVEL_SCALE.value
|
||||
if not env_cfg_override:
|
||||
env_cfg_override = None
|
||||
|
||||
trainer = None
|
||||
if rllib == "rslrl":
|
||||
# RSLRL training flow
|
||||
@@ -111,22 +123,25 @@ def main(argv):
|
||||
assert train_backend == "torch", "RSLRL only supports PyTorch backend"
|
||||
from motrix_rl.rslrl.torch.train import ppo
|
||||
|
||||
trainer = ppo.Trainer(env_name, sim_backend, cfg_override=rl_override, enable_render=enable_render)
|
||||
trainer = ppo.Trainer(env_name, sim_backend, cfg_override=rl_override,
|
||||
enable_render=enable_render, env_cfg_override=env_cfg_override)
|
||||
|
||||
elif train_backend == "jax":
|
||||
from motrix_rl.skrl.jax.train import ppo
|
||||
|
||||
config.jax.backend = "jax" # or "numpy"
|
||||
trainer = ppo.Trainer(env_name, sim_backend, cfg_override=rl_override, enable_render=enable_render)
|
||||
trainer = ppo.Trainer(env_name, sim_backend, cfg_override=rl_override,
|
||||
enable_render=enable_render, env_cfg_override=env_cfg_override)
|
||||
|
||||
elif train_backend == "torch":
|
||||
from motrix_rl.skrl.torch.train import ppo
|
||||
|
||||
trainer = ppo.Trainer(env_name, sim_backend, cfg_override=rl_override, enable_render=enable_render)
|
||||
trainer = ppo.Trainer(env_name, sim_backend, cfg_override=rl_override,
|
||||
enable_render=enable_render, env_cfg_override=env_cfg_override)
|
||||
else:
|
||||
raise Exception(f"Unknown train backend: {train_backend}")
|
||||
|
||||
trainer.train()
|
||||
trainer.train(checkpoint=_CHECKPOINT.value)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
|
||||
Reference in New Issue
Block a user