| 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()) |
|
|
|
|