Azure Databricks中基于VGG19的自定义TableNet模型保存问题
问题概述
基于TableNet+VGG19的表格提取模型,训练数据采用Marmoot,保存路径为映射Azure数据湖的DBFS路径。在Databricks平台尝试多种模型保存方式均出现错误,具体场景如下:
错误场景明细
Pickle保存加载失败
保存代码:import pickle pickle.dump(model, open(filepath, 'wb'))保存时触发未追踪函数警告,加载时执行
loaded_model = pickle.load(open(filepath, 'rb')),抛出ValueError: Unable to restore custom object of type _tf_keras_metric。model.save()进程崩溃
调用model.save(filepath)时,Python内核无响应,进程因段错误终止(exit code 139)。model.save_weights()加载失败
执行model.save_weights(weights_path)后,加载权重时无法恢复模型状态。ModelCheckpoint回调报错
添加tf.keras.callbacks.ModelCheckpoint回调后,首次epoch结束时抛出OSError: Unable to create file (file signature not found)。跨环境加载异常
在非Databricks环境用model.save()保存模型后,加载时出现类似Pickle的自定义对象错误;传入custom_objects参数后,又先后触发ValueError: Unknown layer: table_mask和TypeError: 'KerasTensor' object is not callable。
针对性解决方案
1. 处理自定义层与Metrics的加载问题
TableNet包含自定义层(如table_mask)及可能的自定义Metrics,加载时必须显式注册这些对象:
- 首先确保所有自定义组件的类定义完整(需与训练时的实现一致):
import tensorflow as tf # TableNet自定义table_mask层示例(根据你的实际代码调整) class TableMaskLayer(tf.keras.layers.Layer): def __init__(self, **kwargs): super().__init__(**kwargs) # 层初始化逻辑 def call(self, inputs): # 前向传播逻辑 return outputs # 自定义Metric示例(如果训练时使用了自定义指标) class TableDetectionMetric(tf.keras.metrics.Metric): def __init__(self, name='table_detection_acc', **kwargs): super().__init__(name=name, **kwargs) self.acc = self.add_weight(name='acc', initializer='zeros') def update_state(self, y_true, y_pred, sample_weight=None): # 指标更新逻辑 pass def result(self): return self.acc - 加载模型时将自定义对象传入
custom_objects:from tensorflow.keras.models import load_model custom_objects = { 'TableMaskLayer': TableMaskLayer, 'TableDetectionMetric': TableDetectionMetric } loaded_model = load_model('dbfs:/path/to/your/model', custom_objects=custom_objects)
2. 解决DBFS路径与进程段错误
- 路径规范:使用
dbfs:/开头的绝对路径,避免相对路径导致的文件系统权限或识别问题。 - 内存与版本排查:
- 段错误(exit code 139)多因内存不足或版本不兼容,升级Databricks集群的内存配置;
- 验证TensorFlow、Keras版本与TableNet实现兼容,推荐使用TensorFlow 2.8.x/2.9.x等稳定版本,避免跨大版本的API差异。
3. 修复ModelCheckpoint的OSError
- 权限与路径检查:确认集群服务主体对目标数据湖路径有读写权限,路径中避免特殊字符,层级不要过深;
- 调整Checkpoint参数:先尝试只保存权重,降低内存占用:
checkpoint_callback = tf.keras.callbacks.ModelCheckpoint( filepath='dbfs:/path/to/checkpoint', save_weights_only=True, save_best_only=True, verbose=1 )
4. 稳定保存加载方案(推荐)
优先使用TensorFlow官方推荐的SavedModel格式:
- 保存模型:
model.save('dbfs:/path/to/saved_model', save_format='tf') - 加载模型(配合自定义对象):
loaded_model = tf.keras.models.load_model('dbfs:/path/to/saved_model', custom_objects=custom_objects) - 权重保存加载:需先重建模型结构,再加载权重:
# 重建与训练时一致的模型结构 def build_table_net(): # 你的TableNet+VGG19模型构建代码 return model model = build_table_net() model.load_weights('dbfs:/path/to/weights')
5. 弃用Pickle保存Keras模型
Keras官方明确不推荐用Pickle保存完整模型,因为无法妥善处理TensorFlow内部状态、自定义层及Metrics,建议完全改用model.save()或save_weights()方案。
内容的提问来源于stack exchange,提问作者Lidor Eliyahu Shelef

