集成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参数,和你的分类模型输入格式不匹配,因此触发错误。
修复方案
最直接的解决方法是给微调的分类模型设置唯一变量名,彻底避免冲突:
- 修改加载模型的代码,重命名分类模型:
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)
- 修改
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
相关产品推荐
相关产品推荐

