TensorFlow谱带分割层InvalidArgumentError维度不匹配问题求助
频谱图谱带分割TensorFlow层报错修复
错误原因分析
报错核心是tf.split的分割维度总和不匹配,同时存在多处逻辑错误:
- 输入的频率维度(axis=3)大小是257,但传入
tf.split的self.idxs.numpy()总和是2524,远大于257,完全不符合分割要求 - 混淆了分割点和子带段长度:
self.idxs是计算出的频率分割点,而tf.split需要的是每个子带的长度,也就是后续计算的aux_idxs - call函数中对分割后张量的索引错误(B[i]只有4维,却用了5维索引
[:, :, :, :, 0]) - LayerNormalization的axis参数设置错误,传入了动态计算的
2*i,不符合LayerNormalization的axis要求(必须是固定轴索引)
修复步骤
- 替换分割参数:用
aux_idxs(子带段长度)代替self.idxs作为tf.split的分割依据 - 修正张量索引:分割后的子带张量是4维
(B, T, C, 段长度),取实部虚部应该用B[i][:, :, 0, :]和B[i][:, :, 1, :] - 调整LayerNormalization轴:拼接后的张量形状是
(B, T, 2*段长度),设置axis=-1(最后一维)做层归一化 - 修正张量堆叠逻辑:确保全连接层输出后调整维度,最终堆叠得到目标形状
(B, T, 子带数, subband_dim)
完整修复代码
import tensorflow as tf def select_norm(norm, dim, shape): if norm in ['gln', 'cln', 'ln']: return tf.keras.layers.LayerNormalization(axis=dim, center=True, scale=True) else: return tf.keras.layers.BatchNormalization(axis=dim) class Band_Split(tf.keras.layers.Layer): def __init__(self, temporal_dimension, max_freq_idx, sample_rate, n_fft, subband_dim): super(Band_Split, self).__init__() step_idx_100 = 100 * n_fft / sample_rate step_idx_250 = 250 * n_fft / sample_rate step_idx_500 = 500 * n_fft / sample_rate idx_1k = 1000 * n_fft / sample_rate idx_4k = 4000 * n_fft / sample_rate idx_8k = 8000 * n_fft / sample_rate self.subband_dim = subband_dim # 计算频率分割点 idxs_1 = tf.range(step_idx_100, idx_1k, step_idx_100) idxs_2 = tf.range(idx_1k, idx_4k, step_idx_250) idxs_3 = tf.range(idx_4k, idx_8k, step_idx_500) self.idxs = tf.concat([idxs_1, idxs_2, idxs_3], axis=0) self.idxs = tf.cast(tf.floor(self.idxs), tf.int32) # 计算每个子带的频率段长度 aux_idxs = tf.concat([[0], self.idxs], axis=0) aux_idxs = tf.concat([aux_idxs, [max_freq_idx]], axis=0) self.subband_lengths = aux_idxs[1:] - aux_idxs[:-1] # 初始化层归一化和全连接层 self.layer_norms = [] self.linear_layers = [] for _ in self.subband_lengths: # 层归一化作用在特征维度(最后一维) self.layer_norms.append(tf.keras.layers.LayerNormalization(axis=-1)) self.linear_layers.append(tf.keras.layers.Dense(self.subband_dim)) def call(self, inputs): # 使用子带长度进行分割,axis=3是频率维度 subbands = tf.split(inputs, self.subband_lengths.numpy(), axis=3) Z = [] for i, (layer_norm, linear_layer) in enumerate(zip(self.layer_norms, self.linear_layers)): # 取实部和虚部,拼接成特征维度 real_part = subbands[i][:, :, 0, :] # shape (B, T, 段长度) imag_part = subbands[i][:, :, 1, :] # shape (B, T, 段长度) b_i = tf.concat([real_part, imag_part], axis=-1) # shape (B, T, 2*段长度) # 层归一化+全连接,调整维度方便后续堆叠 normed = layer_norm(b_i) dense_out = linear_layer(normed) # shape (B, T, subband_dim) Z.append(tf.expand_dims(dense_out, axis=2)) # 堆叠所有子带,得到 (B, T, 子带数, subband_dim) Z = tf.concat(Z, axis=2) return Z # 测试运行 B = 4 # batch T = 9 # temporal C = 2 # 实部/虚部 F = 257 # frequency temporal_dimension = T max_freq_idx = F sample_rate = 16000 n_fft = 512 subband_dim = 128 X = tf.random.normal(shape=(B, T, C, F)) band_split = Band_Split(temporal_dimension, max_freq_idx, sample_rate, n_fft, subband_dim) result = band_split(X) print("输入形状:", X.shape) print("输出形状:", result.shape) # 预期 (4, 9, 30, 128)
验证说明
修复后代码运行会输出:
输入形状: (4, 9, 2, 257) 输出形状: (4, 9, 30, 128)
完全符合预期的目标形状。
内容的提问来源于stack exchange,提问作者nic.o
相关产品推荐
相关产品推荐

