File size: 2,617 Bytes
77a4175 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 | 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())
|