MCPcopy Create free account
hub / github.com/SLDGroup/MERIT / TransformerDecoderLayerOptimal

Class TransformerDecoderLayerOptimal

lib/models_timm/layers/ml_decoder.py:35–73  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

33
34
35class TransformerDecoderLayerOptimal(nn.Module):
36 def __init__(self, d_model, nhead=8, dim_feedforward=2048, dropout=0.1, activation="relu",
37 layer_norm_eps=1e-5) -> None:
38 super(TransformerDecoderLayerOptimal, self).__init__()
39 self.norm1 = nn.LayerNorm(d_model, eps=layer_norm_eps)
40 self.dropout = nn.Dropout(dropout)
41 self.dropout1 = nn.Dropout(dropout)
42 self.dropout2 = nn.Dropout(dropout)
43 self.dropout3 = nn.Dropout(dropout)
44
45 self.multihead_attn = nn.MultiheadAttention(d_model, nhead, dropout=dropout)
46
47 # Implementation of Feedforward model
48 self.linear1 = nn.Linear(d_model, dim_feedforward)
49 self.linear2 = nn.Linear(dim_feedforward, d_model)
50
51 self.norm2 = nn.LayerNorm(d_model, eps=layer_norm_eps)
52 self.norm3 = nn.LayerNorm(d_model, eps=layer_norm_eps)
53
54 self.activation = _get_activation_fn(activation)
55
56 def __setstate__(self, state):
57 if 'activation' not in state:
58 state['activation'] = torch.nn.functional.relu
59 super(TransformerDecoderLayerOptimal, self).__setstate__(state)
60
61 def forward(self, tgt: Tensor, memory: Tensor, tgt_mask: Optional[Tensor] = None,
62 memory_mask: Optional[Tensor] = None,
63 tgt_key_padding_mask: Optional[Tensor] = None,
64 memory_key_padding_mask: Optional[Tensor] = None) -> Tensor:
65 tgt = tgt + self.dropout1(tgt)
66 tgt = self.norm1(tgt)
67 tgt2 = self.multihead_attn(tgt, memory, memory)[0]
68 tgt = tgt + self.dropout2(tgt2)
69 tgt = self.norm2(tgt)
70 tgt2 = self.linear2(self.dropout(self.activation(self.linear1(tgt))))
71 tgt = tgt + self.dropout3(tgt2)
72 tgt = self.norm3(tgt)
73 return tgt
74
75
76# @torch.jit.script

Callers 1

__init__Method · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected