xdecoder/Instruct-X-Decoder
163
1# Copyright (c) Facebook, Inc. and its affiliates.2# Modified by Bowen Cheng from: https://github.com/facebookresearch/detr/blob/master/models/transformer.py3"""4Transformer class.5 6Copy-paste from torch.nn.Transformer with modifications:7 * positional encodings are passed in MHattention8 * extra LN at the end of encoder is removed9 * decoder returns a stack of activations from all decoding layers10"""11import copy12from typing import List, Optional13 14import torch15import torch.nn.functional as F16from torch import Tensor, nn17 18 19class Transformer(nn.Module):20 def __init__(21 self,22 d_model=512,23 nhead=8,24 num_encoder_layers=6,25 num_decoder_layers=6,26 dim_feedforward=2048,27 dropout=0.1,28 activation="relu",29 normalize_before=False,30 return_intermediate_dec=False,31 ):32 super().__init__()33 34 encoder_layer = TransformerEncoderLayer(35 d_model, nhead, dim_feedforward, dropout, activation, normalize_before36 )37 encoder_norm = nn.LayerNorm(d_model) if normalize_before else None38 self.encoder = TransformerEncoder(encoder_layer, num_encoder_layers, encoder_norm)39 40 decoder_layer = TransformerDecoderLayer(41 d_model, nhead, dim_feedforward, dropout, activation, normalize_before42 )43 decoder_norm = nn.LayerNorm(d_model)44 self.decoder = TransformerDecoder(45 decoder_layer,46 num_decoder_layers,47 decoder_norm,48 return_intermediate=return_intermediate_dec,49 )50 51 self._reset_parameters()52 53 self.d_model = d_model54 self.nhead = nhead55 56 def _reset_parameters(self):57 for p in self.parameters():58 if p.dim() > 1:59 nn.init.xavier_uniform_(p)60 61 def forward(self, src, mask, query_embed, pos_embed):62 # flatten NxCxHxW to HWxNxC63 bs, c, h, w = src.shape64 src = src.flatten(2).permute(2, 0, 1)65 pos_embed = pos_embed.flatten(2).permute(2, 0, 1)66 query_embed = query_embed.unsqueeze(1).repeat(1, bs, 1)67 if mask is not None:68 mask = mask.flatten(1)69 70 tgt = torch.zeros_like(query_embed)71 memory = self.encoder(src, src_key_padding_mask=mask, pos=pos_embed)72 hs = self.decoder(73 tgt, memory, memory_key_padding_mask=mask, pos=pos_embed, query_pos=query_embed74 )75 return hs.transpose(1, 2), memory.permute(1, 2, 0).view(bs, c, h, w)76 77 78class TransformerEncoder(nn.Module):79 def __init__(self, encoder_layer, num_layers, norm=None):80 super().__init__()81 self.layers = _get_clones(encoder_layer, num_layers)82 self.num_layers = num_layers83 self.norm = norm84 85 def forward(86 self,87 src,88 mask: Optional[Tensor] = None,89 src_key_padding_mask: Optional[Tensor] = None,90 pos: Optional[Tensor] = None,91 ):92 output = src93 94 for layer in self.layers:95 output = layer(96 output, src_mask=mask, src_key_padding_mask=src_key_padding_mask, pos=pos97 )98 99 if self.norm is not None:100 output = self.norm(output)101 102 return output103 104 105class TransformerDecoder(nn.Module):106 def __init__(self, decoder_layer, num_layers, norm=None, return_intermediate=False):107 super().__init__()108 self.layers = _get_clones(decoder_layer, num_layers)109 self.num_layers = num_layers110 self.norm = norm111 self.return_intermediate = return_intermediate112 113 def forward(114 self,115 tgt,116 memory,117 tgt_mask: Optional[Tensor] = None,118 memory_mask: Optional[Tensor] = None,119 tgt_key_padding_mask: Optional[Tensor] = None,120 memory_key_padding_mask: Optional[Tensor] = None,121 pos: Optional[Tensor] = None,122 query_pos: Optional[Tensor] = None,123 ):124 output = tgt125 126 intermediate = []127 128 for layer in self.layers:129 output = layer(130 output,131 memory,132 tgt_mask=tgt_mask,133 memory_mask=memory_mask,134 tgt_key_padding_mask=tgt_key_padding_mask,135 memory_key_padding_mask=memory_key_padding_mask,136 pos=pos,137 query_pos=query_pos,138 )139 if self.return_intermediate:140 intermediate.append(self.norm(output))141 142 if self.norm is not None:143 output = self.norm(output)144 if self.return_intermediate:145 intermediate.pop()146 intermediate.append(output)147 148 if self.return_intermediate:149 return torch.stack(intermediate)150 151 return output.unsqueeze(0)152 153 154class TransformerEncoderLayer(nn.Module):155 def __init__(156 self,157 d_model,158 nhead,159 dim_feedforward=2048,160 dropout=0.1,161 activation="relu",162 normalize_before=False,163 ):164 super().__init__()165 self.self_attn = nn.MultiheadAttention(d_model, nhead, dropout=dropout)166 # Implementation of Feedforward model167 self.linear1 = nn.Linear(d_model, dim_feedforward)168 self.dropout = nn.Dropout(dropout)169 self.linear2 = nn.Linear(dim_feedforward, d_model)170 171 self.norm1 = nn.LayerNorm(d_model)172 self.norm2 = nn.LayerNorm(d_model)173 self.dropout1 = nn.Dropout(dropout)174 self.dropout2 = nn.Dropout(dropout)175 176 self.activation = _get_activation_fn(activation)177 self.normalize_before = normalize_before178 179 def with_pos_embed(self, tensor, pos: Optional[Tensor]):180 return tensor if pos is None else tensor + pos181 182 def forward_post(183 self,184 src,185 src_mask: Optional[Tensor] = None,186 src_key_padding_mask: Optional[Tensor] = None,187 pos: Optional[Tensor] = None,188 ):189 q = k = self.with_pos_embed(src, pos)190 191 src2 = self.self_attn(192 q, k, value=src, attn_mask=src_mask, key_padding_mask=src_key_padding_mask193 )[0]194 src = src + self.dropout1(src2)195 src = self.norm1(src)196 src2 = self.linear2(self.dropout(self.activation(self.linear1(src))))197 src = src + self.dropout2(src2)198 src = self.norm2(src)199 return src200 201 def forward_pre(202 self,203 src,204 src_mask: Optional[Tensor] = None,205 src_key_padding_mask: Optional[Tensor] = None,206 pos: Optional[Tensor] = None,207 ):208 src2 = self.norm1(src)209 q = k = self.with_pos_embed(src2, pos)210 src2 = self.self_attn(211 q, k, value=src2, attn_mask=src_mask, key_padding_mask=src_key_padding_mask212 )[0]213 src = src + self.dropout1(src2)214 src2 = self.norm2(src)215 src2 = self.linear2(self.dropout(self.activation(self.linear1(src2))))216 src = src + self.dropout2(src2)217 return src218 219 def forward(220 self,221 src,222 src_mask: Optional[Tensor] = None,223 src_key_padding_mask: Optional[Tensor] = None,224 pos: Optional[Tensor] = None,225 ):226 if self.normalize_before:227 return self.forward_pre(src, src_mask, src_key_padding_mask, pos)228 return self.forward_post(src, src_mask, src_key_padding_mask, pos)229 230 231class TransformerDecoderLayer(nn.Module):232 def __init__(233 self,234 d_model,235 nhead,236 dim_feedforward=2048,237 dropout=0.1,238 activation="relu",239 normalize_before=False,240 ):241 super().__init__()242 self.self_attn = nn.MultiheadAttention(d_model, nhead, dropout=dropout)243 self.multihead_attn = nn.MultiheadAttention(d_model, nhead, dropout=dropout)244 # Implementation of Feedforward model245 self.linear1 = nn.Linear(d_model, dim_feedforward)246 self.dropout = nn.Dropout(dropout)247 self.linear2 = nn.Linear(dim_feedforward, d_model)248 249 self.norm1 = nn.LayerNorm(d_model)250 self.norm2 = nn.LayerNorm(d_model)251 self.norm3 = nn.LayerNorm(d_model)252 self.dropout1 = nn.Dropout(dropout)253 self.dropout2 = nn.Dropout(dropout)254 self.dropout3 = nn.Dropout(dropout)255 256 self.activation = _get_activation_fn(activation)257 self.normalize_before = normalize_before258 259 def with_pos_embed(self, tensor, pos: Optional[Tensor]):260 return tensor if pos is None else tensor + pos261 262 def forward_post(263 self,264 tgt,265 memory,266 tgt_mask: Optional[Tensor] = None,267 memory_mask: Optional[Tensor] = None,268 tgt_key_padding_mask: Optional[Tensor] = None,269 memory_key_padding_mask: Optional[Tensor] = None,270 pos: Optional[Tensor] = None,271 query_pos: Optional[Tensor] = None,272 ):273 q = k = self.with_pos_embed(tgt, query_pos)274 tgt2 = self.self_attn(275 q, k, value=tgt, attn_mask=tgt_mask, key_padding_mask=tgt_key_padding_mask276 )[0]277 tgt = tgt + self.dropout1(tgt2)278 tgt = self.norm1(tgt)279 tgt2 = self.multihead_attn(280 query=self.with_pos_embed(tgt, query_pos),281 key=self.with_pos_embed(memory, pos),282 value=memory,283 attn_mask=memory_mask,284 key_padding_mask=memory_key_padding_mask,285 )[0]286 tgt = tgt + self.dropout2(tgt2)287 tgt = self.norm2(tgt)288 tgt2 = self.linear2(self.dropout(self.activation(self.linear1(tgt))))289 tgt = tgt + self.dropout3(tgt2)290 tgt = self.norm3(tgt)291 return tgt292 293 def forward_pre(294 self,295 tgt,296 memory,297 tgt_mask: Optional[Tensor] = None,298 memory_mask: Optional[Tensor] = None,299 tgt_key_padding_mask: Optional[Tensor] = None,300 memory_key_padding_mask: Optional[Tensor] = None,301 pos: Optional[Tensor] = None,302 query_pos: Optional[Tensor] = None,303 ):304 tgt2 = self.norm1(tgt)305 q = k = self.with_pos_embed(tgt2, query_pos)306 tgt2 = self.self_attn(307 q, k, value=tgt2, attn_mask=tgt_mask, key_padding_mask=tgt_key_padding_mask308 )[0]309 tgt = tgt + self.dropout1(tgt2)310 tgt2 = self.norm2(tgt)311 tgt2 = self.multihead_attn(312 query=self.with_pos_embed(tgt2, query_pos),313 key=self.with_pos_embed(memory, pos),314 value=memory,315 attn_mask=memory_mask,316 key_padding_mask=memory_key_padding_mask,317 )[0]318 tgt = tgt + self.dropout2(tgt2)319 tgt2 = self.norm3(tgt)320 tgt2 = self.linear2(self.dropout(self.activation(self.linear1(tgt2))))321 tgt = tgt + self.dropout3(tgt2)322 return tgt323 324 def forward(325 self,326 tgt,327 memory,328 tgt_mask: Optional[Tensor] = None,329 memory_mask: Optional[Tensor] = None,330 tgt_key_padding_mask: Optional[Tensor] = None,331 memory_key_padding_mask: Optional[Tensor] = None,332 pos: Optional[Tensor] = None,333 query_pos: Optional[Tensor] = None,334 ):335 if self.normalize_before:336 return self.forward_pre(337 tgt,338 memory,339 tgt_mask,340 memory_mask,341 tgt_key_padding_mask,342 memory_key_padding_mask,343 pos,344 query_pos,345 )346 return self.forward_post(347 tgt,348 memory,349 tgt_mask,350 memory_mask,351 tgt_key_padding_mask,352 memory_key_padding_mask,353 pos,354 query_pos,355 )356 357 358def _get_clones(module, N):359 return nn.ModuleList([copy.deepcopy(module) for i in range(N)])360 361 362def _get_activation_fn(activation):363 """Return an activation function given a string"""364 if activation == "relu":365 return F.relu366 if activation == "gelu":367 return F.gelu368 if activation == "glu":369 return F.glu370 raise RuntimeError(f"activation should be relu/gelu, not {activation}.")371 