mirror of
https://github.com/agentscope-ai/ReMe.git
synced 2026-09-06 08:16:00 +00:00
update registry
This commit is contained in:
parent
43778199da
commit
d21156db9d
21 changed files with 46 additions and 190 deletions
|
|
@ -57,106 +57,4 @@ python -m experiencemaker.em_service \
|
|||
```
|
||||
|
||||
|
||||
|
||||
|
||||
### Step2: Implement AgentWrapper
|
||||
|
||||
In order to utilize the **context generator** and **summarizer** capabilities of experiencemaker, 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 experiencemaker.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
|
||||
experiencemaker.
|
||||
|
||||
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 experiencemaker.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 experiencemaker.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,13 +1,2 @@
|
|||
from experiencemaker.model.base_embedding_model import BaseEmbeddingModel
|
||||
from experiencemaker.model.base_llm import BaseLLM
|
||||
from experiencemaker.model.openai_compatible_embedding_model import OpenAICompatibleEmbeddingModel
|
||||
from experiencemaker.model.openai_compatible_llm import OpenAICompatibleBaseLLM
|
||||
from experiencemaker.utils.registry import Registry
|
||||
|
||||
EMBEDDING_MODEL_REGISTRY = Registry[BaseEmbeddingModel]("embedding_model")
|
||||
|
||||
EMBEDDING_MODEL_REGISTRY.register(OpenAICompatibleEmbeddingModel, "openai_compatible")
|
||||
|
||||
LLM_REGISTRY = Registry[BaseLLM]("llm")
|
||||
|
||||
LLM_REGISTRY.register(OpenAICompatibleBaseLLM, "openai_compatible")
|
||||
|
|
|
|||
|
|
@ -5,7 +5,9 @@ from loguru import logger
|
|||
from pydantic import BaseModel, Field
|
||||
|
||||
from experiencemaker.schema.vector_store_node import VectorStoreNode
|
||||
from experiencemaker.utils.registry import Registry
|
||||
|
||||
EMBEDDING_MODEL_REGISTRY = Registry()
|
||||
|
||||
class BaseEmbeddingModel(BaseModel, ABC):
|
||||
model_name: str = Field(default=..., description="model name")
|
||||
|
|
|
|||
|
|
@ -6,7 +6,9 @@ from pydantic import Field, BaseModel
|
|||
|
||||
from experiencemaker.schema.trajectory import Message, ActionMessage
|
||||
from experiencemaker.tool.base_tool import BaseTool
|
||||
from experiencemaker.utils.registry import Registry
|
||||
|
||||
LLM_REGISTRY = Registry()
|
||||
|
||||
class BaseLLM(BaseModel, ABC):
|
||||
model_name: str = Field(...)
|
||||
|
|
|
|||
|
|
@ -4,9 +4,10 @@ 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
|
||||
|
||||
|
||||
@EMBEDDING_MODEL_REGISTRY.register("openai_compatible")
|
||||
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")
|
||||
|
|
|
|||
|
|
@ -7,11 +7,12 @@ 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
|
||||
|
||||
|
||||
@LLM_REGISTRY.register("openai_compatible")
|
||||
class OpenAICompatibleBaseLLM(BaseLLM):
|
||||
model_name: str = Field(default="qwen3-32b")
|
||||
api_key: str = Field(default_factory=lambda: os.getenv("OPENAI_API_KEY"), description="api key")
|
||||
|
|
|
|||
|
|
@ -1,7 +1 @@
|
|||
from experiencemaker.module.agent_wrapper.agent_wrapper_mixin import AgentWrapperMixin
|
||||
from experiencemaker.module.agent_wrapper.simple_agent_wrapper import SimpleAgentWrapper
|
||||
from experiencemaker.utils.registry import Registry
|
||||
|
||||
AGENT_WRAPPER_REGISTRY = Registry[AgentWrapperMixin]("agent_wrapper")
|
||||
|
||||
AGENT_WRAPPER_REGISTRY.register(SimpleAgentWrapper, "simple")
|
||||
|
|
|
|||
|
|
@ -5,7 +5,9 @@ from pydantic import Field, BaseModel
|
|||
from experiencemaker.model.base_llm import BaseLLM
|
||||
from experiencemaker.module.context_generator.base_context_generator import BaseContextGenerator
|
||||
from experiencemaker.schema.trajectory import Trajectory
|
||||
from experiencemaker.utils.registry import Registry
|
||||
|
||||
AGENT_WRAPPER_REGISTRY = Registry()
|
||||
|
||||
class AgentWrapperMixin(BaseModel, ABC):
|
||||
context_generator: BaseContextGenerator | None = Field(default=None)
|
||||
|
|
|
|||
|
|
@ -1,8 +1,9 @@
|
|||
from experiencemaker.module.agent_wrapper.agent_wrapper_mixin import AgentWrapperMixin
|
||||
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
|
||||
|
||||
|
||||
@AGENT_WRAPPER_REGISTRY.register(name="simple")
|
||||
class SimpleAgentWrapper(SimpleAgent, AgentWrapperMixin):
|
||||
|
||||
def execute(self, query: str, workspace_id: str = None, **kwargs) -> Trajectory:
|
||||
|
|
|
|||
|
|
@ -1,7 +1 @@
|
|||
from experiencemaker.module.context_generator.base_context_generator import BaseContextGenerator
|
||||
from experiencemaker.module.context_generator.simple_context_generator import SimpleContextGenerator
|
||||
from experiencemaker.utils.registry import Registry
|
||||
|
||||
CONTEXT_GENERATOR_REGISTRY = Registry[BaseContextGenerator]("context_generator")
|
||||
|
||||
CONTEXT_GENERATOR_REGISTRY.register(SimpleContextGenerator, "simple")
|
||||
|
|
|
|||
|
|
@ -7,7 +7,9 @@ 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
|
||||
|
||||
CONTEXT_GENERATOR_REGISTRY = Registry()
|
||||
|
||||
class BaseContextGenerator(BaseModel, ABC):
|
||||
vector_store: BaseVectorStore | None = Field(default=None)
|
||||
|
|
|
|||
|
|
@ -1,11 +1,13 @@
|
|||
from typing import List
|
||||
|
||||
from experiencemaker.module.context_generator.base_context_generator import BaseContextGenerator
|
||||
from experiencemaker.module.context_generator.base_context_generator import BaseContextGenerator, \
|
||||
CONTEXT_GENERATOR_REGISTRY
|
||||
from experiencemaker.schema.experience import Experience
|
||||
from experiencemaker.schema.trajectory import Trajectory, ContextMessage
|
||||
from experiencemaker.schema.vector_store_node import VectorStoreNode
|
||||
|
||||
|
||||
@CONTEXT_GENERATOR_REGISTRY.register(name="simple")
|
||||
class SimpleContextGenerator(BaseContextGenerator):
|
||||
|
||||
def _build_retrieve_query(self, trajectory: Trajectory, **kwargs) -> str:
|
||||
|
|
|
|||
|
|
@ -1,6 +1 @@
|
|||
from experiencemaker.module.summarizer.base_summarizer import BaseSummarizer
|
||||
from experiencemaker.module.summarizer.simple_summarizer import SimpleSummarizer
|
||||
from experiencemaker.utils.registry import Registry
|
||||
|
||||
SUMMARIZER_REGISTRY = Registry[BaseSummarizer]("summarizer")
|
||||
SUMMARIZER_REGISTRY.register(SimpleSummarizer, "simple")
|
||||
|
|
|
|||
|
|
@ -8,7 +8,9 @@ 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
|
||||
|
||||
SUMMARIZER_REGISTRY = Registry()
|
||||
|
||||
class BaseSummarizer(BaseModel, ABC):
|
||||
vector_store: BaseVectorStore | None = Field(default=None)
|
||||
|
|
|
|||
|
|
@ -6,12 +6,13 @@ 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
|
||||
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
|
||||
|
||||
|
||||
@SUMMARIZER_REGISTRY.register("simple")
|
||||
class SimpleSummarizer(BaseSummarizer, PromptMixin):
|
||||
max_retries: int = Field(default=5, description="max retries")
|
||||
prompt_file_path: Path = Path(__file__).parent / "simple_summarizer_prompt.yaml"
|
||||
|
|
|
|||
|
|
@ -1,8 +1,2 @@
|
|||
from experiencemaker.storage.base_vector_store import BaseVectorStore
|
||||
from experiencemaker.storage.es_vector_store import EsVectorStore
|
||||
from experiencemaker.storage.file_vector_store import FileVectorStore
|
||||
from experiencemaker.utils.registry import Registry
|
||||
|
||||
VECTOR_STORE_REGISTRY = Registry[BaseVectorStore]("vector_store")
|
||||
VECTOR_STORE_REGISTRY.register(EsVectorStore, "elasticsearch")
|
||||
VECTOR_STORE_REGISTRY.register(FileVectorStore, "local_file")
|
||||
|
|
|
|||
|
|
@ -5,7 +5,9 @@ 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
|
||||
|
||||
VECTOR_STORE_REGISTRY = Registry()
|
||||
|
||||
class BaseVectorStore(BaseModel, ABC):
|
||||
embedding_model: BaseEmbeddingModel = Field(default=...)
|
||||
|
|
|
|||
|
|
@ -8,9 +8,10 @@ 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
|
||||
|
||||
|
||||
@VECTOR_STORE_REGISTRY.register("elasticsearch")
|
||||
class EsVectorStore(BaseVectorStore):
|
||||
hosts: str | List[str] = Field(default_factory=lambda: os.getenv("ES_HOSTS", "http://localhost:9200"))
|
||||
basic_auth: str | Tuple[str, str] | None = Field(default=None)
|
||||
|
|
|
|||
|
|
@ -9,9 +9,10 @@ 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
|
||||
|
||||
|
||||
@VECTOR_STORE_REGISTRY.register("local_file")
|
||||
class FileVectorStore(BaseVectorStore):
|
||||
store_dir: str = Field(default="./file_vector_store")
|
||||
index_path: Path | None = Field(default=None)
|
||||
|
|
|
|||
|
|
@ -1,12 +0,0 @@
|
|||
from experiencemaker.tool.base_tool import BaseTool
|
||||
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
|
||||
|
||||
TOOL_REGISTRY = Registry[BaseTool]("tool")
|
||||
|
||||
TOOL_REGISTRY.register(CodeTool, "code")
|
||||
TOOL_REGISTRY.register(DashscopeSearchTool, "web_search")
|
||||
TOOL_REGISTRY.register(TerminateTool, "terminate")
|
||||
|
|
@ -1,43 +1,27 @@
|
|||
from typing import Dict, List
|
||||
from typing import List
|
||||
|
||||
from typing import TypeVar, Generic
|
||||
|
||||
T = TypeVar('T')
|
||||
from loguru import logger
|
||||
|
||||
|
||||
class Registry(Generic[T]):
|
||||
class Registry(object):
|
||||
def __init__(self):
|
||||
self._registry = {}
|
||||
|
||||
def __init__(self, name: str):
|
||||
self.name: str = name
|
||||
self.module_dict: Dict[str, T] = {}
|
||||
def register(self, name: str = None):
|
||||
|
||||
@property
|
||||
def registered_module_names(self) -> List[str]:
|
||||
return sorted(self.module_dict.keys())
|
||||
def decorator(cls):
|
||||
class_name = name if name is not None else cls.__name__
|
||||
if class_name in self._registry:
|
||||
logger.warning(f"name={class_name} is already registered, will be overwritten.")
|
||||
self._registry[class_name] = cls
|
||||
return cls
|
||||
|
||||
def register(self, module: T, module_name: str = None):
|
||||
if module_name is None:
|
||||
module_name = module.__name__
|
||||
return decorator
|
||||
|
||||
if module_name in self.module_dict:
|
||||
raise KeyError(f'{module_name} is already registered in {self.name}')
|
||||
def __getitem__(self, name: str):
|
||||
if name not in self._registry:
|
||||
raise KeyError(f"name={name} is not registered!")
|
||||
return self._registry[name]
|
||||
|
||||
self.module_dict[module_name] = module
|
||||
|
||||
def batch_register(self, modules: List[T] | Dict[str, T]):
|
||||
if isinstance(modules, list):
|
||||
module_name_dict = {m.__name__: m for m in modules}
|
||||
|
||||
elif isinstance(modules, dict):
|
||||
module_name_dict = modules
|
||||
|
||||
else:
|
||||
raise NotImplementedError("Input must be a list or a dictionary.")
|
||||
self.module_dict.update(module_name_dict)
|
||||
|
||||
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]
|
||||
|
||||
def __contains__(self, module_name: str):
|
||||
return module_name in self.module_dict
|
||||
def list_all(self) -> List[str]:
|
||||
return sorted(self._registry.keys())
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue