feat: 更新框架,新增数据及定义和backbone
This commit is contained in:
@@ -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]
|
||||
Reference in New Issue
Block a user