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

RCAN神经网络权重保存报错及预测匹配优化问题咨询

RCAN模型权重保存报错解决方案及预测结果优化建议

问题描述

  • 部署RCAN神经网络时,MeanShift函数的权重与偏置初始化存在问题,导致每5个周期保存权重时触发报错:

ValueError: Unable to serialize VariableSpec(shape=(1,), dtype=tf.float32, trainable=True, alias_id=None) to JSON, because the TypeSpec class <class 'tensorflow.python.ops.resource_variable_ops.VariableSpec'> has not been registered

  • 升级TensorFlow 2.4/2.9版本无法解决该问题
  • 输入与目标已做百分位数归一化,需要优化scaled_inputs和outputs处的操作,让预测结果更贴合目标

报错原因及解决方案

报错原因

原MeanShift函数中,先通过tf.Variable创建变量,再传入keras.initializers.Constant作为初始化器。这种方式会让Keras尝试序列化未注册的VariableSpec类型,直接导致权重保存失败。

修正后的MeanShift函数

直接使用tf.keras.initializers.Constant初始化权重和偏置,跳过tf.Variable的中间创建步骤,问题即可解决:

def MeanShift(rgb_range, rgb_mean, rgb_std, sign=-1):
    def _func(x):
        # 初始化权重和偏置
        W = tf.keras.initializers.Constant(1.)
        b = tf.keras.initializers.Constant(rgb_mean[0] * rgb_range * sign / rgb_std)

        x = Conv2D(1, 1, padding="same", kernel_initializer=W, bias_initializer=b)(x)
        return x

    return _func

原错误代码对比

原MeanShift函数的错误写法:

def MeanShift(rgb_range, rgb_mean, rgb_std, sign=-1):
    def _func(x):
        # 初始化权重
        weight_initer = tf.constant(value=1., dtype=tf.float32, shape=[1])
        W = tf.Variable(weight_initer, name="Weight", dtype=tf.float32)

        # 初始化偏置
        bias_initer = tf.constant(value=[m * rgb_range * sign / rgb_std for m in rgb_mean], dtype=tf.float32)
        b = tf.Variable(bias_initer, name="Bias", dtype=tf.float32)

        x = Conv2D(1, 1, padding="same", kernel_initializer=keras.initializers.Constant(W), bias_initializer=keras.initializers.Constant(b), activation="relu")(x)
        return x

    return _func

预测结果优化建议(针对百分位数归一化的输入/目标)

针对已做百分位数归一化的数据,可在scaled_inputs和outputs阶段添加以下操作提升匹配度:

  • 输入侧(scaled_inputs):
    • 增加tf.keras.layers.Normalization层,基于训练数据的统计量进一步校准分布,让输入更贴合模型预期
    • 添加tf.keras.layers.GaussianNoise层(如噪声强度设为0.01),增强模型泛化能力,避免过拟合
  • 输出侧(outputs):
    • 增加tf.keras.layers.Activation层,根据目标数据范围选对应激活函数:目标在[0,1]区间用sigmoid,对称区间(如[-1,1])用tanh
    • 搭配tf.keras.layers.Clip层限制输出在合理区间,比如目标归一化到[0,1],就添加tf.keras.layers.Clip(min_value=0, max_value=1)
    • 加入tf.keras.layers.LayerNormalization层,稳定输出分布,让预测结果更贴近目标的统计特性

完整RCAN模型代码

def RCAN(n_RCAB=3, n_RG=5, rgb_mean = [0.4488], rgb_std = 1, **kwargs):
    def _func(input_shape):
        import sys
        sys.setrecursionlimit(25000)

        # 模型起始部分
        inputs = Input(input_shape)
        scaled_inputs = tf.keras.layers.experimental.preprocessing.Rescaling(255)(inputs)
        scaled_inputs = MeanShift(rgb_range=255, rgb_mean=rgb_mean, rgb_std=rgb_std)(inputs)

        # Head模块
        x = Conv2D(64, 3, padding='same')(scaled_inputs)
        _inp = x

        # Body模块
        for i in range(n_RG):
            x = ResidualGroup(x, n_RCAB)
        x = Conv2D(64, 3, padding='same')(x)

        add_global = tf.keras.layers.Add()([x, _inp])

        # Tail模块
        x = UpSampling(add_global)
        x = Conv2D(1, 3, padding='same')(x)

        outputs = MeanShift(rgb_range=255, rgb_mean=rgb_mean, rgb_std=rgb_std, sign=1)(x)
        outputs = tf.keras.layers.experimental.preprocessing.Rescaling(1.0 / 255)(outputs)

        # 定义模型
        model = Model(inputs=inputs, outputs=outputs)
        return model
    return _func

def RCAB(x):
    _res = x

    x = Conv2D(64, 3, padding='same')(x)
    x = tf.keras.layers.LeakyReLU()(x)
    x = Conv2D(64, 3, padding='same')(x)

    x = ChannelAttention(x)
    x = tf.keras.layers.Add()([x, _res])
    return x

def ChannelAttention(x):
    _res = x

    avg_pool = tf.math.reduce_mean(x, axis=[1, 2, 3], keepdims=True)

    feat_mix = Conv2D(4, 1, padding='same', activation='relu')(avg_pool)
    feat_mix = Conv2D(64, 1, padding='same', activation='sigmoid')(feat_mix)

    multi = tf.keras.layers.Multiply()([feat_mix, _res])
    return multi

def ResidualGroup(x,  n_RCAB):
    skip_connection = x

    for i in range(n_RCAB):
        x = RCAB(x)

    x = Conv2D(64, 3, padding='same')(x)
    x = tf.keras.layers.Add()([x, skip_connection])
    return x

# 修正后的MeanShift函数
def MeanShift(rgb_range, rgb_mean, rgb_std, sign=-1):
    def _func(x):
        # 初始化权重和偏置
        W = tf.keras.initializers.Constant(1.)
        b = tf.keras.initializers.Constant(rgb_mean[0] * rgb_range * sign / rgb_std)

        x = Conv2D(1, 1, padding="same", kernel_initializer=W, bias_initializer=b)(x)
        return x

    return _func

def UpSampling(x, act=False):
    features = Conv2D(256, 3, padding='same')(x)
    high_res = tf.nn.depth_to_space(features, 2)

    if act:
        high_res = tf.keras.layers.ReLU(max_value=1)(high_res)

    return high_res

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.09 06:01:45