You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

灰度图扩散模型感知损失函数通道不匹配报错修复咨询

灰度图扩散模型感知损失报错修复方案

核心问题

报错的根本原因是你的输入为单通道灰度图,但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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.08 04:34:52