| 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'] |
| |
| |
| 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() |
|
|
| |
| torch.manual_seed(args.seed) |
| np.random.seed(args.seed) |
|
|
| |
| 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 "")) |
|
|
| |
| if args.auditor_augs: |
| aug_list = [ |
| |
| ] |
| else: |
| aug_list = [] |
|
|
| |
| 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]) |
| ]) |
|
|
| |
| 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))) |
|
|
| |
| 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) |
|
|
| |
| device = torch.device(f"cuda:{args.cuda}" if torch.cuda.is_available() else "cpu") |
|
|
| |
| model = models.resnet50(weights=models.ResNet50_Weights.IMAGENET1K_V1) |
| model.fc = nn.Linear(model.fc.in_features, 5) |
| model = model.to(device) |
|
|
| |
| 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) |
|
|
| |
| optimizer = optim.Adam(model.parameters(), lr=0.001) |
| criterion = nn.BCEWithLogitsLoss(pos_weight=pos_weights) |
|
|
| |
| scaler = torch.amp.GradScaler('cuda') |
|
|
| |
| n_epochs = 10 |
| |
| |
| 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' |
| ) |
|
|
| |
| for epoch in range(n_epochs): |
| |
| 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() |
| |
| |
| 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() |
| |
| |
| 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 |
| |
| |
| 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() |
| |
| |
| 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 |
| |
| |
| 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}') |
| |
| |
| 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() |