OpenMotionLab/MotionGPT
118
1# Copyright (c) Facebook, Inc. and its affiliates. All Rights Reserved2"""3Various positional encodings for the transformer.4"""5import math6from typing import List, Optional7 8import numpy as np9import torch10from torch import Tensor, nn11 12# from util.misc import NestedTensor13 14 15class NestedTensor(object):16 17 def __init__(self, tensors, mask: Optional[Tensor]):18 self.tensors = tensors19 self.mask = mask20 21 def to(self, device):22 # type: (Device) -> NestedTensor # noqa23 cast_tensor = self.tensors.to(device)24 mask = self.mask25 if mask is not None:26 assert mask is not None27 cast_mask = mask.to(device)28 else:29 cast_mask = None30 return NestedTensor(cast_tensor, cast_mask)31 32 def decompose(self):33 return self.tensors, self.mask34 35 def __repr__(self):36 return str(self.tensors)37 38 39class PositionEmbeddingSine(nn.Module):40 """41 This is a more standard version of the position embedding, very similar to the one42 used by the Attention is all you need paper, generalized to work on images.43 """44 45 def __init__(self,46 num_pos_feats=64,47 temperature=10000,48 normalize=False,49 scale=None):50 super().__init__()51 self.num_pos_feats = num_pos_feats52 self.temperature = temperature53 self.normalize = normalize54 if scale is not None and normalize is False:55 raise ValueError("normalize should be True if scale is passed")56 if scale is None:57 scale = 2 * math.pi58 self.scale = scale59 60 def forward(self, tensor_list: NestedTensor):61 x = tensor_list.tensors62 mask = tensor_list.mask63 assert mask is not None64 not_mask = ~mask65 y_embed = not_mask.cumsum(1, dtype=torch.float32)66 x_embed = not_mask.cumsum(2, dtype=torch.float32)67 if self.normalize:68 eps = 1e-669 y_embed = y_embed / (y_embed[:, -1:, :] + eps) * self.scale70 x_embed = x_embed / (x_embed[:, :, -1:] + eps) * self.scale71 72 dim_t = torch.arange(self.num_pos_feats,73 dtype=torch.float32,74 device=x.device)75 dim_t = self.temperature**(2 * (dim_t // 2) / self.num_pos_feats)76 77 pos_x = x_embed[:, :, :, None] / dim_t78 pos_y = y_embed[:, :, :, None] / dim_t79 pos_x = torch.stack(80 (pos_x[:, :, :, 0::2].sin(), pos_x[:, :, :, 1::2].cos()),81 dim=4).flatten(3)82 pos_y = torch.stack(83 (pos_y[:, :, :, 0::2].sin(), pos_y[:, :, :, 1::2].cos()),84 dim=4).flatten(3)85 pos = torch.cat((pos_y, pos_x), dim=3).permute(0, 3, 1, 2)86 return pos87 88 89class PositionEmbeddingLearned(nn.Module):90 """91 Absolute pos embedding, learned.92 """93 94 def __init__(self, num_pos_feats=256):95 super().__init__()96 self.row_embed = nn.Embedding(50, num_pos_feats)97 self.col_embed = nn.Embedding(50, num_pos_feats)98 self.reset_parameters()99 100 def reset_parameters(self):101 nn.init.uniform_(self.row_embed.weight)102 nn.init.uniform_(self.col_embed.weight)103 104 def forward(self, tensor_list: NestedTensor):105 x = tensor_list.tensors106 h, w = x.shape[-2:]107 i = torch.arange(w, device=x.device)108 j = torch.arange(h, device=x.device)109 x_emb = self.col_embed(i)110 y_emb = self.row_embed(j)111 pos = torch.cat([112 x_emb.unsqueeze(0).repeat(h, 1, 1),113 y_emb.unsqueeze(1).repeat(1, w, 1),114 ],115 dim=-1).permute(2, 0, 1).unsqueeze(0).repeat(116 x.shape[0], 1, 1, 1)117 return pos118 119 120class PositionEmbeddingSine1D(nn.Module):121 122 def __init__(self, d_model, max_len=500, batch_first=False):123 super().__init__()124 self.batch_first = batch_first125 126 pe = torch.zeros(max_len, d_model)127 position = torch.arange(0, max_len, dtype=torch.float).unsqueeze(1)128 div_term = torch.exp(129 torch.arange(0, d_model, 2).float() * (-np.log(10000.0) / d_model))130 pe[:, 0::2] = torch.sin(position * div_term)131 pe[:, 1::2] = torch.cos(position * div_term)132 pe = pe.unsqueeze(0).transpose(0, 1)133 134 self.register_buffer('pe', pe)135 136 def forward(self, x):137 # not used in the final model138 if self.batch_first:139 pos = self.pe.permute(1, 0, 2)[:, :x.shape[1], :]140 else:141 pos = self.pe[:x.shape[0], :]142 return pos143 144 145class PositionEmbeddingLearned1D(nn.Module):146 147 def __init__(self, d_model, max_len=500, batch_first=False):148 super().__init__()149 self.batch_first = batch_first150 # self.dropout = nn.Dropout(p=dropout)151 152 self.pe = nn.Parameter(torch.zeros(max_len, 1, d_model))153 # self.pe = pe.unsqueeze(0).transpose(0, 1)154 155 self.reset_parameters()156 157 def reset_parameters(self):158 nn.init.uniform_(self.pe)159 160 def forward(self, x):161 # not used in the final model162 if self.batch_first:163 pos = self.pe.permute(1, 0, 2)[:, :x.shape[1], :]164 else:165 x = x + self.pe[:x.shape[0], :]166 return x167 # return self.dropout(x)168 169 170def build_position_encoding(N_steps,171 position_embedding="sine",172 embedding_dim="1D"):173 # N_steps = hidden_dim // 2174 if embedding_dim == "1D":175 if position_embedding in ('v2', 'sine'):176 position_embedding = PositionEmbeddingSine1D(N_steps)177 elif position_embedding in ('v3', 'learned'):178 position_embedding = PositionEmbeddingLearned1D(N_steps)179 else:180 raise ValueError(f"not supported {position_embedding}")181 elif embedding_dim == "2D":182 if position_embedding in ('v2', 'sine'):183 # TODO find a better way of exposing other arguments184 position_embedding = PositionEmbeddingSine(N_steps, normalize=True)185 elif position_embedding in ('v3', 'learned'):186 position_embedding = PositionEmbeddingLearned(N_steps)187 else:188 raise ValueError(f"not supported {position_embedding}")189 else:190 raise ValueError(f"not supported {embedding_dim}")191 192 return position_embedding193 