Keras自定义梅尔频谱图层构建模型时报'NoneType'不可下标错误
问题原因
- 自定义
MelLayer没有实现compute_output_shape方法,Keras静态构建计算图时无法正确推断层的输出形状,导致后续Conv2D层读取输入通道维度时拿到空值触发报错。 call方法中全局tf.squeeze(audio)逻辑存在隐患:输入形状为(batch_size, 32000, 1),全局squeeze会在batch_size为1的场景下错误压缩批次维度,进一步打乱形状推导逻辑。
修复方案
- 将
call方法中的全局squeeze修改为仅压缩最后一个声道维度,保留批次维度 - 重写
compute_output_shape方法,显式声明层的输出形状,让Keras可以正确做静态形状推断 - (可选优化)将mel滤波器矩阵的初始化移到
build方法中,符合Keras自定义层的规范,避免静态形状推断冲突
修复后的完整MelLayer代码如下:
class MelLayer(tf.keras.layers.Layer): def __init__( self, frame_length=1024, frame_step=256, fft_length=None, sampling_rate=16000, num_mel_channels=80, freq_min=1, freq_max=7600, as_3D_tensor=True, **kwargs, ): super().__init__(**kwargs) self.frame_length = frame_length self.frame_step = frame_step self.fft_length = fft_length if fft_length else frame_length self.sampling_rate = sampling_rate self.num_mel_channels = num_mel_channels self.freq_min = freq_min self.freq_max = freq_max self.as_3D_tensor = as_3D_tensor def build(self, input_shape): # 将mel滤波器初始化移到build方法,避免静态形状冲突 self.mel_filterbank = tf.signal.linear_to_mel_weight_matrix( num_mel_bins=self.num_mel_channels, num_spectrogram_bins=self.fft_length // 2 + 1, sample_rate=self.sampling_rate, lower_edge_hertz=self.freq_min, upper_edge_hertz=self.freq_max, ) self.non_trainable_weights.append(self.mel_filterbank) super().build(input_shape) def call(self, audio, training=True): # 仅压缩最后一维声道,保留批次维度 stft = tf.signal.stft( tf.squeeze(audio, axis=-1), self.frame_length, self.frame_step, self.fft_length, pad_end=True, ) magnitude = tf.abs(stft) mel = tf.matmul(tf.square(magnitude), self.mel_filterbank) log_mel_spec = tfio.audio.dbscale(mel, top_db=80) return tf.expand_dims(log_mel_spec,axis=-1) if self.as_3D_tensor else tf.squeeze(log_mel_spec, axis=-1) def compute_output_shape(self, input_shape): # 显式声明输出形状 batch_size = input_shape[0] if input_shape[1] is None: time_steps = None else: # pad_end=True时时间步向上取整 time_steps = (input_shape[1] + self.frame_step - 1) // self.frame_step if self.as_3D_tensor: return (batch_size, time_steps, self.num_mel_channels, 1) else: return (batch_size, time_steps, self.num_mel_channels) def get_config(self): config = super().get_config() config.update( { "frame_length": self.frame_length, "frame_step": self.frame_step, "fft_length": self.fft_length, "sampling_rate": self.sampling_rate, "num_mel_channels": self.num_mel_channels, "freq_min": self.freq_min, "freq_max": self.freq_max, "as_3D_tensor": self.as_3D_tensor } ) return config
修复后重新构建模型,mel_l的输出形状会被正确推断为(None, 63, 80, 1),后续Conv2D层可以正常读取通道维度,不会再触发报错。
内容的提问来源于stack exchange,提问作者ABaron
相关产品推荐
相关产品推荐

