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

YOTO实现中Loss条件训练异常及矩阵-张量乘法问题排查

问题:YOTO损失条件训练中模型忽略输入参数的异常排查

问题概述

我正在实现YOTO损失条件训练方法,通过带参数的损失函数实现全范围损失训练。训练了一个自动编码器,引入感知损失(VGG)、MSE损失、MAE损失的参数,目标是寻找三者的最优权重。

异常现象

评估训练后的模型时,在图像重建任务下,所有输入参数对应的输出结果完全一致——网络似乎在学习适配三种损失的固定最优解,完全忽略了输入参数的差异。

关键提示:若采用固定损失参数训练(仍使用YOTO和FiLM机制,不进行随机参数采样),模型输出符合预期;仅当训练时随机采样参数时,才会出现上述异常。

技术背景

该方法采用FiLM实现网络对参数的条件控制:通过全连接层输入参数,输出通道级的缩放(scale)和平移(shift)值,再对层激活值执行仿射变换,公式为:
输出 = a * γ + β
其中,a是层激活值,γ、β为通道级的缩放、平移值。

怀疑方向

我怀疑问题出在激活值张量与缩放/平移矩阵的通道级乘法环节。例如,需要将形状为(batchSize, 8, 8, 512)的激活值与形状为(batchSize, 512)的缩放因子相乘,我将矩阵重塑为(batchSize, 1, 1, 512)以实现合法运算,但不确定广播机制是否符合预期,或是存在其他潜在问题。

代码实现

解码器代码

每层依次执行转置卷积、不可训练的批量归一化、基于参数的平移缩放,最后经过ReLU激活:

def FCNet(params, channelCount):
    x = Dense(512)(params)
    x = Activation('relu')(x)
    shift = Dense(channelCount)(x)
    scale = Dense(channelCount)(x)
    return shift, scale

def buildYOTODecoder():
    latent = Input(shape=(8, 8, 512))
    params = Input(shape=(paramCount,))

    depth = [512, 256, 128, 64, 32]
    for d in range(len(depth)):
        if d == 0:
            x = Conv2DTranspose(depth[d], (3, 3), strides = (2, 2), padding = 'same', kernel_initializer=RandomNormal(stddev=0.02))(latent)
        else:
            x = Conv2DTranspose(depth[d], (3, 3), strides = (2, 2), padding = 'same', kernel_initializer=RandomNormal(stddev=0.02))(x)
        x = BatchNormalization(center=False, scale=False)(x)
        m, v = FCNet(params, depth[d])
        rs = [tf.keras.backend.shape(latent)[0], 1, 1, depth[d]]
        m = tf.reshape(m, rs)
        v = tf.reshape(v, rs)
        x = Activation('relu')(x * v + m)

    x = convBlock(3, (3, 3), (1, 1), x, 'none', 'sigmoid')

    model = Model([latent, params], x)
    return model

损失函数代码

使用7个参数加权损失:5个对应VGG各层,1个对应MAE损失,1个对应MSE损失:

def ParamVGGLoss(params):
    def loss(y_true, y_pred):
        paramsReshaped = tf.reshape(params, [tf.keras.backend.shape(y_true)[0], 1, 1, paramCount])
        true = vgg(preprocess_input(y_true * 255))
        pred = vgg(preprocess_input(y_pred * 255))

        vggLoss = 0
        for i in range(len(true)):
            t = normalize_tensor(true[i])
            p = normalize_tensor(pred[i])
            sqDif = tf.math.square(t - p) * paramsReshaped[:, :, :, i : i + 1]
            vggLoss += tf.math.reduce_mean(sqDif)

        maeLoss = tf.math.abs(y_true - y_pred) * paramsReshaped[:, :, :, 5 : 6]
        maeLoss = tf.math.reduce_mean(maeLoss)
        mseLoss = tf.math.square(y_true - y_pred) * paramsReshaped[:, :, :, 6:]
        mseLoss = tf.math.reduce_mean(mseLoss) 
        return vggLoss + maeLoss + mseLoss
    return loss

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.18 14:22:01