:彩色RGB图像去噪实战,从灰度模型升级到真实图片处理)
Pytorch图像去噪实战十五彩色RGB图像去噪实战从灰度模型升级到真实图片处理一、问题场景灰度图跑通了但真实项目都是彩色图前面很多文章为了方便讲清楚模型结构都默认使用灰度图1通道输入但真实项目里大多数图片都是RGB三通道3通道输入比如手机照片商品图片人像图片旅游风景图电商素材图我一开始以为把模型输入通道从 1 改成 3 就完事了结果发现实际效果并不稳定颜色偏移局部发灰饱和度下降去噪后肤色不自然RGB通道噪声不一致这篇文章就完整解决如何把图像去噪模型从灰度图升级到彩色RGB图。二、为什么RGB去噪更难灰度图只有亮度信息。RGB图像包含亮度噪声色彩噪声通道间相关性压缩导致的色块白平衡偏移如果直接把灰度去噪模型改成三通道有时会出现噪声少了但颜色也变脏了。所以RGB去噪不仅要恢复结构还要保持颜色一致。三、RGB去噪的两种思路1. 直接RGB空间去噪输入R, G, B输出clean RGB优点简单直接工程实现方便缺点色彩容易漂移2. 转到YCbCr空间去噪只对Y亮度通道去噪保留CbCr色彩通道。优点色彩更稳定缺点对彩色噪声处理不足本文先实现最通用的RGB端到端去噪。四、工程目录结构rgb_denoise/ ├── data/ │ ├── train/ │ └── val/ ├── models/ │ └── rgb_unet.py ├── dataset.py ├── train.py ├── eval.py └── utils.py五、RGB数据集实现dataset.pyimportosimportrandomimporttorchfromPILimportImagefromtorch.utils.dataimportDatasetimporttorchvision.transformsastransformsclassRGBDenoiseDataset(Dataset):def__init__(self,root_dir,patch_size128):self.paths[os.path.join(root_dir,name)fornameinos.listdir(root_dir)ifname.lower().endswith((.jpg,.jpeg,.png))]self.patch_sizepatch_size self.to_tensortransforms.ToTensor()def__len__(self):returnlen(self.paths)def__getitem__(self,idx):imgImage.open(self.paths[idx]).convert(RGB)w,himg.sizeifwself.patch_sizeandhself.patch_size:xrandom.randint(0,w-self.patch_size)yrandom.randint(0,h-self.patch_size)imgimg.crop((x,y,xself.patch_size,yself.patch_size))else:imgimg.resize((self.patch_size,self.patch_size))cleanself.to_tensor(img)sigmarandom.choice([10,15,25,35,50])noisetorch.randn_like(clean)*sigma/255.0noisytorch.clamp(cleannoise,0.0,1.0)returnnoisy,clean六、RGB UNet模型重点变化in_channels3out_channels3models/rgb_unet.pyimporttorchimporttorch.nnasnnclassConvBlock(nn.Module):def__init__(self,in_channels,out_channels):super().__init__()self.blocknn.Sequential(nn.Conv2d(in_channels,out_channels,3,padding1),nn.GroupNorm(8,out_channels),nn.ReLU(inplaceTrue),nn.Conv2d(out_channels,out_channels,3,padding1),nn.GroupNorm(8,out_channels),nn.ReLU(inplaceTrue))defforward(self,x):returnself.block(x)classRGBUNetDenoise(nn.Module):def__init__(self):super().__init__()self.poolnn.MaxPool2d(2)self.enc1ConvBlock(3,64)self.enc2ConvBlock(64,128)self.enc3ConvBlock(128,256)self.bottleneckConvBlock(256,512)self.up3nn.ConvTranspose2d(512,256,2,2)self.dec3ConvBlock(512,256)self.up2nn.ConvTranspose2d(256,128,2,2)self.dec2ConvBlock(256,128)self.up1nn.ConvTranspose2d(128,64,2,2)self.dec1ConvBlock(128,64)self.outnn.Conv2d(64,3,1)defforward(self,x):e1self.enc1(x)e2self.enc2(self.pool(e1))e3self.enc3(self.pool(e2))bself.bottleneck(self.pool(e3))d3self.up3(b)d3torch.cat([d3,e3],dim1)d3self.dec3(d3)d2self.up2(d3)d2torch.cat([d2,e2],dim1)d2self.dec2(d2)d1self.up1(d2)d1torch.cat([d1,e1],dim1)d1self.dec1(d1)returnself.out(d1)七、为什么这里用GroupNorm而不是BatchNormRGB图像训练时显存占用更高batch size 往往比较小。如果 batch size 2 或 4BatchNorm可能不稳定。因此这里使用 GroupNormnn.GroupNorm(8,channels)它不依赖batch维度更适合小batch训练。八、训练代码train.pyimporttorchfromtorch.utils.dataimportDataLoaderfromdatasetimportRGBDenoiseDatasetfrommodels.rgb_unetimportRGBUNetDenoisedeftrain():devicetorch.device(cudaiftorch.cuda.is_available()elsecpu)datasetRGBDenoiseDataset(data/train,patch_size128)loaderDataLoader(dataset,batch_size4,shuffleTrue,num_workers4)modelRGBUNetDenoise().to(device)optimizertorch.optim.AdamW(model.parameters(),lr1e-4,weight_decay1e-4)criteriontorch.nn.L1Loss()forepochinrange(1,81):model.train()total_loss0fornoisy,cleaninloader:noisynoisy.to(device)cleanclean.to(device)predmodel(noisy)losscriterion(pred,clean)optimizer.zero_grad()loss.backward()torch.nn.utils.clip_grad_norm_(model.parameters(),1.0)optimizer.step()total_lossloss.item()print(fEpoch{epoch}, Loss:{total_loss/len(loader):.6f})ifepoch%100:torch.save(model.state_dict(),frgb_unet_epoch_{epoch}.pth)if__name____main__:train()九、推理代码importtorchfromPILimportImageimporttorchvision.transformsastransformsimporttorchvision.utilsasvutilsfrommodels.rgb_unetimportRGBUNetDenoise devicetorch.device(cudaiftorch.cuda.is_available()elsecpu)modelRGBUNetDenoise().to(device)model.load_state_dict(torch.load(rgb_unet_epoch_80.pth,map_locationdevice))model.eval()imgImage.open(test_rgb.png).convert(RGB)transformtransforms.ToTensor()noisytransform(img).unsqueeze(0).to(device)withtorch.no_grad():predmodel(noisy)predtorch.clamp(pred,0.0,1.0)vutils.save_image(pred.cpu(),rgb_denoised.png)十、颜色偏移问题怎么解决RGB去噪最常见的问题就是颜色偏移。方案1降低学习率lr1e-4不要一开始就用 1e-3。方案2使用L1LossL1Loss比MSE更不容易过度平滑。方案3加入颜色一致性损失defcolor_consistency_loss(pred,target):pred_meanpred.mean(dim[2,3])target_meantarget.mean(dim[2,3])returntorch.abs(pred_mean-target_mean).mean()组合损失lossl1_loss0.05*color_loss十一、加入颜色一致性训练训练中可以改成l1torch.nn.L1Loss()defcolor_consistency_loss(pred,target):pred_meanpred.mean(dim[2,3])target_meantarget.mean(dim[2,3])returntorch.abs(pred_mean-target_mean).mean()lossl1(pred,clean)0.05*color_consistency_loss(pred,clean)这样可以减少整体色偏。十二、踩坑记录坑1直接把灰度模型改成3通道效果很差RGB不是简单多两个通道。色彩一致性需要额外关注。坑2BatchNorm导致颜色不稳定小batch下BatchNorm统计不稳定容易导致输出颜色漂移。建议使用GroupNorm坑3训练图像压缩太严重如果训练集本身JPEG伪影很多模型可能把压缩痕迹也学进去。建议使用高质量图像作为clean数据。十三、效果验证RGB模型相比灰度模型主要提升在彩色噪声去除色块噪声修复商品图增强人像照片降噪但也要注意RGB去噪更容易引入颜色问题评估时不能只看PSNR。十四、适合收藏总结RGB去噪流程图片读取为RGB构造三通道噪声模型输入输出均为3通道使用GroupNorm增强稳定性加颜色一致性损失防止色偏避坑清单不要只改输入通道小batch慎用BatchNorm推理时必须clamp注意颜色偏移clean图质量要高十五、优化建议可以继续升级使用YCbCr空间只对Y通道去噪加感知损失使用Restormer RGB版本加真实噪声微调结尾总结从灰度图到RGB图不只是通道数变化而是任务复杂度明显上升。真实项目中RGB去噪必须同时考虑噪声、纹理、结构、颜色一致性。如果你要做真实图片增强或照片修复RGB去噪是必须跨过去的一步。下一篇预告Pytorch图像去噪实战十六YCbCr颜色空间图像去噪解决RGB去噪色偏问题