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

如何实现适配编码降采样的4倍可学习插值上采样层?

4倍可学习上采样层的正确实现方案

问题核心

你的编码器通过kernel_size=8、stride=4、padding="valid"的卷积层,将输入形状(4n+4, c)降采样为(n, c)。解码器需要将(n, c)上采样回(4n+4, c),但当前实现的输出形状为(4n-4, c),形状不匹配导致模型欠拟合。以下是针对性的修正方案:

实现思路

  1. 对齐降采样逆逻辑:针对编码器的valid卷积降采样,上采样时先补全首尾信息,再进行可学习插值,保证输出长度严格匹配4n+4。
  2. 合理约束插值权重:延续Wave-U-Net的可学习插值思路,为4倍上采样的3个中间插值点设置区间化的权重约束,符合线性插值的位置逻辑。
  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

关键修正点说明

  1. 首尾padding:通过在输入前补1个首元素、后补2个尾元素,确保后续插值后能覆盖编码器降采样时丢失的首尾信息。
  2. 权重区间约束:将三个插值权重分别约束在[0,1/3]、[1/3,2/3]、[2/3,1],符合4倍上采样时的线性插值位置逻辑,避免权重无意义波动。
  3. 索引与长度控制:通过交错拼接原始点和插值点,最后截取到目标长度4n+4,严格匹配编码器的输入形状。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.13 01:42:02