mirror of
https://github.com/agentscope-ai/ReMe.git
synced 2026-09-05 08:06:15 +00:00
[dev] fix active status to new
This commit is contained in:
parent
4d6be46a1e
commit
e758cf41e6
58 changed files with 36 additions and 4423 deletions
2
.flake8
2
.flake8
|
|
@ -2,7 +2,7 @@
|
|||
exclude =
|
||||
scripts/*
|
||||
src/agentscope/rpc/*
|
||||
max-line-length = 79
|
||||
max-line-length = 120
|
||||
inline-quotes = "
|
||||
avoid-escape = no
|
||||
ignore =
|
||||
|
|
|
|||
|
|
@ -66,12 +66,9 @@ class BaseWorker(metaclass=ABCMeta):
|
|||
def set_context(self, key: str, value: Any):
|
||||
if self.is_multi_thread:
|
||||
with self.context_lock:
|
||||
self.context_dict[key] = value
|
||||
self.context[key] = value
|
||||
else:
|
||||
self.context[key] = value
|
||||
|
||||
def has_content(self, key: str):
|
||||
return key in self.context
|
||||
|
||||
def __getattr__(self, key: str):
|
||||
return self.kwargs[key]
|
||||
|
|
|
|||
|
|
@ -5,8 +5,8 @@ from memory_scope.constants.common_constants import CHAT_MESSAGES, CHAT_KWARGS
|
|||
from memory_scope.memory.worker.base_worker import BaseWorker
|
||||
from memory_scope.models.base_model import BaseModel
|
||||
from memory_scope.scheme.message import Message
|
||||
from memory_scope.storage.base_monitor import BaseMonitor
|
||||
from memory_scope.storage.base_memory_store import BaseMemoryStore
|
||||
from memory_scope.storage.base_monitor import BaseMonitor
|
||||
from memory_scope.utils.global_context import G_CONTEXT
|
||||
from memory_scope.utils.prompt_handler import PromptHandler
|
||||
|
||||
|
|
|
|||
|
|
@ -22,7 +22,7 @@ class GetReflectionSubjectWorker(MemoryBaseWorker):
|
|||
meta_data=meta_data,
|
||||
key=insight_key,
|
||||
memory_type=MemoryTypeEnum.INSIGHT.value,
|
||||
status=MemoryNodeStatus.ACTIVE.value)
|
||||
status=MemoryNodeStatus.NEW.value)
|
||||
|
||||
def _run(self):
|
||||
not_reflected_nodes: List[MemoryNode] = self.get_context(NOT_REFLECTED_NODES)
|
||||
|
|
|
|||
|
|
@ -1,46 +1,23 @@
|
|||
from typing import List, Dict
|
||||
from typing import List
|
||||
|
||||
from memory_scope.constants.common_constants import (
|
||||
NEW_INSIGHT_NODES,
|
||||
MODIFIED_MEMORIES,
|
||||
INSIGHT_NODES,
|
||||
NEW_OBS_NODES,
|
||||
NOT_REFLECTED_OBS_NODES,
|
||||
NEW,
|
||||
NOT_REFLECTED_MERGE_NODES,
|
||||
)
|
||||
from memory_scope.scheme.memory_node import MemoryNode
|
||||
from memory_scope.constants.common_constants import NOT_REFLECTED_NODES, INSIGHT_NODES, MERGE_OBS_NODES, \
|
||||
NOT_UPDATED_NODES
|
||||
from memory_scope.memory.worker.memory_base_worker import MemoryBaseWorker
|
||||
from memory_scope.scheme.memory_node import MemoryNode
|
||||
|
||||
|
||||
class SummaryCollectWorker(MemoryBaseWorker):
|
||||
|
||||
def _run(self):
|
||||
insight_nodes: List[MemoryNode] = self.get_context(INSIGHT_NODES)
|
||||
new_insight_nodes: List[MemoryNode] = self.get_context(NEW_INSIGHT_NODES)
|
||||
new_obs_nodes: List[MemoryNode] = self.get_context(NEW_OBS_NODES)
|
||||
not_reflected_nodes: List[MemoryNode] = self.get_context(
|
||||
NOT_REFLECTED_OBS_NODES
|
||||
)
|
||||
not_reflected_merge_nodes: List[MemoryNode] = self.get_context(
|
||||
NOT_REFLECTED_MERGE_NODES
|
||||
)
|
||||
keys = [
|
||||
INSIGHT_NODES,
|
||||
MERGE_OBS_NODES,
|
||||
NOT_UPDATED_NODES,
|
||||
NOT_REFLECTED_NODES,
|
||||
]
|
||||
|
||||
# 合并逻辑,复杂,务必check
|
||||
all_node_dict: Dict[str, MemoryNode] = {}
|
||||
if insight_nodes:
|
||||
all_node_dict.update(
|
||||
{n.id: n for n in insight_nodes if n.obs_updated}
|
||||
)
|
||||
if new_insight_nodes:
|
||||
all_node_dict.update({n.content: n for n in new_insight_nodes})
|
||||
if new_obs_nodes:
|
||||
# 设置为非新
|
||||
for n in new_obs_nodes:
|
||||
n.obs_updated = "0"
|
||||
all_node_dict.update({n.content: n for n in new_obs_nodes})
|
||||
if not_reflected_merge_nodes and not_reflected_nodes:
|
||||
# 进入reflect阶段
|
||||
all_node_dict.update({n.id: n for n in not_reflected_nodes})
|
||||
memory_nodes: List[MemoryNode] = []
|
||||
for key in keys:
|
||||
memory_nodes.extend(self.get_context(key))
|
||||
|
||||
self.set_context(MODIFIED_MEMORIES, list(all_node_dict.values()))
|
||||
self.memory_store.update_memories(memory_nodes)
|
||||
|
|
|
|||
|
|
@ -2,6 +2,7 @@ from typing import List
|
|||
|
||||
from memory_scope.constants.common_constants import INSIGHT_NODES, NOT_UPDATED_NODES, NOT_REFLECTED_NODES
|
||||
from memory_scope.constants.language_constants import COLON_WORD, NONE_WORD, REPEATED_WORD
|
||||
from memory_scope.enumeration.memory_status_enum import MemoryNodeStatus
|
||||
from memory_scope.memory.worker.memory_base_worker import MemoryBaseWorker
|
||||
from memory_scope.scheme.memory_node import MemoryNode
|
||||
from memory_scope.utils.datetime_handler import DatetimeHandler
|
||||
|
|
@ -48,6 +49,8 @@ class UpdateInsightWorker(MemoryBaseWorker):
|
|||
insight_node.meta_data.update({k: str(v) for k, v in dt_handler.dt_info_dict.items()})
|
||||
insight_node.timestamp = dt_handler.timestamp
|
||||
insight_node.dt = dt_handler.datetime_format()
|
||||
if insight_node.status == MemoryNodeStatus.ACTIVE.value:
|
||||
insight_node.status = MemoryNodeStatus.CONTENT_MODIFIED.value
|
||||
self.logger.info(f"after_update_{insight_node.key} value={insight_value}")
|
||||
return insight_node
|
||||
|
||||
|
|
@ -89,6 +92,10 @@ class UpdateInsightWorker(MemoryBaseWorker):
|
|||
self.logger.info(f"update_{insight_node.key} insight_value={insight_value} is invalid.")
|
||||
return insight_node
|
||||
|
||||
if insight_node.value == insight_value:
|
||||
self.logger.info(f"value={insight_value} is same!")
|
||||
return insight_node
|
||||
|
||||
self.update_insight_node(insight_node=insight_node, insight_value=insight_value)
|
||||
return insight_node
|
||||
|
||||
|
|
@ -102,7 +109,7 @@ class UpdateInsightWorker(MemoryBaseWorker):
|
|||
return
|
||||
|
||||
for node in insight_nodes:
|
||||
if node.content:
|
||||
if node.status == MemoryNodeStatus.ACTIVE.value:
|
||||
self.submit_async_task(fn=self.filter_obs_nodes,
|
||||
insight_node=node,
|
||||
not_updated_nodes=not_updated_nodes)
|
||||
|
|
@ -126,3 +133,6 @@ class UpdateInsightWorker(MemoryBaseWorker):
|
|||
|
||||
# get result
|
||||
self.gather_async_result()
|
||||
|
||||
for node in not_updated_nodes:
|
||||
node.obs_updated = True
|
||||
|
|
|
|||
|
|
@ -14,7 +14,7 @@ class StoreMemoryWorker(MemoryBaseWorker):
|
|||
|
||||
if self.has_content(store_key):
|
||||
memory_nodes: List[MemoryNode] = self.get_context(store_key)
|
||||
self.memory_store.update_batch(memory_nodes)
|
||||
self.memory_store.update_memories(memory_nodes)
|
||||
|
||||
elif store_key in self.chat_kwargs:
|
||||
query = self.chat_kwargs[store_key]
|
||||
|
|
@ -27,8 +27,8 @@ class StoreMemoryWorker(MemoryBaseWorker):
|
|||
target_name=self.target_name,
|
||||
content=query,
|
||||
memory_type=MemoryTypeEnum.OBS_CUSTOMIZED.value,
|
||||
status=MemoryNodeStatus.ACTIVE.value,
|
||||
status=MemoryNodeStatus.NEW.value,
|
||||
timestamp=dt_handler.timestamp,
|
||||
obs_reflected=False,
|
||||
obs_updated=False)
|
||||
self.memory_store.update(node)
|
||||
self.memory_store.update_memories(node)
|
||||
|
|
|
|||
|
|
@ -134,6 +134,10 @@ class LlamaIndexEsMemoryStore(BaseMemoryStore):
|
|||
self.es_store.close()
|
||||
|
||||
def update_memories(self, nodes: MemoryNode | List[MemoryNode]):
|
||||
if not nodes:
|
||||
self.logger.warning("empty nodes!")
|
||||
return
|
||||
|
||||
if isinstance(nodes, MemoryNode):
|
||||
nodes = [nodes]
|
||||
|
||||
|
|
|
|||
|
|
@ -1,124 +0,0 @@
|
|||
import json
|
||||
import time
|
||||
from http import HTTPStatus
|
||||
|
||||
import requests
|
||||
|
||||
from utils.logger import Logger
|
||||
from utils.timer import Timer
|
||||
from enumeration.env_type import EnvType
|
||||
|
||||
|
||||
class DashClient(object):
|
||||
|
||||
def __init__(self,
|
||||
request_id: str,
|
||||
dash_scope_uid: str,
|
||||
authorization: str,
|
||||
workspace: str,
|
||||
model_name: str,
|
||||
env_type: EnvType | str = EnvType.DAILY,
|
||||
timeout: int = None,
|
||||
max_retry_count: int = 2,
|
||||
retry_sleep_time: float = 1.0,
|
||||
**kwargs):
|
||||
|
||||
self.model_name: str = model_name
|
||||
self.env_type: EnvType = EnvType(env_type)
|
||||
self.timeout: int = timeout
|
||||
self.max_retry_count: int = max_retry_count
|
||||
self.retry_sleep_time: float = retry_sleep_time
|
||||
self.kwargs: dict = kwargs
|
||||
|
||||
# 20240506 update by 泉雨
|
||||
# if authorization:
|
||||
# workspace = ""
|
||||
# dash_scope_uid = ""
|
||||
|
||||
self.headers = {
|
||||
'Content-Type': 'application/json',
|
||||
'Authorization': authorization,
|
||||
'X-Request-Id': request_id,
|
||||
'X-DashScope-Uid': dash_scope_uid,
|
||||
'X-DashScope-WorkSpace': workspace,
|
||||
}
|
||||
|
||||
self.url: str = ""
|
||||
self.data = {}
|
||||
|
||||
self.logger = Logger.get_logger()
|
||||
|
||||
def before_call(self, model_name: str = None, **kwargs):
|
||||
pass
|
||||
|
||||
def after_call(self, response_obj, **kwargs):
|
||||
pass
|
||||
|
||||
def call_once(self, model_name: str = None, retry_cnt: int = 0, **kwargs):
|
||||
if model_name is None:
|
||||
model_name = self.model_name
|
||||
|
||||
self.before_call(model_name=model_name, **kwargs)
|
||||
|
||||
with Timer(self.__class__.__name__, log_time=False) as t:
|
||||
self.logger.debug(f"url={self.url} header={self.headers} data={self.data} timeout={self.timeout}")
|
||||
response = requests.post(url=self.url,
|
||||
headers=self.headers,
|
||||
data=json.dumps(self.data),
|
||||
timeout=self.timeout)
|
||||
|
||||
if response.status_code == HTTPStatus.OK:
|
||||
response_obj = json.loads(response.text)
|
||||
self.logger.info(f"{self.__class__.__name__} env={self.env_type.value} {t.get_cost_info()}, "
|
||||
f"call model={model_name} success! retry_cnt={retry_cnt}",
|
||||
stacklevel=3)
|
||||
return self.after_call(response_obj, **kwargs), True
|
||||
|
||||
else:
|
||||
self.logger.warning(f"{self.__class__.__name__} env={self.env_type.value} {t.get_cost_info()}, "
|
||||
f"call model={model_name} failed! retry_cnt={retry_cnt} details={response.text}",
|
||||
stacklevel=3)
|
||||
return None, False
|
||||
|
||||
def call(self, model_name: str = None, **kwargs):
|
||||
for i in range(self.max_retry_count):
|
||||
result, flag = self.call_once(model_name=model_name, retry_cnt=i, **kwargs)
|
||||
if flag:
|
||||
return result
|
||||
else:
|
||||
time.sleep(self.retry_sleep_time)
|
||||
|
||||
return None
|
||||
|
||||
|
||||
class LLIClient(object):
|
||||
|
||||
def __init__(self,
|
||||
model_name: str,
|
||||
timeout: int = None,
|
||||
max_retry_count: int = 2,
|
||||
retry_sleep_time: float = 1.0,
|
||||
**kwargs):
|
||||
|
||||
self.model_name: str = model_name
|
||||
self.timeout: int = timeout
|
||||
self.max_retry_count: int = max_retry_count
|
||||
self.retry_sleep_time: float = retry_sleep_time
|
||||
self.kwargs: dict = kwargs
|
||||
|
||||
self.data = {}
|
||||
self.logger = Logger.get_logger()
|
||||
|
||||
|
||||
def before_call(self, **kwargs):
|
||||
pass
|
||||
|
||||
def after_call(self, **kwargs):
|
||||
pass
|
||||
|
||||
def call_once(self, **kwargs):
|
||||
pass
|
||||
|
||||
def call(self, **kwargs):
|
||||
pass
|
||||
|
||||
|
|
@ -1,103 +0,0 @@
|
|||
from typing import List, Dict
|
||||
|
||||
import dashscope
|
||||
import time
|
||||
|
||||
from models import EMB
|
||||
from models.dash_client import DashClient, LLIClient
|
||||
|
||||
from typing import List, Dict
|
||||
from utils.registry import build_from_cfg
|
||||
from utils.timer import Timer
|
||||
|
||||
|
||||
from constants.common_constants import DASH_ENV_URL_DICT, DASH_API_URL_DICT
|
||||
from enumeration.dash_api_enum import DashApiEnum
|
||||
|
||||
|
||||
class DashEmbeddingClient(DashClient):
|
||||
"""
|
||||
url: https://help.aliyun.com/document_detail/2782232.html?spm=a2c4g.2782227.0.0.76195b1d9UeBAk#a6a39590fegqx
|
||||
"""
|
||||
|
||||
def __init__(self, model_name: str = dashscope.TextEmbedding.Models.text_embedding_v2, **kwargs):
|
||||
super(DashEmbeddingClient, self).__init__(model_name=model_name, **kwargs)
|
||||
self.url = DASH_ENV_URL_DICT.get(self.env_type) + DASH_API_URL_DICT.get(DashApiEnum.EMBEDDING)
|
||||
|
||||
def before_call(self, model_name: str = None, **kwargs):
|
||||
text: str | List[str] = kwargs.pop("text", "")
|
||||
# text_type: query or document
|
||||
text_type: str = kwargs.pop("text_type", "query")
|
||||
|
||||
if isinstance(text, str):
|
||||
text = [text]
|
||||
|
||||
self.kwargs["text_type"] = text_type
|
||||
self.data = {
|
||||
"model": model_name,
|
||||
"input": {
|
||||
"texts": text,
|
||||
},
|
||||
"parameters": {**kwargs, **self.kwargs},
|
||||
}
|
||||
|
||||
def after_call(self, response_obj, **kwargs) -> Dict[int, List[float]] | List[float]:
|
||||
embedding_results = {}
|
||||
for emb in response_obj["output"]["embeddings"]:
|
||||
embedding_results[emb["text_index"]] = emb["embedding"]
|
||||
|
||||
if len(embedding_results) == 1:
|
||||
embedding_results = list(embedding_results.values())[0]
|
||||
return embedding_results
|
||||
|
||||
class LLIEmbedding(LLIClient):
|
||||
|
||||
def __init__(self, method, model_name, **kwargs):
|
||||
super(LLIEmbedding, self).__init__(model_name, **kwargs)
|
||||
self.config = {
|
||||
"method": method,
|
||||
"model_name": model_name,
|
||||
**kwargs}
|
||||
self.embedder = build_from_cfg(self.config, EMB)
|
||||
|
||||
def before_call(self, **kwargs):
|
||||
text: str | List[str] = kwargs.pop("text", "")
|
||||
|
||||
if isinstance(text, str):
|
||||
text = [text]
|
||||
self.data = dict(texts=text)
|
||||
|
||||
def after_call(self, emb: Dict[int, List[float]], **kwargs) -> Dict[int, List[float]] | List[float]:
|
||||
embedding_results = {}
|
||||
for idx, e in enumerate(emb):
|
||||
embedding_results[idx] = e
|
||||
|
||||
if len(embedding_results) == 1:
|
||||
embedding_results = list(embedding_results.values())[0]
|
||||
return embedding_results
|
||||
|
||||
|
||||
def call_once(self, model_name: str = None, retry_cnt: int = 0, **kwargs):
|
||||
if model_name is None:
|
||||
model_name = self.model_name
|
||||
|
||||
self.before_call(model_name=model_name, **kwargs)
|
||||
with Timer(self.__class__.__name__, log_time=False) as t:
|
||||
self.logger.debug(f"data={self.data} timeout={self.timeout}")
|
||||
try:
|
||||
results = self.embedder.get_text_embedding_batch(**self.data)
|
||||
results = self.after_call(results)
|
||||
return results, True
|
||||
except Exception as e:
|
||||
self.logger.debug(f"Get Error in Embedding: {e}")
|
||||
return None, False
|
||||
|
||||
|
||||
def call(self, model_name: str = None, **kwargs):
|
||||
for i in range(self.max_retry_count):
|
||||
result, flag = self.call_once(model_name=model_name, retry_cnt=i, **kwargs)
|
||||
if flag:
|
||||
return result
|
||||
else:
|
||||
time.sleep(self.retry_sleep_time)
|
||||
return None
|
||||
|
|
@ -1,129 +0,0 @@
|
|||
from typing import List, Dict
|
||||
|
||||
import dashscope
|
||||
|
||||
from models.dash_client import DashClient, LLIClient
|
||||
from constants.common_constants import DASH_ENV_URL_DICT, DASH_API_URL_DICT
|
||||
from enumeration.dash_api_enum import DashApiEnum
|
||||
|
||||
import time
|
||||
from typing import List, Dict
|
||||
from utils.timer import Timer
|
||||
from models import LLM
|
||||
from utils.registry import build_from_cfg
|
||||
from llama_index.core.base.llms.types import ChatMessage
|
||||
from llama_index.core.base.llms.types import (
|
||||
ChatResponse,
|
||||
CompletionResponse,
|
||||
)
|
||||
|
||||
|
||||
class DashGenerateClient(DashClient):
|
||||
"""
|
||||
url: https://help.aliyun.com/document_detail/2712576.html
|
||||
"""
|
||||
|
||||
def __init__(self, model_name: str = dashscope.Generation.Models.qwen_max, **kwargs):
|
||||
super(DashGenerateClient, self).__init__(model_name=model_name, **kwargs)
|
||||
self.url = DASH_ENV_URL_DICT.get(self.env_type) + DASH_API_URL_DICT.get(DashApiEnum.GENERATION)
|
||||
|
||||
def before_call(self, model_name: str = None, **kwargs):
|
||||
prompt: str = kwargs.pop("prompt", "")
|
||||
messages: List[Dict[str, str]] = kwargs.pop("messages", [])
|
||||
|
||||
input_text = {}
|
||||
if prompt:
|
||||
input_text["prompt"] = prompt
|
||||
elif messages:
|
||||
input_text["messages"] = messages
|
||||
else:
|
||||
raise RuntimeError("prompt and messages is both empty!")
|
||||
|
||||
self.data = {
|
||||
"model": model_name,
|
||||
"input": input_text,
|
||||
"parameters": {**kwargs, **self.kwargs},
|
||||
}
|
||||
|
||||
def after_call(self, response_obj, **kwargs):
|
||||
self.logger.debug(f"response_obj={response_obj}")
|
||||
output = response_obj["output"]
|
||||
if "text" in output:
|
||||
return output["text"]
|
||||
elif "choices" in output:
|
||||
return output["choices"][0]["message"]["content"]
|
||||
else:
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
class LLILLM(LLIClient):
|
||||
|
||||
def __init__(self, method, model_name: str, **kwargs):
|
||||
super(LLILLM, self).__init__(model_name, **kwargs)
|
||||
self.config = {
|
||||
"method": method,
|
||||
"model_name": model_name,
|
||||
**kwargs}
|
||||
self.llm = build_from_cfg(self.config, LLM)
|
||||
|
||||
|
||||
def before_call(self, model_name: str = None, **kwargs):
|
||||
prompt: str = kwargs.pop("prompt", "")
|
||||
messages: List[Dict[str, str]] = kwargs.pop("messages", [])
|
||||
|
||||
if prompt:
|
||||
input_text = prompt
|
||||
input_type = 'prompt'
|
||||
llama_input = input_text
|
||||
elif messages:
|
||||
input_text = messages
|
||||
input_type = 'messages'
|
||||
llama_input = [ChatMessage(
|
||||
role=x['role'], content=x['content']
|
||||
) for x in input_text]
|
||||
else:
|
||||
raise RuntimeError("prompt and messages is both empty!")
|
||||
|
||||
self.data = {
|
||||
input_type: llama_input,
|
||||
}
|
||||
|
||||
def after_call(self, response_obj: ChatResponse | CompletionResponse, **kwargs) -> str:
|
||||
self.logger.debug(f"response_obj={response_obj}")
|
||||
if isinstance(response_obj, CompletionResponse):
|
||||
return response_obj.text
|
||||
elif isinstance(response_obj, ChatResponse):
|
||||
return response_obj.message.content
|
||||
else:
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
def call_once(self, model_name: str = None, retry_cnt: int = 0, **kwargs):
|
||||
if model_name is None:
|
||||
model_name = self.model_name
|
||||
|
||||
self.before_call(model_name=model_name, **kwargs)
|
||||
|
||||
with Timer(self.__class__.__name__, log_time=False) as t:
|
||||
self.logger.debug(f"data={self.data} timeout={self.timeout}")
|
||||
if True:
|
||||
# try:
|
||||
if 'prompt' in self.data:
|
||||
results = self.llm.complete(**self.data)
|
||||
else:
|
||||
results = self.llm.chat(**self.data)
|
||||
results = self.after_call(results)
|
||||
return results, True
|
||||
# except:
|
||||
# return None, False
|
||||
|
||||
|
||||
def call(self, model_name: str = None, **kwargs):
|
||||
for i in range(self.max_retry_count):
|
||||
result, flag = self.call_once(model_name=model_name, retry_cnt=i, **kwargs)
|
||||
print("dashscope llm results:",result)
|
||||
if flag:
|
||||
return result
|
||||
else:
|
||||
time.sleep(self.retry_sleep_time)
|
||||
return None
|
||||
|
|
@ -1,119 +0,0 @@
|
|||
from typing import List
|
||||
|
||||
import dashscope
|
||||
|
||||
from models.dash_client import DashClient, LLIClient
|
||||
from constants.common_constants import DASH_ENV_URL_DICT, DASH_API_URL_DICT
|
||||
from enumeration.dash_api_enum import DashApiEnum
|
||||
|
||||
|
||||
import time
|
||||
from typing import List
|
||||
from models import RERANKER
|
||||
from utils.timer import Timer
|
||||
|
||||
from utils.registry import build_from_cfg
|
||||
from llama_index.core.data_structs import Node
|
||||
from llama_index.core.schema import NodeWithScore # type: ignore
|
||||
|
||||
|
||||
class DashReRankClient(DashClient):
|
||||
"""
|
||||
url: https://help.aliyun.com/document_detail/2780059.html
|
||||
"""
|
||||
|
||||
def __init__(self, model_name: str = dashscope.TextReRank.Models.gte_rerank, **kwargs):
|
||||
super(DashReRankClient, self).__init__(model_name=model_name, **kwargs)
|
||||
self.url = DASH_ENV_URL_DICT.get(self.env_type) + DASH_API_URL_DICT.get(DashApiEnum.RERANK)
|
||||
|
||||
def before_call(self, model_name: str = None, **kwargs):
|
||||
query: str = kwargs.pop("query", "")
|
||||
documents: List[str] = kwargs.pop("documents", [])
|
||||
top_n: int | None = kwargs.pop("top_n", None)
|
||||
return_documents: bool = kwargs.pop("return_documents", False)
|
||||
|
||||
assert query and documents, f"query or documents is empty! query={query}, documents={len(documents)}"
|
||||
if top_n is None:
|
||||
top_n = len(documents)
|
||||
|
||||
self.kwargs.update({
|
||||
"top_n": top_n,
|
||||
"return_documents": return_documents,
|
||||
})
|
||||
self.data = {
|
||||
"model": model_name,
|
||||
"input": {
|
||||
"query": query,
|
||||
"documents": documents,
|
||||
},
|
||||
"parameters": {**kwargs, **self.kwargs},
|
||||
}
|
||||
|
||||
def after_call(self, response_obj, **kwargs):
|
||||
return response_obj["output"]["results"]
|
||||
|
||||
class LLIReRank(LLIClient):
|
||||
|
||||
def __init__(self, method, model_name, **kwargs):
|
||||
super(LLIReRank, self).__init__(model_name, **kwargs)
|
||||
|
||||
self.config = {
|
||||
"method": method,
|
||||
"model_name": model_name,
|
||||
**kwargs}
|
||||
self.reranker = build_from_cfg(self.config, RERANKER)
|
||||
|
||||
|
||||
def before_call(self, model_name: str = None, **kwargs):
|
||||
query: str = kwargs.pop("query", "")
|
||||
documents: List[str] = kwargs.pop("documents", [])
|
||||
top_n: int | None = kwargs.pop("top_n", None)
|
||||
return_documents: bool = kwargs.pop("return_documents", False)
|
||||
|
||||
assert query and documents, f"query or documents is empty! query={query}, documents={len(documents)}"
|
||||
if top_n is None:
|
||||
top_n = len(documents)
|
||||
|
||||
nodes = [NodeWithScore(node=Node(text=text), score=1.0) for text in documents]
|
||||
|
||||
self.data = {
|
||||
"nodes": nodes,
|
||||
"query_str": query,
|
||||
}
|
||||
|
||||
|
||||
def after_call(self, nodes, **kwargs):
|
||||
results = []
|
||||
for idx, node in enumerate(nodes):
|
||||
results.append(dict(index=idx,
|
||||
relevance_score=node.score,
|
||||
document=node.node.text))
|
||||
return results
|
||||
|
||||
|
||||
def call_once(self, model_name: str = None, retry_cnt: int = 0, **kwargs):
|
||||
if model_name is None:
|
||||
model_name = self.model_name
|
||||
|
||||
self.before_call(model_name=model_name, **kwargs)
|
||||
|
||||
with Timer(self.__class__.__name__, log_time=False) as t:
|
||||
self.logger.debug(f"data={self.data} timeout={self.timeout}")
|
||||
try:
|
||||
results = self.reranker.postprocess_nodes(**self.data)
|
||||
results = self.after_call(results)
|
||||
return results, True
|
||||
except Exception as e:
|
||||
self.logger.debug(f"Rerank falls, data={self.data}")
|
||||
# return None, False
|
||||
|
||||
|
||||
def call(self, model_name: str = None, **kwargs):
|
||||
for i in range(self.max_retry_count):
|
||||
result, flag = self.call_once(model_name=model_name, retry_cnt=i, **kwargs)
|
||||
if flag:
|
||||
return result
|
||||
else:
|
||||
time.sleep(self.retry_sleep_time)
|
||||
return None
|
||||
|
||||
|
|
@ -1,419 +0,0 @@
|
|||
from elasticsearch import Elasticsearch
|
||||
from elasticsearch.helpers import bulk
|
||||
|
||||
|
||||
from models.dash_embedding_client import DashEmbeddingClient, LLIEmbedding
|
||||
from common.dash_embedding_client import DashEmbeddingClient
|
||||
from common.logger import Logger
|
||||
|
||||
from constants.common_constants import ES_ENV_URL_DICT
|
||||
from enumeration.env_type import EnvType
|
||||
from utils.logger import Logger
|
||||
from llama_index.core import VectorStoreIndex, StorageContext, ServiceContext
|
||||
from llama_index.vector_stores.elasticsearch import ElasticsearchStore
|
||||
from llama_index.core.schema import TextNode
|
||||
from llama_index.vector_stores.elasticsearch import AsyncDenseVectorStrategy
|
||||
|
||||
class ElasticSearchClient(object):
|
||||
def __init__(self,
|
||||
es_user_name: str,
|
||||
es_password: str,
|
||||
es_index_name: str,
|
||||
embedding_client: DashEmbeddingClient | None = None,
|
||||
env_type: EnvType | str = EnvType.DAILY,
|
||||
content_key: str = "content",
|
||||
vector_key: str = "vector",
|
||||
**kwargs):
|
||||
|
||||
self.es_index_name: str = es_index_name
|
||||
self.embedding_client: DashEmbeddingClient = embedding_client
|
||||
self.content_key: str = content_key
|
||||
self.vector_key: str = vector_key
|
||||
|
||||
self.es_client = Elasticsearch(
|
||||
hosts=[ES_ENV_URL_DICT.get(EnvType(env_type))],
|
||||
basic_auth=(es_user_name, es_password),
|
||||
**kwargs)
|
||||
|
||||
self.logger = Logger.get_logger()
|
||||
self.logger.debug(f"connect es_client info={self.es_client.info()}")
|
||||
|
||||
def log_index_info(self):
|
||||
index_info = self.es_client.indices.get(index=self.es_index_name)
|
||||
self.logger.info(f"index={self.es_index_name} exists. index_info={index_info}")
|
||||
|
||||
def insert(self, _id: str, body: dict):
|
||||
assert body and self.content_key in body, f"body={body} is illegal!"
|
||||
|
||||
# text_type: document
|
||||
content = body[self.content_key]
|
||||
vector = self.embedding_client.call(text=content, text_type="document")
|
||||
if not vector:
|
||||
self.logger.warning(f"embedding_client call failed, stop es insert!")
|
||||
return
|
||||
|
||||
body[self.vector_key] = vector
|
||||
response = self.es_client.index(id=_id, index=self.es_index_name, body=body)
|
||||
self.logger.info(f"insert response={response}")
|
||||
|
||||
def insert_batch(self, doc_list: list):
|
||||
"""
|
||||
doc_list = [
|
||||
{
|
||||
"_id": 2,
|
||||
"_source": {
|
||||
"author": "john",
|
||||
"text": "Elasticsearch: cool.",
|
||||
"timestamp": "2023-03-23T10:00:00"
|
||||
}
|
||||
},
|
||||
{
|
||||
"_id": 3,
|
||||
"_source": {
|
||||
"author": "jane",
|
||||
"text": "Elasticsearch: very cool.",
|
||||
"timestamp": "2023-03-23T11:00:00"
|
||||
}
|
||||
}
|
||||
]
|
||||
"""
|
||||
text_list = []
|
||||
for doc in doc_list:
|
||||
assert "_id" in doc and "_source" in doc
|
||||
content = doc["_source"][self.content_key]
|
||||
text_list.append(content)
|
||||
|
||||
vector_dict = self.embedding_client.call(text=text_list, text_type="document")
|
||||
if not vector_dict:
|
||||
self.logger.warning(f"embedding_client call failed, stop es insert!")
|
||||
return
|
||||
|
||||
# add _index
|
||||
for i, doc in enumerate(doc_list):
|
||||
doc["_index"] = self.es_index_name
|
||||
vector = vector_dict[i]
|
||||
doc["_source"][self.vector_key] = vector
|
||||
|
||||
# 执行批量插入
|
||||
responses = bulk(self.es_client, doc_list)
|
||||
|
||||
# 输出批量插入的响应
|
||||
for response in responses[1]:
|
||||
self.logger.info(f"insert_batch response={response}")
|
||||
|
||||
def print_hits(self, hits: list):
|
||||
for hit in hits:
|
||||
print_kwargs = {
|
||||
"_id": hit['_id'],
|
||||
"_score": hit['_score'],
|
||||
}
|
||||
for k, v in hit['_source'].items():
|
||||
# 不打印vector
|
||||
if k == self.vector_key:
|
||||
v = len(v)
|
||||
print_kwargs[k] = v
|
||||
self.logger.info(" ".join([f"{k}={v}" for k, v in print_kwargs.items()]))
|
||||
|
||||
def exact_search(self,
|
||||
size: int,
|
||||
exact_filters: dict = None,
|
||||
wildcard_filters: dict = None,
|
||||
print_hits: bool = False,
|
||||
exclude_vector: bool = True):
|
||||
"""
|
||||
{
|
||||
"match": {
|
||||
"category": "electronics" # 一级字段过滤
|
||||
}
|
||||
},
|
||||
{
|
||||
"match": {
|
||||
"product.name": "laptop" # 二级字段过滤
|
||||
}
|
||||
}
|
||||
{
|
||||
"terms": {
|
||||
"product.keyA": ["a", "b", "c"] # 二级字段keyA的精确值必须为a、b、c中的一
|
||||
}
|
||||
}
|
||||
"""
|
||||
must_list = []
|
||||
for key, value in exact_filters.items():
|
||||
if not key:
|
||||
continue
|
||||
if isinstance(value, str):
|
||||
must_list.append({"match": {key: value}})
|
||||
elif isinstance(value, list):
|
||||
must_list.append({"terms": {key: value}})
|
||||
|
||||
query = {
|
||||
"size": size,
|
||||
"query": {
|
||||
"bool": {
|
||||
"must": must_list
|
||||
}
|
||||
},
|
||||
# 添加_source配置以排除vector字段
|
||||
"_source": {
|
||||
"excludes": [self.vector_key] if exclude_vector else []
|
||||
}
|
||||
}
|
||||
|
||||
if wildcard_filters:
|
||||
should_list = []
|
||||
for key, value in wildcard_filters.items():
|
||||
if not key:
|
||||
continue
|
||||
if isinstance(value, str):
|
||||
should_list.append({"wildcard": {key: f"*{value}*"}})
|
||||
elif isinstance(value, list):
|
||||
for v in value:
|
||||
should_list.append({"wildcard": {key: f"*{v}*"}})
|
||||
|
||||
query["query"]["bool"].update({
|
||||
"should": should_list,
|
||||
"minimum_should_match": 1,
|
||||
})
|
||||
self.logger.info(f"query={query}")
|
||||
|
||||
response = self.es_client.search(index=self.es_index_name, body=query)
|
||||
hits = response['hits']['hits']
|
||||
|
||||
# 耗时log
|
||||
self.logger.info(f"exact_search cost={response['took']}ms "
|
||||
f"size={len(hits)} "
|
||||
f"timed_out={response['timed_out']} "
|
||||
f"shards={response['_shards']} "
|
||||
f"exact_filters={exact_filters}", stacklevel=2)
|
||||
|
||||
# 每一条结果log一次
|
||||
if print_hits:
|
||||
self.print_hits(hits)
|
||||
|
||||
return hits
|
||||
|
||||
def exact_search_v2(self,
|
||||
size: int,
|
||||
term_filters: dict = None,
|
||||
match_filters: dict = None,
|
||||
print_hits: bool = False,
|
||||
exclude_vector: bool = True):
|
||||
|
||||
"""
|
||||
"bool": {
|
||||
"must": [
|
||||
{"term": {"field1": "固定值"}}, # 一级目录关键字过滤(等于某个值)
|
||||
{"terms": {"field2": ["a", "b", "c"]}} # 二级目录关键字过滤(等于三个中的任意一个)
|
||||
],
|
||||
"should": [ # 至少匹配其中之一
|
||||
{"match": {"key": "ccc"}}, # key包含"ccc"
|
||||
{"match": {"key": "bbb"}} # 或者key包含"bbb"
|
||||
],
|
||||
"minimum_should_match": 1 # 至少有一个`should`条件匹配
|
||||
}
|
||||
"""
|
||||
|
||||
query = {
|
||||
"size": size,
|
||||
"query": {
|
||||
"bool": {
|
||||
|
||||
}
|
||||
},
|
||||
# 添加_source配置以排除vector字段
|
||||
"_source": {
|
||||
"excludes": [self.vector_key] if exclude_vector else []
|
||||
}
|
||||
}
|
||||
|
||||
if term_filters:
|
||||
must_list = []
|
||||
for k, v in term_filters.items():
|
||||
if isinstance(v, list):
|
||||
must_list.append({"terms": {k: v}})
|
||||
elif isinstance(v, str):
|
||||
must_list.append({"term": {k: v}})
|
||||
else:
|
||||
raise NotImplemented
|
||||
query["query"]["bool"]["must"] = must_list
|
||||
|
||||
if match_filters:
|
||||
match_list = []
|
||||
for k, v in match_filters.items():
|
||||
if isinstance(v, list):
|
||||
for v_sub in v:
|
||||
match_list.append({"match": {k: v_sub}})
|
||||
elif isinstance(v, str):
|
||||
match_list.append({"match": {k: v}})
|
||||
else:
|
||||
raise NotImplemented
|
||||
query["query"]["bool"]["should"] = match_list
|
||||
query["query"]["bool"]["minimum_should_match"] = 1
|
||||
|
||||
self.logger.info(query)
|
||||
response = self.es_client.search(index=self.es_index_name, body=query)
|
||||
hits = response['hits']['hits']
|
||||
|
||||
# 耗时log
|
||||
self.logger.info(f"exact_search cost={response['took']}ms "
|
||||
f"size={len(hits)} "
|
||||
f"timed_out={response['timed_out']} "
|
||||
f"shards={response['_shards']}", stacklevel=2)
|
||||
|
||||
# 每一条结果log一次
|
||||
if print_hits:
|
||||
self.print_hits(hits)
|
||||
|
||||
return hits
|
||||
|
||||
def similar_search(self,
|
||||
text: str,
|
||||
size: int,
|
||||
exact_filters: dict = None,
|
||||
print_hits: bool = False,
|
||||
exclude_vector: bool = True):
|
||||
|
||||
if exact_filters is None:
|
||||
exact_filters = {}
|
||||
|
||||
# 过滤or
|
||||
or_filters = {}
|
||||
for k in list(exact_filters.keys()):
|
||||
v = exact_filters[k]
|
||||
if isinstance(v, list):
|
||||
exact_filters.pop(k)
|
||||
or_filters[k] = v
|
||||
|
||||
vector = self.embedding_client.call(text=text)
|
||||
if not vector:
|
||||
self.logger.warning(f"embedding_client call failed, stop select from es!")
|
||||
return
|
||||
|
||||
query = {
|
||||
# 返回最相似的top_k个文档
|
||||
"size": size,
|
||||
"query": {
|
||||
"bool": {
|
||||
"must": {
|
||||
"script_score": {
|
||||
# 对所有文档执行
|
||||
"query": {
|
||||
"match_all": {}
|
||||
},
|
||||
"script": {
|
||||
# 使用余弦相似度+1,es不能返回负数
|
||||
"source": f"cosineSimilarity(params.query_vector, '{self.vector_key}') + 1.0",
|
||||
"params": {"query_vector": vector}
|
||||
}
|
||||
}
|
||||
},
|
||||
"filter": [
|
||||
{"term": {k: v}} for k, v in exact_filters.items()
|
||||
],
|
||||
}
|
||||
},
|
||||
# 添加_source配置以排除vector字段
|
||||
"_source": {
|
||||
"excludes": [self.vector_key] if exclude_vector else []
|
||||
}
|
||||
}
|
||||
|
||||
if or_filters:
|
||||
k_v_pair = []
|
||||
for k, v_list in or_filters.items():
|
||||
for v in v_list:
|
||||
k_v_pair.append((k, v))
|
||||
query["query"]["bool"]["should"] = [{"term": {k: v}} for k, v in k_v_pair]
|
||||
query["query"]["bool"]["minimum_should_match"] = 1
|
||||
|
||||
response = self.es_client.search(index=self.es_index_name, body=query)
|
||||
hits = response['hits']['hits']
|
||||
|
||||
# 耗时log
|
||||
self.logger.info(f"similar_search cost={response['took']}ms "
|
||||
f"size={len(hits)} "
|
||||
f"timed_out={response['timed_out']} "
|
||||
f"shards={response['_shards']} "
|
||||
f"text={text} "
|
||||
f"exact_filters={exact_filters}", stacklevel=2)
|
||||
|
||||
# 还原打分
|
||||
for hit in hits:
|
||||
hit['_score'] -= 1
|
||||
|
||||
# 每一条结果log一次
|
||||
if print_hits:
|
||||
self.print_hits(hits)
|
||||
|
||||
return hits
|
||||
|
||||
|
||||
class LLIElasticSearch(object):
|
||||
def __init__(self,
|
||||
es_index_name: str,
|
||||
embedding_client: LLIEmbedding | None = None,
|
||||
retrieve_topk: int = 3,
|
||||
content_key: str = "text",
|
||||
):
|
||||
self.es_index_name = es_index_name
|
||||
self.content_key = content_key
|
||||
self.embedding_client: LLIEmbedding = embedding_client
|
||||
# using local es for debug convenient
|
||||
self.es_client = ElasticsearchStore(index_name="my_index",
|
||||
es_url="http://localhost:9200",
|
||||
retrieval_strategy=AsyncDenseVectorStrategy(hybrid=True))
|
||||
|
||||
self.service_context = ServiceContext.from_defaults(embed_model=self.embedding_client, llm=None)
|
||||
self.storage_context = StorageContext.from_defaults(vector_store=self.es_client)
|
||||
self.index = VectorStoreIndex(storage_context=self.storage_context,
|
||||
service_context=self.service_context)
|
||||
|
||||
self.retriever = self.index.as_retriever(similarity_top_k=retrieve_topk)
|
||||
self.logger = Logger.get_logger()
|
||||
|
||||
def log_index_info(self, ):
|
||||
pass
|
||||
|
||||
def print_hits(self, hits: list):
|
||||
for hit in hits:
|
||||
print_kwargs = {
|
||||
"_id": hit['_id'],
|
||||
"_score": hit['_score'],
|
||||
}
|
||||
for k, v in hit['_source'].items():
|
||||
# 不打印vector
|
||||
if k == self.vector_key:
|
||||
v = len(v)
|
||||
print_kwargs[k] = v
|
||||
self.logger.info(" ".join([f"{k}={v}" for k, v in print_kwargs.items()]))
|
||||
|
||||
def similar_search(self,
|
||||
text: str,
|
||||
size: int, ):
|
||||
|
||||
ret_nodes = self.retriever.retrieve(text)
|
||||
return ret_nodes
|
||||
|
||||
def insert_batch(self, doc_list:list[str]):
|
||||
node_list = []
|
||||
for doc in doc_list:
|
||||
assert "_id" in doc and "_source" in doc
|
||||
content = doc["_source"]["text"]
|
||||
doc["_source"].pop("text")
|
||||
meta = doc["_source"]
|
||||
node = TextNode(text=content, metadata=meta)
|
||||
node.node_id(doc['_id'])
|
||||
node_list.append(node)
|
||||
self.index.insert_nodes(node_list)
|
||||
|
||||
|
||||
def insert(self, _id: str, body: dict):
|
||||
assert body and self.content_key in body, f"body={body} is illegal!"
|
||||
content = body[self.content_key]
|
||||
body.pop(self.content_key)
|
||||
meta = body
|
||||
node = TextNode(text=content, metadata=meta)
|
||||
self.index.insert_nodes([node])
|
||||
|
||||
|
||||
|
||||
|
|
@ -1,71 +0,0 @@
|
|||
import re
|
||||
from typing import Dict, List
|
||||
|
||||
from pydantic import Field, BaseModel
|
||||
|
||||
|
||||
class MemoryNode(BaseModel):
|
||||
"""
|
||||
除了 content_modified,其他均和数据库字段保持统一
|
||||
根据code判断,如果code是空,则为新增的memoryNode,如果有值,则为更新
|
||||
if content_modified is true,则需要调用embedding服务
|
||||
"""
|
||||
id: str = Field("", description="唯一主键 uuid64")
|
||||
|
||||
code: str = Field("", description="和id保持一致,为空则是新增")
|
||||
|
||||
# 0520新增
|
||||
timeCreated: str = Field("", description="Memory创建时间(算法不关注)")
|
||||
|
||||
# 0520新增
|
||||
timeModified: str = Field("", description="Memory更新时间(算法不关注)")
|
||||
|
||||
content: str = Field("", description="记忆内容")
|
||||
|
||||
memoryId: str = Field("", description="记忆 id,检索区分字段")
|
||||
|
||||
# 0520新增
|
||||
scene: str = Field("", description="source: TONGYI_MAIN_CHAT/TONGYI_CHAR_CHAT/BAILIAN/ASSISTANT")
|
||||
|
||||
# 0520新增
|
||||
# NOTE 百炼服务端只召回observation, insight, profile, obs_customized, profile_customized
|
||||
memoryType: str = Field("", description="conversation, observation, insight, "
|
||||
"profile, obs_customized, profile_customized")
|
||||
|
||||
# 0520新增,但不是数据库字段
|
||||
content_modified: bool = Field(False, description="content是否被更新,if true,则需要调用embedding服务")
|
||||
|
||||
# reflected: 1 is reflected before, 0 has not reflected, 如果是用户自定义,写入空值"".
|
||||
metaData: Dict[str, str] = Field({}, description="元信息: infoScore, algoVersion, datetime, reflected")
|
||||
|
||||
status: str = Field("active", description="active or expired")
|
||||
|
||||
tenantId: str = Field("", description="request id")
|
||||
|
||||
vector: List[float] = Field([], description="content embedding result, return empty")
|
||||
|
||||
def get_time_info(self, time_format: str):
|
||||
pattern = re.compile(r'\{([^}]*)}')
|
||||
keys = pattern.findall(time_format)
|
||||
|
||||
match_flag = True
|
||||
kv_dict = {}
|
||||
for k in keys:
|
||||
if k not in self.metaData:
|
||||
match_flag = False
|
||||
break
|
||||
v = self.metaData[k]
|
||||
if not v:
|
||||
match_flag = False
|
||||
break
|
||||
|
||||
kv_dict[k] = v
|
||||
|
||||
if match_flag:
|
||||
return time_format.format(**kv_dict)
|
||||
return ""
|
||||
|
||||
def to_dict(self):
|
||||
res = {"content": self.content, "memoryId": self.memoryId, "memoryType": self.memoryType,
|
||||
"status": self.status, "metaData": self.metaData}
|
||||
return res
|
||||
|
|
@ -1,33 +0,0 @@
|
|||
from pydantic import Field, BaseModel
|
||||
|
||||
from scheme.memory_node import MemoryNode
|
||||
|
||||
|
||||
class MemoryNode(BaseModel):
|
||||
id: str = Field("", description="uuid64")
|
||||
|
||||
score_similar: float = Field(0, description="相似度打分")
|
||||
|
||||
score_rank: float = Field(0, description="排序打分")
|
||||
|
||||
score_rerank: float = Field(0, description="重排打分")
|
||||
|
||||
memory_node: MemoryNode = Field(None, description="memory node 核心,返回给上游的结构")
|
||||
|
||||
@classmethod
|
||||
def init_from_es(cls, hit: dict):
|
||||
memory_node = MemoryNode(**hit['_source'])
|
||||
return cls(id=hit['_id'], score_similar=hit['_score'], memory_node=memory_node)
|
||||
|
||||
@classmethod
|
||||
def init_from_attrs(cls, **kwargs):
|
||||
_id: str = kwargs.get("_id", "")
|
||||
score_similar: float = kwargs.pop("score_similar", 0)
|
||||
score_rank: float = kwargs.pop("score_rank", 0)
|
||||
score_rerank: float = kwargs.pop("score_rerank", 0)
|
||||
memory_node = MemoryNode(**kwargs)
|
||||
return cls(id=_id,
|
||||
score_similar=score_similar,
|
||||
score_rank=score_rank,
|
||||
score_rerank=score_rerank,
|
||||
memory_node=memory_node)
|
||||
|
|
@ -1,140 +0,0 @@
|
|||
from datetime import datetime
|
||||
from typing import List
|
||||
|
||||
from common.tool_functions import time_to_formatted_str, get_datetime_info_dict
|
||||
from constants.common_constants import NEW_INSIGHT_NODES, DT, NOT_REFLECTED_MERGE_NODES, NEW_INSIGHT_KEYS, INSIGHT_KEY, \
|
||||
INSIGHT_VALUE, REFLECTED
|
||||
from enumeration.memory_status_enum import MemoryNodeStatus
|
||||
from enumeration.memory_type_enum import MemoryTypeEnum
|
||||
from scheme.memory_node import MemoryNode
|
||||
from worker.memory_base_worker import MemoryBaseWorker
|
||||
|
||||
|
||||
class GetInsightWorker(MemoryBaseWorker):
|
||||
def __init__(self, insight_obs_max_cnt, es_insight_similar_top_k, get_insight_model, get_insight_max_token, get_insight_temperature, get_insight_top_k, **kwargs):
|
||||
super(GetInsightWorker,self).__init__(*args,**kwargs)
|
||||
self.insight_obs_max_cnt = insight_obs_max_cnt
|
||||
self.get_insight_model = get_insight_model
|
||||
self.get_insight_max_token = get_insight_max_token
|
||||
self.get_insight_temperature = get_insight_temperature
|
||||
self.get_insight_top_k = get_insight_top_k
|
||||
self.es_insight_similar_top_k = es_insight_similar_top_k
|
||||
|
||||
def new_insight_node(self, insight_key: str, insight_value: str) -> MemoryNode:
|
||||
created_dt = datetime.now()
|
||||
dt = time_to_formatted_str(time=created_dt)
|
||||
|
||||
# 组合meta_data
|
||||
meta_data = {
|
||||
DT: dt,
|
||||
INSIGHT_KEY: insight_key,
|
||||
INSIGHT_VALUE: insight_value,
|
||||
}
|
||||
meta_data.update({k: str(v) for k, v in get_datetime_info_dict(created_dt).items()})
|
||||
|
||||
content = f"用户的{insight_key}:{insight_value}"
|
||||
return MemoryNode.init_from_attrs(content=content,
|
||||
memoryId=self.memory_id,
|
||||
scene=self.scene,
|
||||
memoryType=MemoryTypeEnum.INSIGHT.value,
|
||||
content_modified=True, # 新增的insight需要置为true
|
||||
metaData=meta_data,
|
||||
status=MemoryNodeStatus.ACTIVE.value,
|
||||
tenantId=self.tenant_id)
|
||||
|
||||
def reflect_new_insight_key(self,
|
||||
insight_key: str,
|
||||
not_reflected_merge_nodes: List[MemoryNode]) -> MemoryNode | None:
|
||||
|
||||
# 检索历史memory
|
||||
hits = self.es_client.similar_search(text=insight_key,
|
||||
size=self.es_insight_similar_top_k,
|
||||
exact_filters={
|
||||
"memoryId": self.memory_id,
|
||||
"status": MemoryNodeStatus.ACTIVE.value,
|
||||
"scene": self.scene.lower(),
|
||||
"memoryType": [MemoryTypeEnum.OBSERVATION.value,
|
||||
MemoryTypeEnum.OBS_CUSTOMIZED.value],
|
||||
})
|
||||
|
||||
# 转化成 MemoryNodeWrap 合并新增nodes
|
||||
related_nodes: List[MemoryNode] = [MemoryNode.init_from_es(x) for x in hits]
|
||||
related_nodes.extend(not_reflected_merge_nodes)
|
||||
|
||||
# content去重
|
||||
related_node_dict = {n.memory_node.content: n for n in related_nodes}
|
||||
related_nodes = sorted(list(related_node_dict.values()), key=lambda x: x.memory_node.id)
|
||||
documents = [n.memory_node.content for n in related_nodes]
|
||||
|
||||
# 重排所有记忆
|
||||
result = self.rerank_client.call(query=insight_key, documents=documents)
|
||||
if not result:
|
||||
self.add_run_info(f"reflect insight_key={insight_key} call rerank client failed!")
|
||||
return
|
||||
|
||||
# 根据打分过滤
|
||||
for rank_node in result:
|
||||
index = rank_node["index"]
|
||||
score = rank_node["relevance_score"]
|
||||
related_nodes[index].score_rank = score
|
||||
related_nodes_sorted = sorted(related_nodes, key=lambda x: x.score_rank, reverse=True)[
|
||||
:self.insight_obs_max_cnt]
|
||||
|
||||
# 生成prompt
|
||||
user_query_list = [x.memory_node.content for x in related_nodes_sorted]
|
||||
get_insight_message = self.prompt_to_msg(
|
||||
system_prompt=self.prompt_config.get_insight_system,
|
||||
few_shot=self.prompt_config.get_insight_few_shot,
|
||||
user_query=self.prompt_config.get_insight_user_query.format(
|
||||
insight_key=insight_key, user_query="\n".join(user_query_list)))
|
||||
self.logger.info(f"get_insight_message={get_insight_message}")
|
||||
|
||||
# call LLM, 提取insight
|
||||
response_text = self.gene_client.call(messages=get_insight_message,
|
||||
model_name=self.get_insight_model,
|
||||
max_token=self.get_insight_max_token,
|
||||
temperature=self.get_insight_temperature,
|
||||
top_k=self.get_insight_top_k)
|
||||
|
||||
# return if empty
|
||||
if not response_text:
|
||||
self.add_run_info("reflect_upon_user_attr call llm failed!")
|
||||
return
|
||||
response_text = response_text.strip()
|
||||
if response_text in ["无"]:
|
||||
return
|
||||
return self.new_insight_node(insight_key=insight_key, insight_value=response_text)
|
||||
|
||||
def _run(self):
|
||||
new_insight_keys: List[MemoryNode] = self.get_context(NEW_INSIGHT_KEYS)
|
||||
if not new_insight_keys:
|
||||
self.add_run_info("new_insight_keys is empty! stop insight.")
|
||||
return
|
||||
|
||||
not_reflected_merge_nodes: List[MemoryNode] = self.get_context(NOT_REFLECTED_MERGE_NODES)
|
||||
if not not_reflected_merge_nodes:
|
||||
self.add_run_info("not_reflected_merge_nodes is empty! stop get insight.")
|
||||
return
|
||||
|
||||
# submit insight task
|
||||
for insight_key in new_insight_keys:
|
||||
self.submit_thread(self.reflect_new_insight_key,
|
||||
sleep_time=1,
|
||||
insight_key=insight_key,
|
||||
not_reflected_merge_nodes=not_reflected_merge_nodes)
|
||||
|
||||
# save output
|
||||
new_insight_nodes: List[MemoryNode] = []
|
||||
for result in self.join_threads():
|
||||
if result:
|
||||
new_insight_nodes.append(result)
|
||||
assert isinstance(result, MemoryNode)
|
||||
insight_key = result.memory_node.metaData.get(INSIGHT_KEY, "")
|
||||
insight_value = result.memory_node.metaData.get(INSIGHT_VALUE, "")
|
||||
self.logger.info(f"after_get_insight insight_key={insight_key} insight_value={insight_value}")
|
||||
|
||||
self.set_context(NEW_INSIGHT_NODES, new_insight_nodes)
|
||||
|
||||
# set REFLECTED
|
||||
for node in not_reflected_merge_nodes:
|
||||
scheme.memory_node.metaData[REFLECTED] = "1"
|
||||
|
|
@ -1,80 +0,0 @@
|
|||
from typing import List
|
||||
|
||||
from common.response_text_parser import ResponseTextParser
|
||||
from constants.common_constants import NEW_OBS_NODES, NOT_REFLECTED_OBS_NODES, REFLECTED, INSIGHT_NODES, INSIGHT_KEY, \
|
||||
NEW_INSIGHT_KEYS, NOT_REFLECTED_MERGE_NODES
|
||||
from scheme.memory_node import MemoryNode
|
||||
from worker.memory_base_worker import MemoryBaseWorker
|
||||
|
||||
|
||||
class GetReflectionWorker(MemoryBaseWorker):
|
||||
def __init__(self, reflect_obs_cnt_threshold, reflect_num_questions, reflect_obs_model, reflect_obs_max_token, reflect_obs_temperature, reflect_obs_top_k, *args, **kwargs):
|
||||
super(GetReflectionWorker,self).__init__(*args, **kwargs)
|
||||
self.reflect_obs_cnt_threshold = reflect_obs_cnt_threshold
|
||||
self.reflect_num_questions = reflect_num_questions
|
||||
self.reflect_obs_model = reflect_obs_model
|
||||
self.reflect_obs_max_token = reflect_obs_max_token
|
||||
self.reflect_obs_temperature = reflect_obs_temperature
|
||||
self.reflect_obs_top_k = reflect_obs_top_k
|
||||
|
||||
def _run(self):
|
||||
# 过滤得到 not_reflected_merge_nodes
|
||||
new_obs_nodes: List[MemoryNode] = self.get_context(NEW_OBS_NODES)
|
||||
not_reflected_nodes: List[MemoryNode] = self.get_context(NOT_REFLECTED_OBS_NODES)
|
||||
not_reflected_merge_nodes: List[MemoryNode] = []
|
||||
if new_obs_nodes:
|
||||
not_reflected_merge_nodes.extend(new_obs_nodes)
|
||||
if not_reflected_nodes:
|
||||
not_reflected_merge_nodes.extend(not_reflected_nodes)
|
||||
not_reflected_merge_nodes = [node for node in not_reflected_merge_nodes
|
||||
if scheme.memory_node.metaData.get(REFLECTED, "") == "0"]
|
||||
|
||||
# count
|
||||
not_reflected_count = len(not_reflected_merge_nodes)
|
||||
if not_reflected_count <= self.reflect_obs_cnt_threshold:
|
||||
self.logger.info(f"not_reflected_count={not_reflected_count} is not enough, stop reflect.")
|
||||
return
|
||||
|
||||
# save context
|
||||
self.set_context(NOT_REFLECTED_MERGE_NODES, not_reflected_merge_nodes)
|
||||
|
||||
# get profile_keys
|
||||
exist_keys: List[str] = []
|
||||
profile_keys: List[str] = list(self.user_profile_dict.keys())
|
||||
exist_keys.extend(profile_keys)
|
||||
self.logger.info(f"profile_keys={profile_keys}")
|
||||
|
||||
# get insight_keys
|
||||
insight_nodes: List[MemoryNode] = self.get_context(INSIGHT_NODES)
|
||||
if insight_nodes:
|
||||
insight_keys = [n.memory_node.metaData.get(INSIGHT_KEY) for n in insight_nodes]
|
||||
insight_keys = [x.strip() for x in insight_keys if x]
|
||||
exist_keys.extend(insight_keys)
|
||||
self.logger.info(f"insight_keys={insight_keys}")
|
||||
|
||||
# gen reflect prompt
|
||||
user_query_list = [n.memory_node.content for n in not_reflected_merge_nodes]
|
||||
reflect_message = self.prompt_to_msg(
|
||||
system_prompt=self.prompt_config.get_reflect_system.format(
|
||||
num_questions=self.reflect_num_questions),
|
||||
few_shot=self.prompt_config.get_reflect_few_shot,
|
||||
user_query=self.prompt_config.get_reflect_user_query.format(exist_keys=",".join(exist_keys),
|
||||
user_query="\n".join(user_query_list)))
|
||||
self.logger.info(f"reflect_message={reflect_message}")
|
||||
|
||||
# # call LLM
|
||||
response_text = self.gene_client.call(messages=reflect_message,
|
||||
model_name=self.reflect_obs_model,
|
||||
max_token=self.reflect_obs_max_token,
|
||||
temperature=self.reflect_obs_temperature,
|
||||
top_k=self.reflect_obs_top_k)
|
||||
|
||||
# return if empty
|
||||
if not response_text:
|
||||
self.add_run_info("reflect_obs_questions call llm failed!")
|
||||
return
|
||||
|
||||
# parse text & save
|
||||
new_insight_keys = ResponseTextParser(response_text).parse_v2("get_insight_keys")
|
||||
if new_insight_keys:
|
||||
self.set_context(NEW_INSIGHT_KEYS, new_insight_keys)
|
||||
|
|
@ -1,118 +0,0 @@
|
|||
from typing import List
|
||||
|
||||
from common.response_text_parser import ResponseTextParser
|
||||
from constants.common_constants import NEW_OBS_NODES, TODAY_OBS_NODES, MSG_TIME, NEW_OBS_WITH_TIME_NODES, \
|
||||
MODIFIED_MEMORIES
|
||||
from enumeration.memory_status_enum import MemoryNodeStatus
|
||||
from enumeration.memory_type_enum import MemoryTypeEnum
|
||||
from scheme.memory_node import MemoryNode
|
||||
from worker.memory_base_worker import MemoryBaseWorker
|
||||
|
||||
|
||||
class LongContraRepeatWorker(MemoryBaseWorker):
|
||||
def __init__(es_contra_repeat_similar_top_k, long_contra_repeat_threshold, merge_obs_model, merge_obs_max_token, merge_obs_temperature, merge_obs_top_k, *args, **kwargs):
|
||||
super(LongContraRepeatWorker, self).__init__(*args, **kwargs)
|
||||
self.es_contra_repeat_similar_top_k = es_contra_repeat_similar_top_k
|
||||
self.merge_obs_model = merge_obs_model
|
||||
self.merge_obs_max_token = merge_obs_max_token
|
||||
self.merge_obs_temperature = merge_obs_temperature
|
||||
self.merge_obs_top_k = merge_obs_top_k
|
||||
|
||||
def _run(self):
|
||||
# 合并当前的obs和今日的obs
|
||||
new_obs_nodes: List[MemoryNode] = self.get_context(NEW_OBS_NODES)
|
||||
# new_obs_with_time_nodes: List[MemoryNode] = self.get_context(NEW_OBS_WITH_TIME_NODES)
|
||||
# oday_obs_nodes: List[MemoryNode] = self.get_context(TODAY_OBS_NODES)
|
||||
all_obs_nodes: List[MemoryNode] = []
|
||||
for new_obs_node in new_obs_nodes:
|
||||
text = new_obs_scheme.memory_node.content
|
||||
hits = self.es_client.similar_search(text=text,
|
||||
size=self.es_contra_repeat_similar_top_k,
|
||||
exact_filters={
|
||||
"memoryId": self.memory_id,
|
||||
"status": MemoryNodeStatus.ACTIVE.value,
|
||||
"scene": self.scene.lower(),
|
||||
"memoryType": [MemoryTypeEnum.OBSERVATION.value,
|
||||
MemoryTypeEnum.OBS_CUSTOMIZED.value],
|
||||
})
|
||||
|
||||
related_nodes: List[MemoryNode] = [MemoryNode.init_from_es(x) for x in hits]
|
||||
|
||||
has_match = False
|
||||
for related_node in related_nodes:
|
||||
if related_node.score_similar < self.long_contra_repeat_threshold:
|
||||
continue
|
||||
else:
|
||||
has_match = True
|
||||
all_obs_nodes.append(related_node)
|
||||
if has_match:
|
||||
all_obs_nodes.append(new_obs_node)
|
||||
|
||||
if not all_obs_nodes:
|
||||
self.add_run_info("all_obs_nodes is empty!")
|
||||
return
|
||||
|
||||
# gene prompt
|
||||
user_query_list = []
|
||||
all_obs_nodes = sorted(all_obs_nodes, key=lambda x: x.memory_node.metaData.get(MSG_TIME, ""), reverse=True)
|
||||
for i, n in enumerate(all_obs_nodes):
|
||||
user_query_list.append(f"{i + 1} {n.memory_node.content}")
|
||||
merge_obs_message = self.prompt_to_msg(
|
||||
system_prompt=self.prompt_config.long_contra_repeat_system.format(num_obs=len(user_query_list)),
|
||||
few_shot=self.prompt_config.long_contra_repeat_few_shot,
|
||||
user_query=self.prompt_config.long_contra_repeat_user_query.format(user_query="\n".join(user_query_list)))
|
||||
self.logger.info(f"merge_obs_message={merge_obs_message}")
|
||||
|
||||
# call LLM
|
||||
response_text = self.gene_client.call(messages=merge_obs_message,
|
||||
model_name=self.merge_obs_model,
|
||||
max_token=self.merge_obs_max_token,
|
||||
temperature=self.merge_obs_temperature,
|
||||
top_k=self.merge_obs_top_k)
|
||||
|
||||
# return if empty
|
||||
if not response_text:
|
||||
self.add_run_info("contra repeat call llm failed!")
|
||||
return
|
||||
|
||||
# parse text
|
||||
idx_merge_obs_list = ResponseTextParser(response_text).parse_v1("contra_repeat")
|
||||
if len(idx_merge_obs_list) <= 0:
|
||||
self.add_run_info("idx_merge_obs_list is empty!")
|
||||
return
|
||||
|
||||
# add merged obs
|
||||
merge_obs_nodes: List[MemoryNode] = []
|
||||
for obs_content_list in idx_merge_obs_list:
|
||||
if not obs_content_list:
|
||||
continue
|
||||
|
||||
# [6, 逃课]
|
||||
if len(obs_content_list) != 2:
|
||||
self.logger.warning(f"obs_content_list={obs_content_list} is invalid!")
|
||||
continue
|
||||
|
||||
idx, keep_flag = obs_content_list
|
||||
|
||||
if not idx.isdigit():
|
||||
self.logger.warning(f"idx={idx} is invalid!")
|
||||
continue
|
||||
|
||||
# 序号需要修正-1
|
||||
idx = int(idx) - 1
|
||||
if idx >= len(all_obs_nodes):
|
||||
self.logger.warning(f"idx={idx} is invalid!")
|
||||
continue
|
||||
|
||||
if keep_flag not in ["矛盾", "被包含", "无"]:
|
||||
self.logger.warning(f"keep_flag={keep_flag} is invalid!")
|
||||
continue
|
||||
|
||||
node: MemoryNode = all_obs_nodes[idx]
|
||||
if keep_flag != "无":
|
||||
scheme.memory_node.status = MemoryNodeStatus.EXPIRED.value
|
||||
merge_obs_nodes.append(node)
|
||||
self.logger.info(f"after contra repeat: {scheme.memory_node.content} {scheme.memory_node.status}")
|
||||
|
||||
# save context
|
||||
self.set_context(MODIFIED_MEMORIES, merge_obs_nodes)
|
||||
|
|
@ -1,33 +0,0 @@
|
|||
from typing import List, Dict
|
||||
|
||||
from constants.common_constants import NEW_INSIGHT_NODES, MODIFIED_MEMORIES, INSIGHT_NODES, NEW_OBS_NODES, \
|
||||
NOT_REFLECTED_OBS_NODES, NEW, NOT_REFLECTED_MERGE_NODES
|
||||
from scheme.memory_node import MemoryNode
|
||||
from worker.memory_base_worker import MemoryBaseWorker
|
||||
|
||||
|
||||
class SummaryCollectWorker(MemoryBaseWorker):
|
||||
|
||||
def _run(self):
|
||||
insight_nodes: List[MemoryNode] = self.get_context(INSIGHT_NODES)
|
||||
new_insight_nodes: List[MemoryNode] = self.get_context(NEW_INSIGHT_NODES)
|
||||
new_obs_nodes: List[MemoryNode] = self.get_context(NEW_OBS_NODES)
|
||||
not_reflected_nodes: List[MemoryNode] = self.get_context(NOT_REFLECTED_OBS_NODES)
|
||||
not_reflected_merge_nodes: List[MemoryNode] = self.get_context(NOT_REFLECTED_MERGE_NODES)
|
||||
|
||||
# 合并逻辑,复杂,务必check
|
||||
all_node_dict: Dict[str, MemoryNode] = {}
|
||||
if insight_nodes:
|
||||
all_node_dict.update({n.id: n for n in insight_nodes if n.memory_node.content_modified})
|
||||
if new_insight_nodes:
|
||||
all_node_dict.update({n.memory_node.content: n for n in new_insight_nodes})
|
||||
if new_obs_nodes:
|
||||
# 设置为非新
|
||||
for n in new_obs_nodes:
|
||||
n.memory_node.metaData[NEW] = "0"
|
||||
all_node_dict.update({n.memory_node.content: n for n in new_obs_nodes})
|
||||
if not_reflected_merge_nodes and not_reflected_nodes:
|
||||
# 进入reflect阶段
|
||||
all_node_dict.update({n.id: n for n in not_reflected_nodes})
|
||||
|
||||
self.set_context(MODIFIED_MEMORIES, list(all_node_dict.values()))
|
||||
|
|
@ -1,151 +0,0 @@
|
|||
from typing import List
|
||||
|
||||
from common.response_text_parser import ResponseTextParser
|
||||
from constants.common_constants import INSIGHT_NODES, NEW_OBS_NODES, INSIGHT_KEY, INSIGHT_VALUE
|
||||
from scheme.memory_node import MemoryNode
|
||||
from worker.memory_base_worker import MemoryBaseWorker
|
||||
|
||||
|
||||
class UpdateInsightWorker(MemoryBaseWorker):
|
||||
def __init__(update_insight_threshold, update_insight_max_thread, update_insight_model, update_insight_max_token, update_insight_temperature, update_insight_top_k,*args, **kwargs):
|
||||
super(UpdateInsightWorker, self).__init__(*args, **kwargs)
|
||||
self.update_insight_threshold = update_insight_threshold
|
||||
self.update_insight_max_thread = update_insight_max_thread
|
||||
self.update_insight_model = update_insight_model
|
||||
self.update_insight_max_token = update_insight_max_token
|
||||
self.update_insight_temperature = update_insight_temperature
|
||||
self.update_insight_top_k = update_insight_top_k
|
||||
|
||||
def filter_obs_nodes(self,
|
||||
insight_node: MemoryNode,
|
||||
new_obs_nodes: List[MemoryNode]) -> (MemoryNode, List[MemoryNode], float):
|
||||
max_score: float = 0
|
||||
filtered_nodes: List[MemoryNode] = []
|
||||
|
||||
insight_key = insight_scheme.memory_node.metaData.get(INSIGHT_KEY, "")
|
||||
insight_value = insight_scheme.memory_node.metaData.get(INSIGHT_VALUE, "")
|
||||
if not insight_key or not insight_value:
|
||||
self.logger.warning(f"insight_key={insight_key} insight_value={insight_value} is empty!")
|
||||
return insight_node, filtered_nodes, max_score
|
||||
|
||||
result = self.rerank_client.call(query=insight_key,
|
||||
documents=[x.memory_node.content for x in new_obs_nodes])
|
||||
|
||||
if not result:
|
||||
self.add_run_info(f"update_insight={insight_key} call rerank failed!")
|
||||
return insight_node, filtered_nodes, max_score
|
||||
|
||||
# 找到大于阈值的obs node
|
||||
|
||||
for rank_node in result:
|
||||
index = rank_node["index"]
|
||||
score = rank_node["relevance_score"]
|
||||
node = new_obs_nodes[index]
|
||||
keep_flag = "filtered"
|
||||
if score >= self.update_insight_threshold:
|
||||
filtered_nodes.append(node)
|
||||
keep_flag = "keep"
|
||||
max_score = max(max_score, score)
|
||||
self.logger.info(f"insight_key={insight_key} insight_value={insight_value} "
|
||||
f"score={score} keep_flag={keep_flag}")
|
||||
|
||||
if not filtered_nodes:
|
||||
self.logger.info(f"update_insight={insight_key} filtered_nodes is empty!")
|
||||
|
||||
return insight_node, filtered_nodes, max_score
|
||||
|
||||
def update_insight(self,
|
||||
insight_node: MemoryNode,
|
||||
filtered_nodes: List[MemoryNode]) -> MemoryNode:
|
||||
|
||||
insight_key = insight_scheme.memory_node.metaData.get(INSIGHT_KEY, "")
|
||||
insight_value = insight_scheme.memory_node.metaData.get(INSIGHT_VALUE, "")
|
||||
self.logger.info(f"update_insight insight_key={insight_key} insight_value={insight_value} "
|
||||
f"doc.size={len(filtered_nodes)}")
|
||||
|
||||
# gen prompt
|
||||
user_query_list = []
|
||||
for node in filtered_nodes:
|
||||
user_query_list.append(f"句子:{scheme.memory_node.content}")
|
||||
update_insight_message = self.prompt_to_msg(
|
||||
system_prompt=self.prompt_config.update_insight_system,
|
||||
few_shot=self.prompt_config.update_insight_few_shot,
|
||||
user_query=self.prompt_config.update_insight_user_query.format(
|
||||
user_query="\n".join(user_query_list),
|
||||
insight_key=insight_key,
|
||||
insight_key_value=insight_key + ":" + insight_value))
|
||||
self.logger.info(f"update_insight_message={update_insight_message}")
|
||||
|
||||
# call LLM
|
||||
response_text: str = self.gene_client.call(messages=update_insight_message,
|
||||
model_name=self.update_insight_model,
|
||||
max_token=self.update_insight_max_token,
|
||||
temperature=self.update_insight_temperature,
|
||||
top_k=self.update_insight_top_k)
|
||||
|
||||
# return if empty
|
||||
if not response_text:
|
||||
self.add_run_info(f"update_insight insight_key={insight_key} call llm failed!")
|
||||
return insight_node
|
||||
|
||||
profile_list = ResponseTextParser(response_text).parse_v1(f"update_profile {insight_key}")
|
||||
if not profile_list:
|
||||
self.add_run_info(f"update_insight insight_key={insight_key} profile_list empty 1!")
|
||||
return insight_node
|
||||
profile_list = profile_list[0]
|
||||
if not profile_list:
|
||||
self.add_run_info(f"update_insight insight_key={insight_key} profile_list empty 2")
|
||||
return insight_node
|
||||
insight_value = profile_list[0]
|
||||
|
||||
if not insight_value or insight_value in ["无", "重复"]:
|
||||
self.logger.info(f"insight_value={insight_value}, skip.")
|
||||
return insight_node
|
||||
|
||||
insight_scheme.memory_node.metaData[INSIGHT_VALUE] = insight_value
|
||||
insight_scheme.memory_node.content_modified = True
|
||||
return insight_node
|
||||
|
||||
def _run(self):
|
||||
# 获取新的obs和insight
|
||||
new_obs_nodes: List[MemoryNode] = self.get_context(NEW_OBS_NODES)
|
||||
insight_nodes: List[MemoryNode] = self.get_context(INSIGHT_NODES)
|
||||
if not new_obs_nodes:
|
||||
self.logger.info("new_obs_nodes is empty, stop update sights!")
|
||||
return
|
||||
if not insight_nodes:
|
||||
self.logger.info("insight_nodes is empty, stop update sights!")
|
||||
return
|
||||
|
||||
# 提交打分任务
|
||||
for node in insight_nodes:
|
||||
self.submit_thread(self.filter_obs_nodes,
|
||||
sleep_time=0.1,
|
||||
insight_node=node,
|
||||
new_obs_nodes=new_obs_nodes)
|
||||
|
||||
# 选择topN
|
||||
result_list = []
|
||||
for result in self.join_threads():
|
||||
insight_node, filtered_nodes, max_score = result
|
||||
if not filtered_nodes:
|
||||
continue
|
||||
result_list.append(result)
|
||||
result_sorted = sorted(result_list, key=lambda x: x[2], reverse=True)
|
||||
if len(result_sorted) > self.update_insight_max_thread:
|
||||
result_sorted = result_sorted[:update_insight_max_thread]
|
||||
|
||||
# 提交LLM update任务
|
||||
for insight_node, filtered_nodes, _ in result_sorted:
|
||||
self.submit_thread(self.update_insight,
|
||||
sleep_time=1,
|
||||
insight_node=insight_node,
|
||||
filtered_nodes=filtered_nodes)
|
||||
|
||||
# 等待结果
|
||||
for result in self.join_threads():
|
||||
if result:
|
||||
insight_node: MemoryNode = result
|
||||
insight_key = insight_scheme.memory_node.metaData.get(INSIGHT_KEY, "")
|
||||
insight_value = insight_scheme.memory_node.metaData.get(INSIGHT_VALUE, "")
|
||||
self.logger.info(f"after_update_insight insight_key={insight_key} insight_value={insight_value}")
|
||||
|
|
@ -1,210 +0,0 @@
|
|||
from typing import List
|
||||
|
||||
from common.response_text_parser import ResponseTextParser
|
||||
from constants.common_constants import NEW_OBS_NODES, NEW_USER_PROFILE
|
||||
from enumeration.memory_type_enum import MemoryTypeEnum
|
||||
from scheme.memory_node import MemoryNode
|
||||
from node.user_attribute import UserAttribute
|
||||
from worker.memory_base_worker import MemoryBaseWorker
|
||||
|
||||
|
||||
class UpdateProfileWorker(MemoryBaseWorker):
|
||||
def __init__(update_profile_max_thread, update_profile_threshold, extra_user_attrs, update_profile_model, update_profile_max_token, update_profile_temperature, update_profile_top_k, *args, **kwargs):
|
||||
super(UpdateProfileWorker,self).__init__(*args, **kwargs)
|
||||
self.update_profile_max_thread = update_profile_max_thread
|
||||
self.extra_user_attrs = extra_user_attrs
|
||||
self.update_profile_threshold = update_profile_threshold
|
||||
self.update_profile_model = update_profile_model
|
||||
self.update_profile_max_token = update_profile_max_token
|
||||
self.update_profile_temperature = update_profile_temperature
|
||||
self.update_profile_top_k = update_profile_top_k
|
||||
|
||||
# @property
|
||||
# def extra_user_attrs(self):
|
||||
# return self.request.extra_user_attrs
|
||||
|
||||
def filter_obs_nodes(self,
|
||||
user_attr: UserAttribute,
|
||||
new_obs_nodes: List[MemoryNode]) -> (UserAttribute, List[MemoryNode], float):
|
||||
max_score: float = 0
|
||||
filtered_nodes: List[MemoryNode] = []
|
||||
result = self.rerank_client.call(query=user_attr.description,
|
||||
documents=[x.memory_node.content for x in new_obs_nodes])
|
||||
|
||||
if not result:
|
||||
self.add_run_info(f"update_user_attr={user_attr.memory_key} call rerank failed!")
|
||||
return user_attr, filtered_nodes, max_score
|
||||
|
||||
# 找到大于阈值的obs node
|
||||
filtered_nodes: List[MemoryNode] = []
|
||||
for rank_node in result:
|
||||
index = rank_node["index"]
|
||||
score = rank_node["relevance_score"]
|
||||
node = new_obs_nodes[index]
|
||||
keep_flag = "filtered"
|
||||
if score >= self.update_profile_threshold:
|
||||
filtered_nodes.append(node)
|
||||
keep_flag = "keep"
|
||||
max_score = max(max_score, score)
|
||||
self.logger.info(f"key={user_attr.memory_key} desc={user_attr.description} "
|
||||
f"content={scheme.memory_node.content} score={score} keep_flag={keep_flag}")
|
||||
|
||||
if not filtered_nodes:
|
||||
self.logger.info(f"update_user_attr={user_attr} filtered_nodes is empty!")
|
||||
return user_attr, filtered_nodes, max_score
|
||||
|
||||
def update_user_attr(self, user_attr: UserAttribute, filtered_nodes: List[MemoryNode]) -> UserAttribute:
|
||||
self.logger.info(f"update_user_attr memory_key={user_attr.memory_key} desc={user_attr.description} "
|
||||
f"value={user_attr.value} doc.size={len(filtered_nodes)}")
|
||||
|
||||
# 根据不同的参数类型是否多值,分别给出prompt
|
||||
user_query_list = []
|
||||
for node in filtered_nodes:
|
||||
user_query_list.append(f"句子:{scheme.memory_node.content}")
|
||||
update_profile = f"{user_attr.memory_key}({user_attr.description})"
|
||||
update_profile_value = update_profile + ":" + ",".join(user_attr.value)
|
||||
|
||||
if user_attr.is_unique == 1:
|
||||
update_profile_message = self.prompt_to_msg(
|
||||
system_prompt=self.prompt_config.update_unique_profile_system,
|
||||
few_shot=self.prompt_config.update_unique_profile_few_shot,
|
||||
user_query=self.prompt_config.update_unique_profile_user_query.format(
|
||||
user_query="\n".join(user_query_list),
|
||||
update_profile=update_profile,
|
||||
update_profile_value=update_profile_value))
|
||||
else:
|
||||
update_profile_message = self.prompt_to_msg(
|
||||
system_prompt=self.prompt_config.update_plural_profile_system,
|
||||
few_shot=self.prompt_config.update_plural_profile_few_shot,
|
||||
user_query=self.prompt_config.update_plural_profile_user_query.format(
|
||||
user_query="\n".join(user_query_list),
|
||||
update_profile=update_profile,
|
||||
update_profile_value=update_profile_value))
|
||||
self.logger.info(f"update_profile_message={update_profile_message}")
|
||||
|
||||
# call LLM
|
||||
response_text: str = self.gene_client.call(messages=update_profile_message,
|
||||
model_name=self.update_profile_model,
|
||||
max_token=self.update_profile_max_token,
|
||||
temperature=self.update_profile_temperature,
|
||||
top_k=self.update_profile_top_k)
|
||||
|
||||
# return if empty
|
||||
if not response_text:
|
||||
self.add_run_info(f"update_one_user_attr key={user_attr.memory_key} call llm failed!")
|
||||
return user_attr
|
||||
|
||||
profile_list = ResponseTextParser(response_text).parse_v1(f"update_attr {user_attr.memory_key}")
|
||||
if not profile_list:
|
||||
self.add_run_info(f"update_one_user_attr key={user_attr.memory_key} profile_list empty 1!")
|
||||
return user_attr
|
||||
profile_list = profile_list[0]
|
||||
if not profile_list:
|
||||
self.add_run_info(f"update_one_user_attr key={user_attr.memory_key} profile_list empty 2")
|
||||
return user_attr
|
||||
profile = profile_list[0]
|
||||
|
||||
if not profile or profile in ["无", "重复"]:
|
||||
self.logger.info(f"profile={profile}, skip.")
|
||||
return user_attr
|
||||
|
||||
# check 英文中午逗号
|
||||
if user_attr.is_unique == 1:
|
||||
user_attr.value = [profile.strip()]
|
||||
else:
|
||||
attr_value_list = profile.replace(",", ",").split(",")
|
||||
user_attr.value = [x.strip() for x in sorted(list(set(user_attr.value + attr_value_list)))]
|
||||
return user_attr
|
||||
|
||||
def add_extra_user_attrs(self):
|
||||
# 解析为空返回
|
||||
extra_user_attr_list = [x.strip() for x in self.extra_user_attrs if x.strip()]
|
||||
if not extra_user_attr_list:
|
||||
return
|
||||
|
||||
for user_attr_info in extra_user_attr_list:
|
||||
user_attr_split = user_attr_info.split(":")
|
||||
|
||||
# 格式不对返回
|
||||
if len(user_attr_split) < 1:
|
||||
continue
|
||||
user_attr_key = user_attr_split[0]
|
||||
|
||||
user_attr_desc = ""
|
||||
if len(user_attr_split) >= 2:
|
||||
user_attr_desc = user_attr_split[1]
|
||||
|
||||
user_attr_unique = 0
|
||||
if len(user_attr_split) >= 3:
|
||||
user_attr_unique = int(user_attr_split[2])
|
||||
|
||||
# 已经包含返回
|
||||
if user_attr_key in self.user_profile_dict:
|
||||
user_attr = self.user_profile_dict[user_attr_key]
|
||||
# description为空,补充description
|
||||
if not user_attr.description:
|
||||
user_attr.description = user_attr_desc
|
||||
continue
|
||||
|
||||
# 增加新属性
|
||||
new_attr = UserAttribute(memory_id=self.config.memory_id,
|
||||
scene=self.scene,
|
||||
memory_key=user_attr_key,
|
||||
is_unique=int(user_attr_unique),
|
||||
is_mutable=1,
|
||||
memory_type=MemoryTypeEnum.PROFILE,
|
||||
description=user_attr_desc,
|
||||
status=1)
|
||||
self.user_profile_dict[user_attr_key] = new_attr
|
||||
|
||||
def _run(self):
|
||||
new_obs_nodes: List[MemoryNode] = self.get_context(NEW_OBS_NODES)
|
||||
if not new_obs_nodes:
|
||||
self.logger.info("new_obs_nodes is empty, stop user profile!")
|
||||
self.set_context(NEW_USER_PROFILE, list(self.user_profile_dict.values()))
|
||||
return
|
||||
|
||||
# 增加环境变量配置的属性
|
||||
if self.extra_user_attrs:
|
||||
self.add_extra_user_attrs()
|
||||
|
||||
new_user_profile: List[UserAttribute] = []
|
||||
self.set_context(NEW_USER_PROFILE, new_user_profile)
|
||||
|
||||
for user_attr_key, user_attr in self.user_profile_dict.items():
|
||||
# 不可修改直接跳过
|
||||
if user_attr.is_mutable != 1:
|
||||
new_user_profile.append(user_attr)
|
||||
self.logger.info(f"{user_attr_key} is not mutable! continue")
|
||||
continue
|
||||
|
||||
self.submit_thread(self.filter_obs_nodes,
|
||||
sleep_time=0.1,
|
||||
user_attr=user_attr,
|
||||
new_obs_nodes=new_obs_nodes)
|
||||
|
||||
# 选择topN
|
||||
result_list = []
|
||||
for result in self.join_threads():
|
||||
user_attr, filtered_nodes, max_score = result
|
||||
if not filtered_nodes:
|
||||
continue
|
||||
result_list.append(result)
|
||||
result_sorted = sorted(result_list, key=lambda x: x[2], reverse=True)
|
||||
if len(result_sorted) > self.update_profile_max_thread:
|
||||
result_sorted = result_sorted[:self.update_profile_max_thread]
|
||||
|
||||
# 提交LLM update任务
|
||||
for user_attr, filtered_nodes, _ in result_sorted:
|
||||
self.submit_thread(self.update_user_attr,
|
||||
sleep_time=1,
|
||||
user_attr=user_attr,
|
||||
filtered_nodes=filtered_nodes)
|
||||
|
||||
# collect result & save
|
||||
for result in self.join_threads():
|
||||
if result:
|
||||
user_attribute: UserAttribute = result
|
||||
self.logger.info(f"after_update_profile memory_key={user_attribute.memory_key} "
|
||||
f"desc={user_attribute.description} value={user_attribute.value}")
|
||||
new_user_profile.append(user_attribute)
|
||||
|
|
@ -1,98 +0,0 @@
|
|||
from typing import List
|
||||
|
||||
from common.response_text_parser import ResponseTextParser
|
||||
from constants.common_constants import NEW_OBS_NODES, TODAY_OBS_NODES, MSG_TIME, NEW_OBS_WITH_TIME_NODES, \
|
||||
MODIFIED_MEMORIES
|
||||
from enumeration.memory_status_enum import MemoryNodeStatus
|
||||
from scheme.memory_node import MemoryNode
|
||||
from worker.memory_base_worker import MemoryBaseWorker
|
||||
|
||||
|
||||
class ContraRepeatWorker(MemoryBaseWorker):
|
||||
def __init__(self, merge_obs_model, merge_obs_max_token, merge_obs_temperature, merge_obs_top_k, *args, **kwargs):
|
||||
super(ContraRepeatWorker, self).__init__(*args, **kwargs)
|
||||
self.merge_obs_model = merge_obs_model
|
||||
self.merge_obs_max_token = merge_obs_max_token
|
||||
self.merge_obs_temperature = merge_obs_temperature
|
||||
self.merge_obs_top_k = merge_obs_top_k
|
||||
|
||||
def _run(self):
|
||||
# 合并当前的obs和今日的obs
|
||||
new_obs_nodes: List[MemoryNode] = self.get_context(NEW_OBS_NODES)
|
||||
new_obs_with_time_nodes: List[MemoryNode] = self.get_context(NEW_OBS_WITH_TIME_NODES)
|
||||
today_obs_nodes: List[MemoryNode] = self.get_context(TODAY_OBS_NODES)
|
||||
all_obs_nodes: List[MemoryNode] = []
|
||||
if new_obs_nodes:
|
||||
all_obs_nodes.extend(new_obs_nodes)
|
||||
if new_obs_with_time_nodes:
|
||||
all_obs_nodes.extend(new_obs_with_time_nodes)
|
||||
if today_obs_nodes:
|
||||
all_obs_nodes.extend(today_obs_nodes)
|
||||
if not all_obs_nodes:
|
||||
self.add_run_info("all_obs_nodes is empty!")
|
||||
return
|
||||
|
||||
# gene prompt
|
||||
user_query_list = []
|
||||
all_obs_nodes = sorted(all_obs_nodes, key=lambda x: x.memory_node.metaData.get(MSG_TIME, ""), reverse=True)
|
||||
for i, n in enumerate(all_obs_nodes):
|
||||
user_query_list.append(f"{i + 1} {n.memory_node.content}")
|
||||
merge_obs_message = self.prompt_to_msg(
|
||||
system_prompt=self.prompt_config.contra_repeat_system.format(num_obs=len(user_query_list)),
|
||||
few_shot=self.prompt_config.contra_repeat_few_shot,
|
||||
user_query=self.prompt_config.contra_repeat_user_query.format(user_query="\n".join(user_query_list)))
|
||||
self.logger.info(f"merge_obs_message={merge_obs_message}")
|
||||
|
||||
# call LLM
|
||||
response_text = self.gene_client.call(messages=merge_obs_message,
|
||||
model_name=self.merge_obs_model,
|
||||
max_token=self.merge_obs_max_token,
|
||||
temperature=self.merge_obs_temperature,
|
||||
top_k=self.merge_obs_top_k)
|
||||
|
||||
# return if empty
|
||||
if not response_text:
|
||||
self.add_run_info("contra repeat call llm failed!")
|
||||
return
|
||||
|
||||
# parse text
|
||||
idx_merge_obs_list = ResponseTextParser(response_text).parse_v1("contra_repeat")
|
||||
if len(idx_merge_obs_list) <= 0:
|
||||
self.add_run_info("idx_merge_obs_list is empty!")
|
||||
return
|
||||
|
||||
# add merged obs
|
||||
merge_obs_nodes: List[MemoryNode] = []
|
||||
for obs_content_list in idx_merge_obs_list:
|
||||
if not obs_content_list:
|
||||
continue
|
||||
|
||||
# [6, 逃课]
|
||||
if len(obs_content_list) != 2:
|
||||
self.logger.warning(f"obs_content_list={obs_content_list} is invalid!")
|
||||
continue
|
||||
|
||||
idx, keep_flag = obs_content_list
|
||||
|
||||
if not idx.isdigit():
|
||||
self.logger.warning(f"idx={idx} is invalid!")
|
||||
continue
|
||||
|
||||
# 序号需要修正-1
|
||||
idx = int(idx) - 1
|
||||
if idx >= len(all_obs_nodes):
|
||||
self.logger.warning(f"idx={idx} is invalid!")
|
||||
continue
|
||||
|
||||
if keep_flag not in ["矛盾", "被包含", "无"]:
|
||||
self.logger.warning(f"keep_flag={keep_flag} is invalid!")
|
||||
continue
|
||||
|
||||
node: MemoryNode = all_obs_nodes[idx]
|
||||
if keep_flag != "无":
|
||||
scheme.memory_node.status = MemoryNodeStatus.EXPIRED.value
|
||||
merge_obs_nodes.append(node)
|
||||
self.logger.info(f"after contra repeat: {scheme.memory_node.content} {scheme.memory_node.status}")
|
||||
|
||||
# save context
|
||||
self.set_context(MODIFIED_MEMORIES, merge_obs_nodes)
|
||||
|
|
@ -1,134 +0,0 @@
|
|||
from datetime import datetime
|
||||
from typing import List
|
||||
|
||||
from common.response_text_parser import ResponseTextParser
|
||||
from common.tool_functions import time_to_formatted_str, get_datetime_info_dict, extract_date_parts
|
||||
from constants.common_constants import REFLECTED, DT, TIME_INFER, NEW, MSG_TIME, KEY_WORD, DATATIME_WORD_LIST, \
|
||||
NEW_OBS_WITH_TIME_NODES
|
||||
from enumeration.memory_status_enum import MemoryNodeStatus
|
||||
from enumeration.memory_type_enum import MemoryTypeEnum
|
||||
from scheme.memory_node import MemoryNode
|
||||
from node.message import Message
|
||||
from worker.memory_base_worker import MemoryBaseWorker
|
||||
|
||||
|
||||
class GetObservationWithTimeWorker(MemoryBaseWorker):
|
||||
def __init__(self, summary_messages_model, summary_messages_max_token, summary_messages_temperature, summary_messages_top_k, *args, **kwargs):
|
||||
super(GetObservationWithTimeWorker, self).__init__(*args, **kwargs)
|
||||
self.summary_messages_model = summary_messages_model
|
||||
self.summary_messages_max_token = summary_messages_max_token
|
||||
self.summary_messages_temperature = summary_messages_temperature
|
||||
self.summary_messages_top_k = summary_messages_top_k
|
||||
|
||||
def add_observation(self, message: Message, obs_content: str, time_infer: str, keywords: str):
|
||||
created_dt: datetime = datetime.fromtimestamp(float(message.time_created))
|
||||
dt = time_to_formatted_str(time=created_dt)
|
||||
|
||||
# 组合meta_data
|
||||
meta_data = {
|
||||
MemoryTypeEnum.CONVERSATION.value: message.content, # 原始对话
|
||||
REFLECTED: "0", # reflect标记
|
||||
DT: dt, # 当天标记
|
||||
NEW: "1", # summary-long标记
|
||||
MSG_TIME: message.time_created, # 对话时间
|
||||
TIME_INFER: time_infer, # 推断的时间
|
||||
KEY_WORD: keywords, # 关键词
|
||||
}
|
||||
|
||||
# 事件时间
|
||||
meta_data.update({f"event_{k}": str(v) for k, v in extract_date_parts(time_infer).items()})
|
||||
# 对话时间
|
||||
meta_data.update({f"msg_{k}": str(v) for k, v in get_datetime_info_dict(created_dt).items()})
|
||||
|
||||
return MemoryNode.init_from_attrs(content=obs_content,
|
||||
memoryId=self.memory_id,
|
||||
timeCreated=message.time_created,
|
||||
scene=self.scene,
|
||||
memoryType=MemoryTypeEnum.OBSERVATION.value,
|
||||
content_modified=True, # 新增的obs需要置为true
|
||||
metaData=meta_data,
|
||||
status=MemoryNodeStatus.ACTIVE.value,
|
||||
tenantId=self.tenant_id)
|
||||
|
||||
def _run(self):
|
||||
# gene prompt
|
||||
user_query_list = []
|
||||
i = 1
|
||||
for msg in self.messages:
|
||||
match = False
|
||||
for time_keyword in DATATIME_WORD_LIST:
|
||||
if time_keyword in msg.content:
|
||||
match = True
|
||||
break
|
||||
if match:
|
||||
dt = time_to_formatted_str(time=msg.time_created,
|
||||
date_format="",
|
||||
string_format="{year}年{month}月{day}日{weekday}{hour}点")
|
||||
user_query_list.append(f"{i} {dt} 用户:{msg.content}")
|
||||
i += 1
|
||||
|
||||
if not user_query_list:
|
||||
self.add_run_info(f"get obs with time user_query_list={user_query_list} is empty")
|
||||
return
|
||||
|
||||
obtain_obs_message = self.prompt_to_msg(
|
||||
system_prompt=self.prompt_config.get_observation_with_time_system.format(num_obs=len(user_query_list)),
|
||||
few_shot=self.prompt_config.get_observation_with_time_few_shot,
|
||||
user_query=self.prompt_config.get_observation_with_time_user_query.format(
|
||||
user_query="\n".join(user_query_list)))
|
||||
self.logger.info(f"obtain_obs_message={obtain_obs_message}")
|
||||
|
||||
# call LLM
|
||||
response_text: str = self.gene_client.call(messages=obtain_obs_message,
|
||||
model_name=self.summary_messages_model,
|
||||
max_token=self.summary_messages_max_token,
|
||||
temperature=self.summary_messages_temperature,
|
||||
top_k=self.summary_messages_top_k)
|
||||
|
||||
# return if empty
|
||||
if not response_text:
|
||||
self.add_run_info("summary call llm failed!", continue_run=False)
|
||||
return
|
||||
|
||||
# parse text
|
||||
idx_obs_list = ResponseTextParser(response_text).parse_v1("get_obs_time")
|
||||
if len(idx_obs_list) <= 0:
|
||||
self.add_run_info("idx_obs_list is empty!", continue_run=False)
|
||||
return
|
||||
|
||||
# gene new obs nodes
|
||||
new_obs_nodes: List[MemoryNode] = []
|
||||
for obs_content_list in idx_obs_list:
|
||||
if not obs_content_list:
|
||||
continue
|
||||
|
||||
# [1, 2022年6月, 用户在2022年6月去杭州旅游, 旅游]
|
||||
if len(obs_content_list) != 4:
|
||||
self.logger.warning(f"obs_content_list={obs_content_list} is invalid!")
|
||||
continue
|
||||
|
||||
idx, time_infer, obs_content, keywords = obs_content_list
|
||||
|
||||
if obs_content in ["无", "重复"]:
|
||||
continue
|
||||
|
||||
if not idx.isdigit():
|
||||
self.logger.warning(f"idx={idx} is invalid!")
|
||||
continue
|
||||
|
||||
if time_infer == "无":
|
||||
time_infer = ""
|
||||
|
||||
# 序号需要修正-1
|
||||
idx = int(idx) - 1
|
||||
if idx >= len(self.messages):
|
||||
self.logger.warning(f"idx={idx} is invalid! messages.size={len(self.messages)}")
|
||||
continue
|
||||
|
||||
new_obs_nodes.append(self.add_observation(message=self.messages[idx],
|
||||
obs_content=obs_content,
|
||||
time_infer=time_infer,
|
||||
keywords=keywords))
|
||||
|
||||
# save context
|
||||
self.set_context(NEW_OBS_WITH_TIME_NODES, new_obs_nodes)
|
||||
|
|
@ -1,122 +0,0 @@
|
|||
from datetime import datetime
|
||||
from typing import List
|
||||
|
||||
from common.response_text_parser import ResponseTextParser
|
||||
from common.tool_functions import time_to_formatted_str, get_datetime_info_dict
|
||||
from constants.common_constants import REFLECTED, DT, NEW_OBS_NODES, TIME_INFER, NEW, MSG_TIME, KEY_WORD, \
|
||||
DATATIME_WORD_LIST
|
||||
from enumeration.memory_status_enum import MemoryNodeStatus
|
||||
from enumeration.memory_type_enum import MemoryTypeEnum
|
||||
from scheme.memory_node import MemoryNode
|
||||
from node.message import Message
|
||||
from worker.memory_base_worker import MemoryBaseWorker
|
||||
|
||||
|
||||
class GetObservationWorker(MemoryBaseWorker):
|
||||
def __init__(self, summary_messages_model, summary_messages_max_token, summary_messages_temperature, summary_messages_top_k, *args, **kwargs):
|
||||
super(GetObservationWorker, self).__init__(*args,**kwargs)
|
||||
self.summary_messages_model = summary_messages_model
|
||||
self.summary_messages_max_token = summary_messages_max_token
|
||||
self.summary_messages_temperature = summary_messages_temperature
|
||||
self.summary_messages_top_k = summary_messages_top_k
|
||||
|
||||
def add_observation(self, message: Message, obs_content: str, keywords: str):
|
||||
created_dt: datetime = datetime.fromtimestamp(float(message.time_created))
|
||||
dt = time_to_formatted_str(time=created_dt)
|
||||
|
||||
# 组合meta_data
|
||||
meta_data = {
|
||||
MemoryTypeEnum.CONVERSATION.value: message.content, # 原始对话
|
||||
REFLECTED: "0", # reflect标记
|
||||
DT: dt, # 当天标记
|
||||
NEW: "1", # summary-long标记
|
||||
MSG_TIME: message.time_created, # 对话时间
|
||||
TIME_INFER: "", # 推断的时间
|
||||
KEY_WORD: keywords, # 关键词
|
||||
}
|
||||
meta_data.update({k: str(v) for k, v in get_datetime_info_dict(created_dt).items()})
|
||||
|
||||
return MemoryNode.init_from_attrs(content=obs_content,
|
||||
memoryId=self.memory_id,
|
||||
timeCreated=message.time_created,
|
||||
scene=self.scene,
|
||||
memoryType=MemoryTypeEnum.OBSERVATION.value,
|
||||
content_modified=True, # 新增的obs需要置为true
|
||||
metaData=meta_data,
|
||||
status=MemoryNodeStatus.ACTIVE.value,
|
||||
tenantId=self.tenant_id)
|
||||
|
||||
def _run(self):
|
||||
# gene prompt
|
||||
user_query_list = []
|
||||
i = 1
|
||||
for msg in self.messages:
|
||||
match = False
|
||||
for time_keyword in DATATIME_WORD_LIST:
|
||||
if time_keyword in msg.content:
|
||||
match = True
|
||||
break
|
||||
if not match:
|
||||
user_query_list.append(f"{i} 用户:{msg.content}")
|
||||
i += 1
|
||||
|
||||
if not user_query_list:
|
||||
self.add_run_info(f"get obs user_query_list={user_query_list} is empty")
|
||||
return
|
||||
|
||||
obtain_obs_message = self.prompt_to_msg(
|
||||
system_prompt=self.prompt_config.get_observation_system.format(num_obs=len(user_query_list)),
|
||||
few_shot=self.prompt_config.get_observation_few_shot,
|
||||
user_query=self.prompt_config.get_observation_user_query.format(user_query="\n".join(user_query_list)))
|
||||
self.logger.info(f"obtain_obs_message={obtain_obs_message}")
|
||||
|
||||
# call LLM
|
||||
response_text: str = self.gene_client.call(messages=obtain_obs_message,
|
||||
model_name=self.summary_messages_model,
|
||||
max_token=self.summary_messages_max_token,
|
||||
temperature=self.summary_messages_temperature,
|
||||
top_k=self.summary_messages_top_k)
|
||||
|
||||
# return if empty
|
||||
if not response_text:
|
||||
self.add_run_info("summary call llm failed!", continue_run=False)
|
||||
return
|
||||
|
||||
# parse text
|
||||
idx_obs_list = ResponseTextParser(response_text).parse_v1("get obs")
|
||||
if len(idx_obs_list) <= 0:
|
||||
self.add_run_info("idx_obs_list is empty!", continue_run=False)
|
||||
return
|
||||
|
||||
# gene new obs nodes
|
||||
new_obs_nodes: List[MemoryNode] = []
|
||||
for obs_content_list in idx_obs_list:
|
||||
if not obs_content_list:
|
||||
continue
|
||||
|
||||
# [1, 2022年6月, 用户在2022年6月去杭州旅游, 旅游]
|
||||
if len(obs_content_list) != 4:
|
||||
self.logger.warning(f"obs_content_list={obs_content_list} is invalid!")
|
||||
continue
|
||||
|
||||
idx, time_infer, obs_content, keywords = obs_content_list
|
||||
|
||||
if obs_content in ["无", "重复"]:
|
||||
continue
|
||||
|
||||
if not idx.isdigit():
|
||||
self.logger.warning(f"idx={idx} is invalid!")
|
||||
continue
|
||||
|
||||
# 序号需要修正-1
|
||||
idx = int(idx) - 1
|
||||
if idx >= len(self.messages):
|
||||
self.logger.warning(f"idx={idx} is invalid! messages.size={len(self.messages)}")
|
||||
continue
|
||||
|
||||
new_obs_nodes.append(self.add_observation(message=self.messages[idx],
|
||||
obs_content=obs_content,
|
||||
keywords=keywords))
|
||||
|
||||
# save context
|
||||
self.set_context(NEW_OBS_NODES, new_obs_nodes)
|
||||
|
|
@ -1,64 +0,0 @@
|
|||
from common.response_text_parser import ResponseTextParser
|
||||
from enumeration.message_role_enum import MessageRoleEnum
|
||||
from worker.memory_base_worker import MemoryBaseWorker
|
||||
|
||||
|
||||
class InfoFilterWorker(MemoryBaseWorker):
|
||||
def __init__(self, info_filter_msg_max_size, info_filter_model, info_filter_max_token, info_filter_temperature, info_filter_top_k, *args, **kwargs):
|
||||
super(InfoFilterWorker,self).__init__(*args,**kwargs)
|
||||
self.info_filter_msg_max_size
|
||||
self.info_filter_model = info_filter_model
|
||||
self.info_filter_max_token = info_filter_max_token
|
||||
self.info_filter_temperature = info_filter_temperature
|
||||
self.info_filter_top_k = info_filter_top_k
|
||||
|
||||
def _run(self):
|
||||
# filter user msg
|
||||
info_messages = []
|
||||
for msg in self.messages:
|
||||
if msg.role != MessageRoleEnum.USER.value:
|
||||
continue
|
||||
if len(msg.content) >= self.info_filter_msg_max_size:
|
||||
continue
|
||||
info_messages.append(msg)
|
||||
|
||||
# gene prompt
|
||||
user_query = "\n".join([f"{i + 1} 用户:{msg.content}" for i, msg in enumerate(info_messages)])
|
||||
info_filter_message = self.prompt_to_msg(
|
||||
system_prompt=self.prompt_config.info_filter_system.format(batch_size=len(info_messages)),
|
||||
few_shot=self.prompt_config.info_filter_few_shot,
|
||||
user_query=self.prompt_config.info_filter_user_query.format(user_query=user_query))
|
||||
self.logger.info(f"info_filter_message={info_filter_message}")
|
||||
|
||||
# call llm
|
||||
response_text = self.gene_client.call(messages=info_filter_message,
|
||||
model_name=self.info_filter_model,
|
||||
max_token=self.info_filter_max_token,
|
||||
temperature=self.info_filter_temperature,
|
||||
top_k=self.info_filter_top_k)
|
||||
|
||||
# return if empty
|
||||
if not response_text:
|
||||
self.add_run_info("info score call llm failed!", continue_run=False)
|
||||
return
|
||||
|
||||
# parse text
|
||||
info_score_list = ResponseTextParser(response_text).parse_v1("info_filter")
|
||||
if len(info_score_list) != len(info_messages):
|
||||
self.add_run_info(f"info_score_size != info_messages_size, "
|
||||
f"{len(info_score_list)} vs {len(info_messages)}", continue_run=False)
|
||||
return
|
||||
|
||||
# 过滤value=0的messages
|
||||
filtered_messages = []
|
||||
for msg, info_score in zip(info_messages, info_score_list):
|
||||
if not info_score:
|
||||
continue
|
||||
score = info_score[0]
|
||||
# if score in ("1", "2",):
|
||||
if score in ("2",):
|
||||
msg.info_score = score
|
||||
filtered_messages.append(msg)
|
||||
|
||||
# 后续不会关注为0的msg,直接丢弃
|
||||
self.messages = filtered_messages
|
||||
|
|
@ -1,198 +0,0 @@
|
|||
import re
|
||||
from datetime import datetime
|
||||
from importlib import import_module
|
||||
from typing import Dict, List
|
||||
|
||||
from constants.common_constants import WEEKDAYS
|
||||
|
||||
from enumeration.message_role_enum import MessageRoleEnum
|
||||
|
||||
|
||||
def under_line_to_hump(underline_str):
|
||||
sub = re.sub(r'(_\w)', lambda x: x.group(1)[1].upper(), underline_str)
|
||||
return sub[0:1].upper() + sub[1:]
|
||||
|
||||
|
||||
def parse_response_text_v1(response_text: str) -> dict:
|
||||
"""
|
||||
parse text like:
|
||||
<1> <AAA>
|
||||
<2> <BBB> ddd
|
||||
<4> <CCC> dddd<555>
|
||||
|
||||
result = {1: "AAA", 2: "BBB", 4: "CCC"}
|
||||
"""
|
||||
result_dict: Dict[int, str] = {}
|
||||
|
||||
# 确保第一个数字,后面是string
|
||||
matches = re.findall(r'<(\d+)>\s*<([^>]+)>', response_text.strip())
|
||||
|
||||
# matches 为空返回
|
||||
for key, value in matches:
|
||||
result_dict[int(key)] = value
|
||||
|
||||
return result_dict
|
||||
|
||||
|
||||
def parse_response_text_v2(response_text: str) -> Dict[int, List[str]]:
|
||||
"""
|
||||
parse text like:
|
||||
XXX
|
||||
<1> <AAA> <222>
|
||||
<2> <BBB>
|
||||
<4,5> <CCC>
|
||||
|
||||
result = {1: ["AAA", "222"], 2: "BBB", 4: "CCC"}
|
||||
"""
|
||||
result_dict: Dict[int, List[str]] = {}
|
||||
for line in response_text.strip().split("\n"):
|
||||
if "> <" not in line:
|
||||
continue
|
||||
ll = [x.removeprefix("<").removesuffix(">") for x in line.strip().split("> <")]
|
||||
idx: str = ll[0]
|
||||
values: List[str] = ll[1:]
|
||||
if idx.isdigit():
|
||||
idx_int = int(idx)
|
||||
else:
|
||||
idx_split = idx.split(",")
|
||||
if len(idx_split) == 0:
|
||||
continue
|
||||
idx = idx_split[0]
|
||||
if idx.isdigit():
|
||||
idx_int = int(idx)
|
||||
else:
|
||||
continue
|
||||
if values:
|
||||
result_dict[idx_int] = values
|
||||
|
||||
return result_dict
|
||||
|
||||
|
||||
def parse_response_text_v3(response_text: str) -> List[List[str]]:
|
||||
"""
|
||||
parse text like:
|
||||
XXX
|
||||
<1> <AAA>
|
||||
<2c> <BBB>
|
||||
<41> <CCC> <BBB>
|
||||
|
||||
result = [["1", "AAA"], ["2c", "BBB"], ["41", "CCC", "BBB"]]
|
||||
"""
|
||||
result_list: List[List[str]] = []
|
||||
for line in response_text.strip().split("\n"):
|
||||
if "> <" not in line:
|
||||
continue
|
||||
ll = [x.removeprefix("<").removesuffix(">") for x in line.strip().split("> <")]
|
||||
result_list.append(ll)
|
||||
return result_list
|
||||
|
||||
|
||||
def get_datetime_info_dict(parse_dt: datetime):
|
||||
return {
|
||||
"year": parse_dt.year,
|
||||
"month": parse_dt.month,
|
||||
"day": parse_dt.day,
|
||||
"hour": parse_dt.hour,
|
||||
"minute": parse_dt.minute,
|
||||
"second": parse_dt.second,
|
||||
"week": parse_dt.isocalendar().week,
|
||||
"weekday": WEEKDAYS[parse_dt.isocalendar().weekday - 1],
|
||||
}
|
||||
|
||||
|
||||
def extract_date_parts(input_string: str):
|
||||
# Extending our pattern to handle "每" (every) as a possible value.
|
||||
patterns = {
|
||||
'year': r'(\d+|每)年',
|
||||
'month': r'(\d+|每)月',
|
||||
'day': r'(\d+|每)日',
|
||||
'weekday': r'周([一二三四五六日])?',
|
||||
'hour': r'(\d+)点'
|
||||
}
|
||||
weekday_dict = {"一": 1, "二": 2, "三": 3, "四": 4, "五": 5, "六": 6, "日": 7}
|
||||
extracted_data = {}
|
||||
|
||||
# Search for patterns in the input string and populate the dictionary
|
||||
for key, pattern in patterns.items():
|
||||
match = re.search(pattern, input_string)
|
||||
if match: # If there is a match, include it in the output dictionary
|
||||
if match.group(1) == "每":
|
||||
extracted_data[key] = -1
|
||||
elif match.group(1) in weekday_dict.keys():
|
||||
extracted_data[key] = weekday_dict[match.group(1)]
|
||||
else:
|
||||
extracted_data[key] = int(match.group(1))
|
||||
return extracted_data
|
||||
|
||||
|
||||
def time_to_formatted_str(time: datetime | str | int | float = None,
|
||||
date_format: str = "%Y%m%d", # e.g. %Y%m%d -> "20240528", add %H:%M:%S
|
||||
string_format: str = "") -> str:
|
||||
if isinstance(time, str | int | float):
|
||||
if isinstance(time, str):
|
||||
time = float(time)
|
||||
current_dt = datetime.fromtimestamp(time)
|
||||
elif isinstance(time, datetime):
|
||||
current_dt = time
|
||||
else:
|
||||
current_dt = datetime.now()
|
||||
|
||||
return_str = ""
|
||||
if date_format:
|
||||
return_str = current_dt.strftime(date_format)
|
||||
elif string_format:
|
||||
return_str = string_format.format(**get_datetime_info_dict(current_dt))
|
||||
|
||||
return return_str
|
||||
|
||||
|
||||
def init_instance_by_config(config: dict|object, default_module_path: str = None, try_kwargs: dict = {}, accept_types: type = None):
|
||||
if isinstance(config, accept_types):
|
||||
return config
|
||||
|
||||
import_module(config.pop("path", default_module_path))
|
||||
clazz = getattr(module, config.pop("name"))
|
||||
try:
|
||||
return clazz(**config, **try_kwargs)
|
||||
except:
|
||||
return clazz(**config)
|
||||
|
||||
|
||||
def init_instance_by_config_v2(config: dict, default_clazz_path: str = "", suffix_name: str = "", **kwargs):
|
||||
clazz_path = config.pop("clazz")
|
||||
if not clazz_path:
|
||||
raise RuntimeError("empty clazz_path!")
|
||||
clazz_name_split = clazz_path.split(".")
|
||||
clazz_name: str = clazz_name_split[-1]
|
||||
if suffix_name and not clazz_name.endswith(suffix_name):
|
||||
clazz_name = f"{clazz_name}_{suffix_name}"
|
||||
|
||||
# 构造path
|
||||
clazz_paths = []
|
||||
if default_clazz_path:
|
||||
clazz_paths.append(default_clazz_path)
|
||||
clazz_paths.extend(clazz_name_split[:-1])
|
||||
clazz_paths.append(clazz_name)
|
||||
module = import_module(".".join(clazz_paths))
|
||||
|
||||
cls_name = under_line_to_hump(clazz_name)
|
||||
return getattr(module, cls_name)(**config, **kwargs)
|
||||
|
||||
|
||||
def complete_config_name(config_name: str, suffix: str = ".json"):
|
||||
if not config_name.endswith(suffix):
|
||||
config_name += suffix
|
||||
return config_name
|
||||
|
||||
|
||||
def prompt_to_msg(system_prompt: str, few_shot: str, user_query: str):
|
||||
return [
|
||||
{
|
||||
"role": MessageRoleEnum.SYSTEM.value,
|
||||
"content": system_prompt.strip(),
|
||||
},
|
||||
{
|
||||
"role": MessageRoleEnum.USER.value,
|
||||
"content": "\n".join([x.strip() for x in [few_shot, system_prompt, user_query]])
|
||||
},
|
||||
]
|
||||
|
|
@ -1,33 +0,0 @@
|
|||
from typing import Dict, List
|
||||
|
||||
from pydantic import Field, BaseModel
|
||||
|
||||
|
||||
class UserAttribute(BaseModel):
|
||||
"""
|
||||
用户画像的一条属性,和数据库保持一致,只会选择status为1的属性透传过来。
|
||||
status会透传过来。
|
||||
如果code为空,则为新增,否则是更新。
|
||||
确保请求是10条,返回是原始10条+加上新增的条数(如果可以新增)。只会对正确的请求操作数据库。
|
||||
"""
|
||||
id: str = Field("", description="唯一主键")
|
||||
|
||||
memory_id: str = Field("", description="memory id")
|
||||
|
||||
# 从key改成memory_key
|
||||
memory_key: str = Field("", description="memory key")
|
||||
|
||||
value: List[str] = Field([], description="value")
|
||||
|
||||
is_unique: int = Field(1, description="属性是否唯一,if 1 value只有一个,if 0, value 可以很多个")
|
||||
|
||||
is_mutable: int = Field(1, description="是否可变,if 1,value可变,if 1,不可变(用户定义)")
|
||||
|
||||
memory_type: str = Field("", description="profile, profile_customized")
|
||||
|
||||
description: str = Field("", description="memory id")
|
||||
|
||||
status: int = Field(1,
|
||||
description="0为删除,1为active,状态,算法不感知,只为了保存用户删除的画像,给算法传status为valid的用户画像")
|
||||
|
||||
ext_info: Dict[str, str] = Field({}, description="占位符字典")
|
||||
|
|
@ -1,102 +0,0 @@
|
|||
import json
|
||||
from typing import List, Dict
|
||||
|
||||
from enumeration.memory_status_enum import MemoryNodeStatus
|
||||
from scheme.memory_node import MemoryNode
|
||||
from node.user_attribute import UserAttribute
|
||||
|
||||
|
||||
class UserProfileHandler(object):
|
||||
@classmethod
|
||||
def format_content(cls, key: str, description: str, value: str | List[str] = None):
|
||||
if not key.startswith("用户"):
|
||||
key = f"用户的{key}"
|
||||
|
||||
if not description.startswith("用户"):
|
||||
description = f"用户{description}"
|
||||
|
||||
content = f"{key}({description})"
|
||||
|
||||
if value:
|
||||
if isinstance(value, list):
|
||||
value = ",".join(value)
|
||||
content = f"{content}:{value}"
|
||||
|
||||
return content
|
||||
|
||||
"""
|
||||
提供UserAttribute 和 MemoryNode 的相互转化
|
||||
"""
|
||||
|
||||
@classmethod
|
||||
def to_nodes(cls,
|
||||
user_profile: List[UserAttribute] | Dict[str, UserAttribute] | None = None,
|
||||
split_value: bool = False) -> List[MemoryNode]:
|
||||
|
||||
user_profile_dict: Dict[str, UserAttribute] = {}
|
||||
if user_profile:
|
||||
if isinstance(user_profile, list):
|
||||
for user_attr in user_profile:
|
||||
user_profile_dict[user_attr.memory_key] = user_attr
|
||||
elif isinstance(user_profile, dict):
|
||||
user_profile_dict.update(user_profile)
|
||||
|
||||
user_profile_nodes: List[MemoryNode] = []
|
||||
for _, user_attr in user_profile_dict.items():
|
||||
# 获取id
|
||||
_id = user_attr.code
|
||||
if not _id:
|
||||
_id = f"{user_attr.memory_id}_{user_attr.scene}_profile_{user_attr.memory_key}"
|
||||
|
||||
attr_node = MemoryNode.init_from_attrs(id=_id,
|
||||
code=_id,
|
||||
content="",
|
||||
memoryId=user_attr.memory_id,
|
||||
scene=user_attr.scene,
|
||||
memoryType=user_attr.memory_type,
|
||||
content_modified=True,
|
||||
metaData={
|
||||
"memory_key": user_attr.memory_key,
|
||||
"value": json.dumps(user_attr.value, ensure_ascii=False),
|
||||
"is_unique": str(user_attr.is_unique),
|
||||
"is_mutable": str(user_attr.is_mutable),
|
||||
"description": user_attr.description,
|
||||
"status": MemoryNodeStatus.ACTIVE.value,
|
||||
"ext_info": json.dumps(user_attr.ext_info,
|
||||
ensure_ascii=False),
|
||||
},
|
||||
status=MemoryNodeStatus.ACTIVE.value)
|
||||
|
||||
if split_value:
|
||||
for value in user_attr.value:
|
||||
content = cls.format_content(user_attr.memory_key, user_attr.description, value)
|
||||
attr_node_copy = attr_node.copy(deep=True)
|
||||
attr_node_copy.memory_node.content = content
|
||||
user_profile_nodes.append(attr_node_copy)
|
||||
else:
|
||||
content = cls.format_content(user_attr.memory_key, user_attr.description, user_attr.value)
|
||||
attr_scheme.memory_node.content = content
|
||||
user_profile_nodes.append(attr_node)
|
||||
|
||||
return user_profile_nodes
|
||||
|
||||
@classmethod
|
||||
def to_user_attr(cls, user_profile_nodes: List[MemoryNode]) -> Dict[str, UserAttribute]:
|
||||
user_profile_dict: Dict[str, UserAttribute] = {}
|
||||
|
||||
for node in user_profile_nodes:
|
||||
user_attr = UserAttribute(
|
||||
code=node.id,
|
||||
memory_id=scheme.memory_node.memoryId,
|
||||
scene=scheme.memory_node.scene,
|
||||
memory_key=scheme.memory_node.metaData["memory_key"],
|
||||
value=json.loads(scheme.memory_node.metaData["value"]),
|
||||
is_unique=int(scheme.memory_node.metaData["is_unique"]),
|
||||
is_mutable=int(scheme.memory_node.metaData["is_mutable"]),
|
||||
memory_type=scheme.memory_node.memoryType,
|
||||
description=scheme.memory_node.metaData["description"],
|
||||
status=1 if scheme.memory_node.metaData["status"] == MemoryNodeStatus.ACTIVE.value else 0,
|
||||
ext_info=json.loads(scheme.memory_node.metaData["ext_info"]),
|
||||
)
|
||||
user_profile_dict[user_attr.memory_key] = user_attr
|
||||
return user_profile_dict
|
||||
|
|
@ -1,73 +0,0 @@
|
|||
from typing import Any, Dict
|
||||
|
||||
from ..utils.logger import Logger
|
||||
from ..utils.timer import Timer
|
||||
|
||||
|
||||
class BaseWorker(object):
|
||||
|
||||
def __init__(self, raise_exception: bool = True, **kwargs):
|
||||
super(BaseWorker, self).__init__(**kwargs)
|
||||
# 异常是否继续执行
|
||||
self.raise_exception: bool = raise_exception
|
||||
|
||||
# True 为正常运行,False会结束整个pipeline
|
||||
self.continue_run: bool = True
|
||||
|
||||
# 短name
|
||||
self._name_simple: str = ""
|
||||
|
||||
# 是否多线程环境
|
||||
self.is_multi_thread: bool = False
|
||||
|
||||
# pipeline 上下文
|
||||
self.context_dict: Dict[str, Any] | None = None
|
||||
self.context_lock = None
|
||||
|
||||
# 日志
|
||||
self.logger: Logger = Logger.get_logger()
|
||||
|
||||
# worker 参数
|
||||
self.kwargs: dict = kwargs
|
||||
|
||||
def _run(self):
|
||||
raise NotImplementedError
|
||||
|
||||
def run(self):
|
||||
self.logger.info(f"----- Begin {self.name_simple} -----")
|
||||
with Timer(self.name_simple, log_time=False) as t:
|
||||
if self.raise_exception:
|
||||
self._run()
|
||||
else:
|
||||
try:
|
||||
self._run()
|
||||
except Exception as e:
|
||||
self.logger.exception(f"run {self.name_simple} failed! args={e.args}")
|
||||
|
||||
self.logger.info(f"----- End {self.name_simple} cost={t.cost_str}-----")
|
||||
|
||||
def set_context_dict(self, context_dict: dict, context_lock=None):
|
||||
self.context_dict = context_dict
|
||||
if context_lock is not None:
|
||||
self.context_lock = context_lock
|
||||
self.is_multi_thread = True
|
||||
|
||||
def get_context(self, key: str, default=None):
|
||||
return self.context_dict.get(key, default)
|
||||
|
||||
def set_context(self, key: str, value: Any):
|
||||
if self.is_multi_thread:
|
||||
# add lock to multi thread
|
||||
with self.context_lock:
|
||||
self.context_dict[key] = value
|
||||
else:
|
||||
self.context_dict[key] = value
|
||||
|
||||
def __getattr__(self, key):
|
||||
return self.kwargs[key]
|
||||
|
||||
@property
|
||||
def name_simple(self) -> str:
|
||||
if not self._name_simple:
|
||||
self._name_simple = self.__class__.__name__.replace("Worker", "")
|
||||
return self._name_simple
|
||||
|
|
@ -1,6 +0,0 @@
|
|||
from memory_base_worker import MemoryBaseWorker
|
||||
|
||||
|
||||
class DummyWorker(MemoryBaseWorker):
|
||||
def _run(self):
|
||||
pass
|
||||
|
|
@ -1,22 +0,0 @@
|
|||
from typing import List
|
||||
|
||||
from constants.common_constants import INSIGHT_NODES
|
||||
from enumeration.memory_status_enum import MemoryNodeStatus
|
||||
from enumeration.memory_type_enum import MemoryTypeEnum
|
||||
from scheme.memory_node import MemoryNode
|
||||
from worker.memory_base_worker import MemoryBaseWorker
|
||||
from cli import GLOBAL_CONTEXT
|
||||
|
||||
|
||||
class EsInsightWorker(MemoryBaseWorker):
|
||||
def _run(self):
|
||||
insight_nodes = self.vector_store.retrieve_memories(
|
||||
size=self.kwargs.es_insight_top_k,
|
||||
filter_dict={
|
||||
"memory_id": self.memory_id,
|
||||
"status": MemoryNodeStatus.ACTIVE.value,
|
||||
"memory_type": MemoryTypeEnum.INSIGHT.value,
|
||||
},
|
||||
)
|
||||
self.logger.info(f"insight_nodes.size={len(insight_nodes)}")
|
||||
self.set_context(INSIGHT_NODES, insight_nodes)
|
||||
|
|
@ -1,22 +0,0 @@
|
|||
from typing import List
|
||||
|
||||
from constants.common_constants import NEW, NEW_OBS_NODES
|
||||
from enumeration.memory_status_enum import MemoryNodeStatus
|
||||
from enumeration.memory_type_enum import MemoryTypeEnum
|
||||
from scheme.memory_node import MemoryNode
|
||||
from worker.memory_base_worker import MemoryBaseWorker
|
||||
|
||||
|
||||
class EsNewObsWorker(MemoryBaseWorker):
|
||||
def _run(self):
|
||||
new_obs_nodes = self.vector_store.retrieve_memories(
|
||||
size=self.kwargs.es_new_obs_top_k,
|
||||
filter_dict={
|
||||
"memory_id": self.memory_id,
|
||||
"status": MemoryNodeStatus.ACTIVE.value,
|
||||
"memory_type": MemoryTypeEnum.OBSERVATION.value,
|
||||
f"meta_data.{NEW}": "1",
|
||||
},
|
||||
)
|
||||
self.logger.info(f"es new obs, size={len(new_obs_nodes)}")
|
||||
self.set_context(NEW_OBS_NODES, new_obs_nodes)
|
||||
|
|
@ -1,29 +0,0 @@
|
|||
from typing import List
|
||||
|
||||
from constants.common_constants import REFLECTED, NOT_REFLECTED_OBS_NODES
|
||||
from enumeration.memory_status_enum import MemoryNodeStatus
|
||||
from enumeration.memory_type_enum import MemoryTypeEnum
|
||||
from scheme.memory_node import MemoryNode
|
||||
from worker.memory_base_worker import MemoryBaseWorker
|
||||
|
||||
|
||||
class EsNotReflectedWorker(MemoryBaseWorker):
|
||||
|
||||
def _run(self):
|
||||
|
||||
not_reflected_obs_nodes = self.vector_store.retrieve_memories(
|
||||
size=self.kwargs.es_new_obs_top_k,
|
||||
filter_dict={
|
||||
"memory_id": self.memory_id,
|
||||
"status": MemoryNodeStatus.ACTIVE.value,
|
||||
"memory_type": [
|
||||
MemoryTypeEnum.OBSERVATION.value,
|
||||
MemoryTypeEnum.OBS_CUSTOMIZED.value,
|
||||
],
|
||||
f"meta_data.{REFLECTED}": "0",
|
||||
},
|
||||
)
|
||||
self.logger.info(
|
||||
f"retrieve_not_reflected_obs.size={len(not_reflected_obs_nodes)}"
|
||||
)
|
||||
self.set_context(NOT_REFLECTED_OBS_NODES, not_reflected_obs_nodes)
|
||||
|
|
@ -1,37 +0,0 @@
|
|||
from typing import List
|
||||
|
||||
from constants.common_constants import SIMILAR_OBS_NODES, RECALL_TYPE
|
||||
from enumeration.memory_status_enum import MemoryNodeStatus
|
||||
from enumeration.memory_recall_type import MemoryRecallType
|
||||
from enumeration.memory_type_enum import MemoryTypeEnum
|
||||
from scheme.memory_node import MemoryNode
|
||||
from worker.memory_base_worker import MemoryBaseWorker
|
||||
|
||||
|
||||
class EsSimilarWorker(MemoryBaseWorker):
|
||||
def __init__(self, es_similar_top_k, *args, **kwargs):
|
||||
super(EsSimilarWorker, self).__init__(*args, **kwargs)
|
||||
self.es_similar_top_k = es_similar_top_k
|
||||
|
||||
def _run(self):
|
||||
query = self.messages[-1].content
|
||||
similar_obs_nodes = self.vector_store.retrieve_memories(
|
||||
text=query,
|
||||
size=self.es_similar_top_k,
|
||||
filter_dict={
|
||||
"memory_id": self.memory_id,
|
||||
"status": MemoryNodeStatus.ACTIVE.value,
|
||||
"memory_type": [
|
||||
MemoryTypeEnum.OBSERVATION.value,
|
||||
MemoryTypeEnum.INSIGHT.value,
|
||||
MemoryTypeEnum.OBS_CUSTOMIZED.value,
|
||||
],
|
||||
},
|
||||
)
|
||||
|
||||
for node in similar_obs_nodes:
|
||||
node.meta_data[RECALL_TYPE] = MemoryRecallType.SIMILAR.value
|
||||
self.logger.info(f"similar_obs_nodes.size={len(similar_obs_nodes)}")
|
||||
for node in similar_obs_nodes:
|
||||
self.logger.info(f"node={node.content} score_similar={node.score_similar}")
|
||||
self.set_context(SIMILAR_OBS_NODES, similar_obs_nodes)
|
||||
|
|
@ -1,32 +0,0 @@
|
|||
from typing import List
|
||||
|
||||
from utils.tool_functions import time_to_formatted_str
|
||||
from constants.common_constants import TODAY_OBS_NODES, DT
|
||||
from enumeration.memory_status_enum import MemoryNodeStatus
|
||||
from enumeration.memory_type_enum import MemoryTypeEnum
|
||||
from scheme.memory_node import MemoryNode
|
||||
from worker.memory_base_worker import MemoryBaseWorker
|
||||
|
||||
|
||||
class EsTodayObsWorker(MemoryBaseWorker):
|
||||
def __init__(self, es_today_obs_top_k, *args, **kwargs):
|
||||
super(EsTodayObsWorker, self).__init__(*args, **kwargs)
|
||||
self.es_today_obs_top_k = es_today_obs_top_k
|
||||
|
||||
def _run(self):
|
||||
if not self.messages:
|
||||
self.logger.warning("messages is empty!")
|
||||
return
|
||||
msg_time_created = self.messages[-1].time_created
|
||||
today_obs_nodes = self.vector_store.retrieve_memories(
|
||||
size=self.es_today_obs_top_k,
|
||||
filter_dict={
|
||||
"memory_id": self.memory_id,
|
||||
"status": MemoryNodeStatus.ACTIVE.value,
|
||||
"memory_type": MemoryTypeEnum.OBSERVATION.value,
|
||||
f"meta_Data.{DT}": time_to_formatted_str(msg_time_created),
|
||||
},
|
||||
)
|
||||
|
||||
self.logger.info(f"retrieve_today_obs.size={len(today_obs_nodes)}")
|
||||
self.set_context(TODAY_OBS_NODES, today_obs_nodes)
|
||||
|
|
@ -1,25 +0,0 @@
|
|||
from typing import List, Dict
|
||||
|
||||
from constants import common_constants
|
||||
from enumeration.memory_status_enum import MemoryNodeStatus
|
||||
from enumeration.memory_type_enum import MemoryTypeEnum
|
||||
from scheme.memory_node import MemoryNode
|
||||
from worker.memory_base_worker import MemoryBaseWorker
|
||||
|
||||
|
||||
class LoadProfileWorker(MemoryBaseWorker):
|
||||
|
||||
def _run(self):
|
||||
user_profile_node = self.vector_store(
|
||||
size=10000,
|
||||
filter_dict={
|
||||
"memory_id": self.memory_id,
|
||||
"status": MemoryNodeStatus.ACTIVE.value,
|
||||
"memory_type": [
|
||||
MemoryTypeEnum.PROFILE.value,
|
||||
MemoryTypeEnum.PROFILE_CUSTOMIZED.value,
|
||||
],
|
||||
},
|
||||
)
|
||||
self.set_context(common_constants.USER_PROFILE, user_profile_node)
|
||||
self.logger.info(f"retrieve_user_profile.size={len(user_profile_node)}")
|
||||
|
|
@ -1,74 +0,0 @@
|
|||
import re
|
||||
|
||||
from utils.tool_functions import time_to_formatted_str
|
||||
from constants.common_constants import (
|
||||
DATATIME_WORD_LIST,
|
||||
DATATIME_KEY_MAP,
|
||||
EXTRACT_TIME_DICT,
|
||||
)
|
||||
from worker.memory_base_worker import MemoryBaseWorker
|
||||
|
||||
|
||||
class ExtractTimeWorker(MemoryBaseWorker):
|
||||
# TODO add en version
|
||||
@staticmethod
|
||||
def get_parse_time_prompt(query: str, query_time_str: str):
|
||||
return f"""
|
||||
任务指令:从语句与语句发生的时间,推断并提取语句内容中指向的时间段。回答尽可能完整的时间段。
|
||||
语句:{query}
|
||||
时间:{query_time_str}
|
||||
回答:
|
||||
""".strip()
|
||||
|
||||
def _run(self):
|
||||
# save to context
|
||||
extract_time_dict = {}
|
||||
self.set_context(EXTRACT_TIME_DICT, extract_time_dict)
|
||||
|
||||
# get query & time_created_dt
|
||||
query = self.messages[-1].content
|
||||
time_created = self.messages[-1].time_created
|
||||
|
||||
# find datetime keyword
|
||||
contain_datetime = False
|
||||
for datetime_word in DATATIME_WORD_LIST:
|
||||
if datetime_word in query:
|
||||
contain_datetime = True
|
||||
break
|
||||
if not contain_datetime:
|
||||
self.logger.info(f"contain_datetime={contain_datetime}")
|
||||
return
|
||||
|
||||
# prepare prompt
|
||||
# TODO add en version
|
||||
time_format = "{year}年{month}月{day}日,{year}年第{week}周,{weekday},{hour}时{minute}分{second}秒。"
|
||||
query_time_str = time_to_formatted_str(
|
||||
time=time_created, date_format="", string_format=time_format
|
||||
)
|
||||
extract_time_prompt = self.get_parse_time_prompt(
|
||||
query=query, query_time_str=query_time_str
|
||||
)
|
||||
self.logger.info(f"extract_time_prompt={extract_time_prompt}")
|
||||
|
||||
# call sft model
|
||||
|
||||
response_text = self.generation_model.call(
|
||||
prompt=extract_time_prompt,
|
||||
model_name=self.parse_time_model,
|
||||
max_token=self.parse_time_max_token,
|
||||
temperature=self.parse_time_temperature,
|
||||
top_k=self.parse_time_top_k,
|
||||
)
|
||||
|
||||
# if empty, return
|
||||
if not response_text:
|
||||
return
|
||||
|
||||
# re-match time info to dict
|
||||
pattern = r"-\s*(\S+):(\d+)"
|
||||
matches = re.findall(pattern, response_text)
|
||||
for key, value in matches:
|
||||
if key in DATATIME_KEY_MAP.keys():
|
||||
extract_time_dict[DATATIME_KEY_MAP[key]] = value
|
||||
|
||||
self.logger.info(f"response_text={response_text} filters={extract_time_dict}")
|
||||
|
|
@ -1,120 +0,0 @@
|
|||
from typing import Dict, List
|
||||
|
||||
from constants.common_constants import (
|
||||
RELATED_MEMORIES,
|
||||
EXTRACT_TIME_DICT,
|
||||
ALL_ONLINE_NODES,
|
||||
TIME_MATCHED,
|
||||
)
|
||||
from scheme.memory_node import MemoryNode
|
||||
from worker.memory_base_worker import MemoryBaseWorker
|
||||
|
||||
|
||||
class FuseRerankWorker(MemoryBaseWorker):
|
||||
@staticmethod
|
||||
def format_time_infer(
|
||||
time_infer: str, extract_time_dict: Dict[str, str], meta_data: Dict[str, str]
|
||||
):
|
||||
if time_infer:
|
||||
return time_infer
|
||||
|
||||
time_infer = ""
|
||||
if "year" in extract_time_dict:
|
||||
value = meta_data.get("msg_year")
|
||||
if value:
|
||||
time_infer += f"{value}年"
|
||||
elif value == "-1":
|
||||
time_infer += "每年"
|
||||
|
||||
if "month" in extract_time_dict:
|
||||
value = meta_data.get("msg_month")
|
||||
if value:
|
||||
time_infer += f"{value}月"
|
||||
elif value == "-1":
|
||||
time_infer += "每月"
|
||||
|
||||
if "day" in extract_time_dict:
|
||||
value = meta_data.get("msg_day")
|
||||
if value:
|
||||
time_infer += f"{value}日"
|
||||
elif value == "-1":
|
||||
time_infer += "每日"
|
||||
|
||||
if "weekday" in extract_time_dict:
|
||||
value = meta_data.get("msg_weekday")
|
||||
if value:
|
||||
time_infer += value
|
||||
|
||||
return time_infer
|
||||
|
||||
def _run(self):
|
||||
# 解析时间meta信息
|
||||
extract_time_dict: Dict[str, str] = self.get_context(EXTRACT_TIME_DICT)
|
||||
all_online_nodes: List[MemoryNode] = self.get_context(ALL_ONLINE_NODES)
|
||||
|
||||
if not all_online_nodes:
|
||||
self.add_run_info("all_online_nodes is empty, stop")
|
||||
return
|
||||
|
||||
filtered_nodes = []
|
||||
for node in all_online_nodes:
|
||||
if node.score_rank < self.fuse_score_threshold:
|
||||
continue
|
||||
|
||||
# 根据类型给ratio
|
||||
type_ratio: float = self.fuse_ratio_dict.get(node.memory_type, 0.1)
|
||||
|
||||
# 时间系数,完全匹配才行
|
||||
fuse_time_ratio: float = 1.0
|
||||
match_event_flag = False
|
||||
match_msg_flag = False
|
||||
if extract_time_dict:
|
||||
match_event_flag = True
|
||||
for k, v in extract_time_dict.items():
|
||||
event_value = node.meta_data.get(f"event_{k}", "")
|
||||
if event_value in ["-1", v]:
|
||||
continue
|
||||
else:
|
||||
match_event_flag = False
|
||||
break
|
||||
|
||||
match_msg_flag = True
|
||||
for k, v in extract_time_dict.items():
|
||||
msg_value = node.meta_data.get(f"msg_{k}", "")
|
||||
if msg_value == v:
|
||||
continue
|
||||
else:
|
||||
match_msg_flag = False
|
||||
break
|
||||
|
||||
if match_event_flag or match_msg_flag:
|
||||
fuse_time_ratio = self.fuse_time_ratio
|
||||
node.meta_data[TIME_MATCHED] = "1"
|
||||
|
||||
node.score_rerank = node.score_rank * type_ratio * fuse_time_ratio
|
||||
self.logger.info(
|
||||
f"content={node.content} f_event={int(match_event_flag)} "
|
||||
f"f_msg={int(match_msg_flag)} score_rerank={node.score_rerank}"
|
||||
)
|
||||
filtered_nodes.append(node)
|
||||
|
||||
# get output & save context
|
||||
filtered_nodes = sorted(
|
||||
filtered_nodes, key=lambda x: x.score_rerank, reverse=True
|
||||
)
|
||||
filtered_nodes = filtered_nodes[: self.output_max_count]
|
||||
related_memories: List[str] = []
|
||||
for node in filtered_nodes:
|
||||
content = node.content
|
||||
|
||||
# 如果命中时间逻辑
|
||||
if node.meta_data.get(TIME_MATCHED, "") == "1":
|
||||
time_infer = self.format_time_infer(
|
||||
time_infer="",
|
||||
extract_time_dict=extract_time_dict,
|
||||
meta_data=node.meta_data,
|
||||
)
|
||||
content = f"{time_infer}: {content}"
|
||||
related_memories.append(content)
|
||||
|
||||
self.set_context(RELATED_MEMORIES, related_memories)
|
||||
|
|
@ -1,35 +0,0 @@
|
|||
from typing import List
|
||||
|
||||
from utils.user_profile_handler import UserProfileHandler
|
||||
from constants.common_constants import MODIFIED_MEMORIES, NEW_USER_PROFILE, CONTENT_MODIFIED
|
||||
from scheme.memory_node import MemoryNode
|
||||
from worker.memory_base_worker import MemoryBaseWorker
|
||||
|
||||
|
||||
class MemoryStoreWorker(MemoryBaseWorker):
|
||||
|
||||
def _run(self):
|
||||
modified_memories: List[MemoryNode] | List[MemoryNode] = self.get_context(
|
||||
MODIFIED_MEMORIES
|
||||
)
|
||||
if modified_memories:
|
||||
if isinstance(modified_memories[0], MemoryNode):
|
||||
modified_memories = [n.memory_node for n in modified_memories]
|
||||
|
||||
for n in modified_memories:
|
||||
if not n.id:
|
||||
n.id = f"{n.memory_id}_content_{n.content}"
|
||||
n.code = n.id
|
||||
# TODO add batch insert
|
||||
n.meta_data.pop(CONTENT_MODIFIED)
|
||||
self.vector_store.insert(n)
|
||||
else:
|
||||
self.logger.warning("modified_memories is empty!")
|
||||
|
||||
new_user_profile: List[MemoryNode] = self.get_context(NEW_USER_PROFILE)
|
||||
if new_user_profile:
|
||||
for n in new_user_profile:
|
||||
n.meta_data.pop(CONTENT_MODIFIED)
|
||||
self.vector_store.insert(n)
|
||||
else:
|
||||
self.logger.warning("new_user_profile is empty!")
|
||||
|
|
@ -1,62 +0,0 @@
|
|||
from typing import List, Dict
|
||||
|
||||
from constants.common_constants import SIMILAR_OBS_NODES, RECALL_TYPE, KEYWORD_OBS_NODES, ALL_ONLINE_NODES, \
|
||||
QUERY_KEYWORDS
|
||||
from enumeration.memory_recall_enum import MemoryRecallType
|
||||
from scheme.memory_node import MemoryNode
|
||||
from worker.memory_base_worker import MemoryBaseWorker
|
||||
|
||||
|
||||
class SemanticRankWorker(MemoryBaseWorker):
|
||||
|
||||
def user_profile_to_nodes(self) -> List[MemoryNode]:
|
||||
user_profile_nodes: List[MemoryNode] = self.user_profile_dict
|
||||
for node in user_profile_nodes:
|
||||
# 从画像侧召回
|
||||
node.meta_data[RECALL_TYPE] = MemoryRecallType.PROFILE
|
||||
self.logger.info(f"user profile node={node.content}")
|
||||
return user_profile_nodes
|
||||
|
||||
def _run(self):
|
||||
all_node_dict: Dict[str, MemoryNode] = {}
|
||||
|
||||
# 优先级: similar_obs_nodes < profile_nodes
|
||||
similar_obs_nodes: List[MemoryNode] = self.get_context(SIMILAR_OBS_NODES)
|
||||
if similar_obs_nodes:
|
||||
for node in similar_obs_nodes:
|
||||
all_node_dict[node.content] = node
|
||||
|
||||
profile_nodes: List[MemoryNode] = self.user_profile_to_nodes()
|
||||
if profile_nodes:
|
||||
for node in profile_nodes:
|
||||
all_node_dict[node.content] = node
|
||||
|
||||
if not all_node_dict:
|
||||
self.add_run_info("all_node_dict is empty!", continue_run=False)
|
||||
return
|
||||
|
||||
# call recall model
|
||||
query_keywords = self.get_context(QUERY_KEYWORDS)
|
||||
# TODO 根据效果更改
|
||||
# query: str = "用户:" + self.messages[-1].content
|
||||
query: str = self.messages[-1].content
|
||||
if query_keywords:
|
||||
query_keyword_join = ",".join(query_keywords)
|
||||
query = f"{query} 用户的{query_keyword_join}。"
|
||||
documents = list(all_node_dict.keys())
|
||||
result = self.rank_model.call(query=query, documents=documents)
|
||||
|
||||
if not result:
|
||||
self.add_run_info("semantic call recall model failed!")
|
||||
return
|
||||
|
||||
# set score
|
||||
for index, score in result.rank_scores.items():
|
||||
content = documents[index]
|
||||
node = all_node_dict[content]
|
||||
node.score_rank = score
|
||||
self.logger.info(f"query={query} content={node.content} score_rank={node.score_rank}")
|
||||
|
||||
# save context
|
||||
all_online_nodes: List[MemoryNode] = list(all_node_dict.values())
|
||||
self.set_context(ALL_ONLINE_NODES, all_online_nodes)
|
||||
|
|
@ -1,166 +0,0 @@
|
|||
from datetime import datetime
|
||||
from typing import List
|
||||
|
||||
from ...utils.tool_functions import time_to_formatted_str, get_datetime_info_dict
|
||||
from ...constants.common_constants import (
|
||||
NEW_INSIGHT_NODES,
|
||||
DT,
|
||||
NOT_REFLECTED_MERGE_NODES,
|
||||
NEW_INSIGHT_KEYS,
|
||||
INSIGHT_KEY,
|
||||
INSIGHT_VALUE,
|
||||
REFLECTED,
|
||||
CONTENT_MODIFIED
|
||||
)
|
||||
from ...enumeration.memory_status_enum import MemoryNodeStatus
|
||||
from ...enumeration.memory_type_enum import MemoryTypeEnum
|
||||
from ...scheme.memory_node import MemoryNode
|
||||
from ..memory_base_worker import MemoryBaseWorker
|
||||
from ...prompts.get_insight_prompt import (
|
||||
GET_INSIGHT_FEW_SHOT_PROMPT,
|
||||
GET_INSIGHT_SYSTEM_PROMPT,
|
||||
GET_INSIGHT_USER_QUERY_PROMPT
|
||||
)
|
||||
|
||||
|
||||
class GetInsightWorker(MemoryBaseWorker):
|
||||
def new_insight_node(self, insight_key: str, insight_value: str) -> MemoryNode:
|
||||
created_dt = datetime.now()
|
||||
dt = time_to_formatted_str(time=created_dt)
|
||||
|
||||
# 组合meta_data
|
||||
meta_data = {
|
||||
DT: dt,
|
||||
INSIGHT_KEY: insight_key,
|
||||
INSIGHT_VALUE: insight_value,
|
||||
CONTENT_MODIFIED: True, # 新增的insight需要置为true
|
||||
}
|
||||
meta_data.update(
|
||||
{k: str(v) for k, v in get_datetime_info_dict(created_dt).items()}
|
||||
)
|
||||
|
||||
content = f"用户的{insight_key}:{insight_value}"
|
||||
return MemoryNode(
|
||||
content=content,
|
||||
memory_id=self.memory_id,
|
||||
memory_type=MemoryTypeEnum.INSIGHT.value,
|
||||
meta_data=meta_data,
|
||||
status=MemoryNodeStatus.ACTIVE.value,
|
||||
)
|
||||
|
||||
def reflect_new_insight_key(
|
||||
self, insight_key: str, not_reflected_merge_nodes: List[MemoryNode]
|
||||
) -> MemoryNode | None:
|
||||
|
||||
# 检索历史memory
|
||||
hits = self.vector_store.similar_search(
|
||||
text=insight_key,
|
||||
size=self.es_insight_similar_top_k,
|
||||
exact_filters={
|
||||
"memory_id": self.memory_id,
|
||||
"status": MemoryNodeStatus.ACTIVE.value,
|
||||
"memory_type": [
|
||||
MemoryTypeEnum.OBSERVATION.value,
|
||||
MemoryTypeEnum.OBS_CUSTOMIZED.value,
|
||||
],
|
||||
},
|
||||
)
|
||||
|
||||
# 转化成 MemoryNodeWrap 合并新增nodes
|
||||
related_nodes: List[MemoryNode] = [MemoryNode.init_from_es(x) for x in hits]
|
||||
related_nodes.extend(not_reflected_merge_nodes)
|
||||
|
||||
# content去重
|
||||
related_node_dict = {n.memory_node.content: n for n in related_nodes}
|
||||
related_nodes = sorted(
|
||||
list(related_node_dict.values()), key=lambda x: x.memory_node.id
|
||||
)
|
||||
documents = [n.memory_node.content for n in related_nodes]
|
||||
|
||||
# 重排所有记忆
|
||||
result = self.rank_model.call(query=insight_key, documents=documents)
|
||||
if not result:
|
||||
self.add_run_info(
|
||||
f"reflect insight_key={insight_key} call rerank client failed!"
|
||||
)
|
||||
return
|
||||
|
||||
# 根据打分过滤
|
||||
for rank_node in result:
|
||||
index = rank_node["index"]
|
||||
score = rank_node["relevance_score"]
|
||||
related_nodes[index].score_rank = score
|
||||
related_nodes_sorted = sorted(
|
||||
related_nodes, key=lambda x: x.score_rank, reverse=True
|
||||
)[: self.insight_obs_max_cnt]
|
||||
|
||||
# 生成prompt
|
||||
user_query_list = [x.memory_node.content for x in related_nodes_sorted]
|
||||
get_insight_message = self.prompt_to_msg(
|
||||
system_prompt=self.get_prompt(GET_INSIGHT_SYSTEM_PROMPT),
|
||||
few_shot=self.get_prompt(GET_INSIGHT_FEW_SHOT_PROMPT),
|
||||
user_query=self.get_prompt(GET_INSIGHT_USER_QUERY_PROMPT).format(
|
||||
insight_key=insight_key, user_query="\n".join(user_query_list)
|
||||
),
|
||||
)
|
||||
self.logger.info(f"get_insight_message={get_insight_message}")
|
||||
|
||||
# call LLM, 提取insight
|
||||
response_text = self.generation_model.call(
|
||||
messages=get_insight_message,
|
||||
model_name=self.get_insight_model,
|
||||
max_token=self.get_insight_max_token,
|
||||
temperature=self.get_insight_temperature,
|
||||
top_k=self.get_insight_top_k,
|
||||
)
|
||||
|
||||
# return if empty
|
||||
if not response_text:
|
||||
self.add_run_info("reflect_upon_user_attr call llm failed!")
|
||||
return
|
||||
response_text = response_text.strip()
|
||||
if response_text in ["无"]:
|
||||
return
|
||||
return self.new_insight_node(
|
||||
insight_key=insight_key, insight_value=response_text
|
||||
)
|
||||
|
||||
def _run(self):
|
||||
new_insight_keys: List[MemoryNode] = self.get_context(NEW_INSIGHT_KEYS)
|
||||
if not new_insight_keys:
|
||||
self.add_run_info("new_insight_keys is empty! stop insight.")
|
||||
return
|
||||
|
||||
not_reflected_merge_nodes: List[MemoryNode] = self.get_context(
|
||||
NOT_REFLECTED_MERGE_NODES
|
||||
)
|
||||
if not not_reflected_merge_nodes:
|
||||
self.add_run_info("not_reflected_merge_nodes is empty! stop get insight.")
|
||||
return
|
||||
|
||||
# submit insight task
|
||||
for insight_key in new_insight_keys:
|
||||
self.submit_thread(
|
||||
self.reflect_new_insight_key,
|
||||
sleep_time=1,
|
||||
insight_key=insight_key,
|
||||
not_reflected_merge_nodes=not_reflected_merge_nodes,
|
||||
)
|
||||
|
||||
# save output
|
||||
new_insight_nodes: List[MemoryNode] = []
|
||||
for result in self.join_threads():
|
||||
if result:
|
||||
new_insight_nodes.append(result)
|
||||
assert isinstance(result, MemoryNode)
|
||||
insight_key = result.meta_data.get(INSIGHT_KEY, "")
|
||||
insight_value = result.meta_data.get(INSIGHT_VALUE, "")
|
||||
self.logger.info(
|
||||
f"after_get_insight insight_key={insight_key} insight_value={insight_value}"
|
||||
)
|
||||
|
||||
self.set_context(NEW_INSIGHT_NODES, new_insight_nodes)
|
||||
|
||||
# set REFLECTED
|
||||
for node in not_reflected_merge_nodes:
|
||||
node.meta_data[REFLECTED] = "1"
|
||||
|
|
@ -1,99 +0,0 @@
|
|||
from typing import List
|
||||
|
||||
from ...utilsresponse_text_parser import ResponseTextParser
|
||||
from ...constants.common_constants import (
|
||||
NEW_OBS_NODES,
|
||||
NOT_REFLECTED_OBS_NODES,
|
||||
REFLECTED,
|
||||
INSIGHT_NODES,
|
||||
INSIGHT_KEY,
|
||||
NEW_INSIGHT_KEYS,
|
||||
NOT_REFLECTED_MERGE_NODES,
|
||||
)
|
||||
from ...scheme.memory_node import MemoryNode
|
||||
from ..memory_base_worker import MemoryBaseWorker
|
||||
from ...prompts.get_reflection_prompt import (
|
||||
GET_REFLECTION_FEW_SHOT_PROMPT,
|
||||
GET_REFLECTION_SYSTEM_PROMPT,
|
||||
GET_REFLECTION_USER_QUERY_PROMPT
|
||||
)
|
||||
|
||||
|
||||
class GetReflectionWorker(MemoryBaseWorker):
|
||||
def _run(self):
|
||||
# 过滤得到 not_reflected_merge_nodes
|
||||
new_obs_nodes: List[MemoryNode] = self.get_context(NEW_OBS_NODES)
|
||||
not_reflected_nodes: List[MemoryNode] = self.get_context(
|
||||
NOT_REFLECTED_OBS_NODES
|
||||
)
|
||||
not_reflected_merge_nodes: List[MemoryNode] = []
|
||||
if new_obs_nodes:
|
||||
not_reflected_merge_nodes.extend(new_obs_nodes)
|
||||
if not_reflected_nodes:
|
||||
not_reflected_merge_nodes.extend(not_reflected_nodes)
|
||||
not_reflected_merge_nodes = [
|
||||
node
|
||||
for node in not_reflected_merge_nodes
|
||||
if node.meta_data.get(REFLECTED, "") == "0"
|
||||
]
|
||||
|
||||
# count
|
||||
not_reflected_count = len(not_reflected_merge_nodes)
|
||||
if not_reflected_count <= self.reflect_obs_cnt_threshold:
|
||||
self.logger.info(
|
||||
f"not_reflected_count={not_reflected_count} is not enough, stop reflect."
|
||||
)
|
||||
return
|
||||
|
||||
# save context
|
||||
self.set_context(NOT_REFLECTED_MERGE_NODES, not_reflected_merge_nodes)
|
||||
|
||||
# get profile_keys
|
||||
exist_keys: List[str] = []
|
||||
profile_keys: List[str] = list(self.user_profile_dict.keys())
|
||||
exist_keys.extend(profile_keys)
|
||||
self.logger.info(f"profile_keys={profile_keys}")
|
||||
|
||||
# get insight_keys
|
||||
insight_nodes: List[MemoryNode] = self.get_context(INSIGHT_NODES)
|
||||
if insight_nodes:
|
||||
insight_keys = [
|
||||
n.meta_data.get(INSIGHT_KEY) for n in insight_nodes
|
||||
]
|
||||
insight_keys = [x.strip() for x in insight_keys if x]
|
||||
exist_keys.extend(insight_keys)
|
||||
self.logger.info(f"insight_keys={insight_keys}")
|
||||
|
||||
# gen reflect prompt
|
||||
user_query_list = [n.content for n in not_reflected_merge_nodes]
|
||||
reflect_message = self.prompt_to_msg(
|
||||
system_prompt=self.get_prompt(GET_REFLECTION_SYSTEM_PROMPT).format(
|
||||
num_questions=self.reflect_num_questions
|
||||
),
|
||||
few_shot=self.get_prompt(GET_REFLECTION_FEW_SHOT_PROMPT),
|
||||
user_query=self.get_prompt(GET_REFLECTION_USER_QUERY_PROMPT).format(
|
||||
exist_keys=",".join(exist_keys), user_query="\n".join(user_query_list)
|
||||
),
|
||||
)
|
||||
self.logger.info(f"reflect_message={reflect_message}")
|
||||
|
||||
# # call LLM
|
||||
response_text = self.generation_model.call(
|
||||
messages=reflect_message,
|
||||
model_name=self.reflect_obs_model,
|
||||
max_token=self.reflect_obs_max_token,
|
||||
temperature=self.reflect_obs_temperature,
|
||||
top_k=self.reflect_obs_top_k,
|
||||
)
|
||||
|
||||
# return if empty
|
||||
if not response_text:
|
||||
self.add_run_info("reflect_obs_questions call llm failed!")
|
||||
return
|
||||
|
||||
# parse text & save
|
||||
new_insight_keys = ResponseTextParser(response_text).parse_v2(
|
||||
"get_insight_keys"
|
||||
)
|
||||
if new_insight_keys:
|
||||
self.set_context(NEW_INSIGHT_KEYS, new_insight_keys)
|
||||
|
|
@ -1,129 +0,0 @@
|
|||
from typing import List
|
||||
|
||||
from ...utils.response_text_parser import ResponseTextParser
|
||||
from ...constants.common_constants import (
|
||||
NEW_OBS_NODES,
|
||||
MSG_TIME,
|
||||
MODIFIED_MEMORIES,
|
||||
)
|
||||
from ...enumeration.memory_status_enum import MemoryNodeStatus
|
||||
from ...enumeration.memory_type_enum import MemoryTypeEnum
|
||||
from ...scheme.memory_node import MemoryNode
|
||||
from ..memory_base_worker import MemoryBaseWorker
|
||||
from ...prompts.long_contra_repeat_prompt import (
|
||||
LONG_CONTRA_REPEAT_FEW_SHOT_PROMPT,
|
||||
LONG_CONTRA_REPEAT_SYSTEM_PROMPT,
|
||||
LONG_CONTRA_REPEAT_USER_QUERY_PROMPT,
|
||||
)
|
||||
|
||||
|
||||
class LongContraRepeatWorker(MemoryBaseWorker):
|
||||
|
||||
def _run(self):
|
||||
# 合并当前的obs和今日的obs
|
||||
new_obs_nodes: List[MemoryNode] = self.get_context(NEW_OBS_NODES)
|
||||
all_obs_nodes: List[MemoryNode] = []
|
||||
for new_obs_node in new_obs_nodes:
|
||||
text = new_obs_node.content
|
||||
related_nodes = self.vector_store.similar_search(
|
||||
text=text,
|
||||
size=self.es_contra_repeat_similar_top_k,
|
||||
exact_filters={
|
||||
"memory_id": self.memory_id,
|
||||
"status": MemoryNodeStatus.ACTIVE.value,
|
||||
"memory_type": [
|
||||
MemoryTypeEnum.OBSERVATION.value,
|
||||
MemoryTypeEnum.OBS_CUSTOMIZED.value,
|
||||
],
|
||||
},
|
||||
)
|
||||
|
||||
has_match = False
|
||||
for related_node in related_nodes:
|
||||
if related_node.score_similar < self.long_contra_repeat_threshold:
|
||||
continue
|
||||
else:
|
||||
has_match = True
|
||||
all_obs_nodes.append(related_node)
|
||||
if has_match:
|
||||
all_obs_nodes.append(new_obs_node)
|
||||
|
||||
if not all_obs_nodes:
|
||||
self.add_run_info("all_obs_nodes is empty!")
|
||||
return
|
||||
|
||||
# gene prompt
|
||||
user_query_list = []
|
||||
all_obs_nodes = sorted(
|
||||
all_obs_nodes,
|
||||
key=lambda x: x.meta_data.get(MSG_TIME, ""),
|
||||
reverse=True,
|
||||
)
|
||||
for i, n in enumerate(all_obs_nodes):
|
||||
user_query_list.append(f"{i + 1} {n.content}")
|
||||
merge_obs_message = self.prompt_to_msg(
|
||||
system_prompt=self.get_prompt(LONG_CONTRA_REPEAT_SYSTEM_PROMPT).format(
|
||||
num_obs=len(user_query_list)
|
||||
),
|
||||
few_shot=self.get_prompt(LONG_CONTRA_REPEAT_FEW_SHOT_PROMPT),
|
||||
user_query=self.get_prompt(LONG_CONTRA_REPEAT_USER_QUERY_PROMPT).format(
|
||||
user_query="\n".join(user_query_list)
|
||||
),
|
||||
)
|
||||
self.logger.info(f"merge_obs_message={merge_obs_message}")
|
||||
|
||||
# call LLM
|
||||
response_text = self.generation_model.call(
|
||||
messages=merge_obs_message,
|
||||
model_name=self.merge_obs_model,
|
||||
max_token=self.merge_obs_max_token,
|
||||
temperature=self.merge_obs_temperature,
|
||||
top_k=self.merge_obs_top_k,
|
||||
)
|
||||
|
||||
# return if empty
|
||||
if not response_text:
|
||||
self.add_run_info("contra repeat call llm failed!")
|
||||
return
|
||||
|
||||
# parse text
|
||||
idx_merge_obs_list = ResponseTextParser(response_text).parse_v1("contra_repeat")
|
||||
if len(idx_merge_obs_list) <= 0:
|
||||
self.add_run_info("idx_merge_obs_list is empty!")
|
||||
return
|
||||
|
||||
# add merged obs
|
||||
merge_obs_nodes: List[MemoryNode] = []
|
||||
for obs_content_list in idx_merge_obs_list:
|
||||
if not obs_content_list:
|
||||
continue
|
||||
|
||||
# [6, 逃课]
|
||||
if len(obs_content_list) != 2:
|
||||
self.logger.warning(f"obs_content_list={obs_content_list} is invalid!")
|
||||
continue
|
||||
|
||||
idx, keep_flag = obs_content_list
|
||||
|
||||
if not idx.isdigit():
|
||||
self.logger.warning(f"idx={idx} is invalid!")
|
||||
continue
|
||||
|
||||
# 序号需要修正-1
|
||||
idx = int(idx) - 1
|
||||
if idx >= len(all_obs_nodes):
|
||||
self.logger.warning(f"idx={idx} is invalid!")
|
||||
continue
|
||||
|
||||
if keep_flag not in ["矛盾", "被包含", "无"]:
|
||||
self.logger.warning(f"keep_flag={keep_flag} is invalid!")
|
||||
continue
|
||||
|
||||
node: MemoryNode = all_obs_nodes[idx]
|
||||
if keep_flag != "无":
|
||||
node.status = MemoryNodeStatus.EXPIRED.value
|
||||
merge_obs_nodes.append(node)
|
||||
self.logger.info(f"after contra repeat: {node.content} {node.status}")
|
||||
|
||||
# save context
|
||||
self.set_context(MODIFIED_MEMORIES, merge_obs_nodes)
|
||||
|
|
@ -1,47 +0,0 @@
|
|||
from typing import List, Dict
|
||||
|
||||
from ...constants.common_constants import (
|
||||
NEW_INSIGHT_NODES,
|
||||
MODIFIED_MEMORIES,
|
||||
INSIGHT_NODES,
|
||||
NEW_OBS_NODES,
|
||||
NOT_REFLECTED_OBS_NODES,
|
||||
NEW,
|
||||
NOT_REFLECTED_MERGE_NODES,
|
||||
CONTENT_MODIFIED,
|
||||
)
|
||||
from ...scheme.memory_node import MemoryNode
|
||||
from ..memory_base_worker import MemoryBaseWorker
|
||||
|
||||
|
||||
class SummaryCollectWorker(MemoryBaseWorker):
|
||||
|
||||
def _run(self):
|
||||
insight_nodes: List[MemoryNode] = self.get_context(INSIGHT_NODES)
|
||||
new_insight_nodes: List[MemoryNode] = self.get_context(NEW_INSIGHT_NODES)
|
||||
new_obs_nodes: List[MemoryNode] = self.get_context(NEW_OBS_NODES)
|
||||
not_reflected_nodes: List[MemoryNode] = self.get_context(
|
||||
NOT_REFLECTED_OBS_NODES
|
||||
)
|
||||
not_reflected_merge_nodes: List[MemoryNode] = self.get_context(
|
||||
NOT_REFLECTED_MERGE_NODES
|
||||
)
|
||||
|
||||
# 合并逻辑,复杂,务必check
|
||||
all_node_dict: Dict[str, MemoryNode] = {}
|
||||
if insight_nodes:
|
||||
all_node_dict.update(
|
||||
{n.id: n for n in insight_nodes if n.meta_data.get(CONTENT_MODIFIED, False)}
|
||||
)
|
||||
if new_insight_nodes:
|
||||
all_node_dict.update({n.content: n for n in new_insight_nodes})
|
||||
if new_obs_nodes:
|
||||
# 设置为非新
|
||||
for n in new_obs_nodes:
|
||||
n.meta_data[NEW] = "0"
|
||||
all_node_dict.update({n.content: n for n in new_obs_nodes})
|
||||
if not_reflected_merge_nodes and not_reflected_nodes:
|
||||
# 进入reflect阶段
|
||||
all_node_dict.update({n.id: n for n in not_reflected_nodes})
|
||||
|
||||
self.set_context(MODIFIED_MEMORIES, list(all_node_dict.values()))
|
||||
|
|
@ -1,177 +0,0 @@
|
|||
from typing import List
|
||||
|
||||
from ...utils.response_text_parser import ResponseTextParser
|
||||
from ...constants.common_constants import (
|
||||
INSIGHT_NODES,
|
||||
NEW_OBS_NODES,
|
||||
INSIGHT_KEY,
|
||||
INSIGHT_VALUE,
|
||||
CONTENT_MODIFIED,
|
||||
)
|
||||
from ...scheme.memory_node import MemoryNode
|
||||
from ..memory_base_worker import MemoryBaseWorker
|
||||
from ...prompts.update_insight_prompt import (
|
||||
UPDATE_INSIGHT_FEW_SHOT_PROMPT,
|
||||
UPDATE_INSIGHT_SYSTEM_PROMPT,
|
||||
UPDATE_INSIGHT_USER_QUERY_PROMPT,
|
||||
)
|
||||
|
||||
|
||||
class UpdateInsightWorker(MemoryBaseWorker):
|
||||
|
||||
def filter_obs_nodes(
|
||||
self, insight_node: MemoryNode, new_obs_nodes: List[MemoryNode]
|
||||
) -> (MemoryNode, List[MemoryNode], float):
|
||||
max_score: float = 0
|
||||
filtered_nodes: List[MemoryNode] = []
|
||||
|
||||
insight_key = insight_node.meta_data.get(INSIGHT_KEY, "")
|
||||
insight_value = insight_node.meta_data.get(INSIGHT_VALUE, "")
|
||||
if not insight_key or not insight_value:
|
||||
self.logger.warning(
|
||||
f"insight_key={insight_key} insight_value={insight_value} is empty!"
|
||||
)
|
||||
return insight_node, filtered_nodes, max_score
|
||||
|
||||
result = self.rank_model.call(
|
||||
query=insight_key, documents=[x.content for x in new_obs_nodes]
|
||||
)
|
||||
|
||||
if not result:
|
||||
self.add_run_info(f"update_insight={insight_key} call rerank failed!")
|
||||
return insight_node, filtered_nodes, max_score
|
||||
|
||||
# 找到大于阈值的obs node
|
||||
|
||||
for index, score in result.rank_scores.items():
|
||||
node = new_obs_nodes[index]
|
||||
keep_flag = "filtered"
|
||||
if score >= self.update_insight_threshold:
|
||||
filtered_nodes.append(node)
|
||||
keep_flag = "keep"
|
||||
max_score = max(max_score, score)
|
||||
self.logger.info(
|
||||
f"insight_key={insight_key} insight_value={insight_value} "
|
||||
f"score={score} keep_flag={keep_flag}"
|
||||
)
|
||||
|
||||
if not filtered_nodes:
|
||||
self.logger.info(f"update_insight={insight_key} filtered_nodes is empty!")
|
||||
|
||||
return insight_node, filtered_nodes, max_score
|
||||
|
||||
def update_insight(
|
||||
self, insight_node: MemoryNode, filtered_nodes: List[MemoryNode]
|
||||
) -> MemoryNode:
|
||||
|
||||
insight_key = insight_node.meta_data.get(INSIGHT_KEY, "")
|
||||
insight_value = insight_node.meta_data.get(INSIGHT_VALUE, "")
|
||||
self.logger.info(
|
||||
f"update_insight insight_key={insight_key} insight_value={insight_value} "
|
||||
f"doc.size={len(filtered_nodes)}"
|
||||
)
|
||||
|
||||
# gen prompt
|
||||
user_query_list = []
|
||||
for node in filtered_nodes:
|
||||
user_query_list.append(f"句子:{node.content}")
|
||||
update_insight_message = self.prompt_to_msg(
|
||||
system_prompt=self.get_prompt(UPDATE_INSIGHT_SYSTEM_PROMPT),
|
||||
few_shot=self.get_prompt(UPDATE_INSIGHT_FEW_SHOT_PROMPT),
|
||||
user_query=self.get_prompt(UPDATE_INSIGHT_USER_QUERY_PROMPT).format(
|
||||
user_query="\n".join(user_query_list),
|
||||
insight_key=insight_key,
|
||||
insight_key_value=insight_key + ":" + insight_value,
|
||||
),
|
||||
)
|
||||
self.logger.info(f"update_insight_message={update_insight_message}")
|
||||
|
||||
# call LLM
|
||||
response_text: str = self.generation_model.call(
|
||||
messages=update_insight_message,
|
||||
model_name=self.update_insight_model,
|
||||
max_token=self.update_insight_max_token,
|
||||
temperature=self.update_insight_temperature,
|
||||
top_k=self.update_insight_top_k,
|
||||
)
|
||||
|
||||
# return if empty
|
||||
if not response_text:
|
||||
self.add_run_info(
|
||||
f"update_insight insight_key={insight_key} call llm failed!"
|
||||
)
|
||||
return insight_node
|
||||
|
||||
profile_list = ResponseTextParser(response_text).parse_v1(
|
||||
f"update_profile {insight_key}"
|
||||
)
|
||||
if not profile_list:
|
||||
self.add_run_info(
|
||||
f"update_insight insight_key={insight_key} profile_list empty 1!"
|
||||
)
|
||||
return insight_node
|
||||
profile_list = profile_list[0]
|
||||
if not profile_list:
|
||||
self.add_run_info(
|
||||
f"update_insight insight_key={insight_key} profile_list empty 2"
|
||||
)
|
||||
return insight_node
|
||||
insight_value = profile_list[0]
|
||||
|
||||
if not insight_value or insight_value in ["无", "重复"]:
|
||||
self.logger.info(f"insight_value={insight_value}, skip.")
|
||||
return insight_node
|
||||
|
||||
insight_node.meta_data[INSIGHT_VALUE] = insight_value
|
||||
insight_node.meta_data[CONTENT_MODIFIED] = True
|
||||
return insight_node
|
||||
|
||||
def _run(self):
|
||||
# 获取新的obs和insight
|
||||
new_obs_nodes: List[MemoryNode] = self.get_context(NEW_OBS_NODES)
|
||||
insight_nodes: List[MemoryNode] = self.get_context(INSIGHT_NODES)
|
||||
if not new_obs_nodes:
|
||||
self.logger.info("new_obs_nodes is empty, stop update sights!")
|
||||
return
|
||||
if not insight_nodes:
|
||||
self.logger.info("insight_nodes is empty, stop update sights!")
|
||||
return
|
||||
|
||||
# 提交打分任务
|
||||
for node in insight_nodes:
|
||||
self.submit_thread(
|
||||
self.filter_obs_nodes,
|
||||
sleep_time=0.1,
|
||||
insight_node=node,
|
||||
new_obs_nodes=new_obs_nodes,
|
||||
)
|
||||
|
||||
# 选择topN
|
||||
result_list = []
|
||||
for result in self.join_threads():
|
||||
insight_node, filtered_nodes, max_score = result
|
||||
if not filtered_nodes:
|
||||
continue
|
||||
result_list.append(result)
|
||||
result_sorted = sorted(result_list, key=lambda x: x[2], reverse=True)
|
||||
if len(result_sorted) > self.update_insight_max_thread:
|
||||
result_sorted = result_sorted[: self.update_insight_max_thread]
|
||||
|
||||
# 提交LLM update任务
|
||||
for insight_node, filtered_nodes, _ in result_sorted:
|
||||
self.submit_thread(
|
||||
self.update_insight,
|
||||
sleep_time=1,
|
||||
insight_node=insight_node,
|
||||
filtered_nodes=filtered_nodes,
|
||||
)
|
||||
|
||||
# 等待结果
|
||||
for result in self.join_threads():
|
||||
if result:
|
||||
insight_node: MemoryNode = result
|
||||
insight_key = insight_node.meta_data.get(INSIGHT_KEY, "")
|
||||
insight_value = insight_node.meta_data.get(INSIGHT_VALUE, "")
|
||||
self.logger.info(
|
||||
f"after_update_insight insight_key={insight_key} insight_value={insight_value}"
|
||||
)
|
||||
|
|
@ -1,241 +0,0 @@
|
|||
from typing import List
|
||||
|
||||
from ...utils.response_text_parser import ResponseTextParser
|
||||
from ...constants.common_constants import NEW_OBS_NODES, NEW_USER_PROFILE
|
||||
from ...enumeration.memory_type_enum import MemoryTypeEnum
|
||||
from ....memory_node import MemoryNode
|
||||
from ..memory_base_worker import MemoryBaseWorker
|
||||
from ...prompts.update_profile_prompt import (
|
||||
UPDATE_PLURAL_PROFILE_FEW_SHOT_PROMPT,
|
||||
UPDATE_PLURAL_PROFILE_SYSTEM_PROMPT,
|
||||
UPDATE_PLURAL_PROFILE_USER_QUERY_PROMPT,
|
||||
UPDATE_UNIQUE_PROFILE_FEW_SHOT_PROMPT,
|
||||
UPDATE_UNIQUE_PROFILE_SYSTEM_PROMPT,
|
||||
UPDATE_UNIQUE_PROFILE_USER_QUERY_PROMPT
|
||||
)
|
||||
from ...chat.global_context import GlobalContext
|
||||
|
||||
|
||||
class UpdateProfileWorker(MemoryBaseWorker):
|
||||
@property
|
||||
def extra_user_attrs(self):
|
||||
return GlobalContext.global_configs.get("extra_user_attrs", [])
|
||||
|
||||
def filter_obs_nodes(
|
||||
self, user_attr: MemoryNode, new_obs_nodes: List[MemoryNode]
|
||||
) -> (MemoryNode, List[MemoryNode], float):
|
||||
max_score: float = 0
|
||||
filtered_nodes: List[MemoryNode] = []
|
||||
result = self.rank_model.call(
|
||||
query=user_attr.meta_data.get("description", ""),
|
||||
documents=[x.content for x in new_obs_nodes],
|
||||
)
|
||||
|
||||
if not result:
|
||||
self.add_run_info(
|
||||
f"update_user_attr={user_attr.meta_data.get("memory_key", "")} call rerank failed!"
|
||||
)
|
||||
return user_attr, filtered_nodes, max_score
|
||||
|
||||
# 找到大于阈值的obs node
|
||||
filtered_nodes: List[MemoryNode] = []
|
||||
for index, score in result.rank_scores.items():
|
||||
node = new_obs_nodes[index]
|
||||
keep_flag = "filtered"
|
||||
if score >= self.update_profile_threshold:
|
||||
filtered_nodes.append(node)
|
||||
keep_flag = "keep"
|
||||
max_score = max(max_score, score)
|
||||
self.logger.info(
|
||||
f"key={user_attr.meta_data.get("memory_key", "")} desc={user_attr.meta_data.get("description", "")} "
|
||||
f"content={node.content} score={score} keep_flag={keep_flag}"
|
||||
)
|
||||
|
||||
if not filtered_nodes:
|
||||
self.logger.info(f"update_user_attr={user_attr} filtered_nodes is empty!")
|
||||
return user_attr, filtered_nodes, max_score
|
||||
|
||||
def update_user_attr(
|
||||
self, user_attr: MemoryNode, filtered_nodes: List[MemoryNode]
|
||||
) -> MemoryNode:
|
||||
self.logger.info(
|
||||
f"update_user_attr memory_key={user_attr.meta_data.get("memory_key", "")} desc={user_attr.meta_data.get("description", "")} "
|
||||
f"value={user_attr.meta_data.get("value", "")} doc.size={len(filtered_nodes)}"
|
||||
)
|
||||
|
||||
# 根据不同的参数类型是否多值,分别给出prompt
|
||||
user_query_list = []
|
||||
for node in filtered_nodes:
|
||||
user_query_list.append(f"句子:{node.content}")
|
||||
update_profile = f"{user_attr.meta_data.get("memory_key", "")}({user_attr.meta_data.get("description", "")})"
|
||||
update_profile_value = update_profile + ":" + ",".join(user_attr.meta_data.get("value", ""))
|
||||
|
||||
if user_attr.meta_data.get("is_unique", 0) == 1:
|
||||
update_profile_message = self.prompt_to_msg(
|
||||
system_prompt=self.get_prompt(UPDATE_UNIQUE_PROFILE_SYSTEM_PROMPT),
|
||||
few_shot=self.get_prompt(UPDATE_UNIQUE_PROFILE_FEW_SHOT_PROMPT),
|
||||
user_query=self.get_prompt(UPDATE_UNIQUE_PROFILE_USER_QUERY_PROMPT).format(
|
||||
user_query="\n".join(user_query_list),
|
||||
update_profile=update_profile,
|
||||
update_profile_value=update_profile_value,
|
||||
),
|
||||
)
|
||||
else:
|
||||
update_profile_message = self.prompt_to_msg(
|
||||
system_prompt=self.get_prompt(UPDATE_PLURAL_PROFILE_SYSTEM_PROMPT),
|
||||
few_shot=self.get_prompt(UPDATE_PLURAL_PROFILE_FEW_SHOT_PROMPT),
|
||||
user_query=self.get_prompt(UPDATE_PLURAL_PROFILE_USER_QUERY_PROMPT).format(
|
||||
user_query="\n".join(user_query_list),
|
||||
update_profile=update_profile,
|
||||
update_profile_value=update_profile_value,
|
||||
),
|
||||
)
|
||||
self.logger.info(f"update_profile_message={update_profile_message}")
|
||||
|
||||
# call LLM
|
||||
response_text: str = self.generation_model.call(
|
||||
messages=update_profile_message,
|
||||
model_name=self.update_profile_model,
|
||||
max_token=self.update_profile_max_token,
|
||||
temperature=self.update_profile_temperature,
|
||||
top_k=self.update_profile_top_k,
|
||||
)
|
||||
|
||||
# return if empty
|
||||
if not response_text:
|
||||
self.add_run_info(
|
||||
f"update_one_user_attr key={user_attr.meta_data.get("memory_key", "")} call llm failed!"
|
||||
)
|
||||
return user_attr
|
||||
|
||||
profile_list = ResponseTextParser(response_text).parse_v1(
|
||||
f"update_attr {user_attr.meta_data.get("memory_key", "")}"
|
||||
)
|
||||
if not profile_list:
|
||||
self.add_run_info(
|
||||
f"update_one_user_attr key={user_attr.meta_data.get("memory_key", "")} profile_list empty 1!"
|
||||
)
|
||||
return user_attr
|
||||
profile_list = profile_list[0]
|
||||
if not profile_list:
|
||||
self.add_run_info(
|
||||
f"update_one_user_attr key={user_attr.meta_data.get("memory_key", "")} profile_list empty 2"
|
||||
)
|
||||
return user_attr
|
||||
profile = profile_list[0]
|
||||
|
||||
if not profile or profile in ["无", "重复"]:
|
||||
self.logger.info(f"profile={profile}, skip.")
|
||||
return user_attr
|
||||
|
||||
# check 英文中午逗号
|
||||
if user_attr.meta_data.get("is_unique", 0) == 1:
|
||||
user_attr.meta_data["value"] = [profile.strip()]
|
||||
else:
|
||||
attr_value_list = profile.replace(",", ",").split(",")
|
||||
user_attr.meta_data["value"] = [
|
||||
x.strip() for x in sorted(list(set(user_attr.meta_data.get("value", "") + attr_value_list)))
|
||||
]
|
||||
return user_attr
|
||||
|
||||
def add_extra_user_attrs(self):
|
||||
# 解析为空返回
|
||||
extra_user_attr_list = [x.strip() for x in self.extra_user_attrs if x.strip()]
|
||||
if not extra_user_attr_list:
|
||||
return
|
||||
|
||||
for user_attr_info in extra_user_attr_list:
|
||||
user_attr_split = user_attr_info.split(":")
|
||||
|
||||
# 格式不对返回
|
||||
if len(user_attr_split) < 1:
|
||||
continue
|
||||
user_attr_key = user_attr_split[0]
|
||||
|
||||
user_attr_desc = ""
|
||||
if len(user_attr_split) >= 2:
|
||||
user_attr_desc = user_attr_split[1]
|
||||
|
||||
user_attr_unique = 0
|
||||
if len(user_attr_split) >= 3:
|
||||
user_attr_unique = int(user_attr_split[2])
|
||||
|
||||
# 已经包含返回
|
||||
if user_attr_key in self.user_profile_dict:
|
||||
user_attr = self.user_profile_dict[user_attr_key]
|
||||
# description为空,补充description
|
||||
if not user_attr.meta_data.get("description", ""):
|
||||
user_attr.meta_data["description"] = user_attr_desc
|
||||
continue
|
||||
|
||||
# 增加新属性
|
||||
new_attr = MemoryNode(
|
||||
memory_id=self.memory_id,
|
||||
meta_data={
|
||||
"memory_key": user_attr_key,
|
||||
"is_unique": int(user_attr_unique),
|
||||
"is_mutable": 1,
|
||||
"description": user_attr_desc
|
||||
},
|
||||
memory_type=MemoryTypeEnum.PROFILE,
|
||||
status=1,
|
||||
)
|
||||
self.user_profile_dict[user_attr_key] = new_attr
|
||||
|
||||
def _run(self):
|
||||
new_obs_nodes: List[MemoryNode] = self.get_context(NEW_OBS_NODES)
|
||||
if not new_obs_nodes:
|
||||
self.logger.info("new_obs_nodes is empty, stop user profile!")
|
||||
self.set_context(NEW_USER_PROFILE, list(self.user_profile_dict.values()))
|
||||
return
|
||||
|
||||
# 增加环境变量配置的属性
|
||||
if self.extra_user_attrs:
|
||||
self.add_extra_user_attrs()
|
||||
|
||||
new_user_profile: List[MemoryNode] = []
|
||||
self.set_context(NEW_USER_PROFILE, new_user_profile)
|
||||
|
||||
for user_attr_key, user_attr in self.user_profile_dict.items():
|
||||
# 不可修改直接跳过
|
||||
if user_attr.meta_data.get("is_mutable", 0) != 1:
|
||||
new_user_profile.append(user_attr)
|
||||
self.logger.info(f"{user_attr_key} is not mutable! continue")
|
||||
continue
|
||||
|
||||
self.submit_thread(
|
||||
self.filter_obs_nodes,
|
||||
sleep_time=0.1,
|
||||
user_attr=user_attr,
|
||||
new_obs_nodes=new_obs_nodes,
|
||||
)
|
||||
|
||||
# 选择topN
|
||||
result_list = []
|
||||
for result in self.join_threads():
|
||||
user_attr, filtered_nodes, max_score = result
|
||||
if not filtered_nodes:
|
||||
continue
|
||||
result_list.append(result)
|
||||
result_sorted = sorted(result_list, key=lambda x: x[2], reverse=True)
|
||||
if len(result_sorted) > self.update_profile_max_thread:
|
||||
result_sorted = result_sorted[: self.update_profile_max_thread]
|
||||
|
||||
# 提交LLM update任务
|
||||
for user_attr, filtered_nodes, _ in result_sorted:
|
||||
self.submit_thread(
|
||||
self.update_user_attr,
|
||||
sleep_time=1,
|
||||
user_attr=user_attr,
|
||||
filtered_nodes=filtered_nodes,
|
||||
)
|
||||
|
||||
# collect result & save
|
||||
for result in self.join_threads():
|
||||
if result:
|
||||
user_attribute: MemoryNode = result
|
||||
self.logger.info(
|
||||
f"after_update_profile memory_key={user_attribute.meta_data.get("memory_key", "")} "
|
||||
f"desc={user_attribute.meta_data.get("description", "")} value={user_attribute.meta_data.get("value", "")}"
|
||||
)
|
||||
new_user_profile.append(user_attribute)
|
||||
|
|
@ -1,117 +0,0 @@
|
|||
from typing import List
|
||||
|
||||
from ...utils.response_text_parser import ResponseTextParser
|
||||
from ...constants.common_constants import (
|
||||
NEW_OBS_NODES,
|
||||
TODAY_OBS_NODES,
|
||||
MSG_TIME,
|
||||
NEW_OBS_WITH_TIME_NODES,
|
||||
MODIFIED_MEMORIES,
|
||||
)
|
||||
from ...enumeration.memory_status_enum import MemoryNodeStatus
|
||||
from ...scheme.memory_node import MemoryNode
|
||||
from ..memory_base_worker import MemoryBaseWorker
|
||||
from ...prompts.contra_repeat_prompt import (
|
||||
CONTRA_REPEAT_FEW_SHOT_PROMPT,
|
||||
CONTRA_REPEAT_SYSTEM_PROMPT,
|
||||
CONTRA_REPEAT_USER_QUERY_PROMPT,
|
||||
)
|
||||
|
||||
|
||||
class ContraRepeatWorker(MemoryBaseWorker):
|
||||
|
||||
def _run(self):
|
||||
# 合并当前的obs和今日的obs
|
||||
new_obs_nodes: List[MemoryNode] = self.get_context(NEW_OBS_NODES)
|
||||
new_obs_with_time_nodes: List[MemoryNode] = self.get_context(
|
||||
NEW_OBS_WITH_TIME_NODES
|
||||
)
|
||||
today_obs_nodes: List[MemoryNode] = self.get_context(TODAY_OBS_NODES)
|
||||
all_obs_nodes: List[MemoryNode] = []
|
||||
if new_obs_nodes:
|
||||
all_obs_nodes.extend(new_obs_nodes)
|
||||
if new_obs_with_time_nodes:
|
||||
all_obs_nodes.extend(new_obs_with_time_nodes)
|
||||
if today_obs_nodes:
|
||||
all_obs_nodes.extend(today_obs_nodes)
|
||||
if not all_obs_nodes:
|
||||
self.add_run_info("all_obs_nodes is empty!")
|
||||
return
|
||||
|
||||
# gene prompt
|
||||
user_query_list = []
|
||||
all_obs_nodes = sorted(
|
||||
all_obs_nodes,
|
||||
key=lambda x: x.meta_data.get(MSG_TIME, ""),
|
||||
reverse=True,
|
||||
)
|
||||
for i, n in enumerate(all_obs_nodes):
|
||||
user_query_list.append(f"{i + 1} {n.content}")
|
||||
merge_obs_message = self.prompt_to_msg(
|
||||
system_prompt=self.get_prompt(CONTRA_REPEAT_SYSTEM_PROMPT).format(
|
||||
num_obs=len(user_query_list)
|
||||
),
|
||||
few_shot=self.get_prompt(CONTRA_REPEAT_FEW_SHOT_PROMPT),
|
||||
user_query=self.get_prompt(CONTRA_REPEAT_USER_QUERY_PROMPT).format(
|
||||
user_query="\n".join(user_query_list)
|
||||
),
|
||||
)
|
||||
self.logger.info(f"merge_obs_message={merge_obs_message}")
|
||||
|
||||
# call LLM
|
||||
response_text = self.generation_model.call(
|
||||
messages=merge_obs_message,
|
||||
model_name=self.merge_obs_model,
|
||||
max_token=self.merge_obs_max_token,
|
||||
temperature=self.merge_obs_temperature,
|
||||
top_k=self.merge_obs_top_k,
|
||||
)
|
||||
|
||||
# return if empty
|
||||
if not response_text:
|
||||
self.add_run_info("contra repeat call llm failed!")
|
||||
return
|
||||
|
||||
# parse text
|
||||
idx_merge_obs_list = ResponseTextParser(response_text).parse_v1("contra_repeat")
|
||||
if len(idx_merge_obs_list) <= 0:
|
||||
self.add_run_info("idx_merge_obs_list is empty!")
|
||||
return
|
||||
|
||||
# add merged obs
|
||||
merge_obs_nodes: List[MemoryNode] = []
|
||||
for obs_content_list in idx_merge_obs_list:
|
||||
if not obs_content_list:
|
||||
continue
|
||||
|
||||
# [6, 逃课]
|
||||
if len(obs_content_list) != 2:
|
||||
self.logger.warning(f"obs_content_list={obs_content_list} is invalid!")
|
||||
continue
|
||||
|
||||
idx, keep_flag = obs_content_list
|
||||
|
||||
if not idx.isdigit():
|
||||
self.logger.warning(f"idx={idx} is invalid!")
|
||||
continue
|
||||
|
||||
# 序号需要修正-1
|
||||
idx = int(idx) - 1
|
||||
if idx >= len(all_obs_nodes):
|
||||
self.logger.warning(f"idx={idx} is invalid!")
|
||||
continue
|
||||
|
||||
if keep_flag not in ["矛盾", "被包含", "无"]:
|
||||
self.logger.warning(f"keep_flag={keep_flag} is invalid!")
|
||||
continue
|
||||
|
||||
node: MemoryNode = all_obs_nodes[idx]
|
||||
if keep_flag != "无":
|
||||
node.status = MemoryNodeStatus.EXPIRED.value
|
||||
merge_obs_nodes.append(node)
|
||||
self.logger.info(
|
||||
f"after contra repeat: {node.content} {node.status}"
|
||||
)
|
||||
|
||||
# save context
|
||||
self.set_context(MODIFIED_MEMORIES, merge_obs_nodes)
|
||||
|
|
@ -1,167 +0,0 @@
|
|||
from datetime import datetime
|
||||
from typing import List
|
||||
|
||||
from ...utils.response_text_parser import ResponseTextParser
|
||||
from ...utils.tool_functions import (
|
||||
time_to_formatted_str,
|
||||
get_datetime_info_dict,
|
||||
extract_date_parts,
|
||||
)
|
||||
from ...constants.common_constants import (
|
||||
REFLECTED,
|
||||
DT,
|
||||
TIME_INFER,
|
||||
NEW,
|
||||
MSG_TIME,
|
||||
KEY_WORD,
|
||||
DATATIME_WORD_LIST,
|
||||
NEW_OBS_WITH_TIME_NODES,
|
||||
CONTENT_MODIFIED,
|
||||
)
|
||||
from ...enumeration.memory_status_enum import MemoryNodeStatus
|
||||
from ...enumeration.memory_type_enum import MemoryTypeEnum
|
||||
from ...scheme.memory_node import MemoryNode
|
||||
from ...scheme.message import Message
|
||||
from ..memory_base_worker import MemoryBaseWorker
|
||||
from ...prompts.get_observation_with_time_prompt import (
|
||||
GET_OBSERVATION_WITH_TIME_FEW_SHOT_PROMPT,
|
||||
GET_OBSERVATION_WITH_TIME_SYSTEM_PROMPT,
|
||||
GET_OBSERVATION_WITH_TIME_USER_QUERY_PROMPT,
|
||||
)
|
||||
|
||||
|
||||
class GetObservationWithTimeWorker(MemoryBaseWorker):
|
||||
|
||||
def add_observation(
|
||||
self, message: Message, obs_content: str, time_infer: str, keywords: str
|
||||
):
|
||||
created_dt: datetime = datetime.fromtimestamp(float(message.time_created))
|
||||
dt = time_to_formatted_str(time=created_dt)
|
||||
|
||||
# 组合meta_data
|
||||
meta_data = {
|
||||
MemoryTypeEnum.CONVERSATION.value: message.content, # 原始对话
|
||||
REFLECTED: "0", # reflect标记
|
||||
DT: dt, # 当天标记
|
||||
NEW: "1", # summary-long标记
|
||||
MSG_TIME: message.time_created, # 对话时间
|
||||
TIME_INFER: time_infer, # 推断的时间
|
||||
KEY_WORD: keywords, # 关键词
|
||||
CONTENT_MODIFIED: True, # 新增的obs需要置为true
|
||||
}
|
||||
|
||||
# 事件时间
|
||||
meta_data.update(
|
||||
{f"event_{k}": str(v) for k, v in extract_date_parts(time_infer).items()}
|
||||
)
|
||||
# 对话时间
|
||||
meta_data.update(
|
||||
{f"msg_{k}": str(v) for k, v in get_datetime_info_dict(created_dt).items()}
|
||||
)
|
||||
|
||||
return MemoryNode.init_from_attrs(
|
||||
content=obs_content,
|
||||
memory_id=self.memory_id,
|
||||
memory_type=MemoryTypeEnum.OBSERVATION.value,
|
||||
meta_data=meta_data,
|
||||
status=MemoryNodeStatus.ACTIVE.value,
|
||||
)
|
||||
|
||||
def _run(self):
|
||||
# gene prompt
|
||||
user_query_list = []
|
||||
i = 1
|
||||
for msg in self.messages:
|
||||
match = False
|
||||
for time_keyword in DATATIME_WORD_LIST:
|
||||
if time_keyword in msg.content:
|
||||
match = True
|
||||
break
|
||||
if match:
|
||||
dt = time_to_formatted_str(
|
||||
time=msg.time_created,
|
||||
date_format="",
|
||||
string_format="{year}年{month}月{day}日{weekday}{hour}点",
|
||||
)
|
||||
user_query_list.append(f"{i} {dt} 用户:{msg.content}")
|
||||
i += 1
|
||||
|
||||
if not user_query_list:
|
||||
self.add_run_info(
|
||||
f"get obs with time user_query_list={user_query_list} is empty"
|
||||
)
|
||||
return
|
||||
|
||||
obtain_obs_message = self.prompt_to_msg(
|
||||
system_prompt=self.get_prompt(
|
||||
GET_OBSERVATION_WITH_TIME_SYSTEM_PROMPT
|
||||
).format(num_obs=len(user_query_list)),
|
||||
few_shot=self.get_prompt(GET_OBSERVATION_WITH_TIME_FEW_SHOT_PROMPT),
|
||||
user_query=self.get_prompt(
|
||||
GET_OBSERVATION_WITH_TIME_USER_QUERY_PROMPT
|
||||
).format(user_query="\n".join(user_query_list)),
|
||||
)
|
||||
self.logger.info(f"obtain_obs_message={obtain_obs_message}")
|
||||
|
||||
# call LLM
|
||||
response_text: str = self.generation_model.call(
|
||||
messages=obtain_obs_message,
|
||||
model_name=self.summary_messages_model,
|
||||
max_token=self.summary_messages_max_token,
|
||||
temperature=self.summary_messages_temperature,
|
||||
top_k=self.summary_messages_top_k,
|
||||
)
|
||||
|
||||
# return if empty
|
||||
if not response_text:
|
||||
self.add_run_info("summary call llm failed!", continue_run=False)
|
||||
return
|
||||
|
||||
# parse text
|
||||
idx_obs_list = ResponseTextParser(response_text).parse_v1("get_obs_time")
|
||||
if len(idx_obs_list) <= 0:
|
||||
self.add_run_info("idx_obs_list is empty!", continue_run=False)
|
||||
return
|
||||
|
||||
# gene new obs nodes
|
||||
new_obs_nodes: List[MemoryNode] = []
|
||||
for obs_content_list in idx_obs_list:
|
||||
if not obs_content_list:
|
||||
continue
|
||||
|
||||
# [1, 2022年6月, 用户在2022年6月去杭州旅游, 旅游]
|
||||
if len(obs_content_list) != 4:
|
||||
self.logger.warning(f"obs_content_list={obs_content_list} is invalid!")
|
||||
continue
|
||||
|
||||
idx, time_infer, obs_content, keywords = obs_content_list
|
||||
|
||||
if obs_content in ["无", "重复"]:
|
||||
continue
|
||||
|
||||
if not idx.isdigit():
|
||||
self.logger.warning(f"idx={idx} is invalid!")
|
||||
continue
|
||||
|
||||
if time_infer == "无":
|
||||
time_infer = ""
|
||||
|
||||
# 序号需要修正-1
|
||||
idx = int(idx) - 1
|
||||
if idx >= len(self.messages):
|
||||
self.logger.warning(
|
||||
f"idx={idx} is invalid! messages.size={len(self.messages)}"
|
||||
)
|
||||
continue
|
||||
|
||||
new_obs_nodes.append(
|
||||
self.add_observation(
|
||||
message=self.messages[idx],
|
||||
obs_content=obs_content,
|
||||
time_infer=time_infer,
|
||||
keywords=keywords,
|
||||
)
|
||||
)
|
||||
|
||||
# save context
|
||||
self.set_context(NEW_OBS_WITH_TIME_NODES, new_obs_nodes)
|
||||
|
|
@ -1,144 +0,0 @@
|
|||
from datetime import datetime
|
||||
from typing import List
|
||||
|
||||
from ...utils.response_text_parser import ResponseTextParser
|
||||
from ...utils.tool_functions import time_to_formatted_str, get_datetime_info_dict
|
||||
from ...constants.common_constants import (
|
||||
REFLECTED,
|
||||
DT,
|
||||
NEW_OBS_NODES,
|
||||
TIME_INFER,
|
||||
NEW,
|
||||
MSG_TIME,
|
||||
KEY_WORD,
|
||||
DATATIME_WORD_LIST,
|
||||
CONTENT_MODIFIED,
|
||||
)
|
||||
from ...enumeration.memory_status_enum import MemoryNodeStatus
|
||||
from ...enumeration.memory_type_enum import MemoryTypeEnum
|
||||
from ...scheme.memory_node import MemoryNode
|
||||
from ...scheme.message import Message
|
||||
from ..memory_base_worker import MemoryBaseWorker
|
||||
from ...prompts.get_observation_prompt import (
|
||||
GET_OBSERVATION_FEW_SHOT_PROMPT,
|
||||
GET_OBSERVATION_SYSTEM_PROMPT,
|
||||
GET_OBSERVATION_USER_QUERY_PROMPT,
|
||||
)
|
||||
|
||||
|
||||
class GetObservationWorker(MemoryBaseWorker):
|
||||
|
||||
def add_observation(self, message: Message, obs_content: str, keywords: str):
|
||||
created_dt: datetime = datetime.fromtimestamp(float(message.time_created))
|
||||
dt = time_to_formatted_str(time=created_dt)
|
||||
|
||||
# 组合meta_data
|
||||
meta_data = {
|
||||
MemoryTypeEnum.CONVERSATION.value: message.content, # 原始对话
|
||||
REFLECTED: "0", # reflect标记
|
||||
DT: dt, # 当天标记
|
||||
NEW: "1", # summary-long标记
|
||||
MSG_TIME: message.time_created, # 对话时间
|
||||
TIME_INFER: "", # 推断的时间
|
||||
KEY_WORD: keywords, # 关键词
|
||||
CONTENT_MODIFIED: True, # 新增的obs需要置为true
|
||||
}
|
||||
meta_data.update(
|
||||
{k: str(v) for k, v in get_datetime_info_dict(created_dt).items()}
|
||||
)
|
||||
|
||||
return MemoryNode(
|
||||
content=obs_content,
|
||||
memory_id=self.memory_id,
|
||||
memory_type=MemoryTypeEnum.OBSERVATION.value,
|
||||
meta_data=meta_data,
|
||||
status=MemoryNodeStatus.ACTIVE.value,
|
||||
)
|
||||
|
||||
def _run(self):
|
||||
# gene prompt
|
||||
user_query_list = []
|
||||
i = 1
|
||||
for msg in self.messages:
|
||||
match = False
|
||||
for time_keyword in DATATIME_WORD_LIST:
|
||||
if time_keyword in msg.content:
|
||||
match = True
|
||||
break
|
||||
if not match:
|
||||
user_query_list.append(f"{i} 用户:{msg.content}")
|
||||
i += 1
|
||||
|
||||
if not user_query_list:
|
||||
self.add_run_info(f"get obs user_query_list={user_query_list} is empty")
|
||||
return
|
||||
|
||||
obtain_obs_message = self.prompt_to_msg(
|
||||
system_prompt=self.get_prompt(GET_OBSERVATION_SYSTEM_PROMPT).format(
|
||||
num_obs=len(user_query_list)
|
||||
),
|
||||
few_shot=self.get_prompt(GET_OBSERVATION_FEW_SHOT_PROMPT),
|
||||
user_query=self.get_prompt(GET_OBSERVATION_USER_QUERY_PROMPT).format(
|
||||
user_query="\n".join(user_query_list)
|
||||
),
|
||||
)
|
||||
self.logger.info(f"obtain_obs_message={obtain_obs_message}")
|
||||
|
||||
# call LLM
|
||||
response_text: str = self.generation_model.call(
|
||||
messages=obtain_obs_message,
|
||||
model_name=self.summary_messages_model,
|
||||
max_token=self.summary_messages_max_token,
|
||||
temperature=self.summary_messages_temperature,
|
||||
top_k=self.summary_messages_top_k,
|
||||
)
|
||||
|
||||
# return if empty
|
||||
if not response_text:
|
||||
self.add_run_info("summary call llm failed!", continue_run=False)
|
||||
return
|
||||
|
||||
# parse text
|
||||
idx_obs_list = ResponseTextParser(response_text).parse_v1("get obs")
|
||||
if len(idx_obs_list) <= 0:
|
||||
self.add_run_info("idx_obs_list is empty!", continue_run=False)
|
||||
return
|
||||
|
||||
# gene new obs nodes
|
||||
new_obs_nodes: List[MemoryNode] = []
|
||||
for obs_content_list in idx_obs_list:
|
||||
if not obs_content_list:
|
||||
continue
|
||||
|
||||
# [1, 2022年6月, 用户在2022年6月去杭州旅游, 旅游]
|
||||
if len(obs_content_list) != 4:
|
||||
self.logger.warning(f"obs_content_list={obs_content_list} is invalid!")
|
||||
continue
|
||||
|
||||
idx, time_infer, obs_content, keywords = obs_content_list
|
||||
|
||||
if obs_content in ["无", "重复"]:
|
||||
continue
|
||||
|
||||
if not idx.isdigit():
|
||||
self.logger.warning(f"idx={idx} is invalid!")
|
||||
continue
|
||||
|
||||
# 序号需要修正-1
|
||||
idx = int(idx) - 1
|
||||
if idx >= len(self.messages):
|
||||
self.logger.warning(
|
||||
f"idx={idx} is invalid! messages.size={len(self.messages)}"
|
||||
)
|
||||
continue
|
||||
|
||||
new_obs_nodes.append(
|
||||
self.add_observation(
|
||||
message=self.messages[idx],
|
||||
obs_content=obs_content,
|
||||
keywords=keywords,
|
||||
)
|
||||
)
|
||||
|
||||
# save context
|
||||
self.set_context(NEW_OBS_NODES, new_obs_nodes)
|
||||
|
|
@ -1,70 +0,0 @@
|
|||
from ...utils.response_text_parser import ResponseTextParser
|
||||
from enumeration.message_role_enum import MessageRoleEnum
|
||||
from worker.memory_base_worker import MemoryBaseWorker
|
||||
from ...chat.global_context import GlobalContext
|
||||
from ...prompts.info_filter_prompt import INFO_FILTER_FEW_SHOT_PROMPT, INFO_FILTER_SYSTEM_PROMPT, INFO_FILTER_USER_QUERY_PROMPT
|
||||
|
||||
|
||||
class InfoFilterWorker(MemoryBaseWorker):
|
||||
def _run(self):
|
||||
# filter user msg
|
||||
info_messages = []
|
||||
for msg in self.messages:
|
||||
if msg.role != MessageRoleEnum.USER.value:
|
||||
continue
|
||||
if len(msg.content) >= self.info_filter_msg_max_size:
|
||||
continue
|
||||
info_messages.append(msg)
|
||||
|
||||
# gene prompt
|
||||
user_query = "\n".join(
|
||||
[f"{i + 1} 用户:{msg.content}" for i, msg in enumerate(info_messages)]
|
||||
)
|
||||
info_filter_message = self.prompt_to_msg(
|
||||
system_prompt=self.get_prompt(INFO_FILTER_SYSTEM_PROMPT).format(
|
||||
batch_size=len(info_messages)
|
||||
),
|
||||
few_shot=self.get_prompt(INFO_FILTER_FEW_SHOT_PROMPT),
|
||||
user_query=self.get_prompt(INFO_FILTER_USER_QUERY_PROMPT).format(
|
||||
user_query=user_query
|
||||
),
|
||||
)
|
||||
self.logger.info(f"info_filter_message={info_filter_message}")
|
||||
|
||||
# call llm
|
||||
response_text = self.generation_model.call(
|
||||
messages=info_filter_message,
|
||||
model_name=self.info_filter_model,
|
||||
max_token=self.info_filter_max_token,
|
||||
temperature=self.info_filter_temperature,
|
||||
top_k=self.info_filter_top_k,
|
||||
)
|
||||
|
||||
# return if empty
|
||||
if not response_text:
|
||||
self.add_run_info("info score call llm failed!", continue_run=False)
|
||||
return
|
||||
|
||||
# parse text
|
||||
info_score_list = ResponseTextParser(response_text).parse_v1("info_filter")
|
||||
if len(info_score_list) != len(info_messages):
|
||||
self.add_run_info(
|
||||
f"info_score_size != info_messages_size, "
|
||||
f"{len(info_score_list)} vs {len(info_messages)}",
|
||||
continue_run=False,
|
||||
)
|
||||
return
|
||||
|
||||
# 过滤value=0的messages
|
||||
filtered_messages = []
|
||||
for msg, info_score in zip(info_messages, info_score_list):
|
||||
if not info_score:
|
||||
continue
|
||||
score = info_score[0]
|
||||
# if score in ("1", "2",):
|
||||
if score in ("2",):
|
||||
msg.info_score = score
|
||||
filtered_messages.append(msg)
|
||||
|
||||
# 后续不会关注为0的msg,直接丢弃
|
||||
self.messages = filtered_messages
|
||||
Loading…
Add table
Reference in a new issue