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

加载预训练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()

关键修改点说明

  1. 引入Tokenizer:添加DistilBertTokenizer,将原始文本转换成模型能识别的input_ids和attention_mask张量,确保和训练时使用的tokenizer版本一致。
  2. 实现推理逻辑:替换原来的clf.predict([text]),手动完成tokenization→张量生成→模型前向传播→结果解析的流程。
  3. 模型评估模式:调用model.eval()关闭dropout等训练相关层,避免推理结果波动。
  4. 禁用梯度计算:用torch.no_grad()包裹推理代码,减少内存占用并提升速度。
  5. 类别索引修正:模型输出的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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.10 11:23:14