import os import time import copy import pandas as pd import torch import torch.nn as nn import torch.optim as optim import torchvision.models as models import torchvision.transforms as transforms from torch.utils.data import Dataset, DataLoader from PIL import Image, ImageFile ImageFile.LOAD_TRUNCATED_IMAGES = True class MedicDataset(Dataset): def __init__(self, dataframe, dataset_root, transform=None): self.df = dataframe self.root_dir = dataset_root self.transform = transform self.disaster_mapping = { 'earthquake': 0, 'fire': 1, 'flood': 2, 'hurricane': 3, 'landslide': 4, 'not_disaster': 5 } self.severity_mapping = { 'little_or_none': 0, 'mild': 1, 'severe': 2 } def __len__(self): return len(self.df) def __getitem__(self, idx): relative_path = self.df.iloc[idx]['image_path'] img_path = os.path.join(self.root_dir, relative_path) image = Image.open(img_path).convert('RGBA').convert('RGB') disaster_str = self.df.iloc[idx]['disaster_types'] severity_str = self.df.iloc[idx]['damage_severity'] disaster_label = torch.tensor(self.disaster_mapping[disaster_str], dtype=torch.long) severity_label = torch.tensor(self.severity_mapping[severity_str], dtype=torch.long) if self.transform: image = self.transform(image) return image, disaster_label, severity_label def prepare_dataframes(dataset_root): target_disasters = ['earthquake', 'fire', 'flood', 'hurricane', 'landslide', 'not_disaster'] target_severities = ['severe', 'mild', 'little_or_none'] train_tsv = os.path.join(dataset_root, "MEDIC_train.tsv") df_train_raw = pd.read_csv(train_tsv, sep='\t') df_train = df_train_raw[ df_train_raw['disaster_types'].isin(target_disasters) & df_train_raw['damage_severity'].isin(target_severities) ].copy() df_train.reset_index(drop=True, inplace=True) val_tsv = os.path.join(dataset_root, "MEDIC_dev.tsv") df_val_raw = pd.read_csv(val_tsv, sep='\t') df_val = df_val_raw[ df_val_raw['disaster_types'].isin(target_disasters) & df_val_raw['damage_severity'].isin(target_severities) ].copy() df_val.reset_index(drop=True, inplace=True) return df_train, df_val 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) local_weights_path = "./models/mobilenet_v3_large-8738ca79.pth" if os.path.exists(local_weights_path): mobilenet.load_state_dict(torch.load(local_weights_path)) print("Successfully loaded pre-trained MobileNet weights.") else: print("Warning: Local weights file not found. Training from scratch.") 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) out_disaster = self.disaster_head(x) out_severity = self.severity_head(x) return out_disaster, out_severity if __name__ == '__main__': DATASET_ROOT = "./dataset" device = torch.device("cuda" if torch.cuda.is_available() else "cpu") df_train, df_val = prepare_dataframes(DATASET_ROOT) basic_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]) ]) train_dataset = MedicDataset(df_train, DATASET_ROOT, basic_transform) val_dataset = MedicDataset(df_val, DATASET_ROOT, basic_transform) train_loader = DataLoader( train_dataset, batch_size=64, shuffle=True, num_workers=4, pin_memory=True, drop_last=True, persistent_workers=True ) val_loader = DataLoader( val_dataset, batch_size=64, shuffle=False, num_workers=4, pin_memory=True, persistent_workers=True ) model = MRBEANVisionModel().to(device) criterion_disaster = nn.CrossEntropyLoss() criterion_severity = nn.CrossEntropyLoss() optimizer = optim.AdamW(model.parameters(), lr=1e-4, weight_decay=1e-2) scaler = torch.amp.GradScaler('cuda') EPOCHS = 25 best_val_acc = 0.0 best_model_wts = copy.deepcopy(model.state_dict()) print(f"\nStarting {EPOCHS}-Epoch Training Loop...") for epoch in range(EPOCHS): print(f"\n=== Epoch {epoch + 1}/{EPOCHS} ===") model.train() running_loss = 0.0 start_time = time.time() for batch_idx, (images, labels_d, labels_s) in enumerate(train_loader): images = images.to(device, non_blocking=True) labels_d = labels_d.to(device, non_blocking=True) labels_s = labels_s.to(device, non_blocking=True) optimizer.zero_grad() with torch.amp.autocast('cuda'): out_d, out_s = model(images) loss = criterion_disaster(out_d, labels_d) + criterion_severity(out_s, labels_s) scaler.scale(loss).backward() scaler.step(optimizer) scaler.update() running_loss += loss.item() avg_loss = running_loss / len(train_loader) epoch_time = time.time() - start_time print(f"Train Loss: {avg_loss:.4f} | Epoch Time: {epoch_time:.1f}s") model.eval() correct_d, correct_s, total = 0, 0, 0 with torch.no_grad(): for images, labels_d, labels_s in val_loader: images = images.to(device, non_blocking=True) labels_d = labels_d.to(device, non_blocking=True) labels_s = labels_s.to(device, non_blocking=True) with torch.amp.autocast('cuda'): out_d, out_s = model(images) _, pred_d = torch.max(out_d, 1) _, pred_s = torch.max(out_s, 1) total += labels_d.size(0) correct_d += (pred_d == labels_d).sum().item() correct_s += (pred_s == labels_s).sum().item() acc_d = 100 * correct_d / total acc_s = 100 * correct_s / total avg_val_acc = (acc_d + acc_s) / 2 print(f"Val Disaster Acc: {acc_d:.2f}% | Val Severity Acc: {acc_s:.2f}%") if avg_val_acc > best_val_acc: best_val_acc = avg_val_acc best_model_wts = copy.deepcopy(model.state_dict()) torch.save(best_model_wts, "best_multitask_vision_model.pth") print("Checkpoint: New best model saved to disk!") print(f"\nTraining Complete. Peak Average Validation Accuracy: {best_val_acc:.2f}%")