mirror of
https://github.com/agentscope-ai/ReMe.git
synced 2026-10-09 03:20:54 +00:00
refactor(reme_ai): update app and ops for async execution and new LLM features
- Update app.py to use async service initialization - Refactor multiple ops to use async_execute instead of execute - Add support for stream and use_async flags in config - Update LLM usage to use achat instead of chat - Add new LLM models and update existing ones in config - Improve error handling and logging in several ops - Update dependencies and Python version requirements
This commit is contained in:
parent
1e5250bbfb
commit
49c73b15b3
20 changed files with 135 additions and 63 deletions
|
|
@ -13,7 +13,7 @@ authors = [
|
|||
]
|
||||
license = { file = "LICENSE" }
|
||||
readme = "README.md"
|
||||
requires-python = ">=3.12"
|
||||
requires-python = ">=3.11"
|
||||
|
||||
classifiers = [
|
||||
"Programming Language :: Python :: 3",
|
||||
|
|
@ -24,7 +24,7 @@ classifiers = [
|
|||
keywords = ["llm", "memory", "experience", "memoryscope", "ai", "mcp", "http"]
|
||||
|
||||
dependencies = [
|
||||
"flowllm==0.1.3",
|
||||
"flowllm==0.1.5",
|
||||
]
|
||||
|
||||
[tool.setuptools.packages.find]
|
||||
|
|
|
|||
|
|
@ -7,7 +7,7 @@ from reme_ai.config.config_parser import ConfigParser
|
|||
|
||||
def main():
|
||||
with BaseService.get_service(*sys.argv[1:], parser=ConfigParser) as service:
|
||||
service()
|
||||
service(logo="ReMe")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
|
|
|
|||
|
|
@ -1,4 +1,3 @@
|
|||
# default config.yaml
|
||||
backend: http
|
||||
language: ""
|
||||
thread_pool_max_workers: 32
|
||||
|
|
@ -18,6 +17,9 @@ http:
|
|||
flow:
|
||||
retrieve_task_memory:
|
||||
flow_content: build_query_op >> recall_vector_store_op >> rerank_memory_op >> rewrite_memory_op
|
||||
stream: false
|
||||
use_async: true
|
||||
service_type: http+mcp
|
||||
description: "Retrieves the most relevant top-k memory experiences from historical data based on the current query to enhance task-solving capabilities"
|
||||
input_schema:
|
||||
query:
|
||||
|
|
@ -27,6 +29,9 @@ flow:
|
|||
|
||||
summary_task_memory:
|
||||
flow_content: trajectory_preprocess_op >> (success_extraction_op|failure_extraction_op|comparative_extraction_op) >> memory_validation_op >> update_vector_store_op
|
||||
stream: false
|
||||
use_async: true
|
||||
service_type: http+mcp
|
||||
description: "Summarizes conversation trajectories or messages into structured memory representations for long-term storage"
|
||||
input_schema:
|
||||
trajectories:
|
||||
|
|
@ -36,6 +41,9 @@ flow:
|
|||
|
||||
retrieve_personal_memory:
|
||||
flow_content: set_query_op >> (extract_time_op | (retrieve_memory_op >> semantic_rank_op)) >> fuse_rerank_op
|
||||
stream: false
|
||||
use_async: true
|
||||
service_type: http+mcp
|
||||
description: "Retrieves the most relevant personal memories from historical data based on the query to enhance response quality"
|
||||
input_schema:
|
||||
query:
|
||||
|
|
@ -45,6 +53,9 @@ flow:
|
|||
|
||||
summary_personal_memory:
|
||||
flow_content: info_filter_op >> (get_observation_op | get_observation_with_time_op | load_today_memory_op) >> contra_repeat_op >> update_vector_store_op
|
||||
stream: false
|
||||
use_async: true
|
||||
service_type: http+mcp
|
||||
description: "Consolidates user observations and memories by filtering information and removing redundancies for efficient storage"
|
||||
input_schema:
|
||||
trajectories:
|
||||
|
|
@ -54,6 +65,9 @@ flow:
|
|||
|
||||
retrieve_task_memory_simple:
|
||||
flow_content: build_query_op >> recall_vector_store_op >> merge_memory_op
|
||||
stream: false
|
||||
use_async: true
|
||||
service_type: http+mcp
|
||||
description: "Retrieves the most relevant top-k memory experiences from historical data based on the current query with simplified processing"
|
||||
input_schema:
|
||||
query:
|
||||
|
|
@ -63,6 +77,9 @@ flow:
|
|||
|
||||
summary_task_memory_simple:
|
||||
flow_content: simple_summary_op >> update_vector_store_op
|
||||
stream: false
|
||||
use_async: true
|
||||
service_type: http+mcp
|
||||
description: "Summarizes conversation trajectories or messages into memories using a simplified approach"
|
||||
input_schema:
|
||||
trajectories:
|
||||
|
|
@ -72,6 +89,9 @@ flow:
|
|||
|
||||
vector_store:
|
||||
flow_content: vector_store_action_op
|
||||
stream: false
|
||||
use_async: false
|
||||
service_type: http+mcp
|
||||
description: "Directly operates on the vector store with various management actions"
|
||||
input_schema:
|
||||
action:
|
||||
|
|
@ -142,11 +162,18 @@ op:
|
|||
llm:
|
||||
default:
|
||||
backend: openai_compatible
|
||||
# model_name: qwen3-30b-a3b-thinking-2507
|
||||
model_name: qwen3-30b-a3b-instruct-2507
|
||||
params:
|
||||
temperature: 0.6
|
||||
|
||||
qwen3_30b_instruct:
|
||||
backend: openai_compatible
|
||||
model_name: qwen3-30b-a3b-instruct-2507
|
||||
|
||||
qwen3_30b_thinking:
|
||||
backend: openai_compatible
|
||||
model_name: qwen3-30b-a3b-thinking-2507
|
||||
|
||||
embedding_model:
|
||||
default:
|
||||
backend: openai_compatible
|
||||
|
|
|
|||
|
|
@ -1,21 +1,24 @@
|
|||
import asyncio
|
||||
|
||||
from flowllm import C
|
||||
from flowllm.context.flow_context import FlowContext
|
||||
from flowllm.op.agent.react_v2_op import ReactV2Op
|
||||
from flowllm.op.llm.react_llm_op import ReactLLMOp
|
||||
|
||||
|
||||
@C.register_op()
|
||||
class SimpleReactOp(ReactV2Op):
|
||||
class SimpleReactOp(ReactLLMOp):
|
||||
...
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
async def main():
|
||||
from reme_ai.config.config_parser import ConfigParser
|
||||
|
||||
C.set_default_service_config(parser=ConfigParser).init_by_service_config()
|
||||
C.set_service_config(parser=ConfigParser, config_name="config=default").init_by_service_config()
|
||||
context = FlowContext(query="茅台和五粮现在股价多少?")
|
||||
|
||||
op = SimpleReactOp()
|
||||
op(context=context)
|
||||
# from reme_ai.schema import Message
|
||||
# result = op.llm.chat(messages=[Message(**{"role": "user", "content": "你叫什么名字?"})])
|
||||
# print("!!!", result)
|
||||
await op.async_call(context=context)
|
||||
print(context.response.answer)
|
||||
|
||||
if __name__ == "__main__":
|
||||
asyncio.run(main())
|
||||
|
|
|
|||
|
|
@ -23,10 +23,9 @@ class ExtractTimeOp(BaseLLMOp):
|
|||
"""
|
||||
|
||||
def get_language_value(self, value_dict: dict):
|
||||
|
||||
return value_dict.get(self.language, value_dict.get("en"))
|
||||
|
||||
def execute(self):
|
||||
async def async_execute(self):
|
||||
"""
|
||||
Executes the primary logic of identifying and extracting time data from an LLM's response.
|
||||
|
||||
|
|
@ -59,7 +58,7 @@ class ExtractTimeOp(BaseLLMOp):
|
|||
logger.info(f"Extracting time from query: {query[:100]}...")
|
||||
|
||||
# Invoke the LLM to generate a response
|
||||
response = self.llm.chat([Message(role=Role.USER, content=full_prompt)])
|
||||
response = await self.llm.achat([Message(role=Role.USER, content=full_prompt)])
|
||||
|
||||
# Handle empty or unsuccessful responses
|
||||
if not response or not response.content:
|
||||
|
|
|
|||
|
|
@ -45,7 +45,7 @@ class FuseRerankOp(BaseLLMOp):
|
|||
memory.metadata["match_msg_flag"] = str(int(match_msg_flag))
|
||||
return match_event_flag, match_msg_flag
|
||||
|
||||
def execute(self):
|
||||
async def async_execute(self):
|
||||
"""
|
||||
Executes the reranking process on memories considering their scores, types, and temporal relevance.
|
||||
|
||||
|
|
|
|||
|
|
@ -13,7 +13,7 @@ class PrintMemoryOp(BaseOp):
|
|||
"""
|
||||
file_path: str = __file__
|
||||
|
||||
def execute(self):
|
||||
async def async_execute(self):
|
||||
"""
|
||||
Executes the primary function, it involves:
|
||||
1. Fetches the memories.
|
||||
|
|
|
|||
|
|
@ -12,7 +12,7 @@ class ReadMessageOp(BaseOp):
|
|||
"""
|
||||
file_path: str = __file__
|
||||
|
||||
def execute(self):
|
||||
async def async_execute(self):
|
||||
"""
|
||||
Executes the primary function to fetch unmemorized chat messages.
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -1,5 +1,10 @@
|
|||
from flowllm import C
|
||||
from typing import List
|
||||
|
||||
from flowllm import C
|
||||
from flowllm.schema.vector_node import VectorNode
|
||||
from loguru import logger
|
||||
|
||||
from reme_ai.schema.memory import BaseMemory, vector_node_to_memory
|
||||
from reme_ai.vector_store import RecallVectorStoreOp
|
||||
|
||||
|
||||
|
|
@ -10,4 +15,28 @@ class RetrieveMemoryOp(RecallVectorStoreOp):
|
|||
Processes these memories concurrently, sorts them by similarity, and logs the activity,
|
||||
facilitating efficient memory retrieval operations within a given scope.
|
||||
"""
|
||||
file_path: str = __file__
|
||||
|
||||
async def async_execute(self):
|
||||
recall_key: str = self.op_params.get("recall_key", "query")
|
||||
top_k: int = self.op_params.get("top_k", 3)
|
||||
|
||||
query: str = self.context[recall_key]
|
||||
assert query, "query should be not empty!"
|
||||
|
||||
workspace_id: str = self.context.workspace_id
|
||||
nodes: List[VectorNode] = self.vector_store.search(query=query, workspace_id=workspace_id, top_k=top_k)
|
||||
memory_list: List[BaseMemory] = []
|
||||
memory_content_list: List[str] = []
|
||||
for node in nodes:
|
||||
memory: BaseMemory = vector_node_to_memory(node)
|
||||
if memory.content not in memory_content_list:
|
||||
memory_list.append(memory)
|
||||
memory_content_list.append(memory.content)
|
||||
logger.info(f"retrieve memory.size={len(memory_list)}")
|
||||
|
||||
threshold_score: float | None = self.op_params.get("threshold_score", None)
|
||||
if threshold_score is not None:
|
||||
memory_list = [mem for mem in memory_list if mem.score >= threshold_score or mem.score is None]
|
||||
logger.info(f"after filter by threshold_score size={len(memory_list)}")
|
||||
|
||||
self.context.response.metadata["memory_list"] = memory_list
|
||||
|
|
|
|||
|
|
@ -19,7 +19,7 @@ class SemanticRankOp(BaseLLMOp):
|
|||
"""
|
||||
file_path: str = __file__
|
||||
|
||||
def execute(self):
|
||||
async def async_execute(self):
|
||||
"""
|
||||
Executes the primary workflow of the SemanticRankOp which includes:
|
||||
- Retrieves query and memory list from context.
|
||||
|
|
@ -56,7 +56,7 @@ class SemanticRankOp(BaseLLMOp):
|
|||
logger.info(f"After deduplication: {len(memory_list)} memories")
|
||||
|
||||
# Perform semantic ranking using LLM
|
||||
ranked_memories = self._semantic_rank_memories(query, memory_list)
|
||||
ranked_memories = await self._semantic_rank_memories(query, memory_list)
|
||||
if ranked_memories:
|
||||
memory_list = ranked_memories
|
||||
|
||||
|
|
@ -71,7 +71,7 @@ class SemanticRankOp(BaseLLMOp):
|
|||
# Save ranked memories back to context
|
||||
self.context.response.metadata["memory_list"] = memory_list
|
||||
|
||||
def _semantic_rank_memories(self, query: str, memories: List[BaseMemory]) -> List[BaseMemory]:
|
||||
async def _semantic_rank_memories(self, query: str, memories: List[BaseMemory]) -> List[BaseMemory]:
|
||||
"""
|
||||
Use LLM to semantically rank memories based on relevance to the query
|
||||
"""
|
||||
|
|
@ -93,7 +93,7 @@ Memories:
|
|||
Please respond in JSON format:
|
||||
{{"rankings": [{{"index": 0, "score": 0.8}}, {{"index": 1, "score": 0.6}}, ...]}}"""
|
||||
|
||||
response = self.llm.chat([Message(role=Role.USER, content=prompt)])
|
||||
response = await self.llm.achat([Message(role=Role.USER, content=prompt)])
|
||||
|
||||
if not response or not response.content:
|
||||
logger.warning("LLM ranking failed, using original order")
|
||||
|
|
|
|||
|
|
@ -14,7 +14,7 @@ class SetQueryOp(BaseOp):
|
|||
into the context, utilizing either provided parameters or details from the context.
|
||||
"""
|
||||
|
||||
def execute(self):
|
||||
async def async_execute(self):
|
||||
"""
|
||||
Executes the operation's primary function, which involves determining the query and its
|
||||
timestamp, then storing these values within the context.
|
||||
|
|
|
|||
|
|
@ -1,6 +1,8 @@
|
|||
from typing import List, Tuple, Optional
|
||||
|
||||
from flowllm import C, BaseLLMOp
|
||||
from flowllm.enumeration.role import Role
|
||||
from flowllm.schema.message import Message as FlowMessage
|
||||
from loguru import logger
|
||||
|
||||
from reme_ai.schema import Message, Trajectory
|
||||
|
|
@ -12,7 +14,7 @@ from reme_ai.utils.op_utils import merge_messages_content, parse_json_experience
|
|||
class ComparativeExtractionOp(BaseLLMOp):
|
||||
file_path: str = __file__
|
||||
|
||||
def execute(self):
|
||||
async def async_execute(self):
|
||||
"""Extract comparative task memories by comparing different scoring trajectories"""
|
||||
all_trajectories: List[Trajectory] = self.context.get("all_trajectories", [])
|
||||
success_trajectories: List[Trajectory] = self.context.get("success_trajectories", [])
|
||||
|
|
@ -26,7 +28,7 @@ class ComparativeExtractionOp(BaseLLMOp):
|
|||
if highest_traj and lowest_traj and highest_traj.score > lowest_traj.score:
|
||||
logger.info(
|
||||
f"Extracting soft comparative task memories: highest ({highest_traj.score:.2f}) vs lowest ({lowest_traj.score:.2f})")
|
||||
soft_task_memories = self._extract_soft_comparative_task_memory(highest_traj, lowest_traj)
|
||||
soft_task_memories = await self._extract_soft_comparative_task_memory(highest_traj, lowest_traj)
|
||||
comparative_task_memories.extend(soft_task_memories)
|
||||
|
||||
# Hard comparison: success vs failure (if similarity search is enabled)
|
||||
|
|
@ -37,7 +39,7 @@ class ComparativeExtractionOp(BaseLLMOp):
|
|||
logger.info(f"Found {len(similar_pairs)} similar pairs for hard comparison")
|
||||
|
||||
for success_steps, failure_steps, similarity_score in similar_pairs:
|
||||
hard_task_memories = self._extract_hard_comparative_task_memory(success_steps, failure_steps,
|
||||
hard_task_memories = await self._extract_hard_comparative_task_memory(success_steps, failure_steps,
|
||||
similarity_score)
|
||||
comparative_task_memories.extend(hard_task_memories)
|
||||
|
||||
|
|
@ -73,7 +75,7 @@ class ComparativeExtractionOp(BaseLLMOp):
|
|||
"""Get trajectory score"""
|
||||
return trajectory.score
|
||||
|
||||
def _extract_soft_comparative_task_memory(self, higher_traj: Trajectory, lower_traj: Trajectory) -> List[
|
||||
async def _extract_soft_comparative_task_memory(self, higher_traj: Trajectory, lower_traj: Trajectory) -> List[
|
||||
BaseMemory]:
|
||||
"""Extract soft comparative task memory (high score vs low score)"""
|
||||
higher_steps = self._get_trajectory_steps(higher_traj)
|
||||
|
|
@ -105,9 +107,9 @@ class ComparativeExtractionOp(BaseLLMOp):
|
|||
|
||||
return task_memories
|
||||
|
||||
return self.llm.chat(messages=[Message(content=prompt)], callback_fn=parse_task_memories)
|
||||
return await self.llm.achat(messages=[FlowMessage(role=Role.USER, content=prompt)], callback_fn=parse_task_memories)
|
||||
|
||||
def _extract_hard_comparative_task_memory(self, success_steps: List[Message],
|
||||
async def _extract_hard_comparative_task_memory(self, success_steps: List[Message],
|
||||
failure_steps: List[Message], similarity_score: float) -> List[
|
||||
BaseMemory]:
|
||||
"""Extract hard comparative task memory (success vs failure)"""
|
||||
|
|
@ -134,7 +136,7 @@ class ComparativeExtractionOp(BaseLLMOp):
|
|||
|
||||
return task_memories
|
||||
|
||||
return self.llm.chat(messages=[Message(content=prompt)], callback_fn=parse_task_memories)
|
||||
return await self.llm.achat(messages=[FlowMessage(role=Role.USER, content=prompt)], callback_fn=parse_task_memories)
|
||||
|
||||
@staticmethod
|
||||
def _get_trajectory_steps(trajectory: Trajectory) -> List[Message]:
|
||||
|
|
|
|||
|
|
@ -1,6 +1,8 @@
|
|||
from typing import List
|
||||
|
||||
from flowllm import C, BaseLLMOp
|
||||
from flowllm.enumeration.role import Role
|
||||
from flowllm.schema.message import Message as FlowMessage
|
||||
from loguru import logger
|
||||
|
||||
from reme_ai.schema import Message, Trajectory
|
||||
|
|
@ -12,7 +14,7 @@ from reme_ai.utils.op_utils import merge_messages_content, parse_json_experience
|
|||
class FailureExtractionOp(BaseLLMOp):
|
||||
file_path: str = __file__
|
||||
|
||||
def execute(self):
|
||||
async def async_execute(self):
|
||||
"""Extract task memories from failed trajectories"""
|
||||
failure_trajectories: List[Trajectory] = self.context.get("failure_trajectories", [])
|
||||
|
||||
|
|
@ -29,11 +31,11 @@ class FailureExtractionOp(BaseLLMOp):
|
|||
if hasattr(trajectory, 'segments') and trajectory.segments:
|
||||
# Process segmented step sequences
|
||||
for segment in trajectory.segments:
|
||||
task_memories = self._extract_failure_task_memory_from_steps(segment, trajectory)
|
||||
task_memories = await self._extract_failure_task_memory_from_steps(segment, trajectory)
|
||||
failure_task_memories.extend(task_memories)
|
||||
else:
|
||||
# Process entire trajectory
|
||||
task_memories = self._extract_failure_task_memory_from_steps(trajectory.messages, trajectory)
|
||||
task_memories = await self._extract_failure_task_memory_from_steps(trajectory.messages, trajectory)
|
||||
failure_task_memories.extend(task_memories)
|
||||
|
||||
logger.info(f"Extracted {len(failure_task_memories)} failure task memories")
|
||||
|
|
@ -41,7 +43,7 @@ class FailureExtractionOp(BaseLLMOp):
|
|||
# Add task memories to context
|
||||
self.context.failure_task_memories = failure_task_memories
|
||||
|
||||
def _extract_failure_task_memory_from_steps(self, steps: List[Message], trajectory: Trajectory) -> List[BaseMemory]:
|
||||
async def _extract_failure_task_memory_from_steps(self, steps: List[Message], trajectory: Trajectory) -> List[BaseMemory]:
|
||||
"""Extract task memory from failed step sequences"""
|
||||
step_content = merge_messages_content(steps)
|
||||
context = get_trajectory_context(trajectory, steps)
|
||||
|
|
@ -70,4 +72,4 @@ class FailureExtractionOp(BaseLLMOp):
|
|||
|
||||
return task_memories
|
||||
|
||||
return self.llm.chat(messages=[Message(content=prompt)], callback_fn=parse_task_memories)
|
||||
return await self.llm.achat(messages=[FlowMessage(role=Role.USER, content=prompt)], callback_fn=parse_task_memories)
|
||||
|
|
|
|||
|
|
@ -10,7 +10,7 @@ from reme_ai.schema.memory import BaseMemory
|
|||
class MemoryDeduplicationOp(BaseOp):
|
||||
file_path: str = __file__
|
||||
|
||||
def execute(self):
|
||||
async def async_execute(self):
|
||||
"""Remove duplicate task memories"""
|
||||
# Get task memories to deduplicate
|
||||
task_memories: List[BaseMemory] = self.context.memory_list
|
||||
|
|
|
|||
|
|
@ -3,6 +3,8 @@ import re
|
|||
from typing import List, Dict, Any
|
||||
|
||||
from flowllm import C, BaseLLMOp
|
||||
from flowllm.enumeration.role import Role
|
||||
from flowllm.schema.message import Message as FlowMessage
|
||||
from loguru import logger
|
||||
|
||||
from reme_ai.schema import Message
|
||||
|
|
@ -13,7 +15,7 @@ from reme_ai.schema.memory import BaseMemory
|
|||
class MemoryValidationOp(BaseLLMOp):
|
||||
file_path: str = __file__
|
||||
|
||||
def execute(self):
|
||||
async def async_execute(self):
|
||||
"""Validate quality of extracted task memories"""
|
||||
|
||||
task_memories: List[BaseMemory] = []
|
||||
|
|
@ -31,7 +33,7 @@ class MemoryValidationOp(BaseLLMOp):
|
|||
validated_task_memories = []
|
||||
|
||||
for task_memory in task_memories:
|
||||
validation_result = self._validate_single_task_memory(task_memory)
|
||||
validation_result = await self._validate_single_task_memory(task_memory)
|
||||
if validation_result and validation_result.get("is_valid", False):
|
||||
task_memory.score = validation_result.get("score", 0.0)
|
||||
validated_task_memories.append(task_memory)
|
||||
|
|
@ -45,13 +47,13 @@ class MemoryValidationOp(BaseLLMOp):
|
|||
self.context.response.answer = json.dumps([x.model_dump() for x in validated_task_memories])
|
||||
self.context.response.metadata["memory_list"] = validated_task_memories
|
||||
|
||||
def _validate_single_task_memory(self, task_memory: BaseMemory) -> Dict[str, Any]:
|
||||
async def _validate_single_task_memory(self, task_memory: BaseMemory) -> Dict[str, Any]:
|
||||
"""Validate single task memory"""
|
||||
validation_info = self._llm_validate_task_memory(task_memory)
|
||||
validation_info = await self._llm_validate_task_memory(task_memory)
|
||||
logger.info(f"Validating: {validation_info}")
|
||||
return validation_info
|
||||
|
||||
def _llm_validate_task_memory(self, task_memory: BaseMemory) -> Dict[str, Any]:
|
||||
async def _llm_validate_task_memory(self, task_memory: BaseMemory) -> Dict[str, Any]:
|
||||
"""Validate task memory using LLM"""
|
||||
try:
|
||||
prompt = self.prompt_format(
|
||||
|
|
@ -96,7 +98,7 @@ class MemoryValidationOp(BaseLLMOp):
|
|||
"reason": f"Parse error: {str(e_inner)}"
|
||||
}
|
||||
|
||||
return self.llm.chat(messages=[Message(content=prompt)], callback_fn=parse_validation)
|
||||
return await self.llm.achat(messages=[FlowMessage(role=Role.USER, content=prompt)], callback_fn=parse_validation)
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"LLM validation failed: {e}")
|
||||
|
|
|
|||
|
|
@ -2,6 +2,8 @@ import json
|
|||
from typing import List, Dict
|
||||
|
||||
from flowllm import C, BaseLLMOp
|
||||
from flowllm.enumeration.role import Role
|
||||
from flowllm.schema.message import Message as FlowMessage
|
||||
from loguru import logger
|
||||
|
||||
from reme_ai.schema import Message, Trajectory
|
||||
|
|
@ -13,7 +15,7 @@ from reme_ai.utils.op_utils import merge_messages_content
|
|||
class SimpleComparativeSummaryOp(BaseLLMOp):
|
||||
file_path: str = __file__
|
||||
|
||||
def compare_summary_trajectory(self, trajectory_a: Trajectory, trajectory_b: Trajectory) -> List[BaseMemory]:
|
||||
async def compare_summary_trajectory(self, trajectory_a: Trajectory, trajectory_b: Trajectory) -> List[BaseMemory]:
|
||||
summary_prompt = self.prompt_format(prompt_name="summary_prompt",
|
||||
execution_process_a=merge_messages_content(trajectory_a.messages),
|
||||
execution_process_b=merge_messages_content(trajectory_b.messages),
|
||||
|
|
@ -42,9 +44,9 @@ class SimpleComparativeSummaryOp(BaseLLMOp):
|
|||
logger.exception(f"parse content failed!\n{content}")
|
||||
raise e
|
||||
|
||||
return self.llm.chat(messages=[Message(content=summary_prompt)], callback_fn=parse_content)
|
||||
return await self.llm.achat(messages=[FlowMessage(role=Role.USER, content=summary_prompt)], callback_fn=parse_content)
|
||||
|
||||
def execute(self):
|
||||
async def async_execute(self):
|
||||
trajectories: list = self.context.get("trajectories", [])
|
||||
trajectories: List[Trajectory] = [Trajectory(**x) if isinstance(x, dict) else x for x in trajectories]
|
||||
|
||||
|
|
@ -61,7 +63,7 @@ class SimpleComparativeSummaryOp(BaseLLMOp):
|
|||
continue
|
||||
|
||||
if task_trajectories[0].score > task_trajectories[-1].score:
|
||||
task_memories = self.compare_summary_trajectory(trajectory_a=task_trajectories[0],
|
||||
task_memories = await self.compare_summary_trajectory(trajectory_a=task_trajectories[0],
|
||||
trajectory_b=task_trajectories[-1])
|
||||
memory_list.extend(task_memories)
|
||||
|
||||
|
|
|
|||
|
|
@ -2,6 +2,8 @@ import json
|
|||
from typing import List
|
||||
|
||||
from flowllm import C, BaseLLMOp
|
||||
from flowllm.enumeration.role import Role
|
||||
from flowllm.schema.message import Message as FlowMessage
|
||||
from loguru import logger
|
||||
|
||||
from reme_ai.schema import Message, Trajectory
|
||||
|
|
@ -13,7 +15,7 @@ from reme_ai.utils.op_utils import merge_messages_content
|
|||
class SimpleSummaryOp(BaseLLMOp):
|
||||
file_path: str = __file__
|
||||
|
||||
def summary_trajectory(self, trajectory: Trajectory) -> List[BaseMemory]:
|
||||
async def summary_trajectory(self, trajectory: Trajectory) -> List[BaseMemory]:
|
||||
execution_process = merge_messages_content(trajectory.messages)
|
||||
success_score_threshold: float = self.op_params.get("success_score_threshold", 0.9)
|
||||
logger.info(f"success_score_threshold={success_score_threshold}")
|
||||
|
|
@ -49,15 +51,15 @@ class SimpleSummaryOp(BaseLLMOp):
|
|||
logger.exception(f"parse content failed!\n{content}")
|
||||
raise e
|
||||
|
||||
return self.llm.chat(messages=[Message(content=summary_prompt)], callback_fn=parse_content)
|
||||
return await self.llm.achat(messages=[FlowMessage(role=Role.USER, content=summary_prompt)], callback_fn=parse_content)
|
||||
|
||||
def execute(self):
|
||||
async def async_execute(self):
|
||||
trajectories: list = self.context.trajectories
|
||||
trajectories: List[Trajectory] = [Trajectory(**x) if isinstance(x, dict) else x for x in trajectories]
|
||||
|
||||
memory_list: List[BaseMemory] = []
|
||||
for trajectory in trajectories:
|
||||
memories = self.summary_trajectory(trajectory)
|
||||
memories = await self.summary_trajectory(trajectory)
|
||||
if memories:
|
||||
memory_list.extend(memories)
|
||||
|
||||
|
|
|
|||
|
|
@ -1,6 +1,8 @@
|
|||
from typing import List
|
||||
|
||||
from flowllm import C, BaseLLMOp
|
||||
from flowllm.enumeration.role import Role
|
||||
from flowllm.schema.message import Message as FlowMessage
|
||||
from loguru import logger
|
||||
|
||||
from reme_ai.schema import Message, Trajectory
|
||||
|
|
@ -12,7 +14,7 @@ from reme_ai.utils.op_utils import merge_messages_content, parse_json_experience
|
|||
class SuccessExtractionOp(BaseLLMOp):
|
||||
file_path: str = __file__
|
||||
|
||||
def execute(self):
|
||||
async def async_execute(self):
|
||||
"""Extract task memories from successful trajectories"""
|
||||
success_trajectories: List[Trajectory] = self.context.success_trajectories
|
||||
|
||||
|
|
@ -29,11 +31,11 @@ class SuccessExtractionOp(BaseLLMOp):
|
|||
if "segments" in trajectory.metadata:
|
||||
# Process segmented step sequences
|
||||
for segment in trajectory.metadata["segments"]:
|
||||
task_memories = self._extract_success_task_memory_from_steps(segment, trajectory)
|
||||
task_memories = await self._extract_success_task_memory_from_steps(segment, trajectory)
|
||||
success_task_memories.extend(task_memories)
|
||||
else:
|
||||
# Process entire trajectory
|
||||
task_memories = self._extract_success_task_memory_from_steps(trajectory.messages, trajectory)
|
||||
task_memories = await self._extract_success_task_memory_from_steps(trajectory.messages, trajectory)
|
||||
success_task_memories.extend(task_memories)
|
||||
|
||||
logger.info(f"Extracted {len(success_task_memories)} success task memories")
|
||||
|
|
@ -41,7 +43,7 @@ class SuccessExtractionOp(BaseLLMOp):
|
|||
# Add task memories to context
|
||||
self.context.success_task_memories = success_task_memories
|
||||
|
||||
def _extract_success_task_memory_from_steps(self, steps: List[Message], trajectory: Trajectory) -> List[BaseMemory]:
|
||||
async def _extract_success_task_memory_from_steps(self, steps: List[Message], trajectory: Trajectory) -> List[BaseMemory]:
|
||||
"""Extract task memory from successful step sequences"""
|
||||
step_content = merge_messages_content(steps)
|
||||
context = get_trajectory_context(trajectory, steps)
|
||||
|
|
@ -70,4 +72,4 @@ class SuccessExtractionOp(BaseLLMOp):
|
|||
|
||||
return task_memories
|
||||
|
||||
return self.llm.chat(messages=[Message(content=prompt)], callback_fn=parse_task_memories)
|
||||
return await self.llm.achat(messages=[FlowMessage(role=Role.USER, content=prompt)], callback_fn=parse_task_memories)
|
||||
|
|
|
|||
|
|
@ -10,7 +10,7 @@ from reme_ai.schema import Trajectory
|
|||
class TrajectoryPreprocessOp(BaseOp):
|
||||
file_path: str = __file__
|
||||
|
||||
def execute(self):
|
||||
async def async_execute(self):
|
||||
"""Preprocess trajectories: validate and classify"""
|
||||
trajectories: list = self.context.get("trajectories", [])
|
||||
trajectories: List[Trajectory] = [Trajectory(**x) if isinstance(x, dict) else x for x in trajectories]
|
||||
|
|
|
|||
|
|
@ -3,6 +3,8 @@ import re
|
|||
from typing import List
|
||||
|
||||
from flowllm import C, BaseLLMOp
|
||||
from flowllm.enumeration.role import Role
|
||||
from flowllm.schema.message import Message as FlowMessage
|
||||
from loguru import logger
|
||||
|
||||
from reme_ai.schema import Message, Trajectory
|
||||
|
|
@ -12,7 +14,7 @@ from reme_ai.schema import Message, Trajectory
|
|||
class TrajectorySegmentationOp(BaseLLMOp):
|
||||
file_path: str = __file__
|
||||
|
||||
def execute(self):
|
||||
async def async_execute(self):
|
||||
"""Segment trajectories into meaningful steps"""
|
||||
# Get trajectories from context
|
||||
all_trajectories: List[Trajectory] = self.context.get("all_trajectories", [])
|
||||
|
|
@ -30,7 +32,7 @@ class TrajectorySegmentationOp(BaseLLMOp):
|
|||
# Add segmentation info to trajectories
|
||||
segmented_count = 0
|
||||
for trajectory in target_trajectories:
|
||||
segments = self._llm_segment_trajectory(trajectory)
|
||||
segments = await self._llm_segment_trajectory(trajectory)
|
||||
trajectory.metadata["segments"] = segments
|
||||
segmented_count += 1
|
||||
|
||||
|
|
@ -51,7 +53,7 @@ class TrajectorySegmentationOp(BaseLLMOp):
|
|||
else:
|
||||
return all_trajectories
|
||||
|
||||
def _llm_segment_trajectory(self, trajectory: Trajectory) -> List[List[Message]]:
|
||||
async def _llm_segment_trajectory(self, trajectory: Trajectory) -> List[List[Message]]:
|
||||
"""Use LLM for trajectory segmentation"""
|
||||
trajectory_content = self._format_trajectory_content(trajectory)
|
||||
|
||||
|
|
@ -80,7 +82,7 @@ class TrajectorySegmentationOp(BaseLLMOp):
|
|||
|
||||
return segments if segments else [trajectory.messages]
|
||||
|
||||
return self.llm.chat(messages=[Message(content=prompt)], callback_fn=parse_segmentation,
|
||||
return await self.llm.achat(messages=[FlowMessage(role=Role.USER, content=prompt)], callback_fn=parse_segmentation,
|
||||
default_value=[trajectory.messages])
|
||||
|
||||
@staticmethod
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue