Download model.py from dariussasarman/ROSPIN-Land-Classification: direct link, hf CLI and curl.
- Browser
- Download file 1.14 kB
-
https://huggingface.co/dariussasarman/ROSPIN-Land-Classification/resolve/main/model.py
- Command line
-
hf download hf://dariussasarman/ROSPIN-Land-Classification/model.py
-
curl -L -o model.py https://huggingface.co/dariussasarman/ROSPIN-Land-Classification/resolve/main/model.py
1.14 kB
| import torch | |
| import torch.nn as nn | |
| import numpy as np | |
| import os | |
| from torchvision import transforms | |
| from torchvision.models import resnet18, ResNet18_Weights | |
| class ResNet18_M3(nn.Module): | |
| def __init__(self): | |
| super().__init__() | |
| self.model = resnet18(weights=ResNet18_Weights.DEFAULT) | |
| self.model.fc = nn.Sequential( | |
| nn.Linear(self.model.fc.in_features, 200), | |
| nn.ReLU(), | |
| nn.Dropout(p=0.3), | |
| nn.Linear(200, 100), | |
| nn.ReLU(), | |
| nn.Dropout(p=0.3), | |
| nn.Linear(100, 10) | |
| ) | |
| def forward(self, x): | |
| return self.model(x) | |
| def get_instance(): | |
| DEVICE = torch.device("cuda" if torch.cuda.is_available() else "cpu") | |
| weights_path = os.path.join(os.path.dirname(__file__), "resnet18_m3_best.pth") | |
| model = ResNet18_M3().to(DEVICE) | |
| checkpoint = torch.load( | |
| weights_path, | |
| map_location=DEVICE, | |
| weights_only=False | |
| ) | |
| model.load_state_dict(checkpoint["model_state_dict"]) | |
| model.eval() | |
| return model | |