import torch.nn as nn
from torchvision import models


def create_model(num_classes: int = 2) -> nn.Module:
    weights = models.MobileNet_V3_Large_Weights.DEFAULT
    model = models.mobilenet_v3_large(weights=weights)


    # 替换分类头
    if isinstance(model.classifier, nn.Sequential):
        in_features = model.classifier[-1].in_features
        model.classifier[-1] = nn.Linear(in_features, num_classes)
    else:
        in_features = getattr(model.classifier, "in_features", 1280)
        model.classifier = nn.Linear(in_features, num_classes)
    return model
