如何在LangChain中提升LLM响应速度?(Python开发场景)
解决方案:LangChain LLM响应提速及流式问题处理
一、解决流式模式下的两个问题
1. 统计使用的Token数量
可以通过自定义StreamingCallbackHandler结合Token统计库(如tiktoken)实现。示例代码如下:
from langchain.callbacks.base import BaseCallbackHandler import tiktoken class TokenCountHandler(BaseCallbackHandler): def __init__(self, model_name: str): self.encoder = tiktoken.encoding_for_model(model_name) self.prompt_tokens = 0 self.completion_tokens = 0 def on_llm_start(self, serialized, prompts, **kwargs): # 统计prompt的token数 self.prompt_tokens = sum(len(self.encoder.encode(prompt)) for prompt in prompts) def on_llm_new_token(self, token, **kwargs): # 累计生成的token数 self.completion_tokens += 1 # 使用时实例化handler token_handler = TokenCountHandler(model_name="gpt-3.5-turbo") chain = load_qa_chain( llm=ChatOpenAI(streaming=True, callbacks=[token_handler], temperature=0), chain_type="stuff" ) # 执行后获取token统计 print(f"Prompt tokens: {token_handler.prompt_tokens}, Completion tokens: {token_handler.completion_tokens}")
2. 流式返回至前端
以FastAPI为例,需要将LLM的流式输出包装成生成器,再通过StreamingResponse返回:
from fastapi import FastAPI from fastapi.responses import StreamingResponse from langchain.chat_models import ChatOpenAI from langchain.chains.question_answering import load_qa_chain from langchain.callbacks.base import BaseCallbackHandler from langchain.schema import Document app = FastAPI() class StreamingResponseHandler(BaseCallbackHandler): def __init__(self, queue): self.queue = queue def on_llm_new_token(self, token, **kwargs): self.queue.put(token) def on_llm_end(self, response, **kwargs): self.queue.put(None) # 标记流式结束 @app.post("/qa-stream") async def qa_stream(question: str, docs: list[dict]): # 转换文档格式 langchain_docs = [Document(page_content=doc["content"]) for doc in docs] # 创建队列和handler from asyncio import Queue queue = Queue() handler = StreamingResponseHandler(queue) # 初始化LLM和chain llm = ChatOpenAI(streaming=True, callbacks=[handler], temperature=0) chain = load_qa_chain(llm, chain_type="stuff") # 异步执行chain import asyncio asyncio.create_task(chain.arun(input_documents=langchain_docs, question=question)) # 生成流式响应 async def generate(): while True: token = await queue.get() if token is None: break yield token return StreamingResponse(generate(), media_type="text/plain")
如果使用Flask,可通过Response结合stream_with_context实现类似逻辑。
二、其他提升LLM响应速度的方法
- 优化Chain类型:
stuff模式会将所有文档拼接进Prompt,当文档数量多、内容长时,Prompt Token量巨大,推理速度变慢。可改用map_reduce模式,并行处理每个文档的摘要,再合并结果,能有效减少单轮Prompt长度,提升速度;若文档相关性差异大,也可尝试map_rerank模式,只保留排名靠前的相关片段。 - 减少输入Token量:
- 缩小文档分块的大小,避免无关内容进入Prompt;
- 优化检索逻辑,只返回与问题最相关的Top-N个文档片段,而非全部文档;
- 对文档做预处理,过滤冗余内容(如重复段落、无关格式标记)。
- 更换轻量模型:优先选择小参数模型,如
gpt-3.5-turbo比gpt-4响应速度快2-3倍;若使用开源模型,选择7B/13B量级的量化模型(如4bit量化的Llama 2),推理速度远快于大参数模型。 - 启用缓存机制:使用LangChain的缓存组件(如
InMemoryCache、RedisCache)缓存重复的Prompt请求,避免重复调用LLM。示例:from langchain.cache import InMemoryCache langchain.llm_cache = InMemoryCache() - 优化推理环境:
- 若使用开源模型,用GPU加速推理(如CUDA),或使用vLLM、TensorRT-LLM等推理引擎,这些引擎支持连续批处理和优化的推理路径,能大幅提升吞吐量和响应速度;
- 若使用云服务商的LLM API,选择距离更近的区域部署,减少网络延迟。
- 批量处理请求:如果有多个QA请求,将它们批量发送给LLM,减少网络往返次数,提升整体处理效率。
内容的提问来源于stack exchange,提问作者user22403491
相关产品推荐
相关产品推荐

