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

Method __init__

lib/models_timm/layers/ml_decoder.py:104–134  ·  view source on GitHub ↗
(self, num_classes, num_of_groups=-1, decoder_embedding=768, initial_num_features=2048)

Source from the content-addressed store, hash-verified

102
103class MLDecoder(nn.Module):
104 def __init__(self, num_classes, num_of_groups=-1, decoder_embedding=768, initial_num_features=2048):
105 super(MLDecoder, self).__init__()
106 embed_len_decoder = 100 if num_of_groups < 0 else num_of_groups
107 if embed_len_decoder > num_classes:
108 embed_len_decoder = num_classes
109
110 # switching to 768 initial embeddings
111 decoder_embedding = 768 if decoder_embedding < 0 else decoder_embedding
112 self.embed_standart = nn.Linear(initial_num_features, decoder_embedding)
113
114 # decoder
115 decoder_dropout = 0.1
116 num_layers_decoder = 1
117 dim_feedforward = 2048
118 layer_decode = TransformerDecoderLayerOptimal(d_model=decoder_embedding,
119 dim_feedforward=dim_feedforward, dropout=decoder_dropout)
120 self.decoder = nn.TransformerDecoder(layer_decode, num_layers=num_layers_decoder)
121
122 # non-learnable queries
123 self.query_embed = nn.Embedding(embed_len_decoder, decoder_embedding)
124 self.query_embed.requires_grad_(False)
125
126 # group fully-connected
127 self.num_classes = num_classes
128 self.duplicate_factor = int(num_classes / embed_len_decoder + 0.999)
129 self.duplicate_pooling = torch.nn.Parameter(
130 torch.Tensor(embed_len_decoder, decoder_embedding, self.duplicate_factor))
131 self.duplicate_pooling_bias = torch.nn.Parameter(torch.Tensor(num_classes))
132 torch.nn.init.xavier_normal_(self.duplicate_pooling)
133 torch.nn.init.constant_(self.duplicate_pooling_bias, 0)
134 self.group_fc = GroupFC(embed_len_decoder)
135
136 def forward(self, x):
137 if len(x.shape) == 4: # [bs,2048, 7,7]

Callers

nothing calls this directly

Calls 3

GroupFCClass · 0.85
__init__Method · 0.45

Tested by

no test coverage detected