Skip to content

Commit db09c9a

Browse files
chengruizfan-ziqi
authored andcommitted
fix(cusrl): fix cusrl train and play scripts; update cusrl distillation configuration
1 parent 09f6a9d commit db09c9a

3 files changed

Lines changed: 10 additions & 8 deletions

File tree

scripts/reinforcement_learning/cusrl/play.py

Lines changed: 6 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -11,10 +11,6 @@
1111

1212
from isaaclab.app import AppLauncher
1313

14-
# local imports
15-
sys.path.append(os.path.abspath(os.path.join(os.path.dirname(__file__), "..")))
16-
from rl_utils import camera_follow
17-
1814
# add argparse arguments
1915
parser = argparse.ArgumentParser(description="Evaluate an RL agent with CusRL.")
2016
parser.add_argument("--video", action="store_true", default=False, help="Record videos during training.")
@@ -56,7 +52,6 @@
5652

5753
import cusrl
5854
import gymnasium as gym
59-
import robot_lab.tasks # noqa: F401
6055
import torch
6156
from cusrl.environment.isaaclab import TrainerCfg
6257

@@ -73,6 +68,12 @@
7368

7469
from isaaclab_tasks.utils.hydra import hydra_task_config # noqa: F401
7570

71+
import robot_lab.tasks # noqa: F401 # isort: skip
72+
73+
# local imports
74+
sys.path.append(os.path.abspath(os.path.join(os.path.dirname(__file__), "..")))
75+
from rl_utils import camera_follow
76+
7677

7778
class CameraFollowPlayerHook(cusrl.Player.Hook):
7879
def step(self, step: int, transition: dict, metrics: dict):

scripts/reinforcement_learning/cusrl/train.py

Lines changed: 3 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -58,7 +58,6 @@
5858

5959
import cusrl
6060
import gymnasium as gym
61-
import robot_lab.tasks # noqa: F401
6261
import torch
6362
from cusrl.environment.isaaclab import TrainerCfg
6463

@@ -73,6 +72,8 @@
7372

7473
from isaaclab_tasks.utils.hydra import hydra_task_config # noqa: F401
7574

75+
import robot_lab.tasks # noqa: F401 # isort: skip
76+
7677
torch.backends.cuda.matmul.allow_tf32 = True
7778
torch.backends.cudnn.allow_tf32 = True
7879
torch.backends.cudnn.deterministic = False
@@ -128,7 +129,7 @@ def main(env_cfg: ManagerBasedRLEnvCfg | DirectRLEnvCfg | DirectMARLEnvCfg, agen
128129
agent_factory=agent_cfg.agent_factory.override(
129130
device=args_cli.device, autocast=args_cli.autocast, compile=args_cli.compile
130131
),
131-
logger_factory=cusrl.make_logger_factory(args_cli.logger, log_dir, add_datetime_prefix=False),
132+
logger_factory=cusrl.make_logger_factory(args_cli.logger, log_dir, name=None),
132133
num_iterations=agent_cfg.max_iterations,
133134
save_interval=agent_cfg.save_interval,
134135
checkpoint_path=args_cli.checkpoint,

source/robot_lab/robot_lab/tasks/manager_based/locomotion/velocity/config/quadruped/anymal_d/agents/cusrl_distillation_cfg.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -28,7 +28,7 @@ class AnymalDFlatDistillationTrainerCfg(TrainerCfg):
2828
hooks=[
2929
cusrl.hook.ModuleInitialization(init_actor=False, init_critic=False, distribution_std=0.1),
3030
cusrl.hook.OnPolicyPreparation(),
31-
cusrl.hook.PolicyDistillationLoss(""),
31+
cusrl.hook.PolicyDistillation(""),
3232
cusrl.hook.GradientClipping(1.0),
3333
],
3434
)

0 commit comments

Comments
 (0)