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

TensorFlow模型重启加载权重后测试精度骤降问题排查

问题

我的模型接收二进制编码文本(非one-hot编码)的输入,训练后训练集上的binary accuracy达99.5%、accuracy达85%,训练完成后立即测试测试集精度约75%;但重启内核后,用make_model函数创建模型并加载保存的权重,测试精度仅35%,差异巨大。相关代码如下:

模型定义

# start span identification model definition - increased parameters
def make_model(learn_rate: float):
    inputs_sp_can = Input(shape=(MAX_BIN_SPAN_LEN,), name="Candidate Span")
    sp_embedding_layer = Embedding(
        len(char2idx),
        EMBEDDING_DIM,
        embeddings_initializer=keras.initializers.Constant(embed_matrix),
        trainable=False,
    )(inputs_sp_can)
    x = Bidirectional(LSTM(960))(sp_embedding_layer)
    # x = BatchNormalization()(x)
    x = Dense(1920,activation = 'selu')(x)
    x = Dense(640, activation='selu')(x)
    x = Dropout(0.5)(x)
    x = Dense(1280, activation='selu')(x)
    x = models.Model(inputs=inputs_sp_can, outputs=x)

    inputs_sent = Input(shape=(MAX_BIN_SENT_LEN,), name="Input Sentences")
    sent_embedding_layer = Embedding(
        len(char2idx),
        EMBEDDING_DIM,
        embeddings_initializer=keras.initializers.Constant(embed_matrix),
        trainable=False,
    )(inputs_sent)
    y = Bidirectional(LSTM(1800))(sent_embedding_layer)
    # y = BatchNormalization()(y)
    y = Dense(1920,activation = 'selu')(y)
    y = Dense(640, activation='selu')(y)
    y = Dropout(0.5)(y)
    y = Dense(1280, activation='selu')(y)
    y = models.Model(inputs=inputs_sent, outputs=y)

    combined_layer = Concatenate(axis=1)(
        [x.output, y.output]
    )  # ([x.output,w.output,y.output])
    z = Dense(1024, activation='selu')(combined_layer)
    z = Dense(512, activation='selu')(z)
    z = Dense(512, activation='selu')(z)
    # z = Dropout(0.5)(z)
    # z = Dense(1024, activation='selu')(z)

    pre_span_label = Dense(128, activation="selu", name="pre_span_label")(z)
    span_label = Dense(len(label_idx), activation="softmax", name="span_label")(
        pre_span_label
    )

    model = models.Model(
        inputs=[x.input, y.input], outputs=[span_label]
    )  # [x.input,w.input,y.input], outputs = [span_check,span_label])

    myoptimizer = optimizers.Adam(learning_rate=learn_rate, clipnorm=0.1)

    model.compile(
        optimizer=myoptimizer,
        loss="categorical_crossentropy",
        metrics=["binary_accuracy", "accuracy"],
    )

    return model 

训练代码

filepath = save_path + "Char_Tok_Models/Models/weights.{epoch:02d}.hdf5"
checkpoint = ModelCheckpoint(
    filepath,
    monitor="val_loss",
    verbose=1,
    save_best_only=False,
    mode="max",
    period=5
)

callbacks_list = [checkpoint]

history = model.fit(
    [x_span_train, corr_sent_train],
    corr_label_train,
    validation_data=([x_span_val, corr_sent_val], [corr_label_val]),
    epochs=25,
    callbacks=callbacks_list,
    batch_size=64,
)  #

权重加载与评估代码

model = make_model(0.0002)
model.load_weights(save_path + "Char_Tok_Models/Models_Upd/weights.40.hdf5")

model.evaluate([x_span_test, corr_sent_test], corr_label_test)

已尝试调整学习率、将最后一层激活函数改为sigmoid,期望推理精度至少达到80%,寻求解决办法。


问题分析与解决方案

1. 权重文件与训练epoch不匹配

这是最核心的问题:训练仅设置了25个epoch,但你加载的是weights.40.hdf5,这个文件根本不是当前训练流程生成的,加载错误的权重直接导致精度崩盘。

  • 解决办法:加载训练实际生成的权重文件,比如weights.25.hdf5(因period=5,训练会保存5、10、15、20、25 epoch的权重)。

2. 模型构建参数不一致

训练时使用learn_rate=0.00015创建模型,但加载时用了0.0002。虽然学习率不影响权重存储,但如果模型构建过程中存在隐性依赖参数的逻辑,可能导致层名称或结构细微差异,引发权重映射错误。

  • 解决办法:加载模型时使用与训练完全一致的学习率,即调用make_model(0.00015)。

3. 嵌入层初始化不匹配

重启内核后,若char2idx字典或embed_matrix重新生成时与训练阶段不一致(比如字符顺序变化、嵌入矩阵维度错误),会导致输入编码无法映射到训练时的嵌入向量,直接破坏模型逻辑。

  • 解决办法:训练时将char2idx和embed_matrix用pickle保存为文件,重启后直接加载保存的文件,而非重新生成;验证len(char2idx)和embed_matrix的形状在训练与加载时完全一致。

4. Dropout模式异常

训练时模型处于训练模式,Dropout生效;推理时模型应自动切换到测试模式关闭Dropout,但部分情况下可能出现模式切换失败,导致推理时仍随机丢弃神经元,降低精度。

  • 解决办法:评估前添加model.trainable = False,明确切换到推理模式;或者在加载权重后调用model.compile()确保模式正确。

5. 模型过拟合叠加问题

训练集accuracy 85%、测试集立即测试仅75%,已存在明显过拟合,重启后的精度暴跌是过拟合+其他问题的叠加结果。你的模型参数规模过大(LSTM单元960/1800、Dense层1920/1280),远超任务需求。

  • 解决办法:
    • 缩减模型参数:将LSTM单元数降至256/512,Dense层单元数减半,降低模型容量。
    • 启用代码中注释的BatchNormalization层,缓解过拟合。
    • 增加Dropout比例,或在更多层后添加Dropout。
    • 针对文本任务做数据增强:同义词替换、随机插入/删除字符等,提升训练数据多样性。

6. 评估指标与任务不匹配

binary_accuracy适用于二分类任务,若你是多分类任务,这个指标没有实际意义,还会误导对模型性能的判断。

  • 解决办法:多分类任务改用sparse_categorical_accuracy作为评估指标,确保指标与任务匹配。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.10 13:01:03