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

神经网络2D FFT卷积实现遇维度不匹配错误,求解决方案

自定义2D FFT卷积层维度不匹配问题调试与解决

问题概述

输入张量形状为(100000, 1, 32, 32),遵循(instances, channel, height, width)维度约定,使用自定义FFTConv2D层时触发维度不匹配的ValueError,报错核心信息:

ValueError: Dimensions must be equal, but are 2 and 3 for '{{node fft_conv2d_2/mul}} = Mul[T=DT_COMPLEX64](fft_conv2d_2/FFT2D, fft_conv2d_2/FFT2D_1)' with input shapes: [3,2,16,32], [3,3,32,64].

错误原因分析

  1. FFT维度与核尺寸不匹配:tf.signal.fft2d默认对张量最后两个维度执行傅里叶变换,但输入张量是(batch, channel, h, w),卷积核是(kernel_h, kernel_w, in_channel, out_channel),两者高宽、通道维度完全无法对应,直接元素相乘必然冲突。
  2. 卷积核未填充至输入尺寸:FFT卷积要求核的高宽必须与输入特征图高宽一致,当前核为3x3,而池化后的输入特征图是16x16,尺寸不匹配。
  3. 通道维度广播逻辑错误:输入通道在第二个维度,核的输入通道在第三个维度、输出通道在第四个维度,未做维度扩展支持广播相乘。

解决步骤

  • 调整卷积核维度顺序,将其填充至当前输入特征图的高宽尺寸。
  • 对输入和核扩展维度,确保批量、通道、输出滤波器维度可以正确广播。
  • 频域相乘后对输入通道求和,将结果调整回原维度约定格式。
  • 补充strides的下采样逻辑,对齐标准卷积行为。

修正后的代码

import tensorflow as tf
from tensorflow import keras
from tensorflow.keras import layers, Sequential

class FFTConv2D(keras.layers.Layer):
    def __init__(self, filters, kernel_size, padding="same", strides=1, activation="relu",
                 kernel_initializer='TruncatedNormal', **kwargs):
        super().__init__(**kwargs)
        self.filters = filters
        self.kernel_size = kernel_size
        self.padding = padding
        self.strides = strides
        self.activation = keras.activations.get(activation)
        self.kernel_initializer = keras.initializers.get(kernel_initializer)

    def build(self, input_shape):
        # 输入形状为(batch, in_ch, h, w),input_shape[-3]对应输入通道数
        self.kernel = self.add_weight(
            shape=(self.kernel_size, self.kernel_size, input_shape[-3], self.filters),
            initializer=self.kernel_initializer,
            trainable=True,
            name="kernel"
        )

    def call(self, inputs):
        # 获取输入形状参数
        batch_size, in_ch, h, w = tf.shape(inputs)[0], tf.shape(inputs)[1], tf.shape(inputs)[2], tf.shape(inputs)[3]
        
        # 1. 调整核维度:从(kh, kw, in_ch, out_ch)转为(in_ch, kh, kw, out_ch)
        kernel = tf.transpose(self.kernel, perm=[2, 0, 1, 3])
        
        # 2. 将核填充至输入特征图尺寸(h, w)
        pad_h = h - self.kernel_size
        pad_w = w - self.kernel_size
        kernel_padded = tf.pad(kernel, [[0,0], [0,pad_h], [0,pad_w], [0,0]])
        
        # 3. 对输入和填充后的核执行FFT
        x_fft = tf.signal.fft2d(tf.cast(inputs, tf.complex64))
        kernel_fft = tf.signal.fft2d(tf.cast(kernel_padded, tf.complex64))
        
        # 4. 扩展维度支持广播:输入新增输出滤波器维度,核新增批量维度
        x_fft_expanded = tf.expand_dims(x_fft, axis=-1)
        kernel_fft_expanded = tf.expand_dims(kernel_fft, axis=0)
        
        # 5. 频域相乘后对输入通道求和
        x_kernel_fft = tf.multiply(x_fft_expanded, kernel_fft_expanded)
        x_kernel_fft_sum = tf.reduce_sum(x_kernel_fft, axis=1)
        
        # 6. 逆FFT转回实数域,裁剪避免数值溢出
        x_kernel = tf.math.real(tf.signal.ifft2d(x_kernel_fft_sum))
        x_kernel = tf.clip_by_value(x_kernel, -1e6, 1e6)
        
        # 7. 调整维度回到(batch, out_ch, h, w)
        x_kernel = tf.transpose(x_kernel, perm=[0, 3, 1, 2])
        
        # 处理strides下采样
        if self.strides > 1:
            x_kernel = tf.image.resize(x_kernel, size=(h//self.strides, w//self.strides), method='bilinear')
        
        return self.activation(x_kernel)

    def get_config(self):
        config = super().get_config()
        config.update({
            "filters": self.filters,
            "kernel_size": self.kernel_size,
            "padding": self.padding,
            "strides": self.strides,
            "activation": keras.activations.serialize(self.activation),
            "kernel_initializer": keras.initializers.serialize(self.kernel_initializer)
        })
        return config

# 测试模型
model = Sequential()
model.add(FFTConv2D(32, 3, padding="same", strides=1, activation="relu", kernel_initializer='TruncatedNormal',
                 input_shape=(1,32,32)))
model.add(FFTConv2D(32, 3, padding="same", strides=1, activation="relu", kernel_initializer='TruncatedNormal'))
model.add(layers.MaxPooling2D(pool_size=(2, 2), padding="same"))
model.add(FFTConv2D(64, 3, padding="same", strides=1, activation="relu", kernel_initializer='TruncatedNormal'))
model.add(FFTConv2D(64, 3, padding="same", strides=1, activation="relu", kernel_initializer='TruncatedNormal'))
model.add(layers.Flatten())
model.add(layers.Dense(256, activation='relu', kernel_initializer='TruncatedNormal'))
model.add(layers.Dropout(0.2))
model.add(layers.Dense(128, activation='relu', kernel_initializer='TruncatedNormal'))
model.add(layers.Dropout(0.2))
model.add(layers.Dense(1, activation='sigmoid', kernel_initializer='TruncatedNormal'))

# 验证输出形状
x_test = tf.random.normal((3, 1, 32, 32))
output = model(x_test)
print(output.shape)  # 预期输出: (3, 1)

额外说明

  • 原代码未实现padding的完整逻辑,若需要严格对齐标准卷积的same/valid模式,需在FFT前对输入做对应填充处理。
  • 频域计算易出现数值不稳定,添加tf.clip_by_value可缓解浮点溢出问题。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.25 13:49:57