ReMe/reme_ai/summary/task/trajectory_preprocess_op.py
jinli.yl cceebc631e refactor: rename package and restructure modules
- Rename experiencemaker package to reme_ai
- Move personal modules to new directory structure
- Remove unused classes and imports
- Update module initialization files
2025-08-25 23:59:08 +08:00

46 lines
No EOL
1.6 KiB
Python

from typing import List, Dict
from loguru import logger
from flowllm import C, BaseOp
from reme_ai.schema.message import Trajectory
@C.register_op()
class TrajectoryPreprocessOp(BaseOp):
current_path: str = __file__
def execute(self):
"""Preprocess trajectories: validate and classify"""
trajectories: List[Trajectory] = self.context.get("trajectories", [])
# Classify trajectories
classified = self._classify_trajectories(trajectories)
logger.info(f"Classified trajectories - Success: {len(classified['success'])}, "
f"Failure: {len(classified['failure'])}, All: {len(classified['all'])}")
# Set context for downstream operators
self.context.success_trajectories = classified['success']
self.context.failure_trajectories = classified['failure']
self.context.all_trajectories = classified['all']
def _classify_trajectories(self, trajectories: List[Trajectory]) -> Dict[str, List[Trajectory]]:
"""Classify trajectories based on score threshold"""
success_trajectories = []
failure_trajectories = []
success_threshold = self.op_params.get("success_threshold", 1.0)
for traj in trajectories:
is_success = traj.score >= success_threshold
if is_success:
success_trajectories.append(traj)
else:
failure_trajectories.append(traj)
return {
'success': success_trajectories,
'failure': failure_trajectories,
'all': trajectories
}