-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathmemory_summarizer.py
More file actions
84 lines (76 loc) · 3.79 KB
/
Copy pathmemory_summarizer.py
File metadata and controls
84 lines (76 loc) · 3.79 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
import time
import logging
import os
import cachetools.func
from dotenv import load_dotenv
from typing import Any, Dict, List
from document_summarizer import FlexibleDocumentSummarizer
from langchain_openai import ChatOpenAI
from langchain_qdrant import Qdrant
from qdrant_client.http import models as rest
from qdrant_client.http.models import PayloadSchemaType
from langchain.retrievers import ContextualCompressionRetriever
from qdrant_retriever import QDrantVectorStoreRetriever
from cohere_rerank import CohereRerank
from langchain_openai import OpenAIEmbeddings
from generative_conversation_summarized_memory import GenerativeAgentConversationSummarizedMemory
class MemorySummarizer:
flexible_document_summarizer: FlexibleDocumentSummarizer
def __init__(self, rate_limiter, rate_limiter_sync, flexible_document_summarizer, agent_manager):
load_dotenv() # Load environment variables
os.getenv("COHERE_API_KEY")
self.QDRANT_API_KEY = os.getenv("QDRANT_API_KEY")
self.QDRANT_URL = os.getenv("QDRANT_URL")
self.agent_manager = agent_manager
self.rate_limiter = rate_limiter
self.rate_limiter_sync = rate_limiter_sync
self.flexible_document_summarizer = flexible_document_summarizer
def create_new_conversation_summarizer(self, api_key: str, user_id: str):
"""Create a new vector store retriever unique to the agent."""
collection_name = f"{user_id}_summaries"
# create collection if it doesn't exist
try:
self.agent_manager.client.create_collection(
collection_name=collection_name,
vectors_config=rest.VectorParams(
size=1536,
distance=rest.Distance.COSINE,
),
)
self.agent_manager.client.create_payload_index(
collection_name, "metadata.extra_index", field_schema=PayloadSchemaType.KEYWORD)
except:
print("MemorySummarizer: loaded from cloud...")
finally:
logging.info(
f"MemorySummarizer: Creating memory store with collection {collection_name}")
vectorstore = Qdrant(self.agent_manager.client, collection_name, OpenAIEmbeddings(
model="text-embedding-3-small", openai_api_key=api_key))
compressor = CohereRerank()
compression_retriever = ContextualCompressionRetriever(
base_compressor=compressor, base_retriever=QDrantVectorStoreRetriever(
rate_limiter=self.rate_limiter, rate_limiter_sync=self.rate_limiter_sync, collection_name=collection_name, client=self.agent_manager.client, vectorstore=vectorstore,
)
)
return compression_retriever
def create_summarized_memory(self, api_key: str, user_id: str):
return GenerativeAgentConversationSummarizedMemory(
rate_limiter=self.rate_limiter,
llm=ChatOpenAI(openai_api_key=api_key, temperature=0,
max_tokens=2048, model="gpt-4.1-mini"),
memory_retriever=self.create_new_conversation_summarizer(
api_key, user_id),
verbose=self.agent_manager.verbose
)
@cachetools.func.ttl_cache(maxsize=16384, ttl=36000)
def load(self, api_key: str, user_id: str) -> GenerativeAgentConversationSummarizedMemory:
"""Load existing index data from the cloud."""
start = time.time()
retriever = self.create_summarized_memory(api_key, user_id)
end = time.time()
logging.info(
f"MemorySummarizer: Load operation took {end - start} seconds")
return retriever
async def save(self, api_key: str, user_id: str, outputs: Dict[str, Any]) -> List[str]:
memory = self.load(api_key, user_id)
await memory.save_context(outputs)