import torch.nn as nn from torchvision import models def build_model(num_classes): model = models.mobilenet_v2(weights=models.MobileNet_V2_Weights.DEFAULT) model.classifier[1] = nn.Sequential( nn.Linear(model.last_channel, 128), nn.ReLU(), nn.Dropout(0.5), nn.Linear(128, num_classes) ) return model