import random import torch import torch.nn as nn import torch.optim as optim import torchvision.models as models import torchvision.transforms as transforms from torchvision.transforms import v2 from torch.utils.data import Dataset, DataLoader import pandas as pd import numpy as np from PIL import Image import wandb import argparse from tqdm import tqdm class ChexpertDataset(Dataset): def __init__(self, csv_path, transform=None, is_train=True): self.df = pd.read_csv(csv_path) self.df = self.df.replace(np.nan, 0.0) self.df = self.df.replace(-1.0, 0.0) self.df = self.df[self.df['Frontal/Lateral'] == 'Frontal'] self.df = self.df[self.df['AP/PA'] == 'AP'] # Remove path prefix self.df['Path'] = self.df['Path'].str.replace('CheXpert-v1.0-small/train/', '') if not is_train: self.df['Path'] = self.df['Path'].str.replace('CheXpert-v1.0-small/valid/', '') self.paths = self.df['Path'].to_numpy() self.labels = self.df[['Atelectasis', 'Consolidation', 'Cardiomegaly', 'Pleural Effusion', 'Edema']].to_numpy() self.transform = transform if transform is not None else transforms.ToTensor() self.base_path = 'data/chexpert_resized_224/train/' if is_train else 'data/chexpert_resized_224/valid/' def __len__(self): return len(self.paths) def __getitem__(self, idx): img_path = self.base_path + str(self.paths[idx]) image = Image.open(img_path).convert('RGB') image = self.transform(image) label = torch.tensor(self.labels[idx], dtype=torch.float32) return image, label def main(): parser = argparse.ArgumentParser() parser.add_argument('--resize', type=int, default=224, help='Size to resize images to (default: 224)') parser.add_argument('--seed', type=int, default=1, help='Seed for random number generator (default: 1)') parser.add_argument('--cuda', type=int, default=1, help='CUDA device number (default: 1)') parser.add_argument('--auditor_augs', action='store_true', default=False, help='Enable auditor augmentations (default: False)') parser.add_argument('--auto_aug', action='store_true', default=False, help='Enable auto augmentations (default: False)') parser.add_argument('--subset_len', type=int, default=None, help='Length of subset to use for training (default: None, use full dataset)') args = parser.parse_args() # Set seeds torch.manual_seed(args.seed) np.random.seed(args.seed) # Initialize wandb wandb.init(project="ModelAuditor", name="CheXpert_ResNet50_" + str(args.seed) + "_" + str(args.resize) + ("_AuditorAugs" if args.auditor_augs else "") + ("_AutoAugs" if args.auto_aug else "")) # Define augmentations if args.auditor_augs: aug_list = [ # PUT HERE WHAT THE AUDITOR GIVES YOU ] else: aug_list = [] # Create transforms if args.auto_aug: train_transform = transforms.Compose([ transforms.AutoAugment(transforms.AutoAugmentPolicy.IMAGENET) ] + aug_list + [ transforms.ToTensor(), transforms.Normalize(mean=[0.5], std=[0.5]) ]) else: train_transform = transforms.Compose([ ] + aug_list + [ transforms.Normalize(mean=[0.5], std=[0.5]) ]) val_transform = transforms.Compose([ transforms.ToTensor(), transforms.Normalize(mean=[0.5], std=[0.5]) ]) # Create datasets train_dataset = ChexpertDataset('data/chexpert/train.csv', transform=train_transform, is_train=True) val_dataset = ChexpertDataset('data/chexpert/valid.csv', transform=val_transform, is_train=False) train_subset = torch.utils.data.Subset(train_dataset, torch.arange(max(0, len(train_dataset) - 5000))) # Create data loaders train_loader = DataLoader(train_subset, batch_size=64, shuffle=True, num_workers=1, pin_memory=True, persistent_workers=True) val_loader = DataLoader(val_dataset, batch_size=64, num_workers=1, pin_memory=True, persistent_workers=True) # Set device device = torch.device(f"cuda:{args.cuda}" if torch.cuda.is_available() else "cpu") # Initialize model model = models.resnet50(weights=models.ResNet50_Weights.IMAGENET1K_V1) model.fc = nn.Linear(model.fc.in_features, 5) # 5 classes for CheXpert model = model.to(device) # Calculate positive weights for each class pos_weights = [] for col in range(train_dataset.labels.shape[1]): num_positive = (train_dataset.labels[:, col] == 1.0).sum() num_negative = (train_dataset.labels[:, col] == 0.0).sum() pos_weights.append(num_negative / num_positive) pos_weights = torch.tensor(pos_weights).to(device) # Initialize optimizer and criterion optimizer = optim.Adam(model.parameters(), lr=0.001) criterion = nn.BCEWithLogitsLoss(pos_weight=pos_weights) # Initialize scaler for mixed precision scaler = torch.amp.GradScaler('cuda') # Training parameters n_epochs = 10 # Add learning rate scheduler warmup_epochs = 2 total_steps = len(train_loader) * n_epochs warmup_steps = len(train_loader) * warmup_epochs scheduler = optim.lr_scheduler.OneCycleLR( optimizer, max_lr=0.001, total_steps=total_steps, pct_start=warmup_steps/total_steps, anneal_strategy='cos' ) # Training loop for epoch in range(n_epochs): # Training phase model.train() train_loss = 0 train_correct = 0 train_total = 0 for x, y in tqdm(train_loader, desc=f'Epoch {epoch+1}/{n_epochs}'): x, y = x.to(device), y.to(device) optimizer.zero_grad() # Mixed precision training with torch.amp.autocast('cuda'): outputs = model(x) loss = criterion(outputs, y) scaler.scale(loss).backward() scaler.step(optimizer) scaler.update() scheduler.step() train_loss += loss.item() # Calculate training accuracy preds = (torch.sigmoid(outputs) > 0.5).float() train_correct += (preds == y).sum().item() train_total += y.numel() train_loss /= len(train_loader) train_acc = train_correct / train_total # Validation phase model.eval() val_loss = 0 val_correct = 0 val_total = 0 with torch.no_grad(): for x, y in val_loader: x, y = x.to(device), y.to(device) with torch.amp.autocast('cuda'): outputs = model(x) loss = criterion(outputs, y) val_loss += loss.item() # Calculate validation accuracy preds = (torch.sigmoid(outputs) > 0.5).float() val_correct += (preds == y).sum().item() val_total += y.numel() val_loss /= len(val_loader) val_acc = val_correct / val_total # Log metrics current_lr = scheduler.get_last_lr()[0] wandb.log({ "train_loss": train_loss, "val_loss": val_loss, "train_acc": train_acc, "val_acc": val_acc, "epoch": epoch + 1, "learning_rate": current_lr }) print(f'Epoch {epoch+1}: Train Loss: {train_loss:.4f}, Train Acc: {train_acc:.4f}, Val Loss: {val_loss:.4f}, Val Acc: {val_acc:.4f}') # Save model after each epoch torch.save(model.state_dict(), f'chexpert_resnet50_{args.seed}_{args.resize}' + ("_AuditorAugs" if args.auditor_augs else "") + ("_AutoAugs" if args.auto_aug else "") + '.pt') if __name__ == "__main__": main()