summaryrefslogtreecommitdiff
path: root/models.py
blob: ab21055a588d20a54111f3499264fef2ead9e0fd (plain)
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
import torch.nn as nn
from torchvision.models import (
    resnet50, ResNet50_Weights,
    efficientnet_b0, EfficientNet_B0_Weights,
    vit_b_16, ViT_B_16_Weights,
)


def create_model(model_name: str, num_classes: int = 5, pretrained: bool = True) -> nn.Module:
    """Create a pretrained model with the classification head replaced for num_classes.

    Supported model_name values: 'resnet50', 'efficientnet_b0', 'vit_b_16'
    All backbone layers are trainable (full fine-tuning).
    """
    if model_name == 'resnet50':
        weights = ResNet50_Weights.DEFAULT if pretrained else None
        model = resnet50(weights=weights)
        model.fc = nn.Linear(model.fc.in_features, num_classes)  # 2048 -> num_classes

    elif model_name == 'efficientnet_b0':
        weights = EfficientNet_B0_Weights.DEFAULT if pretrained else None
        model = efficientnet_b0(weights=weights)
        # classifier is Sequential(Dropout(0.2), Linear(1280, 1000))
        model.classifier[1] = nn.Linear(model.classifier[1].in_features, num_classes)

    elif model_name == 'vit_b_16':
        weights = ViT_B_16_Weights.DEFAULT if pretrained else None
        model = vit_b_16(weights=weights)
        # heads is Sequential(head=Linear(768, 1000))
        model.heads.head = nn.Linear(model.heads.head.in_features, num_classes)

    else:
        raise ValueError(f"Unknown model: {model_name!r}. Choose from: resnet50, efficientnet_b0, vit_b_16")

    return model