step summarizer and context generator

This commit is contained in:
鸣山 2025-06-09 16:48:23 +08:00
parent 195d2aef4d
commit 2f519457b3
2 changed files with 973 additions and 0 deletions

View file

@ -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 ""

View 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)