mirror of
https://github.com/agentscope-ai/ReMe.git
synced 2026-08-28 05:25:04 +00:00
add unit test
This commit is contained in:
parent
15da7cc730
commit
79a247e114
30 changed files with 349 additions and 347 deletions
2
.gitignore
vendored
2
.gitignore
vendored
|
|
@ -17,4 +17,4 @@ log/
|
|||
.trash/
|
||||
runs
|
||||
logs
|
||||
|
||||
rag_nodes_index.jsonl
|
||||
|
|
|
|||
|
|
@ -1,5 +0,0 @@
|
|||
from experiencemaker.config.config_handler import ConfigHandler
|
||||
|
||||
agent_wrapper_config = ConfigHandler(module_name="agent_wrapper")
|
||||
summarizer_config = ConfigHandler(module_name="summarizer")
|
||||
context_generator_config = ConfigHandler(module_name="context_generator")
|
||||
|
|
@ -1 +0,0 @@
|
|||
{}
|
||||
|
|
@ -1,32 +0,0 @@
|
|||
from pathlib import Path
|
||||
|
||||
from pydantic import BaseModel, Field, model_validator
|
||||
|
||||
from experiencemaker.utils.file_handler import FileHandler
|
||||
|
||||
|
||||
class ConfigHandler(BaseModel):
|
||||
module_name: str = Field(default=...)
|
||||
config_dict: dict = Field(default_factory=dict)
|
||||
|
||||
@model_validator(mode="after")
|
||||
def register_config(self):
|
||||
module_config_path: Path = Path(__file__).parent / self.module_name
|
||||
for config_path in module_config_path.iterdir():
|
||||
config_name = config_path.stem
|
||||
config = FileHandler(file_path=config_path).load()
|
||||
self.config_dict[config_name] = config
|
||||
return self
|
||||
|
||||
def list_config_names(self):
|
||||
return list(self.config_dict.keys())
|
||||
|
||||
def __getattr__(self, item):
|
||||
if item in self.config_dict:
|
||||
return self.config_dict[item]
|
||||
return super().__getattr__(item)
|
||||
|
||||
def __getitem__(self, item):
|
||||
if item in self.config_dict:
|
||||
return self.config_dict[item]
|
||||
return super().__getitem__(item)
|
||||
|
|
@ -1 +0,0 @@
|
|||
{}
|
||||
|
|
@ -1 +0,0 @@
|
|||
{"a": 1}
|
||||
|
|
@ -10,7 +10,7 @@ from experiencemaker.model.base_embedding_model import BaseEmbeddingModel, EMBED
|
|||
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()
|
||||
|
|
@ -74,3 +74,20 @@ class OpenAICompatibleEmbeddingModel(BaseEmbeddingModel):
|
|||
|
||||
|
||||
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
|
||||
|
|
|
|||
|
|
@ -153,3 +153,25 @@ class OpenAICompatibleBaseLLM(BaseLLM):
|
|||
|
||||
|
||||
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
|
||||
|
|
|
|||
|
|
@ -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,15 +1,15 @@
|
|||
from abc import ABC
|
||||
from typing import List
|
||||
|
||||
from pydantic import Field
|
||||
from pydantic import Field, BaseModel
|
||||
|
||||
from experiencemaker.module.base_module import BaseModule
|
||||
from experiencemaker.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 = Field(default_factory=StateMessage)
|
||||
|
|
@ -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}
|
||||
|
|
@ -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>
|
||||
|
|
@ -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
|
||||
|
|
@ -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
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
@ -4,17 +4,15 @@ import uvicorn
|
|||
from fastapi import FastAPI
|
||||
from loguru import logger
|
||||
|
||||
from experiencemaker.model import LLM_REGISTRY, EMBEDDING_MODEL_REGISTRY
|
||||
from experiencemaker.model.base_embedding_model import BaseEmbeddingModel
|
||||
from experiencemaker.model.base_llm import BaseLLM
|
||||
from experiencemaker.model.base_embedding_model import BaseEmbeddingModel, EMBEDDING_MODEL_REGISTRY
|
||||
from experiencemaker.model.base_llm import BaseLLM, LLM_REGISTRY
|
||||
from experiencemaker.module.agent_wrapper.base_agent_wrapper import BaseAgentWrapper
|
||||
from experiencemaker.module.context_generator.base_context_generator import BaseContextGenerator
|
||||
from experiencemaker.module.summarizer.base_summarizer import BaseSummarizer
|
||||
from experiencemaker.schema.request import AgentWrapperRequest, ContextGeneratorRequest, SummarizerRequest
|
||||
from experiencemaker.schema.response import AgentWrapperResponse, ContextGeneratorResponse, SummarizerResponse
|
||||
from experiencemaker.schema.trajectory import ContextMessage, Trajectory, Sample
|
||||
from experiencemaker.storage import VECTOR_STORE_REGISTRY
|
||||
from experiencemaker.storage.base_vector_store import BaseVectorStore
|
||||
from experiencemaker.storage.base_vector_store import BaseVectorStore, VECTOR_STORE_REGISTRY
|
||||
|
||||
app = FastAPI()
|
||||
from pydantic import BaseModel, Field, model_validator
|
||||
|
|
@ -29,7 +27,6 @@ class ExperienceMakerHttpService(BaseModel):
|
|||
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)
|
||||
|
|
@ -37,7 +34,6 @@ class ExperienceMakerHttpService(BaseModel):
|
|||
llm: BaseLLM | None = Field(default=None)
|
||||
embedding_model: BaseEmbeddingModel | None = Field(default=None)
|
||||
vector_store: BaseVectorStore | None = Field(default=None)
|
||||
|
||||
agent_wrapper: BaseAgentWrapper | None = Field(default=None)
|
||||
context_generator: BaseContextGenerator | None = Field(default=None)
|
||||
summarizer: BaseSummarizer | None = Field(default=None)
|
||||
|
|
|
|||
|
|
@ -6,6 +6,7 @@ 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, VECTOR_STORE_REGISTRY
|
||||
|
||||
|
|
@ -170,3 +171,74 @@ class EsVectorStore(BaseVectorStore):
|
|||
|
||||
|
||||
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,6 +7,7 @@ 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, VECTOR_STORE_REGISTRY
|
||||
|
||||
|
|
@ -133,3 +134,58 @@ class FileVectorStore(BaseVectorStore):
|
|||
|
||||
|
||||
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
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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,72 +0,0 @@
|
|||
import os
|
||||
|
||||
import yaml
|
||||
from loguru import logger
|
||||
from pydantic import BaseModel, Field, model_validator
|
||||
|
||||
|
||||
class PromptHandler(BaseModel):
|
||||
file_path: str = Field(default="")
|
||||
prompt_dict: dict = Field(default_factory=dict)
|
||||
|
||||
@model_validator(mode="after")
|
||||
def init_prompt_dict(self):
|
||||
if self.prompt_dict:
|
||||
logger.info(f"load prompt from dict, keys={self.prompt_dict.keys()}")
|
||||
|
||||
if self.file_path:
|
||||
self.load_file_prompt()
|
||||
logger.info(f"load prompt from file_path, keys={self.prompt_dict.keys()}")
|
||||
|
||||
def load_file_prompt(self):
|
||||
if not os.path.exists(self.file_path):
|
||||
raise RuntimeError(f"file_path={self.file_path} not exists!")
|
||||
|
||||
with open(self.file_path) as f:
|
||||
prompt_dict: dict = yaml.load(f, yaml.FullLoader)
|
||||
self.prompt_dict.update(prompt_dict)
|
||||
|
||||
def __getitem__(self, key: str):
|
||||
return self.prompt_dict[key]
|
||||
|
||||
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)
|
||||
|
||||
def prompt_format(self, prompt_name: str, **kwargs):
|
||||
prompt = self.prompt_dict[prompt_name]
|
||||
|
||||
flag_kwargs = {k: v for k, v in kwargs.items() if isinstance(v, bool)}
|
||||
other_kwargs = {k: v for k, v in kwargs.items() if not isinstance(v, bool)}
|
||||
|
||||
if flag_kwargs:
|
||||
split_prompt = []
|
||||
for line in prompt.strip().split("\n"):
|
||||
hit = False
|
||||
hit_flag = True
|
||||
for key, flag in kwargs.items():
|
||||
if not line.startswith(f"[{key}]"):
|
||||
continue
|
||||
|
||||
else:
|
||||
hit = True
|
||||
hit_flag = flag
|
||||
line = line.strip(f"[{key}]")
|
||||
break
|
||||
|
||||
if not hit:
|
||||
split_prompt.append(line)
|
||||
elif hit_flag:
|
||||
split_prompt.append(line)
|
||||
|
||||
prompt = "\n".join(split_prompt)
|
||||
|
||||
if other_kwargs:
|
||||
prompt = prompt.format(**other_kwargs)
|
||||
|
||||
return prompt
|
||||
|
|
@ -1,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