ARotting's picture
Publish Offline return-conditioned key-door policy
77a4175 verified
Raw
History Blame Contribute Delete
2.62 kB
from __future__ import annotations
import torch
from torch import nn
class DecisionTransformer(nn.Module):
def __init__(self, dimensions: int = 32, maximum_length: int = 20) -> None:
super().__init__()
self.maximum_length = maximum_length
self.position_embedding = nn.Embedding(9, dimensions)
self.key_embedding = nn.Embedding(2, dimensions)
self.previous_action_embedding = nn.Embedding(4, dimensions)
self.return_embedding = nn.Linear(1, dimensions)
self.timestep_embedding = nn.Embedding(maximum_length, dimensions)
layer = nn.TransformerEncoderLayer(
d_model=dimensions,
nhead=4,
dim_feedforward=64,
dropout=0.05,
batch_first=True,
activation="gelu",
norm_first=True,
)
self.transformer = nn.TransformerEncoder(layer, num_layers=2)
self.action_head = nn.Linear(dimensions, 3)
def forward(
self,
positions: torch.Tensor,
keys: torch.Tensor,
returns_to_go: torch.Tensor,
previous_actions: torch.Tensor,
valid: torch.Tensor,
) -> torch.Tensor:
length = positions.shape[1]
timesteps = torch.arange(length, device=positions.device)
hidden = (
self.position_embedding(positions)
+ self.key_embedding(keys)
+ self.previous_action_embedding(previous_actions)
+ self.return_embedding(returns_to_go[..., None])
+ self.timestep_embedding(timesteps)[None]
)
causal_mask = torch.triu(
torch.ones(length, length, device=positions.device, dtype=torch.bool),
diagonal=1,
)
hidden = self.transformer(
hidden,
mask=causal_mask,
src_key_padding_mask=~valid,
)
return self.action_head(hidden)
class BehaviorCloningPolicy(nn.Module):
def __init__(self) -> None:
super().__init__()
self.position_embedding = nn.Embedding(9, 8)
self.key_embedding = nn.Embedding(2, 4)
self.network = nn.Sequential(
nn.Linear(12, 32),
nn.ReLU(),
nn.Linear(32, 3),
)
def forward(self, positions: torch.Tensor, keys: torch.Tensor) -> torch.Tensor:
features = torch.cat(
[self.position_embedding(positions), self.key_embedding(keys)],
dim=-1,
)
return self.network(features)
def parameter_count(model: nn.Module) -> int:
return sum(parameter.numel() for parameter in model.parameters())