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

