基于Transformers+Gradio+FAISS的聊天机器人出现TypeError问题求助
问题分析
你遇到的TypeError: string indices must be integers, not 'tuple'错误,核心原因有3个:
- 直接将HuggingFace原生模型传入FAISS,不符合LangChain的Embeddings接口要求,导致检索返回结果格式异常
- 错误调用
ConversationBufferMemory不存在的update_history方法,且返回值不符合Gradio要求 - 错误访问文档对象的
content属性,LangChain的Document对象实际用page_content存储内容
解决方案
以下是修正后的完整代码,修复点已标注在注释中:
import torch from transformers import ( DistilBertTokenizer, DistilBertModel, DistilBertForQuestionAnswering, ) import gradio as gr from langchain.memory import ConversationBufferMemory from langchain_community.vectorstores import FAISS # 引入HuggingFaceEmbeddings,用于包装原生模型 from langchain_community.embeddings import HuggingFaceEmbeddings import pathlib import logging # Set logging logging.basicConfig(level=logging.INFO) # Initialize tokenizer and models tokenizer = DistilBertTokenizer.from_pretrained("distilbert-base-uncased") embedding_model = DistilBertModel.from_pretrained("distilbert-base-uncased") qa_model = DistilBertForQuestionAnswering.from_pretrained( "distilbert-base-uncased-distilled-squad" ) device = torch.device("cuda" if torch.cuda.is_available() else "cpu") embedding_model.to(device) qa_model.to(device) # 修复点1:用HuggingFaceEmbeddings包装原生模型,符合LangChain接口要求 embeddings = HuggingFaceEmbeddings( model_name="distilbert-base-uncased", model_kwargs={"device": device}, tokenizer_kwargs={"padding": True, "truncation": True, "max_length": 512} ) # Load FAISS index index_path = pathlib.Path("./saved-index-faiss") # Update the path as necessary embeddings_db = FAISS.load_local( index_path, embeddings, allow_dangerous_deserialization=True ) retriever = embeddings_db.as_retriever( search_kwargs={"k": 3} ) # Adjust 'k' as necessary for your context retrieval # Setup memory for conversation memory = ConversationBufferMemory(memory_key="chat_history", return_messages=True) # Define a function to answer questions using the provided context def answer_question(question, context): inputs = tokenizer( question, context, return_tensors="pt", truncation=True, padding="max_length", max_length=512, ) inputs = {k: v.to(device) for k, v in inputs.items()} with torch.no_grad(): outputs = qa_model(**inputs) answer_start = torch.argmax(outputs.start_logits) answer_end = torch.argmax(outputs.end_logits) + 1 answer = tokenizer.decode( inputs["input_ids"][0][answer_start:answer_end], skip_special_tokens=True ) return answer # Define the Gradio function def ask(message, history): if not isinstance(message, str): message = str(message) # Convert to string if it's not already print("Input message:", message) # Fetch context based on the retrieval model results = retriever.invoke(message) # 修复点2:访问Document对象的page_content属性,而非content context = results[0].page_content if results else "No relevant information found." print("Retrieved context:", context) # 修复点3:使用正确的save_context方法更新对话记忆,且不将其作为返回值 answer = answer_question(message, context) memory.save_context({"input": message}, {"output": answer}) # 返回回答字符串,符合Gradio要求 return answer # Create the Gradio interface io = gr.ChatInterface( fn=ask, chatbot=gr.Chatbot(height=400), textbox=gr.Textbox(placeholder="Ask Away! ", container=False, scale=6), title="WikiBuddy", description="Ask me any question", theme="soft", examples=[ "Who is David Beckham?", "Is Bitcoin Dead?", "What club does Christiano Ronaldo plays for?", ], retry_btn=None, undo_btn="Delete Previous", clear_btn="Clear", analytics_enabled=True, fill_height=True, ) # Run the Gradio app io.launch()
关键修复说明
- Embeddings适配:LangChain的FAISS需要使用LangChain定义的Embeddings接口实现,而非直接传入HuggingFace原生模型,用
HuggingFaceEmbeddings包装后,确保索引加载和检索时格式正确。 - 文档属性修正:LangChain检索返回的
Document对象,内容存储在page_content属性中,而非你代码中的content,这是触发字符串索引错误的直接原因。 - 对话记忆修正:
ConversationBufferMemory没有update_history方法,正确更新记忆的方式是使用save_context;同时Gradio的ChatInterface要求函数仅返回回答字符串,不需要返回内存操作结果,避免额外格式错误。
内容的提问来源于stack exchange,提问作者Ghulam Mustafa
相关产品推荐
相关产品推荐

