Spaces:
Sleeping
Sleeping
| 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) | |
| 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() |