| 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}%") |