Team Ai
Apppublic

OpenMotionLab/MotionGPT

sourceHugging Facemitupdated 1y agoView on Hugging Face
118likes
position_encoding.py193 linesDownload Raw Back to utils
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