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

如何基于训练完成的BERT模型文件开展输入测试并查看输出?

用Gradio测试本地训练的BERT模型步骤

1. 安装依赖

先确保安装所需的Python库:

pip install transformers torch gradio

2. 加载本地模型与Tokenizer

你的模型目录包含vocab.json、merges.txt和pytorch_model.bin,可直接用Hugging Face Transformers库加载:

from transformers import AutoTokenizer, AutoModelForSequenceClassification
import torch

# 替换为你的模型实际存放路径
model_path = "./your_model_folder"

# 加载tokenizer和模型
tokenizer = AutoTokenizer.from_pretrained(model_path)
# 根据任务替换模型类:分类用AutoModelForSequenceClassification,NER用AutoModelForTokenClassification
model = AutoModelForSequenceClassification.from_pretrained(model_path)
model.eval()  # 切换到评估模式,关闭训练相关参数更新

3. 编写预测逻辑

根据模型任务(分类、NER、特征提取等)编写预测函数,以下是分类任务示例:

def run_prediction(input_text):
    # 预处理输入文本,适配模型输入格式
    tokenized_input = tokenizer(
        input_text,
        return_tensors="pt",
        truncation=True,
        padding=True,
        max_length=512  # 匹配训练时的最大序列长度
    )
    
    # 无梯度推理,节省资源
    with torch.no_grad():
        model_output = model(**tokenized_input)
    
    # 处理输出(分类任务示例)
    logits = model_output.logits
    predicted_prob = torch.softmax(logits, dim=1).tolist()[0]
    predicted_class_idx = torch.argmax(logits, dim=1).item()
    
    # 替换为你训练时的实际类别标签
    class_labels = ["类别A", "类别B", "类别C"]
    result = {class_labels[i]: f"{prob*100:.2f}%" for i, prob in enumerate(predicted_prob)}
    return result

如果是NER任务,调整输出解码逻辑:

def run_ner_prediction(input_text):
    tokenized_input = tokenizer(input_text, return_tensors="pt", truncation=True, padding=True)
    with torch.no_grad():
        outputs = model(**tokenized_input)
    predictions = torch.argmax(outputs.logits, dim=2)
    # 映射标签ID到实体类型,替换为你的NER标签
    label_map = {0: "O", 1: "PER", 2: "LOC", 3: "ORG"}
    tokens = tokenizer.convert_ids_to_tokens(tokenized_input["input_ids"][0])
    result = [(token, label_map[pred.item()]) for token, pred in zip(tokens, predictions[0])]
    return result

4. 搭建Gradio可视化界面

用Gradio快速创建网页交互界面:

import gradio as gr

with gr.Blocks(title="BERT模型测试") as demo:
    gr.Markdown("### 输入文本测试你的BERT模型")
    input_box = gr.Textbox(label="输入文本", placeholder="请输入测试内容...")
    output_display = gr.Label(label="预测结果")
    predict_btn = gr.Button("开始预测")
    predict_btn.click(fn=run_prediction, inputs=input_box, outputs=output_display)

# 启动服务,默认访问地址为 http://localhost:7860
demo.launch()

关键注意事项

  • 必须替换model_path为你的模型实际存放路径
  • 模型类要和训练时的任务匹配:比如文本分类用AutoModelForSequenceClassification,问答任务用AutoModelForQuestionAnswering
  • 如果训练时用了自定义模型结构,需先导入自定义模型类,再用from_pretrained加载

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.22 11:50:14