MCPcopy Create free account
hub / github.com/AtlasAnalyticsLab/AdaFisher / TransformerClassifier

Class TransformerClassifier

Image_Classification/src/models/cct.py:218–311  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

216
217
218class TransformerClassifier(nn.Module):
219 def __init__(self,
220 seq_pool=True,
221 embedding_dim=768,
222 num_layers=12,
223 num_heads=12,
224 mlp_ratio=4.0,
225 num_classes=1000,
226 dropout_rate=0.1,
227 attention_dropout=0.1,
228 stochastic_depth_rate=0.1,
229 positional_embedding='sine',
230 sequence_length=None,
231 *args, **kwargs):
232 super().__init__()
233 assert positional_embedding in {'sine', 'learnable', 'none'}
234
235 dim_feedforward = int(embedding_dim * mlp_ratio)
236 self.embedding_dim = embedding_dim
237 self.sequence_length = sequence_length
238 self.seq_pool = seq_pool
239
240 assert exists(sequence_length) or positional_embedding == 'none', \
241 f"Positional embedding is set to {positional_embedding} and" \
242 f" the sequence length was not specified."
243
244 if not seq_pool:
245 sequence_length += 1
246 self.class_emb = nn.Parameter(torch.zeros(1, 1, self.embedding_dim), requires_grad=True)
247 else:
248 self.attention_pool = nn.Linear(self.embedding_dim, 1)
249
250 if positional_embedding == 'none':
251 self.positional_emb = None
252 elif positional_embedding == 'learnable':
253 self.positional_emb = nn.Parameter(torch.zeros(1, sequence_length, embedding_dim),
254 requires_grad=True)
255 nn.init.trunc_normal_(self.positional_emb, std=0.2)
256 else:
257 self.positional_emb = nn.Parameter(sinusoidal_embedding(sequence_length, embedding_dim),
258 requires_grad=False)
259
260 self.dropout = nn.Dropout(p=dropout_rate)
261
262 dpr = [x.item() for x in torch.linspace(0, stochastic_depth_rate, num_layers)]
263
264 self.blocks = nn.ModuleList([
265 TransformerEncoderLayer(d_model=embedding_dim, nhead=num_heads,
266 dim_feedforward=dim_feedforward, dropout=dropout_rate,
267 attention_dropout=attention_dropout, drop_path_rate=layer_dpr)
268 for layer_dpr in dpr])
269
270 self.norm = nn.LayerNorm(embedding_dim)
271
272 self.fc = nn.Linear(embedding_dim, num_classes)
273 self.apply(self.init_weight)
274
275 def forward(self, x):

Callers 1

__init__Method · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected