diff --git a/haiku_rag_slim/haiku/rag/embeddings/__init__.py b/haiku_rag_slim/haiku/rag/embeddings/__init__.py index 9b91989b..74794750 100644 --- a/haiku_rag_slim/haiku/rag/embeddings/__init__.py +++ b/haiku_rag_slim/haiku/rag/embeddings/__init__.py @@ -48,15 +48,13 @@ class EmbedderWrapper: return [list(e) for e in result.embeddings] async def embed_image_query(self, image: "Any") -> list[float]: - """Embed a single image as a search query.""" - raise NotImplementedError( - f"{type(self).__name__} does not support image embedding. " - "Configure a multimodal provider (e.g. provider='mlx' or " - "provider='vllm')." - ) + """Embed a single image into the same vector space as text. - async def embed_images(self, images: list["Any"]) -> list[list[float]]: - """Batch-embed images for indexing into the same vector space as text.""" + Multimodal providers override this. Picture embedding is single-image: + vLLM's ``/v1/embeddings`` accepts one image per request via the + ``messages`` superset, and MLX runs forward passes one at a time. + Callers loop when they need many. + """ raise NotImplementedError( f"{type(self).__name__} does not support image embedding. " "Configure a multimodal provider (e.g. provider='mlx' or " diff --git a/haiku_rag_slim/haiku/rag/embeddings/mlx.py b/haiku_rag_slim/haiku/rag/embeddings/mlx.py index 26ad8229..45244454 100644 --- a/haiku_rag_slim/haiku/rag/embeddings/mlx.py +++ b/haiku_rag_slim/haiku/rag/embeddings/mlx.py @@ -77,15 +77,7 @@ class MLXEmbedder(EmbedderWrapper): return await asyncio.to_thread(self._encode_texts, texts) async def embed_image_query(self, image: "bytes | PILImage.Image") -> list[float]: - embeddings = await self.embed_images([image]) - return embeddings[0] - - async def embed_images( - self, images: list["bytes | PILImage.Image"] - ) -> list[list[float]]: - if not images: - return [] - return await asyncio.to_thread(self._encode_images, images) + return await asyncio.to_thread(self._encode_image, image) def _encode_texts(self, texts: list[str]) -> list[list[float]]: import mlx.core as mx # ty: ignore[unresolved-import,unused-ignore-comment] @@ -107,34 +99,27 @@ class MLXEmbedder(EmbedderWrapper): mx.eval(embeddings) return [list(map(float, row)) for row in embeddings] - def _encode_images( - self, images: list["bytes | PILImage.Image"] - ) -> list[list[float]]: + def _encode_image(self, image: "bytes | PILImage.Image") -> list[float]: import mlx.core as mx # ty: ignore[unresolved-import,unused-ignore-comment] - from PIL import Image as PILImageModule model, processor = self._ensure_loaded() - pil_images = [_to_pil(img) for img in images] - out: list[list[float]] = [] - for pil_image in pil_images: - assert isinstance(pil_image, PILImageModule.Image) - inputs = processor( - text=[_DEFAULT_IMAGE_PROMPT], - images=[pil_image], - return_tensors="np", - padding=True, - ) - pixel_values = inputs["pixel_values"] - embedding = model.encode_image( - input_ids=mx.array(inputs["input_ids"]), - pixel_values=mx.array(pixel_values.reshape(-1, pixel_values.shape[-1])), - image_grid_thw=[tuple(r) for r in inputs["image_grid_thw"]], - attention_mask=mx.array(inputs["attention_mask"]), - task="retrieval", - ) - mx.eval(embedding) - out.append([float(x) for x in embedding[0]]) - return out + pil_image = _to_pil(image) + inputs = processor( + text=[_DEFAULT_IMAGE_PROMPT], + images=[pil_image], + return_tensors="np", + padding=True, + ) + pixel_values = inputs["pixel_values"] + embedding = model.encode_image( + input_ids=mx.array(inputs["input_ids"]), + pixel_values=mx.array(pixel_values.reshape(-1, pixel_values.shape[-1])), + image_grid_thw=[tuple(r) for r in inputs["image_grid_thw"]], + attention_mask=mx.array(inputs["attention_mask"]), + task="retrieval", + ) + mx.eval(embedding) + return [float(x) for x in embedding[0]] def _to_pil(image: "bytes | PILImage.Image") -> "PILImage.Image": diff --git a/haiku_rag_slim/haiku/rag/embeddings/vllm.py b/haiku_rag_slim/haiku/rag/embeddings/vllm.py index 0f7f1871..c37c4f4f 100644 --- a/haiku_rag_slim/haiku/rag/embeddings/vllm.py +++ b/haiku_rag_slim/haiku/rag/embeddings/vllm.py @@ -1,15 +1,16 @@ """Multimodal embedder backed by a vLLM OpenAI-compatible HTTP server. -vLLM's ``/v1/embeddings`` endpoint is a superset of OpenAI's: when the -request body uses a ``messages`` array (instead of ``input``), the server -treats it as a chat-style multimodal embedding request and accepts -``image_url`` content parts carrying base64 data URIs. Models like -``Qwen/Qwen3-VL-Embedding-8B`` and ``jinaai/jina-embeddings-v4`` ship with -chat templates that map both text and image inputs into a shared vector -space. +vLLM's ``/v1/embeddings`` endpoint is a superset of OpenAI's: + +- Text inputs use the standard ``input: list[str]`` field — one HTTP call + returns N embeddings. +- Image inputs use a ``messages`` array carrying an ``image_url`` content + part with a base64 data URI. One image per HTTP call. + +Models like ``Qwen/Qwen3-VL-Embedding-8B`` and ``jinaai/jina-embeddings-v4`` +ship with chat templates that map both shapes into a shared vector space. """ -import asyncio import base64 import io from typing import TYPE_CHECKING, Any @@ -45,12 +46,7 @@ class VLLMMultimodalEmbedder(EmbedderWrapper): headers["Authorization"] = f"Bearer {self._api_key}" return headers - async def _post_messages(self, content: list[dict[str, Any]]) -> list[float]: - body = { - "model": self._model_name, - "messages": [{"role": "user", "content": content}], - "encoding_format": "float", - } + async def _post(self, body: dict[str, Any]) -> list[list[float]]: try: async with httpx.AsyncClient(timeout=self._timeout) as client: response = await client.post( @@ -79,27 +75,48 @@ class VLLMMultimodalEmbedder(EmbedderWrapper): data = payload.get("data") or [] if not data: raise ValueError(f"vLLM returned no embeddings: {payload}") - return list(data[0]["embedding"]) + return [list(d["embedding"]) for d in data] async def embed_query(self, text: str) -> list[float]: - return await self._post_messages([{"type": "text", "text": text}]) + rows = await self._post( + { + "model": self._model_name, + "input": [text], + "encoding_format": "float", + } + ) + return rows[0] async def embed_documents(self, texts: list[str]) -> list[list[float]]: if not texts: return [] - return await asyncio.gather(*(self.embed_query(t) for t in texts)) - - async def embed_image_query(self, image: "bytes | PILImage.Image") -> list[float]: - return await self._post_messages( - [{"type": "image_url", "image_url": {"url": _to_data_uri(image)}}] + return await self._post( + { + "model": self._model_name, + "input": texts, + "encoding_format": "float", + } ) - async def embed_images( - self, images: list["bytes | PILImage.Image"] - ) -> list[list[float]]: - if not images: - return [] - return await asyncio.gather(*(self.embed_image_query(img) for img in images)) + async def embed_image_query(self, image: "bytes | PILImage.Image") -> list[float]: + rows = await self._post( + { + "model": self._model_name, + "messages": [ + { + "role": "user", + "content": [ + { + "type": "image_url", + "image_url": {"url": _to_data_uri(image)}, + } + ], + } + ], + "encoding_format": "float", + } + ) + return rows[0] def _to_data_uri(image: "bytes | PILImage.Image") -> str: diff --git a/tests/test_embedder.py b/tests/test_embedder.py index da4da4c3..85c26593 100644 --- a/tests/test_embedder.py +++ b/tests/test_embedder.py @@ -238,12 +238,12 @@ async def test_text_only_embedder_does_not_support_images(): assert embedder.supports_images is False with pytest.raises(NotImplementedError, match="multimodal provider"): await embedder.embed_image_query(b"\x89PNG\r\n\x1a\n") - with pytest.raises(NotImplementedError, match="multimodal provider"): - await embedder.embed_images([b"\x89PNG\r\n\x1a\n"]) async def test_vllm_embed_text_request_shape(monkeypatch): - """vLLM text embedding posts a `messages` array with a single text part.""" + """vLLM text embedding posts a standard OpenAI ``input`` array (real + server-side batching), not the ``messages`` superset (which is reserved + for image inputs).""" from haiku.rag.embeddings.vllm import VLLMMultimodalEmbedder captured: dict = {} @@ -253,7 +253,12 @@ async def test_vllm_embed_text_request_shape(monkeypatch): pass def json(self): - return {"data": [{"embedding": [0.1, 0.2, 0.3]}]} + return { + "data": [ + {"embedding": [0.1, 0.2, 0.3]}, + {"embedding": [0.4, 0.5, 0.6]}, + ] + } class FakeAsyncClient: def __init__(self, *args, **kwargs): @@ -277,14 +282,13 @@ async def test_vllm_embed_text_request_shape(monkeypatch): vector_dim=2048, base_url="http://localhost:8000/v1", ) - vec = await embedder.embed_query("a photo of a cat") - assert vec == [0.1, 0.2, 0.3] + vecs = await embedder.embed_documents(["a photo of a cat", "a sleeping dog"]) + assert vecs == [[0.1, 0.2, 0.3], [0.4, 0.5, 0.6]] assert captured["url"] == "http://localhost:8000/v1/embeddings" body = captured["body"] assert body["model"] == "Qwen/Qwen3-VL-Embedding-2B" - assert body["messages"] == [ - {"role": "user", "content": [{"type": "text", "text": "a photo of a cat"}]} - ] + assert body["input"] == ["a photo of a cat", "a sleeping dog"] + assert "messages" not in body assert body["encoding_format"] == "float"