FastAPI部署带记忆LLM生产环境:多用户请求串答问题解决问询
问题根源
你的代码核心问题是用了全局变量memory和streamer。FastAPI是异步并发框架,多用户请求同时进来时,全局变量会被所有请求共享,导致不同用户的对话记忆串在一起,生成的token也会通过全局streamer混到其他用户的响应里,最终出现答案交叉的情况。
解决方案
要解决这个问题,核心是让每个用户的会话数据完全隔离,具体做这几点:
- 为每个
user_id+session_id维护独立的对话记忆,禁止使用全局变量 - 每个请求生成专属的Streamer实例,避免共享冲突
- 重构函数,移除所有全局状态依赖
修正后的完整代码
from threading import Thread import time from transformers import AutoTokenizer, AutoModelForCausalLM, TextStreamer # 模型和tokenizer全局初始化一次即可,这部分逻辑没问题 tokenizer = AutoTokenizer.from_pretrained("你的模型路径") model = AutoModelForCausalLM.from_pretrained("你的模型路径").to("cuda") # 用字典存储会话记忆,key为user_id+session_id,生产环境建议替换为Redis/数据库 session_memories = {} def get_session_memory(user_id, session_id): """获取指定用户会话的记忆,不存在则初始化""" key = f"{user_id}_{session_id}" if key not in session_memories: session_memories[key] = [] return session_memories[key] def update_memory(question, memory): """原有的记忆更新逻辑,保留即可""" memory.append([question, ""]) # 先占位,后续填充AI回答 def conv_gen(prompt, memory): """原有的对话上下文生成逻辑,保留即可""" conversation = prompt for q, a in memory: conversation += f"用户:{q}\nAI:{a}\n" conversation += f"用户:{memory[-1][0]}\nAI:" return conversation def get_answersf7bsft(user_id, session_id, question, prompt, clean_memory=False, max_new_tokens_=256, temperature_=0.1, top_k_=50, top_p_=.95, typical_p_=1.00, repetition_penalty_=1.2): print("Comienza get_answersf7bsft") start = time.time() # 获取当前会话的专属记忆,不再依赖全局memory memory = get_session_memory(user_id, session_id) if clean_memory: memory.clear() # 仅清空当前会话的记忆 update_memory(question, memory) conversation = conv_gen(prompt, memory) inputs = tokenizer(conversation, return_tensors="pt").to("cuda")["input_ids"] # 每个请求创建独立的Streamer,彻底避免共享冲突 streamer = TextStreamer(tokenizer) generation_kwargs = dict(input_ids=inputs, pad_token_id=tokenizer.eos_token_id, streamer=streamer, do_sample=True, max_new_tokens=max_new_tokens_, temperature=temperature_, top_k=top_k_, top_p=top_p_, typical_p=typical_p_, repetition_penalty=repetition_penalty_, bad_words_ids=[[5150], [12453]]) thread = Thread(target=model.generate, kwargs=generation_kwargs) answer = "" thread.start() for _ in streamer: _ = _.replace("<|endoftext|>", "") answer = answer + _ yield _ print("Finaliza respuesta") # 更新当前会话记忆中的AI回答 memory[-1][1] = answer
生产环境额外优化
- 替换记忆存储:上面的
session_memories是内存字典,服务重启会丢失数据,多实例部署也无法共享。生产环境建议用Redis(设置过期时间自动清理闲置会话)、MongoDB或MySQL存储。 - 参数校验:在FastAPI接口层给
temperature_(比如限制0-2)、max_new_tokens_(比如限制10-2048)等参数加范围校验,防止非法参数导致模型异常。 - 会话清理:给每个会话设置过期时间,比如用户30分钟未操作就删除记忆,避免存储占用过高。
- 异步优化:如果模型支持异步生成,可改用
asyncio替代Thread,更贴合FastAPI的异步特性;非流式场景也可以用Celery做异步任务处理。
内容的提问来源于stack exchange,提问作者conejo
相关产品推荐
相关产品推荐

