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
相关产品推荐
相关产品推荐

