code3939's picture
Upload 765 files
b7e9b58 verified
Raw History Blame Contribute Delete
3.43 kB
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