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 |