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