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

集成Whisper到微调模型聊天机器人时遇TypeError:缺失tokens参数

问题

在Google Colab开发聊天机器人,先通过纯文本测试微调的分类模型运行正常,但集成Whisper音频转录后,调用generate_response函数生成响应时触发错误:

TypeError                                 Traceback (most recent call last)
<ipython-input-40-de822a5197f3> in <cell line: 4>()
     15 
     16         # Generate a response using the chatbot model
---> 17         response = generate_response(transcription)
     18 
     19     else:

1 frames
/usr/local/lib/python3.10/dist-packages/torch/nn/modules/module.py in _call_impl(self, *args, **kwargs)
   1499                 or _global_backward_pre_hooks or _global_backward_hooks
   1500                 or _global_forward_hooks or _global_forward_pre_hooks):
-> 1501             return forward_call(*args, **kwargs)
   1502         # Do not call functions when jit is used
   1503         full_backward_hooks, non_full_backward_hooks = [], []

TypeError: Whisper.forward() missing 1 required positional argument: 'tokens'

加载模型代码

model.save_pretrained("/content/saved_model")
tokenizer.save_pretrained("/content/saved_tokenizer")

model_path = "/content/saved_model"
tokenizer_path = "/content/saved_tokenizer"

model = AutoModelForSequenceClassification.from_pretrained(model_path)
tokenizer = AutoTokenizer.from_pretrained(tokenizer_path)

device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
model.to(device)

纯文本聊天函数(运行正常)

def generate_response(input_text):
    input_ids = tokenizer(input_text, return_tensors="pt").input_ids.to(device)
    with torch.no_grad():
        outputs = model(input_ids)
    predicted_label = torch.argmax(outputs.logits).item()

    # Use the original label name from the dataframe
    predicted_disease_or_symptom = df.loc[df['labels'] == predicted_label, 'label'].iloc[0]

    return predicted_disease_or_symptom 

while True:
    user_input = input("User: ")
    if user_input.lower() == "exit":
        break
    response = generate_response(user_input)
    print(f"Chatbot: {response}")

集成Whisper后的代码(触发错误)

!pip install git+https://github.com/openai/whisper.git 

import whisper
themodel = whisper.load_model('base')

while True:
    user_input = input("User: ")
    if user_input.lower() == "exit":
        break
    if user_input.lower() == "audio":
        audio_path = input("Enter audio file path: ")
        
        # Transcribe audio using Whisper
        text = themodel.transcribe(audio_path)
        transcription = text['text']
        print("Transcribed Text:", transcription)
        
        # Generate a response using the chatbot model
        response = generate_response(transcription)
        
    else:
        response = generate_response(user_input)
    print(f"Chatbot: {response}")

问题原因

核心问题是全局变量名冲突:你加载Whisper模型后,全局环境中的model变量被意外覆盖为Whisper实例。虽然你用themodel存储Whisper模型,但Whisper内部初始化过程可能修改了全局的model变量,导致generate_response函数调用model(input_ids)时,实际调用的是Whisper模型的forward方法——而Whisper的forward需要tokens参数,和你的分类模型输入格式不匹配,因此触发错误。

修复方案

最直接的解决方法是给微调的分类模型设置唯一变量名,彻底避免冲突:

  1. 修改加载模型的代码,重命名分类模型:
model.save_pretrained("/content/saved_model")
tokenizer.save_pretrained("/content/saved_tokenizer")

model_path = "/content/saved_model"
tokenizer_path = "/content/saved_tokenizer"

# 重命名分类模型,避免和Whisper相关变量冲突
classifier_model = AutoModelForSequenceClassification.from_pretrained(model_path)
tokenizer = AutoTokenizer.from_pretrained(tokenizer_path)

device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
classifier_model.to(device)
  1. 修改generate_response函数,调用重命名后的分类模型:
def generate_response(input_text):
    input_ids = tokenizer(input_text, return_tensors="pt").input_ids.to(device)
    with torch.no_grad():
        # 使用重命名后的分类模型
        outputs = classifier_model(input_ids)
    predicted_label = torch.argmax(outputs.logits).item()

    predicted_disease_or_symptom = df.loc[df['labels'] == predicted_label, 'label'].iloc[0]

    return predicted_disease_or_symptom 

替代方案

如果不想修改变量名,也可以把分类模型逻辑封装到类中,通过实例隔离全局变量,代码结构更清晰:

class ChatbotClassifier:
    def __init__(self, model_path, tokenizer_path, df):
        self.device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
        self.model = AutoModelForSequenceClassification.from_pretrained(model_path)
        self.tokenizer = AutoTokenizer.from_pretrained(tokenizer_path)
        self.model.to(self.device)
        self.df = df

    def generate_response(self, input_text):
        input_ids = self.tokenizer(input_text, return_tensors="pt").input_ids.to(self.device)
        with torch.no_grad():
            outputs = self.model(input_ids)
        predicted_label = torch.argmax(outputs.logits).item()
        return self.df.loc[self.df['labels'] == predicted_label, 'label'].iloc[0]

# 初始化分类器实例
chatbot = ChatbotClassifier("/content/saved_model", "/content/saved_tokenizer", df)

# 集成Whisper后的调用逻辑
import whisper
whisper_model = whisper.load_model('base')

while True:
    user_input = input("User: ")
    if user_input.lower() == "exit":
        break
    if user_input.lower() == "audio":
        audio_path = input("Enter audio file path: ")
        text = whisper_model.transcribe(audio_path)
        transcription = text['text']
        print("Transcribed Text:", transcription)
        response = chatbot.generate_response(transcription)
    else:
        response = chatbot.generate_response(user_input)
    print(f"Chatbot: {response}")

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.13 04:22:19