| 216 | |
| 217 | |
| 218 | class 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): |