v1.0.2-rc1; mv rem_cts to moe_cts, moe_cts to moe_no_goal_cts, fix rem params same as moe, add --robogauge to start
This commit is contained in:
@@ -1,3 +1,9 @@
|
|||||||
|
# 20260126
|
||||||
|
## v1.0.2-rc1
|
||||||
|
1. 修改高速移动的训练文件到最终版,删除配置中无用注释
|
||||||
|
2. 将`rem_cts`重命名为`moe_cts`,原`moe_cts`重命名为`moe_no_goal_cts`
|
||||||
|
3. 并发现之前的rem实现中,expert的共享权重网络只有128而之前是256,修改和moe参数量完全一致
|
||||||
|
4. 加入`--robogauge`选择开启robogauge评测的命令行参数,默认关闭,以及`--robogauge_port`选择端口号,默认9973
|
||||||
# 20260117
|
# 20260117
|
||||||
## v1.0.1
|
## v1.0.1
|
||||||
1. 加入日志数据读取脚本`logs_merge.py`,压缩日志`logs_compress.py`工具
|
1. 加入日志数据读取脚本`logs_merge.py`,压缩日志`logs_compress.py`工具
|
||||||
|
|||||||
@@ -1,4 +1,5 @@
|
|||||||
policy_path: "{LEGGED_GYM_ROOT_DIR}/deploy/pre_train/go2/go2_cts_150k.pt"
|
policy_path: "{LEGGED_GYM_ROOT_DIR}/deploy/pre_train/go2/go2_moe_cts_137000_0.6365.pt"
|
||||||
|
# policy_path: "{LEGGED_GYM_ROOT_DIR}/deploy/pre_train/go2/go2_cts_150k.pt"
|
||||||
|
|
||||||
# xml_path: "{LEGGED_GYM_ROOT_DIR}/resources/robots/go2/flat.xml"
|
# xml_path: "{LEGGED_GYM_ROOT_DIR}/resources/robots/go2/flat.xml"
|
||||||
# xml_path: "{LEGGED_GYM_ROOT_DIR}/resources/robots/go2/race_track.xml" # change go2 init pos to pos="-5 2 0.445" in resources/robots/go2/go2.xml
|
# xml_path: "{LEGGED_GYM_ROOT_DIR}/resources/robots/go2/race_track.xml" # change go2 init pos to pos="-5 2 0.445" in resources/robots/go2/go2.xml
|
||||||
|
|||||||
BIN
deploy/pre_train/go2/go2_moe_cts_137000_0.6365.pt
Normal file
BIN
deploy/pre_train/go2/go2_moe_cts_137000_0.6365.pt
Normal file
Binary file not shown.
@@ -1,7 +1,7 @@
|
|||||||
from legged_gym import LEGGED_GYM_ROOT_DIR, LEGGED_GYM_ENVS_DIR
|
from legged_gym import LEGGED_GYM_ROOT_DIR, LEGGED_GYM_ENVS_DIR
|
||||||
|
|
||||||
from legged_gym.envs.go2.go2_env import Go2Robot
|
from legged_gym.envs.go2.go2_env import Go2Robot
|
||||||
from legged_gym.envs.go2.go2_config import GO2Cfg, GO2CfgPPO, GO2CfgCTS, GO2CfgMoECTS, GO2CfgMCPCTS, GO2CfgACMoECTS, GO2CfgDualMoECTS, GO2CfgREMCTS
|
from legged_gym.envs.go2.go2_config import GO2Cfg, GO2CfgPPO, GO2CfgCTS, GO2CfgMoECTS, GO2CfgMoENGCTS, GO2CfgMCPCTS, GO2CfgACMoECTS, GO2CfgDualMoECTS
|
||||||
from .base.legged_robot import LeggedRobot
|
from .base.legged_robot import LeggedRobot
|
||||||
|
|
||||||
from legged_gym.utils.task_registry import task_registry
|
from legged_gym.utils.task_registry import task_registry
|
||||||
@@ -9,7 +9,7 @@ from legged_gym.utils.task_registry import task_registry
|
|||||||
task_registry.register("go2", Go2Robot, GO2Cfg(), GO2CfgPPO())
|
task_registry.register("go2", Go2Robot, GO2Cfg(), GO2CfgPPO())
|
||||||
task_registry.register("go2_cts", Go2Robot, GO2Cfg(), GO2CfgCTS())
|
task_registry.register("go2_cts", Go2Robot, GO2Cfg(), GO2CfgCTS())
|
||||||
task_registry.register("go2_moe_cts", Go2Robot, GO2Cfg(), GO2CfgMoECTS())
|
task_registry.register("go2_moe_cts", Go2Robot, GO2Cfg(), GO2CfgMoECTS())
|
||||||
|
task_registry.register("go2_moe_ng_cts", Go2Robot, GO2Cfg(), GO2CfgMoENGCTS())
|
||||||
task_registry.register("go2_mcp_cts", Go2Robot, GO2Cfg(), GO2CfgMCPCTS())
|
task_registry.register("go2_mcp_cts", Go2Robot, GO2Cfg(), GO2CfgMCPCTS())
|
||||||
task_registry.register("go2_ac_moe_cts", Go2Robot, GO2Cfg(), GO2CfgACMoECTS())
|
task_registry.register("go2_ac_moe_cts", Go2Robot, GO2Cfg(), GO2CfgACMoECTS())
|
||||||
task_registry.register("go2_dual_moe_cts", Go2Robot, GO2Cfg(), GO2CfgDualMoECTS())
|
task_registry.register("go2_dual_moe_cts", Go2Robot, GO2Cfg(), GO2CfgDualMoECTS())
|
||||||
task_registry.register("go2_rem_cts", Go2Robot, GO2Cfg(), GO2CfgREMCTS())
|
|
||||||
|
|||||||
@@ -302,6 +302,10 @@ class LeggedRobotCfgPPO(BaseConfig):
|
|||||||
checkpoint = -1 # -1 = last saved model
|
checkpoint = -1 # -1 = last saved model
|
||||||
resume_path = None # updated from load_run and chkpt
|
resume_path = None # updated from load_run and chkpt
|
||||||
|
|
||||||
|
class robogauge:
|
||||||
|
enabled = False
|
||||||
|
port = 9973
|
||||||
|
|
||||||
class LeggedRobotCfgCTS(BaseConfig):
|
class LeggedRobotCfgCTS(BaseConfig):
|
||||||
seed = 0
|
seed = 0
|
||||||
runner_class_name = "OnPolicyRunnerCTS"
|
runner_class_name = "OnPolicyRunnerCTS"
|
||||||
@@ -350,7 +354,11 @@ class LeggedRobotCfgCTS(BaseConfig):
|
|||||||
checkpoint = -1 # -1 = last saved model
|
checkpoint = -1 # -1 = last saved model
|
||||||
resume_path = None # updated from load_run and chkpt
|
resume_path = None # updated from load_run and chkpt
|
||||||
|
|
||||||
class LeggedRobotCfgMoECTS(LeggedRobotCfgCTS):
|
class robogauge:
|
||||||
|
enabled = False
|
||||||
|
port = 9973
|
||||||
|
|
||||||
|
class LeggedRobotCfgMoENGCTS(LeggedRobotCfgCTS):
|
||||||
class policy(LeggedRobotCfgCTS.policy):
|
class policy(LeggedRobotCfgCTS.policy):
|
||||||
obs_no_goal_mask = None # mask for observation without goal inputs
|
obs_no_goal_mask = None # mask for observation without goal inputs
|
||||||
student_expert_num = 8 # number of experts in the student model
|
student_expert_num = 8 # number of experts in the student model
|
||||||
@@ -359,8 +367,8 @@ class LeggedRobotCfgMoECTS(LeggedRobotCfgCTS):
|
|||||||
load_balance_coef = 0.01 # coefficient for load balance loss
|
load_balance_coef = 0.01 # coefficient for load balance loss
|
||||||
|
|
||||||
class runner(LeggedRobotCfgCTS.runner):
|
class runner(LeggedRobotCfgCTS.runner):
|
||||||
policy_class_name = 'ActorCriticMoECTS'
|
policy_class_name = 'ActorCriticMoENGCTS'
|
||||||
algorithm_class_name = 'MoECTS'
|
algorithm_class_name = 'MoENGCTS'
|
||||||
|
|
||||||
class LeggedRobotCfgMCPCTS(LeggedRobotCfgCTS):
|
class LeggedRobotCfgMCPCTS(LeggedRobotCfgCTS):
|
||||||
class policy(LeggedRobotCfgCTS.policy):
|
class policy(LeggedRobotCfgCTS.policy):
|
||||||
@@ -382,20 +390,20 @@ class LeggedRobotCfgACMoECTS(LeggedRobotCfgCTS):
|
|||||||
class LeggedRobotCfgDualMoECTS(LeggedRobotCfgCTS):
|
class LeggedRobotCfgDualMoECTS(LeggedRobotCfgCTS):
|
||||||
class policy(LeggedRobotCfgCTS.policy):
|
class policy(LeggedRobotCfgCTS.policy):
|
||||||
expert_num = 8 # number of experts in the student model
|
expert_num = 8 # number of experts in the student model
|
||||||
student_encoder_hidden_dims = [512, 256, 128]
|
student_encoder_hidden_dims = [512, 256, 256]
|
||||||
|
|
||||||
class runner(LeggedRobotCfgCTS.runner):
|
class runner(LeggedRobotCfgCTS.runner):
|
||||||
policy_class_name = 'ActorCriticDualMoECTS'
|
policy_class_name = 'ActorCriticDualMoECTS'
|
||||||
algorithm_class_name = 'DualMoECTS'
|
algorithm_class_name = 'DualMoECTS'
|
||||||
|
|
||||||
class LeggedRobotCfgREMCTS(LeggedRobotCfgCTS):
|
class LeggedRobotCfgMoECTS(LeggedRobotCfgCTS):
|
||||||
class policy(LeggedRobotCfgCTS.policy):
|
class policy(LeggedRobotCfgCTS.policy):
|
||||||
expert_num = 8 # number of experts in the student model
|
expert_num = 8 # number of experts in the student model
|
||||||
student_encoder_hidden_dims = [512, 256, 128]
|
student_encoder_hidden_dims = [512, 256, 256]
|
||||||
|
|
||||||
class algorithm(LeggedRobotCfgCTS.algorithm):
|
class algorithm(LeggedRobotCfgCTS.algorithm):
|
||||||
load_balance_coef = 0.01 # coefficient for load balance loss
|
load_balance_coef = 0.01 # coefficient for load balance loss
|
||||||
|
|
||||||
class runner(LeggedRobotCfgCTS.runner):
|
class runner(LeggedRobotCfgCTS.runner):
|
||||||
policy_class_name = 'ActorCriticREMCTS'
|
policy_class_name = 'ActorCriticMoECTS'
|
||||||
algorithm_class_name = 'REMCTS'
|
algorithm_class_name = 'MoECTS'
|
||||||
|
|||||||
@@ -1,5 +1,5 @@
|
|||||||
import math
|
import math
|
||||||
from legged_gym.envs.base.legged_robot_config import LeggedRobotCfg, LeggedRobotCfgPPO, LeggedRobotCfgCTS, LeggedRobotCfgMoECTS, LeggedRobotCfgMoECTS, LeggedRobotCfgMCPCTS, LeggedRobotCfgACMoECTS, LeggedRobotCfgDualMoECTS, LeggedRobotCfgREMCTS
|
from legged_gym.envs.base.legged_robot_config import LeggedRobotCfg, LeggedRobotCfgPPO, LeggedRobotCfgCTS, LeggedRobotCfgMoENGCTS, LeggedRobotCfgMoENGCTS, LeggedRobotCfgMCPCTS, LeggedRobotCfgACMoECTS, LeggedRobotCfgDualMoECTS, LeggedRobotCfgMoECTS
|
||||||
|
|
||||||
class GO2Cfg(LeggedRobotCfg):
|
class GO2Cfg(LeggedRobotCfg):
|
||||||
class init_state(LeggedRobotCfg.init_state):
|
class init_state(LeggedRobotCfg.init_state):
|
||||||
@@ -228,17 +228,17 @@ class GO2CfgCTS(LeggedRobotCfgCTS):
|
|||||||
latent_dim = 32
|
latent_dim = 32
|
||||||
norm_type = 'l2norm'
|
norm_type = 'l2norm'
|
||||||
|
|
||||||
class GO2CfgMoECTS(LeggedRobotCfgMoECTS):
|
class GO2CfgMoENGCTS(LeggedRobotCfgMoENGCTS):
|
||||||
class policy(LeggedRobotCfgMoECTS.policy):
|
class policy(LeggedRobotCfgMoENGCTS.policy):
|
||||||
obs_no_goal_mask = [True] * 6 + [False] * 3 + [True] * 36 # mask for obs without command info
|
obs_no_goal_mask = [True] * 6 + [False] * 3 + [True] * 36 # mask for obs without command info
|
||||||
student_expert_num = 8 # number of experts in the student model
|
student_expert_num = 8 # number of experts in the student model
|
||||||
|
|
||||||
class algorithm(LeggedRobotCfgMoECTS.algorithm):
|
class algorithm(LeggedRobotCfgMoENGCTS.algorithm):
|
||||||
load_balance_coef = 0.01
|
load_balance_coef = 0.01
|
||||||
|
|
||||||
class runner(LeggedRobotCfgMoECTS.runner):
|
class runner(LeggedRobotCfgMoENGCTS.runner):
|
||||||
run_name = ''
|
run_name = ''
|
||||||
experiment_name = 'go2_moe_cts'
|
experiment_name = 'go2_moe_no_goal_cts'
|
||||||
max_iterations = 150000
|
max_iterations = 150000
|
||||||
save_interval = 500
|
save_interval = 500
|
||||||
|
|
||||||
@@ -273,12 +273,12 @@ class GO2CfgDualMoECTS(LeggedRobotCfgDualMoECTS):
|
|||||||
max_iterations = 150000
|
max_iterations = 150000
|
||||||
save_interval = 500
|
save_interval = 500
|
||||||
|
|
||||||
class GO2CfgREMCTS(LeggedRobotCfgREMCTS):
|
class GO2CfgMoECTS(LeggedRobotCfgMoECTS):
|
||||||
class policy(LeggedRobotCfgREMCTS.policy):
|
class policy(LeggedRobotCfgMoECTS.policy):
|
||||||
expert_num = 8 # number of experts in the student model
|
expert_num = 8 # number of experts in the student model
|
||||||
|
|
||||||
class runner(LeggedRobotCfgREMCTS.runner):
|
class runner(LeggedRobotCfgMoECTS.runner):
|
||||||
run_name = ''
|
run_name = ''
|
||||||
experiment_name = 'go2_rem_cts'
|
experiment_name = 'go2_moe_cts'
|
||||||
max_iterations = 150000
|
max_iterations = 150000
|
||||||
save_interval = 500
|
save_interval = 500
|
||||||
|
|||||||
@@ -9,7 +9,7 @@
|
|||||||
Change command_range_curriculum, init command range
|
Change command_range_curriculum, init command range
|
||||||
'''
|
'''
|
||||||
import math
|
import math
|
||||||
from legged_gym.envs.base.legged_robot_config import LeggedRobotCfg, LeggedRobotCfgPPO, LeggedRobotCfgCTS, LeggedRobotCfgMoECTS, LeggedRobotCfgMoECTS, LeggedRobotCfgMCPCTS, LeggedRobotCfgACMoECTS, LeggedRobotCfgDualMoECTS, LeggedRobotCfgREMCTS
|
from legged_gym.envs.base.legged_robot_config import LeggedRobotCfg, LeggedRobotCfgPPO, LeggedRobotCfgCTS, LeggedRobotCfgMoENGCTS, LeggedRobotCfgMoENGCTS, LeggedRobotCfgMCPCTS, LeggedRobotCfgACMoECTS, LeggedRobotCfgDualMoECTS, LeggedRobotCfgMoECTS
|
||||||
|
|
||||||
class GO2Cfg(LeggedRobotCfg):
|
class GO2Cfg(LeggedRobotCfg):
|
||||||
class init_state(LeggedRobotCfg.init_state):
|
class init_state(LeggedRobotCfg.init_state):
|
||||||
@@ -259,17 +259,17 @@ class GO2CfgCTS(LeggedRobotCfgCTS):
|
|||||||
latent_dim = 32
|
latent_dim = 32
|
||||||
norm_type = 'l2norm'
|
norm_type = 'l2norm'
|
||||||
|
|
||||||
class GO2CfgMoECTS(LeggedRobotCfgMoECTS):
|
class GO2CfgMoENGCTS(LeggedRobotCfgMoENGCTS):
|
||||||
class policy(LeggedRobotCfgMoECTS.policy):
|
class policy(LeggedRobotCfgMoENGCTS.policy):
|
||||||
obs_no_goal_mask = [True] * 6 + [False] * 3 + [True] * 36 # mask for obs without command info
|
obs_no_goal_mask = [True] * 6 + [False] * 3 + [True] * 36 # mask for obs without command info
|
||||||
student_expert_num = 8 # number of experts in the student model
|
student_expert_num = 8 # number of experts in the student model
|
||||||
|
|
||||||
class algorithm(LeggedRobotCfgMoECTS.algorithm):
|
class algorithm(LeggedRobotCfgMoENGCTS.algorithm):
|
||||||
load_balance_coef = 0.01
|
load_balance_coef = 0.01
|
||||||
|
|
||||||
class runner(LeggedRobotCfgMoECTS.runner):
|
class runner(LeggedRobotCfgMoENGCTS.runner):
|
||||||
run_name = ''
|
run_name = ''
|
||||||
experiment_name = 'go2_moe_cts'
|
experiment_name = 'go2_moe_no_goal_cts'
|
||||||
max_iterations = 150000
|
max_iterations = 150000
|
||||||
save_interval = 500
|
save_interval = 500
|
||||||
|
|
||||||
@@ -304,12 +304,12 @@ class GO2CfgDualMoECTS(LeggedRobotCfgDualMoECTS):
|
|||||||
max_iterations = 150000
|
max_iterations = 150000
|
||||||
save_interval = 500
|
save_interval = 500
|
||||||
|
|
||||||
class GO2CfgREMCTS(LeggedRobotCfgREMCTS):
|
class GO2CfgMoECTS(LeggedRobotCfgMoECTS):
|
||||||
class policy(LeggedRobotCfgREMCTS.policy):
|
class policy(LeggedRobotCfgMoECTS.policy):
|
||||||
expert_num = 8 # number of experts in the student model
|
expert_num = 8 # number of experts in the student model
|
||||||
|
|
||||||
class runner(LeggedRobotCfgREMCTS.runner):
|
class runner(LeggedRobotCfgMoECTS.runner):
|
||||||
run_name = ''
|
run_name = ''
|
||||||
experiment_name = 'go2_rem_cts'
|
experiment_name = 'go2_moe_cts'
|
||||||
max_iterations = 150000
|
max_iterations = 150000
|
||||||
save_interval = 500
|
save_interval = 500
|
||||||
|
|||||||
@@ -1,17 +1,17 @@
|
|||||||
# -*- coding: utf-8 -*-
|
# -*- coding: utf-8 -*-
|
||||||
'''
|
'''
|
||||||
@File : go2_config_vanilla.py
|
@File : go2_config_vanilla.py
|
||||||
@Time : 2026/01/10 02:26:04
|
@Time : 2026/01/10 02:27:28
|
||||||
@Author : wty-yy
|
@Author : wty-yy
|
||||||
@Version : 1.0
|
@Version : 1.0
|
||||||
@Blog : https://wty-yy.github.io/
|
@Blog : https://wty-yy.github.io/
|
||||||
@Desc : Go2 vanilla training config
|
@Desc : Go2 vanilla training config, same as unitree rl gym except domain randomization and rewards
|
||||||
episode length 25, resample commands 5 sec,
|
episode length 20, resample commands 10 sec,
|
||||||
open move_down_by_accumulated_xy_command, dynamic_resample_commands
|
open heading_command
|
||||||
close heading_command, zero_command_curriculum, limit_vel_prob, command_range_curriculum, dynamic_sigma
|
close move_down_by_accumulated_xy_command, dynamic_resample_commands, zero_command_curriculum, limit_vel_prob, command_range_curriculum, dynamic_sigma
|
||||||
'''
|
'''
|
||||||
import math
|
import math
|
||||||
from legged_gym.envs.base.legged_robot_config import LeggedRobotCfg, LeggedRobotCfgPPO, LeggedRobotCfgCTS, LeggedRobotCfgMoECTS, LeggedRobotCfgMoECTS, LeggedRobotCfgMCPCTS, LeggedRobotCfgACMoECTS, LeggedRobotCfgDualMoECTS, LeggedRobotCfgREMCTS
|
from legged_gym.envs.base.legged_robot_config import LeggedRobotCfg, LeggedRobotCfgPPO, LeggedRobotCfgCTS, LeggedRobotCfgMoENGCTS, LeggedRobotCfgMoENGCTS, LeggedRobotCfgMCPCTS, LeggedRobotCfgACMoECTS, LeggedRobotCfgDualMoECTS, LeggedRobotCfgMoECTS
|
||||||
|
|
||||||
class GO2Cfg(LeggedRobotCfg):
|
class GO2Cfg(LeggedRobotCfg):
|
||||||
class init_state(LeggedRobotCfg.init_state):
|
class init_state(LeggedRobotCfg.init_state):
|
||||||
@@ -48,7 +48,7 @@ class GO2Cfg(LeggedRobotCfg):
|
|||||||
num_privileged_obs = 45 + 3 + 4 + 12 + 12 + 187 # 263
|
num_privileged_obs = 45 + 3 + 4 + 12 + 12 + 187 # 263
|
||||||
# num_privileged_obs = 45 + 3 + 187 # 235
|
# num_privileged_obs = 45 + 3 + 187 # 235
|
||||||
# num_privileged_obs = 48 # without height measurements
|
# num_privileged_obs = 48 # without height measurements
|
||||||
episode_length_s = 25
|
episode_length_s = 20
|
||||||
|
|
||||||
class domain_rand(LeggedRobotCfg.domain_rand):
|
class domain_rand(LeggedRobotCfg.domain_rand):
|
||||||
### Robot properties ###
|
### Robot properties ###
|
||||||
@@ -105,14 +105,14 @@ class GO2Cfg(LeggedRobotCfg):
|
|||||||
# terrain_proportions = [0.0, 0.0, 0.0, 0.0, 1.0, 0.0, 0.0, 0.0, 0.0]
|
# terrain_proportions = [0.0, 0.0, 0.0, 0.0, 1.0, 0.0, 0.0, 0.0, 0.0]
|
||||||
# terrain_proportions = [0.3, 0.3, 0.3, 0.0, 0.0, 0.0, 0.0, 0.0, 0.1]
|
# terrain_proportions = [0.3, 0.3, 0.3, 0.0, 0.0, 0.0, 0.0, 0.0, 0.1]
|
||||||
# terrain_proportions = [0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 1.0]
|
# terrain_proportions = [0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 1.0]
|
||||||
move_down_by_accumulated_xy_command = True # move down the terrain curriculum based on accumulated xy command distance instead of absolute distance
|
move_down_by_accumulated_xy_command = False # move down the terrain curriculum based on accumulated xy command distance instead of absolute distance
|
||||||
|
|
||||||
class commands(LeggedRobotCfg.commands):
|
class commands(LeggedRobotCfg.commands):
|
||||||
curriculum = False
|
curriculum = False
|
||||||
max_curriculum = 1.
|
max_curriculum = 1.
|
||||||
num_commands = 4 # default: lin_vel_x, lin_vel_y, ang_vel_yaw (in heading mode ang_vel_yaw is recomputed from heading error)
|
num_commands = 4 # default: lin_vel_x, lin_vel_y, ang_vel_yaw (in heading mode ang_vel_yaw is recomputed from heading error)
|
||||||
resampling_time = 5. # time before command are changed[s]
|
resampling_time = 10. # time before command are changed[s]
|
||||||
heading_command = False # if true: compute ang vel command from heading error
|
heading_command = True # if true: compute ang vel command from heading error
|
||||||
# start training with zero commands and then gradually increase zero command probability
|
# start training with zero commands and then gradually increase zero command probability
|
||||||
zero_command_curriculum = None
|
zero_command_curriculum = None
|
||||||
# zero_command_curriculum = {'start_iter': 0, 'end_iter': 1500, 'start_value': 0.0, 'end_value': 0.1}
|
# zero_command_curriculum = {'start_iter': 0, 'end_iter': 1500, 'start_value': 0.0, 'end_value': 0.1}
|
||||||
@@ -121,7 +121,7 @@ class GO2Cfg(LeggedRobotCfg):
|
|||||||
limit_vel_invert_when_continuous = True # invert the limit logic when using continuous sample limit velocity commands
|
limit_vel_invert_when_continuous = True # invert the limit logic when using continuous sample limit velocity commands
|
||||||
limit_vel = {"lin_vel_x": [-1, 1], "lin_vel_y": [-1, 1], "ang_vel_yaw": [-1, 0, 1]} # sample vel commands from min [-1] or zero [0] or max [1] range only
|
limit_vel = {"lin_vel_x": [-1, 1], "lin_vel_y": [-1, 1], "ang_vel_yaw": [-1, 0, 1]} # sample vel commands from min [-1] or zero [0] or max [1] range only
|
||||||
stop_heading_at_limit = True # stop heading updates when vel is limited
|
stop_heading_at_limit = True # stop heading updates when vel is limited
|
||||||
dynamic_resample_commands = True # sample commands with low bounds
|
dynamic_resample_commands = False # sample commands with low bounds
|
||||||
command_range_curriculum = []
|
command_range_curriculum = []
|
||||||
# command_range_curriculum = [{ # list for command range curriculums at specific training iterations
|
# command_range_curriculum = [{ # list for command range curriculums at specific training iterations
|
||||||
# 'iter': 20000, # training iteration at which the command ranges are updated
|
# 'iter': 20000, # training iteration at which the command ranges are updated
|
||||||
@@ -243,17 +243,17 @@ class GO2CfgCTS(LeggedRobotCfgCTS):
|
|||||||
latent_dim = 32
|
latent_dim = 32
|
||||||
norm_type = 'l2norm'
|
norm_type = 'l2norm'
|
||||||
|
|
||||||
class GO2CfgMoECTS(LeggedRobotCfgMoECTS):
|
class GO2CfgMoENGCTS(LeggedRobotCfgMoENGCTS):
|
||||||
class policy(LeggedRobotCfgMoECTS.policy):
|
class policy(LeggedRobotCfgMoENGCTS.policy):
|
||||||
obs_no_goal_mask = [True] * 6 + [False] * 3 + [True] * 36 # mask for obs without command info
|
obs_no_goal_mask = [True] * 6 + [False] * 3 + [True] * 36 # mask for obs without command info
|
||||||
student_expert_num = 8 # number of experts in the student model
|
student_expert_num = 8 # number of experts in the student model
|
||||||
|
|
||||||
class algorithm(LeggedRobotCfgMoECTS.algorithm):
|
class algorithm(LeggedRobotCfgMoENGCTS.algorithm):
|
||||||
load_balance_coef = 0.01
|
load_balance_coef = 0.01
|
||||||
|
|
||||||
class runner(LeggedRobotCfgMoECTS.runner):
|
class runner(LeggedRobotCfgMoENGCTS.runner):
|
||||||
run_name = ''
|
run_name = ''
|
||||||
experiment_name = 'go2_moe_cts'
|
experiment_name = 'go2_moe_no_goal_cts'
|
||||||
max_iterations = 150000
|
max_iterations = 150000
|
||||||
save_interval = 500
|
save_interval = 500
|
||||||
|
|
||||||
@@ -288,12 +288,12 @@ class GO2CfgDualMoECTS(LeggedRobotCfgDualMoECTS):
|
|||||||
max_iterations = 150000
|
max_iterations = 150000
|
||||||
save_interval = 500
|
save_interval = 500
|
||||||
|
|
||||||
class GO2CfgREMCTS(LeggedRobotCfgREMCTS):
|
class GO2CfgMoECTS(LeggedRobotCfgMoECTS):
|
||||||
class policy(LeggedRobotCfgREMCTS.policy):
|
class policy(LeggedRobotCfgMoECTS.policy):
|
||||||
expert_num = 8 # number of experts in the student model
|
expert_num = 8 # number of experts in the student model
|
||||||
|
|
||||||
class runner(LeggedRobotCfgREMCTS.runner):
|
class runner(LeggedRobotCfgMoECTS.runner):
|
||||||
run_name = ''
|
run_name = ''
|
||||||
experiment_name = 'go2_rem_cts'
|
experiment_name = 'go2_moe_cts'
|
||||||
max_iterations = 150000
|
max_iterations = 150000
|
||||||
save_interval = 500
|
save_interval = 500
|
||||||
|
|||||||
@@ -1,17 +1,17 @@
|
|||||||
# -*- coding: utf-8 -*-
|
# -*- coding: utf-8 -*-
|
||||||
'''
|
'''
|
||||||
@File : go2_config_vanilla2.py
|
@File : go2_config_vanilla_with_dynamic_cmd.py
|
||||||
@Time : 2026/01/10 02:27:28
|
@Time : 2026/01/10 02:26:04
|
||||||
@Author : wty-yy
|
@Author : wty-yy
|
||||||
@Version : 1.0
|
@Version : 1.0
|
||||||
@Blog : https://wty-yy.github.io/
|
@Blog : https://wty-yy.github.io/
|
||||||
@Desc : Go2 vanilla2 training config, same as unitree rl gym except domain randomization and rewards
|
@Desc : Go2 vanilla training config
|
||||||
episode length 20, resample commands 10 sec,
|
episode length 25, resample commands 5 sec,
|
||||||
open heading_command
|
open move_down_by_accumulated_xy_command, dynamic_resample_commands
|
||||||
close move_down_by_accumulated_xy_command, dynamic_resample_commands, zero_command_curriculum, limit_vel_prob, command_range_curriculum, dynamic_sigma
|
close heading_command, zero_command_curriculum, limit_vel_prob, command_range_curriculum, dynamic_sigma
|
||||||
'''
|
'''
|
||||||
import math
|
import math
|
||||||
from legged_gym.envs.base.legged_robot_config import LeggedRobotCfg, LeggedRobotCfgPPO, LeggedRobotCfgCTS, LeggedRobotCfgMoECTS, LeggedRobotCfgMoECTS, LeggedRobotCfgMCPCTS, LeggedRobotCfgACMoECTS, LeggedRobotCfgDualMoECTS, LeggedRobotCfgREMCTS
|
from legged_gym.envs.base.legged_robot_config import LeggedRobotCfg, LeggedRobotCfgPPO, LeggedRobotCfgCTS, LeggedRobotCfgMoENGCTS, LeggedRobotCfgMoENGCTS, LeggedRobotCfgMCPCTS, LeggedRobotCfgACMoECTS, LeggedRobotCfgDualMoECTS, LeggedRobotCfgMoECTS
|
||||||
|
|
||||||
class GO2Cfg(LeggedRobotCfg):
|
class GO2Cfg(LeggedRobotCfg):
|
||||||
class init_state(LeggedRobotCfg.init_state):
|
class init_state(LeggedRobotCfg.init_state):
|
||||||
@@ -48,7 +48,7 @@ class GO2Cfg(LeggedRobotCfg):
|
|||||||
num_privileged_obs = 45 + 3 + 4 + 12 + 12 + 187 # 263
|
num_privileged_obs = 45 + 3 + 4 + 12 + 12 + 187 # 263
|
||||||
# num_privileged_obs = 45 + 3 + 187 # 235
|
# num_privileged_obs = 45 + 3 + 187 # 235
|
||||||
# num_privileged_obs = 48 # without height measurements
|
# num_privileged_obs = 48 # without height measurements
|
||||||
episode_length_s = 20
|
episode_length_s = 25
|
||||||
|
|
||||||
class domain_rand(LeggedRobotCfg.domain_rand):
|
class domain_rand(LeggedRobotCfg.domain_rand):
|
||||||
### Robot properties ###
|
### Robot properties ###
|
||||||
@@ -105,14 +105,14 @@ class GO2Cfg(LeggedRobotCfg):
|
|||||||
# terrain_proportions = [0.0, 0.0, 0.0, 0.0, 1.0, 0.0, 0.0, 0.0, 0.0]
|
# terrain_proportions = [0.0, 0.0, 0.0, 0.0, 1.0, 0.0, 0.0, 0.0, 0.0]
|
||||||
# terrain_proportions = [0.3, 0.3, 0.3, 0.0, 0.0, 0.0, 0.0, 0.0, 0.1]
|
# terrain_proportions = [0.3, 0.3, 0.3, 0.0, 0.0, 0.0, 0.0, 0.0, 0.1]
|
||||||
# terrain_proportions = [0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 1.0]
|
# terrain_proportions = [0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 1.0]
|
||||||
move_down_by_accumulated_xy_command = False # move down the terrain curriculum based on accumulated xy command distance instead of absolute distance
|
move_down_by_accumulated_xy_command = True # move down the terrain curriculum based on accumulated xy command distance instead of absolute distance
|
||||||
|
|
||||||
class commands(LeggedRobotCfg.commands):
|
class commands(LeggedRobotCfg.commands):
|
||||||
curriculum = False
|
curriculum = False
|
||||||
max_curriculum = 1.
|
max_curriculum = 1.
|
||||||
num_commands = 4 # default: lin_vel_x, lin_vel_y, ang_vel_yaw (in heading mode ang_vel_yaw is recomputed from heading error)
|
num_commands = 4 # default: lin_vel_x, lin_vel_y, ang_vel_yaw (in heading mode ang_vel_yaw is recomputed from heading error)
|
||||||
resampling_time = 10. # time before command are changed[s]
|
resampling_time = 5. # time before command are changed[s]
|
||||||
heading_command = True # if true: compute ang vel command from heading error
|
heading_command = False # if true: compute ang vel command from heading error
|
||||||
# start training with zero commands and then gradually increase zero command probability
|
# start training with zero commands and then gradually increase zero command probability
|
||||||
zero_command_curriculum = None
|
zero_command_curriculum = None
|
||||||
# zero_command_curriculum = {'start_iter': 0, 'end_iter': 1500, 'start_value': 0.0, 'end_value': 0.1}
|
# zero_command_curriculum = {'start_iter': 0, 'end_iter': 1500, 'start_value': 0.0, 'end_value': 0.1}
|
||||||
@@ -121,7 +121,7 @@ class GO2Cfg(LeggedRobotCfg):
|
|||||||
limit_vel_invert_when_continuous = True # invert the limit logic when using continuous sample limit velocity commands
|
limit_vel_invert_when_continuous = True # invert the limit logic when using continuous sample limit velocity commands
|
||||||
limit_vel = {"lin_vel_x": [-1, 1], "lin_vel_y": [-1, 1], "ang_vel_yaw": [-1, 0, 1]} # sample vel commands from min [-1] or zero [0] or max [1] range only
|
limit_vel = {"lin_vel_x": [-1, 1], "lin_vel_y": [-1, 1], "ang_vel_yaw": [-1, 0, 1]} # sample vel commands from min [-1] or zero [0] or max [1] range only
|
||||||
stop_heading_at_limit = True # stop heading updates when vel is limited
|
stop_heading_at_limit = True # stop heading updates when vel is limited
|
||||||
dynamic_resample_commands = False # sample commands with low bounds
|
dynamic_resample_commands = True # sample commands with low bounds
|
||||||
command_range_curriculum = []
|
command_range_curriculum = []
|
||||||
# command_range_curriculum = [{ # list for command range curriculums at specific training iterations
|
# command_range_curriculum = [{ # list for command range curriculums at specific training iterations
|
||||||
# 'iter': 20000, # training iteration at which the command ranges are updated
|
# 'iter': 20000, # training iteration at which the command ranges are updated
|
||||||
@@ -243,17 +243,17 @@ class GO2CfgCTS(LeggedRobotCfgCTS):
|
|||||||
latent_dim = 32
|
latent_dim = 32
|
||||||
norm_type = 'l2norm'
|
norm_type = 'l2norm'
|
||||||
|
|
||||||
class GO2CfgMoECTS(LeggedRobotCfgMoECTS):
|
class GO2CfgMoENGCTS(LeggedRobotCfgMoENGCTS):
|
||||||
class policy(LeggedRobotCfgMoECTS.policy):
|
class policy(LeggedRobotCfgMoENGCTS.policy):
|
||||||
obs_no_goal_mask = [True] * 6 + [False] * 3 + [True] * 36 # mask for obs without command info
|
obs_no_goal_mask = [True] * 6 + [False] * 3 + [True] * 36 # mask for obs without command info
|
||||||
student_expert_num = 8 # number of experts in the student model
|
student_expert_num = 8 # number of experts in the student model
|
||||||
|
|
||||||
class algorithm(LeggedRobotCfgMoECTS.algorithm):
|
class algorithm(LeggedRobotCfgMoENGCTS.algorithm):
|
||||||
load_balance_coef = 0.01
|
load_balance_coef = 0.01
|
||||||
|
|
||||||
class runner(LeggedRobotCfgMoECTS.runner):
|
class runner(LeggedRobotCfgMoENGCTS.runner):
|
||||||
run_name = ''
|
run_name = ''
|
||||||
experiment_name = 'go2_moe_cts'
|
experiment_name = 'go2_moe_no_goal_cts'
|
||||||
max_iterations = 150000
|
max_iterations = 150000
|
||||||
save_interval = 500
|
save_interval = 500
|
||||||
|
|
||||||
@@ -288,12 +288,12 @@ class GO2CfgDualMoECTS(LeggedRobotCfgDualMoECTS):
|
|||||||
max_iterations = 150000
|
max_iterations = 150000
|
||||||
save_interval = 500
|
save_interval = 500
|
||||||
|
|
||||||
class GO2CfgREMCTS(LeggedRobotCfgREMCTS):
|
class GO2CfgMoECTS(LeggedRobotCfgMoECTS):
|
||||||
class policy(LeggedRobotCfgREMCTS.policy):
|
class policy(LeggedRobotCfgMoECTS.policy):
|
||||||
expert_num = 8 # number of experts in the student model
|
expert_num = 8 # number of experts in the student model
|
||||||
|
|
||||||
class runner(LeggedRobotCfgREMCTS.runner):
|
class runner(LeggedRobotCfgMoECTS.runner):
|
||||||
run_name = ''
|
run_name = ''
|
||||||
experiment_name = 'go2_rem_cts'
|
experiment_name = 'go2_moe_cts'
|
||||||
max_iterations = 150000
|
max_iterations = 150000
|
||||||
save_interval = 500
|
save_interval = 500
|
||||||
@@ -81,9 +81,9 @@ class _TorchPolicyExporter(torch.nn.Module):
|
|||||||
self.obs_no_goal_mask = copy.deepcopy(policy.obs_no_goal_mask).cpu()
|
self.obs_no_goal_mask = copy.deepcopy(policy.obs_no_goal_mask).cpu()
|
||||||
self.history_length = policy.history.shape[1]
|
self.history_length = policy.history.shape[1]
|
||||||
self.history = torch.zeros([1, policy.history.shape[1], policy.history.shape[2]], device='cpu')
|
self.history = torch.zeros([1, policy.history.shape[1], policy.history.shape[2]], device='cpu')
|
||||||
self.forward = self.forward_moe_cts
|
self.forward = self.forward_moe_no_goal_cts
|
||||||
if not hasattr(policy, "obs_no_goal_mask"):
|
if not hasattr(policy, "obs_no_goal_mask"):
|
||||||
self.forward = self.forward_rem_cts
|
self.forward = self.forward_moe_cts
|
||||||
if hasattr(policy, "actor_mcp"):
|
if hasattr(policy, "actor_mcp"):
|
||||||
self.actor = copy.deepcopy(policy.actor_mcp)
|
self.actor = copy.deepcopy(policy.actor_mcp)
|
||||||
self.obs_no_goal_mask = copy.deepcopy(policy.obs_no_goal_mask).cpu()
|
self.obs_no_goal_mask = copy.deepcopy(policy.obs_no_goal_mask).cpu()
|
||||||
@@ -134,7 +134,7 @@ class _TorchPolicyExporter(torch.nn.Module):
|
|||||||
x = torch.cat([latent, x], dim=1)
|
x = torch.cat([latent, x], dim=1)
|
||||||
return self.actor(x), (None, latent)
|
return self.actor(x), (None, latent)
|
||||||
|
|
||||||
def forward_moe_cts(self, x): # x is single observations
|
def forward_moe_no_goal_cts(self, x): # x is single observations
|
||||||
x = self.normalizer(x)
|
x = self.normalizer(x)
|
||||||
self.history = torch.cat([self.history[:, 1:], x.unsqueeze(1)], dim=1)
|
self.history = torch.cat([self.history[:, 1:], x.unsqueeze(1)], dim=1)
|
||||||
history_no_goal = self.history.reshape(1, self.history_length, -1)[:, :, self.obs_no_goal_mask].reshape(1, -1)
|
history_no_goal = self.history.reshape(1, self.history_length, -1)[:, :, self.obs_no_goal_mask].reshape(1, -1)
|
||||||
@@ -142,7 +142,7 @@ class _TorchPolicyExporter(torch.nn.Module):
|
|||||||
x = torch.cat([latent, x], dim=1)
|
x = torch.cat([latent, x], dim=1)
|
||||||
return self.actor(x), (weights, latent)
|
return self.actor(x), (weights, latent)
|
||||||
|
|
||||||
def forward_rem_cts(self, x): # x is single observations
|
def forward_moe_cts(self, x): # x is single observations
|
||||||
x = self.normalizer(x)
|
x = self.normalizer(x)
|
||||||
self.history = torch.cat([self.history[:, 1:], x.unsqueeze(1)], dim=1)
|
self.history = torch.cat([self.history[:, 1:], x.unsqueeze(1)], dim=1)
|
||||||
latent, weights = self.student_moe_encoder(self.history.flatten(1))
|
latent, weights = self.student_moe_encoder(self.history.flatten(1))
|
||||||
@@ -211,12 +211,12 @@ class _OnnxPolicyExporter(torch.nn.Module):
|
|||||||
elif hasattr(policy, "student_moe_encoder"):
|
elif hasattr(policy, "student_moe_encoder"):
|
||||||
self.student_moe_encoder = copy.deepcopy(policy.student_moe_encoder)
|
self.student_moe_encoder = copy.deepcopy(policy.student_moe_encoder)
|
||||||
self.history_length = policy.history.shape[1]
|
self.history_length = policy.history.shape[1]
|
||||||
self.forward = self.forward_moe_cts
|
self.forward = self.forward_moe_no_goal_cts
|
||||||
self.input_dim = self.history_length * policy.history.shape[2]
|
self.input_dim = self.history_length * policy.history.shape[2]
|
||||||
if hasattr(policy, "obs_no_goal_mask"):
|
if hasattr(policy, "obs_no_goal_mask"):
|
||||||
self.obs_no_goal_mask = copy.deepcopy(policy.obs_no_goal_mask).cpu()
|
self.obs_no_goal_mask = copy.deepcopy(policy.obs_no_goal_mask).cpu()
|
||||||
else:
|
else:
|
||||||
self.forward = self.forward_rem_cts
|
self.forward = self.forward_moe_cts
|
||||||
|
|
||||||
else: # PPO
|
else: # PPO
|
||||||
self.forward = self.forward_ppo
|
self.forward = self.forward_ppo
|
||||||
@@ -277,7 +277,7 @@ class _OnnxPolicyExporter(torch.nn.Module):
|
|||||||
|
|
||||||
return self.actor(x)
|
return self.actor(x)
|
||||||
|
|
||||||
def forward_moe_cts(self, x):
|
def forward_moe_no_goal_cts(self, x):
|
||||||
x = self.normalizer(x)
|
x = self.normalizer(x)
|
||||||
history, obs_dim = self.flatten_obs(x)
|
history, obs_dim = self.flatten_obs(x)
|
||||||
|
|
||||||
@@ -290,7 +290,7 @@ class _OnnxPolicyExporter(torch.nn.Module):
|
|||||||
|
|
||||||
return self.actor(x), weights, latent
|
return self.actor(x), weights, latent
|
||||||
|
|
||||||
def forward_rem_cts(self, x):
|
def forward_moe_cts(self, x):
|
||||||
x = self.normalizer(x)
|
x = self.normalizer(x)
|
||||||
history, obs_dim = self.flatten_obs(x)
|
history, obs_dim = self.flatten_obs(x)
|
||||||
|
|
||||||
@@ -319,7 +319,7 @@ class _OnnxPolicyExporter(torch.nn.Module):
|
|||||||
obs = torch.zeros(1, self.input_dim)
|
obs = torch.zeros(1, self.input_dim)
|
||||||
|
|
||||||
output_names = ["actions"]
|
output_names = ["actions"]
|
||||||
if self.forward == self.forward_moe_cts:
|
if self.forward == self.forward_moe_no_goal_cts:
|
||||||
output_names.append("weights")
|
output_names.append("weights")
|
||||||
output_names.append("latent")
|
output_names.append("latent")
|
||||||
if self.forward == self.forward_mcp_cts:
|
if self.forward == self.forward_mcp_cts:
|
||||||
|
|||||||
@@ -118,6 +118,10 @@ def update_cfg_from_args(env_cfg, cfg_train, args):
|
|||||||
cfg_train.runner.load_run = args.load_run
|
cfg_train.runner.load_run = args.load_run
|
||||||
if args.checkpoint is not None:
|
if args.checkpoint is not None:
|
||||||
cfg_train.runner.checkpoint = args.checkpoint
|
cfg_train.runner.checkpoint = args.checkpoint
|
||||||
|
if args.robogauge is not None:
|
||||||
|
cfg_train.robogauge.enabled = args.robogauge
|
||||||
|
if args.robogauge_port is not None:
|
||||||
|
cfg_train.robogauge.port = args.robogauge_port
|
||||||
|
|
||||||
return env_cfg, cfg_train
|
return env_cfg, cfg_train
|
||||||
|
|
||||||
@@ -136,6 +140,9 @@ def get_args():
|
|||||||
{"name": "--num_envs", "type": int, "help": "Number of environments to create. Overrides config file if provided."},
|
{"name": "--num_envs", "type": int, "help": "Number of environments to create. Overrides config file if provided."},
|
||||||
{"name": "--seed", "type": int, "help": "Random seed. Overrides config file if provided."},
|
{"name": "--seed", "type": int, "help": "Random seed. Overrides config file if provided."},
|
||||||
{"name": "--max_iterations", "type": int, "help": "Maximum number of training iterations. Overrides config file if provided."},
|
{"name": "--max_iterations", "type": int, "help": "Maximum number of training iterations. Overrides config file if provided."},
|
||||||
|
|
||||||
|
{"name": "--robogauge", "action": "store_true", "default": False, "help": "Enable robogauge evaluation interface."},
|
||||||
|
{"name": "--robogauge_port", "type": int, "default": 9973, "help": "Port for robogauge evaluation interface."},
|
||||||
]
|
]
|
||||||
# parse arguments
|
# parse arguments
|
||||||
args = gymutil.parse_arguments(
|
args = gymutil.parse_arguments(
|
||||||
|
|||||||
@@ -30,8 +30,8 @@
|
|||||||
|
|
||||||
from .ppo import PPO
|
from .ppo import PPO
|
||||||
from .cts import CTS
|
from .cts import CTS
|
||||||
from .moe_cts import MoECTS
|
from .moe_ng_cts import MoENGCTS
|
||||||
from .mcp_cts import MCPCTS
|
from .mcp_cts import MCPCTS
|
||||||
from .ac_moe_cts import ACMoECTS
|
from .ac_moe_cts import ACMoECTS
|
||||||
from .dual_moe_cts import DualMoECTS
|
from .dual_moe_cts import DualMoECTS
|
||||||
from .rem_cts import REMCTS
|
from .moe_cts import MoECTS
|
||||||
@@ -202,7 +202,7 @@ class MoECTS(CTS):
|
|||||||
hid_states_batch, masks_batch
|
hid_states_batch, masks_batch
|
||||||
) = sample
|
) = sample
|
||||||
# Student encoder update
|
# Student encoder update
|
||||||
student_latent, gating_weights = self.model.get_student_latent_and_weights(history_batch[teacher_samples:])
|
student_latent, gating_weights = self.model.student_moe_encoder(history_batch[teacher_samples:])
|
||||||
with torch.no_grad():
|
with torch.no_grad():
|
||||||
teacher_latent = self.model.teacher_encoder(privileged_obs_batch[teacher_samples:])
|
teacher_latent = self.model.teacher_encoder(privileged_obs_batch[teacher_samples:])
|
||||||
latent_loss = (teacher_latent - student_latent).pow(2).mean()
|
latent_loss = (teacher_latent - student_latent).pow(2).mean()
|
||||||
|
|||||||
@@ -33,12 +33,12 @@ import torch.nn as nn
|
|||||||
import torch.optim as optim
|
import torch.optim as optim
|
||||||
|
|
||||||
import itertools
|
import itertools
|
||||||
from rsl_rl.modules import ActorCriticMoECTS
|
from rsl_rl.modules import ActorCriticMoENGCTS
|
||||||
from rsl_rl.storage import RolloutStorageCTS
|
from rsl_rl.storage import RolloutStorageCTS
|
||||||
from rsl_rl.algorithms.cts import CTS
|
from rsl_rl.algorithms.cts import CTS
|
||||||
|
|
||||||
class REMCTS(CTS):
|
class MoENGCTS(CTS):
|
||||||
model: ActorCriticMoECTS
|
model: ActorCriticMoENGCTS
|
||||||
def __init__(self,
|
def __init__(self,
|
||||||
model,
|
model,
|
||||||
num_envs,
|
num_envs,
|
||||||
@@ -202,7 +202,7 @@ class REMCTS(CTS):
|
|||||||
hid_states_batch, masks_batch
|
hid_states_batch, masks_batch
|
||||||
) = sample
|
) = sample
|
||||||
# Student encoder update
|
# Student encoder update
|
||||||
student_latent, gating_weights = self.model.student_moe_encoder(history_batch[teacher_samples:])
|
student_latent, gating_weights = self.model.get_student_latent_and_weights(history_batch[teacher_samples:])
|
||||||
with torch.no_grad():
|
with torch.no_grad():
|
||||||
teacher_latent = self.model.teacher_encoder(privileged_obs_batch[teacher_samples:])
|
teacher_latent = self.model.teacher_encoder(privileged_obs_batch[teacher_samples:])
|
||||||
latent_loss = (teacher_latent - student_latent).pow(2).mean()
|
latent_loss = (teacher_latent - student_latent).pow(2).mean()
|
||||||
@@ -31,8 +31,8 @@
|
|||||||
from .actor_critic import ActorCritic
|
from .actor_critic import ActorCritic
|
||||||
from .actor_critic_recurrent import ActorCriticRecurrent
|
from .actor_critic_recurrent import ActorCriticRecurrent
|
||||||
from .actor_critic_cts import ActorCriticCTS
|
from .actor_critic_cts import ActorCriticCTS
|
||||||
from .actor_critic_moe_cts import ActorCriticMoECTS
|
from .actor_critic_moe_ng_cts import ActorCriticMoENGCTS
|
||||||
from .actor_critic_mcp_cts import ActorCriticMCPCTS
|
from .actor_critic_mcp_cts import ActorCriticMCPCTS
|
||||||
from .actor_critic_ac_moe_cts import ActorCriticACMoECTS
|
from .actor_critic_ac_moe_cts import ActorCriticACMoECTS
|
||||||
from .actor_critic_dual_moe_cts import ActorCriticDualMoECTS
|
from .actor_critic_dual_moe_cts import ActorCriticDualMoECTS
|
||||||
from .actor_critic_rem_cts import ActorCriticREMCTS
|
from .actor_critic_moe_cts import ActorCriticMoECTS
|
||||||
@@ -27,7 +27,7 @@ class ActorCriticDualMoECTS(nn.Module):
|
|||||||
actor_hidden_dims=[512, 256, 128],
|
actor_hidden_dims=[512, 256, 128],
|
||||||
critic_hidden_dims=[512, 256, 128],
|
critic_hidden_dims=[512, 256, 128],
|
||||||
teacher_encoder_hidden_dims=[512, 256],
|
teacher_encoder_hidden_dims=[512, 256],
|
||||||
student_encoder_hidden_dims=[512, 256, 128], # last dim is expert hidden dim
|
student_encoder_hidden_dims=[512, 256, 256], # last dim is expert hidden dim
|
||||||
expert_num=8,
|
expert_num=8,
|
||||||
activation='elu',
|
activation='elu',
|
||||||
init_noise_std=1.0,
|
init_noise_std=1.0,
|
||||||
|
|||||||
@@ -15,6 +15,8 @@ import torch.nn as nn
|
|||||||
import torch.nn.functional as F
|
import torch.nn.functional as F
|
||||||
from torch.distributions import Normal
|
from torch.distributions import Normal
|
||||||
|
|
||||||
|
from rsl_rl.modules.utils import L2Norm, SimNorm, StudentMoEEncoder, MLP
|
||||||
|
|
||||||
class ActorCriticMoECTS(nn.Module):
|
class ActorCriticMoECTS(nn.Module):
|
||||||
is_recurrent = False
|
is_recurrent = False
|
||||||
def __init__(self, num_obs,
|
def __init__(self, num_obs,
|
||||||
@@ -22,12 +24,11 @@ class ActorCriticMoECTS(nn.Module):
|
|||||||
num_actions,
|
num_actions,
|
||||||
num_envs,
|
num_envs,
|
||||||
history_length,
|
history_length,
|
||||||
obs_no_goal_mask,
|
|
||||||
actor_hidden_dims=[512, 256, 128],
|
actor_hidden_dims=[512, 256, 128],
|
||||||
critic_hidden_dims=[512, 256, 128],
|
critic_hidden_dims=[512, 256, 128],
|
||||||
teacher_encoder_hidden_dims=[512, 256],
|
teacher_encoder_hidden_dims=[512, 256],
|
||||||
student_encoder_hidden_dims=[512, 256],
|
student_encoder_hidden_dims=[512, 256, 256],
|
||||||
student_expert_num=8,
|
expert_num=8,
|
||||||
activation='elu',
|
activation='elu',
|
||||||
init_noise_std=1.0,
|
init_noise_std=1.0,
|
||||||
latent_dim=32,
|
latent_dim=32,
|
||||||
@@ -36,17 +37,12 @@ class ActorCriticMoECTS(nn.Module):
|
|||||||
if kwargs:
|
if kwargs:
|
||||||
print("ActorCritic.__init__ got unexpected arguments, which will be ignored: " + str([key for key in kwargs.keys()]))
|
print("ActorCritic.__init__ got unexpected arguments, which will be ignored: " + str([key for key in kwargs.keys()]))
|
||||||
assert norm_type in ['l2norm', 'simnorm'], f"Normalization type {norm_type} not supported!"
|
assert norm_type in ['l2norm', 'simnorm'], f"Normalization type {norm_type} not supported!"
|
||||||
super(ActorCriticMoECTS, self).__init__()
|
super().__init__()
|
||||||
self.num_actions = num_actions
|
self.num_actions = num_actions
|
||||||
self.history_length = history_length
|
self.history_length = history_length
|
||||||
self.register_buffer("obs_no_goal_mask", torch.tensor(obs_no_goal_mask, dtype=torch.bool), persistent=False)
|
|
||||||
|
|
||||||
activation_str = activation
|
|
||||||
activation = get_activation(activation)
|
|
||||||
|
|
||||||
mlp_input_dim_t = num_critic_obs
|
mlp_input_dim_t = num_critic_obs
|
||||||
mlp_input_dim_e = torch.sum(self.obs_no_goal_mask).item() * history_length # exclude command inputs for expert
|
mlp_input_dim_s = num_obs * history_length
|
||||||
mlp_input_dim_g = num_obs * history_length # all obs for gating
|
|
||||||
mlp_input_dim_a = latent_dim + num_obs
|
mlp_input_dim_a = latent_dim + num_obs
|
||||||
mlp_input_dim_c = latent_dim + num_critic_obs
|
mlp_input_dim_c = latent_dim + num_critic_obs
|
||||||
|
|
||||||
@@ -54,54 +50,26 @@ class ActorCriticMoECTS(nn.Module):
|
|||||||
self.register_buffer("history", torch.zeros((num_envs, history_length, num_obs)), persistent=False)
|
self.register_buffer("history", torch.zeros((num_envs, history_length, num_obs)), persistent=False)
|
||||||
|
|
||||||
# Teacher encoder
|
# Teacher encoder
|
||||||
encoder_layers = []
|
self.teacher_encoder = nn.Sequential(
|
||||||
encoder_layers.append(nn.Linear(mlp_input_dim_t, teacher_encoder_hidden_dims[0]))
|
MLP([mlp_input_dim_t, *teacher_encoder_hidden_dims, latent_dim], activation=activation),
|
||||||
encoder_layers.append(activation)
|
L2Norm() if norm_type == 'l2norm' else SimNorm()
|
||||||
for l in range(len(teacher_encoder_hidden_dims)):
|
)
|
||||||
if l == len(teacher_encoder_hidden_dims) - 1:
|
|
||||||
encoder_layers.append(nn.Linear(teacher_encoder_hidden_dims[l], latent_dim))
|
|
||||||
if norm_type == 'l2norm':
|
|
||||||
encoder_layers.append(L2Norm())
|
|
||||||
elif norm_type == 'simnorm':
|
|
||||||
encoder_layers.append(SimNorm())
|
|
||||||
else:
|
|
||||||
encoder_layers.append(nn.Linear(teacher_encoder_hidden_dims[l], teacher_encoder_hidden_dims[l + 1]))
|
|
||||||
encoder_layers.append(activation)
|
|
||||||
self.teacher_encoder = nn.Sequential(*encoder_layers)
|
|
||||||
|
|
||||||
# Student MoE encoder
|
# Student MoE encoder
|
||||||
self.student_moe_encoder = StudentMoEEncoder(
|
self.student_moe_encoder = StudentMoEEncoder(
|
||||||
expert_dim=mlp_input_dim_e,
|
expert_num=expert_num,
|
||||||
gating_dim=mlp_input_dim_g,
|
input_dim=mlp_input_dim_s,
|
||||||
hidden_dims=student_encoder_hidden_dims,
|
hidden_dims=student_encoder_hidden_dims,
|
||||||
expert_num=student_expert_num,
|
output_dim=latent_dim,
|
||||||
latent_dim=latent_dim,
|
activation=activation,
|
||||||
activation=activation_str
|
norm_type=norm_type,
|
||||||
)
|
)
|
||||||
|
|
||||||
# Policy
|
# Policy
|
||||||
actor_layers = []
|
self.actor = MLP([mlp_input_dim_a, *actor_hidden_dims, num_actions], activation=activation)
|
||||||
actor_layers.append(nn.Linear(mlp_input_dim_a, actor_hidden_dims[0]))
|
|
||||||
actor_layers.append(activation)
|
|
||||||
for l in range(len(actor_hidden_dims)):
|
|
||||||
if l == len(actor_hidden_dims) - 1:
|
|
||||||
actor_layers.append(nn.Linear(actor_hidden_dims[l], num_actions))
|
|
||||||
else:
|
|
||||||
actor_layers.append(nn.Linear(actor_hidden_dims[l], actor_hidden_dims[l + 1]))
|
|
||||||
actor_layers.append(activation)
|
|
||||||
self.actor = nn.Sequential(*actor_layers)
|
|
||||||
|
|
||||||
# Value function
|
# Value function
|
||||||
critic_layers = []
|
self.critic = MLP([mlp_input_dim_c, *critic_hidden_dims, 1], activation=activation)
|
||||||
critic_layers.append(nn.Linear(mlp_input_dim_c, critic_hidden_dims[0]))
|
|
||||||
critic_layers.append(activation)
|
|
||||||
for l in range(len(critic_hidden_dims)):
|
|
||||||
if l == len(critic_hidden_dims) - 1:
|
|
||||||
critic_layers.append(nn.Linear(critic_hidden_dims[l], 1))
|
|
||||||
else:
|
|
||||||
critic_layers.append(nn.Linear(critic_hidden_dims[l], critic_hidden_dims[l + 1]))
|
|
||||||
critic_layers.append(activation)
|
|
||||||
self.critic = nn.Sequential(*critic_layers)
|
|
||||||
|
|
||||||
print(f"Actor MLP: {self.actor}")
|
print(f"Actor MLP: {self.actor}")
|
||||||
print(f"Critic MLP: {self.critic}")
|
print(f"Critic MLP: {self.critic}")
|
||||||
@@ -114,10 +82,6 @@ class ActorCriticMoECTS(nn.Module):
|
|||||||
# disable args validation for speedup
|
# disable args validation for speedup
|
||||||
Normal.set_default_validate_args = False
|
Normal.set_default_validate_args = False
|
||||||
|
|
||||||
# seems that we get better performance without init
|
|
||||||
# self.init_memory_weights(self.memory_a, 0.001, 0.)
|
|
||||||
# self.init_memory_weights(self.memory_c, 0.001, 0.)
|
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
# not used at the moment
|
# not used at the moment
|
||||||
def init_weights(sequential, scales):
|
def init_weights(sequential, scales):
|
||||||
@@ -152,7 +116,7 @@ class ActorCriticMoECTS(nn.Module):
|
|||||||
latent = self.teacher_encoder(privileged_obs)
|
latent = self.teacher_encoder(privileged_obs)
|
||||||
else:
|
else:
|
||||||
with torch.no_grad():
|
with torch.no_grad():
|
||||||
latent, _ = self.get_student_latent_and_weights(history)
|
latent, _ = self.student_moe_encoder(history)
|
||||||
x = torch.cat([latent, obs], dim=1)
|
x = torch.cat([latent, obs], dim=1)
|
||||||
self.update_distribution(x)
|
self.update_distribution(x)
|
||||||
return self.distribution.sample()
|
return self.distribution.sample()
|
||||||
@@ -162,7 +126,7 @@ class ActorCriticMoECTS(nn.Module):
|
|||||||
|
|
||||||
def act_inference(self, obs):
|
def act_inference(self, obs):
|
||||||
self.history = torch.cat([self.history[:, 1:], obs.unsqueeze(1)], dim=1)
|
self.history = torch.cat([self.history[:, 1:], obs.unsqueeze(1)], dim=1)
|
||||||
latent, _ = self.get_student_latent_and_weights(self.history.flatten(1))
|
latent, _ = self.student_moe_encoder(self.history.flatten(1))
|
||||||
x = torch.cat([latent, obs], dim=1)
|
x = torch.cat([latent, obs], dim=1)
|
||||||
actions_mean = self.actor(x)
|
actions_mean = self.actor(x)
|
||||||
return actions_mean
|
return actions_mean
|
||||||
@@ -171,117 +135,7 @@ class ActorCriticMoECTS(nn.Module):
|
|||||||
if is_teacher:
|
if is_teacher:
|
||||||
latent = self.teacher_encoder(privileged_obs)
|
latent = self.teacher_encoder(privileged_obs)
|
||||||
else:
|
else:
|
||||||
latent, _ = self.get_student_latent_and_weights(history)
|
latent, _ = self.student_moe_encoder(history)
|
||||||
x = torch.cat([latent.detach(), privileged_obs], dim=1)
|
x = torch.cat([latent.detach(), privileged_obs], dim=1)
|
||||||
value = self.critic(x)
|
value = self.critic(x)
|
||||||
return value
|
return value
|
||||||
|
|
||||||
def get_student_latent_and_weights(self, history):
|
|
||||||
B = history.shape[0]
|
|
||||||
history_no_goal = history.reshape(B, self.history_length, -1)[:, :, self.obs_no_goal_mask].reshape(B, -1)
|
|
||||||
return self.student_moe_encoder(history, history_no_goal)
|
|
||||||
|
|
||||||
class StudentMoEEncoder(nn.Module):
|
|
||||||
def __init__(
|
|
||||||
self,
|
|
||||||
expert_dim,
|
|
||||||
gating_dim,
|
|
||||||
hidden_dims=[512, 256],
|
|
||||||
expert_num=8,
|
|
||||||
expert_hidden_dim=256,
|
|
||||||
latent_dim=32,
|
|
||||||
activation='elu',
|
|
||||||
norm_type='l2norm',
|
|
||||||
):
|
|
||||||
super().__init__()
|
|
||||||
self.expert_num = expert_num
|
|
||||||
self.latent_dim = latent_dim
|
|
||||||
self.norm_layer = L2Norm() if norm_type == 'l2norm' else SimNorm()
|
|
||||||
activation = get_activation(activation)
|
|
||||||
|
|
||||||
# Expert networks
|
|
||||||
experts_layers = []
|
|
||||||
last_dim = expert_dim
|
|
||||||
for l in hidden_dims:
|
|
||||||
experts_layers.append(nn.Linear(last_dim, l))
|
|
||||||
experts_layers.append(activation)
|
|
||||||
last_dim = l
|
|
||||||
self.experts_backbone = nn.Sequential(*experts_layers)
|
|
||||||
self.experts_hidden = nn.Sequential(
|
|
||||||
nn.Linear(last_dim, expert_num * expert_hidden_dim),
|
|
||||||
activation
|
|
||||||
)
|
|
||||||
self.experts_out = nn.Conv1d(
|
|
||||||
in_channels=expert_num*expert_hidden_dim,
|
|
||||||
out_channels=expert_num*latent_dim,
|
|
||||||
kernel_size=1,
|
|
||||||
groups=expert_num
|
|
||||||
)
|
|
||||||
|
|
||||||
# Gating network
|
|
||||||
gating_layers = []
|
|
||||||
last_dim = gating_dim
|
|
||||||
for l in hidden_dims:
|
|
||||||
gating_layers.append(nn.Linear(last_dim, l))
|
|
||||||
gating_layers.append(activation)
|
|
||||||
last_dim = l
|
|
||||||
gating_layers.append(nn.Linear(last_dim, expert_num))
|
|
||||||
gating_layers.append(nn.Softmax(dim=-1))
|
|
||||||
self.gating_network = nn.Sequential(*gating_layers)
|
|
||||||
|
|
||||||
def forward(self, obs, obs_no_goal):
|
|
||||||
weights = self.gating_network(obs) # (batch, expert_num)
|
|
||||||
shared_features = self.experts_backbone(obs_no_goal)
|
|
||||||
expert_hidden = self.experts_hidden(shared_features)
|
|
||||||
expert_hidden = expert_hidden.unsqueeze(-1)
|
|
||||||
expert_latent_flat = self.experts_out(expert_hidden) # (batch, expert_num * latent_dim, 1)
|
|
||||||
expert_latent = expert_latent_flat.reshape(-1, self.expert_num, self.latent_dim)
|
|
||||||
latent = torch.sum(weights.unsqueeze(-1) * expert_latent, dim=1) # (batch, latent_dim)
|
|
||||||
latent = self.norm_layer(latent)
|
|
||||||
return latent, weights
|
|
||||||
|
|
||||||
def get_activation(act_name):
|
|
||||||
if act_name == "elu":
|
|
||||||
return nn.ELU()
|
|
||||||
elif act_name == "selu":
|
|
||||||
return nn.SELU()
|
|
||||||
elif act_name == "relu":
|
|
||||||
return nn.ReLU()
|
|
||||||
elif act_name == "crelu":
|
|
||||||
return nn.ReLU()
|
|
||||||
elif act_name == "lrelu":
|
|
||||||
return nn.LeakyReLU()
|
|
||||||
elif act_name == "tanh":
|
|
||||||
return nn.Tanh()
|
|
||||||
elif act_name == "sigmoid":
|
|
||||||
return nn.Sigmoid()
|
|
||||||
else:
|
|
||||||
print("invalid activation function!")
|
|
||||||
return None
|
|
||||||
|
|
||||||
class L2Norm(nn.Module):
|
|
||||||
|
|
||||||
def __init__(self):
|
|
||||||
super().__init__()
|
|
||||||
|
|
||||||
def forward(self, x):
|
|
||||||
return F.normalize(x, p=2.0, dim=-1)
|
|
||||||
|
|
||||||
class SimNorm(nn.Module):
|
|
||||||
"""
|
|
||||||
Simplicial normalization.
|
|
||||||
Adapted from https://arxiv.org/abs/2204.00616.
|
|
||||||
"""
|
|
||||||
|
|
||||||
def __init__(self):
|
|
||||||
super().__init__()
|
|
||||||
self.dim = 8 # for latent dim 512
|
|
||||||
|
|
||||||
def forward(self, x):
|
|
||||||
shp = x.shape
|
|
||||||
x = x.view(*shp[:-1], -1, self.dim)
|
|
||||||
x = F.softmax(x, dim=-1)
|
|
||||||
return x.view(*shp)
|
|
||||||
|
|
||||||
def __repr__(self):
|
|
||||||
return f"SimNorm(dim={self.dim})"
|
|
||||||
|
|||||||
288
rsl_rl/rsl_rl/modules/actor_critic_moe_ng_cts.py
Normal file
288
rsl_rl/rsl_rl/modules/actor_critic_moe_ng_cts.py
Normal file
@@ -0,0 +1,288 @@
|
|||||||
|
# -*- coding: utf-8 -*-
|
||||||
|
'''
|
||||||
|
@File : actor_critic_moe_ng_cts.py
|
||||||
|
@Time : 2025/12/30 21:06:46
|
||||||
|
@Author : wty-yy
|
||||||
|
@Version : 1.0
|
||||||
|
@Blog : https://wty-yy.github.io/
|
||||||
|
@Desc : Mixture of Experts (experts without goal) Concurrent Teacher Student Network
|
||||||
|
@Refer : CTS https://arxiv.org/abs/2405.10830, Switch Transformers https://arxiv.org/abs/2101.03961
|
||||||
|
'''
|
||||||
|
import numpy as np
|
||||||
|
|
||||||
|
import torch
|
||||||
|
import torch.nn as nn
|
||||||
|
import torch.nn.functional as F
|
||||||
|
from torch.distributions import Normal
|
||||||
|
|
||||||
|
class ActorCriticMoENGCTS(nn.Module):
|
||||||
|
is_recurrent = False
|
||||||
|
def __init__(self, num_obs,
|
||||||
|
num_critic_obs,
|
||||||
|
num_actions,
|
||||||
|
num_envs,
|
||||||
|
history_length,
|
||||||
|
obs_no_goal_mask,
|
||||||
|
actor_hidden_dims=[512, 256, 128],
|
||||||
|
critic_hidden_dims=[512, 256, 128],
|
||||||
|
teacher_encoder_hidden_dims=[512, 256],
|
||||||
|
student_encoder_hidden_dims=[512, 256],
|
||||||
|
student_expert_num=8,
|
||||||
|
activation='elu',
|
||||||
|
init_noise_std=1.0,
|
||||||
|
latent_dim=32,
|
||||||
|
norm_type='l2norm',
|
||||||
|
**kwargs):
|
||||||
|
if kwargs:
|
||||||
|
print("ActorCritic.__init__ got unexpected arguments, which will be ignored: " + str([key for key in kwargs.keys()]))
|
||||||
|
assert norm_type in ['l2norm', 'simnorm'], f"Normalization type {norm_type} not supported!"
|
||||||
|
super(ActorCriticMoENGCTS, self).__init__()
|
||||||
|
self.num_actions = num_actions
|
||||||
|
self.history_length = history_length
|
||||||
|
self.register_buffer("obs_no_goal_mask", torch.tensor(obs_no_goal_mask, dtype=torch.bool), persistent=False)
|
||||||
|
|
||||||
|
activation_str = activation
|
||||||
|
activation = get_activation(activation)
|
||||||
|
|
||||||
|
mlp_input_dim_t = num_critic_obs
|
||||||
|
mlp_input_dim_e = torch.sum(self.obs_no_goal_mask).item() * history_length # exclude command inputs for expert
|
||||||
|
mlp_input_dim_g = num_obs * history_length # all obs for gating
|
||||||
|
mlp_input_dim_a = latent_dim + num_obs
|
||||||
|
mlp_input_dim_c = latent_dim + num_critic_obs
|
||||||
|
|
||||||
|
# History
|
||||||
|
self.register_buffer("history", torch.zeros((num_envs, history_length, num_obs)), persistent=False)
|
||||||
|
|
||||||
|
# Teacher encoder
|
||||||
|
encoder_layers = []
|
||||||
|
encoder_layers.append(nn.Linear(mlp_input_dim_t, teacher_encoder_hidden_dims[0]))
|
||||||
|
encoder_layers.append(activation)
|
||||||
|
for l in range(len(teacher_encoder_hidden_dims)):
|
||||||
|
if l == len(teacher_encoder_hidden_dims) - 1:
|
||||||
|
encoder_layers.append(nn.Linear(teacher_encoder_hidden_dims[l], latent_dim))
|
||||||
|
if norm_type == 'l2norm':
|
||||||
|
encoder_layers.append(L2Norm())
|
||||||
|
elif norm_type == 'simnorm':
|
||||||
|
encoder_layers.append(SimNorm())
|
||||||
|
else:
|
||||||
|
encoder_layers.append(nn.Linear(teacher_encoder_hidden_dims[l], teacher_encoder_hidden_dims[l + 1]))
|
||||||
|
encoder_layers.append(activation)
|
||||||
|
self.teacher_encoder = nn.Sequential(*encoder_layers)
|
||||||
|
|
||||||
|
# Student MoE no goal encoder
|
||||||
|
self.student_moe_encoder = StudentMoEEncoder(
|
||||||
|
expert_dim=mlp_input_dim_e,
|
||||||
|
gating_dim=mlp_input_dim_g,
|
||||||
|
hidden_dims=student_encoder_hidden_dims,
|
||||||
|
expert_num=student_expert_num,
|
||||||
|
latent_dim=latent_dim,
|
||||||
|
activation=activation_str
|
||||||
|
)
|
||||||
|
|
||||||
|
# Policy
|
||||||
|
actor_layers = []
|
||||||
|
actor_layers.append(nn.Linear(mlp_input_dim_a, actor_hidden_dims[0]))
|
||||||
|
actor_layers.append(activation)
|
||||||
|
for l in range(len(actor_hidden_dims)):
|
||||||
|
if l == len(actor_hidden_dims) - 1:
|
||||||
|
actor_layers.append(nn.Linear(actor_hidden_dims[l], num_actions))
|
||||||
|
else:
|
||||||
|
actor_layers.append(nn.Linear(actor_hidden_dims[l], actor_hidden_dims[l + 1]))
|
||||||
|
actor_layers.append(activation)
|
||||||
|
self.actor = nn.Sequential(*actor_layers)
|
||||||
|
|
||||||
|
# Value function
|
||||||
|
critic_layers = []
|
||||||
|
critic_layers.append(nn.Linear(mlp_input_dim_c, critic_hidden_dims[0]))
|
||||||
|
critic_layers.append(activation)
|
||||||
|
for l in range(len(critic_hidden_dims)):
|
||||||
|
if l == len(critic_hidden_dims) - 1:
|
||||||
|
critic_layers.append(nn.Linear(critic_hidden_dims[l], 1))
|
||||||
|
else:
|
||||||
|
critic_layers.append(nn.Linear(critic_hidden_dims[l], critic_hidden_dims[l + 1]))
|
||||||
|
critic_layers.append(activation)
|
||||||
|
self.critic = nn.Sequential(*critic_layers)
|
||||||
|
|
||||||
|
print(f"Actor MLP: {self.actor}")
|
||||||
|
print(f"Critic MLP: {self.critic}")
|
||||||
|
print(f"Teacher Encoder: {self.teacher_encoder}")
|
||||||
|
print(f"Student MoE no goal Encoder: {self.student_moe_encoder}")
|
||||||
|
|
||||||
|
|
||||||
|
# Action noise
|
||||||
|
self.std = nn.Parameter(init_noise_std * torch.ones(num_actions))
|
||||||
|
self.distribution = None
|
||||||
|
# disable args validation for speedup
|
||||||
|
Normal.set_default_validate_args = False
|
||||||
|
|
||||||
|
# seems that we get better performance without init
|
||||||
|
# self.init_memory_weights(self.memory_a, 0.001, 0.)
|
||||||
|
# self.init_memory_weights(self.memory_c, 0.001, 0.)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
# not used at the moment
|
||||||
|
def init_weights(sequential, scales):
|
||||||
|
[torch.nn.init.orthogonal_(module.weight, gain=scales[idx]) for idx, module in
|
||||||
|
enumerate(mod for mod in sequential if isinstance(mod, nn.Linear))]
|
||||||
|
|
||||||
|
|
||||||
|
def reset(self, dones=None):
|
||||||
|
self.history[dones > 0] = 0.0
|
||||||
|
|
||||||
|
def forward(self):
|
||||||
|
raise NotImplementedError
|
||||||
|
|
||||||
|
@property
|
||||||
|
def action_mean(self):
|
||||||
|
return self.distribution.mean
|
||||||
|
|
||||||
|
@property
|
||||||
|
def action_std(self):
|
||||||
|
return self.distribution.stddev
|
||||||
|
|
||||||
|
@property
|
||||||
|
def entropy(self):
|
||||||
|
return self.distribution.entropy().sum(dim=-1)
|
||||||
|
|
||||||
|
def update_distribution(self, latent_and_obs):
|
||||||
|
mean = self.actor(latent_and_obs)
|
||||||
|
self.distribution = Normal(mean, mean*0. + self.std)
|
||||||
|
|
||||||
|
def act(self, obs, privileged_obs, history, is_teacher, **kwargs):
|
||||||
|
if is_teacher:
|
||||||
|
latent = self.teacher_encoder(privileged_obs)
|
||||||
|
else:
|
||||||
|
with torch.no_grad():
|
||||||
|
latent, _ = self.get_student_latent_and_weights(history)
|
||||||
|
x = torch.cat([latent, obs], dim=1)
|
||||||
|
self.update_distribution(x)
|
||||||
|
return self.distribution.sample()
|
||||||
|
|
||||||
|
def get_actions_log_prob(self, actions):
|
||||||
|
return self.distribution.log_prob(actions).sum(dim=-1)
|
||||||
|
|
||||||
|
def act_inference(self, obs):
|
||||||
|
self.history = torch.cat([self.history[:, 1:], obs.unsqueeze(1)], dim=1)
|
||||||
|
latent, _ = self.get_student_latent_and_weights(self.history.flatten(1))
|
||||||
|
x = torch.cat([latent, obs], dim=1)
|
||||||
|
actions_mean = self.actor(x)
|
||||||
|
return actions_mean
|
||||||
|
|
||||||
|
def evaluate(self, privileged_obs, history, is_teacher, **kwargs):
|
||||||
|
if is_teacher:
|
||||||
|
latent = self.teacher_encoder(privileged_obs)
|
||||||
|
else:
|
||||||
|
latent, _ = self.get_student_latent_and_weights(history)
|
||||||
|
x = torch.cat([latent.detach(), privileged_obs], dim=1)
|
||||||
|
value = self.critic(x)
|
||||||
|
return value
|
||||||
|
|
||||||
|
def get_student_latent_and_weights(self, history):
|
||||||
|
B = history.shape[0]
|
||||||
|
history_no_goal = history.reshape(B, self.history_length, -1)[:, :, self.obs_no_goal_mask].reshape(B, -1)
|
||||||
|
return self.student_moe_encoder(history, history_no_goal)
|
||||||
|
|
||||||
|
class StudentMoEEncoder(nn.Module):
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
expert_dim,
|
||||||
|
gating_dim,
|
||||||
|
hidden_dims=[512, 256],
|
||||||
|
expert_num=8,
|
||||||
|
expert_hidden_dim=256,
|
||||||
|
latent_dim=32,
|
||||||
|
activation='elu',
|
||||||
|
norm_type='l2norm',
|
||||||
|
):
|
||||||
|
super().__init__()
|
||||||
|
self.expert_num = expert_num
|
||||||
|
self.latent_dim = latent_dim
|
||||||
|
self.norm_layer = L2Norm() if norm_type == 'l2norm' else SimNorm()
|
||||||
|
activation = get_activation(activation)
|
||||||
|
|
||||||
|
# Expert networks
|
||||||
|
experts_layers = []
|
||||||
|
last_dim = expert_dim
|
||||||
|
for l in hidden_dims:
|
||||||
|
experts_layers.append(nn.Linear(last_dim, l))
|
||||||
|
experts_layers.append(activation)
|
||||||
|
last_dim = l
|
||||||
|
self.experts_backbone = nn.Sequential(*experts_layers)
|
||||||
|
self.experts_hidden = nn.Sequential(
|
||||||
|
nn.Linear(last_dim, expert_num * expert_hidden_dim),
|
||||||
|
activation
|
||||||
|
)
|
||||||
|
self.experts_out = nn.Conv1d(
|
||||||
|
in_channels=expert_num*expert_hidden_dim,
|
||||||
|
out_channels=expert_num*latent_dim,
|
||||||
|
kernel_size=1,
|
||||||
|
groups=expert_num
|
||||||
|
)
|
||||||
|
|
||||||
|
# Gating network
|
||||||
|
gating_layers = []
|
||||||
|
last_dim = gating_dim
|
||||||
|
for l in hidden_dims:
|
||||||
|
gating_layers.append(nn.Linear(last_dim, l))
|
||||||
|
gating_layers.append(activation)
|
||||||
|
last_dim = l
|
||||||
|
gating_layers.append(nn.Linear(last_dim, expert_num))
|
||||||
|
gating_layers.append(nn.Softmax(dim=-1))
|
||||||
|
self.gating_network = nn.Sequential(*gating_layers)
|
||||||
|
|
||||||
|
def forward(self, obs, obs_no_goal):
|
||||||
|
weights = self.gating_network(obs) # (batch, expert_num)
|
||||||
|
shared_features = self.experts_backbone(obs_no_goal)
|
||||||
|
expert_hidden = self.experts_hidden(shared_features)
|
||||||
|
expert_hidden = expert_hidden.unsqueeze(-1)
|
||||||
|
expert_latent_flat = self.experts_out(expert_hidden) # (batch, expert_num * latent_dim, 1)
|
||||||
|
expert_latent = expert_latent_flat.reshape(-1, self.expert_num, self.latent_dim)
|
||||||
|
latent = torch.sum(weights.unsqueeze(-1) * expert_latent, dim=1) # (batch, latent_dim)
|
||||||
|
latent = self.norm_layer(latent)
|
||||||
|
return latent, weights
|
||||||
|
|
||||||
|
def get_activation(act_name):
|
||||||
|
if act_name == "elu":
|
||||||
|
return nn.ELU()
|
||||||
|
elif act_name == "selu":
|
||||||
|
return nn.SELU()
|
||||||
|
elif act_name == "relu":
|
||||||
|
return nn.ReLU()
|
||||||
|
elif act_name == "crelu":
|
||||||
|
return nn.ReLU()
|
||||||
|
elif act_name == "lrelu":
|
||||||
|
return nn.LeakyReLU()
|
||||||
|
elif act_name == "tanh":
|
||||||
|
return nn.Tanh()
|
||||||
|
elif act_name == "sigmoid":
|
||||||
|
return nn.Sigmoid()
|
||||||
|
else:
|
||||||
|
print("invalid activation function!")
|
||||||
|
return None
|
||||||
|
|
||||||
|
class L2Norm(nn.Module):
|
||||||
|
|
||||||
|
def __init__(self):
|
||||||
|
super().__init__()
|
||||||
|
|
||||||
|
def forward(self, x):
|
||||||
|
return F.normalize(x, p=2.0, dim=-1)
|
||||||
|
|
||||||
|
class SimNorm(nn.Module):
|
||||||
|
"""
|
||||||
|
Simplicial normalization.
|
||||||
|
Adapted from https://arxiv.org/abs/2204.00616.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self):
|
||||||
|
super().__init__()
|
||||||
|
self.dim = 8 # for latent dim 512
|
||||||
|
|
||||||
|
def forward(self, x):
|
||||||
|
shp = x.shape
|
||||||
|
x = x.view(*shp[:-1], -1, self.dim)
|
||||||
|
x = F.softmax(x, dim=-1)
|
||||||
|
return x.view(*shp)
|
||||||
|
|
||||||
|
def __repr__(self):
|
||||||
|
return f"SimNorm(dim={self.dim})"
|
||||||
@@ -1,141 +0,0 @@
|
|||||||
# -*- coding: utf-8 -*-
|
|
||||||
'''
|
|
||||||
@File : actor_critic_moe_cts.py
|
|
||||||
@Time : 2025/12/30 21:06:46
|
|
||||||
@Author : wty-yy
|
|
||||||
@Version : 1.0
|
|
||||||
@Blog : https://wty-yy.github.io/
|
|
||||||
@Desc : Mixture of Experts Concurrent Teacher Student Network
|
|
||||||
@Refer : CTS https://arxiv.org/abs/2405.10830, Switch Transformers https://arxiv.org/abs/2101.03961
|
|
||||||
'''
|
|
||||||
import numpy as np
|
|
||||||
|
|
||||||
import torch
|
|
||||||
import torch.nn as nn
|
|
||||||
import torch.nn.functional as F
|
|
||||||
from torch.distributions import Normal
|
|
||||||
|
|
||||||
from rsl_rl.modules.utils import L2Norm, SimNorm, StudentMoEEncoder, MLP
|
|
||||||
|
|
||||||
class ActorCriticREMCTS(nn.Module):
|
|
||||||
is_recurrent = False
|
|
||||||
def __init__(self, num_obs,
|
|
||||||
num_critic_obs,
|
|
||||||
num_actions,
|
|
||||||
num_envs,
|
|
||||||
history_length,
|
|
||||||
actor_hidden_dims=[512, 256, 128],
|
|
||||||
critic_hidden_dims=[512, 256, 128],
|
|
||||||
teacher_encoder_hidden_dims=[512, 256],
|
|
||||||
student_encoder_hidden_dims=[512, 256, 128],
|
|
||||||
expert_num=8,
|
|
||||||
activation='elu',
|
|
||||||
init_noise_std=1.0,
|
|
||||||
latent_dim=32,
|
|
||||||
norm_type='l2norm',
|
|
||||||
**kwargs):
|
|
||||||
if kwargs:
|
|
||||||
print("ActorCritic.__init__ got unexpected arguments, which will be ignored: " + str([key for key in kwargs.keys()]))
|
|
||||||
assert norm_type in ['l2norm', 'simnorm'], f"Normalization type {norm_type} not supported!"
|
|
||||||
super().__init__()
|
|
||||||
self.num_actions = num_actions
|
|
||||||
self.history_length = history_length
|
|
||||||
|
|
||||||
mlp_input_dim_t = num_critic_obs
|
|
||||||
mlp_input_dim_s = num_obs * history_length
|
|
||||||
mlp_input_dim_a = latent_dim + num_obs
|
|
||||||
mlp_input_dim_c = latent_dim + num_critic_obs
|
|
||||||
|
|
||||||
# History
|
|
||||||
self.register_buffer("history", torch.zeros((num_envs, history_length, num_obs)), persistent=False)
|
|
||||||
|
|
||||||
# Teacher encoder
|
|
||||||
self.teacher_encoder = nn.Sequential(
|
|
||||||
MLP([mlp_input_dim_t, *teacher_encoder_hidden_dims, latent_dim], activation=activation),
|
|
||||||
L2Norm() if norm_type == 'l2norm' else SimNorm()
|
|
||||||
)
|
|
||||||
|
|
||||||
# Student MoE encoder
|
|
||||||
self.student_moe_encoder = StudentMoEEncoder(
|
|
||||||
expert_num=expert_num,
|
|
||||||
input_dim=mlp_input_dim_s,
|
|
||||||
hidden_dims=student_encoder_hidden_dims,
|
|
||||||
output_dim=latent_dim,
|
|
||||||
activation=activation,
|
|
||||||
norm_type=norm_type,
|
|
||||||
)
|
|
||||||
|
|
||||||
# Policy
|
|
||||||
self.actor = MLP([mlp_input_dim_a, *actor_hidden_dims, num_actions], activation=activation)
|
|
||||||
|
|
||||||
# Value function
|
|
||||||
self.critic = MLP([mlp_input_dim_c, *critic_hidden_dims, 1], activation=activation)
|
|
||||||
|
|
||||||
print(f"Actor MLP: {self.actor}")
|
|
||||||
print(f"Critic MLP: {self.critic}")
|
|
||||||
print(f"Teacher Encoder: {self.teacher_encoder}")
|
|
||||||
print(f"Student MoE Encoder: {self.student_moe_encoder}")
|
|
||||||
|
|
||||||
# Action noise
|
|
||||||
self.std = nn.Parameter(init_noise_std * torch.ones(num_actions))
|
|
||||||
self.distribution = None
|
|
||||||
# disable args validation for speedup
|
|
||||||
Normal.set_default_validate_args = False
|
|
||||||
|
|
||||||
@staticmethod
|
|
||||||
# not used at the moment
|
|
||||||
def init_weights(sequential, scales):
|
|
||||||
[torch.nn.init.orthogonal_(module.weight, gain=scales[idx]) for idx, module in
|
|
||||||
enumerate(mod for mod in sequential if isinstance(mod, nn.Linear))]
|
|
||||||
|
|
||||||
|
|
||||||
def reset(self, dones=None):
|
|
||||||
self.history[dones > 0] = 0.0
|
|
||||||
|
|
||||||
def forward(self):
|
|
||||||
raise NotImplementedError
|
|
||||||
|
|
||||||
@property
|
|
||||||
def action_mean(self):
|
|
||||||
return self.distribution.mean
|
|
||||||
|
|
||||||
@property
|
|
||||||
def action_std(self):
|
|
||||||
return self.distribution.stddev
|
|
||||||
|
|
||||||
@property
|
|
||||||
def entropy(self):
|
|
||||||
return self.distribution.entropy().sum(dim=-1)
|
|
||||||
|
|
||||||
def update_distribution(self, latent_and_obs):
|
|
||||||
mean = self.actor(latent_and_obs)
|
|
||||||
self.distribution = Normal(mean, mean*0. + self.std)
|
|
||||||
|
|
||||||
def act(self, obs, privileged_obs, history, is_teacher, **kwargs):
|
|
||||||
if is_teacher:
|
|
||||||
latent = self.teacher_encoder(privileged_obs)
|
|
||||||
else:
|
|
||||||
with torch.no_grad():
|
|
||||||
latent, _ = self.student_moe_encoder(history)
|
|
||||||
x = torch.cat([latent, obs], dim=1)
|
|
||||||
self.update_distribution(x)
|
|
||||||
return self.distribution.sample()
|
|
||||||
|
|
||||||
def get_actions_log_prob(self, actions):
|
|
||||||
return self.distribution.log_prob(actions).sum(dim=-1)
|
|
||||||
|
|
||||||
def act_inference(self, obs):
|
|
||||||
self.history = torch.cat([self.history[:, 1:], obs.unsqueeze(1)], dim=1)
|
|
||||||
latent, _ = self.student_moe_encoder(self.history.flatten(1))
|
|
||||||
x = torch.cat([latent, obs], dim=1)
|
|
||||||
actions_mean = self.actor(x)
|
|
||||||
return actions_mean
|
|
||||||
|
|
||||||
def evaluate(self, privileged_obs, history, is_teacher, **kwargs):
|
|
||||||
if is_teacher:
|
|
||||||
latent = self.teacher_encoder(privileged_obs)
|
|
||||||
else:
|
|
||||||
latent, _ = self.student_moe_encoder(history)
|
|
||||||
x = torch.cat([latent.detach(), privileged_obs], dim=1)
|
|
||||||
value = self.critic(x)
|
|
||||||
return value
|
|
||||||
@@ -115,7 +115,7 @@ class MoE(nn.Module):
|
|||||||
|
|
||||||
# Gating network
|
# Gating network
|
||||||
self.gating_network = nn.Sequential(
|
self.gating_network = nn.Sequential(
|
||||||
MLP([input_dim, *hidden_dims, expert_num], activation),
|
MLP([input_dim, *hidden_dims[:-1], expert_num], activation),
|
||||||
nn.Softmax(dim=-1)
|
nn.Softmax(dim=-1)
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
@@ -102,9 +102,12 @@ class OnPolicyRunner:
|
|||||||
|
|
||||||
# robogauge client
|
# robogauge client
|
||||||
try:
|
try:
|
||||||
|
if not train_cfg['robogauge']['enabled']:
|
||||||
|
raise ImportError("config disabled")
|
||||||
from robogauge.scripts.client import RoboGaugeClient
|
from robogauge.scripts.client import RoboGaugeClient
|
||||||
self.robogauge_client = RoboGaugeClient()
|
self.robogauge_client = RoboGaugeClient(f"http://127.0.0.1:{train_cfg['robogauge']['port']}")
|
||||||
except:
|
except Exception as e:
|
||||||
|
print(f"[INFO] RoboGauge client could not be initialized: {e}, disabling RoboGauge interface.")
|
||||||
self.robogauge_client = None
|
self.robogauge_client = None
|
||||||
|
|
||||||
def learn(self, num_learning_iterations, init_at_random_ep_len=False):
|
def learn(self, num_learning_iterations, init_at_random_ep_len=False):
|
||||||
|
|||||||
@@ -36,8 +36,8 @@ import statistics
|
|||||||
from torch.utils.tensorboard import SummaryWriter
|
from torch.utils.tensorboard import SummaryWriter
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
from rsl_rl.algorithms import CTS, MoECTS, MCPCTS, ACMoECTS, DualMoECTS, REMCTS
|
from rsl_rl.algorithms import CTS, MoENGCTS, MCPCTS, ACMoECTS, DualMoECTS, MoECTS
|
||||||
from rsl_rl.modules import ActorCriticCTS, ActorCriticMoECTS, ActorCriticMCPCTS, ActorCriticACMoECTS, ActorCriticDualMoECTS, ActorCriticREMCTS
|
from rsl_rl.modules import ActorCriticCTS, ActorCriticMoENGCTS, ActorCriticMCPCTS, ActorCriticACMoECTS, ActorCriticDualMoECTS, ActorCriticMoECTS
|
||||||
from rsl_rl.env import VecEnv
|
from rsl_rl.env import VecEnv
|
||||||
|
|
||||||
import yaml
|
import yaml
|
||||||
@@ -79,7 +79,7 @@ class OnPolicyRunnerCTS:
|
|||||||
num_critic_obs = self.env.num_obs
|
num_critic_obs = self.env.num_obs
|
||||||
history_length = train_cfg["history_length"]
|
history_length = train_cfg["history_length"]
|
||||||
actor_critic_class = eval(self.cfg["policy_class_name"])
|
actor_critic_class = eval(self.cfg["policy_class_name"])
|
||||||
model: Union[ActorCriticCTS, ActorCriticMoECTS, ActorCriticMCPCTS, ActorCriticACMoECTS, ActorCriticDualMoECTS, ActorCriticREMCTS] = actor_critic_class(
|
model: Union[ActorCriticCTS, ActorCriticMoENGCTS, ActorCriticMCPCTS, ActorCriticACMoECTS, ActorCriticDualMoECTS, ActorCriticMoECTS] = actor_critic_class(
|
||||||
self.env.num_obs,
|
self.env.num_obs,
|
||||||
num_critic_obs,
|
num_critic_obs,
|
||||||
self.env.num_actions,
|
self.env.num_actions,
|
||||||
@@ -87,7 +87,7 @@ class OnPolicyRunnerCTS:
|
|||||||
history_length,
|
history_length,
|
||||||
**self.policy_cfg).to(self.device)
|
**self.policy_cfg).to(self.device)
|
||||||
alg_class = eval(self.cfg["algorithm_class_name"])
|
alg_class = eval(self.cfg["algorithm_class_name"])
|
||||||
self.alg: Union[CTS, MoECTS, MCPCTS, ACMoECTS, DualMoECTS, REMCTS] = alg_class(model, self.env.num_envs, history_length, device=self.device, **self.alg_cfg)
|
self.alg: Union[CTS, MoENGCTS, MCPCTS, ACMoECTS, DualMoECTS, MoECTS] = alg_class(model, self.env.num_envs, history_length, device=self.device, **self.alg_cfg)
|
||||||
self.num_steps_per_env = self.cfg["num_steps_per_env"]
|
self.num_steps_per_env = self.cfg["num_steps_per_env"]
|
||||||
self.save_interval = self.cfg["save_interval"]
|
self.save_interval = self.cfg["save_interval"]
|
||||||
|
|
||||||
@@ -112,9 +112,12 @@ class OnPolicyRunnerCTS:
|
|||||||
|
|
||||||
# robogauge client
|
# robogauge client
|
||||||
try:
|
try:
|
||||||
|
if not train_cfg['robogauge']['enabled']:
|
||||||
|
raise ImportError("config disabled")
|
||||||
from robogauge.scripts.client import RoboGaugeClient
|
from robogauge.scripts.client import RoboGaugeClient
|
||||||
self.robogauge_client = RoboGaugeClient("http://127.0.0.1:9973") # Change PORT to your server port if needed, default is 9973
|
self.robogauge_client = RoboGaugeClient(f"http://127.0.0.1:{train_cfg['robogauge']['port']}")
|
||||||
except:
|
except Exception as e:
|
||||||
|
print(f"[INFO] RoboGauge client could not be initialized: {e}, disabling RoboGauge interface.")
|
||||||
self.robogauge_client = None
|
self.robogauge_client = None
|
||||||
|
|
||||||
def learn(self, num_learning_iterations, init_at_random_ep_len=False):
|
def learn(self, num_learning_iterations, init_at_random_ep_len=False):
|
||||||
@@ -183,7 +186,7 @@ class OnPolicyRunnerCTS:
|
|||||||
|
|
||||||
if self.cfg["algorithm_class_name"] in ["CTS", "MCPCTS"]:
|
if self.cfg["algorithm_class_name"] in ["CTS", "MCPCTS"]:
|
||||||
mean_value_loss, mean_surrogate_loss, mean_entropy_loss, mean_latent_loss = self.alg.update()
|
mean_value_loss, mean_surrogate_loss, mean_entropy_loss, mean_latent_loss = self.alg.update()
|
||||||
elif self.cfg["algorithm_class_name"] in ["MoECTS", "ACMoECTS", "REMCTS"]:
|
elif self.cfg["algorithm_class_name"] in ["MoECTS", "MoENGCTS", "ACMoECTS"]:
|
||||||
mean_value_loss, mean_surrogate_loss, mean_entropy_loss, mean_latent_loss, mean_load_balance_loss = self.alg.update()
|
mean_value_loss, mean_surrogate_loss, mean_entropy_loss, mean_latent_loss, mean_load_balance_loss = self.alg.update()
|
||||||
elif self.cfg["algorithm_class_name"] == "DualMoECTS":
|
elif self.cfg["algorithm_class_name"] == "DualMoECTS":
|
||||||
mean_value_loss, mean_surrogate_loss, mean_entropy_loss, mean_latent_loss, mean_load_balance_loss, mean_actor_load_balance_loss = self.alg.update()
|
mean_value_loss, mean_surrogate_loss, mean_entropy_loss, mean_latent_loss, mean_load_balance_loss, mean_actor_load_balance_loss = self.alg.update()
|
||||||
|
|||||||
2
setup.py
2
setup.py
@@ -2,7 +2,7 @@ from setuptools import find_packages
|
|||||||
from distutils.core import setup
|
from distutils.core import setup
|
||||||
|
|
||||||
setup(name='go2_rl_gym',
|
setup(name='go2_rl_gym',
|
||||||
version='1.0.1',
|
version='1.0.2',
|
||||||
author='Wu Tianyang',
|
author='Wu Tianyang',
|
||||||
license="MIT",
|
license="MIT",
|
||||||
packages=find_packages(),
|
packages=find_packages(),
|
||||||
|
|||||||
Reference in New Issue
Block a user