import torch import torch.nn as nn from torchvision import models def load_model(model_path, num_classes=5, device="cpu"): model = models.alexnet(weights=None) model.classifier[6] = nn.Linear(4096, num_classes) state_dict = torch.load(model_path, map_location=device) model.load_state_dict(state_dict) model.eval() model.to(device) return model