import torch.nn as nn from torchvision.models import resnet50, ResNet50_Weights class MLPHead(nn.Module): def __init__(self, in_features, hidden_dim, num_classes): super().__init__() self.net = nn.Sequential( nn.Linear(in_features, hidden_dim), nn.BatchNorm1d(hidden_dim), nn.ReLU(inplace=True), nn.Linear(hidden_dim, hidden_dim), nn.BatchNorm1d(hidden_dim), nn.ReLU(inplace=True), nn.Linear(hidden_dim, num_classes), ) def forward(self, x): return self.net(x) class ResNetMLP(nn.Module): def __init__(self, num_classes=10, freeze_backbone=True, hidden_dim=512): super().__init__() weights = ResNet50_Weights.DEFAULT self.resnet = resnet50(weights=weights) if freeze_backbone: for p in self.resnet.parameters(): p.requires_grad = False for p in self.resnet.layer3.parameters(): p.requires_grad = True for p in self.resnet.layer4.parameters(): p.requires_grad = True num_features = self.resnet.fc.in_features self.resnet.fc = nn.Identity() self.mlp_head = MLPHead( in_features=num_features, hidden_dim=hidden_dim, num_classes=num_classes, ) def forward(self, x): x = self.resnet(x) x = x.view(x.size(0), -1) x = self.mlp_head(x) return x