神经网络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].
错误原因分析
- FFT维度与核尺寸不匹配:
tf.signal.fft2d默认对张量最后两个维度执行傅里叶变换,但输入张量是(batch, channel, h, w),卷积核是(kernel_h, kernel_w, in_channel, out_channel),两者高宽、通道维度完全无法对应,直接元素相乘必然冲突。 - 卷积核未填充至输入尺寸:FFT卷积要求核的高宽必须与输入特征图高宽一致,当前核为3x3,而池化后的输入特征图是16x16,尺寸不匹配。
- 通道维度广播逻辑错误:输入通道在第二个维度,核的输入通道在第三个维度、输出通道在第四个维度,未做维度扩展支持广播相乘。
解决步骤
- 调整卷积核维度顺序,将其填充至当前输入特征图的高宽尺寸。
- 对输入和核扩展维度,确保批量、通道、输出滤波器维度可以正确广播。
- 频域相乘后对输入通道求和,将结果调整回原维度约定格式。
- 补充
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
相关产品推荐
相关产品推荐

