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()