fix: fix local inference on list of prefetches (#965)

This commit is contained in:
George
2025-04-24 17:28:17 +03:00
committed by George Panchuk
parent eef5e2560e
commit b2af969c15
4 changed files with 23 additions and 6 deletions

View File

@@ -527,7 +527,9 @@ class AsyncQdrantClient(AsyncQdrantFastembedMixin):
"""
assert len(kwargs) == 0, f"Unknown arguments: {list(kwargs.keys())}"
query = self._resolve_query(query)
requires_inference = self._inference_inspector.inspect([query, prefetch])
requires_inference = self._inference_inspector.inspect(query)
if not requires_inference:
requires_inference = self._inference_inspector.inspect(prefetch)
if requires_inference and (not self.cloud_inference):
query = (
next(
@@ -691,7 +693,9 @@ class AsyncQdrantClient(AsyncQdrantFastembedMixin):
"""
assert len(kwargs) == 0, f"Unknown arguments: {list(kwargs.keys())}"
query = self._resolve_query(query)
requires_inference = self._inference_inspector.inspect([query, prefetch])
requires_inference = self._inference_inspector.inspect(query)
if not requires_inference:
requires_inference = self._inference_inspector.inspect(prefetch)
if requires_inference and (not self.cloud_inference):
query = (
next(

View File

@@ -45,6 +45,8 @@ class Inspector:
self.parser.parse_model(point.__class__)
if self._inspect_model(point):
return True
else:
return False
return False
def _inspect_model(self, model: BaseModel, paths: Optional[list[FieldPath]] = None) -> bool:

View File

@@ -556,7 +556,9 @@ class QdrantClient(QdrantFastembedMixin):
# If the query contains unprocessed documents, we need to embed them and
# replace the original query with the embedded vectors.
query = self._resolve_query(query)
requires_inference = self._inference_inspector.inspect([query, prefetch])
requires_inference = self._inference_inspector.inspect(query)
if not requires_inference:
requires_inference = self._inference_inspector.inspect(prefetch)
if requires_inference and not self.cloud_inference:
query = (
next(
@@ -725,7 +727,9 @@ class QdrantClient(QdrantFastembedMixin):
# If the query contains unprocessed documents, we need to embed them and
# replace the original query with the embedded vectors.
query = self._resolve_query(query)
requires_inference = self._inference_inspector.inspect([query, prefetch])
requires_inference = self._inference_inspector.inspect(query)
if not requires_inference:
requires_inference = self._inference_inspector.inspect(prefetch)
if requires_inference and not self.cloud_inference:
query = (
next(

View File

@@ -246,6 +246,13 @@ def test_inspect_prefetch_types():
paths = inspector_embed.inspect(doc_prefetch)
assert len(paths) == 1 and paths[0].as_str_list() == ["query"]
no_query_list_prefetch_with_doc = models.Prefetch(
query=None, prefetch=[models.Prefetch(query=None), models.Prefetch(query=doc)]
)
assert inspector.inspect(no_query_list_prefetch_with_doc)
paths = inspector_embed.inspect(no_query_list_prefetch_with_doc)
assert len(paths) == 1 and paths[0].as_str_list() == ["prefetch.query"]
nested_prefetch = models.Prefetch(
query=None,
prefetch=models.Prefetch(query=doc),
@@ -276,8 +283,8 @@ def test_inspect_prefetch_types():
"prefetch.prefetch.query",
}
assert inspector.inspect([None, deep_nested_prefetch])
paths = inspector_embed.inspect([None, deep_nested_prefetch])
assert inspector.inspect([none_prefetch, deep_nested_prefetch])
paths = inspector_embed.inspect([none_prefetch, deep_nested_prefetch])
assert len(paths) == 1 and set(paths[0].as_str_list()) == {
"prefetch.prefetch.prefetch.query",
"prefetch.prefetch.query",