TFF运行报错RuntimeError: Attempting to capture an EagerTensor问题求解
问题原因
- 核心错误是
loss和metrics实例在全局作用域创建,被model_fn捕获时,其内部持有的Eager模式张量无法被TFF的计算图追踪流程识别 - 预处理函数中的
reshape未明确使用tf.reshape,如果引入的是numpy的reshape方法,也会生成Eager张量触发报错
修复方案
- 将损失函数、评价指标的实例化逻辑移入
model_fn内部,保证每次调用model_fn都生成全新的实例,不会携带外部Eager上下文 - 预处理函数中的reshape操作明确使用
tf.reshape - (可选)如果仍有问题,可将
input_spec改为显式声明的tf.TensorSpec,避免从样本数据集捕获隐含的Eager属性
修正后的核心代码片段
首先修改预处理函数:
def preprocess(dataset): NUM_EPOCHS = 5 BATCH_SIZE = 32 PREFETCH_BUFFER = 10 def batch_format_fn(element): return collections.OrderedDict( x=tf.reshape(element['x'], [-1, 13055]), y=tf.reshape(element['y'], [-1, 2])) return dataset.repeat(NUM_EPOCHS).batch(BATCH_SIZE).map( batch_format_fn).prefetch(PREFETCH_BUFFER)
删除全局的losses、metric定义,修改model_fn如下:
def model_fn(): keras_model = CNN() losses = tf.keras.losses.CategoricalCrossentropy() metric = [tf.keras.metrics.CategoricalAccuracy()] return tff.learning.from_keras_model( keras_model, input_spec=preprocessed_sample_dataset.element_spec, loss=losses, metrics=metric)
其余代码保持原有逻辑即可正常运行。
内容的提问来源于stack exchange,提问作者user10985800
相关产品推荐
相关产品推荐

