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

ResNet实现中类内调用带**kwargs的conv_bn函数报参数错误

Fixing "TypeError: forward() got an unexpected keyword argument 'bias'" in ResNet conv_bn Implementation

Let's break down why this error is happening and how to fix it step by step:

What's Causing the Error?

Your conv_bn function works fine when called directly because you're passing nn.Conv1d (the class) as the third argument. When conv(in_channels, out_channels, *args, **kwargs) runs inside conv_bn, you're creating a new instance of the Conv1d class, and parameters like bias=False or stride are valid for initializing a Conv1d layer.

But inside your ResNetBasicBlock class, you're passing conv=self.conv — here's the problem: self.conv is almost certainly an already-instantiated Conv1d module object, not the class itself. When you call this object like conv(...), you're actually invoking its forward() method, which only accepts input tensors as arguments, not initialization parameters like bias or stride. That's why you get the "unexpected keyword argument" error.

Step-by-Step Fixes

1. Ensure self.conv Stores the Convolution Class, Not an Instance

First, check your parent ResNetResidualBlock class. It should accept a convolution class (like nn.Conv1d) as a parameter and store it, rather than storing an initialized module. For example:

import torch.nn as nn
import torch

class ResNetResidualBlock(nn.Module):
    expansion = 1
    def __init__(self, in_channels, out_channels, conv=nn.Conv1d, activation=nn.ReLU, stride=1, *args, **kwargs):
        super().__init__()
        self.in_channels = in_channels
        self.out_channels = out_channels
        # Store the convolution CLASS, not an instance
        self.conv = conv
        self.activation = activation
        self.downsampling = stride
        self.expanded_channels = out_channels * self.expansion

2. Fix the conv_bn Call in ResNetBasicBlock

Since conv_bn expects conv as a positional parameter (third argument), you don't need to use the conv= keyword when calling it. Remove that prefix to pass self.conv correctly:

class ResNetBasicBlock(ResNetResidualBlock):
    def __init__(self, in_channels, out_channels, *args, **kwargs):
        super().__init__(in_channels, out_channels, *args, **kwargs)
        self.blocks = nn.Sequential(
            # Remove `conv=` — self.conv is the third positional argument
            conv_bn(self.in_channels, self.out_channels, self.conv, bias=False, stride=self.downsampling),
            self.activation(),
            conv_bn(self.out_channels, self.expanded_channels, self.conv, bias=False),
        )
        # Add shortcut for downsampling/channel mismatch (critical for ResNet)
        self.shortcut = nn.Identity()
        if self.downsampling != 1 or self.in_channels != self.expanded_channels:
            self.shortcut = conv_bn(self.in_channels, self.expanded_channels, self.conv, stride=self.downsampling, bias=False)
    
    def forward(self, x):
        residual = self.shortcut(x)
        out = self.blocks(x)
        out += residual
        return self.activation()(out)

3. Fix Your Test Tensor Shape (Bonus)

Your dummy tensor is 4D ((1, 32, 224, 224)), but nn.Conv1d expects 3D input ((batch_size, channels, sequence_length)). Adjust it to match:

# 3D tensor for Conv1d: batch=1, channels=32, sequence length=224
dummy = torch.ones((1, 32, 224))
block = ResNetBasicBlock(32, 64, stride=2)
print(block(dummy).shape)  # Output: torch.Size([1, 64, 112])
print(block)

Why This Works

Now when conv_bn runs conv(in_channels, out_channels, *args, **kwargs), it's using the stored convolution class (like nn.Conv1d) to create a new layer instance, and the bias/stride parameters are correctly passed to the layer's initialization (not the forward method).

内容的提问来源于stack exchange,提问作者Carlos Moraes

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.11 08:45:27