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

