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:
@@ -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
|
||||
@@ -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()
|
||||
|
||||
@@ -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()
|
||||
@@ -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
|
||||
@@ -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,
|
||||
|
||||
@@ -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})"
|
||||
|
||||
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
|
||||
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)
|
||||
)
|
||||
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user