Cocoyawn32's picture
Add files using upload-large-folder tool
3f6e26d verified
Raw
History Blame Contribute Delete
2.38 kB
import torch
import torch.nn as nn
class ActionTransformerProjector(nn.Module):
def __init__(self, action_dim, hidden_size, depth=2, num_heads=4, mlp_ratio=4.0, max_len=64):
super().__init__()
self.input_proj = nn.Linear(action_dim, hidden_size)
self.pos_embed = nn.Parameter(torch.randn(1, max_len, hidden_size) * 0.02)
encoder_layer = nn.TransformerEncoderLayer(
d_model=hidden_size,
nhead=num_heads,
dim_feedforward=int(hidden_size * mlp_ratio),
activation="gelu",
batch_first=True,
norm_first=True
)
self.encoder = nn.TransformerEncoder(encoder_layer, num_layers=depth)
self.norm = nn.LayerNorm(hidden_size)
def forward(self, x):
x = self.input_proj(x)
seq_len = x.shape[1]
eff_len = min(seq_len, self.pos_embed.shape[1])
pos_embed_slice = self.pos_embed[:, :eff_len, :]
if seq_len > eff_len:
x[:, :eff_len, :] = x[:, :eff_len, :] + pos_embed_slice
else:
x = x + pos_embed_slice
x = self.encoder(x)
x = self.norm(x)
return x
class ActionTransformerDecoder(nn.Module):
def __init__(self, action_dim, hidden_size, depth=1, num_heads=4, mlp_ratio=4.0, max_len=64):
super().__init__()
self.pos_embed = nn.Parameter(torch.randn(1, max_len, hidden_size) * 0.02)
decoder_layer = nn.TransformerEncoderLayer(
d_model=hidden_size,
nhead=num_heads,
dim_feedforward=int(hidden_size * mlp_ratio),
activation="gelu",
batch_first=True,
norm_first=True
)
self.decoder = nn.TransformerEncoder(decoder_layer, num_layers=depth)
self.norm = nn.LayerNorm(hidden_size)
self.output_proj = nn.Linear(hidden_size, action_dim)
def forward(self, x):
seq_len = x.shape[1]
eff_len = min(seq_len, self.pos_embed.shape[1])
pos_embed_slice = self.pos_embed[:, :eff_len, :]
if seq_len > eff_len:
x[:, :eff_len, :] = x[:, :eff_len, :] + pos_embed_slice
else:
x = x + pos_embed_slice
x = self.decoder(x)
x = self.norm(x)
x = self.output_proj(x)
return x