如何基于训练完成的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
相关产品推荐
相关产品推荐

