MRBEAN_Vision / test.py
VishavGupta01's picture
Upload 17 files
3e852ae verified
Raw
History Blame Contribute Delete
2.08 kB
import torch
import torchvision.transforms as transforms
from PIL import Image
import torch.nn as nn
import torchvision.models as models
class MRBEANVisionModel(nn.Module):
def __init__(self, num_disaster_classes=6, num_severity_classes=3):
super(MRBEANVisionModel, self).__init__()
mobilenet = models.mobilenet_v3_large(weights=None)
self.features = mobilenet.features
self.pool = nn.AdaptiveAvgPool2d(1)
self.flatten = nn.Flatten()
self.disaster_head = nn.Sequential(
nn.Linear(960, 256), nn.ReLU(), nn.Dropout(0.3), nn.Linear(256, num_disaster_classes)
)
self.severity_head = nn.Sequential(
nn.Linear(960, 128), nn.ReLU(), nn.Dropout(0.3), nn.Linear(128, num_severity_classes)
)
def forward(self, x):
x = self.features(x)
x = self.pool(x)
x = self.flatten(x)
return self.disaster_head(x), self.severity_head(x)
disaster_labels = {0: 'Earthquake', 1: 'Fire', 2: 'Flood', 3: 'Hurricane', 4: 'Landslide', 5: 'Not a Disaster'}
severity_labels = {0: 'Little or None', 1: 'Mild', 2: 'Severe'}
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
model = MRBEANVisionModel().to(device)
print("Loading trained weights...")
model.load_state_dict(torch.load("best_multitask_vision_model.pth", map_location=device))
model.eval()
transform = transforms.Compose([
transforms.Resize((224, 224)),
transforms.ToTensor(),
transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])
])
test_image_path = "download.jpg"
image = Image.open(test_image_path).convert('RGB')
input_tensor = transform(image).unsqueeze(0).to(device)
with torch.no_grad():
out_disaster, out_severity = model(input_tensor)
_, pred_d = torch.max(out_disaster, 1)
_, pred_s = torch.max(out_severity, 1)
print("\n--- MRBEAN ANALYSIS RESULTS ---")
print(f"Disaster Type: {disaster_labels[pred_d.item()]}")
print(f"Damage Severity: {severity_labels[pred_s.item()]}")