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

如何使用TensorFlow版BERT Transformer执行预测操作?

解决BERT微调后无法调用model.predict()的问题

首先,你的代码缺少模型加载与微调的关键环节,这是实现预测的基础。先补全这部分流程,再完成预测功能:

1. 补全模型加载与微调代码

from transformers import TFAutoModelForSequenceClassification

# 加载预训练模型,指定Yelp的5分类任务
model = TFAutoModelForSequenceClassification.from_pretrained("bert-base-cased", num_labels=5)

# 编译模型
model.compile(
    optimizer=tf.keras.optimizers.Adam(learning_rate=5e-5),
    loss=tf.keras.losses.SparseCategoricalCrossentropy(from_logits=True),
    metrics=tf.metrics.SparseCategoricalAccuracy(),
)

# 执行微调
model.fit(
    tf_train_dataset,
    validation_data=tf_validation_dataset,
    epochs=3
)

# 保存微调后的模型(可选,方便后续直接加载)
model.save_pretrained("./fine-tuned-bert-yelp")

2. 正确实现预测功能

无论处理单个字符串还是字符串列表,都需要先通过tokenizer将输入转换成模型可接受的格式,再调用model.predict(),最后解析结果:

预测单个字符串

# 若加载已保存的微调模型,替换为以下代码
# model = TFAutoModelForSequenceClassification.from_pretrained("./fine-tuned-bert-yelp")

test_text = "This restaurant has amazing food and friendly staff!"

# 预处理输入:参数需与训练时完全一致
inputs = tokenizer(
    test_text,
    padding="max_length",
    truncation=True,
    return_tensors="tf"
)

# 执行预测
predictions = model.predict(inputs)

# 解析结果:Yelp标签为1-5,模型输出索引为0-4,需+1匹配实际评分
predicted_label = tf.argmax(predictions.logits, axis=1).numpy()[0] + 1
print(f"预测评分:{predicted_label}")

预测字符串列表

test_texts = [
    "Worst experience ever, the service was terrible.",
    "The food was okay, nothing special.",
    "Will definitely come back again, highly recommend!"
]

# 批量预处理输入
inputs = tokenizer(
    test_texts,
    padding="max_length",
    truncation=True,
    return_tensors="tf"
)

# 批量预测
predictions = model.predict(inputs)

# 解析每个结果
predicted_labels = tf.argmax(predictions.logits, axis=1).numpy() + 1
for text, label in zip(test_texts, predicted_labels):
    print(f"文本:{text}\n预测评分:{label}\n")

关键注意事项

  • 预处理参数(padding="max_length"、truncation=True)必须与训练时完全一致,否则模型无法正确处理输入。
  • Yelp数据集的标签是1-5分,模型输出的logits索引对应0-4,因此需给预测结果加1才能匹配实际评分。
  • 加载微调后的模型时,要确保tokenizer与模型的预训练版本匹配(均为bert-base-cased或对应微调模型路径)。

内容的提问来源于stack exchange,提问作者pavis

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.24 18:54:26