mirror of
https://github.com/agentscope-ai/ReMe.git
synced 2026-08-28 05:25:04 +00:00
Merge branch 'main' of http://gitlab.alibaba-inc.com/OpenRepo/ExperienceMaker
This commit is contained in:
commit
5ec394bccb
37 changed files with 680 additions and 237 deletions
3
.gitignore
vendored
3
.gitignore
vendored
|
|
@ -17,5 +17,4 @@ log/
|
|||
.trash/
|
||||
runs
|
||||
logs
|
||||
alfworld_data
|
||||
beyondagent/dataset/appworld/data
|
||||
|
||||
|
|
|
|||
148
doc/quick_start.md
Normal file
148
doc/quick_start.md
Normal file
|
|
@ -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}")
|
||||
```
|
||||
5
experiencemaker/config/__init__.py
Normal file
5
experiencemaker/config/__init__.py
Normal file
|
|
@ -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")
|
||||
1
experiencemaker/config/agent_wrapper/simple.json
Normal file
1
experiencemaker/config/agent_wrapper/simple.json
Normal file
|
|
@ -0,0 +1 @@
|
|||
{}
|
||||
32
experiencemaker/config/config_handler.py
Normal file
32
experiencemaker/config/config_handler.py
Normal file
|
|
@ -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)
|
||||
1
experiencemaker/config/context_generator/simple.json
Normal file
1
experiencemaker/config/context_generator/simple.json
Normal file
|
|
@ -0,0 +1 @@
|
|||
{}
|
||||
1
experiencemaker/config/summarizer/simple.json
Normal file
1
experiencemaker/config/summarizer/simple.json
Normal file
|
|
@ -0,0 +1 @@
|
|||
{"a": 1}
|
||||
|
|
@ -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")
|
||||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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(...)
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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<error>{chunk}</error>", end="")
|
||||
|
||||
|
||||
LLM_REGISTRY.register(OpenAICompatibleBaseLLM, "openai_compatible")
|
||||
|
|
|
|||
17
experiencemaker/module/agent_wrapper/agent_wrapper_mixin.py
Normal file
17
experiencemaker/module/agent_wrapper/agent_wrapper_mixin.py
Normal file
|
|
@ -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")
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -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)
|
||||
97
experiencemaker/module/agent_wrapper/simple_agent.py
Normal file
97
experiencemaker/module/agent_wrapper/simple_agent.py
Normal file
|
|
@ -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
|
||||
|
|
@ -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}
|
||||
|
||||
21
experiencemaker/module/agent_wrapper/simple_agent_wrapper.py
Normal file
21
experiencemaker/module/agent_wrapper/simple_agent_wrapper.py
Normal file
|
|
@ -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")
|
||||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
@ -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):
|
||||
|
|
|
|||
0
experiencemaker/module/prompt/__init__.py
Normal file
0
experiencemaker/module/prompt/__init__.py
Normal file
66
experiencemaker/module/prompt/prompt_mixin.py
Normal file
66
experiencemaker/module/prompt/prompt_mixin.py
Normal file
|
|
@ -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
|
||||
|
|
@ -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 []
|
||||
|
|
@ -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")
|
||||
|
|
|
|||
122
experiencemaker/service.py
Normal file
122
experiencemaker/service.py
Normal file
|
|
@ -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
|
||||
|
|
@ -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
|
||||
|
|
@ -0,0 +1,3 @@
|
|||
from experiencemaker.utils.registry import Registry
|
||||
|
||||
VECTOR_STORE_REGISTRY = Registry("vector_store")
|
||||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue