File size: 3,433 Bytes
b7e9b58
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
80
81
82
83
84
85
import torch
from torch import nn
import math

class DecisionTransformer(nn.Module):
    def __init__(self, obs_dim, act_dim, hidden=256, n_layers=4, n_heads=4, max_len=1024):
        super().__init__()
        self.hidden = hidden
        self.n_heads = n_heads
        self.max_len = max_len

        # ๊ฐ ์ž…๋ ฅ์„ hidden ์ฐจ์›์œผ๋กœ ์ž„๋ฒ ๋”ฉ
        self.obs_embed = nn.Linear(obs_dim, hidden)
        self.act_embed = nn.Linear(act_dim, hidden)
        self.rtg_embed = nn.Linear(1, hidden)
        self.time_embed = nn.Embedding(max_len, hidden)

        # Transformer Encoder
        layer = nn.TransformerEncoderLayer(
            d_model=hidden,
            nhead=n_heads,
            dim_feedforward=hidden * 4,
            batch_first=True,
            dropout=0.1
        )
        self.transformer = nn.TransformerEncoder(layer, num_layers=n_layers)

        # ํ–‰๋™ ์˜ˆ์ธก ํ—ค๋“œ
        self.act_head = nn.Linear(hidden, act_dim)

        # Causal Mask ์ƒ์„ฑ์„ ์œ„ํ•œ ๋ฒ„ํผ ๋“ฑ๋ก (ONNX export์‹œ ์‚ฌ์šฉ ์•ˆ ํ•  ์ˆ˜๋„ ์žˆ์ง€๋งŒ ํ˜ธํ™˜์„ฑ ์œ„ํ•ด ์œ ์ง€)
        # self.register_buffer("mask", torch.tril(torch.ones(max_len * 3, max_len * 3)))

    def forward(self, obs, act, rtg, timesteps):
        # obs: (B, T, obs_dim)
        # act: (B, T, act_dim)
        # rtg: (B, T, 1)
        # timesteps: (B, T)

        B, T, _ = obs.shape

        # 1. ์ž„๋ฒ ๋”ฉ (Embedding)
        obs_emb = self.obs_embed(obs)          # (B, T, hidden)
        act_emb = self.act_embed(act)          # (B, T, hidden)
        rtg_emb = self.rtg_embed(rtg)          # (B, T, hidden)
        time_emb = self.time_embed(timesteps)  # (B, T, hidden)

        # 2. Timestep Embedding ๋”ํ•˜๊ธฐ
        # ๋…ผ๋ฌธ์—์„œ๋Š” R, s, a ๋ชจ๋‘์— timestep embedding์„ ๋”ํ•จ
        obs_emb = obs_emb + time_emb
        act_emb = act_emb + time_emb
        rtg_emb = rtg_emb + time_emb

        # 3. Stacking (R_t, s_t, a_t) ์ˆœ์„œ๋กœ ์Œ“๊ธฐ
        # (B, T, 3, hidden) -> (B, 3*T, hidden)
        # dim=2์— stack ํ›„ flatten
        stacked_inputs = torch.stack((rtg_emb, obs_emb, act_emb), dim=2)
        stacked_inputs = stacked_inputs.view(B, T * 3, self.hidden)

        # 4. Causal Masking
        # ํ˜„์žฌ ์‹œํ€€์Šค ๊ธธ์ด(3*T)์— ๋งž๋Š” ๋งˆ์Šคํฌ ๋™์  ์ƒ์„ฑ
        seq_len = T * 3
        # (seq_len, seq_len) ํฌ๊ธฐ์˜ ๋งˆ์Šคํฌ ์ƒ์„ฑ
        # ๋Œ€๊ฐ์„  ์œ„์ชฝ(๋ฏธ๋ž˜)์„ -inf๋กœ ์ฑ„์›€ (Attention์—์„œ ๋ฌด์‹œ๋จ)
        # ๋Œ€๊ฐ์„  ํฌํ•จ ์•„๋ž˜์ชฝ(๊ณผ๊ฑฐ+ํ˜„์žฌ)์€ 0์œผ๋กœ ์œ ์ง€
        causal_mask = torch.triu(torch.full((seq_len, seq_len), float('-inf'), device=obs.device), diagonal=1)

        # 5. Transformer Forward
        # is_causal=True๋ฅผ ๋ช…์‹œํ•˜์—ฌ ๋‚ด๋ถ€์ ์ธ ๋งˆ์Šคํฌ ๊ฒ€์‚ฌ(data-dependent check)๋ฅผ ์šฐํšŒ
        x = self.transformer(stacked_inputs, mask=causal_mask, is_causal=True)

        # 6. Action Prediction
        # ์ž…๋ ฅ ์ˆœ์„œ๊ฐ€ (R_t, s_t, a_t) ์ด๋ฏ€๋กœ,
        # s_t์˜ ์ถœ๋ ฅ(index 1, 4, 7...)์„ ์‚ฌ์šฉํ•˜์—ฌ a_t๋ฅผ ์˜ˆ์ธกํ•ด์•ผ ํ•จ
        # x: (B, 3*T, hidden)
        # reshape -> (B, T, 3, hidden)
        x = x.view(B, T, 3, self.hidden)
        
        # s_t์— ํ•ด๋‹นํ•˜๋Š” ์ž„๋ฒ ๋”ฉ ์ถ”์ถœ (index 1)
        # R_t(0), s_t(1), a_t(2)
        state_preds = x[:, :, 1, :]  # (B, T, hidden)

        action_preds = self.act_head(state_preds)  # (B, T, act_dim)

        return action_preds