加载预训练Distil-BERT至Gradio时出现BERT_Arch无predict属性错误
问题描述
开发快餐店语音助手时,基于Distil-BERT训练了多分类意图分类模型,以pickle格式保存后,在Gradio应用中加载模型识别意图时触发属性错误:'BERT_Arch' object has no attribute 'predict'。尝试使用Trainer.predict方法时,又提示不存在Trainer属性。已在Gradio代码中添加BERT_Arch类(因加载pickle模型时无法找到该类)。
问题代码
import gradio as gr import pickle import torch.nn as nn class BERT_Arch(nn.Module): def __init__(self, bert): super(BERT_Arch, self).__init__() self.bert = bert # dropout layer self.dropout = nn.Dropout(0.2) # relu activation function self.relu = nn.ReLU() # dense layer self.fc1 = nn.Linear(768,512) self.fc2 = nn.Linear(512,256) self.fc3 = nn.Linear(256,5) #softmax activation function self.softmax = nn.LogSoftmax(dim=1) #define the forward pass def forward(self, sent_id, mask): #pass the inputs to the model cls_hs = self.bert(sent_id, attention_mask=mask)[0][:,0] x = self.fc1(cls_hs) x = self.relu(x) x = self.dropout(x) x = self.fc2(x) x = self.relu(x) x = self.dropout(x) # output layer x = self.fc3(x) # apply softmax activation x = self.softmax(x) return x def make_prediction(text): with open("model.pickle", "rb") as f: clf = pickle.load(f) preds = clf.predict([text]) # Assuming your model accepts a single text input if preds == 1: return "Sure I will add your order, that will be $30, anything else?" elif preds == 2: return "I have added newer items in the list as well that will be $55." elif preds == 3: return "That's a great choice, you want in the veg or non-veg section?" elif preds == 4: return "Your order will be ready in 20 mins" elif preds == 5: return "Okay, I will make the quantity according to the specified number of people" output = gr.Textbox() app = gr.Interface(fn=make_prediction, inputs="text", outputs=output) app.launch()
错误信息
Traceback (most recent call last): File "c:\Users\hp\AppData\Local\Programs\Python\Python39\lib\site-packages\gradio\routes.py", line 516, in predict output = await route_utils.call_process_api( File "c:\Users\hp\AppData\Local\Programs\Python\Python39\lib\site-packages\gradio\route_utils.py", line 219, in call_process_api output = await app.get_blocks().process_api( File "c:\Users\hp\AppData\Local\Programs\Python\Python39\lib\site-packages\gradio\blocks.py", line 1437, in process_api result = await self.call_function( File "c:\Users\hp\AppData\Local\Programs\Python\Python39\lib\site-packages\gradio\blocks.py", line 1109, in call_function prediction = await anyio.to_thread.run_sync( File "c:\Users\hp\AppData\Local\Programs\Python\Python39\lib\site-packages\anyio\to_thread.py", line 33, in run_sync return await get_asynclib().run_sync_in_worker_thread( File "c:\Users\hp\AppData\Local\Programs\Python\Python39\lib\site-packages\anyio\_backends\_asyncio.py", line 877, in run_sync_in_worker_thread return await future File "c:\Users\hp\AppData\Local\Programs\Python\Python39\lib\site-packages\anyio\_backends\_asyncio.py", line 807, in run result = context.run(func, *args) File "c:\Users\hp\AppData\Local\Programs\Python\Python39\lib\site-packages\gradio\utils.py", line 650, in wrapper response = f(*args, **kwargs) File "C:\Users\hp\AppData\Local\Temp\ipykernel_14968\3045401095.py", line 24, in make_prediction preds = clf.predict([text]) # Assuming your model accepts a single text input File "c:\Users\hp\AppData\Local\Programs\Python\Python39\lib\site-packages\torch\nn\modules\module.py", line 1614, in __getattr__ raise AttributeError("'{}' object has no attribute '{}'".format( AttributeError: 'BERT_Arch' object has no attribute 'predict'
解决方案
核心问题
BERT_Arch是PyTorch的nn.Module子类,本身没有predict方法(该方法是scikit-learn模型的标准方法)。PyTorch模型需要手动实现完整的推理流程,包括文本转模型输入张量、前向传播、结果解析三个核心步骤。另外,你还缺少文本tokenization环节——Distil-BERT无法直接处理原始文本,必须用对应的tokenizer转换成模型需要的input_ids和attention_mask。
修改后的完整代码
import gradio as gr import pickle import torch import torch.nn as nn from transformers import DistilBertTokenizer, DistilBertModel class BERT_Arch(nn.Module): def __init__(self, bert): super(BERT_Arch, self).__init__() self.bert = bert self.dropout = nn.Dropout(0.2) self.relu = nn.ReLU() self.fc1 = nn.Linear(768,512) self.fc2 = nn.Linear(512,256) self.fc3 = nn.Linear(256,5) self.softmax = nn.LogSoftmax(dim=1) def forward(self, sent_id, mask): cls_hs = self.bert(sent_id, attention_mask=mask)[0][:,0] x = self.fc1(cls_hs) x = self.relu(x) x = self.dropout(x) x = self.fc2(x) x = self.relu(x) x = self.dropout(x) x = self.fc3(x) x = self.softmax(x) return x # 初始化tokenizer(和训练时用的Distil-BERT版本一致) tokenizer = DistilBertTokenizer.from_pretrained('distilbert-base-uncased') def make_prediction(text): # 加载模型 with open("model.pickle", "rb") as f: model = pickle.load(f) # 切换到评估模式,关闭dropout等训练层 model.eval() # 处理输入文本:tokenize并生成模型需要的张量 encoding = tokenizer.encode_plus( text, add_special_tokens=True, max_length=64, padding='max_length', truncation=True, return_attention_mask=True, return_tensors='pt' ) input_ids = encoding['input_ids'] attention_mask = encoding['attention_mask'] # 模型推理(禁用梯度计算,提升速度) with torch.no_grad(): outputs = model(input_ids, attention_mask) # 取logits的argmax得到预测类别索引(注意:索引是0-4,对应你的1-5类别,所以+1) pred_idx = torch.argmax(outputs, dim=1).item() + 1 # 根据预测结果返回对应回复 if pred_idx == 1: return "Sure I will add your order, that will be $30, anything else?" elif pred_idx == 2: return "I have added newer items in the list as well that will be $55." elif pred_idx == 3: return "That's a great choice, you want in the veg or non-veg section?" elif pred_idx == 4: return "Your order will be ready in 20 mins" elif pred_idx == 5: return "Okay, I will make the quantity according to the specified number of people" output = gr.Textbox() app = gr.Interface(fn=make_prediction, inputs="text", outputs=output) app.launch()
关键修改点说明
- 引入Tokenizer:添加
DistilBertTokenizer,将原始文本转换成模型能识别的input_ids和attention_mask张量,确保和训练时使用的tokenizer版本一致。 - 实现推理逻辑:替换原来的
clf.predict([text]),手动完成tokenization→张量生成→模型前向传播→结果解析的流程。 - 模型评估模式:调用
model.eval()关闭dropout等训练相关层,避免推理结果波动。 - 禁用梯度计算:用
torch.no_grad()包裹推理代码,减少内存占用并提升速度。 - 类别索引修正:模型输出的logits索引是0-4,对应你定义的1-5类别,所以需要+1匹配后续的回复逻辑。
长期优化建议
不要用pickle保存PyTorch模型,pickle容易导致兼容性问题。正确做法是保存模型的state_dict:
# 训练时保存 torch.save(model.state_dict(), 'model_state_dict.pt') # 加载时 bert = DistilBertModel.from_pretrained('distilbert-base-uncased') model = BERT_Arch(bert) model.load_state_dict(torch.load('model_state_dict.pt'))
内容的提问来源于stack exchange,提问作者Isha
相关产品推荐
相关产品推荐

