Expose filter arg in get_state_history

- Combine list and search methods in Checkpointer
This commit is contained in:
Nuno Campos
2024-06-05 17:14:35 -07:00
parent 052284f21d
commit d366bacf03
11 changed files with 154 additions and 348 deletions
+4 -4
View File
@@ -66,15 +66,15 @@ class FewShotExamples(ManagedValue[Sequence[V]], Generic[V]):
return self.metadata_filter
def iter(self, score: int = 1) -> Iterator[V]:
for example in self.graph.checkpointer.search(
{"score": score, **self.metadata_filter_dict}, limit=self.k
for example in self.graph.checkpointer.list(
None, filter={"score": score, **self.metadata_filter_dict}, limit=self.k
):
with ChannelsManager(self.graph.channels, example.checkpoint) as channels:
yield read_channels(channels, self.graph.output_channels)
async def aiter(self, score: int = 1) -> AsyncIterator[V]:
async for example in self.graph.checkpointer.asearch(
{"score": score, **self.metadata_filter_dict}, limit=self.k
async for example in self.graph.checkpointer.alist(
None, filter={"score": score, **self.metadata_filter_dict}, limit=self.k
):
async with AsyncChannelsManager(
self.graph.channels, example.checkpoint