如何在Flask API中实现HuggingFace LLM的LangChain流式响应?
问题描述
我开发了一个基于Flask的API,用来流式返回由LangChain封装的LLM响应。使用OpenAI模型时流式功能正常,但切换到HuggingFace模型(如Llama2-13B)后流式输出失效。以下是相关代码片段,求解决办法。
model.py 代码片段
def load_llama2_13b(): model_str = "meta-llama/Llama-2-13b-chat-hf" access_token = "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" tokenizer = AutoTokenizer.from_pretrained(model_str, use_auth_token=access_token) model = AutoModelForCausalLM.from_pretrained( model_str, device_map="auto", # quantization_config=bnb_config, trust_remote_code=True, use_auth_token=access_token) streamer = TextStreamer(tokenizer) llm_pipeline = pipeline( "text-generation", # task model=model, tokenizer=tokenizer, trust_remote_code=True, device_map="auto", do_sample=True, max_new_tokens=300, streamer=streamer, eos_token_id=tokenizer.eos_token_id, model_kwargs={"temperature": 0.01, "repetition_penalty": 2.5} ) llm = HuggingFacePipeline(pipeline=llm_pipeline) return llm def query_llm(llm, query, task, prompt_template=None, vectordb=None): valid_tasks = ["instructive", "answer"] if task not in valid_tasks: raise ValueError(f"Invalid task '{task}'. Allowed values are {valid_tasks}") if task == "instructive": template = """ You are an intelligent chatbot. Answer the question posed by user. Question: {question} Answer:""" if task == "answer": if prompt_template: template = prompt_template else: template = """ You are an intelligent chatbot. Given the context below, answer the question given at the end: Context: {context} QUESTION: {question} Answer:""" PROMPT = PromptTemplate(template=template, input_variables=["context", "question"]) llm_chain = RetrievalQA.from_chain_type(llm, chain_type="stuff", retriever=vectordb.as_retriever(), return_source_documents=True, chain_type_kwargs={"prompt": PROMPT}) res = llm_chain({'query': query}) # import pdb;pdb.set_trace() source_docs = [t.__dict__ for t in res["source_documents"]] return {"result": res["result"], "source_documents": source_docs} PROMPT = PromptTemplate(template=template, input_variables=["question"]) llm_chain = LLMChain(llm=llm, prompt=PROMPT, verbose=True) return {"result": llm_chain.run(question=query)}
api.py 调用代码片段
@app.route('/llmgeneration_stream', methods=['POST']) def llm_generation_stream(): json_data = {"status": 0, "message": "Failed"} start = time.time() data = json.loads(flask.request.data) model = data.get("model", "falcon-7b") query = data.get("query") task = data.get("task", "instructive") prompt_template = data.get("prompt_template", None) valid_models = ["falcon-7b", "openai", "falcon-40b", "llama2-13b", "llama2-7b", "llama2-70b"] if model not in valid_models: raise ValueError(f"Invalid model '{model}'. Allowed values are {valid_models}") try: if model == "falcon-7b": llm = falcon_7b_llm return Response(stream_with_context(query_llm(llm, query, task, prompt_template, vectordb=vectordb)), mimetype='application/json') json_data["data"] = result json_data["message"] = "Passed" json_data["status"] = 1 end = time.time() logger.info("llm API: " + "Time: " + str(end - start) + "Model: " + model + "Query" + query + "text_generated" + str(result)) except Exception as e: logger.error(e, model, query) return flask.jsonify(json_data)
问题分析与解决办法
核心问题
- 原代码用
TextStreamer仅会将token打印到控制台,无法将输出传递给Flask响应流。 query_llm直接返回完整结果字典,没有生成器来逐块输出内容,无法对接Flask的流式响应机制。
修正步骤
1. 自定义流式处理类替换TextStreamer
创建基于队列的流式处理类,将生成的token暂存,供Flask逐块读取:
from transformers import BaseStreamer from queue import Queue from threading import Thread class QueueStreamer(BaseStreamer): def __init__(self, tokenizer, skip_prompt=True): self.tokenizer = tokenizer self.skip_prompt = skip_prompt self.queue = Queue() self.stop_signal = None def put(self, value): if self.skip_prompt and self.stop_signal is None: self.stop_signal = value elif value != self.stop_signal: self.queue.put(self.tokenizer.decode(value, skip_special_tokens=True)) def end(self): self.queue.put(None) def __iter__(self): while True: item = self.queue.get() if item is None: break yield item
2. 修改模型加载函数,绑定自定义Streamer
替换原TextStreamer,并将Streamer绑定到LLM对象上:
def load_llama2_13b(): # 原代码中tokenizer、model加载逻辑不变 streamer = QueueStreamer(tokenizer) llm_pipeline = pipeline( "text-generation", model=model, tokenizer=tokenizer, trust_remote_code=True, device_map="auto", do_sample=True, max_new_tokens=300, streamer=streamer, eos_token_id=tokenizer.eos_token_id, model_kwargs={"temperature": 0.01, "repetition_penalty": 2.5}, return_full_text=False # 避免返回prompt内容 ) llm = HuggingFacePipeline(pipeline=llm_pipeline) llm.streamer = streamer # 绑定streamer供后续调用 return llm
3. 重构query_llm返回生成器
将函数改为返回流式生成器,逐块输出内容:
def query_llm(llm, query, task, prompt_template=None, vectordb=None): valid_tasks = ["instructive", "answer"] if task not in valid_tasks: raise ValueError(f"Invalid task '{task}'. Allowed values are {valid_tasks}") if task == "instructive": template = """ You are an intelligent chatbot. Answer the question posed by user. Question: {question} Answer:""" PROMPT = PromptTemplate(template=template, input_variables=["question"]) llm_chain = LLMChain(llm=llm, prompt=PROMPT, verbose=True) # 启动线程执行生成,避免阻塞主线程 def generate(): llm_chain.run(question=query) Thread(target=generate).start() # 逐块返回生成的token,包装为JSON格式 for token in llm.streamer: yield f'{{"result": "{token}"}}\n' elif task == "answer": # RetrievalQA暂不支持原生流式,如需流式需额外重构,此处先返回完整结果 if prompt_template: template = prompt_template else: template = """ You are an intelligent chatbot. Given the context below, answer the question given at the end: Context: {context} QUESTION: {question} Answer:""" PROMPT = PromptTemplate(template=template, input_variables=["context", "question"]) llm_chain = RetrievalQA.from_chain_type(llm, chain_type="stuff", retriever=vectordb.as_retriever(), return_source_documents=True, chain_type_kwargs={"prompt": PROMPT}) res = llm_chain({'query': query}) source_docs = [t.__dict__ for t in res["source_documents"]] yield f'{{"result": "{res["result"]}", "source_documents": {source_docs}}}\n'
4. 修正Flask响应配置
调整响应类型为text/event-stream(更适合流式输出),并补充模型加载逻辑:
@app.route('/llmgeneration_stream', methods=['POST']) def llm_generation_stream(): data = json.loads(flask.request.data) model = data.get("model", "falcon-7b") query = data.get("query") task = data.get("task", "instructive") prompt_template = data.get("prompt_template", None) valid_models = ["falcon-7b", "openai", "falcon-40b", "llama2-13b", "llama2-7b", "llama2-70b"] if model not in valid_models: return flask.jsonify({"status":0, "message": f"Invalid model '{model}'. Allowed values are {valid_models}"}) try: # 补充其他模型的加载逻辑 if model == "falcon-7b": llm = falcon_7b_llm elif model == "llama2-13b": llm = load_llama2_13b() # 其他模型同理补充 return Response(stream_with_context(query_llm(llm, query, task, prompt_template, vectordb=vectordb)), mimetype='text/event-stream') except Exception as e: logger.error(e, model, query) return flask.jsonify({"status":0, "message": str(e)})
额外注意事项
- 确保
return_full_text=False,避免将prompt内容混入输出。 - 客户端需支持逐块接收响应(如使用Fetch API的ReadableStream)。
- 若需要RetrievalQA的流式输出,需自行实现检索后用LLMChain逐段生成的逻辑。
内容的提问来源于stack exchange,提问作者xerxes01
相关产品推荐
相关产品推荐

