mirror of
https://github.com/qdrant/qdrant-client.git
synced 2026-07-23 11:11:01 -05:00
fix: fix local inference on list of prefetches (#965)
This commit is contained in:
@@ -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(
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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",
|
||||
|
||||
Reference in New Issue
Block a user