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