Keras2.8.0运行BERT遇AttributeError:'Functional'无_jit_compile属性求助
解决Keras 2.8.0中'Functional' object has no attribute '_jit_compile'错误
问题原因
Keras从2.3.0升级到2.8.0后,train_function的实现逻辑发生了变化——它不再是独立的函数,而是与模型实例的内部属性(比如_jit_compile)绑定的方法。直接用自定义函数覆盖model.train_function会丢失原方法关联的模型属性,从而触发报错。
解决方案
方法1:禁用JIT编译(快速适配旧代码)
在模型编译阶段添加jit_compile=False参数,跳过JIT相关的属性检查,直接复用你的自定义训练逻辑:
# 修改模型编译代码,新增jit_compile参数 model.compile( optimizer=your_optimizer, loss=your_loss, metrics=your_metrics, jit_compile=False ) # 保留原自定义训练函数逻辑 old_train_function = model.train_function def train_function(inputs): # 重新定义训练函数 grads = embedding_gradients(inputs)[0] # Embedding梯度 delta = epsilon * grads / (np.sqrt((grads**2).sum()) + 1e-8) # 计算扰动 K.set_value(embeddings, K.eval(embeddings) + delta) # 注入扰动 outputs = old_train_function(inputs) # 梯度下降 K.set_value(embeddings, K.eval(embeddings) - delta) # 删除扰动 return outputs model.train_function = train_function # 覆盖原训练函数
方法2:使用自定义训练循环(符合新版本规范)
Keras 2.6+推荐用自定义训练循环替代直接覆盖train_function,这种方式兼容性更强,也更易维护:
import tensorflow as tf # 保存模型原有的优化器、损失和指标 optimizer = model.optimizer loss_fn = model.loss metrics = model.metrics # 定义自定义训练步骤 @tf.function def train_step(inputs): x, y = inputs with tf.GradientTape() as tape: # 注入Embedding扰动 grads = embedding_gradients(inputs)[0] delta = epsilon * grads / (tf.sqrt(tf.reduce_sum(grads**2)) + 1e-8) embeddings.assign_add(delta) # 前向传播计算损失 y_pred = model(x, training=True) loss = loss_fn(y, y_pred) loss += sum(model.losses) # 加入正则化损失 # 计算梯度并更新模型参数 trainable_vars = model.trainable_variables gradients = tape.gradient(loss, trainable_vars) optimizer.apply_gradients(zip(gradients, trainable_vars)) # 移除Embedding扰动 embeddings.assign_sub(delta) # 更新评估指标 for metric in metrics: metric.update_state(y, y_pred) return {m.name: m.result() for m in metrics} # 执行训练循环 for epoch in range(epochs): for batch in train_dataset: metrics_result = train_step(batch) # 打印epoch结果 print(f"Epoch {epoch+1}: {metrics_result}") # 重置指标 for metric in metrics: metric.reset_states()
说明
- 方法1改动最小,适合快速迁移旧代码,但会丢失JIT编译带来的性能提升。
- 方法2是Keras新版本的标准用法,兼容性和可维护性更好。
内容的提问来源于stack exchange,提问作者jolin
相关产品推荐
相关产品推荐

