mirror of
https://github.com/agentscope-ai/ReMe.git
synced 2026-09-12 23:01:15 +00:00
fix grammer bugs
This commit is contained in:
parent
590cbf65c4
commit
a5b45afa8c
29 changed files with 157 additions and 351 deletions
|
|
@ -1,148 +0,0 @@
|
|||
global_config:
|
||||
language: cn
|
||||
max_workers: 5
|
||||
|
||||
memory_chat:
|
||||
cli_memory_chat:
|
||||
class: chat.cli_memory_chat
|
||||
stream: false
|
||||
memory_service: memory_chat_service
|
||||
generation_model: dashscope_generation
|
||||
|
||||
memory_service:
|
||||
memory_chat_service:
|
||||
class: memory.service.chat_memory_service
|
||||
history_msg_count: 32
|
||||
contextual_msg_count: 6
|
||||
memory_operations:
|
||||
read_message:
|
||||
class: memory.operation.read_message
|
||||
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
|
||||
description: "read related memories of the user"
|
||||
list_memory:
|
||||
class: memory.operation.read_memory
|
||||
workflow: set_query,retrieve_memory2,print_memory
|
||||
description: "read all memories of the user"
|
||||
write_memory:
|
||||
class: memory.operation.write_memory
|
||||
workflow: info_filter,load_memory1,[get_observation|get_observation_with_time],contra_repeat,store_memory
|
||||
description: "write observation memories of the user"
|
||||
interval_time: 5
|
||||
# summary_memory:
|
||||
# class: memory.operation.summary_memory
|
||||
# workflow: load_memory2,get_reflection_subject,update_insight,long_contra_repeat,store_memory
|
||||
# description: "summary observation memories of the user"
|
||||
# interval_time: 60
|
||||
|
||||
worker:
|
||||
dummy:
|
||||
class: memory.worker.dummy_worker
|
||||
generation_model: dashscope_generation
|
||||
embedding_model: dashscope_embedding
|
||||
rank_model: dashscope_rank
|
||||
set_query:
|
||||
class: memory.worker.read.set_query_worker
|
||||
retrieve_memory1:
|
||||
class: memory.worker.read.retrieve_memory_worker
|
||||
retrieve_obs_top_k: 100
|
||||
retrieve_ins_pf_top_k: 100
|
||||
retrieve_expired_top_k: 0
|
||||
extract_time:
|
||||
class: memory.worker.read.extract_time_worker
|
||||
generation_model: dashscope_generation
|
||||
generation_model_top_k: 1
|
||||
semantic_rank:
|
||||
class: memory.worker.read.semantic_rank_worker
|
||||
fuse_rerank:
|
||||
class: memory.worker.read.fuse_rerank_worker
|
||||
fuse_score_threshold: 0.1
|
||||
fuse_ratio_dict:
|
||||
conversation: 0.5
|
||||
observation: 1
|
||||
obs_customized: 1.2
|
||||
insight: 2.0
|
||||
fuse_time_ratio: 2.0
|
||||
fuse_rerank_top_k: 10
|
||||
retrieve_memory2:
|
||||
class: memory.worker.read.retrieve_memory_worker
|
||||
retrieve_obs_top_k: 100
|
||||
retrieve_ins_pf_top_k: 100
|
||||
retrieve_expired_top_k: 100
|
||||
print_memory:
|
||||
class: memory.worker.read.print_memory_worker
|
||||
info_filter:
|
||||
class: memory.worker.write.info_filter_worker
|
||||
generation_model: dashscope_generation
|
||||
info_filter_msg_max_size: 200
|
||||
generation_model_top_k: 1
|
||||
load_memory1:
|
||||
class: memory.worker.write.load_memory_worker
|
||||
retrieve_not_reflected_top_k: 0
|
||||
retrieve_not_updated_top_k: 0
|
||||
retrieve_insight_top_k: 0
|
||||
today_obs_top_k: 100
|
||||
get_observation:
|
||||
class: memory.worker.write.get_observation_worker
|
||||
generation_model: dashscope_generation
|
||||
generation_model_top_k: 1
|
||||
get_observation_with_time:
|
||||
class: memory.worker.write.get_observation_with_time_worker
|
||||
generation_model: dashscope_generation
|
||||
generation_model_top_k: 1
|
||||
contra_repeat:
|
||||
class: memory.worker.write.contra_repeat_worker
|
||||
generation_model: dashscope_generation
|
||||
generation_model_top_k: 1
|
||||
retrieve_top_k: 30
|
||||
contra_repeat_max_count: 50
|
||||
store_memory:
|
||||
class: memory.worker.write.store_memory_worker
|
||||
store_key: all
|
||||
load_memory2:
|
||||
class: memory.worker.write.load_memory_worker
|
||||
retrieve_not_reflected_top_k: 100
|
||||
retrieve_not_updated_top_k: 100
|
||||
retrieve_insight_top_k: 100
|
||||
today_obs_top_k: 0
|
||||
get_reflection_subject:
|
||||
class: memory.worker.summary.get_reflection_subject_worker
|
||||
retrieve_top_k: 100
|
||||
reflect_obs_cnt_threshold: 32
|
||||
generation_model_top_k: 1
|
||||
update_insight:
|
||||
class: memory.worker.summary.update_insight_worker
|
||||
update_insight_threshold: 0.1
|
||||
generation_model_top_k: 1
|
||||
update_insight_max_thread: 10
|
||||
long_contra_repeat:
|
||||
class: memory.worker.summary.long_contra_repeat_worker
|
||||
long_contra_repeat_top_k: 2
|
||||
long_contra_repeat_threshold: 0.1
|
||||
generation_model_top_k: 1
|
||||
|
||||
models:
|
||||
dashscope_generation:
|
||||
class: models.llama_index_generation_model
|
||||
module_name: dashscope_generation
|
||||
model_name: qwen-max
|
||||
dashscope_embedding:
|
||||
class: models.llama_index_embedding_model
|
||||
module_name: dashscope_embedding
|
||||
model_name: text-embedding-v2
|
||||
dashscope_rank:
|
||||
class: models.llama_index_rank_model
|
||||
module_name: dashscope_rank
|
||||
model_name: gte-rerank
|
||||
|
||||
memory_store:
|
||||
class: storage.llama_index_es_memory_store
|
||||
embedding_model: dashscope_embedding
|
||||
index_name: memory_index
|
||||
es_url: http://localhost:9200
|
||||
use_hybrid: false
|
||||
|
||||
monitor:
|
||||
class: storage.dummy_monitor
|
||||
|
|
@ -1,83 +0,0 @@
|
|||
global_config:
|
||||
language: cn
|
||||
max_workers: 5
|
||||
dash_scope_apikey:
|
||||
open_ai_apikey:
|
||||
|
||||
memory_chat:
|
||||
cli_memory_chat:
|
||||
class: chat.cli_memory_chat # select class
|
||||
memory_service: memory_chat_service
|
||||
generation_model: dashscope_generation
|
||||
human_name: human
|
||||
assistant_name: assistant
|
||||
|
||||
memory_service:
|
||||
memory_chat_service:
|
||||
class: memory.service.chat_memory_service # select class
|
||||
history_msg_count: 32
|
||||
contextual_msg_count: 6
|
||||
read_memory_key: read_memory
|
||||
memory_operations:
|
||||
read_message: # define operation
|
||||
class: memory.operation.read_memory
|
||||
workflow: dummy_workflow # select workflow
|
||||
description: "read session messages of the user"
|
||||
read_memory:
|
||||
class: memory.operation.read_memory
|
||||
workflow: dummy_workflow
|
||||
description: "read related memories of the user"
|
||||
list_memory:
|
||||
class: memory.operation.read_memory
|
||||
workflow: dummy_workflow
|
||||
description: "read all memories of the user"
|
||||
write_memory:
|
||||
class: memory.operation.write_memory
|
||||
workflow: dummy_workflow
|
||||
description: "write observation memories of the user"
|
||||
interval_time: 60
|
||||
summary_memory:
|
||||
class: memory.operation.summary_memory
|
||||
workflow: dummy_workflow
|
||||
description: "summary observation memories of the user"
|
||||
interval_time: 300
|
||||
|
||||
models:
|
||||
dashscope_generation:
|
||||
class: models.llama_index_generation_model # select class
|
||||
module_name: dashscope_generation
|
||||
model_name: qwen-max
|
||||
dashscope_embedding:
|
||||
class: models.llama_index_embedding_model # select class
|
||||
module_name: dashscope_embedding
|
||||
model_name: text-embedding-v2
|
||||
dashscope_rank:
|
||||
class: models.llama_index_rank_model # select class
|
||||
module_name: dashscope_rank
|
||||
model_name: gte-rerank
|
||||
|
||||
vector_store:
|
||||
class: storage.dummy_vector_store # select class
|
||||
embedding_model: dashscope_embedding
|
||||
|
||||
monitor:
|
||||
class: storage.dummy_monitor # select class
|
||||
|
||||
worker:
|
||||
dummy_workflow:
|
||||
class: memory.worker.dummy_worker
|
||||
generation_model: dashscope_generation
|
||||
embedding_model: dashscope_embedding
|
||||
rank_model: dashscope_rank
|
||||
retrieve_store_worker:
|
||||
class: memory.worker.read.retrieve_store_worker
|
||||
retrieve_obs_top_k: 100
|
||||
retrieve_ins_pf_top_k: 100
|
||||
fuse_rerank_worker:
|
||||
class: memory.worker.read.fuse_rerank_worker
|
||||
fuse_score_threshold: 0.1
|
||||
fuse_ratio_dict:
|
||||
observation: 1
|
||||
fuse_time_ratio: 2.0
|
||||
fuse_rerank_top_k: 10
|
||||
|
||||
0
examples/docker/__init__.py
Normal file
0
examples/docker/__init__.py
Normal file
0
examples/docker/docker_config.yaml
Normal file
0
examples/docker/docker_config.yaml
Normal file
|
|
@ -3,6 +3,7 @@ global_config:
|
|||
thread_pool_max_workers: 5
|
||||
logger_name: memoryscope
|
||||
logger_name_time_suffix: %Y%m%d_%H%M%S
|
||||
use_dummy_ranker: true
|
||||
|
||||
memory_chat:
|
||||
cli_memory_chat:
|
||||
|
|
|
|||
|
|
@ -45,6 +45,10 @@ class MemoryscopeArguments(object):
|
|||
|
||||
embedding_params: dict = field(default_factory=lambda: {})
|
||||
|
||||
use_dummy_ranker: bool = field(default=True, metadata={
|
||||
"help": "If a semantic ranking model is not available, MemoryScope will use cosine similarity scoring as a "
|
||||
"substitute. However, the ranking effectiveness will be somewhat compromised."})
|
||||
|
||||
rank_backend: str = field(default="dashscope_rank", metadata={"help": "global rank backend: dashscope_rank, etc."})
|
||||
|
||||
rank_model: str = field(default="gte-rerank", metadata={"help": "global rank model: gte-rerank, etc."})
|
||||
|
|
@ -58,4 +62,5 @@ class MemoryscopeArguments(object):
|
|||
retrieve_mode: str = field(default="dense", metadata={
|
||||
"help": "retrieve_mode: dense, sparse(not implemented), hybrid(not implemented)"})
|
||||
|
||||
hybrid_alpha: float | None = field(default=1.0, metadata={"help": ""})
|
||||
hybrid_alpha: float | None = field(default=1.0, metadata={
|
||||
"help": "fuse alpha params used in hybrid mode(not implemented)"})
|
||||
|
|
|
|||
|
|
@ -7,8 +7,9 @@ class ModelEnum(str, Enum):
|
|||
|
||||
Members:
|
||||
GENERATION_MODEL: Represents a model responsible for generating content.
|
||||
EMBEDDING_MODEL: Represents a model tasked with creating embeddings, typically used for transforming data into a numerical form suitable for machine learning tasks.
|
||||
RANK_MODEL: Denotes a model that specializes in ranking, often used to order items based on relevance or importance.
|
||||
EMBEDDING_MODEL: Represents a model tasked with creating embeddings, typically used for transforming data into a
|
||||
numerical form suitable for machine learning tasks.
|
||||
RANK_MODEL: Denotes a model that specializes in ranking, often used to order items based on relevance.
|
||||
"""
|
||||
GENERATION_MODEL = "generation_model"
|
||||
|
||||
|
|
|
|||
|
|
@ -12,7 +12,6 @@ class BaseOperation(metaclass=ABCMeta):
|
|||
operation_type (OPERATION_TYPE): Specifies the type of operation, defaulting to "frontend".
|
||||
name (str): The name of the operation.
|
||||
description (str): A description of the operation.
|
||||
kwargs (dict): Additional keyword arguments for operation configuration.
|
||||
"""
|
||||
|
||||
operation_type: OPERATION_TYPE = "frontend"
|
||||
|
|
|
|||
|
|
@ -54,7 +54,7 @@ class BaseMemoryService(metaclass=ABCMeta):
|
|||
def start_backend_service(self):
|
||||
pass
|
||||
|
||||
def stop_backend_service(self):
|
||||
def stop_backend_service(self, wait_service_end: bool = False):
|
||||
pass
|
||||
|
||||
def do_operation(self, name: str, **kwargs):
|
||||
|
|
|
|||
|
|
@ -80,7 +80,7 @@ class ContraRepeatWorker(MemoryBaseWorker):
|
|||
response_text = response.message.content
|
||||
|
||||
# parse text
|
||||
idx_merge_obs_list = ResponseTextParser(response_text, self.__class__.__name__).parse_v1()
|
||||
idx_merge_obs_list = ResponseTextParser(response_text, self.language, self.__class__.__name__).parse_v1()
|
||||
if len(idx_merge_obs_list) <= 0:
|
||||
self.logger.warning("idx_merge_obs_list is empty!")
|
||||
return
|
||||
|
|
|
|||
|
|
@ -26,7 +26,7 @@ class GetObservationWithTimeWorker(GetObservationWorker):
|
|||
filter_messages = []
|
||||
for msg in self.chat_messages:
|
||||
# Checks if the message content has any time reference words
|
||||
if DatetimeHandler.has_time_word(query=msg.content):
|
||||
if DatetimeHandler.has_time_word(query=msg.content, language=self.language):
|
||||
filter_messages.append(msg)
|
||||
return filter_messages
|
||||
|
||||
|
|
@ -49,7 +49,7 @@ class GetObservationWithTimeWorker(GetObservationWorker):
|
|||
for i, msg in enumerate(filter_messages):
|
||||
# Create a DatetimeHandler instance for each message's timestamp and format it
|
||||
dt_handler = DatetimeHandler(dt=msg.time_created)
|
||||
dt = dt_handler.string_format(self.prompt_handler.time_string_format)
|
||||
dt = dt_handler.string_format(string_format=self.prompt_handler.time_string_format, language=self.language)
|
||||
# Append formatted timestamp-query pairs to the user_query_list
|
||||
user_query_list.append(f"{i + 1} {dt} {self.target_name}{self.get_language_value(COLON_WORD)}{msg.content}")
|
||||
|
||||
|
|
|
|||
|
|
@ -139,7 +139,7 @@ class GetObservationWorker(MemoryBaseWorker):
|
|||
response_text = response.message.content
|
||||
|
||||
# Parses the generated text to extract observation indices, times, contents, and keywords
|
||||
idx_obs_list = ResponseTextParser(response_text, self.__class__.__name__).parse_v1()
|
||||
idx_obs_list = ResponseTextParser(response_text, self.language, self.__class__.__name__).parse_v1()
|
||||
if len(idx_obs_list) <= 0:
|
||||
self.logger.warning("idx_obs_list is empty!")
|
||||
return
|
||||
|
|
|
|||
|
|
@ -36,7 +36,7 @@ class GetReflectionSubjectWorker(MemoryBaseWorker):
|
|||
"""
|
||||
dt_handler = DatetimeHandler()
|
||||
# Prepare metadata with current datetime info
|
||||
meta_data = {k: str(v) for k, v in dt_handler.get_dt_info_dict.items()}
|
||||
meta_data = {k: str(v) for k, v in dt_handler.get_dt_info_dict(self.language).items()}
|
||||
|
||||
return MemoryNode(user_name=self.user_name,
|
||||
target_name=self.target_name,
|
||||
|
|
@ -101,7 +101,8 @@ class GetReflectionSubjectWorker(MemoryBaseWorker):
|
|||
return
|
||||
|
||||
# Parse LLM response for new insight keys and update memory
|
||||
new_insight_keys = ResponseTextParser(response.message.content, self.__class__.__name__).parse_v2()
|
||||
new_insight_keys = ResponseTextParser(response.message.content, self.language,
|
||||
self.__class__.__name__).parse_v2()
|
||||
if new_insight_keys:
|
||||
for insight_key in new_insight_keys:
|
||||
self.memory_manager.add_memories(INSIGHT_NODES, self.new_insight_node(insight_key))
|
||||
|
|
|
|||
|
|
@ -76,7 +76,7 @@ class InfoFilterWorker(MemoryBaseWorker):
|
|||
response_text = response.message.content
|
||||
|
||||
# parse text
|
||||
info_score_list = ResponseTextParser(response_text, self.__class__.__name__).parse_v1()
|
||||
info_score_list = ResponseTextParser(response_text, self.language, self.__class__.__name__).parse_v1()
|
||||
if len(info_score_list) != len(info_messages):
|
||||
self.logger.warning(f"score_size != messages_size, {len(info_score_list)} vs {len(info_messages)}")
|
||||
|
||||
|
|
|
|||
|
|
@ -111,7 +111,8 @@ class LongContraRepeatWorker(MemoryBaseWorker):
|
|||
return
|
||||
|
||||
# Parses the model's response text to identify updates for memory nodes
|
||||
idx_obs_info_list = ResponseTextParser(response.message.content, self.__class__.__name__).parse_v1()
|
||||
idx_obs_info_list = ResponseTextParser(response.message.content, self.language,
|
||||
self.__class__.__name__).parse_v1()
|
||||
if len(idx_obs_info_list) <= 0:
|
||||
self.logger.warning("idx_obs_info_list is empty!")
|
||||
return
|
||||
|
|
|
|||
|
|
@ -8,7 +8,7 @@ from memoryscope.memory.worker.memory_base_worker import MemoryBaseWorker
|
|||
from memoryscope.scheme.memory_node import MemoryNode
|
||||
from memoryscope.utils.datetime_handler import DatetimeHandler
|
||||
from memoryscope.utils.response_text_parser import ResponseTextParser
|
||||
from memoryscope.utils.tool_functions import prompt_to_msg
|
||||
from memoryscope.utils.tool_functions import prompt_to_msg, cosine_similarity
|
||||
|
||||
|
||||
class UpdateInsightWorker(MemoryBaseWorker):
|
||||
|
|
@ -27,13 +27,15 @@ class UpdateInsightWorker(MemoryBaseWorker):
|
|||
|
||||
def filter_obs_nodes(self,
|
||||
insight_node: MemoryNode,
|
||||
obs_nodes: List[MemoryNode]) -> (MemoryNode, List[MemoryNode], float):
|
||||
obs_nodes: List[MemoryNode],
|
||||
use_dummy_ranker: bool) -> (MemoryNode, List[MemoryNode], float):
|
||||
"""
|
||||
Filters observed nodes based on their relevance to a given insight node using a ranking model.
|
||||
|
||||
Args:
|
||||
insight_node (MemoryNode): The insight node used as the basis for filtering.
|
||||
obs_nodes (List[MemoryNode]): A list of observed nodes to be filtered.
|
||||
use_dummy_ranker (bool): Global parameters, whether to use rank model or not.
|
||||
|
||||
Returns:
|
||||
tuple: A tuple containing:
|
||||
|
|
@ -53,24 +55,47 @@ class UpdateInsightWorker(MemoryBaseWorker):
|
|||
self.logger.warning("obs_nodes is empty!")
|
||||
return insight_node, filtered_nodes, max_score
|
||||
|
||||
# Call the ranking model to get scores for each observed node's content against the insight key
|
||||
documents = [x.content for x in obs_nodes]
|
||||
self.logger.debug(f"update.insight.rank key={insight_node.key} \n docs={'|'.join(documents)}")
|
||||
response = self.rank_model.call(query=insight_node.key, documents=documents)
|
||||
if not response.status:
|
||||
return insight_node, filtered_nodes, max_score
|
||||
if use_dummy_ranker:
|
||||
key_vector: List[float] = self.embedding_model.call(text=insight_node.key).embedding_results
|
||||
if not key_vector:
|
||||
self.logger.warning(f"embedding call {insight_node.key} failed!")
|
||||
return insight_node, filtered_nodes, max_score
|
||||
|
||||
# Iterate over the ranked scores to filter nodes
|
||||
for index, score in response.rank_scores.items():
|
||||
node = obs_nodes[index]
|
||||
# Determine if the node should be kept based on the threshold
|
||||
keep_flag = score >= self.update_insight_threshold
|
||||
if keep_flag:
|
||||
filtered_nodes.append(node)
|
||||
max_score = max(max_score, score)
|
||||
# Log information about each node's processing
|
||||
self.logger.info(f"insight_key={insight_node.key} content={node.content} "
|
||||
f"score={score} keep_flag={keep_flag}")
|
||||
insight_node.key_vector = key_vector
|
||||
documents_vector = [x.vector for x in obs_nodes]
|
||||
score_recall_list = cosine_similarity(key_vector, documents_vector)
|
||||
assert len(score_recall_list) == len(obs_nodes), \
|
||||
f"size is not as excepted. {len(score_recall_list)} v.s. {len(obs_nodes)}"
|
||||
|
||||
for score, node in zip(score_recall_list, obs_nodes):
|
||||
keep_flag = score >= self.update_insight_threshold
|
||||
if keep_flag:
|
||||
filtered_nodes.append(node)
|
||||
max_score = max(max_score, score)
|
||||
|
||||
# Log information about each node's processing
|
||||
self.logger.info(f"insight_key={insight_node.key} content={node.content} "
|
||||
f"score={score} keep_flag={keep_flag}")
|
||||
|
||||
else:
|
||||
# Call the ranking model to get scores for each observed node's content against the insight key
|
||||
documents = [x.content for x in obs_nodes]
|
||||
self.logger.debug(f"update.insight.rank key={insight_node.key} \n docs={'|'.join(documents)}")
|
||||
response = self.rank_model.call(query=insight_node.key, documents=documents)
|
||||
if not response.status:
|
||||
return insight_node, filtered_nodes, max_score
|
||||
|
||||
# Iterate over the ranked scores to filter nodes
|
||||
for index, score in response.rank_scores.items():
|
||||
node = obs_nodes[index]
|
||||
# Determine if the node should be kept based on the threshold
|
||||
keep_flag = score >= self.update_insight_threshold
|
||||
if keep_flag:
|
||||
filtered_nodes.append(node)
|
||||
max_score = max(max_score, score)
|
||||
# Log information about each node's processing
|
||||
self.logger.info(f"insight_key={insight_node.key} content={node.content} "
|
||||
f"score={score} keep_flag={keep_flag}")
|
||||
|
||||
# Warn if no nodes were filtered
|
||||
if not filtered_nodes:
|
||||
|
|
@ -95,7 +120,7 @@ class UpdateInsightWorker(MemoryBaseWorker):
|
|||
content = f"{key}{self.get_language_value(COLON_WORD)} {insight_value}"
|
||||
insight_node.content = content
|
||||
insight_node.value = insight_value
|
||||
insight_node.meta_data.update({k: str(v) for k, v in dt_handler.get_dt_info_dict.items()})
|
||||
insight_node.meta_data.update({k: str(v) for k, v in dt_handler.get_dt_info_dict(self.language).items()})
|
||||
insight_node.timestamp = dt_handler.timestamp
|
||||
insight_node.dt = dt_handler.datetime_format()
|
||||
if insight_node.action_status == ActionStatusEnum.NONE.value:
|
||||
|
|
@ -136,7 +161,7 @@ class UpdateInsightWorker(MemoryBaseWorker):
|
|||
if not response.status or not response.message.content:
|
||||
return insight_node
|
||||
|
||||
insight_value_list = ResponseTextParser(response.message.content,
|
||||
insight_value_list = ResponseTextParser(response.message.content, self.language,
|
||||
f"update_{insight_node.key}").parse_v1()
|
||||
if not insight_value_list:
|
||||
self.logger.warning(f"update_{insight_node.key} insight_value_list is empty!")
|
||||
|
|
@ -184,17 +209,21 @@ class UpdateInsightWorker(MemoryBaseWorker):
|
|||
self.logger.warning("insight_nodes is empty, stopping processing.")
|
||||
return
|
||||
|
||||
use_dummy_ranker: bool = self.memoryscope_context.meta_data["use_dummy_ranker"]
|
||||
|
||||
# Process active insight nodes with corresponding not updated nodes
|
||||
for node in insight_nodes:
|
||||
time.sleep(1)
|
||||
if node.action_status == ActionStatusEnum.NEW.value:
|
||||
self.submit_thread_task(fn=self.filter_obs_nodes,
|
||||
insight_node=node,
|
||||
obs_nodes=not_reflected_nodes)
|
||||
obs_nodes=not_reflected_nodes,
|
||||
use_dummy_ranker=use_dummy_ranker)
|
||||
else:
|
||||
self.submit_thread_task(fn=self.filter_obs_nodes,
|
||||
insight_node=node,
|
||||
obs_nodes=not_updated_nodes)
|
||||
obs_nodes=not_updated_nodes,
|
||||
use_dummy_ranker=use_dummy_ranker)
|
||||
|
||||
# select top n
|
||||
result_list = []
|
||||
|
|
@ -221,7 +250,3 @@ class UpdateInsightWorker(MemoryBaseWorker):
|
|||
for node in not_updated_nodes:
|
||||
node.obs_updated = 1
|
||||
node.action_status = ActionStatusEnum.MODIFIED
|
||||
|
||||
# for node in not_reflected_nodes:
|
||||
# node.obs_updated = 1
|
||||
# node.action_status = ActionStatusEnum.MODIFIED
|
||||
|
|
|
|||
|
|
@ -33,13 +33,14 @@ class ExtractTimeWorker(MemoryBaseWorker):
|
|||
query, query_timestamp = self.get_context(QUERY_WITH_TS)
|
||||
|
||||
# Identify if the query contains datetime keywords
|
||||
contain_datetime = DatetimeHandler.has_time_word(query)
|
||||
contain_datetime = DatetimeHandler.has_time_word(query, self.language)
|
||||
if not contain_datetime:
|
||||
self.logger.info(f"contain_datetime={contain_datetime}")
|
||||
return
|
||||
|
||||
# Prepare the prompt with necessary contextual details
|
||||
query_time_str = DatetimeHandler(dt=query_timestamp).string_format(self.prompt_handler.time_string_format)
|
||||
query_time_str = DatetimeHandler(dt=query_timestamp).string_format(self.prompt_handler.time_string_format,
|
||||
self.language)
|
||||
system_prompt = self.prompt_handler.extract_time_system
|
||||
few_shot = self.prompt_handler.extract_time_few_shot
|
||||
user_query = self.prompt_handler.extract_time_user_query.format(query=query, query_time_str=query_time_str)
|
||||
|
|
|
|||
|
|
@ -34,6 +34,13 @@ class SemanticRankWorker(MemoryBaseWorker):
|
|||
self.logger.warning("Retrieve memory nodes is empty!")
|
||||
return
|
||||
|
||||
use_dummy_ranker: bool = self.memoryscope_context.meta_data["use_dummy_ranker"]
|
||||
if use_dummy_ranker:
|
||||
for node in memory_node_list:
|
||||
node.score_rank = node.score_recall
|
||||
self.logger.warning("use score_recall instead of score_rank!")
|
||||
return
|
||||
|
||||
# drop repeated
|
||||
memory_node_dict: Dict[str, MemoryNode] = {n.content.strip(): n for n in memory_node_list if n.content.strip()}
|
||||
memory_node_list = list(memory_node_dict.values())
|
||||
|
|
|
|||
|
|
@ -52,6 +52,7 @@ class MemoryScope(object):
|
|||
"thread_pool_max_workers": arguments.thread_pool_max_workers,
|
||||
"logger_name": arguments.logger_name,
|
||||
"logger_name_time_suffix": arguments.logger_name_time_suffix,
|
||||
"use_dummy_ranker": arguments.use_dummy_ranker,
|
||||
}
|
||||
|
||||
# prepare memory chat
|
||||
|
|
@ -151,6 +152,7 @@ class MemoryScope(object):
|
|||
# set global config
|
||||
self.context.language = LanguageEnum(self.global_conf["language"])
|
||||
self.context.thread_pool = ThreadPoolExecutor(max_workers=self.global_conf["max_workers"])
|
||||
self.context.meta_data["use_dummy_ranker"] = self.global_conf["use_dummy_ranker"]
|
||||
|
||||
# init memory_chat
|
||||
if self.memory_chat_conf_dict:
|
||||
|
|
@ -181,8 +183,9 @@ class MemoryScope(object):
|
|||
self.context.worker_config = self.worker_conf_dict
|
||||
|
||||
def close(self):
|
||||
# wait service to stop
|
||||
for _, service in self.context.memory_service_dict.items():
|
||||
service.stop_backend_service()
|
||||
service.stop_backend_service(wait_service_end=True)
|
||||
self.context.memory_store.close()
|
||||
self.context.thread_pool.shutdown()
|
||||
|
||||
|
|
|
|||
|
|
@ -19,35 +19,37 @@ class DummyGenerationModel(BaseModel):
|
|||
"""
|
||||
m_type: ModelEnum = ModelEnum.GENERATION_MODEL
|
||||
|
||||
class DummyModel:
|
||||
pass
|
||||
MODEL_REGISTRY.register("dummy_generation", object)
|
||||
|
||||
MODEL_REGISTRY.register("dummy_generation", DummyModel)
|
||||
|
||||
def before_call(self, **kwargs):
|
||||
def before_call(self, model_response: ModelResponse, **kwargs):
|
||||
"""
|
||||
Prepares the input data before making a call to the model's generate function.
|
||||
Accepts either a 'prompt' or a list of 'messages'. If both are provided or missing,
|
||||
a RuntimeError is raised. Transforms the input into a standardized format for processing.
|
||||
Prepares the input data before making a call to the language model.
|
||||
It accepts either a 'prompt' directly or a list of 'messages'.
|
||||
If 'prompt' is provided, it sets the data accordingly.
|
||||
If 'messages' are provided, it constructs a list of ChatMessage objects from the list.
|
||||
Raises an error if neither 'prompt' nor 'messages' are supplied.
|
||||
|
||||
Args:
|
||||
**kwargs: Arbitrary keyword arguments including 'prompt' or 'messages'.
|
||||
|
||||
model_response: model_response
|
||||
**kwargs: Arbitrary keyword arguments including 'prompt' and 'messages'.
|
||||
|
||||
Raises:
|
||||
RuntimeError: If neither 'prompt' nor 'messages' is provided, or both are provided.
|
||||
RuntimeError: When both 'prompt' and 'messages' inputs are not provided.
|
||||
"""
|
||||
prompt: str = kwargs.pop("prompt", "")
|
||||
messages: List[Message] | List[dict] = kwargs.pop("messages", [])
|
||||
|
||||
if prompt:
|
||||
self.data = {"prompt": prompt}
|
||||
data = {"prompt": prompt}
|
||||
elif messages:
|
||||
if isinstance(messages[0], dict):
|
||||
self.data = {"messages": [ChatMessage(role=msg["role"], content=msg["content"]) for msg in messages]}
|
||||
data = {"messages": [ChatMessage(role=msg["role"], content=msg["content"]) for msg in messages]}
|
||||
else:
|
||||
self.data = {"messages": [ChatMessage(role=msg.role, content=msg.content) for msg in messages]}
|
||||
data = {"messages": [ChatMessage(role=msg.role, content=msg.content) for msg in messages]}
|
||||
else:
|
||||
raise RuntimeError("Both 'prompt' and 'messages' are empty!")
|
||||
raise RuntimeError("prompt and messages are both empty!")
|
||||
data.update(**kwargs)
|
||||
model_response.meta_data["data"] = data
|
||||
|
||||
def after_call(self,
|
||||
model_response: ModelResponse,
|
||||
|
|
@ -84,35 +86,8 @@ class DummyGenerationModel(BaseModel):
|
|||
model_response.message.content = "".join(call_result)
|
||||
return model_response
|
||||
|
||||
def _call(self, stream: bool = False, **kwargs) -> ModelResponse | ModelResponseGen:
|
||||
"""
|
||||
Generates a dummy response based on the input data, supporting both immediate
|
||||
and streamed response types.
|
||||
def _call(self, model_response: ModelResponse, stream: bool = False, **kwargs):
|
||||
return model_response
|
||||
|
||||
Args:
|
||||
stream (bool, optional): If True, indicates the response should be generated
|
||||
in a streaming manner. Defaults to False.
|
||||
**kwargs: Additional keyword arguments not used in this dummy implementation.
|
||||
|
||||
Returns:
|
||||
Union[ModelResponse, ModelResponseGen]: A dummy response object or a generator
|
||||
object capable of streaming responses.
|
||||
"""
|
||||
assert "prompt" in self.data or "messages" in self.data
|
||||
results = ModelResponse(m_type=self.m_type)
|
||||
return results
|
||||
|
||||
async def _async_call(self, **kwargs) -> ModelResponse:
|
||||
"""
|
||||
Asynchronous version of `_call`, providing the same functionality but designed
|
||||
to be used in asynchronous contexts.
|
||||
|
||||
Args:
|
||||
**kwargs: Additional keyword arguments not used in this dummy implementation.
|
||||
|
||||
Returns:
|
||||
ModelResponse: A dummy response object suitable for asynchronous use.
|
||||
"""
|
||||
assert "prompt" in self.data or "messages" in self.data
|
||||
results = ModelResponse(m_type=self.m_type)
|
||||
return results
|
||||
async def _async_call(self, model_response: ModelResponse, **kwargs):
|
||||
return model_response
|
||||
|
|
|
|||
|
|
@ -23,6 +23,8 @@ class MemoryNode(BaseModel):
|
|||
|
||||
key: str = Field("", description="memory key")
|
||||
|
||||
key_vector: List[float] = Field([], description="memory key embedding result")
|
||||
|
||||
value: str = Field("", description="memory value")
|
||||
|
||||
score_recall: float = Field(0, description="embedding similarity score used in recall stage")
|
||||
|
|
@ -37,7 +39,7 @@ class MemoryNode(BaseModel):
|
|||
|
||||
store_status: str = Field("valid", description="store_status: valid / expired")
|
||||
|
||||
vector: List[float] = Field([], description="content embedding result, return empty")
|
||||
vector: List[float] = Field([], description="content embedding result")
|
||||
|
||||
timestamp: int = Field(default_factory=lambda: int(datetime.datetime.now().timestamp()),
|
||||
description="timestamp of the memory node")
|
||||
|
|
|
|||
|
|
@ -1,5 +1,5 @@
|
|||
import random
|
||||
from typing import Dict, List, Optional
|
||||
from typing import Dict, List
|
||||
|
||||
from llama_index.core import VectorStoreIndex
|
||||
from llama_index.core.schema import TextNode, NodeWithScore, QueryBundle
|
||||
|
|
@ -24,10 +24,10 @@ class LlamaIndexEsMemoryStore(BaseMemoryStore):
|
|||
self.emb_dims = None
|
||||
self.index_name = index_name
|
||||
self.embedding_model: BaseModel = embedding_model
|
||||
retrieval_strategy = ESCombinedRetrieveStrategy(retrieve_mode=retrieve_mode, hybrid_alpha=hybrid_alpha)
|
||||
self.es_store = SyncElasticsearchStore(index_name=index_name,
|
||||
es_url=es_url,
|
||||
retrieval_strategy=ESCombinedRetrieveStrategy(retrieve_mode=retrieve_mode,
|
||||
hybrid_alpha=hybrid_alpha),
|
||||
retrieval_strategy=retrieval_strategy,
|
||||
**kwargs)
|
||||
|
||||
# TODO The llamaIndex utilizes some deprecated functions, hence langchain logs warning messages. By
|
||||
|
|
@ -144,7 +144,11 @@ class LlamaIndexEsMemoryStore(BaseMemoryStore):
|
|||
return TextNode(id_=memory_node.memory_id,
|
||||
text=memory_node.content,
|
||||
embedding=embedding,
|
||||
metadata=memory_node.model_dump(exclude={"content", "vector", "score_recall", "score_rank", "score_rerank"}))
|
||||
metadata=memory_node.model_dump(exclude={"content",
|
||||
"vector",
|
||||
"score_recall",
|
||||
"score_rank",
|
||||
"score_rerank"}))
|
||||
|
||||
@staticmethod
|
||||
def _text_node_2_memory_node(text_node: NodeWithScore) -> MemoryNode:
|
||||
|
|
|
|||
|
|
@ -133,7 +133,6 @@ def _mode_must_match_retrieval_strategy(
|
|||
|
||||
class ESCombinedRetrieveStrategy(AsyncDenseVectorStrategy):
|
||||
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
|
|
@ -613,7 +612,6 @@ class SyncElasticsearchStore(BasePydanticVectorStore):
|
|||
] = None,
|
||||
es_filter: Optional[List[Dict]] = None,
|
||||
fields: List[str] = [],
|
||||
**kwargs: Any,
|
||||
) -> VectorStoreQueryResult:
|
||||
"""
|
||||
Asynchronously queries the Elasticsearch index for the top k most similar nodes
|
||||
|
|
@ -626,6 +624,7 @@ class SyncElasticsearchStore(BasePydanticVectorStore):
|
|||
A custom function to modify the Elasticsearch query body. Defaults to None.
|
||||
es_filter (List[Dict], optional): Additional filters to apply during the query.
|
||||
If filters are present in the query, these filters will not be used. Defaults to None.
|
||||
fields (List[str], optional): .
|
||||
|
||||
Returns:
|
||||
VectorStoreQueryResult: The result of the query, including nodes, their IDs,
|
||||
|
|
@ -700,7 +699,7 @@ class SyncElasticsearchStore(BasePydanticVectorStore):
|
|||
isinstance(self.retrieval_strategy, AsyncDenseVectorStrategy)
|
||||
and self.retrieval_strategy.hybrid
|
||||
):
|
||||
total_rank = sum(top_k_scores)
|
||||
# total_rank = sum(top_k_scores)
|
||||
top_k_scores = [rank for rank in top_k_scores]
|
||||
# top_k_scores = [(total_rank - rank) / total_rank for rank in top_k_scores]
|
||||
# top_k_scores = [total_rank - rank / total_rank for rank in top_k_scores]
|
||||
|
|
|
|||
|
|
@ -302,6 +302,7 @@ class DatetimeHandler(object):
|
|||
|
||||
Args:
|
||||
string_format (str): A format string where placeholders are keys from `dt_info_dict`.
|
||||
language (str): current language.
|
||||
|
||||
Returns:
|
||||
str: A formatted datetime string.
|
||||
|
|
|
|||
|
|
@ -26,7 +26,7 @@ class Logger(logging.Logger):
|
|||
max_bytes: int = 1024 * 1024 * 1024,
|
||||
backup_count: int = 10):
|
||||
"""
|
||||
Initializes the Logger instance, setting up handlers for console and/or file logging based on provided parameters.
|
||||
Initializes the Logger instance, setting up handlers for console and file logging based on provided parameters.
|
||||
|
||||
Args:
|
||||
name (str): Identifier for the logger.
|
||||
|
|
|
|||
|
|
@ -12,7 +12,8 @@ class Registry(object):
|
|||
|
||||
Attributes:
|
||||
name (str): The name of the registry.
|
||||
module_dict (Dict[str, Any]): A dictionary holding registered modules where keys are module names and values are the modules themselves.
|
||||
module_dict (Dict[str, Any]): A dictionary holding registered modules where keys are module names and values are
|
||||
the modules themselves.
|
||||
"""
|
||||
|
||||
def __init__(self, name: str):
|
||||
|
|
@ -31,7 +32,7 @@ class Registry(object):
|
|||
|
||||
Args:
|
||||
module_name (str): The name of module to be registered.
|
||||
modules (List[Any] | Dict[str, Any]): The module to be registered.
|
||||
module (List[Any] | Dict[str, Any]): The module to be registered.
|
||||
|
||||
Raises:
|
||||
NotImplementedError: If the input is already registered.
|
||||
|
|
@ -46,7 +47,8 @@ class Registry(object):
|
|||
|
||||
def batch_register(self, modules: List[Any] | Dict[str, Any]):
|
||||
"""
|
||||
Registers multiple modules in the registry in a single call. Accepts either a list of modules or a dictionary mapping names to modules.
|
||||
Registers multiple modules in the registry in a single call. Accepts either a list of modules or a dictionary
|
||||
mapping names to modules.
|
||||
|
||||
Args:
|
||||
modules (List[Any] | Dict[str, Any]): A list of modules or a dictionary mapping module names to the modules.
|
||||
|
|
|
|||
|
|
@ -2,28 +2,22 @@ import re
|
|||
from typing import List
|
||||
|
||||
from memoryscope.constants.language_constants import NONE_WORD
|
||||
from memoryscope.utils.global_context import G_CONTEXT
|
||||
from memoryscope.enumeration.language_enum import LanguageEnum
|
||||
from memoryscope.utils.logger import Logger
|
||||
|
||||
|
||||
class ResponseTextParser(object):
|
||||
"""
|
||||
The `ResponseTextParser` class is designed to parse and process response texts. It provides methods to extract specific
|
||||
The `ResponseTextParser` class is designed to parse and process response texts. It provides methods to extract
|
||||
patterns from the text and filter out unnecessary information, while also logging the processing steps and outcomes.
|
||||
"""
|
||||
|
||||
pattern_v1 = re.compile(r"<(.*?)>") # Regular expression pattern to match content within angle brackets
|
||||
|
||||
def __init__(self, response_text: str, logger_prefix: str = ""):
|
||||
"""
|
||||
Initializes the `ResponseTextParser` instance with the provided response text and sets up a logger.
|
||||
|
||||
Args:
|
||||
response_text (str): The raw response text that needs to be parsed and processed.
|
||||
"""
|
||||
PATTERN_V1 = re.compile(r"<(.*?)>") # Regular expression pattern to match content within angle brackets
|
||||
|
||||
def __init__(self, response_text: str, language: LanguageEnum, logger_prefix: str = ""):
|
||||
# Strips leading and trailing whitespace from the response text
|
||||
self.response_text: str = response_text.strip()
|
||||
self.language: LanguageEnum = language
|
||||
|
||||
# The prefix of log. Defaults to "".
|
||||
self.logger_prefix: str = logger_prefix
|
||||
|
|
@ -43,7 +37,7 @@ class ResponseTextParser(object):
|
|||
line = line.strip()
|
||||
if not line:
|
||||
continue
|
||||
matches = [match.group(1) for match in self.pattern_v1.finditer(line)]
|
||||
matches = [match.group(1) for match in self.PATTERN_V1.finditer(line)]
|
||||
if matches:
|
||||
result.append(matches)
|
||||
self.logger.info(f"{self.logger_prefix} response_text={self.response_text} result={result}", stacklevel=2)
|
||||
|
|
@ -51,18 +45,15 @@ class ResponseTextParser(object):
|
|||
|
||||
def parse_v2(self) -> List[str]:
|
||||
"""
|
||||
Extract lines which contain NONE_WORD in Chinese or English.
|
||||
Extract lines which contain NONE_WORD.
|
||||
|
||||
Args:
|
||||
prefix (str): The prefix of log. Defaults to "".
|
||||
|
||||
Returns:
|
||||
Contents match the specific patterns.
|
||||
"""
|
||||
result = []
|
||||
for line in self.response_text.split("\n"):
|
||||
line = line.strip()
|
||||
if not line or line.lower() == NONE_WORD.get(G_CONTEXT.language):
|
||||
if not line or line.lower() == NONE_WORD.get(self.language):
|
||||
continue
|
||||
result.append(line)
|
||||
self.logger.info(f"{self.logger_prefix} response_text={self.response_text} result={result}", stacklevel=2)
|
||||
|
|
|
|||
|
|
@ -26,7 +26,7 @@ class Timer(object):
|
|||
Args:
|
||||
name (str): The log name.
|
||||
time_log_type (str): The log type. Defaults to 'End'.
|
||||
use_ms (bool): Use 'ms' as the time scale or not. Defaults to True.
|
||||
use_ms (bool): Use 'ms' as the timescale or not. Defaults to True.
|
||||
stack_level (int): The stack level of log. Defaults to 2.
|
||||
float_precision (int): The precision of cost time. Defaults to 4.
|
||||
|
||||
|
|
|
|||
|
|
@ -6,6 +6,7 @@ from copy import deepcopy
|
|||
from importlib import import_module
|
||||
from typing import List
|
||||
|
||||
import numpy as np
|
||||
import pyfiglet
|
||||
from termcolor import colored
|
||||
|
||||
|
|
@ -18,7 +19,7 @@ ALL_COLORS = ["red", "green", "yellow", "blue", "magenta", "cyan", "light_grey",
|
|||
|
||||
def underscore_to_camelcase(name: str, is_first_title: bool = True) -> str:
|
||||
"""
|
||||
Converts a underscore_notation string to CamelCase.
|
||||
Converts an underscore_notation string to CamelCase.
|
||||
|
||||
Args:
|
||||
name (str): The underscore_notation string to be converted.
|
||||
|
|
@ -188,3 +189,21 @@ def contains_keyword(text, keywords) -> bool:
|
|||
escaped_keywords = map(re.escape, keywords)
|
||||
pattern = re.compile('|'.join(escaped_keywords), re.IGNORECASE)
|
||||
return pattern.search(text) is not None
|
||||
|
||||
|
||||
def cosine_similarity(query: List[float], documents: List[List[float]]):
|
||||
query = np.array(query)
|
||||
documents = np.array(documents)
|
||||
|
||||
query_norm = np.linalg.norm(query)
|
||||
if query_norm == 0:
|
||||
raise ValueError("Query vector norm is zero, which will result in a division by zero")
|
||||
|
||||
documents_norm = np.linalg.norm(documents, axis=1)
|
||||
if np.any(documents_norm == 0):
|
||||
raise ValueError("One of the document vectors has zero norm, which will result in a division by zero")
|
||||
|
||||
dot_product = np.dot(documents, query)
|
||||
|
||||
cosine_similarities = dot_product / (query_norm * documents_norm)
|
||||
return cosine_similarities.tolist()
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue