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

PyTorch conv2d参数无效错误排查:Unet模型训练异常求助

问题分析与解决

报错核心原因

报错提示conv2d() received an invalid combination of arguments - got (tuple, ...),本质是模型forward传递过程中错误地将元组作为卷积层输入,同时Unet核心的跳连接逻辑完全缺失。

具体代码问题点

  • 编码器输出处理错误:encoder_block的forward返回(x, p)(卷积特征图+池化后特征图),但你在unet的forward中直接把元组结果传给下一层编码器,导致卷积层收到的是元组而非Tensor,触发类型错误。
  • 解码器调用参数缺失:decoder_block的forward需要两个参数(当前输入+编码器跳连接特征),但你调用时只传了一个参数,且内部拼接逻辑错误——应该用上采样后的特征和跳连接特征拼接,而非原始输入。
  • 跳连接逻辑未实现:Unet的核心是编码器特征与解码器上采样特征的拼接,你的代码完全没保存编码器的特征用于跳连接。

修正后的完整代码

import torch
import torch.nn as nn
import torch.nn.functional as F
import pytorch_lightning as L

class double_conv(nn.Sequential):
    def __init__(self, input_dim, output_dim):
        super(double_conv, self).__init__(
            nn.Conv2d(input_dim, output_dim, kernel_size=3, padding=1, stride=1),
            nn.BatchNorm2d(output_dim),
            nn.ReLU(inplace=True),
            nn.Conv2d(output_dim, output_dim, kernel_size=3, padding=1, stride=1),
            nn.BatchNorm2d(output_dim),
            nn.ReLU(inplace=True)
        )

class encoder_block(nn.Module):
    def __init__(self, in_c, out_c):
        super().__init__()
        self.conv = double_conv(in_c, out_c)
        self.pool = nn.MaxPool2d((2, 2))

    def forward(self, inputs):
        x = self.conv(inputs)
        p = self.pool(x)
        return x, p  # 返回卷积特征(用于跳连接)和池化后的特征(下一层输入)

class decoder_block(nn.Module):
    def __init__(self, input_dim, output_dim):
        super().__init__()
        self.up = nn.ConvTranspose2d(input_dim, output_dim, kernel_size=3, stride=2, padding=1, output_padding=1)
        self.conv = double_conv(output_dim * 2, output_dim)

    def forward(self, inputs, skip_features):
        x = self.up(inputs)
        # 拼接上采样后的特征和编码器传来的跳连接特征
        x = torch.cat([x, skip_features], dim=1)
        x = self.conv(x)
        return x

class unet(L.LightningModule):
    def __init__(self):
        super().__init__()
        """ Encoder """
        self.e1 = encoder_block(1, 64)
        self.e2 = encoder_block(64, 128)
        self.e3 = encoder_block(128, 256)
        self.e4 = encoder_block(256, 512)

        """ Bottleneck """
        self.b = double_conv(512, 1024)

        """ Decoder """
        self.d1 = decoder_block(1024, 512)
        self.d2 = decoder_block(512, 256)
        self.d3 = decoder_block(256, 128)
        self.d4 = decoder_block(128, 64)

        """ Output """
        self.output = nn.Conv2d(64, 1, kernel_size=1, padding=0)

    def forward(self, x):
        # 编码器:拆分输出,保存跳连接特征
        skip1, p1 = self.e1(x)
        skip2, p2 = self.e2(p1)
        skip3, p3 = self.e3(p2)
        skip4, p4 = self.e4(p3)

        # 瓶颈层
        bottleneck = self.b(p4)

        # 解码器:传入当前输入+对应编码器的跳连接特征
        d1 = self.d1(bottleneck, skip4)
        d2 = self.d2(d1, skip3)
        d3 = self.d3(d2, skip2)
        d4 = self.d4(d3, skip1)

        # 输出层
        output = self.output(d4)
        return output

    def training_step(self, batch, batch_number):
        x, y = batch
        y_pred = self.forward(x)
        loss = F.binary_cross_entropy_with_logits(y_pred, y)
        self.log('train_loss', loss)
        return loss

    def configure_optimizers(self):
        optimizer = torch.optim.Adam(self.parameters(), lr=1e-3)
        return optimizer

unet_deneme = unet()

关键修正说明

  • 拆分编码器返回的元组,分别保存跳连接特征和下一层输入,避免将元组传入卷积层。
  • 调用解码器时补全跳连接特征参数,修正拼接逻辑为上采样特征+跳连接特征。
  • 替换手动日志为Lightning内置的self.log,符合框架规范。

内容的提问来源于stack exchange,提问作者gill bates

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.12 15:12:06