[dev] change dummy_generation to dashscope_generation

This commit is contained in:
jinli.yl 2024-07-11 15:44:42 +08:00
parent 0e07add988
commit 8282350ab4
15 changed files with 80 additions and 74 deletions

View file

@ -6,7 +6,7 @@ memory_chat:
cli_memory_chat:
class: chat.cli_memory_chat
memory_service: memory_chat_service
generation_model: dummy_generation
generation_model: dashscope_generation
memory_service:
memory_chat_service:
@ -19,8 +19,7 @@ memory_service:
description: "read session messages of the user"
read_memory:
class: memory.operation.read_memory
# workflow: set_query,retrieve_memory1,[extract_time|semantic_rank],fuse_rerank
workflow: dummy
workflow: set_query,[extract_time|retrieve_memory1,semantic_rank],fuse_rerank
description: "read related memories of the user"
list_memory:
class: memory.operation.read_memory
@ -56,6 +55,7 @@ worker:
generation_model_top_k: 1
semantic_rank:
class: memory.worker.read.semantic_rank_worker
rank_model: dashscope_rank
fuse_rerank:
class: memory.worker.read.fuse_rerank_worker
fuse_score_threshold: 0.1
@ -142,7 +142,7 @@ models:
model_name: dummy_generation_model
memory_store:
class: storage.llama_index_es_memory_store
class: storage.llama_index_es_memory_store_sync
embedding_model: dashscope_embedding
index_name: memory_index
es_url: http://localhost:9200

View file

@ -154,11 +154,12 @@ class CliMemoryChat(BaseMemoryChat):
if refresh_time and refresh_time.isdigit():
refresh_time = int(refresh_time)
while True:
time.sleep(refresh_time)
result = self.memory_service.do_operation(op_name=command, **kwargs)
os.system("clear")
self.print_logo()
questionary.print(result)
time.sleep(refresh_time)
else:
result = self.memory_service.do_operation(op_name=command, **kwargs)
questionary.print(result)

View file

@ -1,3 +1,4 @@
import datetime
import sys
sys.path.append(".") # noqa: E402
@ -22,7 +23,8 @@ class MemoryScope(object):
def __init__(self):
self.config: Dict[str, Any] = {}
self.logger: Logger = Logger.get_logger("cli_job", to_stream=False)
datetime_suffix = datetime.datetime.now().strftime('%Y%m%d_%H%M%S')
self.logger: Logger = Logger.get_logger(f"cli_job_{datetime_suffix}", to_stream=False)
def load_config(self, path: str):
with open(path) as f:
@ -36,7 +38,8 @@ class MemoryScope(object):
atexit.register(self.shutdown) # register clean up function
return self
def shutdown(self):
@staticmethod
def shutdown():
print('Gracefully executing the shutdown function...')
G_CONTEXT.memory_store.close()
G_CONTEXT.monitor.close()
@ -68,9 +71,8 @@ class MemoryScope(object):
if "memory_store" not in self.config:
raise RuntimeError("memory_store config is required!")
memory_store_config = self.config["memory_store"]
# embedding_model = G_CONTEXT.model_dict[memory_store_config[ModelEnum.EMBEDDING_MODEL.value]]
embedding_model_conf = self.config["models"][memory_store_config[ModelEnum.EMBEDDING_MODEL.value]]
G_CONTEXT.memory_store = init_instance_by_config(memory_store_config, embedding_model_conf=embedding_model_conf)
embedding_model = G_CONTEXT.model_dict[memory_store_config[ModelEnum.EMBEDDING_MODEL.value]]
G_CONTEXT.memory_store = init_instance_by_config(memory_store_config, embedding_model=embedding_model)
# init monitor
G_CONTEXT.monitor = init_instance_by_config(self.config["monitor"])
@ -78,10 +80,12 @@ class MemoryScope(object):
# set worker config
G_CONTEXT.worker_config = self.config["worker"]
def get_default_service(self):
@property
def default_service(self):
return list(G_CONTEXT.memory_service_dict.values())[0]
def get_default_chat_handle(self):
@property
def default_chat_handle(self):
return list(G_CONTEXT.memory_chat_dict.values())[0]

View file

@ -25,10 +25,10 @@ class ReadMemory(BaseWorkflow, BaseOperation):
self.init_workers(**kwargs)
def run_operation(self, **kwargs):
self.context.clear()
max_count = 1 + self.his_msg_count
self.context[CHAT_MESSAGES] = [x.copy(deep=True) for x in self.chat_messages[-max_count:]]
self.context[CHAT_KWARGS] = kwargs
self.run_workflow()
result = self.context.get(RESULT)
self.context.clear()
return result

View file

@ -11,7 +11,7 @@ class ReadMessage(BaseOperation):
name: str,
description: str,
chat_messages: List[Message],
contextual_msg_count: int = 6, # for the current context dialogue
contextual_msg_count: int, # for the current context dialogue
**kwargs):
super().__init__(name=name, description=description, **kwargs)
self.chat_messages: List[Message] = chat_messages

View file

@ -32,14 +32,18 @@ class BaseWorker(metaclass=ABCMeta):
self.logger: Logger = Logger.get_logger()
def submit_async_task(self, fn, *args, **kwargs):
# if self.is_multi_thread:
# raise RuntimeError(f"async_task is not allowed in multi_thread condition")
if self.is_multi_thread:
raise RuntimeError(f"async_task is not allowed in multi_thread condition")
self.async_task_list.append((fn, args, kwargs))
async def _async_gather(self):
return await asyncio.gather(*[fn(*args, **kwargs) for fn, args, kwargs in self.async_task_list])
def gather_async_result(self):
if self.is_multi_thread:
raise RuntimeError(f"async_task is not allowed in multi_thread condition")
results = asyncio.run(self._async_gather())
self.async_task_list.clear()
return results
@ -50,6 +54,7 @@ class BaseWorker(metaclass=ABCMeta):
def gather_thread_result(self):
for future in as_completed(self.thread_task_list):
yield future.result()
self.thread_task_list.clear()
@abstractmethod
def _run(self):

View file

@ -95,6 +95,8 @@ class MemoryBaseWorker(BaseWorker, metaclass=ABCMeta):
if node.memory_id in self.contex_memory_dict:
continue
self.contex_memory_dict[node.memory_id] = node
self.logger.info(f"add to memory context memory id={node.memory_id} content={node.content} "
f"status={node.status}")
self.set_context(key, [n.memory_id for n in nodes])
def save_memories(self, keys: str | List[str] = None):

View file

@ -40,7 +40,7 @@ class FuseRerankWorker(MemoryBaseWorker):
def _run(self):
# parse input
extract_time_dict: Dict[str, str] = self.get_context(EXTRACT_TIME_DICT)
memory_node_list: List[MemoryNode] = self.get_context(RANKED_MEMORY_NODES)
memory_node_list: List[MemoryNode] = self.get_memories(RANKED_MEMORY_NODES)
if not memory_node_list:
self.logger.warning(f"ranked memory nodes is empty!")
return

View file

@ -26,7 +26,7 @@ class PrintMemoryWorker(MemoryBaseWorker):
line = f" {i} {node.content}"
expired_content_list.append(line)
if MemoryTypeEnum(node.memory_type) in [MemoryTypeEnum.OBSERVATION, MemoryTypeEnum.OBS_CUSTOMIZED]:
elif MemoryTypeEnum(node.memory_type) in [MemoryTypeEnum.OBSERVATION, MemoryTypeEnum.OBS_CUSTOMIZED]:
j += 1
dt_handler = DatetimeHandler(node.timestamp)
dt = dt_handler.datetime_format("%Y%m%d %H:%M:%S")

View file

@ -11,7 +11,7 @@ from memory_scope.utils.timer import timer
class RetrieveMemoryWorker(MemoryBaseWorker):
@timer
async def retrieve_from_observation(self, query: str) -> List[MemoryNode]:
def retrieve_from_observation(self, query: str) -> List[MemoryNode]:
if not self.retrieve_obs_top_k:
return []
@ -22,11 +22,11 @@ class RetrieveMemoryWorker(MemoryBaseWorker):
"memory_type": [MemoryTypeEnum.OBSERVATION.value, MemoryTypeEnum.OBS_CUSTOMIZED.value],
}
return self.memory_store.retrieve_memories(query=query,
top_k=self.retrieve_obs_top_k,
filter_dict=filter_dict)
top_k=self.retrieve_obs_top_k,
filter_dict=filter_dict)
@timer
async def retrieve_from_insight_and_profile(self, query: str) -> List[MemoryNode]:
def retrieve_from_insight_and_profile(self, query: str) -> List[MemoryNode]:
if not self.retrieve_ins_pf_top_k:
return []
@ -37,11 +37,11 @@ class RetrieveMemoryWorker(MemoryBaseWorker):
"memory_type": MemoryTypeEnum.INSIGHT.value,
}
return self.memory_store.retrieve_memories(query=query,
top_k=self.retrieve_ins_pf_top_k,
filter_dict=filter_dict)
top_k=self.retrieve_ins_pf_top_k,
filter_dict=filter_dict)
@timer
async def retrieve_expired_memory(self, query: str) -> List[MemoryNode]:
def retrieve_expired_memory(self, query: str) -> List[MemoryNode]:
if not self.retrieve_expired_top_k:
return []
@ -52,23 +52,23 @@ class RetrieveMemoryWorker(MemoryBaseWorker):
"memory_type": [MemoryTypeEnum.OBSERVATION.value, MemoryTypeEnum.OBS_CUSTOMIZED.value],
}
return self.memory_store.retrieve_memories(query=query,
top_k=self.retrieve_expired_top_k,
filter_dict=filter_dict)
top_k=self.retrieve_expired_top_k,
filter_dict=filter_dict)
def _run(self):
query, _ = self.get_context(QUERY_WITH_TS)
self.submit_async_task(self.retrieve_from_observation, query=query)
self.submit_async_task(self.retrieve_from_insight_and_profile, query=query)
self.submit_async_task(self.retrieve_expired_memory, query=query)
self.submit_thread_task(self.retrieve_from_observation, query=query)
self.submit_thread_task(self.retrieve_from_insight_and_profile, query=query)
self.submit_thread_task(self.retrieve_expired_memory, query=query)
memory_node_list: List[MemoryNode] = []
for result in self.gather_async_result():
for result in self.gather_thread_result():
if result:
memory_node_list.extend(result)
self.logger.info(f"memory_node_list.size={len(memory_node_list)}")
memory_node_list = sorted(memory_node_list, key=lambda x: x.score_similar, reverse=True)
for node in memory_node_list:
self.logger.info(f"recall_stage: content={node.content} score={node.score_similar} type={node.memory_type}"
f"status={node.status}")
self.logger.info(f"recall_stage: content={node.content} score={node.score_similar} "
f"type={node.memory_type} status={node.status}")
self.set_memories(RETRIEVE_MEMORY_NODES, memory_node_list)

View file

@ -13,7 +13,7 @@ from memory_scope.utils.timer import timer
class LoadMemoryWorker(MemoryBaseWorker):
@timer
async def retrieve_not_reflected_memory(self, query: str):
def retrieve_not_reflected_memory(self, query: str):
if not self.retrieve_not_reflected_top_k:
return
@ -24,13 +24,13 @@ class LoadMemoryWorker(MemoryBaseWorker):
"memory_type": [MemoryTypeEnum.OBSERVATION.value, MemoryTypeEnum.OBS_CUSTOMIZED.value],
"obs_reflected": False,
}
nodes: List[MemoryNode] = await self.memory_store.a_retrieve_memories(query=query,
nodes: List[MemoryNode] = self.memory_store.retrieve_memories(query=query,
top_k=self.retrieve_not_reflected_top_k,
filter_dict=filter_dict)
self.set_memories(NOT_REFLECTED_NODES, nodes)
@timer
async def retrieve_not_updated_memory(self, query: str):
def retrieve_not_updated_memory(self, query: str):
if not self.retrieve_not_updated_top_k:
return
@ -41,13 +41,13 @@ class LoadMemoryWorker(MemoryBaseWorker):
"memory_type": [MemoryTypeEnum.OBSERVATION.value, MemoryTypeEnum.OBS_CUSTOMIZED.value],
"obs_updated": False,
}
nodes: List[MemoryNode] = await self.memory_store.a_retrieve_memories(query=query,
nodes: List[MemoryNode] = self.memory_store.retrieve_memories(query=query,
top_k=self.retrieve_not_updated_top_k,
filter_dict=filter_dict)
self.set_memories(NOT_UPDATED_NODES, nodes)
@timer
async def retrieve_insight_memory(self, query: str):
def retrieve_insight_memory(self, query: str):
if not self.retrieve_insight_top_k:
return
@ -57,13 +57,13 @@ class LoadMemoryWorker(MemoryBaseWorker):
"status": MemoryNodeStatus.ACTIVE.value,
"memory_type": MemoryTypeEnum.INSIGHT.value,
}
nodes: List[MemoryNode] = await self.memory_store.a_retrieve_memories(query=query,
nodes: List[MemoryNode] = self.memory_store.retrieve_memories(query=query,
top_k=self.retrieve_insight_top_k,
filter_dict=filter_dict)
self.set_memories(INSIGHT_NODES, nodes)
@timer
async def retrieve_today_memory(self):
def retrieve_today_memory(self):
if not self.today_obs_top_k:
return
@ -80,7 +80,7 @@ class LoadMemoryWorker(MemoryBaseWorker):
"memory_type": [MemoryTypeEnum.OBSERVATION.value, MemoryTypeEnum.OBS_CUSTOMIZED.value],
"dt": dt_handler.datetime_format(),
}
nodes: List[MemoryNode] = await self.memory_store.a_retrieve_memories(query=message.content,
nodes: List[MemoryNode] = self.memory_store.retrieve_memories(query=message.content,
top_k=self.today_obs_top_k,
filter_dict=filter_dict)
@ -88,14 +88,8 @@ class LoadMemoryWorker(MemoryBaseWorker):
def _run(self):
mock_query = "-"
# self.submit_thread_task(self.retrieve_not_reflected_memory, query=mock_query)
# self.submit_thread_task(self.retrieve_not_updated_memory, query=mock_query)
# self.submit_thread_task(self.retrieve_insight_memory, query=mock_query)
# self.submit_thread_task(self.retrieve_today_memory)
# self.gather_thread_result()
self.submit_async_task(self.retrieve_not_reflected_memory, query=mock_query)
self.submit_async_task(self.retrieve_not_updated_memory, query=mock_query)
self.submit_async_task(self.retrieve_insight_memory, query=mock_query)
self.submit_async_task(self.retrieve_today_memory)
self.gather_async_result()
self.submit_thread_task(self.retrieve_not_reflected_memory, query=mock_query)
self.submit_thread_task(self.retrieve_not_updated_memory, query=mock_query)
self.submit_thread_task(self.retrieve_insight_memory, query=mock_query)
self.submit_thread_task(self.retrieve_today_memory)
self.gather_thread_result()

View file

@ -2,15 +2,15 @@ from typing import Dict, List, Any, Optional, cast
from llama_index.core import VectorStoreIndex
from llama_index.core.schema import TextNode, NodeWithScore
from llama_index.vector_stores.elasticsearch import ElasticsearchStore, AsyncDenseVectorStrategy
from llama_index.vector_stores.elasticsearch import AsyncDenseVectorStrategy
from memory_scope.enumeration.memory_status_enum import MemoryNodeStatus
from memory_scope.models.base_model import BaseModel
from memory_scope.scheme.memory_node import MemoryNode
from memory_scope.storage.base_memory_store import BaseMemoryStore
from memory_scope.storage.llama_index_sync_elasticsearch import SyncElasticsearchStore
from memory_scope.utils.logger import Logger
from memory_scope.storage.llama_index_sync_elasticsearch import SyncElasticsearchStore
class _AsyncDenseVectorStrategy(AsyncDenseVectorStrategy):
def _hybrid(self, query: str, knn: Dict[str, Any], filter: List[Dict[str, Any]], top_k: int) -> Dict[str, Any]:
@ -124,7 +124,7 @@ def _to_elasticsearch_filter(standard_filters: Dict[str, List[str]]) -> Dict[str
return result
class LlamaIndexEsMemoryStore(BaseMemoryStore):
class LlamaIndexEsMemoryStoreSync(BaseMemoryStore):
def __init__(self,
embedding_model: BaseModel,
index_name: str,
@ -134,11 +134,13 @@ class LlamaIndexEsMemoryStore(BaseMemoryStore):
self.embedding_model: BaseModel = embedding_model
self.es_store = SyncElasticsearchStore(index_name=index_name,
es_url=es_url,
retrieval_strategy=_AsyncDenseVectorStrategy(hybrid=use_hybrid),
**kwargs)
self.index = VectorStoreIndex.from_vector_store(vector_store=self.es_store,
embed_model=self.embedding_model.model)
es_url=es_url,
retrieval_strategy=_AsyncDenseVectorStrategy(hybrid=use_hybrid),
**kwargs)
# use /dev/null
with open(os.devnull, 'w') as devnull:
self.index = VectorStoreIndex.from_vector_store(vector_store=self.es_store,
embed_model=self.embedding_model.model)
self.index.build_index_from_nodes([TextNode(text="text")])
self.logger = Logger.get_logger()
@ -192,13 +194,14 @@ class LlamaIndexEsMemoryStore(BaseMemoryStore):
if not nodes:
self.logger.warning("empty nodes!")
return
self.logger.info(f"update_memories nodes={nodes}")
if isinstance(nodes, MemoryNode):
nodes = [nodes]
# emb & insert new memories
# TODO batch insert
new_memories = [n for n in nodes if n.status == MemoryNodeStatus.NEW]
new_memories = [n for n in nodes if n.status == MemoryNodeStatus.NEW.value]
if new_memories:
for n in new_memories:
n.status = MemoryNodeStatus.ACTIVE.value
@ -206,7 +209,7 @@ class LlamaIndexEsMemoryStore(BaseMemoryStore):
# emb & update new memories
# TODO insert overwrite
c_modified_memories = [n for n in nodes if n.status == MemoryNodeStatus.CONTENT_MODIFIED]
c_modified_memories = [n for n in nodes if n.status == MemoryNodeStatus.CONTENT_MODIFIED.value]
if c_modified_memories:
for n in c_modified_memories:
n.status = MemoryNodeStatus.ACTIVE.value
@ -215,7 +218,7 @@ class LlamaIndexEsMemoryStore(BaseMemoryStore):
# update new memories
# TODO no emb
modified_memories = [n for n in nodes if n.status == MemoryNodeStatus.MODIFIED]
modified_memories = [n for n in nodes if n.status == MemoryNodeStatus.MODIFIED.value]
if modified_memories:
for n in modified_memories:
n.status = MemoryNodeStatus.ACTIVE.value
@ -223,10 +226,9 @@ class LlamaIndexEsMemoryStore(BaseMemoryStore):
self.insert(n)
# set memories expired
expired_memories = [n for n in nodes if n.status == MemoryNodeStatus.EXPIRED]
expired_memories = [n for n in nodes if n.status == MemoryNodeStatus.EXPIRED.value]
if expired_memories:
for n in expired_memories:
n.status = MemoryNodeStatus.ACTIVE.value
self.delete(n)
self.insert(n)

View file

@ -1,11 +1,9 @@
import init_test
from memory_scope.cli import MemoryScope
from memory_scope.enumeration.message_role_enum import MessageRoleEnum
from memory_scope.scheme.message import Message
ms = MemoryScope().load_config("config/demo_config_no_stream.yaml")
memory_service = ms.get_default_service()
memory_chat = ms.get_default_chat_handle()
memory_service = ms.default_service
memory_chat = ms.default_chat_handle
# new_message: Message = Message(role=MessageRoleEnum.USER.value, role_name="我", content="我的爱好是弹琴并且喜欢看电影。")
# memory_service.add_messages(new_message)

View file

@ -2,7 +2,8 @@ import unittest
from memory_scope.models.llama_index_embedding_model import LlamaIndexEmbeddingModel
from memory_scope.scheme.memory_node import MemoryNode
from memory_scope.storage.llama_index_es_memory_store_sync import LlamaIndexEsMemoryStore
from memory_scope.storage.llama_index_es_memory_store_sync import LlamaIndexEsMemoryStoreSync
class TestLlamaIndexElasticSearchStore(unittest.TestCase):
"""Tests for LLIEmbedding"""
@ -22,7 +23,7 @@ class TestLlamaIndexElasticSearchStore(unittest.TestCase):
"use_hybrid": True
}
self.es_store = LlamaIndexEsMemoryStore(**config)
self.es_store = LlamaIndexEsMemoryStoreSync(**config)
self.data = [
MemoryNode(
content="The lives of two mob hitmen, a boxer, a gangster and his wife, "

View file

@ -2,13 +2,12 @@ import sys
sys.path.append(".")
import asyncio
from concurrent.futures import ThreadPoolExecutor
from memory_scope.models.llama_index_embedding_model import LlamaIndexEmbeddingModel
from memory_scope.storage.llama_index_es_memory_store import LlamaIndexEsMemoryStore
from memory_scope.storage.llama_index_es_memory_store_sync import LlamaIndexEsMemoryStore as SyncLlamaIndexEsMemoryStore
from memory_scope.storage.llama_index_es_memory_store_sync import \
LlamaIndexEsMemoryStoreSync as SyncLlamaIndexEsMemoryStore
from memory_scope.utils.logger import Logger
logger = Logger.get_logger("default")