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:
wty-yy
2026-01-26 17:48:34 +08:00
parent 2aed91e7ce
commit 1798e67c29
23 changed files with 438 additions and 409 deletions

View File

@@ -30,8 +30,8 @@
from .ppo import PPO
from .cts import CTS
from .moe_cts import MoECTS
from .moe_ng_cts import MoENGCTS
from .mcp_cts import MCPCTS
from .ac_moe_cts import ACMoECTS
from .dual_moe_cts import DualMoECTS
from .rem_cts import REMCTS
from .moe_cts import MoECTS

View File

@@ -202,7 +202,7 @@ class MoECTS(CTS):
hid_states_batch, masks_batch
) = sample
# 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():
teacher_latent = self.model.teacher_encoder(privileged_obs_batch[teacher_samples:])
latent_loss = (teacher_latent - student_latent).pow(2).mean()

View File

@@ -33,12 +33,12 @@ import torch.nn as nn
import torch.optim as optim
import itertools
from rsl_rl.modules import ActorCriticMoECTS
from rsl_rl.modules import ActorCriticMoENGCTS
from rsl_rl.storage import RolloutStorageCTS
from rsl_rl.algorithms.cts import CTS
class REMCTS(CTS):
model: ActorCriticMoECTS
class MoENGCTS(CTS):
model: ActorCriticMoENGCTS
def __init__(self,
model,
num_envs,
@@ -202,7 +202,7 @@ class REMCTS(CTS):
hid_states_batch, masks_batch
) = sample
# 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():
teacher_latent = self.model.teacher_encoder(privileged_obs_batch[teacher_samples:])
latent_loss = (teacher_latent - student_latent).pow(2).mean()

View File

@@ -31,8 +31,8 @@
from .actor_critic import ActorCritic
from .actor_critic_recurrent import ActorCriticRecurrent
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_ac_moe_cts import ActorCriticACMoECTS
from .actor_critic_dual_moe_cts import ActorCriticDualMoECTS
from .actor_critic_rem_cts import ActorCriticREMCTS
from .actor_critic_moe_cts import ActorCriticMoECTS

View File

@@ -27,7 +27,7 @@ class ActorCriticDualMoECTS(nn.Module):
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], # last dim is expert hidden dim
student_encoder_hidden_dims=[512, 256, 256], # last dim is expert hidden dim
expert_num=8,
activation='elu',
init_noise_std=1.0,

View File

@@ -15,6 +15,8 @@ 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 ActorCriticMoECTS(nn.Module):
is_recurrent = False
def __init__(self, num_obs,
@@ -22,12 +24,11 @@ class ActorCriticMoECTS(nn.Module):
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,
student_encoder_hidden_dims=[512, 256, 256],
expert_num=8,
activation='elu',
init_noise_std=1.0,
latent_dim=32,
@@ -36,17 +37,12 @@ class ActorCriticMoECTS(nn.Module):
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(ActorCriticMoECTS, self).__init__()
super().__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_s = num_obs * history_length
mlp_input_dim_a = latent_dim + num_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)
# 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)
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_dim=mlp_input_dim_e,
gating_dim=mlp_input_dim_g,
expert_num=expert_num,
input_dim=mlp_input_dim_s,
hidden_dims=student_encoder_hidden_dims,
expert_num=student_expert_num,
latent_dim=latent_dim,
activation=activation_str
output_dim=latent_dim,
activation=activation,
norm_type=norm_type,
)
# 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)
self.actor = MLP([mlp_input_dim_a, *actor_hidden_dims, num_actions], activation=activation)
# 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)
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}")
@@ -113,10 +81,6 @@ class ActorCriticMoECTS(nn.Module):
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
@@ -152,7 +116,7 @@ class ActorCriticMoECTS(nn.Module):
latent = self.teacher_encoder(privileged_obs)
else:
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)
self.update_distribution(x)
return self.distribution.sample()
@@ -162,7 +126,7 @@ class ActorCriticMoECTS(nn.Module):
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))
latent, _ = self.student_moe_encoder(self.history.flatten(1))
x = torch.cat([latent, obs], dim=1)
actions_mean = self.actor(x)
return actions_mean
@@ -171,117 +135,7 @@ class ActorCriticMoECTS(nn.Module):
if is_teacher:
latent = self.teacher_encoder(privileged_obs)
else:
latent, _ = self.get_student_latent_and_weights(history)
latent, _ = self.student_moe_encoder(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})"

View 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})"

View File

@@ -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

View File

@@ -115,7 +115,7 @@ class MoE(nn.Module):
# Gating network
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)
)

View File

@@ -102,9 +102,12 @@ class OnPolicyRunner:
# robogauge client
try:
if not train_cfg['robogauge']['enabled']:
raise ImportError("config disabled")
from robogauge.scripts.client import RoboGaugeClient
self.robogauge_client = RoboGaugeClient()
except:
self.robogauge_client = RoboGaugeClient(f"http://127.0.0.1:{train_cfg['robogauge']['port']}")
except Exception as e:
print(f"[INFO] RoboGauge client could not be initialized: {e}, disabling RoboGauge interface.")
self.robogauge_client = None
def learn(self, num_learning_iterations, init_at_random_ep_len=False):

View File

@@ -36,8 +36,8 @@ import statistics
from torch.utils.tensorboard import SummaryWriter
import torch
from rsl_rl.algorithms import CTS, MoECTS, MCPCTS, ACMoECTS, DualMoECTS, REMCTS
from rsl_rl.modules import ActorCriticCTS, ActorCriticMoECTS, ActorCriticMCPCTS, ActorCriticACMoECTS, ActorCriticDualMoECTS, ActorCriticREMCTS
from rsl_rl.algorithms import CTS, MoENGCTS, MCPCTS, ACMoECTS, DualMoECTS, MoECTS
from rsl_rl.modules import ActorCriticCTS, ActorCriticMoENGCTS, ActorCriticMCPCTS, ActorCriticACMoECTS, ActorCriticDualMoECTS, ActorCriticMoECTS
from rsl_rl.env import VecEnv
import yaml
@@ -79,7 +79,7 @@ class OnPolicyRunnerCTS:
num_critic_obs = self.env.num_obs
history_length = train_cfg["history_length"]
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,
num_critic_obs,
self.env.num_actions,
@@ -87,7 +87,7 @@ class OnPolicyRunnerCTS:
history_length,
**self.policy_cfg).to(self.device)
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.save_interval = self.cfg["save_interval"]
@@ -112,9 +112,12 @@ class OnPolicyRunnerCTS:
# robogauge client
try:
if not train_cfg['robogauge']['enabled']:
raise ImportError("config disabled")
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
except:
self.robogauge_client = RoboGaugeClient(f"http://127.0.0.1:{train_cfg['robogauge']['port']}")
except Exception as e:
print(f"[INFO] RoboGauge client could not be initialized: {e}, disabling RoboGauge interface.")
self.robogauge_client = None
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"]:
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()
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()