灰度图扩散模型感知损失函数通道不匹配报错修复咨询
灰度图扩散模型感知损失报错修复方案
核心问题
报错的根本原因是你的输入为单通道灰度图,但VGG模型及配套的MeanShift层都是针对3通道RGB图像设计的,导致输入通道数(1)与模型要求的通道数(3)不匹配,触发维度错误。
修复步骤及代码修改
1. 修正感知损失类(VGGPerceptualLoss)
主要修改点:
- 替换过时的
pretrained参数,消除版本警告 - 添加单通道转3通道的逻辑,匹配VGG输入要求
- 修正batch处理逻辑,不再只处理单个样本
修改后的代码:
from .custom_loss import common import torch import torch.nn as nn import torch.nn.functional as F import torchvision.models as models from torchvision.models import VGG19_Weights class VGGPerceptualLoss(nn.Module): def __init__(self, conv_index, rgb_range=1): super(VGGPerceptualLoss, self).__init__() # 替换过时的pretrained参数,使用官方推荐的weights参数 vgg_features = models.vgg19(weights=VGG19_Weights.IMAGENET1K_V1).features modules = [m for m in vgg_features] if conv_index.find('22') >= 0: self.vgg = nn.Sequential(*modules[:8]) elif conv_index.find('54') >= 0: self.vgg = nn.Sequential(*modules[:35]) vgg_mean = (0.485, 0.456, 0.406) vgg_std = (0.229 * rgb_range, 0.224 * rgb_range, 0.225 * rgb_range) self.sub_mean = common.MeanShift(rgb_range, vgg_mean, vgg_std) for p in self.parameters(): p.requires_grad = False def forward(self, sr, hr): def _forward(x): # 将单通道灰度图复制为3通道,匹配VGG的输入要求 if x.size(1) == 1: x = x.repeat(1, 3, 1, 1) x = self.sub_mean(x) x = self.vgg(x) return x # 处理整个batch,而非仅第一个样本 vgg_sr = _forward(sr) with torch.no_grad(): vgg_hr = _forward(hr.detach()) loss = F.mse_loss(vgg_sr, vgg_hr) return loss
2. 修正扩散模型中的感知损失调用
原代码中用整个大batch与当前microbatch的sample计算损失,会导致维度不匹配,需替换为当前的micro真实图像:
修改forward_backward函数中的对应行:
# 原代码 # perceptual_loss = perceptual_loss_fn(sample, batch) # 修改为 perceptual_loss = perceptual_loss_fn(sample, micro)
关键说明
- 单通道转3通道:通过
repeat(1,3,1,1)将单通道张量在通道维度复制3次,模拟灰度图的RGB表示,符合VGG预训练时的输入格式。 - 批量处理:直接传入整个batch张量,保留批量训练的梯度信息,提升训练效率。
- 维度匹配:确保生成图像(sample)与真实图像(micro)的batch、尺寸完全一致,避免因维度不匹配引发的错误。
内容的提问来源于stack exchange,提问作者yun_91D
相关产品推荐
相关产品推荐

