mirror of
https://github.com/agentscope-ai/ReMe.git
synced 2026-09-06 08:16:00 +00:00
feat: Resolve conflict, auto committed by CodeFlow
This commit is contained in:
commit
ae0d9956f3
43 changed files with 212 additions and 194 deletions
|
|
@ -2,9 +2,9 @@
|
|||
"thread_pool_max_count": "",
|
||||
"worker": "config/worker.json",
|
||||
"pipeline": {
|
||||
"summary_short": "",
|
||||
"summary_long": "",
|
||||
"retrieve": ""
|
||||
"summary_short": "parse_params,[load_profile|es_new_obs|es_insight],[update_insight|get_reflection,get_insight|update_profile],summary_collect,memory_store",
|
||||
"summary_long": "parse_params,info_filter,[es_today_obs|get_observation|get_observation_with_time],contra_repeat,memory_store",
|
||||
"retrieve": "parse_params,load_profile,[extract_time|es_similar|es_keyword],semantic_rank,fuse_rerank"
|
||||
},
|
||||
"model_embedding": "config/model/dash_embedding.json",
|
||||
"model_rerank": "config/model/dash_rerank.json",
|
||||
|
|
|
|||
|
|
@ -1,52 +1,45 @@
|
|||
{
|
||||
"EsInsightWorker": {
|
||||
"module_name": "EsInsightWorker",
|
||||
"module_path": "memory_scope/worker",
|
||||
"kwargs": {
|
||||
"es_insight_top_k": 128
|
||||
}
|
||||
"es_insight": {
|
||||
"name": "EsInsightWorker",
|
||||
"path": "memory_scope/worker",
|
||||
"es_insight_top_k": 128
|
||||
},
|
||||
"es_keyword": {
|
||||
"name": "EsKeywordWorker",
|
||||
"path": "memory_scope/worker",
|
||||
"es_keyword_top_k": 10
|
||||
},
|
||||
"es_keyword2": {
|
||||
"module_name": "EsKeywordWorker",
|
||||
"module_path": "memory_scope/worker",
|
||||
"es_keyword_top_k": 10
|
||||
},
|
||||
"EsNewObsWorker": {
|
||||
"module_name": "EsNewObsWorker",
|
||||
"module_path": "memory_scope/worker",
|
||||
"name": "EsNewObsWorker",
|
||||
"path": "memory_scope/worker",
|
||||
"kwargs": {
|
||||
"es_new_obs_top_k": 256
|
||||
}
|
||||
},
|
||||
"EsNotReflectedWorker": {
|
||||
"module_name": "EsNotReflectedWorker",
|
||||
"module_path": "memory_scope/worker",
|
||||
"name": "EsNotReflectedWorker",
|
||||
"path": "memory_scope/worker",
|
||||
"kwargs": {
|
||||
"es_not_reflected_top_k": 256
|
||||
}
|
||||
},
|
||||
"EsSimilarWorker": {
|
||||
"module_name": "EsSimilarWorker",
|
||||
"module_path": "memory_scope/worker",
|
||||
"name": "EsSimilarWorker",
|
||||
"path": "memory_scope/worker",
|
||||
"kwargs": {
|
||||
"es_similar_top_k": 128
|
||||
}
|
||||
},
|
||||
"EsTodayObsWorker": {
|
||||
"module_name": "EsTodayObsWorker",
|
||||
"module_path": "memory_scope/worker",
|
||||
"name": "EsTodayObsWorker",
|
||||
"path": "memory_scope/worker",
|
||||
"kwargs": {
|
||||
"es_today_obs_top_k": 128
|
||||
}
|
||||
},
|
||||
"GetInsightWorker": {
|
||||
"module_name": "GetInsightWorker",
|
||||
"module_path": "memory_scope/worker",
|
||||
"name": "GetInsightWorker",
|
||||
"path": "memory_scope/worker",
|
||||
"kwargs": {
|
||||
"es_insight_similar_top_k": 128,
|
||||
"insight_obs_max_cnt": 10,
|
||||
|
|
@ -57,18 +50,16 @@
|
|||
}
|
||||
},
|
||||
"ExtractTimeWorker": {
|
||||
"module_name": "ExtractTimeWorker",
|
||||
"module_path": "memory_scope/worker",
|
||||
"kwargs": {
|
||||
"parse_time_model": "qwen_1_8_parse_time_service",
|
||||
"parse_time_max_token": 100,
|
||||
"parse_time_temperature": 0.6,
|
||||
"parse_time_top_k": 1
|
||||
}
|
||||
"name": "ExtractTimeWorker",
|
||||
"path": "memory_scope/worker",
|
||||
"parse_time_model": "qwen_1_8_parse_time_service",
|
||||
"parse_time_max_token": 100,
|
||||
"parse_time_temperature": 0.6,
|
||||
"parse_time_top_k": 1
|
||||
},
|
||||
"InfoFilterWorker": {
|
||||
"module_name": "InfoFilterWorker",
|
||||
"module_path": "memory_scope/worker",
|
||||
"name": "InfoFilterWorker",
|
||||
"path": "memory_scope/worker",
|
||||
"kwargs": {
|
||||
"info_filter_msg_max_size": 200,
|
||||
"info_filter_model": "qwen_max",
|
||||
|
|
@ -78,8 +69,8 @@
|
|||
}
|
||||
},
|
||||
"GetObservationWithTimeWorker": {
|
||||
"module_name": "GetObservationWithTimeWorker",
|
||||
"module_path": "memory_scope/worker",
|
||||
"name": "GetObservationWithTimeWorker",
|
||||
"path": "memory_scope/worker",
|
||||
"kwargs": {
|
||||
"summary_messages_model": "qwen_max",
|
||||
"summary_messages_max_token": 500,
|
||||
|
|
@ -88,8 +79,8 @@
|
|||
}
|
||||
},
|
||||
"GetObservationWorker": {
|
||||
"module_name": "GetObservationWorker",
|
||||
"module_path": "memory_scope/worker",
|
||||
"name": "GetObservationWorker",
|
||||
"path": "memory_scope/worker",
|
||||
"kwargs": {
|
||||
"summary_messages_model": "qwen_max",
|
||||
"summary_messages_max_token": 500,
|
||||
|
|
@ -98,8 +89,8 @@
|
|||
}
|
||||
},
|
||||
"ContraRepeatWorker": {
|
||||
"module_name": "ContraRepeatWorker",
|
||||
"module_path": "memory_scope/worker",
|
||||
"name": "ContraRepeatWorker",
|
||||
"path": "memory_scope/worker",
|
||||
"kwargs": {
|
||||
"merge_obs_model": "qwen_max",
|
||||
"merge_obs_max_token": 500,
|
||||
|
|
@ -108,8 +99,8 @@
|
|||
}
|
||||
},
|
||||
"FuseRerankWorker": {
|
||||
"module_name": "FuseRerankWorker",
|
||||
"module_path": "memory_scope/worker",
|
||||
"name": "FuseRerankWorker",
|
||||
"path": "memory_scope/worker",
|
||||
"kwargs": {
|
||||
"fuse_score_threshold": 0.1,
|
||||
"fuse_ratio_dict": {
|
||||
|
|
@ -124,8 +115,8 @@
|
|||
}
|
||||
},
|
||||
"UpdateProfileWorker": {
|
||||
"module_name": "UpdateProfileWorker",
|
||||
"module_path": "memory_scope/worker",
|
||||
"name": "UpdateProfileWorker",
|
||||
"path": "memory_scope/worker",
|
||||
"kwargs": {
|
||||
"update_profile_threshold": 0.1,
|
||||
"update_profile_model": "qwen_max",
|
||||
|
|
@ -136,8 +127,8 @@
|
|||
}
|
||||
},
|
||||
"GetReflectionWorker": {
|
||||
"module_name": "GetReflectionWorker",
|
||||
"module_path": "memory_scope/worker",
|
||||
"name": "GetReflectionWorker",
|
||||
"path": "memory_scope/worker",
|
||||
"kwargs": {
|
||||
"reflect_obs_cnt_threshold": 40,
|
||||
"reflect_num_questions": 3,
|
||||
|
|
@ -148,8 +139,8 @@
|
|||
}
|
||||
},
|
||||
"UpdateInsightWorker": {
|
||||
"module_name": "UpdateInsightWorker",
|
||||
"module_path": "memory_scope/worker",
|
||||
"name": "UpdateInsightWorker",
|
||||
"path": "memory_scope/worker",
|
||||
"kwargs": {
|
||||
"update_insight_threshold": 0.1,
|
||||
"update_insight_model": "qwen_max",
|
||||
|
|
|
|||
14
memory_scope/cli.py
Normal file
14
memory_scope/cli.py
Normal file
|
|
@ -0,0 +1,14 @@
|
|||
# 使用argparse库的示例
|
||||
import argparse
|
||||
import fire
|
||||
from config import C, init
|
||||
from chat.memory_chat import MemoryChat
|
||||
|
||||
def main(config_path:str):
|
||||
init(config_path)
|
||||
|
||||
agent = MemoryChat()
|
||||
agent.run()
|
||||
|
||||
if __name__ == "__main__":
|
||||
fire.Fire(main)
|
||||
|
|
@ -1,15 +0,0 @@
|
|||
# 使用argparse库的示例
|
||||
import argparse
|
||||
|
||||
|
||||
def main():
|
||||
parser = argparse.ArgumentParser(description="示例CLI程序")
|
||||
parser.add_argument('--echo', help="输出传入的消息")
|
||||
|
||||
args = parser.parse_args()
|
||||
if args.echo:
|
||||
print(f"收到的消息: {args.echo}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
|
|
@ -1 +0,0 @@
|
|||
# runtime params apikey or overwrite config params
|
||||
34
memory_scope/config.py
Normal file
34
memory_scope/config.py
Normal file
|
|
@ -0,0 +1,34 @@
|
|||
|
||||
import json
|
||||
|
||||
from pipeline.memory_service import MemoryService
|
||||
from utils.tool_functions import init_instance_by_config
|
||||
|
||||
|
||||
class Wrapper:
|
||||
"""Wrapper class for anything that needs to set up during init"""
|
||||
|
||||
def __init__(self):
|
||||
self._provider = None
|
||||
|
||||
def register(self, provider):
|
||||
self._provider = provider
|
||||
|
||||
def __getattr__(self, key: str):
|
||||
if self.__dict__.get("_provider", None) is None:
|
||||
raise AttributeError("Please run __init__ first!")
|
||||
return getattr(self._provider, key)
|
||||
|
||||
C = Wrapper()
|
||||
|
||||
def init(config_path: str):
|
||||
config = json.loads(config_path)
|
||||
C.register(config)
|
||||
|
||||
## register workers
|
||||
C.worker = json.loads(C.worker)
|
||||
|
||||
## register services
|
||||
for k,v in C.pipeline.items():
|
||||
C.pipeline[k] = MemoryService(k)
|
||||
|
||||
|
|
@ -1,7 +1,11 @@
|
|||
from elasticsearch import Elasticsearch
|
||||
from elasticsearch.helpers import bulk
|
||||
|
||||
|
||||
from memory_scope.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 memory_scope.utils.logger import Logger
|
||||
|
|
|
|||
|
|
@ -4,8 +4,8 @@ from http import HTTPStatus
|
|||
|
||||
import requests
|
||||
|
||||
from common.logger import Logger
|
||||
from common.timer import Timer
|
||||
from utils.logger import Logger
|
||||
from utils.timer import Timer
|
||||
from enumeration.env_type import EnvType
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
from pydantic import Field, BaseModel
|
||||
|
||||
from model.memory_node import MemoryNode
|
||||
from node.memory_node import MemoryNode
|
||||
|
||||
|
||||
class MemoryWrapNode(BaseModel):
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
import re
|
||||
|
||||
from common.logger import Logger
|
||||
from utils.logger import Logger
|
||||
|
||||
|
||||
class ResponseTextParser(object):
|
||||
|
|
|
|||
|
|
@ -1,15 +1,18 @@
|
|||
from typing import List, Dict
|
||||
|
||||
from pydantic import Field
|
||||
from pydantic import Field, BaseModel
|
||||
|
||||
from model.message import Message
|
||||
from model.user_attribute import UserAttribute
|
||||
from request.base_model import RequestBaseModel
|
||||
from node.message import Message
|
||||
from node.user_attribute import UserAttribute
|
||||
|
||||
class MemoryServiceRequestModel(RequestBaseModel):
|
||||
class MemoryServiceRequestModel(BaseModel):
|
||||
user: UserConfig = None
|
||||
|
||||
messages: List[Message] = Field(...,
|
||||
description="summary: 多轮对话的list,默认按照时间正序,最后一条是最新的; retrieve: 最后一条是query")
|
||||
|
||||
messages_pick_n: int = Field(1, description="summary:传需要总结的msg的个数;retrieve:不传")
|
||||
user_profile: List[UserAttribute] = Field([], description="user_profile")
|
||||
|
||||
ext_info: Dict[str, str] = Field({}, description="extra information")
|
||||
|
||||
extra_user_attrs: List = []
|
||||
|
|
@ -6,17 +6,16 @@ from importlib import import_module
|
|||
from itertools import zip_longest
|
||||
from typing import Dict, Any
|
||||
|
||||
from worker.memory.base_worker import BaseWorker
|
||||
|
||||
from common.context_handler import ContextHandler
|
||||
from common.logger import Logger
|
||||
from common.timer import timer, Timer
|
||||
from worker.base_worker import BaseWorker
|
||||
from utils.context_handler import ContextHandler
|
||||
from utils.logger import Logger
|
||||
from utils.timer import timer, Timer
|
||||
from common.tool_functions import under_line_to_hump
|
||||
from constants import common_constants
|
||||
from constants.common_constants import RESPONSE_EXT_INFO, MAX_WORKERS, PIPELINE
|
||||
from enumeration.memory_method_enum import MemoryMethodEnum
|
||||
from request.memory import MemoryServiceRequestModel
|
||||
from config.env_config import C, Workers
|
||||
from pipeline.memory import MemoryServiceRequestModel
|
||||
from cli.cli_config import C
|
||||
from utils.tool_functions import init_instance_by_config
|
||||
|
||||
|
||||
|
|
@ -26,7 +25,7 @@ class MemoryService(object):
|
|||
self.context_handler = ContextHandler()
|
||||
|
||||
# 线程池
|
||||
self.thread_pool = ThreadPoolExecutor(max_workers=C.THREAD_POOL_MAX_COUNT)
|
||||
self.thread_pool = ThreadPoolExecutor(max_workers=C.thread_pool_max_count)
|
||||
|
||||
# 全部初始化的worker
|
||||
self.worker_dict: Dict[str, BaseWorker] = {}
|
||||
|
|
@ -40,7 +39,7 @@ class MemoryService(object):
|
|||
|
||||
def get_worker(self, worker_name: str, is_multi_thread: bool = False) -> BaseWorker:
|
||||
return init_instance_by_config(
|
||||
config = W.get(worker_name),
|
||||
config = C.worker.get(worker_name),
|
||||
try_kwargs={
|
||||
"is_multi_thread": is_multi_thread,
|
||||
"thread_pool": self.thread_pool
|
||||
|
|
@ -2,7 +2,7 @@ import os
|
|||
import threading
|
||||
from typing import Dict, Any
|
||||
|
||||
from common.logger import Logger
|
||||
from utils.logger import Logger
|
||||
|
||||
class ContextHandler(object):
|
||||
def __init__(self):
|
||||
|
|
|
|||
|
|
@ -6,7 +6,7 @@ date: 20221106
|
|||
|
||||
import time
|
||||
|
||||
from common.logger import Logger
|
||||
from utils.logger import Logger
|
||||
|
||||
|
||||
class Timer(object):
|
||||
|
|
|
|||
|
|
@ -4,29 +4,8 @@ from datetime import datetime
|
|||
from typing import Dict, List
|
||||
|
||||
from constants.common_constants import WEEKDAYS
|
||||
from enumeration.env_type import EnvType
|
||||
from importlib import import_module
|
||||
|
||||
global_env_type = None
|
||||
|
||||
|
||||
def get_global_env_type():
|
||||
global global_env_type
|
||||
|
||||
if global_env_type is None:
|
||||
env = os.environ.get("APP_ENV", "")
|
||||
if env is None or not env:
|
||||
raise EnvironmentError("Environment variable APP_ENV must be set")
|
||||
env = env.split("-")[-1]
|
||||
|
||||
if env not in EnvType.__members__.values():
|
||||
global_env_type = EnvType.DAILY
|
||||
else:
|
||||
global_env_type = EnvType(env)
|
||||
|
||||
return global_env_type
|
||||
|
||||
|
||||
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:]
|
||||
|
|
@ -165,10 +144,13 @@ def time_to_formatted_str(time: datetime | str | int | float = None,
|
|||
return return_str
|
||||
|
||||
|
||||
def init_instance_by_config(config, default_module_path=None, try_kwargs={}):
|
||||
import_module(config.get("module_path", default_module_path))
|
||||
clazz = getattr(module, config.get("module_name"))
|
||||
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.get("kwargs"),**try_kwargs)
|
||||
return clazz(**config, **try_kwargs)
|
||||
except:
|
||||
return clazz(**config.get("kwargs"))
|
||||
return clazz(**config)
|
||||
|
|
@ -2,8 +2,8 @@ import json
|
|||
from typing import List, Dict
|
||||
|
||||
from enumeration.memory_node_status import MemoryNodeStatus
|
||||
from model.memory_wrap_node import MemoryWrapNode
|
||||
from model.user_attribute import UserAttribute
|
||||
from node.memory_wrap_node import MemoryWrapNode
|
||||
from node.user_attribute import UserAttribute
|
||||
|
||||
|
||||
class UserProfileHandler(object):
|
||||
|
|
|
|||
|
|
@ -2,9 +2,9 @@ import time
|
|||
from concurrent.futures import ThreadPoolExecutor, as_completed
|
||||
from typing import Any, List
|
||||
|
||||
from common.context_handler import ContextHandler
|
||||
from common.logger import Logger
|
||||
from common.timer import Timer
|
||||
from utils.context_handler import ContextHandler
|
||||
from utils.logger import Logger
|
||||
from utils.timer import Timer
|
||||
|
||||
|
||||
class BaseWorker(object):
|
||||
|
|
|
|||
|
|
@ -3,8 +3,8 @@ from typing import List
|
|||
from constants.common_constants import INSIGHT_NODES
|
||||
from enumeration.memory_node_status import MemoryNodeStatus
|
||||
from enumeration.memory_type_enum import MemoryTypeEnum
|
||||
from model.memory.memory_wrap_node import MemoryWrapNode
|
||||
from worker.memory.memory_base_worker import MemoryBaseWorker
|
||||
from node.memory_wrap_node import MemoryWrapNode
|
||||
from worker.memory_base_worker import MemoryBaseWorker
|
||||
|
||||
|
||||
class EsInsightWorker(MemoryBaseWorker):
|
||||
|
|
|
|||
|
|
@ -4,8 +4,8 @@ from constants.common_constants import KEY_WORD, KEYWORD_OBS_NODES, RECALL_TYPE,
|
|||
from enumeration.memory_node_status import MemoryNodeStatus
|
||||
from enumeration.memory_recall_type import MemoryRecallType
|
||||
from enumeration.memory_type_enum import MemoryTypeEnum
|
||||
from model.memory_wrap_node import MemoryWrapNode
|
||||
from worker.bailian.memory_base_worker import MemoryBaseWorker
|
||||
from node.memory_wrap_node import MemoryWrapNode
|
||||
from worker.memory_base_worker import MemoryBaseWorker
|
||||
|
||||
|
||||
class EsKeywordWorker(MemoryBaseWorker):
|
||||
|
|
|
|||
|
|
@ -3,8 +3,8 @@ from typing import List
|
|||
from constants.common_constants import NEW, NEW_OBS_NODES
|
||||
from enumeration.memory_node_status import MemoryNodeStatus
|
||||
from enumeration.memory_type_enum import MemoryTypeEnum
|
||||
from model.memory.memory_wrap_node import MemoryWrapNode
|
||||
from worker.memory.memory_base_worker import MemoryBaseWorker
|
||||
from node.memory_wrap_node import MemoryWrapNode
|
||||
from worker.memory_base_worker import MemoryBaseWorker
|
||||
|
||||
|
||||
class EsNewObsWorker(MemoryBaseWorker):
|
||||
|
|
|
|||
|
|
@ -3,8 +3,8 @@ from typing import List
|
|||
from constants.common_constants import REFLECTED, NOT_REFLECTED_OBS_NODES
|
||||
from enumeration.memory_node_status import MemoryNodeStatus
|
||||
from enumeration.memory_type_enum import MemoryTypeEnum
|
||||
from model.memory.memory_wrap_node import MemoryWrapNode
|
||||
from worker.memory.memory_base_worker import MemoryBaseWorker
|
||||
from node.memory_wrap_node import MemoryWrapNode
|
||||
from worker.memory_base_worker import MemoryBaseWorker
|
||||
|
||||
|
||||
class EsNotReflectedWorker(MemoryBaseWorker):
|
||||
|
|
|
|||
|
|
@ -2,8 +2,8 @@ from typing import List
|
|||
|
||||
from constants.common_constants import ALL_NODES, ALL_MEMORIES
|
||||
from enumeration.memory_node_status import MemoryNodeStatus
|
||||
from model.memory.memory_wrap_node import MemoryWrapNode
|
||||
from worker.memory.memory_base_worker import MemoryBaseWorker
|
||||
from node.memory_wrap_node import MemoryWrapNode
|
||||
from worker.memory_base_worker import MemoryBaseWorker
|
||||
|
||||
|
||||
class EsRetrieveAllWorker(MemoryBaseWorker):
|
||||
|
|
|
|||
|
|
@ -4,8 +4,8 @@ from constants.common_constants import SIMILAR_OBS_NODES, RECALL_TYPE
|
|||
from enumeration.memory_node_status import MemoryNodeStatus
|
||||
from enumeration.memory_recall_type import MemoryRecallType
|
||||
from enumeration.memory_type_enum import MemoryTypeEnum
|
||||
from model.memory.memory_wrap_node import MemoryWrapNode
|
||||
from worker.memory.memory_base_worker import MemoryBaseWorker
|
||||
from node.memory_wrap_node import MemoryWrapNode
|
||||
from worker.memory_base_worker import MemoryBaseWorker
|
||||
|
||||
|
||||
class EsSimilarWorker(MemoryBaseWorker):
|
||||
|
|
|
|||
|
|
@ -4,9 +4,8 @@ from common.tool_functions import time_to_formatted_str
|
|||
from constants.common_constants import TODAY_OBS_NODES, DT
|
||||
from enumeration.memory_node_status import MemoryNodeStatus
|
||||
from enumeration.memory_type_enum import MemoryTypeEnum
|
||||
from model.memory.memory_wrap_node import MemoryWrapNode
|
||||
from worker.memory.memory_base_worker import MemoryBaseWorker
|
||||
|
||||
from node.memory_wrap_node import MemoryWrapNode
|
||||
from worker.memory_base_worker import MemoryBaseWorker
|
||||
|
||||
class EsTodayObsWorker(MemoryBaseWorker):
|
||||
def __init__(self, es_today_obs_top_k, *args, **kwargs):
|
||||
|
|
|
|||
|
|
@ -1,13 +1,13 @@
|
|||
from typing import List, Dict
|
||||
|
||||
from common.user_profile_handler import UserProfileHandler
|
||||
from utils.user_profile_handler import UserProfileHandler
|
||||
from constants import common_constants
|
||||
from enumeration.memory_node_status import MemoryNodeStatus
|
||||
from enumeration.memory_type_enum import MemoryTypeEnum
|
||||
from model.memory.memory_wrap_node import MemoryWrapNode
|
||||
from model.user_attribute import UserAttribute
|
||||
from request.memory import MemoryServiceRequestModel
|
||||
from worker.memory.memory_base_worker import MemoryBaseWorker
|
||||
from node.memory_wrap_node import MemoryWrapNode
|
||||
from node.user_attribute import UserAttribute
|
||||
from pipeline.memory import MemoryServiceRequestModel
|
||||
from worker.memory_base_worker import MemoryBaseWorker
|
||||
|
||||
|
||||
class LoadProfileWorker(MemoryBaseWorker):
|
||||
|
|
|
|||
|
|
@ -3,11 +3,12 @@ from typing import List, Dict, Optional
|
|||
from constants import common_constants
|
||||
from constants.common_constants import CONFIG, MESSAGES, PROMPT_CONFIG
|
||||
from enumeration.message_role_enum import MessageRoleEnum
|
||||
from model.message import Message
|
||||
from model.user_attribute import UserAttribute
|
||||
from request.memory import MemoryServiceRequestModel
|
||||
from worker.memory.base_worker import BaseWorker
|
||||
from config.env_config import EnvConfig
|
||||
from node.message import Message
|
||||
from node.user_attribute import UserAttribute
|
||||
from pipeline.memory import MemoryServiceRequestModel
|
||||
from worker.base_worker import BaseWorker
|
||||
from cli.cli_config import C
|
||||
from utils.tool_functions import init_instance_by_config
|
||||
|
||||
class MemoryBaseWorker(BaseWorker):
|
||||
def __init__(self, **kwargs):
|
||||
|
|
@ -21,7 +22,7 @@ class MemoryBaseWorker(BaseWorker):
|
|||
def messages(self) -> List[Message]:
|
||||
messages: List[Message] = self.context_handler.get_context(MESSAGES)
|
||||
if messages is None:
|
||||
messages = self.request.messages[-self.request.messages_pick_n:]
|
||||
messages = self.request.messages
|
||||
self.context_handler.set_context(MESSAGES, messages)
|
||||
return messages
|
||||
|
||||
|
|
@ -51,20 +52,32 @@ class MemoryBaseWorker(BaseWorker):
|
|||
return self.request.user.prompt
|
||||
|
||||
@property
|
||||
def emb_client(self):
|
||||
return C.emb_client
|
||||
def client(self, model_type: str, model_name:str):
|
||||
models = C.get(model_type)
|
||||
models["model_name"] = init_instance_by_config(
|
||||
config = models.get(model_name),
|
||||
try_kwargs={
|
||||
"is_multi_thread": is_multi_thread,
|
||||
"thread_pool": self.thread_pool
|
||||
}
|
||||
)
|
||||
return models["model_name"]
|
||||
|
||||
@property
|
||||
def emb_client(self, model_name: str):
|
||||
self.client("model_embedding", model_name)
|
||||
|
||||
@property
|
||||
def gene_client(self):
|
||||
return C.gene_client
|
||||
self.client("model_generate", model_name)
|
||||
|
||||
@property
|
||||
def rerank_client(self):
|
||||
return C.rerank_client
|
||||
self.client("model_rerank", model_name)
|
||||
|
||||
@property
|
||||
def es_client(self):
|
||||
return C.es_client
|
||||
self.client("db", model_name)
|
||||
|
||||
@property
|
||||
def tenant_id(self):
|
||||
|
|
|
|||
|
|
@ -1,12 +1,11 @@
|
|||
from typing import List
|
||||
|
||||
from common.user_profile_handler import UserProfileHandler
|
||||
from utils.user_profile_handler import UserProfileHandler
|
||||
from constants.common_constants import MODIFIED_MEMORIES, NEW_USER_PROFILE
|
||||
from model.memory_node import MemoryNode
|
||||
from model.memory.memory_wrap_node import MemoryWrapNode
|
||||
from model.user_attribute import UserAttribute
|
||||
from worker.memory.memory_base_worker import MemoryBaseWorker
|
||||
|
||||
from node.memory_node import MemoryNode
|
||||
from node.memory_wrap_node import MemoryWrapNode
|
||||
from node.user_attribute import UserAttribute
|
||||
from worker.memory_base_worker import MemoryBaseWorker
|
||||
|
||||
class MemoryStoreWorker(MemoryBaseWorker):
|
||||
|
||||
|
|
|
|||
|
|
@ -2,8 +2,8 @@ import json
|
|||
|
||||
from config.bailian_memory_config import BailianMemoryConfig
|
||||
from constants.common_constants import REQUEST, CONFIG
|
||||
from request.memory import MemoryServiceRequestModel
|
||||
from worker.memory.base_worker import BaseWorker
|
||||
from pipeline.memory import MemoryServiceRequestModel
|
||||
from worker.base_worker import BaseWorker
|
||||
|
||||
|
||||
class ParseParamsWorker(BaseWorker):
|
||||
|
|
|
|||
|
|
@ -3,7 +3,7 @@ import re
|
|||
from common.tool_functions import time_to_formatted_str
|
||||
from constants.common_constants import DATATIME_WORD_LIST, DATATIME_KEY_MAP
|
||||
from constants.common_constants import EXTRACT_TIME_DICT
|
||||
from worker.memory.memory_base_worker import MemoryBaseWorker
|
||||
from worker.memory_base_worker import MemoryBaseWorker
|
||||
|
||||
|
||||
class ExtractTimeWorker(MemoryBaseWorker):
|
||||
|
|
|
|||
|
|
@ -2,8 +2,8 @@ from typing import Dict, List
|
|||
|
||||
from constants.common_constants import RELATED_MEMORIES, EXTRACT_TIME_DICT, ALL_ONLINE_NODES, \
|
||||
TIME_MATCHED
|
||||
from model.memory.memory_wrap_node import MemoryWrapNode
|
||||
from worker.memory.memory_base_worker import MemoryBaseWorker
|
||||
from node.memory_wrap_node import MemoryWrapNode
|
||||
from worker.memory_base_worker import MemoryBaseWorker
|
||||
|
||||
|
||||
class FuseRerankWorker(MemoryBaseWorker):
|
||||
|
|
|
|||
|
|
@ -1,11 +1,11 @@
|
|||
from typing import List, Dict
|
||||
|
||||
from common.user_profile_handler import UserProfileHandler
|
||||
from utils.user_profile_handler import UserProfileHandler
|
||||
from constants.common_constants import SIMILAR_OBS_NODES, RECALL_TYPE, KEYWORD_OBS_NODES, ALL_ONLINE_NODES, \
|
||||
QUERY_KEYWORDS
|
||||
from enumeration.memory_recall_type import MemoryRecallType
|
||||
from model.memory.memory_wrap_node import MemoryWrapNode
|
||||
from worker.memory.memory_base_worker import MemoryBaseWorker
|
||||
from node.memory_wrap_node import MemoryWrapNode
|
||||
from worker.memory_base_worker import MemoryBaseWorker
|
||||
|
||||
|
||||
class SemanticRankWorker(MemoryBaseWorker):
|
||||
|
|
|
|||
|
|
@ -6,8 +6,8 @@ from constants.common_constants import NEW_INSIGHT_NODES, DT, NOT_REFLECTED_MERG
|
|||
INSIGHT_VALUE, REFLECTED
|
||||
from enumeration.memory_node_status import MemoryNodeStatus
|
||||
from enumeration.memory_type_enum import MemoryTypeEnum
|
||||
from model.memory.memory_wrap_node import MemoryWrapNode
|
||||
from worker.memory.memory_base_worker import MemoryBaseWorker
|
||||
from node.memory_wrap_node import MemoryWrapNode
|
||||
from worker.memory_base_worker import MemoryBaseWorker
|
||||
|
||||
|
||||
class GetInsightWorker(MemoryBaseWorker):
|
||||
|
|
|
|||
|
|
@ -3,8 +3,8 @@ 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 model.memory.memory_wrap_node import MemoryWrapNode
|
||||
from worker.memory.memory_base_worker import MemoryBaseWorker
|
||||
from node.memory_wrap_node import MemoryWrapNode
|
||||
from worker.memory_base_worker import MemoryBaseWorker
|
||||
|
||||
|
||||
class GetReflectionWorker(MemoryBaseWorker):
|
||||
|
|
|
|||
|
|
@ -5,8 +5,8 @@ from constants.common_constants import NEW_OBS_NODES, TODAY_OBS_NODES, MSG_TIME,
|
|||
MODIFIED_MEMORIES
|
||||
from enumeration.memory_node_status import MemoryNodeStatus
|
||||
from enumeration.memory_type_enum import MemoryTypeEnum
|
||||
from model.memory_wrap_node import MemoryWrapNode
|
||||
from worker.bailian.memory_base_worker import MemoryBaseWorker
|
||||
from node.memory_wrap_node import MemoryWrapNode
|
||||
from worker.memory_base_worker import MemoryBaseWorker
|
||||
|
||||
|
||||
class LongContraRepeatWorker(MemoryBaseWorker):
|
||||
|
|
|
|||
|
|
@ -2,8 +2,8 @@ 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 model.memory.memory_wrap_node import MemoryWrapNode
|
||||
from worker.memory.memory_base_worker import MemoryBaseWorker
|
||||
from node.memory_wrap_node import MemoryWrapNode
|
||||
from worker.memory_base_worker import MemoryBaseWorker
|
||||
|
||||
|
||||
class SummaryCollectWorker(MemoryBaseWorker):
|
||||
|
|
|
|||
|
|
@ -2,8 +2,8 @@ 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 model.memory.memory_wrap_node import MemoryWrapNode
|
||||
from worker.memory.memory_base_worker import MemoryBaseWorker
|
||||
from node.memory_wrap_node import MemoryWrapNode
|
||||
from worker.memory_base_worker import MemoryBaseWorker
|
||||
|
||||
|
||||
class UpdateInsightWorker(MemoryBaseWorker):
|
||||
|
|
|
|||
|
|
@ -3,9 +3,9 @@ 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 model.memory.memory_wrap_node import MemoryWrapNode
|
||||
from model.user_attribute import UserAttribute
|
||||
from worker.memory.memory_base_worker import MemoryBaseWorker
|
||||
from node.memory_wrap_node import MemoryWrapNode
|
||||
from node.user_attribute import UserAttribute
|
||||
from worker.memory_base_worker import MemoryBaseWorker
|
||||
|
||||
|
||||
class UpdateProfileWorker(MemoryBaseWorker):
|
||||
|
|
|
|||
|
|
@ -4,8 +4,8 @@ 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_node_status import MemoryNodeStatus
|
||||
from model.memory.memory_wrap_node import MemoryWrapNode
|
||||
from worker.memory.memory_base_worker import MemoryBaseWorker
|
||||
from node.memory_wrap_node import MemoryWrapNode
|
||||
from worker.memory_base_worker import MemoryBaseWorker
|
||||
|
||||
|
||||
class ContraRepeatWorker(MemoryBaseWorker):
|
||||
|
|
|
|||
|
|
@ -7,9 +7,9 @@ from constants.common_constants import REFLECTED, DT, TIME_INFER, NEW, MSG_TIME,
|
|||
NEW_OBS_WITH_TIME_NODES
|
||||
from enumeration.memory_node_status import MemoryNodeStatus
|
||||
from enumeration.memory_type_enum import MemoryTypeEnum
|
||||
from model.memory.memory_wrap_node import MemoryWrapNode
|
||||
from model.message import Message
|
||||
from worker.memory.memory_base_worker import MemoryBaseWorker
|
||||
from node.memory_wrap_node import MemoryWrapNode
|
||||
from node.message import Message
|
||||
from worker.memory_base_worker import MemoryBaseWorker
|
||||
|
||||
|
||||
class GetObservationWithTimeWorker(MemoryBaseWorker):
|
||||
|
|
|
|||
|
|
@ -7,9 +7,9 @@ from constants.common_constants import REFLECTED, DT, NEW_OBS_NODES, TIME_INFER,
|
|||
DATATIME_WORD_LIST
|
||||
from enumeration.memory_node_status import MemoryNodeStatus
|
||||
from enumeration.memory_type_enum import MemoryTypeEnum
|
||||
from model.memory.memory_wrap_node import MemoryWrapNode
|
||||
from model.message import Message
|
||||
from worker.memory.memory_base_worker import MemoryBaseWorker
|
||||
from node.memory_wrap_node import MemoryWrapNode
|
||||
from node.message import Message
|
||||
from worker.memory_base_worker import MemoryBaseWorker
|
||||
|
||||
|
||||
class GetObservationWorker(MemoryBaseWorker):
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
from common.response_text_parser import ResponseTextParser
|
||||
from enumeration.message_role_enum import MessageRoleEnum
|
||||
from worker.memory.memory_base_worker import MemoryBaseWorker
|
||||
from worker.memory_base_worker import MemoryBaseWorker
|
||||
|
||||
|
||||
class InfoFilterWorker(MemoryBaseWorker):
|
||||
|
|
|
|||
|
|
@ -3,13 +3,13 @@ import os
|
|||
import time
|
||||
from typing import List, Dict
|
||||
|
||||
from common.logger import Logger
|
||||
from constants.common_constants import NEW_USER_PROFILE, MODIFIED_MEMORIES, RELATED_MEMORIES
|
||||
from enumeration.memory_method_enum import MemoryMethodEnum
|
||||
from model.memory_node import MemoryNode
|
||||
from model.user_attribute import UserAttribute
|
||||
from request.memory import MemoryServiceRequestModel
|
||||
from service.memory_service_bailian import MemoryServiceBailian
|
||||
from memory_scope.utils.logger import Logger
|
||||
from memory_scope.constants.common_constants import NEW_USER_PROFILE, MODIFIED_MEMORIES, RELATED_MEMORIES
|
||||
from memory_scope.enumeration.memory_method_enum import MemoryMethodEnum
|
||||
from memory_scope.node.memory_node import MemoryNode
|
||||
from memory_scope.node.user_attribute import UserAttribute
|
||||
from memory_scope.pipeline.memory import MemoryServiceRequestModel
|
||||
from memory_scope.pipeline.memory_service import MemoryService
|
||||
|
||||
"""
|
||||
任务:随机生成一个用户的画像,随机种子0,并根据用户的画像虚拟一段用户和AI的对话。
|
||||
|
|
@ -139,10 +139,8 @@ for i, msg in enumerate(messages3):
|
|||
|
||||
|
||||
def summary_short(messages):
|
||||
messages_pick_n = len(messages)
|
||||
request: MemoryServiceRequestModel = MemoryServiceRequestModel(
|
||||
messages=messages,
|
||||
messages_pick_n=messages_pick_n,
|
||||
memory_id=memory_id,
|
||||
workspace_id=workspace_id,
|
||||
api_key=api_key,
|
||||
|
|
@ -171,7 +169,6 @@ def summary_short(messages):
|
|||
def summary_long():
|
||||
request: MemoryServiceRequestModel = MemoryServiceRequestModel(
|
||||
messages=[],
|
||||
messages_pick_n=0,
|
||||
memory_id=memory_id,
|
||||
workspace_id=workspace_id,
|
||||
api_key=api_key,
|
||||
|
|
@ -204,7 +201,6 @@ def summary_long():
|
|||
def retrieve(messages):
|
||||
request: MemoryServiceRequestModel = MemoryServiceRequestModel(
|
||||
messages=messages,
|
||||
messages_pick_n=1,
|
||||
memory_id=memory_id,
|
||||
workspace_id=workspace_id,
|
||||
api_key=api_key,
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue