fix fix_commands bug in play.py
This commit is contained in:
@@ -74,6 +74,7 @@ from isaaclab.envs import (
|
|||||||
ManagerBasedRLEnvCfg,
|
ManagerBasedRLEnvCfg,
|
||||||
multi_agent_to_single_agent,
|
multi_agent_to_single_agent,
|
||||||
)
|
)
|
||||||
|
from isaaclab.envs.mdp import UniformVelocityCommandCfg
|
||||||
from isaaclab.utils.assets import retrieve_file_path
|
from isaaclab.utils.assets import retrieve_file_path
|
||||||
from isaaclab.utils.dict import print_dict
|
from isaaclab.utils.dict import print_dict
|
||||||
# from isaaclab.utils.pretrained_checkpoint import get_published_pretrained_checkpoint
|
# from isaaclab.utils.pretrained_checkpoint import get_published_pretrained_checkpoint
|
||||||
@@ -83,12 +84,35 @@ from isaaclab_tasks.utils.hydra import hydra_task_config
|
|||||||
import robot_lab.tasks # noqa: F401
|
import robot_lab.tasks # noqa: F401
|
||||||
|
|
||||||
def fix_commands(env_cfg: ManagerBasedRLEnvCfg | DirectRLEnvCfg | DirectMARLEnvCfg):
|
def fix_commands(env_cfg: ManagerBasedRLEnvCfg | DirectRLEnvCfg | DirectMARLEnvCfg):
|
||||||
from isaaclab.managers import ObservationTermCfg as ObsTerm
|
"""Fix commanded velocity to a constant target.
|
||||||
num_envs = env_cfg.scene.num_envs
|
|
||||||
def fix(env):
|
Prefer locking the command generator to keep policy observations consistent
|
||||||
return torch.tensor([1.0, 0.0, 0.0], device=env_cfg.sim.device).repeat(num_envs, 1)
|
with environment internal command state.
|
||||||
env_cfg.observations.policy.velocity_commands = ObsTerm(func=fix)
|
"""
|
||||||
env_cfg.observations.critic.velocity_commands = ObsTerm(func=fix)
|
fixed_lin_x, fixed_lin_y, fixed_ang_z = 1.0, 0.0, 0.0
|
||||||
|
|
||||||
|
base_velocity_cfg = getattr(getattr(env_cfg, "commands", None), "base_velocity", None)
|
||||||
|
if base_velocity_cfg is None:
|
||||||
|
return
|
||||||
|
|
||||||
|
fixed_cfg = UniformVelocityCommandCfg(
|
||||||
|
asset_name=getattr(base_velocity_cfg, "asset_name", "robot"),
|
||||||
|
heading_command=False,
|
||||||
|
rel_standing_envs=0.0,
|
||||||
|
rel_heading_envs=0.0,
|
||||||
|
resampling_time_range=(5.0, 5.0),
|
||||||
|
ranges=UniformVelocityCommandCfg.Ranges(
|
||||||
|
lin_vel_x=(fixed_lin_x, fixed_lin_x),
|
||||||
|
lin_vel_y=(fixed_lin_y, fixed_lin_y),
|
||||||
|
ang_vel_z=(fixed_ang_z, fixed_ang_z),
|
||||||
|
heading=None,
|
||||||
|
),
|
||||||
|
)
|
||||||
|
env_cfg.commands.base_velocity = fixed_cfg
|
||||||
|
|
||||||
|
# terrain_levels_vel_gym expects custom Go2RLGymCommand fields.
|
||||||
|
if hasattr(env_cfg, "curriculum") and hasattr(env_cfg.curriculum, "terrain_levels"):
|
||||||
|
env_cfg.curriculum.terrain_levels = None
|
||||||
|
|
||||||
@hydra_task_config(args_cli.task, args_cli.agent)
|
@hydra_task_config(args_cli.task, args_cli.agent)
|
||||||
def main(env_cfg: ManagerBasedRLEnvCfg | DirectRLEnvCfg | DirectMARLEnvCfg, agent_cfg: RslRlBaseRunnerCfg):
|
def main(env_cfg: ManagerBasedRLEnvCfg | DirectRLEnvCfg | DirectMARLEnvCfg, agent_cfg: RslRlBaseRunnerCfg):
|
||||||
@@ -109,7 +133,7 @@ def main(env_cfg: ManagerBasedRLEnvCfg | DirectRLEnvCfg | DirectMARLEnvCfg, agen
|
|||||||
env_cfg.observations.policy.enable_corruption = False
|
env_cfg.observations.policy.enable_corruption = False
|
||||||
# remove random pushing
|
# remove random pushing
|
||||||
env_cfg.events.randomize_apply_external_force_torque = None
|
env_cfg.events.randomize_apply_external_force_torque = None
|
||||||
env_cfg.events.push_robot = None
|
env_cfg.events.randomize_push_robot = None
|
||||||
env_cfg.curriculum.command_levels_lin_vel = None
|
env_cfg.curriculum.command_levels_lin_vel = None
|
||||||
env_cfg.curriculum.command_levels_ang_vel = None
|
env_cfg.curriculum.command_levels_ang_vel = None
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user