fix grammer bugs

This commit is contained in:
jinli.yl 2024-07-27 21:54:19 +08:00
parent 590cbf65c4
commit a5b45afa8c
29 changed files with 157 additions and 351 deletions

View file

@ -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

View file

@ -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

View file

View file

View 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:

View file

@ -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)"})

View file

@ -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"

View file

@ -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"

View file

@ -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):

View file

@ -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

View file

@ -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}")

View file

@ -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

View file

@ -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))

View file

@ -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)}")

View file

@ -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

View file

@ -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

View file

@ -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)

View file

@ -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())

View file

@ -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()

View file

@ -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

View file

@ -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")

View file

@ -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:

View file

@ -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]

View file

@ -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.

View file

@ -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.

View file

@ -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.

View file

@ -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)

View file

@ -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.

View file

@ -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()