era-resnet / train.py
Arnab Sinha
Initial commit
59563c2
Raw
History Blame Contribute Delete
12.9 kB
import argparse
import math
import os
import torch
import torch.nn as nn
import torch.nn.functional as F
from torch.optim import SGD
from torch.optim.lr_scheduler import OneCycleLR
from torch.cuda.amp import autocast, GradScaler
import torchvision
import torchvision.transforms as T
from torchvision.utils import save_image, make_grid
import numpy as np
from tqdm import tqdm
from model import resnet18_cifar
# ---------- Utilities ----------
CIFAR100_MEAN = (0.5071, 0.4867, 0.4408)
CIFAR100_STD = (0.2675, 0.2565, 0.2761)
def set_seed(seed=42):
import random
random.seed(seed)
np.random.seed(seed)
torch.manual_seed(seed)
torch.cuda.manual_seed_all(seed)
torch.backends.cudnn.deterministic = False
torch.backends.cudnn.benchmark = True
def get_dataloaders(data_dir, batch_size, num_workers=2):
normalize = T.Normalize(CIFAR100_MEAN, CIFAR100_STD)
train_tfms = T.Compose([
T.RandomCrop(32, padding=4, padding_mode="reflect"),
T.RandomHorizontalFlip(),
T.ColorJitter(0.2, 0.2, 0.2, 0.1),
T.ToTensor(),
normalize,
])
test_tfms = T.Compose([
T.ToTensor(),
normalize,
])
train_set = torchvision.datasets.CIFAR100(root=data_dir, train=True, download=True, transform=train_tfms)
test_set = torchvision.datasets.CIFAR100(root=data_dir, train=False, download=True, transform=test_tfms)
train_loader = torch.utils.data.DataLoader(
train_set, batch_size=batch_size, shuffle=True, num_workers=num_workers, pin_memory=True, drop_last=True
)
test_loader = torch.utils.data.DataLoader(
test_set, batch_size=batch_size, shuffle=False, num_workers=num_workers, pin_memory=True
)
return train_loader, test_loader, train_set, test_set
def accuracy(output, target, topk=(1,)):
with torch.no_grad():
maxk = max(topk)
batch_size = target.size(0)
_, pred = output.topk(maxk, 1, True, True)
pred = pred.t()
correct = pred.eq(target.view(1, -1).expand_as(pred))
res = []
for k in topk:
correct_k = correct[:k].reshape(-1).float().sum(0, keepdim=True)
res.append((correct_k.mul_(100.0 / batch_size)).item())
return res
def denormalize(imgs):
'''imgs: (N,3,H,W) normalized -> return in [0,1].'''
mean = torch.tensor(CIFAR100_MEAN, device=imgs.device).view(1,3,1,1)
std = torch.tensor(CIFAR100_STD, device=imgs.device).view(1,3,1,1)
return (imgs * std) + mean
# ---------- Grad-CAM ----------
class GradCAM:
'''
Minimal Grad-CAM for CNNs.
Hooks target_layer (e.g., model.layer4[-1].conv2) to capture activations & gradients.
'''
def __init__(self, model, target_layer):
self.model = model
self.target_layer = target_layer
self.activations = None
self.gradients = None
def fwd_hook(module, inp, out):
self.activations = out.detach()
def bwd_hook(module, grad_in, grad_out):
# grad_out is a tuple; take grad w.r.t. output of the layer
self.gradients = grad_out[0].detach()
self.h1 = self.target_layer.register_forward_hook(fwd_hook)
try:
self.h2 = self.target_layer.register_full_backward_hook(bwd_hook)
except AttributeError:
self.h2 = self.target_layer.register_backward_hook(bwd_hook)
def __del__(self):
try:
self.h1.remove()
self.h2.remove()
except Exception:
pass
def generate(self, images, class_idx):
'''
images: (N,3,H,W) normalized
class_idx: (N,) tensor of ints (target class per image)
Returns: cams in shape (N, H, W) normalized to [0,1]
'''
self.model.zero_grad(set_to_none=True)
self.gradients = None
self.activations = None
outputs = self.model(images) # (N, num_classes)
scores = outputs.gather(1, class_idx.view(-1, 1)).sum()
scores.backward()
acts = self.activations # (N, C, h, w)
grads = self.gradients # (N, C, h, w)
weights = grads.mean(dim=(2, 3), keepdim=True) # (N, C, 1, 1)
cam = (weights * acts).sum(dim=1) # (N, h, w)
cam = F.relu(cam)
cam_min = cam.view(cam.size(0), -1).min(dim=1)[0].view(-1, 1, 1)
cam_max = cam.view(cam.size(0), -1).max(dim=1)[0].view(-1, 1, 1)
cam = (cam - cam_min) / (cam_max - cam_min + 1e-8)
return cam
def cam_overlay_batch(images_norm, cams):
'''
images_norm: (N,3,H,W) normalized images in model space
cams: (N,H,W) in [0,1]
Returns: (N,3,H,W) overlay images in [0,1]
'''
imgs = denormalize(images_norm).clamp(0, 1)
cams_up = F.interpolate(cams.unsqueeze(1), size=imgs.shape[-2:], mode="bilinear", align_corners=False).squeeze(1)
heat_rgb = cams_up.unsqueeze(1).repeat(1, 3, 1, 1) # (N,3,H,W)
overlay = (0.6 * imgs + 0.4 * heat_rgb).clamp(0, 1)
return overlay
# ---------- Training / Eval ----------
def train_one_epoch(model, loader, criterion, optimizer, scheduler, device, scaler, clip_grad=None, amp_dtype=torch.float16):
model.train()
running_loss = 0.0
n = 0
top1_meter = 0.0
top5_meter = 0.0
pbar = tqdm(loader, desc="Train", leave=False)
for images, targets in pbar:
images = images.to(device, non_blocking=True)
targets = targets.to(device, non_blocking=True)
optimizer.zero_grad(set_to_none=True)
with autocast(dtype=amp_dtype):
outputs = model(images)
loss = criterion(outputs, targets)
scaler.scale(loss).backward()
if clip_grad is not None and clip_grad > 0:
scaler.unscale_(optimizer)
torch.nn.utils.clip_grad_norm_(model.parameters(), clip_grad)
scaler.step(optimizer)
scaler.update()
if scheduler is not None:
scheduler.step()
bs = images.size(0)
running_loss += loss.item() * bs
n += bs
acc1, acc5 = accuracy(outputs, targets, topk=(1, 5))
top1_meter += acc1 * bs / 100.0
top5_meter += acc5 * bs / 100.0
pbar.set_postfix(loss=running_loss/n, acc1=100*top1_meter/n, acc5=100*top5_meter/n)
return running_loss / n, (100*top1_meter/n), (100*top5_meter/n)
@torch.no_grad()
def evaluate(model, loader, criterion, device, amp_dtype=torch.float16):
model.eval()
running_loss = 0.0
n = 0
top1_meter = 0.0
top5_meter = 0.0
for images, targets in tqdm(loader, desc="Eval", leave=False):
images = images.to(device, non_blocking=True)
targets = targets.to(device, non_blocking=True)
with autocast(dtype=amp_dtype):
outputs = model(images)
loss = criterion(outputs, targets)
bs = images.size(0)
running_loss += loss.item() * bs
n += bs
acc1, acc5 = accuracy(outputs, targets, topk=(1, 5))
top1_meter += acc1 * bs / 100.0
top5_meter += acc5 * bs / 100.0
return running_loss / n, (100*top1_meter/n), (100*top5_meter/n)
def save_checkpoint(state, is_best, work_dir):
os.makedirs(work_dir, exist_ok=True)
torch.save(state, os.path.join(work_dir, "last.pth"))
if is_best:
torch.save(state, os.path.join(work_dir, "best.pth"))
def run_gradcam_and_save(model, images_norm, classes, device, out_path, amp_dtype=torch.float16):
model.eval()
images = images_norm.to(device)
with torch.no_grad(), autocast(dtype=amp_dtype):
logits = model(images)
preds = logits.argmax(dim=1)
target_layer = model.layer4[-1].conv2
cam = GradCAM(model, target_layer)
with torch.cuda.amp.autocast(enabled=False):
cams = cam.generate(images, preds)
overlays = cam_overlay_batch(images, cams)
grid = make_grid(overlays, nrow=int(math.sqrt(overlays.size(0))), padding=2)
save_image(grid, out_path)
with open(out_path.replace(".png", ".txt"), "w") as f:
names = [classes[int(i)] for i in preds.cpu()]
f.write("\\n".join(names))
# ---------- Main ----------
def parse_args():
p = argparse.ArgumentParser(description="Train ResNet-18 on CIFAR-100 with OneCycle and Grad-CAM")
p.add_argument("--data-dir", type=str, default="./data", help="dataset directory")
p.add_argument("--work-dir", type=str, default="./outputs", help="where to save logs/checkpoints")
p.add_argument("--epochs", type=int, default=100)
p.add_argument("--batch-size", type=int, default=128)
p.add_argument("--max-lr", type=float, default=0.1, help="OneCycle max LR")
p.add_argument("--weight-decay", type=float, default=5e-4)
p.add_argument("--momentum", type=float, default=0.9)
p.add_argument("--seed", type=int, default=42)
p.add_argument("--num-workers", type=int, default=2)
p.add_argument("--label-smoothing", type=float, default=0.1)
p.add_argument("--width", type=int, default=64, help="base channels (64 is standard ResNet-18)")
p.add_argument("--gradcam-interval", type=int, default=10, help="save Grad-CAM every N epochs")
p.add_argument("--gradcam-samples", type=int, default=16, help="how many val images to visualize")
p.add_argument("--clip-grad", type=float, default=0.5, help="gradient clipping (0 to disable)")
p.add_argument("--amp-dtype", type=str, default="float16", choices=["float16","bfloat16"], help="mixed precision dtype")
return p.parse_args()
def main():
args = parse_args()
set_seed(args.seed)
device = "cuda" if torch.cuda.is_available() else "cpu"
amp_dtype = torch.float16 if args.amp_dtype == "float16" else torch.bfloat16
train_loader, test_loader, train_set, test_set = get_dataloaders(args.data_dir, args.batch_size, args.num_workers)
num_classes = 100
model = resnet18_cifar(num_classes=num_classes, width=args.width).to(device)
criterion = nn.CrossEntropyLoss(label_smoothing=args.label_smoothing).to(device)
optimizer = SGD(model.parameters(), lr=args.max_lr/25.0, momentum=args.momentum, nesterov=True, weight_decay=args.weight_decay)
total_steps = args.epochs * len(train_loader)
scheduler = OneCycleLR(
optimizer,
max_lr=args.max_lr,
total_steps=total_steps,
pct_start=0.3,
anneal_strategy="cos",
div_factor=25.0,
final_div_factor=1e4,
three_phase=False,
)
scaler = GradScaler(enabled=(device=="cuda"))
os.makedirs(args.work_dir, exist_ok=True)
fixed_images = []
fixed_targets = []
for images, targets in test_loader:
fixed_images.append(images[:args.gradcam_samples])
fixed_targets.append(targets[:args.gradcam_samples])
break
fixed_images = torch.cat(fixed_images, dim=0)[:args.gradcam_samples]
fixed_targets = torch.cat(fixed_targets, dim=0)[:args.gradcam_samples]
best_acc1 = 0.0
history = []
for epoch in range(1, args.epochs + 1):
print(f"Epoch {epoch}/{args.epochs}")
train_loss, train_acc1, train_acc5 = train_one_epoch(
model, train_loader, criterion, optimizer, scheduler, device, scaler,
clip_grad=args.clip_grad, amp_dtype=amp_dtype
)
val_loss, val_acc1, val_acc5 = evaluate(model, test_loader, criterion, device, amp_dtype=amp_dtype)
print(f" Train: loss={train_loss:.4f}, acc1={train_acc1:.2f}%, acc5={train_acc5:.2f}%")
print(f" Val : loss={val_loss:.4f}, acc1={val_acc1:.2f}%, acc5={val_acc5:.2f}%")
history.append((epoch, train_loss, train_acc1, val_loss, val_acc1))
is_best = val_acc1 > best_acc1
best_acc1 = max(best_acc1, val_acc1)
save_checkpoint({
"epoch": epoch,
"model_state": model.state_dict(),
"optimizer_state": optimizer.state_dict(),
"scheduler_state": scheduler.state_dict(),
"scaler_state": scaler.state_dict(),
"best_acc1": best_acc1,
"args": vars(args),
}, is_best, args.work_dir)
if (epoch % args.gradcam_interval) == 0:
out_cam_path = os.path.join(args.work_dir, f"gradcam_epoch_{epoch:03d}.png")
print(f" Saving Grad-CAM overlays -> {out_cam_path}")
run_gradcam_and_save(model, fixed_images.to(device), test_set.classes, device, out_cam_path, amp_dtype=amp_dtype)
try:
import csv
with open(os.path.join(args.work_dir, "history.csv"), "w", newline="") as f:
w = csv.writer(f)
w.writerow(["epoch","train_loss","train_acc1","val_loss","val_acc1"])
for row in history:
w.writerow(row)
except Exception as e:
print("Could not save history.csv:", e)
print(f"Training complete. Best Top-1 Val Acc: {best_acc1:.2f}%")
if __name__ == "__main__":
main()