Download Upload/01_Source_Code/Python_Training/model_dt.py from code3939/DecisionTransformer-Unity-Sim: direct link, hf CLI and curl.
- Browser
- Download file 3.43 kB
-
https://huggingface.co/code3939/DecisionTransformer-Unity-Sim/resolve/main/Upload/01_Source_Code/Python_Training/model_dt.py
- Command line
-
hf download hf://code3939/DecisionTransformer-Unity-Sim/Upload/01_Source_Code/Python_Training/model_dt.py
-
curl -L -o model_dt.py https://huggingface.co/code3939/DecisionTransformer-Unity-Sim/resolve/main/Upload/01_Source_Code/Python_Training/model_dt.py
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 |