From c8945745759774fc7d863939ccca27e69b01263e Mon Sep 17 00:00:00 2001 From: "jinli.yl" Date: Mon, 9 Jun 2025 16:28:02 +0800 Subject: [PATCH] add experience schema --- .gitignore | 3 +- doc/quick_start.md | 148 ++++++++++++++++++ .../model_service_client.py => client.py} | 0 experiencemaker/config/__init__.py | 5 + .../config/agent_wrapper/simple.json | 1 + experiencemaker/config/config_handler.py | 32 ++++ .../config/context_generator/simple.json | 1 + experiencemaker/config/summarizer/simple.json | 1 + experiencemaker/model/__init__.py | 10 -- experiencemaker/model/base_embedding_model.py | 3 + experiencemaker/model/base_llm.py | 2 + .../openai_compatible_embedding_model.py | 5 +- .../model/openai_compatible_llm.py | 5 +- .../agent_wrapper/agent_wrapper_mixin.py | 17 ++ .../agent_wrapper/base_agent_wrapper.py | 29 +--- .../agent_wrapper/base_agent_wrapper_mixin.py | 27 ---- .../agent_wrapper/naive_agent_wrapper.py | 62 -------- .../module/agent_wrapper/simple_agent.py | 97 ++++++++++++ .../agent_wrapper/simple_agent_prompt.yaml | 28 ++++ .../agent_wrapper/simple_agent_wrapper.py | 21 +++ experiencemaker/module/base_module.py | 2 +- .../base_context_generator.py | 11 +- .../simple_context_generator.py | 40 +++++ .../module/environment/base_environment.py | 2 +- experiencemaker/module/prompt/__init__.py | 0 experiencemaker/module/prompt/prompt_mixin.py | 66 ++++++++ .../module/summarizer/base_summarizer.py | 29 ++-- experiencemaker/schema/trajectory.py | 27 ++-- experiencemaker/service.py | 122 +++++++++++++++ experiencemaker/service/model_service.py | 42 ----- experiencemaker/storage/__init__.py | 3 + experiencemaker/storage/es_vector_store.py | 6 +- experiencemaker/storage/file_vector_store.py | 12 +- experiencemaker/tool/__init__.py | 6 +- experiencemaker/utils/file_handler.py | 2 +- experiencemaker/utils/prompt_handler.py | 30 ++-- experiencemaker/utils/registry.py | 20 ++- 37 files changed, 680 insertions(+), 237 deletions(-) create mode 100644 doc/quick_start.md rename experiencemaker/{service/model_service_client.py => client.py} (100%) create mode 100644 experiencemaker/config/__init__.py create mode 100644 experiencemaker/config/agent_wrapper/simple.json create mode 100644 experiencemaker/config/config_handler.py create mode 100644 experiencemaker/config/context_generator/simple.json create mode 100644 experiencemaker/config/summarizer/simple.json create mode 100644 experiencemaker/module/agent_wrapper/agent_wrapper_mixin.py delete mode 100644 experiencemaker/module/agent_wrapper/base_agent_wrapper_mixin.py delete mode 100644 experiencemaker/module/agent_wrapper/naive_agent_wrapper.py create mode 100644 experiencemaker/module/agent_wrapper/simple_agent.py create mode 100644 experiencemaker/module/agent_wrapper/simple_agent_prompt.yaml create mode 100644 experiencemaker/module/agent_wrapper/simple_agent_wrapper.py create mode 100644 experiencemaker/module/context_generator/simple_context_generator.py create mode 100644 experiencemaker/module/prompt/__init__.py create mode 100644 experiencemaker/module/prompt/prompt_mixin.py create mode 100644 experiencemaker/service.py delete mode 100644 experiencemaker/service/model_service.py diff --git a/.gitignore b/.gitignore index 69361c03..8424f1ef 100644 --- a/.gitignore +++ b/.gitignore @@ -17,5 +17,4 @@ log/ .trash/ runs logs -alfworld_data -beyondagent/dataset/appworld/data + diff --git a/doc/quick_start.md b/doc/quick_start.md new file mode 100644 index 00000000..ce963070 --- /dev/null +++ b/doc/quick_start.md @@ -0,0 +1,148 @@ +# Quick Start + +Here is a simple user guide. + +### Step0: Configuration +- start es +- set env APIKEY / host +- python -m model_service -port 8000 -config simple/path + +### Step1: Own an agent + +Assume you have a runnable [agent code](./mxc_agent.py). +Here, we use a basic LLM combined with a simple react framework including three tools(code, web_search, terminate) as an example. + +```python +class MxcAgent(BaseModel): + llm: BaseLLM | None = Field(default=None) + max_steps: int = Field(default=10) + tools: List[BaseTool] = Field(default_factory=list) + + def think(self, query: str, **kwargs) -> bool: + ... + + def act(self, **kwargs): + ... + + def run(self, query: str, **kwargs) -> List[Message]: + messages: List[Message] = [] + + for i in range(self.max_steps): + should_act: bool = self.think(query, messages=messages, **kwargs) + if should_act: + self.act(messages=messages, **kwargs) + else: + break + return messages + + +query = "Analyze Xiaomi Corporation +agent = MxcAgent(llm=OpenAICompatibleBaseLLM(model_name="qwen3-32b", temperature=0.0001), + max_steps=10, + tools=[CodeTool(), DashscopeSearchTool(), TerminateTool()]) +messages = agent.run(query=query) +answer = messages[-1].content +print(answer) +``` + +### Step2: Implement AgentWrapper + +In order to utilize the **context generator** and **summarizer** capabilities of beyondagent, please inherit from **MxcAgent** and **BaseAgentWrapperMixin** to implement the AgentWrapper. + +Here, you need to customize two parts: +- how to integrate the content message(insight) generated by the `self.context_generator` into the context. +- implement the execute function to output the trajectory. + +Below is a simple example of integrating **trajectory-level insight** into the context. + +```python +from beyondagent.core.module.agent_wrapper.base_agent_wrapper import BaseAgentWrapperMixin + +class MxcAgentWrapper(MxcAgent, BaseAgentWrapperMixin): + def execute(self, query: str, **kwargs) -> Trajectory: + trajectory = Trajectory(steps=messages, query=query) + context_msg = self.context_generator.execute(trajectory=trajectory) + new_query = f""" +previous insight: +{context_msg.content} +Please consider the helpful parts from these in answering the question, to make the response more comprehensive and substantial. + +user query: +{query} + """.strip() + + messages = self.run(new_query, **kwargs) + return Trajectory(query=query, steps=messages, answer=messages[-1].content, done=True) + +``` + +### Step3: Run AgentRunner with insight + +Once you have completed the implementation of the AgentWrapper class, you will be able to utilize the capabilities of +beyondagent. + +Here is an example using **SimpleAgentRunner**. +We first executed two historical tasks, then summarized the experience and made it persistent. +Finally, we utilized the historical experience in a new task. + +[insights demo](./insight.json) + + +```python +from beyondagent.core.module.runner.simple_agent_runner import SimpleAgentRunner + + + +mxc_agent_wrapper = MxcAgentWrapper(llm=OpenAICompatibleBaseLLM(model_name="qwen3-32b", temperature=0.0001), + max_steps=10, + tools=[CodeTool(), DashscopeSearchTool(), TerminateTool()]) +agent_runner = SimpleAgentRunner(agent_wrapper=mxc_agent_wrapper, summarizer="default", context_generator="default") + +# historical tasks +agent_runner.rollout_trajectory(query="Analyze the company Tesla.") +agent_runner.rollout_trajectory(query="Analyze the company Apple.") + +# summary insights and store them +agent_runner.summary_and_store() + +# run agent with historical insights +trajectory = agent_runner.rollout_trajectory(query="Analyze the company Xiaomi Corporation.") +``` + + +### Step4: Evaluation(Optional) + +If we have a reward function that allows us to compare the performance before and after adding context, we can try this +part. + +Use `run_agent` to obtain the answer from the original agent (answer1), and use `run_agent_wrapper` to get the answer +with added insights and experience (answer2). + +Here, the reward function is used to compare and score the two answers. The `reward.reward_value` indicates the win rate +of answer2. + +```python +# task +query = "Analyze Xiaomi Corporation." + +# run agent +agent = MxcAgent(llm=OpenAICompatibleBaseLLM(model_name="qwen3-32b", temperature=0.0001), + max_steps=10, + tools=[CodeTool(), DashscopeSearchTool(), TerminateTool()]) +messages = agent.run(query=query) +answer1 = messages[-1].content + +# agent runner: Assume we already have some historical experience. +mxc_agent_wrapper = MxcAgentWrapper(llm=OpenAICompatibleBaseLLM(model_name="qwen3-32b", temperature=0.0001), + max_steps=10, + tools=[CodeTool(), DashscopeSearchTool(), TerminateTool()]) +agent_runner = SimpleAgentRunner(agent_wrapper=mxc_agent_wrapper, summarizer="default", context_generator="default") +trajectory = agent_runner.rollout_trajectory(query=query) +answer2 = trajectory.answer + +# pair-wise LLM evaluation +from beyondagent.core.module.reward_fn.simple_reward_fn import SimpleRewardFn +reward_fn = SimpleRewardFn(llm=OpenAICompatibleBaseLLM(model_name="qwen3-32b", temperature=0.0001)) +reward = reward_fn.execute(query=query, answer1=answer1, answer2=answer2, eval_times=5) +print(f"final reward={reward.reward_value}") +``` \ No newline at end of file diff --git a/experiencemaker/service/model_service_client.py b/experiencemaker/client.py similarity index 100% rename from experiencemaker/service/model_service_client.py rename to experiencemaker/client.py diff --git a/experiencemaker/config/__init__.py b/experiencemaker/config/__init__.py new file mode 100644 index 00000000..8eea6cd5 --- /dev/null +++ b/experiencemaker/config/__init__.py @@ -0,0 +1,5 @@ +from experiencemaker.config.config_handler import ConfigHandler + +agent_wrapper_config = ConfigHandler(module_name="agent_wrapper") +summarizer_config = ConfigHandler(module_name="summarizer") +context_generator_config = ConfigHandler(module_name="context_generator") diff --git a/experiencemaker/config/agent_wrapper/simple.json b/experiencemaker/config/agent_wrapper/simple.json new file mode 100644 index 00000000..9e26dfee --- /dev/null +++ b/experiencemaker/config/agent_wrapper/simple.json @@ -0,0 +1 @@ +{} \ No newline at end of file diff --git a/experiencemaker/config/config_handler.py b/experiencemaker/config/config_handler.py new file mode 100644 index 00000000..1e471de3 --- /dev/null +++ b/experiencemaker/config/config_handler.py @@ -0,0 +1,32 @@ +from pathlib import Path + +from pydantic import BaseModel, Field, model_validator + +from experiencemaker.utils.file_handler import FileHandler + + +class ConfigHandler(BaseModel): + module_name: str = Field(default=...) + config_dict: dict = Field(default_factory=dict) + + @model_validator(mode="after") + def register_config(self): + module_config_path: Path = Path(__file__).parent / self.module_name + for config_path in module_config_path.iterdir(): + config_name = config_path.stem + config = FileHandler(file_path=config_path).load() + self.config_dict[config_name] = config + return self + + def list_config_names(self): + return list(self.config_dict.keys()) + + def __getattr__(self, item): + if item in self.config_dict: + return self.config_dict[item] + return super().__getattr__(item) + + def __getitem__(self, item): + if item in self.config_dict: + return self.config_dict[item] + return super().__getitem__(item) diff --git a/experiencemaker/config/context_generator/simple.json b/experiencemaker/config/context_generator/simple.json new file mode 100644 index 00000000..9e26dfee --- /dev/null +++ b/experiencemaker/config/context_generator/simple.json @@ -0,0 +1 @@ +{} \ No newline at end of file diff --git a/experiencemaker/config/summarizer/simple.json b/experiencemaker/config/summarizer/simple.json new file mode 100644 index 00000000..4fec8f75 --- /dev/null +++ b/experiencemaker/config/summarizer/simple.json @@ -0,0 +1 @@ +{"a": 1} \ No newline at end of file diff --git a/experiencemaker/model/__init__.py b/experiencemaker/model/__init__.py index 5afcdb48..e69de29b 100644 --- a/experiencemaker/model/__init__.py +++ b/experiencemaker/model/__init__.py @@ -1,10 +0,0 @@ -from experiencemaker.model.openai_compatible_embedding_model import OpenAICompatibleEmbeddingModel -from experiencemaker.model.openai_compatible_llm import OpenAICompatibleBaseLLM - -from experiencemaker.utils.registry import Registry - -LLM_REGISTRY = Registry("llm") -LLM_REGISTRY.register(OpenAICompatibleBaseLLM, "openai_compatible") - -EMBEDDING_MODEL_REGISTRY = Registry("embedding_model") -EMBEDDING_MODEL_REGISTRY.register(OpenAICompatibleEmbeddingModel, "openai_compatible") diff --git a/experiencemaker/model/base_embedding_model.py b/experiencemaker/model/base_embedding_model.py index 3cde45c3..0391f5d8 100644 --- a/experiencemaker/model/base_embedding_model.py +++ b/experiencemaker/model/base_embedding_model.py @@ -5,6 +5,9 @@ from loguru import logger from pydantic import BaseModel, Field from experiencemaker.schema.vector_store_node import VectorStoreNode +from experiencemaker.utils.registry import Registry + +EMBEDDING_MODEL_REGISTRY = Registry("embedding_model") class BaseEmbeddingModel(BaseModel, ABC): diff --git a/experiencemaker/model/base_llm.py b/experiencemaker/model/base_llm.py index f78dc6e2..f5632aa0 100644 --- a/experiencemaker/model/base_llm.py +++ b/experiencemaker/model/base_llm.py @@ -6,7 +6,9 @@ from pydantic import Field, BaseModel from experiencemaker.schema.trajectory import Message, ActionMessage from experiencemaker.tool.base_tool import BaseTool +from experiencemaker.utils.registry import Registry +LLM_REGISTRY = Registry("llm") class BaseLLM(BaseModel, ABC): model_name: str = Field(...) diff --git a/experiencemaker/model/openai_compatible_embedding_model.py b/experiencemaker/model/openai_compatible_embedding_model.py index 79eedee6..bec37d61 100644 --- a/experiencemaker/model/openai_compatible_embedding_model.py +++ b/experiencemaker/model/openai_compatible_embedding_model.py @@ -4,7 +4,7 @@ from typing import Literal, List from openai import OpenAI from pydantic import Field, PrivateAttr, model_validator -from experiencemaker.model.base_embedding_model import BaseEmbeddingModel +from experiencemaker.model.base_embedding_model import BaseEmbeddingModel, EMBEDDING_MODEL_REGISTRY class OpenAICompatibleEmbeddingModel(BaseEmbeddingModel): @@ -71,3 +71,6 @@ class OpenAICompatibleEmbeddingModel(BaseEmbeddingModel): else: # If the input type is neither a string nor a list of strings, throw an exception raise RuntimeError(f"unsupported type={type(input_text)}") + + +EMBEDDING_MODEL_REGISTRY.register(OpenAICompatibleEmbeddingModel, "openai_compatible") diff --git a/experiencemaker/model/openai_compatible_llm.py b/experiencemaker/model/openai_compatible_llm.py index 03814862..d4f5d0a3 100644 --- a/experiencemaker/model/openai_compatible_llm.py +++ b/experiencemaker/model/openai_compatible_llm.py @@ -7,7 +7,7 @@ from openai.types import CompletionUsage from pydantic import Field, PrivateAttr, model_validator from experiencemaker.enumeration.chunk_enum import ChunkEnum -from experiencemaker.model.base_llm import BaseLLM +from experiencemaker.model.base_llm import BaseLLM, LLM_REGISTRY from experiencemaker.schema.trajectory import Message, ActionMessage, ToolCall from experiencemaker.tool.base_tool import BaseTool @@ -150,3 +150,6 @@ class OpenAICompatibleBaseLLM(BaseLLM): elif chunk_enum is ChunkEnum.ERROR: print(f"\n{chunk}", end="") + + +LLM_REGISTRY.register(OpenAICompatibleBaseLLM, "openai_compatible") diff --git a/experiencemaker/module/agent_wrapper/agent_wrapper_mixin.py b/experiencemaker/module/agent_wrapper/agent_wrapper_mixin.py new file mode 100644 index 00000000..5bfc2ae4 --- /dev/null +++ b/experiencemaker/module/agent_wrapper/agent_wrapper_mixin.py @@ -0,0 +1,17 @@ +from abc import ABC + +from pydantic import Field, BaseModel + +from experiencemaker.module.context_generator.base_context_generator import BaseContextGenerator +from experiencemaker.schema.trajectory import Trajectory +from experiencemaker.utils.registry import Registry + + +class AgentWrapperMixin(BaseModel, ABC): + context_generator: BaseContextGenerator | None = Field(default=None) + + def execute(self, query: str, **kwargs) -> Trajectory: + raise NotImplementedError + + +AGENT_WRAPPER_REGISTRY = Registry[AgentWrapperMixin]("agent_wrapper") diff --git a/experiencemaker/module/agent_wrapper/base_agent_wrapper.py b/experiencemaker/module/agent_wrapper/base_agent_wrapper.py index 47844d04..93f6303f 100644 --- a/experiencemaker/module/agent_wrapper/base_agent_wrapper.py +++ b/experiencemaker/module/agent_wrapper/base_agent_wrapper.py @@ -1,26 +1,21 @@ -from abc import ABC from typing import List -from loguru import logger from pydantic import Field -from experiencemaker.module.agent_wrapper.base_agent_wrapper_mixin import BaseAgentWrapperMixin +from experiencemaker.module.agent_wrapper.agent_wrapper_mixin import AgentWrapperMixin from experiencemaker.module.environment.base_environment import BaseEnvironment from experiencemaker.schema.trajectory import Trajectory, Message, StateMessage, ActionMessage, ContextMessage -class BaseAgentWrapper(BaseAgentWrapperMixin, ABC): +class BaseAgentWrapper(AgentWrapperMixin): max_steps: int = Field(default=10) enable_exploration: bool = Field(default=False) trajectory: Trajectory | None = Field(default_factory=Trajectory) def reset(self): - self.trajectory = Trajectory() + self.trajectory.reset() - def before_execute_hook(self, query: str, **kwargs): - self.trajectory.query = query - - def after_step_hook(self, action_msg: ActionMessage, next_state: StateMessage, **kwargs): + def after_step_hook(self, **kwargs): raise NotImplementedError def build_messages(self, @@ -32,7 +27,7 @@ class BaseAgentWrapper(BaseAgentWrapperMixin, ABC): def explore_messages(self, messages: List[Message], **kwargs) -> List[Message]: raise NotImplementedError - def action_parser(self, action_msg: ActionMessage, **kwargs) -> ActionMessage: + def action_parser(self, action_msg: ActionMessage) -> ActionMessage: # noqa return action_msg def generate_action(self, @@ -48,13 +43,10 @@ class BaseAgentWrapper(BaseAgentWrapperMixin, ABC): action_msg: ActionMessage = self.llm.chat(messages, tools=env.tools) return self.action_parser(action_msg) - def after_execute_hook(self, **kwargs): - return self.trajectory - def execute(self, query: str, env: BaseEnvironment = None, **kwargs) -> Trajectory: - self.before_execute_hook(query=query, **kwargs) - + self.trajectory.query = query current_state = env.current_state + for i in range(self.max_steps): self.trajectory.current_step = i @@ -62,17 +54,12 @@ class BaseAgentWrapper(BaseAgentWrapperMixin, ABC): context_msg: ContextMessage | None = None if self.context_generator: context_msg = self.context_generator.execute(trajectory=self.trajectory, **kwargs) - if context_msg.content: - logger.info(f"step{i}.context_msg={context_msg.content}") # generate action action_msg = self.generate_action(state=current_state, context_msg=context_msg, env=env, **kwargs) - logger.info(f"step{i} ====== reasoning_content ======\n{action_msg.reasoning_content}\n\n" - f"====== content ======\n{action_msg.content}\ntool_calls={action_msg.tool_calls}") # generate next state next_state, reward, done, info = env.step(action_msg, trajectory=self.trajectory, **kwargs) - logger.info(f"step{i}.next_state={next_state.content} reward={reward.reward_value} done={done} info={info}") self.after_step_hook(step_index=i, current_state=current_state, @@ -88,4 +75,4 @@ class BaseAgentWrapper(BaseAgentWrapperMixin, ABC): current_state = next_state - return self.after_execute_hook(**kwargs) + return self.trajectory diff --git a/experiencemaker/module/agent_wrapper/base_agent_wrapper_mixin.py b/experiencemaker/module/agent_wrapper/base_agent_wrapper_mixin.py deleted file mode 100644 index 9dc7b82a..00000000 --- a/experiencemaker/module/agent_wrapper/base_agent_wrapper_mixin.py +++ /dev/null @@ -1,27 +0,0 @@ -from abc import ABC - -from pydantic import Field - -from experiencemaker.module.base_module import BaseModule -from experiencemaker.module.context_generator.base_context_generator import BaseContextGenerator -from experiencemaker.schema.trajectory import Trajectory, ActionMessage, Message - - -class BaseAgentWrapperMixin(BaseModule, ABC): - context_generator: BaseContextGenerator | None = Field(default=None) - - def execute(self, query: str, **kwargs) -> Trajectory: - raise NotImplementedError - - -class MockAgentWrapper(BaseAgentWrapperMixin): - - def execute(self, query: str, **kwargs) -> Trajectory: - user_message = Message(content=query) - answer_message = ActionMessage(content="hello world") - - traj = Trajectory(steps=[user_message, answer_message], - done=True, - query=query, - answer=answer_message.content) - return traj diff --git a/experiencemaker/module/agent_wrapper/naive_agent_wrapper.py b/experiencemaker/module/agent_wrapper/naive_agent_wrapper.py deleted file mode 100644 index f5e769bf..00000000 --- a/experiencemaker/module/agent_wrapper/naive_agent_wrapper.py +++ /dev/null @@ -1,62 +0,0 @@ -import datetime - -from experiencemaker.module.agent_wrapper.base_agent_wrapper_v2 import BaseAgentWrapperV2 -from loguru import logger - -from experiencemaker.module.environment.base_environment import BaseEnvironment -from experiencemaker.schema.trajectory import Message, StateMessage, ContextMessage, ActionMessage - - -class NaiveAgentWrapper(BaseAgentWrapperV2): - - def generate_action(self, - state: StateMessage, - context_msg: ContextMessage | None, - env: BaseEnvironment, - **kwargs) -> ActionMessage: - - tool_names = [x.name for x in env.tools] - if self.trajectory.current_step == 0: - now_time = datetime.datetime.now().strftime('%Y-%m-%d %H:%M:%S') - insight_tag: bool = True if context_msg is not None and context_msg.content else False - user_prompt = self.prompt_handler.prompt_format( - prompt_name="role_prompt", - insight_tag=insight_tag, - time=now_time, - tools=", ".join(tool_names), - previous_insight=context_msg.content if insight_tag else "", - query=self.trajectory.query) - # When using the reasoning models of Qwen3 or DeepSeek R1, it is not recommended to use system prompt. - self.trajectory.steps.append(Message(content=user_prompt)) - - elif self.trajectory.metadata.get("has_terminate_tool") is True: - user_prompt = self.prompt_handler.final_prompt.format(query=self.trajectory.query) - self.trajectory.steps.append(Message(content=user_prompt)) - - else: - user_prompt = self.prompt_handler.next_prompt.format(query=self.trajectory.query) - self.trajectory.steps.append(Message(content=user_prompt)) - - if self.trajectory.metadata.get("has_terminate_tool") is True: - action_msg: ActionMessage = self.llm.chat(self.trajectory.steps) - logger.info(f"step{self.trajectory.current_step} size={len(self.trajectory.steps)} user_prompt={user_prompt}") - - else: - action_msg: ActionMessage = self.llm.chat(messages, tools=env.tools) - logger.info(f"step{self.trajectory.current_step} size={len(messages)} user_prompt={user_prompt} " - f"tool_names={tool_names}") - - for tool in action_msg.tool_calls: - if tool.name == "terminate": - self.trajectory.metadata["has_terminate_tool"] = True - break - - self.trajectory.add_step(action_msg) - return action_msg - - - def after_step_hook(self, action_msg: ActionMessage, next_state: StateMessage, done: bool = False, **kwargs): - if done: - self.trajectory.answer = action_msg.content - else: - self.trajectory.add_step(next_state) diff --git a/experiencemaker/module/agent_wrapper/simple_agent.py b/experiencemaker/module/agent_wrapper/simple_agent.py new file mode 100644 index 00000000..2e0ced3c --- /dev/null +++ b/experiencemaker/module/agent_wrapper/simple_agent.py @@ -0,0 +1,97 @@ +import datetime +from pathlib import Path +from typing import List + +from loguru import logger +from pydantic import Field, BaseModel + +from experiencemaker.model.base_llm import BaseLLM +from experiencemaker.module.prompt.prompt_mixin import PromptMixin +from experiencemaker.schema.trajectory import Message, ActionMessage, ToolCall, StateMessage +from experiencemaker.tool.base_tool import BaseTool + + +class SimpleAgentContext(BaseModel): + current_step: int = Field(default=-1) + query: str = Field(default="") + previous_experience: str = Field(default="") + messages: List[Message] = Field(default_factory=list) + metadata: dict = Field(default_factory=dict) + has_terminate_tool: bool = Field(default=False) + + +class SimpleAgent(PromptMixin): + llm: BaseLLM | None = Field(default=None) + max_steps: int = Field(default=10) + tools: List[BaseTool] = Field(default_factory=list) + prompt_file_path: Path = Path(__file__).parent / "simple_agent_prompt.yaml" + + def think(self, context: SimpleAgentContext): + now_time = datetime.datetime.now().strftime('%Y-%m-%d %H:%M:%S') + tool_names = [x.name for x in self.tools] + + if context.current_step == 0: + user_prompt = self.prompt_format(prompt_name="role_prompt", + experience_tag=False if context.previous_experience else True, + time=now_time, + tools=", ".join(tool_names), + previous_insight=context.previous_experience, + query=context.query) + + elif context.has_terminate_tool: + user_prompt = self.prompt_format(prompt_name="final_prompt", query=context.query) + + else: + user_prompt = self.prompt_format(prompt_name="next_prompt", query=context.query) + + context.messages.append(Message(content=user_prompt)) + logger.info(f"step.{context.current_step} user_prompt={user_prompt}") + + if context.has_terminate_tool: + action_msg: ActionMessage = self.llm.chat(context.messages) + + else: + action_msg: ActionMessage = self.llm.chat(context.messages, tools=self.tools) + for tool in action_msg.tool_calls: + if tool.name == "terminate": + context.has_terminate_tool = True + break + + context.messages.append(action_msg) + action_msg_context: str = action_msg.content + "\n\n" + action_msg.reasoning_content + logger.info(f"step.{context.current_step} action_msg_context={action_msg_context} " + f"tool_calls={action_msg.tool_calls}") + return True if action_msg.tool_calls else False + + def act(self, context: SimpleAgentContext): + action_msg = context.messages[-1] + assert isinstance(action_msg, ActionMessage) + + tool_dict = {tool.name: tool for tool in self.tools} + + new_tool_calls: List[ToolCall] = [] + for tool_call in action_msg.tool_calls: + if tool_call.name not in tool_dict: + continue + + new_tool_call = tool_call.model_copy(deep=True) + tool = tool_dict[tool_call.name] + new_tool_call.result = tool.execute(**tool_call.argument_dict) + new_tool_calls.append(new_tool_call) + + state_msg = StateMessage(tool_calls=new_tool_calls) + context.messages.append(state_msg) + logger.info(f"step.{context.current_step} state_msg_context={state_msg.content}") + + def run(self, query: str, previous_experience: str) -> List[Message]: + context: SimpleAgentContext = SimpleAgentContext(query=query, previous_experience=previous_experience) + + for i in range(self.max_steps): + context.current_step = i + + should_act: bool = self.think(context) + if should_act: + self.act(context) + else: + break + return context.messages diff --git a/experiencemaker/module/agent_wrapper/simple_agent_prompt.yaml b/experiencemaker/module/agent_wrapper/simple_agent_prompt.yaml new file mode 100644 index 00000000..514c2ebe --- /dev/null +++ b/experiencemaker/module/agent_wrapper/simple_agent_prompt.yaml @@ -0,0 +1,28 @@ +role_prompt: | + You are a helpful assistant named BeyondAgent. + The current time is {time}. + + Please proactively choose the most suitable tool or combination of tools based on the user's question, including {tools} etc. + For complex tasks, you can break down the problem step by step and use different tools to solve it incrementally. + Please determine the response language based on the language of the user's question. + [experience_tag] + [experience_tag]Previous Experience + [experience_tag]{previous_experience} + [experience_tag]Please consider the helpful parts from these in answering the question, to make the response more comprehensive and substantial. + + User's question + {query} + +next_prompt: | + User's question + {query} + + Plan the most suitable tool or combination of tools based on the context and the user's question. + For complex tasks, break down the problem and use different tools step by step to solve it, don't give up easily. + Try calling the same tool multiple times with different parameters to obtain information from various perspectives. + If the task is completed and the user's question can now be answered, use the **terminate** tool. + +final_prompt: | + Please integrate the context and provide a complete answer to the user's question: + {query} + diff --git a/experiencemaker/module/agent_wrapper/simple_agent_wrapper.py b/experiencemaker/module/agent_wrapper/simple_agent_wrapper.py new file mode 100644 index 00000000..124ced7c --- /dev/null +++ b/experiencemaker/module/agent_wrapper/simple_agent_wrapper.py @@ -0,0 +1,21 @@ +from experiencemaker.module.agent_wrapper.agent_wrapper_mixin import AgentWrapperMixin, AGENT_WRAPPER_REGISTRY +from experiencemaker.module.agent_wrapper.simple_agent import SimpleAgent +from experiencemaker.schema.trajectory import Trajectory + + +class SimpleAgentWrapper(SimpleAgent, AgentWrapperMixin): + + def execute(self, query: str, **kwargs) -> Trajectory: + trajectory = Trajectory(query=query) + context_msg = self.context_generator.execute(trajectory=trajectory) + previous_experience = context_msg.content + + messages = self.run(query, previous_experience) + + trajectory.steps = messages + trajectory.answer = messages[-1].content + trajectory.done = True + return trajectory + + +AGENT_WRAPPER_REGISTRY.register(SimpleAgentWrapper, "simple") diff --git a/experiencemaker/module/base_module.py b/experiencemaker/module/base_module.py index 4cad23c9..e45efcba 100644 --- a/experiencemaker/module/base_module.py +++ b/experiencemaker/module/base_module.py @@ -12,8 +12,8 @@ from experiencemaker.utils.prompt_handler import PromptHandler class BaseModule(BaseModel, ABC): prompt_dir: str | None = Field(default=None) prompt_file: str | None = Field(default=None) - prompt_handler: PromptHandler | None = Field(default=None) + prompt_handler: PromptHandler | None = Field(default=None) llm: BaseLLM | None = Field(default=None) embedding_model: BaseEmbeddingModel | None = Field(default=None) diff --git a/experiencemaker/module/context_generator/base_context_generator.py b/experiencemaker/module/context_generator/base_context_generator.py index 200b46de..4934e372 100644 --- a/experiencemaker/module/context_generator/base_context_generator.py +++ b/experiencemaker/module/context_generator/base_context_generator.py @@ -1,15 +1,15 @@ from abc import ABC from typing import List -from pydantic import Field +from pydantic import Field, BaseModel -from experiencemaker.module.base_module import BaseModule from experiencemaker.schema.trajectory import Trajectory, ContextMessage from experiencemaker.schema.vector_store_node import VectorStoreNode from experiencemaker.storage.base_vector_store import BaseVectorStore +from experiencemaker.utils.registry import Registry -class BaseContextGenerator(BaseModule, ABC): +class BaseContextGenerator(BaseModel, ABC): vector_store: BaseVectorStore | None = Field(default=None) def _build_retrieve_query(self, trajectory: Trajectory, **kwargs) -> str: @@ -31,7 +31,4 @@ class BaseContextGenerator(BaseModule, ABC): return context_msg -class MockContextGenerator(BaseContextGenerator): - - def execute(self, trajectory: Trajectory, **kwargs) -> ContextMessage: - return ContextMessage(content="mock context") +CONTEXT_GENERATOR_REGISTRY = Registry[BaseContextGenerator]("context_generator") diff --git a/experiencemaker/module/context_generator/simple_context_generator.py b/experiencemaker/module/context_generator/simple_context_generator.py new file mode 100644 index 00000000..4beb15a0 --- /dev/null +++ b/experiencemaker/module/context_generator/simple_context_generator.py @@ -0,0 +1,40 @@ +from typing import List + +from experiencemaker.module.context_generator.base_context_generator import BaseContextGenerator, \ + CONTEXT_GENERATOR_REGISTRY +from experiencemaker.schema.trajectory import Trajectory, ContextMessage +from experiencemaker.schema.vector_store_node import VectorStoreNode + + +class SimpleContextGenerator(BaseContextGenerator): + + def _build_retrieve_query(self, trajectory: Trajectory, **kwargs) -> str: + query = "" + if trajectory.current_step == 0: + query = trajectory.query + return query + + def _retrieve_by_query(self, trajectory: Trajectory, query: str, **kwargs) -> List[VectorStoreNode]: + if not query: + return [] + + return self.vector_store.retrieve_by_query(query=query, top_k=self.vector_store_top_k) + + def _generate_context_message(self, + trajectory: Trajectory, + nodes: List[VectorStoreNode], + **kwargs) -> ContextMessage: + if not nodes: + return ContextMessage(content="") + + content = "" + for node in nodes: + experience = node.metadata.get("experience", "") + if not experience: + continue + + content += f"- {node.content} {experience}\n" + return ContextMessage(content=content.strip()) + + +CONTEXT_GENERATOR_REGISTRY.register(SimpleContextGenerator, "simple") diff --git a/experiencemaker/module/environment/base_environment.py b/experiencemaker/module/environment/base_environment.py index 9e52d3c1..a4efd655 100644 --- a/experiencemaker/module/environment/base_environment.py +++ b/experiencemaker/module/environment/base_environment.py @@ -12,7 +12,7 @@ from experiencemaker.tool.base_tool import BaseTool class BaseEnvironment(BaseModule): tools: List[BaseTool] = Field(default_factory=list) reward_fns: List[BaseRewardFn] = Field(default_factory=list) - current_state: StateMessage | None = Field(default=None) + current_state: StateMessage = Field(default_factory=StateMessage) metadata: dict = Field(default_factory=dict, description="add query / answer and etc for reward calculating!") def reset(self): diff --git a/experiencemaker/module/prompt/__init__.py b/experiencemaker/module/prompt/__init__.py new file mode 100644 index 00000000..e69de29b diff --git a/experiencemaker/module/prompt/prompt_mixin.py b/experiencemaker/module/prompt/prompt_mixin.py new file mode 100644 index 00000000..56fef91e --- /dev/null +++ b/experiencemaker/module/prompt/prompt_mixin.py @@ -0,0 +1,66 @@ +from pathlib import Path + +import yaml +from loguru import logger +from pydantic import BaseModel, Field, model_validator + + +class PromptMixin(BaseModel): + prompt_file_path: Path | str = Field(default=None) + prompt_dict: dict = Field(default_factory=dict) + + @model_validator(mode="after") + def init_prompt(self): + if self.prompt_dict: + logger.info(f"load prompt from dict, keys={self.prompt_dict.keys()}") + + if self.prompt_file_path is not None: + if isinstance(self.prompt_file_path, str): + self.prompt_file_path = Path(self.prompt_file_path) + + if not self.prompt_file_path.exists(): + logger.warning(f"prompt_file_path={self.prompt_file_path} not exists!") + + else: + with self.prompt_file_path.open("r") as f: + for k, v in yaml.load(f, yaml.FullLoader): + if k not in self.prompt_dict: + self.prompt_dict[k] = v + logger.info(f"add prompt_dict key={k}") + else: + logger.warning(f"key={k} is already exists in prompt_dict!") + + return self + + def prompt_format(self, prompt_name: str, **kwargs): + prompt = self.prompt_dict[prompt_name] + + flag_kwargs = {k: v for k, v in kwargs.items() if isinstance(v, bool)} + other_kwargs = {k: v for k, v in kwargs.items() if not isinstance(v, bool)} + + if flag_kwargs: + split_prompt = [] + for line in prompt.strip().split("\n"): + hit = False + hit_flag = True + for key, flag in kwargs.items(): + if not line.startswith(f"[{key}]"): + continue + + else: + hit = True + hit_flag = flag + line = line.strip(f"[{key}]") + break + + if not hit: + split_prompt.append(line) + elif hit_flag: + split_prompt.append(line) + + prompt = "\n".join(split_prompt) + + if other_kwargs: + prompt = prompt.format(**other_kwargs) + + return prompt diff --git a/experiencemaker/module/summarizer/base_summarizer.py b/experiencemaker/module/summarizer/base_summarizer.py index 2ccb2c6c..b06889d2 100644 --- a/experiencemaker/module/summarizer/base_summarizer.py +++ b/experiencemaker/module/summarizer/base_summarizer.py @@ -1,37 +1,26 @@ +from abc import ABC from typing import List -from pydantic import Field +from pydantic import Field, BaseModel -from experiencemaker.module.base_module import BaseModule -from experiencemaker.schema.trajectory import Trajectory, Sample, SummaryMessage +from experiencemaker.schema.trajectory import Trajectory, Sample from experiencemaker.storage.base_vector_store import BaseVectorStore -class BaseSummarizer(BaseModule): +class BaseSummarizer(BaseModel, ABC): vector_store: BaseVectorStore | None = Field(default=None) - def extract_samples(self, trajectories: List[Trajectory], **kwargs) -> List[Sample]: + def _extract_samples(self, trajectories: List[Trajectory], **kwargs) -> List[Sample]: raise NotImplementedError - def insert_into_vector_store(self, samples: List[Sample], **kwargs): + def _insert_into_database(self, samples: List[Sample], **kwargs): raise NotImplementedError def execute(self, trajectories: List[Trajectory], return_samples: bool = True, **kwargs) -> List[Sample]: - samples: List[Sample] = self.extract_samples(trajectories, **kwargs) - self.insert_into_vector_store(samples, **kwargs) + samples: List[Sample] = self._extract_samples(trajectories, **kwargs) + self._insert_into_database(samples, **kwargs) if return_samples: return samples - return [] - - -class MockSummarizer(BaseSummarizer): - - def execute(self, trajectories: List[Trajectory], return_samples: bool = True, **kwargs) -> List[Sample]: - tip_message = SummaryMessage(content="I am a mock summarizer.") - - if return_samples: - return [Sample(steps=[tip_message])] - - return [] + return [] \ No newline at end of file diff --git a/experiencemaker/schema/trajectory.py b/experiencemaker/schema/trajectory.py index 8842eb08..64dd5b09 100644 --- a/experiencemaker/schema/trajectory.py +++ b/experiencemaker/schema/trajectory.py @@ -49,19 +49,12 @@ class Message(BaseModel): content: str | bytes = Field(default="") reasoning_content: str = Field(default="") tool_calls: List[ToolCall] = Field(default_factory=list) - timestamp: str = Field( - default_factory=lambda: datetime.datetime.now().strftime( - "%Y-%m-%d %H:%M:%S.%f", - ), - ) + timestamp: str = Field(default_factory=lambda: datetime.datetime.now().strftime("%Y-%m-%d %H:%M:%S.%f")) metadata: dict = Field(default_factory=dict) @property def simple_dict(self) -> dict: - result = { - "role": self.role.value, - "content": self.content, - } + result = {"role": self.role.value, "content": self.content} if self.tool_calls: result["tool_calls"] = [x.simple_dict for x in self.tool_calls] return result @@ -77,10 +70,7 @@ class StateMessage(Message): @property def simple_dict(self) -> dict: - result = { - "role": self.role.value, - "content": self.content, - } + result = super().simple_dict if self.tool_call_id: result["tool_call_id"] = self.tool_call_id return result @@ -110,6 +100,7 @@ class Sample(BaseModel): class Trajectory(BaseModel): id: str = Field(default_factory=lambda: uuid4().hex) steps: List[Message] = Field(default_factory=list) + current_step: int = Field(default=0) done: bool = Field(default=False) query: str = Field(default="") @@ -120,8 +111,18 @@ class Trajectory(BaseModel): self.steps.append(step) def reset(self): + self.id = uuid4().hex self.steps.clear() + self.current_step = 0 self.done = False self.query = "" self.answer = "" self.metadata.clear() + + +class Experience(BaseModel): + experience_id: str = Field(default_factory=lambda: uuid4().hex, description="experience unique id") + experience_desc: str = Field(default="", description="use condition or use purpose for vector matching") + experience_content: str | bytes = Field(default="", description="content of the experience") + experience_score: float = Field(default=0.0, description="score of the experience") + metadata: dict = Field(default_factory=dict, description="additional metadata") diff --git a/experiencemaker/service.py b/experiencemaker/service.py new file mode 100644 index 00000000..8c4a93e2 --- /dev/null +++ b/experiencemaker/service.py @@ -0,0 +1,122 @@ +from typing import List + +import uvicorn +from fastapi import FastAPI +from loguru import logger + +from experiencemaker.model import LLM_REGISTRY, EMBEDDING_MODEL_REGISTRY +from experiencemaker.model.base_embedding_model import BaseEmbeddingModel +from experiencemaker.model.base_llm import BaseLLM +from experiencemaker.module.agent_wrapper.base_agent_wrapper import BaseAgentWrapper +from experiencemaker.module.context_generator.base_context_generator import BaseContextGenerator +from experiencemaker.module.summarizer.base_summarizer import BaseSummarizer +from experiencemaker.schema.request import AgentWrapperRequest, ContextGeneratorRequest, SummarizerRequest +from experiencemaker.schema.response import AgentWrapperResponse, ContextGeneratorResponse, SummarizerResponse +from experiencemaker.schema.trajectory import ContextMessage, Trajectory, Sample +from experiencemaker.storage import VECTOR_STORE_REGISTRY +from experiencemaker.storage.base_vector_store import BaseVectorStore + +app = FastAPI() +from pydantic import BaseModel, Field, model_validator + + +class ExperienceMakerHttpService(BaseModel): + host: str = Field(default="0.0.0.0") + port: int = Field(default=8001) + timeout_keep_alive: int = Field(default=600000) + limit_concurrency: int = Field(default=32) + + llm_config: dict = Field(default_factory=dict) + embedding_model_config: dict = Field(default_factory=dict) + vector_store_config: dict = Field(default_factory=dict) + + agent_wrapper_config: dict = Field(default_factory=dict) + context_generator_config: dict = Field(default_factory=dict) + summarizer_config: dict = Field(default_factory=dict) + + llm: BaseLLM | None = Field(default=None) + embedding_model: BaseEmbeddingModel | None = Field(default=None) + vector_store: BaseVectorStore | None = Field(default=None) + + agent_wrapper: BaseAgentWrapper | None = Field(default=None) + context_generator: BaseContextGenerator | None = Field(default=None) + summarizer: BaseSummarizer | None = Field(default=None) + + @staticmethod + def init_llm(llm_config: dict): + backend = llm_config.pop("backend", None) + assert backend is not None, "llm must have a backend like `openai_compatible`." + assert backend in LLM_REGISTRY, f"llm backend={backend} not supported. supported backend={LLM_REGISTRY.registered_modules}" + llm = LLM_REGISTRY[backend](**llm_config) + logger.info(f"llm is inited with backend={backend} params={llm_config}") + return llm + + @staticmethod + def init_embedding_model(embedding_model_config: dict): + backend = embedding_model_config.pop("backend", None) + assert backend is not None, "embedding_model must have a backend like `openai_compatible`." + assert backend in EMBEDDING_MODEL_REGISTRY, f"embedding_model backend={backend} not supported. supported backend={EMBEDDING_MODEL_REGISTRY.registered_modules}" + embedding_model = EMBEDDING_MODEL_REGISTRY[backend](**embedding_model_config) + logger.info(f"embedding_model is inited with backend={backend} params={embedding_model_config}") + return embedding_model + + @staticmethod + def init_vector_store(vector_store_config: dict): + backend = vector_store_config.pop("backend", None) + assert backend is not None, "vector_store must have a backend like `elasticsearch`." + assert backend in VECTOR_STORE_REGISTRY, f"vector_store backend={backend} not supported. supported backend={VECTOR_STORE_REGISTRY.registered_modules}" + vector_store = VECTOR_STORE_REGISTRY[backend](**vector_store_config) + logger.info(f"vector_store is inited with backend={backend} params={vector_store_config}") + return vector_store + + @model_validator(mode="after") + def init_modules(self): + if self.llm_config: + self.llm = self.init_llm(self.llm_config) + + if self.embedding_model_config: + self.embedding_model = self.init_embedding_model(self.embedding_model_config) + + if self.vector_store_config: + self.vector_store = self.init_vector_store(self.vector_store_config) + + + + + +@app.post('/agent_wrapper', response_model=AgentWrapperResponse) +def call_agent_wrapper(request: AgentWrapperRequest): + module: BaseAgentWrapper = request.load_from_path() + trajectory: Trajectory = module.execute(request.query, **request.metadata) + return AgentWrapperResponse(trajectory=trajectory) + + +@app.post('/context_generator', response_model=ContextGeneratorResponse) +def call_context_generator(request: ContextGeneratorRequest): + module: BaseContextGenerator = request.load_from_path() + context_msg: ContextMessage = module.execute(request.trajectory, **request.metadata) + return ContextGeneratorResponse(context_msg=context_msg) + + +@app.post('/summarizer', response_model=SummarizerResponse) +def call_summarizer(request: SummarizerRequest): + module: BaseSummarizer = request.load_from_path() + samples: List[Sample] = module.execute(request.trajectories, request.return_samples, **request.metadata) + return SummarizerResponse(extract_samples=samples) + + +if __name__ == '__main__': + uvicorn.run(app, host="0.0.0.0", port=8000, timeout_keep_alive=600000, limit_concurrency=32) + # from experiencemaker.config import summarizer_config + # print(summarizer_config.simple) + + from experiencemaker.config.config_handler import ConfigHandler + + summarizer_config = ConfigHandler(module_name="summarizer") + context_generator_config = ConfigHandler(module_name="context_generator") + print(context_generator_config.config_dict) + + + +# launch with: +# python -m experiencemaker.service.model_service diff --git a/experiencemaker/service/model_service.py b/experiencemaker/service/model_service.py deleted file mode 100644 index 3f00afe0..00000000 --- a/experiencemaker/service/model_service.py +++ /dev/null @@ -1,42 +0,0 @@ -from experiencemaker.utils.logger import init_logger -init_logger() - -from typing import List -from fastapi import FastAPI -from experiencemaker.module.agent_wrapper.base_agent_wrapper import BaseAgentWrapper -from experiencemaker.module.context_generator.base_context_generator import BaseContextGenerator -from experiencemaker.module.summarizer.base_summarizer import BaseSummarizer -from experiencemaker.schema.request import AgentWrapperRequest, ContextGeneratorRequest, SummarizerRequest -from experiencemaker.schema.response import AgentWrapperResponse, ContextGeneratorResponse, SummarizerResponse -from experiencemaker.schema.trajectory import ContextMessage, Trajectory, Sample -import uvicorn - -app = FastAPI() - - -@app.post('/agent_wrapper', response_model=AgentWrapperResponse) -def call_agent_wrapper(request: AgentWrapperRequest): - module: BaseAgentWrapper = request.load_from_path() - trajectory: Trajectory = module.execute(request.query, **request.metadata) - return AgentWrapperResponse(trajectory=trajectory) - - -@app.post('/context_generator', response_model=ContextGeneratorResponse) -def call_context_generator(request: ContextGeneratorRequest): - module: BaseContextGenerator = request.load_from_path() - context_msg: ContextMessage = module.execute(request.trajectory, **request.metadata) - return ContextGeneratorResponse(context_msg=context_msg) - - -@app.post('/summarizer', response_model=SummarizerResponse) -def call_summarizer(request: SummarizerRequest): - module: BaseSummarizer = request.load_from_path() - samples: List[Sample] = module.execute(request.trajectories, request.return_samples, **request.metadata) - return SummarizerResponse(extract_samples=samples) - - -if __name__ == '__main__': - uvicorn.run(app, host="0.0.0.0", port=8000, timeout_keep_alive=600000, limit_concurrency=32) - -# launch with: -# python -m experiencemaker.service.model_service diff --git a/experiencemaker/storage/__init__.py b/experiencemaker/storage/__init__.py index e69de29b..5df9193b 100644 --- a/experiencemaker/storage/__init__.py +++ b/experiencemaker/storage/__init__.py @@ -0,0 +1,3 @@ +from experiencemaker.utils.registry import Registry + +VECTOR_STORE_REGISTRY = Registry("vector_store") diff --git a/experiencemaker/storage/es_vector_store.py b/experiencemaker/storage/es_vector_store.py index 92526f9b..b0de2d85 100644 --- a/experiencemaker/storage/es_vector_store.py +++ b/experiencemaker/storage/es_vector_store.py @@ -7,6 +7,7 @@ from loguru import logger from pydantic import Field, PrivateAttr, model_validator from experiencemaker.schema.vector_store_node import VectorStoreNode +from experiencemaker.storage import VECTOR_STORE_REGISTRY from experiencemaker.storage.base_vector_store import BaseVectorStore @@ -60,7 +61,7 @@ class EsVectorStore(BaseVectorStore): node = VectorStoreNode(**doc["_source"]) node.unique_id = doc["_id"] if "_score" in doc: - node.metadata["score"] = doc["_score"] - 1 + node.metadata["_score"] = doc["_score"] - 1 return node def exist_id(self, doc_id: str): @@ -167,3 +168,6 @@ class EsVectorStore(BaseVectorStore): self.retrieve_filters.clear() return nodes + + +VECTOR_STORE_REGISTRY.register(EsVectorStore, "elasticsearch") diff --git a/experiencemaker/storage/file_vector_store.py b/experiencemaker/storage/file_vector_store.py index 4544a7c9..8d8bf772 100644 --- a/experiencemaker/storage/file_vector_store.py +++ b/experiencemaker/storage/file_vector_store.py @@ -8,6 +8,7 @@ from loguru import logger from pydantic import Field, model_validator, PrivateAttr from experiencemaker.schema.vector_store_node import VectorStoreNode +from experiencemaker.storage import VECTOR_STORE_REGISTRY from experiencemaker.storage.base_vector_store import BaseVectorStore @@ -39,22 +40,22 @@ class FileVectorStore(BaseVectorStore): self.index_path.touch(exist_ok=True) def load(self) -> List[VectorStoreNode]: + nodes = [] with self._thread_lock: - nodes = [] with open(self.index_path) as f: for line in f: if line.strip(): nodes.append(VectorStoreNode(**json.loads(line))) - return nodes + return nodes def _load(self) -> List[VectorStoreNode]: + nodes = [] with self._thread_lock: - nodes = [] with open(self.index_path) as f: for line in f: if line.strip(): nodes.append(VectorStoreNode(**json.loads(line))) - return nodes + return nodes def _dump(self, nodes: List[VectorStoreNode]): with self._thread_lock: @@ -130,3 +131,6 @@ class FileVectorStore(BaseVectorStore): nodes = sorted(nodes, key=lambda x: x.metadata["score"], reverse=True) return nodes[:top_k] + + +VECTOR_STORE_REGISTRY.register(FileVectorStore, "local_file") diff --git a/experiencemaker/tool/__init__.py b/experiencemaker/tool/__init__.py index f342c5ec..4df01cc7 100644 --- a/experiencemaker/tool/__init__.py +++ b/experiencemaker/tool/__init__.py @@ -1,6 +1,6 @@ -from experiencemaker.tool.python_tools.code_tool import CodeTool -from experiencemaker.tool.python_tools.dashscope_search_tool import DashscopeSearchTool -from experiencemaker.tool.python_tools.terminate_tool import TerminateTool +from experiencemaker.tool.code_tool import CodeTool +from experiencemaker.tool.dashscope_search_tool import DashscopeSearchTool +from experiencemaker.tool.terminate_tool import TerminateTool from experiencemaker.utils.registry import Registry diff --git a/experiencemaker/utils/file_handler.py b/experiencemaker/utils/file_handler.py index 66850ac9..8ca8ac5c 100644 --- a/experiencemaker/utils/file_handler.py +++ b/experiencemaker/utils/file_handler.py @@ -7,7 +7,7 @@ from pydantic import BaseModel, Field, PrivateAttr class FileHandler(BaseModel): - file_path: str = Field(default=...) + file_path: str | Path = Field(default=...) _obj: Any = PrivateAttr() def __init__(self, **kwargs): diff --git a/experiencemaker/utils/prompt_handler.py b/experiencemaker/utils/prompt_handler.py index fe68956e..30c48b87 100644 --- a/experiencemaker/utils/prompt_handler.py +++ b/experiencemaker/utils/prompt_handler.py @@ -2,27 +2,29 @@ import os import yaml from loguru import logger -from pydantic import BaseModel, Field +from pydantic import BaseModel, Field, model_validator class PromptHandler(BaseModel): - dir_path: str = Field(default="") + file_path: str = Field(default="") prompt_dict: dict = Field(default_factory=dict) - def add_prompt_file(self, file_name: str): - prompt_path = os.path.join(self.dir_path, file_name + ".yaml") - self._add_prompt_file(prompt_path) + @model_validator(mode="after") + def init_prompt_dict(self): + if self.prompt_dict: + logger.info(f"load prompt from dict, keys={self.prompt_dict.keys()}") - def _add_prompt_file(self, prompt_path: str): - if os.path.exists(prompt_path): - with open(prompt_path) as f: - prompt_dict: dict = yaml.load(f, yaml.FullLoader) - self.update_prompt_dict(prompt_dict) - else: - logger.warning(f"prompt_path={prompt_path} not exists!") + if self.file_path: + self.load_file_prompt() + logger.info(f"load prompt from file_path, keys={self.prompt_dict.keys()}") - def update_prompt_dict(self, prompt_dict: dict): - self.prompt_dict.update(prompt_dict) + def load_file_prompt(self): + if not os.path.exists(self.file_path): + raise RuntimeError(f"file_path={self.file_path} not exists!") + + with open(self.file_path) as f: + prompt_dict: dict = yaml.load(f, yaml.FullLoader) + self.prompt_dict.update(prompt_dict) def __getitem__(self, key: str): return self.prompt_dict[key] diff --git a/experiencemaker/utils/registry.py b/experiencemaker/utils/registry.py index 65199cda..da9e8634 100644 --- a/experiencemaker/utils/registry.py +++ b/experiencemaker/utils/registry.py @@ -1,13 +1,21 @@ -from typing import Dict, Any, List +from typing import Dict, List + +from typing import TypeVar, Generic + +T = TypeVar('T') -class Registry(object): +class Registry(Generic[T]): def __init__(self, name: str): self.name: str = name - self.module_dict: Dict[str, Any] = {} + self.module_dict: Dict[str, T] = {} - def register(self, module, module_name: str = None): + @property + def registered_modules(self) -> List[str]: + return sorted(self.module_dict.keys()) + + def register(self, module: T, module_name: str = None): if module_name is None: module_name = module.__name__ @@ -16,7 +24,7 @@ class Registry(object): self.module_dict[module_name] = module - def batch_register(self, modules: List[Any] | Dict[str, Any]): + def batch_register(self, modules: List[T] | Dict[str, T]): if isinstance(modules, list): module_name_dict = {m.__name__: m for m in modules} @@ -27,6 +35,6 @@ class Registry(object): raise NotImplementedError("Input must be a list or a dictionary.") self.module_dict.update(module_name_dict) - def __getitem__(self, module_name: str): + def __getitem__(self, module_name: str) -> T: assert module_name in self.module_dict, f"{module_name} not found in {self.name}" return self.module_dict[module_name]