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

添加ModelCheckpoint回调与验证数据时遇EagerTensor序列化错误

问题

复现Keras官方《小数据集下的Vision Transformer》示例时,添加了ModelCheckpoint和EarlyStopping回调并传入验证数据,代码如下:

cbs = [ModelCheckpoint("ViT-Model-new-small-dataset.h5", save_best_only=True), 
       EarlyStopping(patience=50, monitor='val_Accuracy', mode='max', restore_best_weights=True)]

history = model.fit(
    x=x_train,
    y=y_train,
    validation_data=(x_valid, y_valid),
    batch_size=BATCH_SIZE,
    epochs=EPOCHS, 
    callbacks=cbs
)

第一个训练epoch结束后触发如下错误:

File "vit_test_small_datasets.py", line 359, in <module>
    history = run_experiment(vit_sl)
  File "vit_test_small_datasets.py", line 335, in run_experiment
    history = model.fit(
  File "/usr/local/lib/python3.8/dist-packages/keras/utils/traceback_utils.py", line 67, in error_handler
    raise e.with_traceback(filtered_tb) from None
  File "/usr/lib/python3.8/json/__init__.py", line 234, in dumps
    return cls(
  File "/usr/lib/python3.8/json/encoder.py", line 199, in encode
    chunks = self.iterencode(o, _one_shot=True)
  File "/usr/lib/python3.8/json/encoder.py", line 257, in iterencode
    return _iterencode(o, 0)
TypeError: Unable to serialize [  0   1   2   3   4   5   6   7   8   9  10  11  12  13  14  15  16  17
  18  19  20  21  22  23  24  25  26  27  28  29  30  31  32  33  34  35
  36  37  38  39  40  41  42  43  44  45  46  47  48  49  50  51  52  53
  54  55  56  57  58  59  60  61  62  63  64  65  66  67  68  69  70  71
  72  73  74  75  76  77  78  79  80  81  82  83  84  85  86  87  88  89
  90  91  92  93  94  95  96  97  98  99 100 101 102 103 104 105 106 107
 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125
 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143] to JSON. Unrecognized type <class 'tensorflow.python.framework.ops.EagerTensor'>.

已尝试降级TensorFlow到2.9.0/2.9.1,问题仍存在;移除验证数据和ModelCheckpoint则无错误。询问是否能在该ViT示例中使用上述回调。

解决方案

可以在该ViT示例中使用ModelCheckpoint和EarlyStopping回调,问题根源是EagerTensor无法被JSON序列化,可通过以下方法解决:

  • 修正监控指标名称:Keras默认准确率指标名称为val_accuracy(小写a),而非val_Accuracy,名称不匹配会导致内部张量处理异常,修改回调参数:

    EarlyStopping(patience=50, monitor='val_accuracy', mode='max', restore_best_weights=True)
    
  • 转换验证数据类型:如果x_valid或y_valid是EagerTensor类型,转换为NumPy数组后再传入:

    x_valid = x_valid.numpy()
    y_valid = y_valid.numpy()
    
  • 改用SavedModel格式保存模型:.h5格式对复杂模型(如ViT)的序列化支持有限,改用TensorFlow原生的SavedModel格式:

    ModelCheckpoint("ViT-Model-new-small-dataset", save_best_only=True, save_format="tf")
    

    注意此处无需添加.h5后缀,TensorFlow会自动生成模型文件夹。

  • 明确指定编译阶段的指标:在模型编译时显式声明准确率指标,避免自动推断引发的张量问题:

    model.compile(
        optimizer=optimizer,
        loss=keras.losses.SparseCategoricalCrossentropy(from_logits=True),
        metrics=[keras.metrics.SparseCategoricalAccuracy(name="accuracy")]
    )
    

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.17 09:37:02