v0.1.5; fix student moe encoder bug

This commit is contained in:
wty-yy
2026-01-07 16:34:38 +08:00
parent e4aa714eab
commit 858314fcfa
2 changed files with 89 additions and 62 deletions

View File

@@ -15,7 +15,7 @@ import numpy as np
import torch import torch
import torch.nn as nn import torch.nn as nn
from torch.distributions import Normal from torch.distributions import Normal
from rsl_rl.modules.utils import MLP, MoE, Experts, L2Norm, SimNorm from rsl_rl.modules.utils import MLP, MoE, StudentMoEEncoder, Experts, L2Norm, SimNorm
class ActorCriticDualMoECTS(nn.Module): class ActorCriticDualMoECTS(nn.Module):
is_recurrent = False is_recurrent = False
@@ -56,12 +56,13 @@ class ActorCriticDualMoECTS(nn.Module):
) )
# Student encoder # Student encoder
self.student_moe_encoder = MoE( self.student_moe_encoder = StudentMoEEncoder(
expert_num=expert_num, expert_num=expert_num,
input_dim=mlp_input_dim_s, input_dim=mlp_input_dim_s,
hidden_dims=student_encoder_hidden_dims, hidden_dims=student_encoder_hidden_dims,
output_dim=latent_dim, output_dim=latent_dim,
activation=activation, activation=activation,
norm_type=norm_type,
) )
# MCP Actor # MCP Actor

View File

@@ -2,6 +2,70 @@ import torch
import torch.nn as nn import torch.nn as nn
import torch.nn.functional as F import torch.nn.functional as F
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})"
class MLP(nn.Module):
def __init__(self, dims, activation='elu', last_activation=False):
super().__init__()
activation = get_activation(activation)
layers = []
last_dim = dims[0]
for h_dim in dims[1:-1]:
layers.append(nn.Linear(last_dim, h_dim))
layers.append(activation)
last_dim = h_dim
layers.append(nn.Linear(last_dim, dims[-1]))
if last_activation:
layers.append(activation)
self.network = nn.Sequential(*layers)
def forward(self, x):
return self.network(x)
class Experts(nn.Module): class Experts(nn.Module):
def __init__(self, def __init__(self,
expert_num, expert_num,
@@ -61,66 +125,28 @@ class MoE(nn.Module):
output = torch.sum(weights.unsqueeze(-1) * expert_outs, dim=1) # (B, output_dim) output = torch.sum(weights.unsqueeze(-1) * expert_outs, dim=1) # (B, output_dim)
return output, weights return output, weights
class MLP(nn.Module): class StudentMoEEncoder(nn.Module):
def __init__(self, dims, activation='elu', last_activation=False): def __init__(
self,
expert_num,
input_dim,
hidden_dims,
output_dim,
activation='elu',
norm_type='l2norm',
):
super().__init__() super().__init__()
activation = get_activation(activation) self.norm_layer = L2Norm() if norm_type == 'l2norm' else SimNorm()
layers = [] self.moe = MoE(
last_dim = dims[0] expert_num=expert_num,
for h_dim in dims[1:-1]: input_dim=input_dim,
layers.append(nn.Linear(last_dim, h_dim)) hidden_dims=hidden_dims,
layers.append(activation) output_dim=output_dim,
last_dim = h_dim activation=activation,
layers.append(nn.Linear(last_dim, dims[-1])) )
if last_activation:
layers.append(activation)
self.network = nn.Sequential(*layers)
def forward(self, x):
return self.network(x)
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): def forward(self, obs):
super().__init__() latent, weights = self.moe(obs)
latent = self.norm_layer(latent)
return latent, weights
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})"