From a83f83506edfaf257d3dea4061d9fd9b09521cfa Mon Sep 17 00:00:00 2001 From: qingxu Date: Fri, 6 Jun 2025 16:12:11 +0800 Subject: [PATCH] create traj summarizer --- .../module/summarizer/traj_summarizer.py | 26 +++++++++++++++++++ 1 file changed, 26 insertions(+) create mode 100644 experiencemaker/module/summarizer/traj_summarizer.py diff --git a/experiencemaker/module/summarizer/traj_summarizer.py b/experiencemaker/module/summarizer/traj_summarizer.py new file mode 100644 index 00000000..7e6575f5 --- /dev/null +++ b/experiencemaker/module/summarizer/traj_summarizer.py @@ -0,0 +1,26 @@ +from typing import List + +from pydantic import Field + +from experiencemaker.schema.trajectory import Trajectory, Sample, SummaryMessage +from experiencemaker.storage.base_vector_store import BaseVectorStore +from experiencemaker.module.summarizer.base_summarizer import BaseSummarizer + +class TrajectorySummarizer(BaseSummarizer): + vector_store: BaseVectorStore | None = Field(default=None) + + def extract_samples(self, trajectories: List[Trajectory], **kwargs) -> List[Sample]: + raise NotImplementedError + + def insert_into_vector_store(self, samples: List[Sample], **kwargs): + raise NotImplementedError + + def execute(self, trajectories: List[Trajectory], return_samples: bool = True, **kwargs) -> List[Sample]: + samples: List[Sample] = self.extract_samples(trajectories, **kwargs) + self.insert_into_vector_store(samples, **kwargs) + + if return_samples: + return samples + + return [] +