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