Merge branch 'master' of https://github.com/wty-yy/RoboGauge
This commit is contained in:
@@ -1,10 +1,11 @@
|
||||
from robogauge.utils.task_register import task_register
|
||||
from robogauge.tasks.simulator.mujoco_config import MujocoConfig
|
||||
from robogauge.tasks.robots import RobotConfig, Go2Config
|
||||
from robogauge.tasks.robots import RobotConfig, Go2Config, Go2MoEConfig
|
||||
from robogauge.tasks.pipeline import BasePipeline
|
||||
from robogauge.tasks.gauge import BaseGaugeConfig
|
||||
|
||||
from robogauge.tasks.custom.go2_flat_task import Go2FlatGaugeConfig
|
||||
from robogauge.tasks.custom.go2_flat_task import Go2FlatGaugeConfig, Go2FlatConfig, Go2MoEFlatConfig
|
||||
|
||||
task_register.register('base', BasePipeline, MujocoConfig, BaseGaugeConfig, RobotConfig)
|
||||
task_register.register('go2_flat', BasePipeline, MujocoConfig, Go2FlatGaugeConfig, Go2Config)
|
||||
task_register.register('go2_flat', BasePipeline, MujocoConfig, Go2FlatGaugeConfig, Go2FlatConfig)
|
||||
task_register.register('go2_moe_flat', BasePipeline, MujocoConfig, Go2FlatGaugeConfig, Go2MoEFlatConfig)
|
||||
|
||||
@@ -1,4 +1,4 @@
|
||||
from robogauge.tasks.robots import Go2Config
|
||||
from robogauge.tasks.robots import Go2Config, Go2MoEConfig
|
||||
from robogauge.tasks.gauge import FlatGaugeConfig
|
||||
|
||||
class Go2FlatGaugeConfig(FlatGaugeConfig):
|
||||
@@ -16,3 +16,19 @@ class Go2FlatGaugeConfig(FlatGaugeConfig):
|
||||
class diagonal_velocity(FlatGaugeConfig.goals.diagonal_velocity):
|
||||
enabled = True
|
||||
cmd_duration = 6.0
|
||||
|
||||
class Go2FlatConfig(Go2Config):
|
||||
class commands(Go2Config.commands):
|
||||
lin_vel_x = [-2.0, 2.0] # min max [m/s]
|
||||
lin_vel_y = [-2.0, 2.0] # min max [m/s]
|
||||
ang_vel_yaw = [-2.0, 2.0] # min max [rad/s]
|
||||
|
||||
class Go2MoEFlatConfig(Go2MoEConfig):
|
||||
class commands(Go2Config.commands):
|
||||
lin_vel_x = [-2.0, 2.0] # min max [m/s]
|
||||
lin_vel_y = [-2.0, 2.0] # min max [m/s]
|
||||
ang_vel_yaw = [-2.0, 2.0] # min max [rad/s]
|
||||
|
||||
class control(Go2Config.control):
|
||||
# model_path = "{ROBOGAUGE_ROOT_DIR}/resources/models/go2/go2_moe_cts_124k.pt"
|
||||
model_path = "/home/xfy/Coding/kaiwu2025/rob_finals/sim2real/models/v6-2_106503/kaiwu_script_v6-2_106503.pt"
|
||||
|
||||
@@ -12,7 +12,9 @@ from pathlib import Path
|
||||
|
||||
from robogauge.utils.logger import logger
|
||||
from robogauge.tasks.simulator import MujocoSimulator, MujocoConfig
|
||||
from robogauge.tasks.robots import BaseRobot, RobotConfig, Go2Config, Go2
|
||||
from robogauge.tasks.robots import (
|
||||
BaseRobot, RobotConfig, Go2Config, Go2, Go2MoEConfig, Go2MoE
|
||||
)
|
||||
from robogauge.tasks.gauge import BaseGauge, BaseGaugeConfig
|
||||
from robogauge.utils.helpers import class_to_dict
|
||||
|
||||
|
||||
@@ -2,3 +2,5 @@ from .base_robot_config import RobotConfig
|
||||
from .base_robot import BaseRobot
|
||||
from .go2.go2_config import Go2Config
|
||||
from .go2.go2 import Go2
|
||||
from .go2.go2_moe_config import Go2MoEConfig
|
||||
from .go2.go2_moe import Go2MoE
|
||||
|
||||
28
robogauge/tasks/robots/go2/go2_moe.py
Normal file
28
robogauge/tasks/robots/go2/go2_moe.py
Normal file
@@ -0,0 +1,28 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
'''
|
||||
@File : go2_moe.py
|
||||
@Time : 2025/12/05 17:29:16
|
||||
@Author : wty-yy
|
||||
@Version : 1.0
|
||||
@Blog : https://wty-yy.github.io/
|
||||
@Desc : None
|
||||
'''
|
||||
import torch
|
||||
import numpy as np
|
||||
|
||||
from robogauge.tasks.robots.base_robot import BaseRobot, get_projected_gravity
|
||||
from robogauge.tasks.robots.go2.go2_config import Go2Config
|
||||
from robogauge.tasks.robots.go2.go2 import Go2
|
||||
from robogauge.tasks.simulator.sim_data import SimData
|
||||
from robogauge.tasks.gauge.goal_data import GoalData
|
||||
from robogauge.utils.logger import logger
|
||||
|
||||
class Go2MoE(Go2):
|
||||
def get_action(self, obs: np.ndarray):
|
||||
obs_tensor = torch.tensor(obs, dtype=torch.float32).unsqueeze(0).to(self.device)
|
||||
action, weights = self.model(obs_tensor)
|
||||
action = action.detach().cpu().numpy().squeeze(0)[self.model2mj_idx]
|
||||
weights = weights.detach().cpu().numpy().squeeze(0)
|
||||
self.last_action = action
|
||||
target_dof_pos = action * self.action_scale + self.default_dof_pos
|
||||
return target_dof_pos, self.p_gains, self.d_gains, self.control_type
|
||||
16
robogauge/tasks/robots/go2/go2_moe_config.py
Normal file
16
robogauge/tasks/robots/go2/go2_moe_config.py
Normal file
@@ -0,0 +1,16 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
'''
|
||||
@File : go2_moe_config.py
|
||||
@Time : 2025/12/05 17:32:27
|
||||
@Author : wty-yy
|
||||
@Version : 1.0
|
||||
@Blog : https://wty-yy.github.io/
|
||||
@Desc : None
|
||||
'''
|
||||
from robogauge.tasks.robots.go2.go2_config import Go2Config
|
||||
|
||||
class Go2MoEConfig(Go2Config):
|
||||
robot_class = 'Go2MoE'
|
||||
|
||||
class control(Go2Config.control):
|
||||
model_path = "{ROBOGAUGE_ROOT_DIR}/resources/models/go2/go2_moe_cts_124k.pt"
|
||||
Reference in New Issue
Block a user