""" ablation_full200_lambda01.py — single-config, full 200-epoch run for lambda_det=0.1, to test whether it beats the main result's 19.83 dB. """ import os import json import random import time from pathlib import Path import torch import torch.nn as nn import torch.nn.functional as F from torch.optim import Adam from torch.optim.lr_scheduler import CosineAnnealingLR from torch.utils.data import Dataset, DataLoader import torchvision.transforms as T import torchvision.transforms.functional as TF import torchvision.models as models from PIL import Image import numpy as np from ultralytics import YOLO DEVICE = ( torch.device('mps') if torch.backends.mps.is_available() else torch.device('cuda') if torch.cuda.is_available() else torch.device('cpu') ) print(f"Device: {DEVICE}") CFG = dict( lol_root = './LOL', crop_size = 256, base_ch = 32, yolo_weights = 'yolov8n.pt', batch_size = 4, lr = 1e-4, lambda_perc = 0.1, num_workers = 0, ckpt_dir = './checkpoints_ablation', ) Path(CFG['ckpt_dir']).mkdir(parents=True, exist_ok=True) ABLATION_LAMBDAS = [0.1] ABLATION_EPOCHS = 200 ABLATION_SEED = 42 RESULTS_PATH = 'ablation_full200_lambda01.json' class LOLDataset(Dataset): SPLIT_DIRS = {'train': 'our485', 'val': 'eval15'} _IMG_EXTS = {'.png', '.jpg', '.jpeg', '.bmp', '.tif', '.tiff'} def __init__(self, root, split='train', crop_size=256, augment=True): assert split in self.SPLIT_DIRS self.crop_size = crop_size self.augment = augment and (split == 'train') base = Path(root) / self.SPLIT_DIRS[split] self.low_dir = base / 'low' self.high_dir = base / 'high' for d in (self.low_dir, self.high_dir): if not d.is_dir(): raise FileNotFoundError(f"Not found: {d}") self.filenames = sorted( f for f in os.listdir(self.low_dir) if Path(f).suffix.lower() in self._IMG_EXTS ) missing = [f for f in self.filenames if not (self.high_dir / f).exists()] if missing: raise FileNotFoundError(f"Missing high images: {missing[:3]}") self.to_tensor = T.ToTensor() def __len__(self): return len(self.filenames) def __getitem__(self, idx): fname = self.filenames[idx] low_img = Image.open(self.low_dir / fname).convert('RGB') high_img = Image.open(self.high_dir / fname).convert('RGB') if self.crop_size is not None: low_img, high_img = self._paired_crop(low_img, high_img) if self.augment and random.random() > 0.5: low_img = TF.hflip(low_img) high_img = TF.hflip(high_img) return self.to_tensor(low_img), self.to_tensor(high_img), fname def _paired_crop(self, low, high): w, h = low.size c = self.crop_size if w < c or h < c: return TF.resize(low, [c, c]), TF.resize(high, [c, c]) top = random.randint(0, h - c) left = random.randint(0, w - c) return TF.crop(low, top, left, c, c), TF.crop(high, top, left, c, c) def get_loaders(cfg): train_set = LOLDataset(cfg['lol_root'], 'train', cfg['crop_size'], augment=True) val_set = LOLDataset(cfg['lol_root'], 'val', None, augment=False) train_loader = DataLoader(train_set, batch_size=cfg['batch_size'], shuffle=True, num_workers=cfg['num_workers'], drop_last=True, pin_memory=False) val_loader = DataLoader(val_set, batch_size=1, shuffle=False, num_workers=cfg['num_workers'], drop_last=False, pin_memory=False) return train_loader, val_loader class ResBlock(nn.Module): def __init__(self, ch): super().__init__() self.block = nn.Sequential( nn.Conv2d(ch, ch, 3, padding=1, bias=False), nn.BatchNorm2d(ch), nn.ReLU(True), nn.Conv2d(ch, ch, 3, padding=1, bias=False), nn.BatchNorm2d(ch), ) self.relu = nn.ReLU(True) def forward(self, x): return self.relu(x + self.block(x)) class EncoderBlock(nn.Module): def __init__(self, in_ch, out_ch): super().__init__() self.conv = nn.Sequential( nn.Conv2d(in_ch, out_ch, 3, padding=1, bias=False), nn.BatchNorm2d(out_ch), nn.ReLU(True), ) self.res = ResBlock(out_ch) self.down = nn.MaxPool2d(2) def forward(self, x): x = self.res(self.conv(x)) skip = x return self.down(x), skip class DecoderBlock(nn.Module): def __init__(self, in_ch, skip_ch, out_ch, feedback_ch=32): super().__init__() self.merge = nn.Sequential( nn.Conv2d(in_ch + skip_ch, out_ch, 1, bias=False), nn.BatchNorm2d(out_ch), nn.ReLU(True), ) self.res = ResBlock(out_ch) if feedback_ch is not None: self.feedback_proj = nn.Conv2d(feedback_ch, out_ch, 1, bias=False) else: self.feedback_proj = None def forward(self, x, skip, feedback=None): x = F.interpolate(x, size=skip.shape[-2:], mode='bilinear', align_corners=False) x = self.merge(torch.cat([x, skip], dim=1)) if feedback is not None and self.feedback_proj is not None: if feedback.shape[-2:] != x.shape[-2:]: feedback = F.interpolate(feedback, size=x.shape[-2:], mode='bilinear', align_corners=False) x = x + self.feedback_proj(feedback) return self.res(x) class LLEN(nn.Module): def __init__(self, base_ch=32): super().__init__() c = base_ch self.enc1 = EncoderBlock(3, c) self.enc2 = EncoderBlock(c, c * 2) self.enc3 = EncoderBlock(c * 2, c * 4) self.enc4 = EncoderBlock(c * 4, c * 8) self.bottleneck = nn.Sequential(ResBlock(c * 8), ResBlock(c * 8)) self.dec4 = DecoderBlock(c * 8, c * 8, c * 4, feedback_ch=c * 4) self.dec3 = DecoderBlock(c * 4, c * 4, c * 2, feedback_ch=c * 2) self.dec2 = DecoderBlock(c * 2, c * 2, c, feedback_ch=c) self.dec1 = DecoderBlock(c, c, c, feedback_ch=None) self.head = nn.Sequential( nn.Conv2d(c, c, 3, padding=1, bias=False), nn.ReLU(True), nn.Conv2d(c, 3, 1), nn.Sigmoid(), ) def forward(self, x, dgff_p5=None, dgff_p4=None, dgff_p3=None): x, s1 = self.enc1(x) x, s2 = self.enc2(x) x, s3 = self.enc3(x) x, s4 = self.enc4(x) x = self.bottleneck(x) x = self.dec4(x, s4, feedback=dgff_p5) x = self.dec3(x, s3, feedback=dgff_p4) x = self.dec2(x, s2, feedback=dgff_p3) x = self.dec1(x, s1) return self.head(x) def count_parameters(self): return sum(p.numel() for p in self.parameters() if p.requires_grad) class GatedAdapter(nn.Module): def __init__(self, in_ch: int, out_ch: int): super().__init__() self.proj = nn.Sequential( nn.Conv2d(in_ch, out_ch, kernel_size=1, bias=False), nn.BatchNorm2d(out_ch), nn.ReLU(inplace=True), ) self.gate = nn.Sequential( nn.Conv2d(in_ch, out_ch, kernel_size=1, bias=False), nn.Sigmoid(), ) def forward(self, f_det, target_size): f_det = nn.functional.interpolate( f_det, size=target_size, mode='bilinear', align_corners=False ) return self.gate(f_det) * self.proj(f_det) class DGFFModule(nn.Module): def __init__(self, yolo_weights='yolov8n.pt', llen_base_ch=32, freeze_backbone=True): super().__init__() yolo = YOLO(yolo_weights) self.backbone = yolo.model.model[:10] if freeze_backbone: for p in self.backbone.parameters(): p.requires_grad_(False) with torch.no_grad(): probe = torch.zeros(1, 3, 64, 64) out = probe for i, layer in enumerate(self.backbone): out = layer(out) if i == 4: actual_p3 = out.shape[1] if i == 6: actual_p4 = out.shape[1] if i == 9: actual_p5 = out.shape[1] self.p3_ch = actual_p3 self.p4_ch = actual_p4 self.p5_ch = actual_p5 c = llen_base_ch self._llen_base_ch = c self.adapter_p3 = GatedAdapter(self.p3_ch, c) self.adapter_p5 = GatedAdapter(self.p5_ch, c * 4) self.adapter_p4 = GatedAdapter(self.p4_ch, c * 2) def forward(self, enhanced_img): x = enhanced_img p4_raw = p5_raw = p3_raw = None for i, layer in enumerate(self.backbone): x = layer(x) if i == 4: p3_raw = x if i == 6: p4_raw = x if i == 9: p5_raw = x c = self._llen_base_ch H_input, W_input = enhanced_img.shape[-2:] p5_size = (H_input // 8, W_input // 8) p4_size = (H_input // 4, W_input // 4) p3_size = (H_input // 2, W_input // 2) p5_feedback = self.adapter_p5(p5_raw, p5_size) p4_feedback = self.adapter_p4(p4_raw, p4_size) p3_feedback = self.adapter_p3(p3_raw, p3_size) return p4_feedback, p5_feedback, p3_feedback def count_parameters(self): total = sum(p.numel() for p in self.parameters()) trainable = sum(p.numel() for p in self.parameters() if p.requires_grad) return total, trainable class PerceptualLoss(nn.Module): def __init__(self, device): super().__init__() vgg = models.vgg16(weights=models.VGG16_Weights.DEFAULT).features self.slice1 = nn.Sequential(*list(vgg)[:10]).to(device).eval() self.slice2 = nn.Sequential(*list(vgg)[:17]).to(device).eval() for p in self.parameters(): p.requires_grad_(False) mean = torch.tensor([0.485, 0.456, 0.406], device=device).view(1, 3, 1, 1) std = torch.tensor([0.229, 0.224, 0.225], device=device).view(1, 3, 1, 1) self.register_buffer('mean', mean) self.register_buffer('std', std) def forward(self, pred, target): pred = (pred - self.mean) / self.std target = (target - self.mean) / self.std return F.mse_loss(self.slice1(pred), self.slice1(target)) + \ F.mse_loss(self.slice2(pred), self.slice2(target)) class DetectionFeatureLoss(nn.Module): def __init__(self, dgff_module): super().__init__() self.backbone = dgff_module.backbone def _extract(self, x, no_grad=False): feats = [] ctx = torch.no_grad() if no_grad else torch.enable_grad() with ctx: for i, layer in enumerate(self.backbone): x = layer(x) if i in (4, 6, 9): feats.append(x) return feats def forward(self, enhanced_guided, high): feats_enh = self._extract(enhanced_guided, no_grad=False) feats_clean = self._extract(high, no_grad=True) return sum(F.mse_loss(fe, fc.detach()) for fe, fc in zip(feats_enh, feats_clean)) def compute_psnr(pred, target): mse = F.mse_loss(pred, target).item() return float('inf') if mse == 0 else 10 * torch.log10(torch.tensor(1.0 / mse)).item() def compute_ssim(pred, target): try: from pytorch_msssim import ssim return ssim(pred, target, data_range=1.0).item() except ImportError: mu1, mu2 = pred.mean(), target.mean() s1, s2 = pred.std(), target.std() s12 = ((pred - mu1) * (target - mu2)).mean() C1, C2 = 0.01 ** 2, 0.03 ** 2 return (((2 * mu1 * mu2 + C1) * (2 * s12 + C2)) / ((mu1 ** 2 + mu2 ** 2 + C1) * (s1 ** 2 + s2 ** 2 + C2))).item() def get_lambda_det(epoch, total_epochs, max_lambda): warmup = total_epochs // 2 return max_lambda * min(1.0, epoch / warmup) @torch.no_grad() def validate(llen, dgff, val_loader, device): llen.eval() dgff.eval() total_psnr = total_ssim = 0.0 for low, high, _ in val_loader: low, high = low.to(device), high.to(device) enhanced = llen(low) p4, p5, p3 = dgff(enhanced) enhanced_guided = llen(low, dgff_p5=p5, dgff_p4=p4, dgff_p3=p3) total_psnr += compute_psnr(enhanced_guided, high) total_ssim += compute_ssim(enhanced_guided, high) n = len(val_loader) return total_psnr / n, total_ssim / n def load_results(): if os.path.exists(RESULTS_PATH): with open(RESULTS_PATH) as f: return json.load(f) return [] def save_result(result): results = load_results() results.append(result) with open(RESULTS_PATH, 'w') as f: json.dump(results, f, indent=2) def run_config(lambda_max, epochs, seed, train_loader, val_loader): torch.manual_seed(seed) llen = LLEN(base_ch=CFG['base_ch']).to(DEVICE) dgff = DGFFModule(yolo_weights=CFG['yolo_weights'], llen_base_ch=CFG['base_ch'], freeze_backbone=True).to(DEVICE) perc_loss = PerceptualLoss(DEVICE) det_loss = DetectionFeatureLoss(dgff).to(DEVICE) l1_loss = nn.L1Loss() trainable = (list(llen.parameters()) + list(dgff.adapter_p3.parameters()) + list(dgff.adapter_p4.parameters()) + list(dgff.adapter_p5.parameters())) optimiser = Adam(trainable, lr=CFG['lr'], betas=(0.9, 0.999)) scheduler = CosineAnnealingLR(optimiser, T_max=epochs, eta_min=1e-6) best_psnr, best_ssim = 0.0, 0.0 t_start = time.time() for epoch in range(1, epochs + 1): llen.train() dgff.train() lam_det = get_lambda_det(epoch, epochs, lambda_max) ep_loss = 0.0 for low, high, _ in train_loader: low, high = low.to(DEVICE), high.to(DEVICE) enhanced = llen(low) p4, p5, p3 = dgff(enhanced) enhanced_guided = llen(low, dgff_p5=p5, dgff_p4=p4, dgff_p3=p3) loss_l1 = l1_loss(enhanced_guided, high) loss_perc = perc_loss(enhanced_guided, high) loss_det = det_loss(enhanced_guided, high) if lam_det > 0 else torch.tensor(0.0, device=DEVICE) loss = loss_l1 + CFG['lambda_perc'] * loss_perc + lam_det * loss_det optimiser.zero_grad() loss.backward() torch.nn.utils.clip_grad_norm_(trainable, max_norm=1.0) optimiser.step() ep_loss += loss.item() scheduler.step() if epoch % 5 == 0 or epoch == epochs: psnr, ssim = validate(llen, dgff, val_loader, DEVICE) best_psnr, best_ssim = max(best_psnr, psnr), max(best_ssim, ssim) elapsed = time.time() - t_start print(f" [lambda_det={lambda_max}] epoch {epoch:03d}/{epochs} " f"loss={ep_loss/len(train_loader):.4f} PSNR={psnr:.2f}dB " f"SSIM={ssim:.4f} ({elapsed/60:.1f} min elapsed)", flush=True) torch.save({'llen_state': llen.state_dict(), 'dgff_state': dgff.state_dict(), 'best_psnr': best_psnr, 'best_ssim': best_ssim}, Path(CFG['ckpt_dir']) / f'lambda_{lambda_max}_full200.pt') return {'lambda_det': lambda_max, 'epochs': epochs, 'seed': seed, 'best_psnr': best_psnr, 'best_ssim': best_ssim, 'minutes': (time.time() - t_start) / 60} def main(): print(f"Loading LOL dataset from {CFG['lol_root']}...") train_loader, val_loader = get_loaders(CFG) print(f"Train: {len(train_loader.dataset)} pairs | Val: {len(val_loader.dataset)} pairs") done_lambdas = {r['lambda_det'] for r in load_results()} print(f"\nAlready completed: {sorted(done_lambdas) or 'none'}") for lam in ABLATION_LAMBDAS: if lam in done_lambdas: print(f"\nSkipping lambda_det={lam} (already in {RESULTS_PATH})") continue print(f"\n=== lambda_det = {lam} (full {ABLATION_EPOCHS}-epoch schedule) ===", flush=True) result = run_config(lam, ABLATION_EPOCHS, ABLATION_SEED, train_loader, val_loader) save_result(result) print(f" Done: PSNR={result['best_psnr']:.2f}dB SSIM={result['best_ssim']:.4f} " f"({result['minutes']:.1f} min) -- saved to {RESULTS_PATH}") print("\n=== Result ===") results = load_results() for r in results: print(f"lambda_det={r['lambda_det']}: PSNR={r['best_psnr']:.2f}dB SSIM={r['best_ssim']:.4f} " f"(paper's main result was 19.83dB / 0.9048 at lambda_det=0.5)") if __name__ == '__main__': main()