| 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 | |