如何实现适配编码降采样的4倍可学习插值上采样层?
4倍可学习上采样层的正确实现方案
问题核心
你的编码器通过kernel_size=8、stride=4、padding="valid"的卷积层,将输入形状(4n+4, c)降采样为(n, c)。解码器需要将(n, c)上采样回(4n+4, c),但当前实现的输出形状为(4n-4, c),形状不匹配导致模型欠拟合。以下是针对性的修正方案:
实现思路
- 对齐降采样逆逻辑:针对编码器的
valid卷积降采样,上采样时先补全首尾信息,再进行可学习插值,保证输出长度严格匹配4n+4。 - 合理约束插值权重:延续Wave-U-Net的可学习插值思路,为4倍上采样的3个中间插值点设置区间化的权重约束,符合线性插值的位置逻辑。
- 修正索引拼接逻辑:重新设计插值点与原始点的拼接顺序,确保输出长度精准匹配需求。
修正后的代码实现
import tensorflow as tf class LearnableUpsample4xLayer(tf.keras.layers.Layer): def __init__(self, level=0): super(LearnableUpsample4xLayer, self).__init__() self.level = level def build(self, input_shape): features = input_shape[-1] # 为三个插值点定义可学习权重 self.w1 = self.add_weight( name=f"interp_{self.level}_w1", shape=(features,), initializer="random_normal", trainable=True, dtype=tf.float32 ) self.w2 = self.add_weight( name=f"interp_{self.level}_w2", shape=(features,), initializer="random_normal", trainable=True, dtype=tf.float32 ) self.w3 = self.add_weight( name=f"interp_{self.level}_w3", shape=(features,), initializer="random_normal", trainable=True, dtype=tf.float32 ) def call(self, inputs): # 输入形状:[batch_size, n, c] batch_size, n, c = tf.shape(inputs)[0], tf.shape(inputs)[1], tf.shape(inputs)[2] # 补全首尾:前补1个首元素,后补2个尾元素,得到长度n+3的输入 first_val = inputs[:, 0:1, :] last_val = inputs[:, -1:, :] padded_input = tf.concat([first_val, inputs, last_val, last_val], axis=1) # 形状:[batch_size, n+3, c] # 约束权重到对应插值区间:[0, 1/3], [1/3, 2/3], [2/3, 1] w1_scaled = tf.nn.sigmoid(self.w1) * (1/3) w2_scaled = tf.nn.sigmoid(self.w2) * (1/3) + (1/3) w3_scaled = tf.nn.sigmoid(self.w3) * (1/3) + (2/3) c1_weights = 1.0 - w1_scaled c2_weights = 1.0 - w2_scaled c3_weights = 1.0 - w3_scaled # 构造卷积权重:每个权重是2xc的对角矩阵,对应相邻两个特征的加权组合 conv_w1 = tf.concat([tf.expand_dims(tf.linalg.diag(c1_weights), 0), tf.expand_dims(tf.linalg.diag(w1_scaled), 0)], axis=0) conv_w2 = tf.concat([tf.expand_dims(tf.linalg.diag(c2_weights), 0), tf.expand_dims(tf.linalg.diag(w2_scaled), 0)], axis=0) conv_w3 = tf.concat([tf.expand_dims(tf.linalg.diag(c3_weights), 0), tf.expand_dims(tf.linalg.diag(w3_scaled), 0)], axis=0) # 计算三个插值序列:每个序列长度为 (n+3) -1 =n+2 interp1 = tf.nn.conv1d(padded_input, conv_w1, stride=1, padding="VALID") # [batch_size, n+2, c] interp2 = tf.nn.conv1d(padded_input, conv_w2, stride=1, padding="VALID") interp3 = tf.nn.conv1d(padded_input, conv_w3, stride=1, padding="VALID") # 转置以便按时间轴拼接:[seq_len, batch_size, c] padded_input_t = tf.transpose(padded_input, [1, 0, 2]) interp1_t = tf.transpose(interp1, [1, 0, 2]) interp2_t = tf.transpose(interp2, [1, 0, 2]) interp3_t = tf.transpose(interp3, [1, 0, 2]) # 拼接所有序列:原始序列 + 三个插值序列 concat_all = tf.concat([padded_input_t, interp1_t, interp2_t, interp3_t], axis=0) # 构造索引,按 [原始点, 插值1, 插值2, 插值3] 的顺序交错排列 indices = [] original_len = tf.shape(padded_input)[1] # n+3 interp_len = tf.shape(interp1)[1] # n+2 for i in range(original_len -1): indices.append(i) # 原始点i indices.append(original_len + i) # 插值1的第i个点 indices.append(original_len + interp_len + i) # 插值2的第i个点 indices.append(original_len + 2*interp_len + i) # 插值3的第i个点 indices.append(original_len -1) # 最后一个原始点 # 截取到目标长度4n+4 target_len = 4 * n +4 indices = indices[:target_len] # 按索引取元素并转置回原形状 output_t = tf.gather(concat_all, indices, axis=0) output = tf.transpose(output_t, [1, 0, 2]) # [batch_size, 4n+4, c] return output
关键修正点说明
- 首尾padding:通过在输入前补1个首元素、后补2个尾元素,确保后续插值后能覆盖编码器降采样时丢失的首尾信息。
- 权重区间约束:将三个插值权重分别约束在
[0,1/3]、[1/3,2/3]、[2/3,1],符合4倍上采样时的线性插值位置逻辑,避免权重无意义波动。 - 索引与长度控制:通过交错拼接原始点和插值点,最后截取到目标长度
4n+4,严格匹配编码器的输入形状。
内容的提问来源于stack exchange,提问作者Bisnu Sarkar
相关产品推荐
相关产品推荐

