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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.10 09:20:25