File size: 610 Bytes
45461c9 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 | import torch
import torch.nn as nn
class Classifier(nn.Module):
def __init__(self, encoder, num_classes, bottleneck_dim=256):
super().__init__()
self.encoder = encoder
self.embed_dim = self.encoder.embed_dim
self.head = torch.nn.Sequential(
nn.Linear(self.embed_dim, bottleneck_dim),
nn.BatchNorm1d(bottleneck_dim),
nn.ReLU(),
nn.Linear(bottleneck_dim, num_classes)
)
def forward(self, x):
x = self.encoder(x)
if type(x) == tuple:
x = x[0]
x = self.head(x)
return x
|