feat: 更新框架,新增数据及定义和backbone

This commit is contained in:
gouhanke
2026-02-05 01:37:55 +08:00
parent 92660562fb
commit dd2749cb12
10 changed files with 224 additions and 134 deletions
+10 -10
View File
@@ -30,17 +30,17 @@ class MLP(nn.Module):
class SinusoidalPositionalEncoding(nn.Module):
def __init__(
self,
emb_dim
embed_dim
):
super().__init__()
self.emb_dim = emb_dim
self.embed_dim = embed_dim
def forward(self, timesteps):
timesteps = timesteps.float()
B, T = timesteps.shape
device = timesteps.device
half_dim = self.emb_dim // 2
half_dim = self.embed_dim // 2
exponent = -torch.arange(half_dim, dtype=torch.float, device=device) * (
torch.log(torch.tensor(10000.0)) / half_dim
@@ -58,14 +58,14 @@ class ActionEncoder(nn.Module):
def __init__(
self,
action_dim,
emb_dim,
embed_dim,
):
super().__init__()
self.W1 = nn.Linear(action_dim, emb_dim)
self.W1 = nn.Linear(action_dim, embed_dim)
self.W2 = nn.Linear(2 * action_dim, action_dim)
self.W3 = nn.Linear(emb_dim, emb_dim)
self.pos_encoder = SinusoidalPositionalEncoding(emb_dim)
self.W3 = nn.Linear(embed_dim, embed_dim)
self.pos_encoder = SinusoidalPositionalEncoding(embed_dim)
def forward(
self,
@@ -89,13 +89,13 @@ class StateEncoder(nn.Module):
self,
state_dim,
hidden_dim,
emb_dim
embed_dim
):
super().__init__()
self.mlp = MLP(
state_dim,
hidden_dim,
emb_dim
embed_dim
)
def forward(
@@ -103,4 +103,4 @@ class StateEncoder(nn.Module):
states
):
state_emb = self.mlp(states)
return state_emb # [B, 1, emb_dim]
return state_emb # [B, 1, embed_dim]