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
9cf061a7ca
56 changed files with 2185 additions and 491 deletions
3
.gitignore
vendored
3
.gitignore
vendored
|
|
@ -17,5 +17,4 @@ log/
|
|||
.trash/
|
||||
runs
|
||||
logs
|
||||
alfworld_data
|
||||
beyondagent/dataset/appworld/data
|
||||
rag_nodes_index.jsonl
|
||||
|
|
|
|||
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}")
|
||||
```
|
||||
|
|
@ -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,7 @@ from loguru import logger
|
|||
from pydantic import BaseModel, Field
|
||||
|
||||
from experiencemaker.schema.vector_store_node import VectorStoreNode
|
||||
from experiencemaker.utils.registry import Registry
|
||||
|
||||
|
||||
class BaseEmbeddingModel(BaseModel, ABC):
|
||||
|
|
@ -84,3 +85,6 @@ class BaseEmbeddingModel(BaseModel, ABC):
|
|||
|
||||
else:
|
||||
raise RuntimeError(f"unsupported type={type(nodes)}")
|
||||
|
||||
|
||||
EMBEDDING_MODEL_REGISTRY = Registry[BaseEmbeddingModel]("embedding_model")
|
||||
|
|
|
|||
|
|
@ -6,6 +6,7 @@ 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
|
||||
|
||||
|
||||
class BaseLLM(BaseModel, ABC):
|
||||
|
|
@ -108,3 +109,6 @@ class BaseLLM(BaseModel, ABC):
|
|||
raise e
|
||||
|
||||
return None
|
||||
|
||||
|
||||
LLM_REGISTRY = Registry[BaseLLM]("llm")
|
||||
|
|
|
|||
|
|
@ -4,13 +4,13 @@ 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):
|
||||
api_key: str = Field(default_factory=lambda: os.getenv("OPENAI_API_KEY"), description="api key")
|
||||
base_url: str = Field(default_factory=lambda: os.getenv("OPENAI_BASE_URL"), description="base url")
|
||||
model_name: str = Field(default="text-embedding-v3", description="model name")
|
||||
model_name: str = Field(default="text-embedding-v4", description="model name")
|
||||
dimensions: int = Field(default=1024, description="dimensions")
|
||||
encoding_format: Literal["float", "base64"] = Field(default="float", description="encoding_format")
|
||||
_client: OpenAI = PrivateAttr()
|
||||
|
|
@ -71,3 +71,23 @@ 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")
|
||||
|
||||
|
||||
def main():
|
||||
from experiencemaker.utils.util_function import load_env_keys
|
||||
load_env_keys()
|
||||
|
||||
model = OpenAICompatibleEmbeddingModel(dimensions=64, model_name="text-embedding-v4")
|
||||
res1 = model.get_embeddings(
|
||||
"The clothes are of good quality and look good, definitely worth the wait. I love them.")
|
||||
res2 = model.get_embeddings(["aa", "bb"])
|
||||
print(res1)
|
||||
print(res2)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
# launch with: python -m experiencemaker.model.openai_compatible_embedding_model
|
||||
|
|
|
|||
|
|
@ -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,28 @@ class OpenAICompatibleBaseLLM(BaseLLM):
|
|||
|
||||
elif chunk_enum is ChunkEnum.ERROR:
|
||||
print(f"\n<error>{chunk}</error>", end="")
|
||||
|
||||
|
||||
LLM_REGISTRY.register(OpenAICompatibleBaseLLM, "openai_compatible")
|
||||
|
||||
|
||||
def main():
|
||||
from experiencemaker.utils.util_function import load_env_keys
|
||||
from experiencemaker.tool.dashscope_search_tool import DashscopeSearchTool
|
||||
from experiencemaker.tool.code_tool import CodeTool
|
||||
from experiencemaker.enumeration.role import Role
|
||||
|
||||
load_env_keys()
|
||||
model_name = "qwen-max-2025-01-25"
|
||||
# model_name = "qwen3-32b"
|
||||
llm = OpenAICompatibleBaseLLM(model_name=model_name)
|
||||
tools: List[BaseTool] = [DashscopeSearchTool(), CodeTool()]
|
||||
|
||||
llm.stream_print([Message(role=Role.USER, content="hello")], [])
|
||||
print("=" * 20)
|
||||
llm.stream_print([Message(role=Role.USER, content="What's the weather like in Beijing today?")], tools)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
# launch with: python -m experiencemaker.model.openai_compatible_llm
|
||||
|
|
|
|||
18
experiencemaker/module/agent_wrapper/agent_wrapper_mixin.py
Normal file
18
experiencemaker/module/agent_wrapper/agent_wrapper_mixin.py
Normal file
|
|
@ -0,0 +1,18 @@
|
|||
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)
|
||||
workspace_id: str = Field(default="")
|
||||
|
||||
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)
|
||||
98
experiencemaker/module/agent_wrapper/simple_agent.py
Normal file
98
experiencemaker/module/agent_wrapper/simple_agent.py
Normal file
|
|
@ -0,0 +1,98 @@
|
|||
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 import CodeTool, DashscopeSearchTool, TerminateTool
|
||||
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] = [CodeTool(), DashscopeSearchTool(), TerminateTool()]
|
||||
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")
|
||||
|
|
@ -1,48 +0,0 @@
|
|||
from abc import ABC
|
||||
|
||||
from loguru import logger
|
||||
from pydantic import BaseModel, Field, model_validator
|
||||
|
||||
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.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)
|
||||
|
||||
llm: BaseLLM | None = Field(default=None)
|
||||
embedding_model: BaseEmbeddingModel | None = Field(default=None)
|
||||
|
||||
@model_validator(mode="before") # noqa
|
||||
@classmethod
|
||||
def init_model(cls, data: dict):
|
||||
if "llm" in data and isinstance(data["llm"], dict):
|
||||
backend = data["llm"].pop("backend", None)
|
||||
assert backend is not None, "llm must have a backend"
|
||||
module = LLM_REGISTRY[backend]
|
||||
params = data["llm"]
|
||||
data["llm"] = module(**params)
|
||||
logger.info(f"{cls.__name__} load llm.backend={backend} params={params}")
|
||||
|
||||
if "embedding_model" in data and isinstance(data["embedding_model"], dict):
|
||||
backend = data["embedding_model"].pop("backend", None)
|
||||
assert backend is not None, "embedding_model must have a backend"
|
||||
module = EMBEDDING_MODEL_REGISTRY[backend]
|
||||
params = data["embedding_model"]
|
||||
data["embedding_model"] = module(**params)
|
||||
logger.info(f"{cls.__name__} load embedding_model.backend={backend} params={params}")
|
||||
|
||||
if "prompt_dir" in data:
|
||||
handler = PromptHandler(dir_path=data.get("prompt_dir"))
|
||||
data["prompt_handler"] = handler
|
||||
if data.get("prompt_file"):
|
||||
handler.add_prompt_file(data.get("prompt_file"))
|
||||
|
||||
return data
|
||||
|
||||
def execute(self, **kwargs):
|
||||
raise NotImplementedError
|
||||
|
|
@ -1,16 +1,20 @@
|
|||
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.model.base_embedding_model import BaseEmbeddingModel
|
||||
from experiencemaker.model.base_llm import BaseLLM
|
||||
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)
|
||||
llm: BaseLLM | None = Field(default=None)
|
||||
workspace_id: str = Field(default="")
|
||||
|
||||
def _build_retrieve_query(self, trajectory: Trajectory, **kwargs) -> str:
|
||||
raise NotImplementedError
|
||||
|
|
@ -31,7 +35,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")
|
||||
|
|
@ -0,0 +1,393 @@
|
|||
import json
|
||||
import re
|
||||
from typing import List, Dict, Any, Optional
|
||||
from loguru import logger
|
||||
from pydantic import Field, model_validator
|
||||
|
||||
from experiencemaker.enumeration.role import Role
|
||||
from experiencemaker.module.context_generator.base_context_generator import BaseContextGenerator
|
||||
from experiencemaker.schema.trajectory import Trajectory, ContextMessage, Message
|
||||
from experiencemaker.schema.vector_store_node import VectorStoreNode
|
||||
from experiencemaker.storage.es_vector_store import EsVectorStore
|
||||
from experiencemaker.storage.file_vector_store import FileVectorStore
|
||||
|
||||
|
||||
class StepContextGenerator(BaseContextGenerator):
|
||||
"""
|
||||
Step-level context generator that retrieves and utilizes step-level experiences
|
||||
from the experience store to provide relevant context for agent execution
|
||||
"""
|
||||
|
||||
# Vector Store Configuration
|
||||
vector_store_type: str = Field(default="file_vector_store")
|
||||
vector_store_hosts: str | List[str] = Field(default="http://localhost:9200")
|
||||
vector_store_index_name: str = Field(default="step_experience_store")
|
||||
store_dir: str = Field(default="./step_experiences/")
|
||||
|
||||
# Retrieval Configuration
|
||||
vector_retrieve_top_k: int = Field(default=15)
|
||||
final_top_k: int = Field(default=5)
|
||||
min_score_threshold: float = Field(default=0.3)
|
||||
|
||||
# Feature Switches
|
||||
enable_llm_rerank: bool = Field(default=True)
|
||||
enable_context_rewrite: bool = Field(default=True)
|
||||
enable_score_filter: bool = Field(default=True)
|
||||
|
||||
@model_validator(mode="after")
|
||||
def init_vector_store(self):
|
||||
"""Initialize vector store based on configuration"""
|
||||
if self.vector_store_type == "file_vector_store":
|
||||
self.vector_store = FileVectorStore(
|
||||
embedding_model=self.embedding_model,
|
||||
index_name=self.vector_store_index_name,
|
||||
store_dir=self.store_dir
|
||||
)
|
||||
elif self.vector_store_type == "es_vector_store":
|
||||
self.vector_store = EsVectorStore(
|
||||
embedding_model=self.embedding_model,
|
||||
index_name=self.vector_store_index_name,
|
||||
hosts=self.vector_store_hosts
|
||||
)
|
||||
else:
|
||||
raise ValueError(f"Unknown vector store type: {self.vector_store_type}")
|
||||
|
||||
return self
|
||||
|
||||
def _build_retrieve_query(self, trajectory: Trajectory, **kwargs) -> str:
|
||||
"""Build retrieval query from trajectory"""
|
||||
# Use the original query as base
|
||||
base_query = trajectory.query
|
||||
|
||||
# Optionally enhance with current step context if available
|
||||
current_context = kwargs.get("current_context", "")
|
||||
if current_context:
|
||||
base_query = f"{base_query} {current_context}"
|
||||
|
||||
return base_query
|
||||
|
||||
def vector_retrieve(self, query: str, top_k: int = 10) -> List[VectorStoreNode]:
|
||||
"""Vector similarity retrieval from experience store"""
|
||||
if not query:
|
||||
logger.warning("Empty query provided for vector retrieval")
|
||||
return []
|
||||
|
||||
try:
|
||||
retrieved_nodes = self.vector_store.retrieve_by_query(
|
||||
query=query,
|
||||
top_k=top_k
|
||||
)
|
||||
logger.info(f"Vector retrieval found {len(retrieved_nodes)} candidates")
|
||||
return retrieved_nodes
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Error in vector retrieval: {e}")
|
||||
return []
|
||||
|
||||
def llm_rerank(self, query: str, candidates: List[VectorStoreNode]) -> List[VectorStoreNode]:
|
||||
"""LLM-based reranking of candidate experiences"""
|
||||
if not self.enable_llm_rerank or not candidates:
|
||||
return candidates
|
||||
|
||||
try:
|
||||
# Format candidates for LLM evaluation
|
||||
candidates_text = self._format_candidates_for_rerank(candidates)
|
||||
|
||||
prompt = self.prompt_handler.experience_rerank_prompt.format(
|
||||
query=query,
|
||||
candidates=candidates_text,
|
||||
num_candidates=len(candidates)
|
||||
)
|
||||
|
||||
response = self.llm.chat([Message(role=Role.USER, content=prompt)])
|
||||
|
||||
# Parse reranking results
|
||||
reranked_indices = self._parse_rerank_response(response.content)
|
||||
|
||||
# Reorder candidates based on LLM ranking
|
||||
if reranked_indices:
|
||||
reranked_candidates = []
|
||||
for idx in reranked_indices:
|
||||
if 0 <= idx < len(candidates):
|
||||
reranked_candidates.append(candidates[idx])
|
||||
return reranked_candidates
|
||||
|
||||
return candidates
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Error in LLM reranking: {e}")
|
||||
return candidates
|
||||
|
||||
def llm_rewrite_context(self, query: str, context_content: str, trajectory: Trajectory) -> str:
|
||||
"""LLM-based context rewriting to make experiences more relevant and actionable for current task"""
|
||||
if not self.enable_query_rewrite or not context_content:
|
||||
return context_content
|
||||
|
||||
try:
|
||||
# Extract current trajectory context
|
||||
current_context = self._extract_trajectory_context(trajectory)
|
||||
|
||||
prompt = self.prompt_handler.context_rewrite_prompt.format(
|
||||
current_query=query,
|
||||
current_context=current_context,
|
||||
original_context=context_content
|
||||
)
|
||||
|
||||
response = self.llm.chat([Message(role=Role.USER, content=prompt)])
|
||||
|
||||
# Extract rewritten context from JSON
|
||||
rewritten_context = self._parse_json_response(response.content, "rewritten_context")
|
||||
|
||||
if rewritten_context and rewritten_context.strip():
|
||||
logger.info("Context successfully rewritten for current task")
|
||||
return rewritten_context.strip()
|
||||
|
||||
return context_content
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Error in context rewriting: {e}")
|
||||
return context_content
|
||||
|
||||
def score_based_filter(self, experiences: List[VectorStoreNode],
|
||||
min_score: float) -> List[VectorStoreNode]:
|
||||
"""Filter experiences based on quality scores"""
|
||||
if not self.enable_score_filter:
|
||||
return experiences
|
||||
|
||||
filtered_experiences = []
|
||||
|
||||
for exp in experiences:
|
||||
# Get confidence score from metadata
|
||||
confidence = exp.metadata.get("confidence", 0.5)
|
||||
validation_score = exp.metadata.get("validation_score", 0.5)
|
||||
|
||||
# Calculate combined score
|
||||
combined_score = (confidence + validation_score) / 2
|
||||
|
||||
if combined_score >= min_score:
|
||||
filtered_experiences.append(exp)
|
||||
else:
|
||||
logger.debug(f"Filtered out experience with score {combined_score:.2f}")
|
||||
|
||||
logger.info(f"Score filtering: {len(filtered_experiences)}/{len(experiences)} experiences retained")
|
||||
return filtered_experiences
|
||||
|
||||
def hybrid_retrieve(self, query: str, trajectory: Trajectory, top_k: int = 5) -> List[VectorStoreNode]:
|
||||
"""Hybrid retrieval strategy combining multiple approaches"""
|
||||
logger.info(f"Starting hybrid retrieval for query: '{query}'")
|
||||
|
||||
# Step 1: Vector retrieval to get candidates
|
||||
candidates = self.vector_retrieve(query, self.vector_retrieve_top_k)
|
||||
|
||||
if not candidates:
|
||||
logger.warning("No candidates found in vector retrieval")
|
||||
return []
|
||||
|
||||
# Step 2: LLM reranking (optional)
|
||||
reranked = self.llm_rerank(query, candidates)
|
||||
|
||||
# Step 3: Score-based filtering (optional)
|
||||
filtered = self.score_based_filter(reranked, self.min_score_threshold)
|
||||
|
||||
# Step 4: Return top-k results
|
||||
final_results = filtered[:top_k]
|
||||
logger.info(f"Hybrid retrieval completed: {len(final_results)} experiences selected")
|
||||
|
||||
return final_results
|
||||
|
||||
def retrieve_by_query(self, trajectory: Trajectory, query: str, **kwargs) -> List[VectorStoreNode]:
|
||||
"""Retrieve experiences by query (implements base class method)"""
|
||||
return self.hybrid_retrieve(query, trajectory, self.final_top_k)
|
||||
|
||||
def generate_context_message(self,
|
||||
trajectory: Trajectory,
|
||||
nodes: List[VectorStoreNode],
|
||||
**kwargs) -> ContextMessage:
|
||||
"""Generate context message from retrieved experiences"""
|
||||
if not nodes:
|
||||
return ContextMessage(content="")
|
||||
|
||||
try:
|
||||
# Format retrieved experiences
|
||||
formatted_experiences = self._format_experiences_for_context(nodes)
|
||||
|
||||
prompt = self.prompt_handler.context_generation_prompt.format(
|
||||
query=trajectory.query,
|
||||
current_step=kwargs.get("current_step", ""),
|
||||
retrieved_experiences=formatted_experiences,
|
||||
num_experiences=len(nodes)
|
||||
)
|
||||
|
||||
response = self.llm.chat([Message(role=Role.USER, content=prompt)])
|
||||
|
||||
# Extract generated context from JSON
|
||||
context_content = self._parse_json_response(response.content, "context")
|
||||
|
||||
if not context_content:
|
||||
# Fallback to simple formatting
|
||||
context_content = self._create_context(nodes)
|
||||
|
||||
return ContextMessage(content=context_content)
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Error generating context message: {e}")
|
||||
return ContextMessage(content=self._create_context(nodes))
|
||||
|
||||
def build_context_messages(self, task: str, experiences: List[VectorStoreNode], trajectory: Trajectory) -> List[
|
||||
Message]:
|
||||
"""Build context messages from experiences for agent consumption"""
|
||||
if not experiences:
|
||||
return []
|
||||
|
||||
messages = []
|
||||
|
||||
# Create initial context content with experiences
|
||||
system_content = "You have access to the following relevant experiences from previous executions:\n\n"
|
||||
|
||||
for i, exp in enumerate(experiences, 1):
|
||||
condition = exp.content
|
||||
experience_content = exp.metadata.get("experience", "")
|
||||
tags = exp.metadata.get("tags", [])
|
||||
|
||||
system_content += f"**Experience {i}:**\n"
|
||||
system_content += f"When to use: {condition}\n"
|
||||
system_content += f"Experience: {experience_content}\n"
|
||||
system_content += f"Tags: {', '.join(tags)}\n\n"
|
||||
|
||||
system_content += "Consider these experiences when planning and executing your approach."
|
||||
|
||||
# Rewrite the complete context to make it more relevant to current task
|
||||
if self.enable_context_rewrite:
|
||||
system_content = self.llm_rewrite_context(task, system_content, trajectory)
|
||||
|
||||
messages.append(Message(role=Role.SYSTEM, content=system_content))
|
||||
|
||||
return messages
|
||||
|
||||
def get_best_experiences(self, task: str, trajectory: Trajectory, max_count: int = 3) -> List[Message]:
|
||||
"""Get the best relevant experiences for a task as formatted messages"""
|
||||
experiences = self.hybrid_retrieve(task, trajectory, max_count)
|
||||
return self.build_context_messages(task, experiences, trajectory)
|
||||
|
||||
def _extract_trajectory_context(self, trajectory: Trajectory) -> str:
|
||||
"""Extract relevant context from trajectory for query enhancement"""
|
||||
context_parts = []
|
||||
|
||||
# Add recent steps if available
|
||||
if trajectory.steps:
|
||||
recent_steps = trajectory.steps[-3:] # Last 3 steps
|
||||
step_summaries = []
|
||||
for step in recent_steps:
|
||||
step_summary = step.content[:100] + "..." if len(step.content) > 100 else step.content
|
||||
step_summaries.append(f"- {step.role.value}: {step_summary}")
|
||||
|
||||
if step_summaries:
|
||||
context_parts.append("Recent steps:\n" + "\n".join(step_summaries))
|
||||
|
||||
# Add metadata if available
|
||||
if trajectory.metadata:
|
||||
relevant_metadata = {k: v for k, v in trajectory.metadata.items()
|
||||
if k in ["domain", "task_type", "difficulty"]}
|
||||
if relevant_metadata:
|
||||
context_parts.append(f"Task metadata: {relevant_metadata}")
|
||||
|
||||
return "\n\n".join(context_parts)
|
||||
|
||||
def _format_candidates_for_rerank(self, candidates: List[VectorStoreNode]) -> str:
|
||||
"""Format candidates for LLM reranking"""
|
||||
formatted_candidates = []
|
||||
|
||||
for i, candidate in enumerate(candidates):
|
||||
condition = candidate.content
|
||||
experience = candidate.metadata.get("experience", "")
|
||||
tags = candidate.metadata.get("tags", [])
|
||||
confidence = candidate.metadata.get("confidence", 0.5)
|
||||
|
||||
candidate_text = f"Candidate {i}:\n"
|
||||
candidate_text += f"Condition: {condition}\n"
|
||||
candidate_text += f"Experience: {experience}\n"
|
||||
candidate_text += f"Tags: {', '.join(tags)}\n"
|
||||
candidate_text += f"Confidence: {confidence}\n"
|
||||
|
||||
formatted_candidates.append(candidate_text)
|
||||
|
||||
return "\n---\n".join(formatted_candidates)
|
||||
|
||||
def _parse_rerank_response(self, response: str) -> List[int]:
|
||||
"""Parse LLM reranking response to extract ranked indices"""
|
||||
try:
|
||||
# Try to extract JSON format
|
||||
json_pattern = r'```json\s*([\s\S]*?)\s*```'
|
||||
json_blocks = re.findall(json_pattern, response)
|
||||
|
||||
if json_blocks:
|
||||
parsed = json.loads(json_blocks[0])
|
||||
if isinstance(parsed, dict) and "ranked_indices" in parsed:
|
||||
return parsed["ranked_indices"]
|
||||
elif isinstance(parsed, list):
|
||||
return parsed
|
||||
|
||||
# Try to extract numbers from text
|
||||
numbers = re.findall(r'\b\d+\b', response)
|
||||
return [int(num) for num in numbers]
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Error parsing rerank response: {e}")
|
||||
return []
|
||||
|
||||
def _format_experiences_for_context(self, experiences: List[VectorStoreNode]) -> str:
|
||||
"""Format experiences for context generation"""
|
||||
formatted_experiences = []
|
||||
|
||||
for i, exp in enumerate(experiences, 1):
|
||||
condition = exp.content
|
||||
experience_content = exp.metadata.get("experience", "")
|
||||
experience_type = exp.metadata.get("experience_type", "general")
|
||||
tags = exp.metadata.get("tags", [])
|
||||
|
||||
exp_text = f"Experience {i} ({experience_type}):\n"
|
||||
exp_text += f"When to use: {condition}\n"
|
||||
exp_text += f"Experience: {experience_content}\n"
|
||||
exp_text += f"Tags: {', '.join(tags)}"
|
||||
|
||||
formatted_experiences.append(exp_text)
|
||||
|
||||
return "\n\n---\n\n".join(formatted_experiences)
|
||||
|
||||
def _create_context(self, experiences: List[VectorStoreNode]) -> str:
|
||||
"""Create simple context when LLM generation fails"""
|
||||
if not experiences:
|
||||
return ""
|
||||
|
||||
context = "Here are some relevant experiences that might help:\n\n"
|
||||
|
||||
for i, exp in enumerate(experiences, 1):
|
||||
condition = exp.content
|
||||
experience_content = exp.metadata.get("experience", "")
|
||||
|
||||
context += f"{i}. **When**: {condition}\n"
|
||||
context += f" **Experience**: {experience_content}\n\n"
|
||||
|
||||
return context
|
||||
|
||||
def _parse_json_response(self, response: str, key: str) -> str:
|
||||
"""Parse JSON response to extract specific key"""
|
||||
try:
|
||||
# Try to extract JSON blocks
|
||||
json_pattern = r'```json\s*([\s\S]*?)\s*```'
|
||||
json_blocks = re.findall(json_pattern, response)
|
||||
|
||||
if json_blocks:
|
||||
parsed = json.loads(json_blocks[0])
|
||||
if isinstance(parsed, dict) and key in parsed:
|
||||
return parsed[key]
|
||||
|
||||
# Fallback: try to parse the entire response as JSON
|
||||
parsed = json.loads(response)
|
||||
if isinstance(parsed, dict) and key in parsed:
|
||||
return parsed[key]
|
||||
|
||||
except json.JSONDecodeError:
|
||||
logger.warning(f"Failed to parse JSON response for key '{key}'")
|
||||
|
||||
return ""
|
||||
|
|
@ -1,18 +1,18 @@
|
|||
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.module.reward_fn.base_reward_fn import BaseRewardFn
|
||||
from experiencemaker.schema.reward import Reward
|
||||
from experiencemaker.schema.trajectory import StateMessage, ActionMessage, ToolCall
|
||||
from experiencemaker.tool.base_tool import BaseTool
|
||||
|
||||
|
||||
class BaseEnvironment(BaseModule):
|
||||
class BaseEnvironment(BaseModel, ABC):
|
||||
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):
|
||||
|
|
@ -54,15 +54,3 @@ class BaseEnvironment(BaseModule):
|
|||
|
||||
def build_info(self, **kwargs):
|
||||
return {}
|
||||
|
||||
def get_tool_info(self,tool_name):
|
||||
tool_dict = {tool.name: tool for tool in self.tools}
|
||||
if tool_name in tool_dict:
|
||||
|
||||
return f'tool \'{tool_name}\' description is: {tool_dict[tool_name].description}\t' + f'parameters: {str(tool_dict[tool_name].input_schema)}'
|
||||
else:
|
||||
return ''
|
||||
|
||||
def get_tools_info(self):
|
||||
tool_dict = {tool.name: tool for tool in self.tools}
|
||||
return {tool_name:self.get_tool_info(tool_name=tool_name) for tool_name in tool_dict}
|
||||
|
|
@ -2,7 +2,7 @@ from abc import ABC
|
|||
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
from experiencemaker.module.agent_wrapper.base_agent_wrapper import BaseAgentWrapper
|
||||
from experiencemaker.module.agent_wrapper.agent_wrapper_mixin import AgentWrapperMixin
|
||||
from experiencemaker.module.context_generator.base_context_generator import BaseContextGenerator
|
||||
from experiencemaker.module.environment.base_environment import BaseEnvironment
|
||||
from experiencemaker.module.summarizer.base_summarizer import BaseSummarizer
|
||||
|
|
@ -10,7 +10,7 @@ from experiencemaker.module.summarizer.base_summarizer import BaseSummarizer
|
|||
|
||||
class BaseEvaluator(BaseModel, ABC):
|
||||
data_path: str = Field(default="")
|
||||
agent_wrapper: BaseAgentWrapper | None = Field(default=None)
|
||||
agent_wrapper: AgentWrapperMixin | None = Field(default=None)
|
||||
context_generator: BaseContextGenerator | None = Field(default=None)
|
||||
summarizer: BaseSummarizer | None = Field(default=None)
|
||||
env: BaseEnvironment | None = Field(default=None)
|
||||
|
|
|
|||
0
experiencemaker/module/prompt/__init__.py
Normal file
0
experiencemaker/module/prompt/__init__.py
Normal file
|
|
@ -1,40 +1,36 @@
|
|||
import os
|
||||
from pathlib import Path
|
||||
|
||||
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="")
|
||||
class PromptMixin(BaseModel):
|
||||
prompt_file_path: Path | str = Field(default=None)
|
||||
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(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.prompt_file_path is not None:
|
||||
if isinstance(self.prompt_file_path, str):
|
||||
self.prompt_file_path = Path(self.prompt_file_path)
|
||||
|
||||
def update_prompt_dict(self, prompt_dict: dict):
|
||||
self.prompt_dict.update(prompt_dict)
|
||||
if not self.prompt_file_path.exists():
|
||||
logger.warning(f"prompt_file_path={self.prompt_file_path} not exists!")
|
||||
|
||||
def __getitem__(self, key: str):
|
||||
return self.prompt_dict[key]
|
||||
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!")
|
||||
|
||||
def __setitem__(self, key: str, value: str):
|
||||
self.prompt_dict[key] = value
|
||||
|
||||
def __getattr__(self, key: str):
|
||||
if key in self.prompt_dict:
|
||||
return self.prompt_dict[key]
|
||||
|
||||
return super().__getattr__(key)
|
||||
return self
|
||||
|
||||
def prompt_format(self, prompt_name: str, **kwargs):
|
||||
prompt = self.prompt_dict[prompt_name]
|
||||
|
|
@ -1,11 +1,12 @@
|
|||
from abc import ABC
|
||||
|
||||
from experiencemaker.module.base_module import BaseModule
|
||||
from pydantic import BaseModel
|
||||
|
||||
from experiencemaker.schema.reward import Reward
|
||||
from experiencemaker.schema.trajectory import Trajectory
|
||||
|
||||
|
||||
class BaseRewardFn(BaseModule, ABC):
|
||||
class BaseRewardFn(BaseModel, ABC):
|
||||
|
||||
def execute(self, trajectory: Trajectory, ground_truth=None, **kwargs) -> Reward:
|
||||
def execute(self, trajectory: Trajectory = None, ground_truth=None, **kwargs) -> Reward:
|
||||
raise NotImplementedError
|
||||
|
|
|
|||
61
experiencemaker/module/reward_fn/simple_compare_reward_fn.py
Normal file
61
experiencemaker/module/reward_fn/simple_compare_reward_fn.py
Normal file
|
|
@ -0,0 +1,61 @@
|
|||
import datetime
|
||||
from pathlib import Path
|
||||
|
||||
from loguru import logger
|
||||
from pydantic import Field
|
||||
|
||||
from experiencemaker.model.base_llm import BaseLLM
|
||||
from experiencemaker.module.prompt.prompt_mixin import PromptMixin
|
||||
from experiencemaker.module.reward_fn.base_reward_fn import BaseRewardFn
|
||||
from experiencemaker.schema.reward import Reward
|
||||
from experiencemaker.schema.trajectory import Trajectory, Message
|
||||
from experiencemaker.utils.util_function import get_html_match_content
|
||||
|
||||
|
||||
class SimpleCompareRewardFn(BaseRewardFn, PromptMixin):
|
||||
llm: BaseLLM | None = Field(default=None)
|
||||
eval_times: int = Field(default=5)
|
||||
prompt_file_path: Path = Path(__file__).parent / "simple_compare_reward_fn_prompt.yaml"
|
||||
|
||||
def execute(self, trajectory: Trajectory = None, comp_traj: Trajectory = None, **kwargs) -> Reward:
|
||||
query = trajectory.query
|
||||
answer1 = trajectory.answer
|
||||
answer2 = comp_traj.answer
|
||||
logger.info("=" * 10 + f"answer1\n{answer1}\n" + "=" * 10 + f"answer2\n{answer2}\n")
|
||||
|
||||
valid_cnt = 0
|
||||
better_cnt = 0
|
||||
for i in range(self.eval_times):
|
||||
now_time = datetime.datetime.now().strftime('%Y-%m-%d %H:%M:%S')
|
||||
user_prompt = self.prompt_handler.compare_prompt.format(
|
||||
now_time=now_time,
|
||||
query=query,
|
||||
answer1=answer1,
|
||||
answer2=answer2)
|
||||
|
||||
messages = [Message(content=user_prompt)]
|
||||
action_msg = self.llm.chat(messages=messages)
|
||||
|
||||
rule: str = get_html_match_content(action_msg.content, "rule")
|
||||
rule_based_comparison: str = get_html_match_content(action_msg.content, "rule_based_comparison")
|
||||
result: str | None = get_html_match_content(action_msg.content, "result")
|
||||
logger.info(f"round.{i} rule={rule} rule_based_comparison={rule_based_comparison} result={result}")
|
||||
if result:
|
||||
result = result.lower()
|
||||
if "plan1" in result and "plan2" in result:
|
||||
logger.warning(f"both plan exists in result={result}")
|
||||
|
||||
elif "plan1" in result:
|
||||
valid_cnt += 1
|
||||
better_cnt += 1
|
||||
|
||||
elif "plan2" in result:
|
||||
valid_cnt += 1
|
||||
|
||||
else:
|
||||
logger.warning(f"no plan exists in result={result}")
|
||||
|
||||
reward_value = 0
|
||||
if valid_cnt > 0:
|
||||
reward_value = better_cnt / valid_cnt
|
||||
return Reward(reward_value=reward_value)
|
||||
|
|
@ -0,0 +1,30 @@
|
|||
compare_prompt: |
|
||||
# Role
|
||||
You are a helpful assistant named BeyondAgent.
|
||||
current time: {now_time}
|
||||
|
||||
# User Question
|
||||
{query}
|
||||
|
||||
# Plan1 Answer
|
||||
{answer1}
|
||||
|
||||
# Plan2 Answer
|
||||
{answer2}
|
||||
|
||||
# Task
|
||||
Based on the **User Question**, determine which plan provides a better answer. Response steps:
|
||||
1. Consider which comparison rules apply, the rules for comparison can be macro-level dimensions or micro-level details.
|
||||
2. Conduct a step-by-step comparison according to the rules.
|
||||
3. Combine all results to arrive at a final answer.
|
||||
|
||||
# Output Example
|
||||
<rule>
|
||||
List the rules that could be used for comparison...
|
||||
</rule>
|
||||
<rule_based_comparison>
|
||||
Conduct a step-by-step comparison according to the rules...
|
||||
</rule_based_comparison>
|
||||
<result>
|
||||
Output only the name of the better plan, either **Plan1** or **Plan2**.
|
||||
</result>
|
||||
|
|
@ -2,7 +2,7 @@ from typing import List
|
|||
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
from experiencemaker.module.agent_wrapper.base_agent_wrapper import BaseAgentWrapper
|
||||
from experiencemaker.module.agent_wrapper.agent_wrapper_mixin import AgentWrapperMixin
|
||||
from experiencemaker.module.context_generator.base_context_generator import BaseContextGenerator
|
||||
from experiencemaker.module.environment.base_environment import BaseEnvironment
|
||||
from experiencemaker.module.summarizer.base_summarizer import BaseSummarizer
|
||||
|
|
@ -10,7 +10,7 @@ from experiencemaker.schema.trajectory import Trajectory
|
|||
|
||||
|
||||
class BaseRunner(BaseModel):
|
||||
agent_wrapper: BaseAgentWrapper | None = Field(default=None)
|
||||
agent_wrapper: AgentWrapperMixin | None = Field(default=None)
|
||||
context_generator: BaseContextGenerator | None = Field(default=None)
|
||||
summarizer: BaseSummarizer | None = Field(default=None)
|
||||
env: BaseEnvironment | None = Field(default=None)
|
||||
|
|
@ -18,12 +18,10 @@ class BaseRunner(BaseModel):
|
|||
|
||||
def reset(self):
|
||||
self.traj_buffer.clear()
|
||||
self.env.reset()
|
||||
|
||||
def rollout_trajectory(self, user_query: str, **kwargs):
|
||||
def rollout_trajectory(self, query: str, **kwargs):
|
||||
raise NotImplementedError
|
||||
|
||||
def summary(self):
|
||||
raise NotImplementedError
|
||||
|
||||
def start_backend_summary(self):
|
||||
raise NotImplementedError
|
||||
def summary(self, **kwargs):
|
||||
raise NotImplementedError
|
||||
|
|
@ -1,37 +1,33 @@
|
|||
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.model.base_llm import BaseLLM
|
||||
from experiencemaker.schema.experience import Experience
|
||||
from experiencemaker.schema.trajectory import Trajectory
|
||||
from experiencemaker.schema.vector_store_node import VectorStoreNode
|
||||
from experiencemaker.storage.base_vector_store import BaseVectorStore
|
||||
from experiencemaker.utils.registry import Registry
|
||||
|
||||
|
||||
class BaseSummarizer(BaseModule):
|
||||
class BaseSummarizer(BaseModel, ABC):
|
||||
vector_store: BaseVectorStore | None = Field(default=None)
|
||||
llm: BaseLLM | None = Field(default=None)
|
||||
workspace_id: str = Field(default="")
|
||||
|
||||
def extract_samples(self, trajectories: List[Trajectory], **kwargs) -> List[Sample]:
|
||||
def _extract_experiences(self, trajectories: List[Trajectory], **kwargs) -> List[Experience]:
|
||||
raise NotImplementedError
|
||||
|
||||
def insert_into_vector_store(self, samples: List[Sample], **kwargs):
|
||||
raise NotImplementedError
|
||||
def execute(self, trajectories: List[Trajectory], return_experience: bool = True, **kwargs) -> List[Experience]:
|
||||
experiences: List[Experience] = self._extract_experiences(trajectories, **kwargs)
|
||||
|
||||
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)
|
||||
|
||||
if return_samples:
|
||||
return samples
|
||||
nodes: List[VectorStoreNode] = [x.to_vector_store_node() for x in experiences]
|
||||
self.vector_store.insert(nodes, **kwargs)
|
||||
|
||||
if return_experience:
|
||||
return experiences
|
||||
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 []
|
||||
SUMMARIZER_REGISTRY = Registry[BaseSummarizer]("summarizer")
|
||||
|
|
|
|||
68
experiencemaker/module/summarizer/simple_summarizer.py
Normal file
68
experiencemaker/module/summarizer/simple_summarizer.py
Normal file
|
|
@ -0,0 +1,68 @@
|
|||
from pathlib import Path
|
||||
from typing import List
|
||||
|
||||
from loguru import logger
|
||||
from pydantic import Field
|
||||
|
||||
from experiencemaker.enumeration.role import Role
|
||||
from experiencemaker.module.prompt.prompt_mixin import PromptMixin
|
||||
from experiencemaker.module.summarizer.base_summarizer import BaseSummarizer, SUMMARIZER_REGISTRY
|
||||
from experiencemaker.schema.experience import Experience
|
||||
from experiencemaker.schema.trajectory import Trajectory, Message, ActionMessage
|
||||
from experiencemaker.utils.util_function import get_html_match_content
|
||||
|
||||
|
||||
class SimpleSummarizer(BaseSummarizer, PromptMixin):
|
||||
max_retries: int = Field(default=5, description="max retries")
|
||||
prompt_file_path: Path = Path(__file__).parent / "simple_summarizer_prompt.yaml"
|
||||
|
||||
def _extract_trajectory_experience(self, trajectory: Trajectory) -> Experience | None:
|
||||
step_content_collector: List[str] = []
|
||||
|
||||
for step in trajectory.steps:
|
||||
step_index = len(step_content_collector)
|
||||
|
||||
if step.role is Role.ASSISTANT:
|
||||
line = f"### step.{step_index} role={step.role.value} content=\n{step.content}\n{step.reasoning_content}\n"
|
||||
if step.tool_calls:
|
||||
for tool_call in step.tool_calls:
|
||||
line += f" - tool call={tool_call.name}\n params={tool_call.arguments}\n"
|
||||
step_content_collector.append(line)
|
||||
|
||||
elif step.role is Role.USER:
|
||||
line = f"### step.{step_index} role={step.role.value} content=\n{step.content}\n"
|
||||
step_content_collector.append(line)
|
||||
|
||||
elif step.role is Role.TOOL:
|
||||
line = f"### step.{step_index} role={step.role.value} tool call result=\n{step.content}\n"
|
||||
step_content_collector.append(line)
|
||||
|
||||
prompt = self.prompt_format(prompt_name="summary_prompt",
|
||||
query=trajectory.query,
|
||||
execution_process="\n".join(step_content_collector).strip(),
|
||||
answer=trajectory.answer)
|
||||
|
||||
for i in range(self.max_retries):
|
||||
action_message: ActionMessage = self.llm.chat(messages=[Message(content=prompt)])
|
||||
experience_str = get_html_match_content(action_message.content, key="experience")
|
||||
condition_str = get_html_match_content(action_message.content, key="condition")
|
||||
if experience_str and condition_str:
|
||||
return Experience(experience_workspace_id=self.workspace_id,
|
||||
experience_role=self.llm.model_name,
|
||||
experience_desc=condition_str,
|
||||
experience_content=experience_str)
|
||||
else:
|
||||
logger.warning(f"action_message.content={action_message.content} re.search failed.")
|
||||
|
||||
return None
|
||||
|
||||
def _extract_experiences(self, trajectories: List[Trajectory], **kwargs) -> List[Experience]:
|
||||
experiences: List[Experience] = []
|
||||
for trajectory in trajectories:
|
||||
experience: Experience = self._extract_trajectory_experience(trajectory)
|
||||
if experience:
|
||||
experiences.append(experience)
|
||||
return experiences
|
||||
|
||||
|
||||
SUMMARIZER_REGISTRY.register(SimpleSummarizer, "simple")
|
||||
|
|
@ -0,0 +1,22 @@
|
|||
summary_prompt: |
|
||||
# Role
|
||||
You are a helpful assistant named BeyondAgent.
|
||||
|
||||
# User Question
|
||||
{query}
|
||||
|
||||
# Execution Process
|
||||
{execution_process}
|
||||
|
||||
# Answer
|
||||
{answer}
|
||||
|
||||
# Task
|
||||
Reflect on the strengths and weaknesses of the **Execution Process** and the **Answer** based on the **User Question**.
|
||||
Finally, summarize generalized experience from handling such problems to accumulate experience for future similar tasks.
|
||||
The experience should be broadly applicable, such as how to use tools effectively or approaches to solving certain types of problems.
|
||||
Also, specify the conditions or scenarios in which these experience are applicable.
|
||||
|
||||
# Output Format
|
||||
<condition> Output the scenarios or conditions in which applying this experience would be particularly effective... </condition>
|
||||
<experience> Output generalized experience, concise content is required... </experience>
|
||||
580
experiencemaker/module/summarizer/step_summarizer.py
Normal file
580
experiencemaker/module/summarizer/step_summarizer.py
Normal file
|
|
@ -0,0 +1,580 @@
|
|||
import re
|
||||
import uuid
|
||||
import json
|
||||
from typing import List, Dict, Any, Optional, Tuple
|
||||
from datetime import datetime
|
||||
from loguru import logger
|
||||
from pydantic import Field, model_validator
|
||||
|
||||
from experiencemaker.enumeration.role import Role
|
||||
from experiencemaker.module.summarizer.base_summarizer import BaseSummarizer
|
||||
from experiencemaker.schema.trajectory import Trajectory, Sample, SummaryMessage, Message
|
||||
from experiencemaker.schema.vector_store_node import VectorStoreNode
|
||||
from experiencemaker.storage.es_vector_store import EsVectorStore
|
||||
from experiencemaker.storage.file_vector_store import FileVectorStore
|
||||
|
||||
|
||||
class StepSummarizer(BaseSummarizer):
|
||||
"""
|
||||
Step-level experience extractor that focuses on extracting reusable experiences
|
||||
from individual steps or step sequences in trajectories
|
||||
"""
|
||||
|
||||
# Vector Store 配置
|
||||
vector_store_type: str = Field(default="file_vector_store")
|
||||
vector_store_hosts: str | List[str] = Field(default="http://localhost:9200")
|
||||
vector_store_index_name: str = Field(default="step_experience_store")
|
||||
store_dir: str = Field(default="./step_experiences/")
|
||||
|
||||
# 功能开关
|
||||
enable_step_segmentation: bool = Field(default=False)
|
||||
enable_similarity_search: bool = Field(default=False)
|
||||
enable_experience_validation: bool = Field(default=True)
|
||||
|
||||
# llm retries
|
||||
max_retries: int = Field(default=3)
|
||||
|
||||
@model_validator(mode="after")
|
||||
def init_vector_store(self):
|
||||
"""initialize"""
|
||||
if self.vector_store_type == "file_vector_store":
|
||||
self.vector_store = FileVectorStore(
|
||||
embedding_model=self.embedding_model,
|
||||
index_name=self.vector_store_index_name,
|
||||
store_dir=self.store_dir
|
||||
)
|
||||
elif self.vector_store_type == "es_vector_store":
|
||||
self.vector_store = EsVectorStore(
|
||||
embedding_model=self.embedding_model,
|
||||
index_name=self.vector_store_index_name,
|
||||
hosts=self.vector_store_hosts
|
||||
)
|
||||
else:
|
||||
raise ValueError(f"Unknown vector store type: {self.vector_store_type}")
|
||||
|
||||
return self
|
||||
|
||||
def extract_step_experiences_from_success(self, trajectories: List[Trajectory], **kwargs) -> List[SummaryMessage]:
|
||||
"""Extract step-level experiences from successful samples"""
|
||||
logger.info(f"Extracting step experiences from {len(trajectories)} successful trajectories")
|
||||
|
||||
all_experiences = []
|
||||
for trajectory in trajectories:
|
||||
step_sequences = self._segment_trajectory_into_steps(trajectory)
|
||||
|
||||
for step_seq in step_sequences:
|
||||
try:
|
||||
prompt = self.prompt_handler.success_step_experience_prompt.format(
|
||||
query=trajectory.query,
|
||||
step_sequence=self._format_step_sequence(step_seq),
|
||||
context=self._get_trajectory_context(trajectory, step_seq),
|
||||
outcome="successful"
|
||||
)
|
||||
|
||||
experience = self._extract_with_llm(prompt, "success")
|
||||
if experience:
|
||||
all_experiences.extend(experience)
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Error extracting success experience: {e}")
|
||||
continue
|
||||
|
||||
return all_experiences
|
||||
|
||||
def extract_step_experiences_from_failure(self, trajectories: List[Trajectory], **kwargs) -> List[SummaryMessage]:
|
||||
"""Extract step-level experiences from failed samples"""
|
||||
logger.info(f"Extracting step experiences from {len(trajectories)} failed trajectories")
|
||||
|
||||
all_experiences = []
|
||||
for trajectory in trajectories:
|
||||
step_sequences = self._segment_trajectory_into_steps(trajectory)
|
||||
|
||||
for step_seq in step_sequences:
|
||||
try:
|
||||
prompt = self.prompt_handler.failure_step_experience_prompt.format(
|
||||
query=trajectory.query,
|
||||
step_sequence=self._format_step_sequence(step_seq),
|
||||
context=self._get_trajectory_context(trajectory, step_seq),
|
||||
outcome="failed"
|
||||
)
|
||||
|
||||
experience = self._extract_with_llm(prompt, "failure")
|
||||
if experience:
|
||||
all_experiences.extend(experience)
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Error extracting failure experience: {e}")
|
||||
continue
|
||||
|
||||
return all_experiences
|
||||
|
||||
def extract_step_experiences_from_comparison(self,
|
||||
success_trajectories: List[Trajectory],
|
||||
failure_trajectories: List[Trajectory],
|
||||
**kwargs) -> List[SummaryMessage]:
|
||||
"""Extract step-level experiences from comparative samples"""
|
||||
logger.info(f"Extracting comparative step experiences from {len(success_trajectories)} success "
|
||||
f"and {len(failure_trajectories)} failure trajectories")
|
||||
|
||||
all_experiences = []
|
||||
|
||||
# Find similar step sequences for comparison
|
||||
similar_step_pairs = self._find_similar_step_sequences(success_trajectories, failure_trajectories)
|
||||
|
||||
for success_steps, failure_steps, similarity_score in similar_step_pairs:
|
||||
try:
|
||||
prompt = self.prompt_handler.comparative_step_experience_prompt.format(
|
||||
success_steps=self._format_step_sequence(success_steps),
|
||||
failure_steps=self._format_step_sequence(failure_steps),
|
||||
similarity_score=similarity_score
|
||||
)
|
||||
|
||||
experience = self._extract_with_llm(prompt, "comparative")
|
||||
if experience:
|
||||
all_experiences.extend(experience)
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Error extracting comparative experience: {e}")
|
||||
continue
|
||||
|
||||
return all_experiences
|
||||
|
||||
def extract_step_experiences_general(self, trajectories: List[Trajectory], **kwargs) -> List[SummaryMessage]:
|
||||
"""Extract general step experiences when no labels are provided"""
|
||||
logger.info(f"Extracting general step experiences from {len(trajectories)} trajectories")
|
||||
|
||||
all_experiences = []
|
||||
|
||||
for trajectory in trajectories:
|
||||
step_sequences = self._segment_trajectory_into_steps(trajectory)
|
||||
|
||||
for step_seq in step_sequences:
|
||||
try:
|
||||
prompt = self.prompt_handler.general_step_experience_prompt.format(
|
||||
query=trajectory.query,
|
||||
step_sequence=self._format_step_sequence(step_seq),
|
||||
context=self._get_trajectory_context(trajectory, step_seq)
|
||||
)
|
||||
|
||||
experience = self._extract_with_llm(prompt, "general")
|
||||
if experience:
|
||||
all_experiences.extend(experience)
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Error extracting general experience: {e}")
|
||||
continue
|
||||
|
||||
return all_experiences
|
||||
|
||||
def validate_experiences(self, experiences: List[SummaryMessage], **kwargs) -> List[SummaryMessage]:
|
||||
"""Validate the quality and validity of extracted experiences"""
|
||||
if not self.enable_experience_validation:
|
||||
return experiences
|
||||
|
||||
logger.info(f"Validating {len(experiences)} extracted experiences")
|
||||
|
||||
validated_experiences = []
|
||||
|
||||
for experience in experiences:
|
||||
try:
|
||||
validation_result = self._validate_single_experience(experience)
|
||||
|
||||
if validation_result["is_valid"]:
|
||||
# Add validation info to metadata
|
||||
experience.metadata.update({
|
||||
"validation_score": validation_result["score"],
|
||||
"validation_feedback": validation_result["feedback"],
|
||||
"validated_at": datetime.now().isoformat()
|
||||
})
|
||||
validated_experiences.append(experience)
|
||||
else:
|
||||
logger.warning(f"Experience validation failed: {validation_result['reason']}")
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Error validating experience: {e}")
|
||||
continue
|
||||
|
||||
logger.info(f"Validated {len(validated_experiences)} out of {len(experiences)} experiences")
|
||||
return validated_experiences
|
||||
|
||||
def store_experiences(self, experiences: List[SummaryMessage], **kwargs):
|
||||
"""Store experiences into vector storage"""
|
||||
if not experiences:
|
||||
logger.warning("No experiences to store")
|
||||
return
|
||||
|
||||
# Deduplication
|
||||
unique_experiences = self._deduplicate_experiences(experiences)
|
||||
logger.info(f"Storing {len(unique_experiences)} unique experiences (deduplicated from {len(experiences)})")
|
||||
|
||||
# Convert to storage nodes
|
||||
nodes = []
|
||||
for exp in unique_experiences:
|
||||
node = VectorStoreNode(
|
||||
content=exp.content,
|
||||
metadata={
|
||||
**exp.metadata,
|
||||
"stored_at": datetime.now().isoformat(),
|
||||
"experience_type": "step_level"
|
||||
}
|
||||
)
|
||||
nodes.append(node)
|
||||
|
||||
# Store to vector database
|
||||
refresh_index = kwargs.get("refresh_index", True)
|
||||
self.vector_store.insert(nodes, refresh_index=refresh_index)
|
||||
logger.info(f"Successfully stored {len(nodes)} step experiences")
|
||||
|
||||
def execute(self, trajectories: List[Trajectory], **kwargs) -> List[Sample]:
|
||||
"""Execute complete step-level experience extraction pipeline"""
|
||||
logger.info(f"Starting step-level experience extraction pipeline for {len(trajectories)} trajectories")
|
||||
|
||||
all_experiences = []
|
||||
|
||||
# Classify trajectories based on trajectory.done
|
||||
success_trajectories = [traj for traj in trajectories if traj.done]
|
||||
failure_trajectories = [traj for traj in trajectories if not traj.done]
|
||||
|
||||
# Process success and failure samples separately
|
||||
if success_trajectories:
|
||||
success_experiences = self.extract_step_experiences_from_success(success_trajectories, **kwargs)
|
||||
all_experiences.extend(success_experiences)
|
||||
|
||||
if failure_trajectories:
|
||||
failure_experiences = self.extract_step_experiences_from_failure(failure_trajectories, **kwargs)
|
||||
all_experiences.extend(failure_experiences)
|
||||
|
||||
# Comparative analysis (if similarity search is enabled)
|
||||
if success_trajectories and failure_trajectories and self.enable_similarity_search:
|
||||
comparative_experiences = self.extract_step_experiences_from_comparison(
|
||||
success_trajectories, failure_trajectories, **kwargs
|
||||
)
|
||||
all_experiences.extend(comparative_experiences)
|
||||
|
||||
# Validate experiences
|
||||
if self.enable_experience_validation:
|
||||
validated_experiences = self.validate_experiences(all_experiences, **kwargs)
|
||||
else:
|
||||
validated_experiences = all_experiences
|
||||
|
||||
# Store experiences
|
||||
if validated_experiences:
|
||||
self.store_experiences(validated_experiences, **kwargs)
|
||||
|
||||
# Construct return result
|
||||
return [Sample(steps=validated_experiences)]
|
||||
|
||||
# ========== Helper Methods ==========
|
||||
|
||||
def _segment_trajectory_into_steps(self, trajectory: Trajectory) -> List[List[Message]]:
|
||||
"""Segment trajectory into meaningful step sequences"""
|
||||
if not self.enable_step_segmentation:
|
||||
# If segmentation is not enabled, return the entire trajectory as one step sequence
|
||||
return [trajectory.steps]
|
||||
|
||||
try:
|
||||
# Use LLM for segmentation
|
||||
trajectory_content = self._format_trajectory_content(trajectory)
|
||||
|
||||
prompt = self.prompt_handler.step_segmentation_prompt.format(
|
||||
query=trajectory.query,
|
||||
trajectory_content=trajectory_content,
|
||||
total_steps=len(trajectory.steps)
|
||||
)
|
||||
|
||||
response = self.llm.chat([Message(role=Role.USER, content=prompt)])
|
||||
|
||||
# Parse segmentation points
|
||||
segment_points = self._parse_segmentation_response(response.content)
|
||||
|
||||
# Segment trajectory based on split points
|
||||
step_sequences = []
|
||||
start_idx = 0
|
||||
|
||||
for end_idx in segment_points:
|
||||
if start_idx < end_idx <= len(trajectory.steps):
|
||||
step_sequences.append(trajectory.steps[start_idx:end_idx])
|
||||
start_idx = end_idx
|
||||
|
||||
# Add remaining steps
|
||||
if start_idx < len(trajectory.steps):
|
||||
step_sequences.append(trajectory.steps[start_idx:])
|
||||
|
||||
return step_sequences if step_sequences else [trajectory.steps]
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Error in step segmentation: {e}, falling back to whole trajectory")
|
||||
return [trajectory.steps]
|
||||
|
||||
def _parse_segmentation_response(self, response: str) -> List[int]:
|
||||
"""Parse segmentation response to extract split point positions"""
|
||||
segment_points = []
|
||||
|
||||
# Try to extract JSON format split points
|
||||
json_pattern = r'```json\s*([\s\S]*?)\s*```'
|
||||
json_blocks = re.findall(json_pattern, response)
|
||||
|
||||
if json_blocks:
|
||||
try:
|
||||
parsed = json.loads(json_blocks[0])
|
||||
if isinstance(parsed, dict) and "segment_points" in parsed:
|
||||
segment_points = parsed["segment_points"]
|
||||
elif isinstance(parsed, list):
|
||||
segment_points = parsed
|
||||
except json.JSONDecodeError:
|
||||
pass
|
||||
|
||||
# If JSON parsing fails, try to extract numbers
|
||||
if not segment_points:
|
||||
numbers = re.findall(r'\b\d+\b', response)
|
||||
segment_points = [int(num) for num in numbers if int(num) > 0]
|
||||
|
||||
return sorted(list(set(segment_points))) # Remove duplicates and sort
|
||||
|
||||
def _format_step_sequence(self, step_sequence: List[Message]) -> str:
|
||||
"""Format step sequence to string"""
|
||||
formatted_steps = []
|
||||
for i, step in enumerate(step_sequence):
|
||||
step_info = f"Step {i + 1} [{step.role.value}]:"
|
||||
|
||||
if hasattr(step, 'reasoning_content') and step.reasoning_content:
|
||||
step_info += f"\nReasoning: {step.reasoning_content}"
|
||||
|
||||
step_info += f"\nContent: {step.content}"
|
||||
|
||||
if hasattr(step, 'tool_calls') and step.tool_calls:
|
||||
for tool_call in step.tool_calls:
|
||||
step_info += f"\nTool: {tool_call.name}({tool_call.arguments})"
|
||||
|
||||
formatted_steps.append(step_info)
|
||||
|
||||
return "\n\n".join(formatted_steps)
|
||||
|
||||
def _get_trajectory_context(self, trajectory: Trajectory, step_sequence: List[Message]) -> str:
|
||||
"""Get context of step sequence within trajectory"""
|
||||
# Find position of step sequence in trajectory
|
||||
start_idx = 0
|
||||
for i, step in enumerate(trajectory.steps):
|
||||
if step == step_sequence[0]:
|
||||
start_idx = i
|
||||
break
|
||||
|
||||
# Extract before and after context
|
||||
context_before = trajectory.steps[max(0, start_idx - 2):start_idx]
|
||||
context_after = trajectory.steps[start_idx + len(step_sequence):start_idx + len(step_sequence) + 2]
|
||||
|
||||
context = f"Query: {trajectory.query}\n"
|
||||
|
||||
if context_before:
|
||||
context += "Previous steps:\n" + "\n".join([f"- {step.content[:100]}..." for step in context_before]) + "\n"
|
||||
|
||||
if context_after:
|
||||
context += "Following steps:\n" + "\n".join([f"- {step.content[:100]}..." for step in context_after])
|
||||
|
||||
return context
|
||||
|
||||
def _format_trajectory_content(self, trajectory: Trajectory) -> str:
|
||||
"""Format trajectory content to string"""
|
||||
content = ""
|
||||
for i, step in enumerate(trajectory.steps):
|
||||
content += f"Step {i + 1} ({step.role.value}):\n{step.content}\n\n"
|
||||
return content
|
||||
|
||||
def _find_similar_step_sequences(self, success_trajectories: List[Trajectory],
|
||||
failure_trajectories: List[Trajectory]) -> List[Tuple]:
|
||||
"""Use embedding model to find similar step sequences for comparison"""
|
||||
if not self.enable_similarity_search:
|
||||
return []
|
||||
|
||||
try:
|
||||
similar_pairs = []
|
||||
|
||||
# Get step sequences from success and failure trajectories
|
||||
success_step_sequences = []
|
||||
for traj in success_trajectories:
|
||||
sequences = self._segment_trajectory_into_steps(traj)
|
||||
success_step_sequences.extend(sequences)
|
||||
|
||||
failure_step_sequences = []
|
||||
for traj in failure_trajectories:
|
||||
sequences = self._segment_trajectory_into_steps(traj)
|
||||
failure_step_sequences.extend(sequences)
|
||||
|
||||
# Limit comparison count to avoid computation overload
|
||||
max_sequences = 5
|
||||
success_step_sequences = success_step_sequences[:max_sequences]
|
||||
failure_step_sequences = failure_step_sequences[:max_sequences]
|
||||
|
||||
if not success_step_sequences or not failure_step_sequences:
|
||||
return []
|
||||
|
||||
# Generate text representations of step sequences for embedding
|
||||
success_texts = [self._format_step_sequence(seq) for seq in success_step_sequences]
|
||||
failure_texts = [self._format_step_sequence(seq) for seq in failure_step_sequences]
|
||||
|
||||
# Get embeddings
|
||||
success_embeddings = self.embedding_model.get_embeddings(success_texts)
|
||||
failure_embeddings = self.embedding_model.get_embeddings(failure_texts)
|
||||
|
||||
# Calculate similarity and find most similar pairs
|
||||
for i, s_emb in enumerate(success_embeddings):
|
||||
for j, f_emb in enumerate(failure_embeddings):
|
||||
similarity = self._calculate_cosine_similarity(s_emb, f_emb)
|
||||
|
||||
if similarity > 0.3: # Similarity threshold
|
||||
similar_pairs.append((
|
||||
success_step_sequences[i],
|
||||
failure_step_sequences[j],
|
||||
similarity
|
||||
))
|
||||
|
||||
# Return top 3 most similar pairs
|
||||
return sorted(similar_pairs, key=lambda x: x[2], reverse=True)[:3]
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Error finding similar step sequences: {e}")
|
||||
return []
|
||||
|
||||
def _calculate_cosine_similarity(self, embedding1: List[float], embedding2: List[float]) -> float:
|
||||
"""Calculate cosine similarity between two embedding vectors"""
|
||||
try:
|
||||
import numpy as np
|
||||
|
||||
vec1 = np.array(embedding1)
|
||||
vec2 = np.array(embedding2)
|
||||
|
||||
# Calculate cosine similarity
|
||||
dot_product = np.dot(vec1, vec2)
|
||||
norm1 = np.linalg.norm(vec1)
|
||||
norm2 = np.linalg.norm(vec2)
|
||||
|
||||
if norm1 == 0 or norm2 == 0:
|
||||
return 0.0
|
||||
|
||||
return dot_product / (norm1 * norm2)
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Error calculating cosine similarity: {e}")
|
||||
return 0.0
|
||||
import json
|
||||
|
||||
def _extract_with_llm(self, prompt: str, experience_type: str) -> List[SummaryMessage]:
|
||||
for attempt in range(self.max_retries):
|
||||
try:
|
||||
response = self.llm.chat([Message(role=Role.USER, content=prompt)])
|
||||
experiences = self._parse_experience_response(response.content, experience_type)
|
||||
|
||||
if experiences:
|
||||
return experiences
|
||||
|
||||
except Exception as e:
|
||||
logger.warning(f"Attempt {attempt + 1} failed for experience extraction: {e}")
|
||||
|
||||
logger.error(f"Failed to extract experience after {self.max_retries} attempts")
|
||||
return []
|
||||
|
||||
def _parse_experience_response(self, response: str, experience_type: str) -> List[SummaryMessage]:
|
||||
"""解析经验抽取响应"""
|
||||
experiences = []
|
||||
|
||||
try:
|
||||
# 尝试提取JSON格式的经验
|
||||
json_pattern = r'```json\s*([\s\S]*?)\s*```'
|
||||
json_blocks = re.findall(json_pattern, response)
|
||||
|
||||
for block in json_blocks:
|
||||
try:
|
||||
parsed = json.loads(block)
|
||||
if isinstance(parsed, list):
|
||||
for exp_data in parsed:
|
||||
experience = self._create_experience_message(exp_data, experience_type)
|
||||
if experience:
|
||||
experiences.append(experience)
|
||||
else:
|
||||
experience = self._create_experience_message(parsed, experience_type)
|
||||
if experience:
|
||||
experiences.append(experience)
|
||||
except json.JSONDecodeError:
|
||||
continue
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Error parsing experience response: {e}")
|
||||
|
||||
return experiences
|
||||
|
||||
def _create_experience_message(self, exp_data: Dict[str, Any], experience_type: str) -> Optional[SummaryMessage]:
|
||||
"""创建经验消息对象"""
|
||||
try:
|
||||
condition = exp_data.get("when_to_use", exp_data.get("condition", ""))
|
||||
experience_content = exp_data.get("experience", exp_data.get("tip_content", exp_data.get("tips", "")))
|
||||
|
||||
if not condition or not experience_content:
|
||||
return None
|
||||
|
||||
metadata = {
|
||||
"experience": experience_content,
|
||||
"experience_type": experience_type,
|
||||
"tags": exp_data.get("tags", []),
|
||||
"confidence": exp_data.get("confidence", 0.5),
|
||||
"extracted_at": datetime.now().isoformat(),
|
||||
"experience_id": str(uuid.uuid4())
|
||||
}
|
||||
|
||||
return SummaryMessage(content=condition, metadata=metadata)
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Error creating experience message: {e}")
|
||||
return None
|
||||
|
||||
def _validate_single_experience(self, experience: SummaryMessage) -> Dict[str, Any]:
|
||||
"""验证单个经验的有效性"""
|
||||
try:
|
||||
prompt = self.prompt_handler.experience_validation_prompt.format(
|
||||
condition=experience.content,
|
||||
experience_content=experience.metadata.get("experience", ""),
|
||||
experience_type=experience.metadata.get("experience_type", ""),
|
||||
tags=experience.metadata.get("tags", [])
|
||||
)
|
||||
|
||||
response = self.llm.chat([Message(role=Role.USER, content=prompt)])
|
||||
|
||||
# 解析验证结果
|
||||
is_valid = "valid" in response.content.lower() and "invalid" not in response.content.lower()
|
||||
score_match = re.search(r'score[:\s]*([0-9.]+)', response.content.lower())
|
||||
score = float(score_match.group(1)) if score_match else 0.5
|
||||
|
||||
return {
|
||||
"is_valid": is_valid and score > 0.3,
|
||||
"score": score,
|
||||
"feedback": response.content,
|
||||
"reason": "" if is_valid else "Low validation score or marked as invalid"
|
||||
}
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Error validating experience: {e}")
|
||||
return {"is_valid": False, "score": 0.0, "feedback": "", "reason": str(e)}
|
||||
|
||||
def _deduplicate_experiences(self, experiences: List[SummaryMessage]) -> List[SummaryMessage]:
|
||||
unique_experiences = []
|
||||
seen_contents = set()
|
||||
|
||||
for exp in experiences:
|
||||
content_hash = hash(exp.content)
|
||||
|
||||
if content_hash not in seen_contents:
|
||||
seen_contents.add(content_hash)
|
||||
unique_experiences.append(exp)
|
||||
|
||||
return unique_experiences
|
||||
|
||||
def extract_samples(self, trajectories: List[Trajectory], **kwargs) -> List[Sample]:
|
||||
experiences = self.execute(trajectories, **kwargs)
|
||||
return [Sample(steps=experiences)] if experiences else []
|
||||
|
||||
def insert_into_vector_store(self, samples: List[Sample], **kwargs):
|
||||
all_experiences = []
|
||||
for sample in samples:
|
||||
all_experiences.extend(sample.steps)
|
||||
|
||||
if all_experiences:
|
||||
self.store_experiences(all_experiences, **kwargs)
|
||||
|
|
@ -7,19 +7,8 @@ from experiencemaker.schema.trajectory import Trajectory
|
|||
|
||||
|
||||
class BaseTrainner(BaseModel, ABC):
|
||||
"""
|
||||
load model/prompt
|
||||
data
|
||||
env
|
||||
-> 新cpt
|
||||
|
||||
off policy/onpolicy
|
||||
"""
|
||||
traj_buffer: List[Trajectory] = Field()
|
||||
|
||||
def __init__(self, **kwargs):
|
||||
super().__init__(**kwargs)
|
||||
|
||||
def fit(self):
|
||||
return
|
||||
|
||||
|
|
@ -29,17 +18,17 @@ class BaseTrainner(BaseModel, ABC):
|
|||
|
||||
class BaseContextTrainner(BaseTrainner):
|
||||
"""
|
||||
load model/prompt/db/buffer -> 新cpt
|
||||
load model/prompt/db/buffer -> new cpt
|
||||
"""
|
||||
|
||||
|
||||
class BaseSummaryTrainner(BaseTrainner):
|
||||
"""
|
||||
load model/prompt/db/buffer -> 新cpt
|
||||
load model/prompt/db/buffer -> new cpt
|
||||
"""
|
||||
|
||||
|
||||
class BasePolicyTrainner(BaseTrainner):
|
||||
"""
|
||||
load model/prompt/db/buffer -> 新cpt
|
||||
load model/prompt/db/buffer -> new cpt
|
||||
"""
|
||||
|
|
|
|||
83
experiencemaker/schema/experience.py
Normal file
83
experiencemaker/schema/experience.py
Normal file
|
|
@ -0,0 +1,83 @@
|
|||
import datetime
|
||||
from typing import List
|
||||
from uuid import uuid4
|
||||
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
from experiencemaker.schema.vector_store_node import VectorStoreNode
|
||||
|
||||
|
||||
class ExperienceFunctionArg(BaseModel):
|
||||
arg_name: str = Field(default=..., description="argument name")
|
||||
arg_type: str = Field(default=..., description="argument type, like: 'str', 'int', 'bool'")
|
||||
required: bool = Field(default=True, description="whether the argument is required")
|
||||
|
||||
|
||||
class ExperienceFunction(BaseModel):
|
||||
func_code: str = Field(default=..., description="function code")
|
||||
func_name: str = Field(default=..., description="function name")
|
||||
func_args: List[ExperienceFunctionArg] = Field(default_factory=list, description="function arguments")
|
||||
|
||||
|
||||
class Experience(BaseModel):
|
||||
experience_id: str = Field(default_factory=lambda: uuid4().hex, description="experience unique id")
|
||||
experience_workspace_id: str = Field(default="", description="unique workspace id")
|
||||
experience_role: str = Field(default="", description="experience role")
|
||||
experience_desc: str = Field(default="", description="use condition/purpose. It will be used in vector matching")
|
||||
experience_content: str | bytes = Field(default="", description="content of the experience")
|
||||
experience_function: ExperienceFunction | None = Field(default=None, description="experience function(optional)")
|
||||
experience_score: float = Field(default=0.0, description="score of the experience")
|
||||
experience_created_time: str = Field(default_factory=lambda: datetime.datetime.now().strftime("%Y-%m-%d %H:%M:%S"))
|
||||
experience_modified_time: str = Field(default_factory=lambda: datetime.datetime.now().strftime("%Y-%m-%d %H:%M:%S"))
|
||||
metadata: dict = Field(default_factory=dict, description="additional metadata")
|
||||
|
||||
def to_vector_store_node(self) -> VectorStoreNode:
|
||||
metadata: dict = {
|
||||
"experience_role": self.experience_role,
|
||||
"experience_content": self.experience_content,
|
||||
"experience_function": self.experience_function.model_dump(),
|
||||
"experience_score": self.experience_score,
|
||||
"experience_created_time": self.experience_created_time,
|
||||
"experience_modified_time": self.experience_modified_time,
|
||||
"metadata": self.metadata,
|
||||
}
|
||||
return VectorStoreNode(
|
||||
unique_id=self.experience_id,
|
||||
workspace_id=self.experience_workspace_id,
|
||||
content=self.experience_desc,
|
||||
metadata=metadata)
|
||||
|
||||
@classmethod
|
||||
def from_vector_store_node(cls, node: VectorStoreNode) -> "Experience":
|
||||
return cls(
|
||||
experience_id=node.unique_id,
|
||||
experience_workspace_id=node.workspace_id,
|
||||
experience_role=node.metadata.get("experience_role", ""),
|
||||
experience_desc=node.content,
|
||||
experience_content=node.metadata.get("experience_content", ""),
|
||||
experience_function=node.metadata.get("experience_function", None),
|
||||
experience_score=node.metadata.get("experience_score", 0.0),
|
||||
experience_created_time=node.metadata.get("experience_created_time", ""),
|
||||
experience_modified_time=node.metadata.get("experience_modified_time", ""),
|
||||
metadata=node.metadata.get("metadata", {}))
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
e1 = Experience(
|
||||
experience_workspace_id="w_1024",
|
||||
experience_role="qwen3",
|
||||
experience_desc="test desc",
|
||||
experience_content="test content",
|
||||
experience_function=ExperienceFunction(
|
||||
func_code="def a():\n return",
|
||||
func_name="a",
|
||||
func_args=[ExperienceFunctionArg(arg_name="x", arg_type="str", required=True)]
|
||||
),
|
||||
experience_score=0.99,
|
||||
metadata={"haha": 1}
|
||||
)
|
||||
print(e1.model_dump_json(indent=2))
|
||||
v1 = e1.to_vector_store_node()
|
||||
print(v1.model_dump_json(indent=2))
|
||||
e2 = Experience.from_vector_store_node(v1)
|
||||
print(e2.model_dump_json(indent=2))
|
||||
|
|
@ -1,19 +0,0 @@
|
|||
from importlib import import_module
|
||||
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
from experiencemaker.utils.file_handler import FileHandler
|
||||
|
||||
|
||||
class ModuleLoader(BaseModel):
|
||||
class_path: str = Field(default=...)
|
||||
class_name: str = Field(default=...)
|
||||
config_path: str = Field(default="")
|
||||
config: dict = Field(default_factory=dict)
|
||||
|
||||
def load_from_config(self):
|
||||
return getattr(import_module(self.class_path), self.class_name)(**self.config)
|
||||
|
||||
def load_from_path(self):
|
||||
self.config = FileHandler(file_path=self.config_path).load()
|
||||
return self.load_from_config()
|
||||
|
|
@ -1,13 +1,12 @@
|
|||
from abc import ABC
|
||||
from typing import List
|
||||
|
||||
from pydantic import Field
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
from experiencemaker.schema.module_loader import ModuleLoader
|
||||
from experiencemaker.schema.trajectory import Trajectory
|
||||
|
||||
|
||||
class BaseRequest(ModuleLoader, ABC):
|
||||
class BaseRequest(BaseModel, ABC):
|
||||
metadata: dict = Field(default_factory=dict)
|
||||
|
||||
|
||||
|
|
@ -21,4 +20,4 @@ class ContextGeneratorRequest(BaseRequest):
|
|||
|
||||
class SummarizerRequest(BaseRequest):
|
||||
trajectories: List[Trajectory] = Field(default_factory=dict)
|
||||
return_samples: bool = Field(default=False)
|
||||
return_experience: bool = Field(default=True)
|
||||
|
|
|
|||
|
|
@ -3,7 +3,8 @@ from typing import List
|
|||
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
from experiencemaker.schema.trajectory import Trajectory, ContextMessage, Sample
|
||||
from experiencemaker.schema.experience import Experience
|
||||
from experiencemaker.schema.trajectory import Trajectory, ContextMessage
|
||||
|
||||
|
||||
class BaseResponse(BaseModel, ABC):
|
||||
|
|
@ -20,4 +21,4 @@ class ContextGeneratorResponse(BaseResponse):
|
|||
|
||||
|
||||
class SummarizerResponse(BaseResponse):
|
||||
extract_samples: List[Sample] = Field(default_factory=list)
|
||||
experiences: List[Experience] = Field(default_factory=list)
|
||||
|
|
|
|||
|
|
@ -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,7 +111,9 @@ 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 = ""
|
||||
|
|
|
|||
164
experiencemaker/service/experience_maker_service.py
Normal file
164
experiencemaker/service/experience_maker_service.py
Normal file
|
|
@ -0,0 +1,164 @@
|
|||
from typing import List
|
||||
|
||||
from loguru import logger
|
||||
from pydantic import BaseModel, Field, model_validator
|
||||
|
||||
from experiencemaker.model.base_embedding_model import BaseEmbeddingModel, EMBEDDING_MODEL_REGISTRY
|
||||
from experiencemaker.model.base_llm import BaseLLM, LLM_REGISTRY
|
||||
from experiencemaker.module.agent_wrapper.agent_wrapper_mixin import AGENT_WRAPPER_REGISTRY, AgentWrapperMixin
|
||||
from experiencemaker.module.context_generator.base_context_generator import BaseContextGenerator, \
|
||||
CONTEXT_GENERATOR_REGISTRY
|
||||
from experiencemaker.module.summarizer.base_summarizer import BaseSummarizer, SUMMARIZER_REGISTRY
|
||||
from experiencemaker.schema.experience import Experience
|
||||
from experiencemaker.schema.request import AgentWrapperRequest, ContextGeneratorRequest, SummarizerRequest
|
||||
from experiencemaker.schema.response import AgentWrapperResponse, ContextGeneratorResponse, SummarizerResponse
|
||||
from experiencemaker.schema.trajectory import Trajectory, ContextMessage
|
||||
from experiencemaker.storage.base_vector_store import BaseVectorStore, VECTOR_STORE_REGISTRY
|
||||
|
||||
|
||||
class ExperienceMakerService(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: AgentWrapperMixin | None = Field(default=None)
|
||||
context_generator: BaseContextGenerator | None = Field(default=None)
|
||||
summarizer: BaseSummarizer | None = Field(default=None)
|
||||
|
||||
@staticmethod
|
||||
def init_llm(llm_config: dict) -> BaseLLM:
|
||||
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. " \
|
||||
f"supported={LLM_REGISTRY.registered_modules}"
|
||||
llm = LLM_REGISTRY[backend](**llm_config)
|
||||
logger.info(f"llm is inited with backend={backend} params={llm_config}")
|
||||
return llm
|
||||
|
||||
def get_llm(self, config: dict, llm: BaseLLM = None) -> BaseLLM:
|
||||
if "llm" in config:
|
||||
llm_config = config.pop("llm")
|
||||
llm = self.init_llm(llm_config)
|
||||
elif llm is None:
|
||||
raise RuntimeError("llm must be provided.")
|
||||
return llm
|
||||
|
||||
@staticmethod
|
||||
def init_embedding_model(embedding_model_config: dict) -> BaseEmbeddingModel:
|
||||
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. " \
|
||||
f"supported={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
|
||||
|
||||
def get_embedding_model(self, config: dict, embedding_model: BaseEmbeddingModel = None) -> BaseEmbeddingModel:
|
||||
if "embedding_model" in config:
|
||||
embedding_model_config = config.pop("embedding_model")
|
||||
embedding_model = self.init_embedding_model(embedding_model_config)
|
||||
elif embedding_model is None:
|
||||
raise RuntimeError("embedding_model must be provided.")
|
||||
return embedding_model
|
||||
|
||||
def init_vector_store(self, vector_store_config: dict) -> BaseVectorStore:
|
||||
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. " \
|
||||
f"supported={VECTOR_STORE_REGISTRY.registered_modules}"
|
||||
embedding_model = self.get_embedding_model(vector_store_config, embedding_model=self.embedding_model)
|
||||
vector_store = VECTOR_STORE_REGISTRY[backend](**vector_store_config, embedding_model=embedding_model)
|
||||
logger.info(f"vector_store is inited with backend={backend} params={vector_store_config}")
|
||||
return vector_store
|
||||
|
||||
def get_vector_store(self, config: dict, vector_store: BaseVectorStore = None) -> BaseVectorStore:
|
||||
if "vector_store" in config:
|
||||
vector_store_config = config.pop("vector_store")
|
||||
vector_store = self.init_vector_store(vector_store_config)
|
||||
elif vector_store is None:
|
||||
raise RuntimeError("vector_store must be provided.")
|
||||
return vector_store
|
||||
|
||||
def init_context_generator(self, context_generator_config: dict) -> BaseContextGenerator:
|
||||
backend = context_generator_config.pop("backend", None)
|
||||
assert backend is not None, "context_generator must have a backend like `simple`."
|
||||
assert backend in CONTEXT_GENERATOR_REGISTRY, f"context_generator backend={backend} not supported. " \
|
||||
f"supported={CONTEXT_GENERATOR_REGISTRY.registered_modules}"
|
||||
llm = self.get_llm(context_generator_config, llm=self.llm)
|
||||
vector_store = self.get_vector_store(context_generator_config, vector_store=self.vector_store)
|
||||
context_generator: BaseContextGenerator = CONTEXT_GENERATOR_REGISTRY[backend](
|
||||
**context_generator_config, llm=llm, vector_store=vector_store)
|
||||
logger.info(f"context_generator is inited with backend={backend} params={context_generator_config}")
|
||||
return context_generator
|
||||
|
||||
def init_summarizer(self, summarizer_config: dict) -> BaseSummarizer:
|
||||
backend = summarizer_config.pop("backend", None)
|
||||
assert backend is not None, "summarizer must have a backend like `simple`."
|
||||
assert backend in SUMMARIZER_REGISTRY, f"summarizer backend={backend} not supported. " \
|
||||
f"supported={SUMMARIZER_REGISTRY.registered_modules}"
|
||||
llm = self.get_llm(summarizer_config, llm=self.llm)
|
||||
vector_store = self.get_vector_store(summarizer_config, vector_store=self.vector_store)
|
||||
summarizer: BaseSummarizer = SUMMARIZER_REGISTRY[backend](**summarizer_config,
|
||||
llm=llm, vector_store=vector_store)
|
||||
logger.info(f"summarizer is inited with backend={backend} params={summarizer_config}")
|
||||
return summarizer
|
||||
|
||||
def init_agent_wrapper(self, agent_wrapper_config: dict) -> AgentWrapperMixin:
|
||||
backend = agent_wrapper_config.pop("backend", None)
|
||||
assert backend is not None, "agent_wrapper must have a backend like `simple`."
|
||||
assert backend in AGENT_WRAPPER_REGISTRY, f"agent_wrapper backend={backend} not supported. " \
|
||||
f"supported={AGENT_WRAPPER_REGISTRY.registered_modules}"
|
||||
|
||||
llm = self.get_llm(agent_wrapper_config, llm=self.llm)
|
||||
|
||||
agent_wrapper: AgentWrapperMixin = AGENT_WRAPPER_REGISTRY[backend](
|
||||
**agent_wrapper_config, llm=llm, context_generator=self.context_generator)
|
||||
logger.info(f"agent_wrapper is inited with backend={backend} params={agent_wrapper_config}")
|
||||
return agent_wrapper
|
||||
|
||||
@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)
|
||||
|
||||
if self.context_generator_config:
|
||||
self.context_generator = self.init_context_generator(self.context_generator_config)
|
||||
|
||||
if self.summarizer_config:
|
||||
self.summarizer = self.init_summarizer(self.summarizer_config)
|
||||
|
||||
if self.agent_wrapper_config:
|
||||
self.agent_wrapper = self.init_agent_wrapper(self.agent_wrapper_config)
|
||||
|
||||
def call_agent_wrapper(self, request: AgentWrapperRequest) -> AgentWrapperResponse:
|
||||
assert self.agent_wrapper is not None, "agent_wrapper must be provided."
|
||||
trajectory: Trajectory = self.agent_wrapper.execute(request.query, **request.metadata)
|
||||
return AgentWrapperResponse(trajectory=trajectory)
|
||||
|
||||
def call_context_generator(self, request: ContextGeneratorRequest) -> ContextGeneratorResponse:
|
||||
assert self.context_generator is not None, "context_generator must be provided."
|
||||
context_msg: ContextMessage = self.context_generator.execute(request.trajectory, **request.metadata)
|
||||
return ContextGeneratorResponse(context_msg=context_msg)
|
||||
|
||||
def call_summarizer(self, request: SummarizerRequest) -> SummarizerResponse:
|
||||
assert self.summarizer is not None, "summarizer must be provided."
|
||||
experiences: List[Experience] = self.summarizer.execute(request.trajectories, request.return_experience,
|
||||
**request.metadata)
|
||||
return SummarizerResponse(experiences=experiences)
|
||||
50
experiencemaker/service/http_service.py
Normal file
50
experiencemaker/service/http_service.py
Normal file
|
|
@ -0,0 +1,50 @@
|
|||
import argparse
|
||||
import json
|
||||
|
||||
import uvicorn
|
||||
from fastapi import FastAPI
|
||||
|
||||
from experiencemaker.schema.request import AgentWrapperRequest, ContextGeneratorRequest, SummarizerRequest
|
||||
from experiencemaker.schema.response import AgentWrapperResponse, ContextGeneratorResponse, SummarizerResponse
|
||||
from experiencemaker.service.experience_maker_service import ExperienceMakerService
|
||||
from experiencemaker.utils.file_handler import FileHandler
|
||||
|
||||
app = FastAPI()
|
||||
service: ExperienceMakerService | None = None
|
||||
|
||||
@app.post('/agent_wrapper', response_model=AgentWrapperResponse)
|
||||
def call_agent_wrapper(request: AgentWrapperRequest):
|
||||
return service.call_agent_wrapper(request)
|
||||
|
||||
|
||||
@app.post('/context_generator', response_model=ContextGeneratorResponse)
|
||||
def call_context_generator(request: ContextGeneratorRequest):
|
||||
return service.call_context_generator(request)
|
||||
|
||||
|
||||
@app.post('/summarizer', response_model=SummarizerResponse)
|
||||
def call_summarizer(request: SummarizerRequest):
|
||||
return service.call_summarizer(request)
|
||||
|
||||
|
||||
# launch with: python -m experiencemaker.service.http_service
|
||||
|
||||
if __name__ == '__main__':
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument('--config', type=str, help='config dict')
|
||||
parser.add_argument('--config_path', type=str, help='config load path')
|
||||
args = parser.parse_args()
|
||||
|
||||
if args.config_path:
|
||||
config = FileHandler(file_path=args.config_path).load()
|
||||
elif args.config:
|
||||
config = json.loads(args.config)
|
||||
else:
|
||||
raise RuntimeError("both config and config_path are not specified")
|
||||
|
||||
service = ExperienceMakerService(**config)
|
||||
uvicorn.run(app,
|
||||
host=service.host,
|
||||
port=service.port,
|
||||
timeout_keep_alive=service.timeout_keep_alive,
|
||||
limit_concurrency=service.limit_concurrency)
|
||||
|
|
@ -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
|
||||
|
|
@ -5,6 +5,7 @@ from pydantic import BaseModel, Field
|
|||
|
||||
from experiencemaker.model.base_embedding_model import BaseEmbeddingModel
|
||||
from experiencemaker.schema.vector_store_node import VectorStoreNode
|
||||
from experiencemaker.utils.registry import Registry
|
||||
|
||||
|
||||
class BaseVectorStore(BaseModel, ABC):
|
||||
|
|
@ -24,3 +25,6 @@ class BaseVectorStore(BaseModel, ABC):
|
|||
|
||||
def retrieve_by_query(self, query: str, top_k: int = 3, **kwargs) -> List[VectorStoreNode]:
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
VECTOR_STORE_REGISTRY = Registry[BaseVectorStore]("vector_store")
|
||||
|
|
|
|||
|
|
@ -6,8 +6,9 @@ from elasticsearch.helpers import bulk
|
|||
from loguru import logger
|
||||
from pydantic import Field, PrivateAttr, model_validator
|
||||
|
||||
from experiencemaker.model.openai_compatible_embedding_model import OpenAICompatibleEmbeddingModel
|
||||
from experiencemaker.schema.vector_store_node import VectorStoreNode
|
||||
from experiencemaker.storage.base_vector_store import BaseVectorStore
|
||||
from experiencemaker.storage.base_vector_store import BaseVectorStore, VECTOR_STORE_REGISTRY
|
||||
|
||||
|
||||
class EsVectorStore(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):
|
||||
|
|
@ -103,7 +104,7 @@ class EsVectorStore(BaseVectorStore):
|
|||
self.retrieve_filters.clear()
|
||||
return self
|
||||
|
||||
def insert(self, nodes: VectorStoreNode | List[VectorStoreNode], refresh_index: bool = False, **kwargs):
|
||||
def insert(self, nodes: VectorStoreNode | List[VectorStoreNode], refresh_index: bool = True, **kwargs):
|
||||
if isinstance(nodes, VectorStoreNode):
|
||||
nodes = [nodes]
|
||||
|
||||
|
|
@ -118,7 +119,7 @@ class EsVectorStore(BaseVectorStore):
|
|||
if refresh_index:
|
||||
self.refresh_index()
|
||||
|
||||
def update(self, nodes: VectorStoreNode | List[VectorStoreNode], refresh_index: bool = False, **kwargs):
|
||||
def update(self, nodes: VectorStoreNode | List[VectorStoreNode], refresh_index: bool = True, **kwargs):
|
||||
if isinstance(nodes, VectorStoreNode):
|
||||
nodes = [nodes]
|
||||
|
||||
|
|
@ -167,3 +168,77 @@ class EsVectorStore(BaseVectorStore):
|
|||
|
||||
self.retrieve_filters.clear()
|
||||
return nodes
|
||||
|
||||
|
||||
VECTOR_STORE_REGISTRY.register(EsVectorStore, "elasticsearch")
|
||||
|
||||
|
||||
def main():
|
||||
from experiencemaker.utils.util_function import load_env_keys
|
||||
load_env_keys()
|
||||
|
||||
embedding_model = OpenAICompatibleEmbeddingModel(dimensions=1024)
|
||||
index_name = "rag_nodes_index"
|
||||
hosts = "http://11.160.132.46:8200"
|
||||
es = EsVectorStore(hosts=hosts, embedding_model=embedding_model, index_name=index_name)
|
||||
es.delete_index()
|
||||
es.create_index()
|
||||
|
||||
sample_nodes = [
|
||||
VectorStoreNode(
|
||||
workspace_id="w1",
|
||||
content="Artificial intelligence is a technology that simulates human intelligence.",
|
||||
metadata={
|
||||
"node_type": "n1",
|
||||
}
|
||||
),
|
||||
VectorStoreNode(
|
||||
workspace_id="w1",
|
||||
content="AI is the future of mankind.",
|
||||
metadata={
|
||||
"node_type": "n1",
|
||||
}
|
||||
),
|
||||
VectorStoreNode(
|
||||
workspace_id="w1",
|
||||
content="I want to eat fish!",
|
||||
metadata={
|
||||
"node_type": "n2",
|
||||
}
|
||||
),
|
||||
VectorStoreNode(
|
||||
workspace_id="w2",
|
||||
content="The bigger the storm, the more expensive the fish.",
|
||||
metadata={
|
||||
"node_type": "n1",
|
||||
}
|
||||
),
|
||||
]
|
||||
|
||||
es.insert(sample_nodes, refresh_index=True)
|
||||
|
||||
logger.info("=" * 20)
|
||||
results = es.add_term_filter(key="workspace_id", value="w1") \
|
||||
.add_term_filter(key="metadata.node_type", value="n1") \
|
||||
.retrieve_by_query("What is AI?", top_k=5)
|
||||
for r in results:
|
||||
logger.info(r.model_dump(exclude={"vector"}))
|
||||
logger.info("=" * 20)
|
||||
|
||||
logger.info("=" * 20)
|
||||
results = es.add_term_filter(key="workspace_id", value="w1") \
|
||||
.retrieve_by_query("What is AI?", top_k=5)
|
||||
for r in results:
|
||||
logger.info(r.model_dump(exclude={"vector"}))
|
||||
logger.info("=" * 20)
|
||||
|
||||
logger.info("=" * 20)
|
||||
results = es.retrieve_by_query("What is AI?", top_k=5)
|
||||
for r in results:
|
||||
logger.info(r.model_dump(exclude={"vector"}))
|
||||
logger.info("=" * 20)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
# launch with: python -m experiencemaker.storage.es_vector_store
|
||||
|
|
|
|||
|
|
@ -7,8 +7,9 @@ from typing import List, Any
|
|||
from loguru import logger
|
||||
from pydantic import Field, model_validator, PrivateAttr
|
||||
|
||||
from experiencemaker.model.openai_compatible_embedding_model import OpenAICompatibleEmbeddingModel
|
||||
from experiencemaker.schema.vector_store_node import VectorStoreNode
|
||||
from experiencemaker.storage.base_vector_store import BaseVectorStore
|
||||
from experiencemaker.storage.base_vector_store import BaseVectorStore, VECTOR_STORE_REGISTRY
|
||||
|
||||
|
||||
class FileVectorStore(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,61 @@ 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")
|
||||
|
||||
|
||||
def main():
|
||||
from experiencemaker.utils.util_function import load_env_keys
|
||||
load_env_keys()
|
||||
|
||||
embedding_model = OpenAICompatibleEmbeddingModel(dimensions=1024)
|
||||
index_name = "rag_nodes_index"
|
||||
client = FileVectorStore(embedding_model=embedding_model, index_name=index_name)
|
||||
client.delete_index()
|
||||
client.create_index()
|
||||
|
||||
sample_nodes = [
|
||||
VectorStoreNode(
|
||||
workspace_id="w1",
|
||||
content="Artificial intelligence is a technology that simulates human intelligence.",
|
||||
metadata={
|
||||
"node_type": "n1",
|
||||
}
|
||||
),
|
||||
VectorStoreNode(
|
||||
workspace_id="w1",
|
||||
content="AI is the future of mankind.",
|
||||
metadata={
|
||||
"node_type": "n1",
|
||||
}
|
||||
),
|
||||
VectorStoreNode(
|
||||
workspace_id="w1",
|
||||
content="I want to eat fish!",
|
||||
metadata={
|
||||
"node_type": "n2",
|
||||
}
|
||||
),
|
||||
VectorStoreNode(
|
||||
workspace_id="w2",
|
||||
content="The bigger the storm, the more expensive the fish.",
|
||||
metadata={
|
||||
"node_type": "n1",
|
||||
}
|
||||
),
|
||||
]
|
||||
|
||||
client.insert(sample_nodes)
|
||||
|
||||
logger.info("=" * 20)
|
||||
results = client.retrieve_by_query("What is AI?", top_k=5)
|
||||
for r in results:
|
||||
logger.info(r.model_dump(exclude={"vector"}))
|
||||
logger.info("=" * 20)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
# launch with: python -m experiencemaker.storage.file_vector_store
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -3,6 +3,8 @@ from abc import ABC
|
|||
from loguru import logger
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
from experiencemaker.utils.registry import Registry
|
||||
|
||||
|
||||
class BaseTool(BaseModel, ABC):
|
||||
tool_id: str = Field(default="")
|
||||
|
|
@ -77,3 +79,6 @@ class BaseTool(BaseModel, ABC):
|
|||
|
||||
def get_cache_id(self, **kwargs) -> str:
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
TOOL_REGISTRY = Registry[BaseTool]("tool")
|
||||
|
|
|
|||
|
|
@ -1,7 +1,7 @@
|
|||
import sys
|
||||
from io import StringIO
|
||||
|
||||
from experiencemaker.tool.base_tool import BaseTool
|
||||
from experiencemaker.tool.base_tool import BaseTool, TOOL_REGISTRY
|
||||
|
||||
|
||||
class CodeTool(BaseTool):
|
||||
|
|
@ -35,6 +35,8 @@ class CodeTool(BaseTool):
|
|||
return result
|
||||
|
||||
|
||||
TOOL_REGISTRY.register(CodeTool, "code")
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
tool = CodeTool()
|
||||
|
|
|
|||
|
|
@ -6,7 +6,7 @@ from dashscope.api_entities.dashscope_response import Message
|
|||
from loguru import logger
|
||||
from pydantic import Field
|
||||
|
||||
from experiencemaker.tool.base_tool import BaseTool
|
||||
from experiencemaker.tool.base_tool import BaseTool, TOOL_REGISTRY
|
||||
|
||||
|
||||
class DashscopeSearchTool(BaseTool):
|
||||
|
|
@ -140,9 +140,13 @@ Extract the original content related to the user's question directly from the co
|
|||
else:
|
||||
return result
|
||||
|
||||
|
||||
TOOL_REGISTRY.register(DashscopeSearchTool, "web_search")
|
||||
|
||||
|
||||
def main():
|
||||
from experiencemaker.utils.test_key import set_key
|
||||
set_key()
|
||||
from experiencemaker.utils.util_function import load_env_keys
|
||||
load_env_keys()
|
||||
query = "What is artificial intelligence?"
|
||||
|
||||
tool = DashscopeSearchTool(stream_print=True)
|
||||
|
|
|
|||
|
|
@ -1,11 +1,10 @@
|
|||
|
||||
import asyncio
|
||||
from typing import List, Optional
|
||||
from typing import List
|
||||
|
||||
from loguru import logger
|
||||
from mcp import ClientSession
|
||||
from mcp.client.sse import sse_client
|
||||
from pydantic import Field
|
||||
from pydantic import Field, model_validator
|
||||
|
||||
from experiencemaker.tool.base_tool import BaseTool
|
||||
|
||||
|
|
@ -14,35 +13,11 @@ class MCPTool(BaseTool):
|
|||
server_url: str = Field(..., description="MCP server URL")
|
||||
tool_name_list: List[str] = Field(default_factory=list)
|
||||
cache_tools: dict = Field(default_factory=dict, alias="cache_tools")
|
||||
cache_tools_info: Optional[dict] = Field(default=None, alias="cache_tools_info")
|
||||
|
||||
class Config:
|
||||
underscore_attrs_are_private = True
|
||||
|
||||
def __init__(self, **data):
|
||||
super().__init__(**data)
|
||||
@model_validator(mode="after")
|
||||
def refresh_tools(self):
|
||||
self.refresh()
|
||||
|
||||
def get_tool_name_list(self) -> List[str]:
|
||||
return self.tool_name_list
|
||||
|
||||
def get_server_info(self):
|
||||
return self.cache_tools_info
|
||||
|
||||
def refresh(self):
|
||||
self.cache_tools.clear()
|
||||
self.tool_name_list.clear()
|
||||
|
||||
if "sse" in self.server_url:
|
||||
original_tool_list = asyncio.run(self._get_tools())
|
||||
self.cache_tools_info = original_tool_list.tools
|
||||
|
||||
for tool in self.cache_tools_info:
|
||||
self.cache_tools[tool.name] = tool
|
||||
self.tool_name_list.append(tool.name)
|
||||
else:
|
||||
# TODO: Implement non-SSE refresh logic
|
||||
logger.warning("Non-SSE refresh not implemented yet")
|
||||
return self
|
||||
|
||||
async def _get_tools(self):
|
||||
async with sse_client(url=self.server_url) as streams:
|
||||
|
|
@ -51,41 +26,51 @@ class MCPTool(BaseTool):
|
|||
tools = await session.list_tools()
|
||||
return tools
|
||||
|
||||
def input_schema(self, tool_name: str) -> dict:
|
||||
return self.cache_tools.get(tool_name, {}).inputSchema
|
||||
def refresh(self):
|
||||
self.tool_name_list.clear()
|
||||
self.cache_tools.clear()
|
||||
|
||||
def output_schema(self, tool_name: str) -> dict:
|
||||
# TODO: Implement output schema logic
|
||||
return {}
|
||||
if "sse" in self.server_url:
|
||||
original_tool_list = asyncio.run(self._get_tools())
|
||||
for tool in original_tool_list.tools:
|
||||
self.cache_tools[tool.name] = tool
|
||||
self.tool_name_list.append(tool.name)
|
||||
else:
|
||||
raise NotImplementedError("Non-SSE refresh not implemented yet")
|
||||
|
||||
@property
|
||||
def input_schema(self) -> dict:
|
||||
return {x: self.cache_tools[x].inputSchema for x in self.cache_tools}
|
||||
|
||||
@property
|
||||
def output_schema(self) -> dict:
|
||||
raise NotImplementedError("Output schema not implemented yet")
|
||||
|
||||
def get_tool_description(self, tool_name: str, schema: bool = False) -> str:
|
||||
if tool_name not in self.cache_tools:
|
||||
raise RuntimeError(f"Tool {tool_name} not found")
|
||||
|
||||
tool = self.cache_tools.get(tool_name)
|
||||
if not tool:
|
||||
return ""
|
||||
|
||||
description = f'tool \'{tool_name}\' description is:'+ tool.description
|
||||
description = f"tool={tool_name} description={tool.description}\n"
|
||||
if schema:
|
||||
description += f"\nInput Schema: {self.input_schema(tool_name)}"
|
||||
description += f"\nOutput Schema: {self.output_schema(tool_name)}"
|
||||
return description
|
||||
|
||||
async def _execute(self, **kwargs):
|
||||
tool_name = kwargs.get('tool_name')
|
||||
args = kwargs.get('args', {})
|
||||
description += f"input_schema={self.input_schema[tool_name]}\n" \
|
||||
f"output_schema={self.output_schema[tool_name]}\n"
|
||||
return description.strip()
|
||||
|
||||
async def async_execute(self, tool_name: str, **kwargs):
|
||||
if "sse" in self.server_url:
|
||||
async with sse_client(url=self.server_url) as streams:
|
||||
async with ClientSession(streams[0], streams[1]) as session:
|
||||
await session.initialize()
|
||||
results = await session.call_tool(tool_name, args)
|
||||
results = await session.call_tool(tool_name, kwargs)
|
||||
return results.content[0].text, results.isError
|
||||
else:
|
||||
return "Failed to connect to the tool", False
|
||||
|
||||
def execute(self, **kwargs):
|
||||
return asyncio.run(self._execute(**kwargs))
|
||||
else:
|
||||
raise NotImplementedError("Non-SSE execute not implemented yet")
|
||||
|
||||
def _execute(self, **kwargs):
|
||||
return asyncio.run(self.async_execute(**kwargs))
|
||||
|
||||
def get_cache_id(self, **kwargs) -> str:
|
||||
# Implement a method to generate a unique cache ID based on the input
|
||||
return f"{kwargs.get('tool_name')}_{hash(frozenset(kwargs.get('args', {}).items()))}"
|
||||
|
||||
|
|
|
|||
|
|
@ -1,4 +1,4 @@
|
|||
from experiencemaker.tool.base_tool import BaseTool
|
||||
from experiencemaker.tool.base_tool import BaseTool, TOOL_REGISTRY
|
||||
|
||||
|
||||
class TerminateTool(BaseTool):
|
||||
|
|
@ -21,3 +21,4 @@ class TerminateTool(BaseTool):
|
|||
return f"The interaction has been completed with status: {status}"
|
||||
|
||||
|
||||
TOOL_REGISTRY.register(TerminateTool, "terminate")
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -111,6 +111,8 @@ class HttpClient(BaseModel):
|
|||
retry_sleep_time *= self.retry_time_multiplier
|
||||
time.sleep(retry_sleep_time)
|
||||
|
||||
return None
|
||||
|
||||
def request_stream(self,
|
||||
data: str = None,
|
||||
json_data: dict = None,
|
||||
|
|
@ -137,7 +139,7 @@ class HttpClient(BaseModel):
|
|||
http_enum=http_enum,
|
||||
**kwargs)
|
||||
|
||||
return
|
||||
return None
|
||||
|
||||
except Exception as e:
|
||||
logger.exception(f"{self.__class__.__name__} {i}th request failed with args={e.args}")
|
||||
|
|
@ -150,3 +152,5 @@ class HttpClient(BaseModel):
|
|||
|
||||
retry_sleep_time *= self.retry_time_multiplier
|
||||
time.sleep(retry_sleep_time)
|
||||
|
||||
return None
|
||||
|
|
|
|||
|
|
@ -1,11 +0,0 @@
|
|||
from best_logger import register_logger
|
||||
|
||||
def init_logger():
|
||||
register_logger(
|
||||
mods=["agent", "context", "summary"],
|
||||
non_console_mods=[],
|
||||
auto_clean_mods=[],
|
||||
base_log_path=f"logs/default"
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -1,33 +0,0 @@
|
|||
import json
|
||||
import os
|
||||
|
||||
from experiencemaker.schema.module_loader import ModuleLoader
|
||||
|
||||
|
||||
def load_env_keys():
|
||||
if os.path.exists(".env"):
|
||||
with open(".env") as f:
|
||||
config = json.load(f)
|
||||
for k, v in config.items():
|
||||
os.environ[k] = v
|
||||
|
||||
|
||||
agent_wrapper_loader = ModuleLoader(
|
||||
class_path="experiencemaker.module.agent_wrapper.naive_agent_wrapper",
|
||||
class_name="NaiveAgentWrapper",
|
||||
config_path="beyondagent/config/agent_wrapper/naive_agent_wrapper.json")
|
||||
|
||||
context_generator_loader = ModuleLoader(
|
||||
class_path="experiencemaker.module.context_generator.simple_context_generator",
|
||||
class_name="SimpleContextGenerator",
|
||||
config_path="beyondagent/config/context_generator/simple_context_generator.json")
|
||||
|
||||
summarizer_loader = ModuleLoader(
|
||||
class_path="experiencemaker.module.summarizer.simple_summarizer",
|
||||
class_name="SimpleSummarizer",
|
||||
config_path="beyondagent/config/summarizer/simple_summarizer.json")
|
||||
|
||||
env_loader = ModuleLoader(
|
||||
class_path="experiencemaker.module.environment.simple_environment",
|
||||
class_name="SimpleEnvironment",
|
||||
config_path="beyondagent/config/environment/simple_environment.json")
|
||||
|
|
@ -1,18 +0,0 @@
|
|||
from typing import List
|
||||
|
||||
from experiencemaker.schema.trajectory import Message, StateMessage, ActionMessage
|
||||
|
||||
|
||||
def format_trajectory_steps(steps: List[Message]) -> str:
|
||||
format_steps = []
|
||||
step_idx = 0
|
||||
single_step = []
|
||||
for idx, step in enumerate(steps):
|
||||
if isinstance(step, ActionMessage):
|
||||
step_idx += 1
|
||||
single_step.append(f"** STEP {step_idx} **\n{step.content}")
|
||||
elif isinstance(step, StateMessage):
|
||||
single_step.append(f"{step.content}")
|
||||
format_steps.append("\n".join(single_step))
|
||||
single_step = []
|
||||
return "\n\n".join(format_steps)
|
||||
|
|
@ -1,5 +1,7 @@
|
|||
import json
|
||||
import os
|
||||
import re
|
||||
|
||||
from loguru import logger
|
||||
|
||||
def get_html_match_content(content: str, key: str):
|
||||
pattern = rf"<{key}>(.*?)</{key}>"
|
||||
|
|
@ -7,3 +9,13 @@ def get_html_match_content(content: str, key: str):
|
|||
if match:
|
||||
return match.group(1)
|
||||
return None
|
||||
|
||||
|
||||
def load_env_keys():
|
||||
if os.path.exists(".env"):
|
||||
with open(".env") as f:
|
||||
config = json.load(f)
|
||||
for k, v in config.items():
|
||||
os.environ[k] = v
|
||||
else:
|
||||
logger.warning(".env file not found~")
|
||||
Loading…
Add table
Reference in a new issue