RAG Pipeline内存泄漏:Memo AI上下文切换后向量嵌入未释放
RAG架构记忆增强AI系统的内存泄漏与会话隔离问题
问题现象
- 上下文切换后向量嵌入未被垃圾回收,内存占用持续上升(从初始500MB增长至2.3GB且无下降)
- 前会话的嵌入内容泄露到新对话中:即使设置了
session_id过滤,新会话仍能检索到旧会话的内容 - FAISS索引在约1000次检索后性能下降
当前实现代码
class MemoAI: def __init__(self): self.vector_store = FAISS.load_local("./embeddings", embeddings) self.memory_buffer = ConversationSummaryBufferMemory( llm=llm, max_token_limit=2000 ) def add_memory(self, text, metadata): chunks = self.recursive_splitter.split_text(text) embeddings = self.embedder.embed_documents(chunks) # 问题:这些嵌入在会话结束后仍持续存在 self.vector_store.add_embeddings( [(chunk, embedding) for chunk, embedding in zip(chunks, embeddings)], metadatas=[metadata] * len(chunks) ) def retrieve_context(self, query, k=5): # 问题:返回来自前会话的陈旧片段 return self.vector_store.similarity_search_with_score( query, k=k, filter={"session_id": self.current_session} )
可复现步骤
# Session 1 memo_ai.add_memory("User likes Python", {"session_id": "session_1"})
# Session 2 (新用户) memo_ai.switch_session("session_2") result = memo_ai.retrieve_context("What programming language?") # BUG:尽管设置了过滤,仍返回session_1中的"likes Python" # 内存占用:从初始500MB增长至2.3GB且持续上升
已尝试方案
- 使用
del self.vector_store和gc.collect()手动清理,但内存未释放 - 为每个会话创建独立FAISS索引,但实时性不足
- 元数据过滤,但结果不一致
环境信息
- LangChain 0.1.0、FAISS-GPU 1.7.2、Python 3.10
- 硬件:32GB内存、RTX 3090显卡
解决方案
一、正确实现会话隔离(无需重建整个向量存储)
修复元数据过滤逻辑
LangChain 0.1.0中FAISS的filter参数依赖底层元数据索引,需在初始化时显式指定要索引的元数据字段,确保session_id被正确识别:# 初始化时指定元数据索引字段 self.vector_store = FAISS.load_local( "./embeddings", embeddings, index_metadata_fields=["session_id"] # 显式索引session_id )同时检查
switch_session方法是否正确赋值self.current_session,避免未初始化或赋值错误。全局+会话分区的存储设计
将通用知识与会话专属内容分离存储,既保证全局知识复用,又实现会话隔离:class MemoAI: def __init__(self): # 全局通用向量存储(长期保留) self.global_store = FAISS.load_local("./global_embeddings", embeddings) # 会话临时存储容器(键为session_id,值为FAISS索引) self.session_stores = {} self.current_session = None self.memory_buffer = ConversationSummaryBufferMemory(llm=llm, max_token_limit=2000) def switch_session(self, session_id): self.current_session = session_id # 会话不存在则创建临时存储 if session_id not in self.session_stores: self.session_stores[session_id] = FAISS.from_texts([], embeddings) def add_memory(self, text, metadata): if not self.current_session: raise ValueError("No active session") chunks = self.recursive_splitter.split_text(text) embeddings = self.embedder.embed_documents(chunks) # 将会话专属嵌入添加到会话临时存储 self.session_stores[self.current_session].add_embeddings( [(chunk, embedding) for chunk, embedding in zip(chunks, embeddings)], metadatas=[metadata] * len(chunks) ) def retrieve_context(self, query, k=5): if not self.current_session: raise ValueError("No active session") # 合并全局存储和会话存储的检索结果 global_results = self.global_store.similarity_search_with_score(query, k=k) session_results = self.session_stores[self.current_session].similarity_search_with_score(query, k=k) # 按相似度排序后返回前k条 all_results = global_results + session_results all_results.sort(key=lambda x: x[1]) return all_results[:k]
二、生产级嵌入垃圾回收模式
GPU内存手动释放
FAISS-GPU的内存不会被Python GC自动回收,需调用原生API+CUDA缓存清理:def cleanup_session(self, session_id): if session_id in self.session_stores: # 释放FAISS索引的GPU内存 self.session_stores[session_id].index.reset() del self.session_stores[session_id] # 触发Python GC并清理CUDA缓存 import gc gc.collect() import torch torch.cuda.empty_cache()定时过期清理机制
为会话设置超时时间,定期扫描并清理过期会话的临时存储:import threading import time class MemoAI: def __init__(self): # ... 其他初始化代码 ... self.session_expiry = 3600 # 会话超时1小时 self.session_last_active = {} # 启动后台定时清理线程 self.cleanup_thread = threading.Thread(target=self._periodic_cleanup, daemon=True) self.cleanup_thread.start() def switch_session(self, session_id): self.current_session = session_id self.session_last_active[session_id] = time.time() # ... 会话存储创建逻辑 ... def _periodic_cleanup(self): while True: current_time = time.time() # 筛选过期会话 expired_sessions = [sid for sid, last_time in self.session_last_active.items() if current_time - last_time > self.session_expiry] for sid in expired_sessions: self.cleanup_session(sid) del self.session_last_active[sid] time.sleep(300) # 每5分钟执行一次清理FAISS索引性能优化
针对检索次数过多后的性能下降,定期对索引进行重构优化:def optimize_index(self, session_id=None): # 优化全局索引 self.global_store.index.train(self.global_store.index.xb) self.global_store.index.reset() # 优化指定会话索引 if session_id and session_id in self.session_stores: self.session_stores[session_id].index.train(self.session_stores[session_id].index.xb) self.session_stores[session_id].index.reset()
内容的提问来源于stack exchange,提问作者TensorMind
相关产品推荐
相关产品推荐

