MRBEAN_Vision / train.py
VishavGupta01's picture
Upload 17 files
3e852ae verified
Raw
History Blame Contribute Delete
7.51 kB
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}%")