Skip to content

Commit 189d9bc

Browse files
feat: add tests for OnnxTextEmbedding (#267)
* test: add missing tests for onnx embedding Co-authored-by: n24q02m <135627235+n24q02m@users.noreply.github.com> * test: add missing tests for onnx embedding Co-authored-by: n24q02m <135627235+n24q02m@users.noreply.github.com> * test: fix OnnxTextEmbeddingWorker test to appease type checker Co-authored-by: n24q02m <135627235+n24q02m@users.noreply.github.com> * test: fix OnnxTextEmbeddingWorker test to appease type checker Co-authored-by: n24q02m <135627235+n24q02m@users.noreply.github.com> --------- Co-authored-by: google-labs-jules[bot] <161369871+google-labs-jules[bot]@users.noreply.github.com>
1 parent 629721e commit 189d9bc

1 file changed

Lines changed: 184 additions & 0 deletions

File tree

tests/test_onnx_embedding.py

Lines changed: 184 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,184 @@
1+
from pathlib import Path
2+
from unittest.mock import MagicMock, patch
3+
4+
import numpy as np
5+
6+
from qwen3_embed.common.model_description import DenseModelDescription, ModelSource
7+
from qwen3_embed.common.onnx_model import OnnxOutputContext
8+
from qwen3_embed.text.onnx_embedding import OnnxTextEmbedding, OnnxTextEmbeddingWorker
9+
10+
_MODEL_NAME = "test-org/test-model"
11+
_MODEL_DESC = DenseModelDescription(
12+
model=_MODEL_NAME,
13+
sources=ModelSource(hf=_MODEL_NAME),
14+
model_file="model.onnx",
15+
description="Test model",
16+
license="MIT",
17+
size_in_GB=0.1,
18+
dim=4,
19+
)
20+
21+
22+
@patch("qwen3_embed.text.onnx_embedding.OnnxTextEmbedding._select_exposed_session_options")
23+
@patch("qwen3_embed.text.onnx_embedding.OnnxTextEmbedding._get_model_description")
24+
@patch("qwen3_embed.text.onnx_embedding.OnnxTextEmbedding.download_model")
25+
@patch("qwen3_embed.text.onnx_embedding.OnnxTextEmbedding.load_onnx_model")
26+
def test_onnx_text_embedding_init_lazy_load(
27+
mock_load_onnx_model: MagicMock,
28+
mock_download_model: MagicMock,
29+
mock_get_model_description: MagicMock,
30+
mock_select_exposed_session_options: MagicMock,
31+
) -> None:
32+
mock_get_model_description.return_value = _MODEL_DESC
33+
mock_download_model.return_value = Path("/tmp/model")
34+
mock_select_exposed_session_options.return_value = {}
35+
36+
embedding = OnnxTextEmbedding(model_name=_MODEL_NAME, lazy_load=True)
37+
38+
mock_load_onnx_model.assert_not_called()
39+
assert embedding.lazy_load is True
40+
assert embedding.model_name == _MODEL_NAME
41+
mock_get_model_description.assert_called_once_with(_MODEL_NAME)
42+
43+
44+
@patch("qwen3_embed.text.onnx_embedding.OnnxTextEmbedding._select_exposed_session_options")
45+
@patch("qwen3_embed.text.onnx_embedding.OnnxTextEmbedding._get_model_description")
46+
@patch("qwen3_embed.text.onnx_embedding.OnnxTextEmbedding.download_model")
47+
@patch("qwen3_embed.text.onnx_embedding.OnnxTextEmbedding.load_onnx_model")
48+
def test_onnx_text_embedding_init_no_lazy_load(
49+
mock_load_onnx_model: MagicMock,
50+
mock_download_model: MagicMock,
51+
mock_get_model_description: MagicMock,
52+
mock_select_exposed_session_options: MagicMock,
53+
) -> None:
54+
mock_get_model_description.return_value = _MODEL_DESC
55+
mock_download_model.return_value = Path("/tmp/model")
56+
mock_select_exposed_session_options.return_value = {}
57+
58+
embedding = OnnxTextEmbedding(model_name=_MODEL_NAME, lazy_load=False)
59+
60+
mock_load_onnx_model.assert_called_once()
61+
assert embedding.lazy_load is False
62+
63+
64+
@patch("qwen3_embed.text.onnx_embedding.OnnxTextEmbedding._select_exposed_session_options")
65+
@patch("qwen3_embed.text.onnx_embedding.OnnxTextEmbedding._get_model_description")
66+
@patch("qwen3_embed.text.onnx_embedding.OnnxTextEmbedding.download_model")
67+
@patch("qwen3_embed.text.onnx_embedding.OnnxTextEmbedding._embed_documents")
68+
def test_onnx_text_embedding_embed(
69+
mock_embed_documents: MagicMock,
70+
mock_download_model: MagicMock,
71+
mock_get_model_description: MagicMock,
72+
mock_select_exposed_session_options: MagicMock,
73+
) -> None:
74+
mock_get_model_description.return_value = _MODEL_DESC
75+
mock_download_model.return_value = Path("/tmp/model")
76+
mock_select_exposed_session_options.return_value = {}
77+
78+
# Return empty iterator
79+
mock_embed_documents.return_value = iter([])
80+
81+
embedding = OnnxTextEmbedding(model_name=_MODEL_NAME, lazy_load=True, cache_dir="/tmp/cache")
82+
83+
docs = ["doc1", "doc2"]
84+
list(embedding.embed(documents=docs, batch_size=32, parallel=4))
85+
86+
mock_embed_documents.assert_called_once()
87+
kwargs = mock_embed_documents.call_args.kwargs
88+
assert kwargs["model_name"] == _MODEL_NAME
89+
assert (
90+
kwargs["cache_dir"] == str(Path("/tmp/cache").absolute())
91+
if Path("/tmp/cache").is_absolute()
92+
else str(Path("/tmp/cache").resolve())
93+
)
94+
assert kwargs["documents"] == docs
95+
assert kwargs["batch_size"] == 32
96+
assert kwargs["parallel"] == 4
97+
98+
99+
@patch("qwen3_embed.text.onnx_embedding.OnnxTextEmbedding._select_exposed_session_options")
100+
@patch("qwen3_embed.text.onnx_embedding.OnnxTextEmbedding._get_model_description")
101+
@patch("qwen3_embed.text.onnx_embedding.OnnxTextEmbedding.download_model")
102+
def test_onnx_text_embedding_preprocess_input(
103+
mock_download_model: MagicMock,
104+
mock_get_model_description: MagicMock,
105+
mock_select_exposed_session_options: MagicMock,
106+
) -> None:
107+
mock_get_model_description.return_value = _MODEL_DESC
108+
mock_download_model.return_value = Path("/tmp/model")
109+
mock_select_exposed_session_options.return_value = {}
110+
111+
embedding = OnnxTextEmbedding(model_name=_MODEL_NAME, lazy_load=True)
112+
113+
input_dict = {"input_ids": np.array([1, 2, 3])}
114+
output_dict = embedding._preprocess_onnx_input(input_dict)
115+
116+
assert output_dict is input_dict
117+
118+
119+
@patch("qwen3_embed.text.onnx_embedding.normalize")
120+
@patch("qwen3_embed.text.onnx_embedding.OnnxTextEmbedding._select_exposed_session_options")
121+
@patch("qwen3_embed.text.onnx_embedding.OnnxTextEmbedding._get_model_description")
122+
@patch("qwen3_embed.text.onnx_embedding.OnnxTextEmbedding.download_model")
123+
def test_onnx_text_embedding_postprocess_2d(
124+
mock_download_model: MagicMock,
125+
mock_get_model_description: MagicMock,
126+
mock_select_exposed_session_options: MagicMock,
127+
mock_normalize: MagicMock,
128+
) -> None:
129+
mock_get_model_description.return_value = _MODEL_DESC
130+
mock_download_model.return_value = Path("/tmp/model")
131+
mock_select_exposed_session_options.return_value = {}
132+
mock_normalize.side_effect = lambda x: x
133+
134+
embedding = OnnxTextEmbedding(model_name=_MODEL_NAME, lazy_load=True)
135+
136+
model_output = np.array([[1.0, 2.0], [3.0, 4.0]])
137+
output_context = OnnxOutputContext(model_output=model_output, attention_mask=None)
138+
139+
embedding._post_process_onnx_output(output_context)
140+
141+
mock_normalize.assert_called_once()
142+
assert np.array_equal(mock_normalize.call_args[0][0], model_output)
143+
144+
145+
@patch("qwen3_embed.text.onnx_embedding.normalize")
146+
@patch("qwen3_embed.text.onnx_embedding.OnnxTextEmbedding._select_exposed_session_options")
147+
@patch("qwen3_embed.text.onnx_embedding.OnnxTextEmbedding._get_model_description")
148+
@patch("qwen3_embed.text.onnx_embedding.OnnxTextEmbedding.download_model")
149+
def test_onnx_text_embedding_postprocess_3d(
150+
mock_download_model: MagicMock,
151+
mock_get_model_description: MagicMock,
152+
mock_select_exposed_session_options: MagicMock,
153+
mock_normalize: MagicMock,
154+
) -> None:
155+
mock_get_model_description.return_value = _MODEL_DESC
156+
mock_download_model.return_value = Path("/tmp/model")
157+
mock_select_exposed_session_options.return_value = {}
158+
mock_normalize.side_effect = lambda x: x
159+
160+
embedding = OnnxTextEmbedding(model_name=_MODEL_NAME, lazy_load=True)
161+
162+
# 3D array: (batch_size, seq_len, dim)
163+
model_output = np.array([[[1.0, 2.0], [9.0, 9.0]], [[3.0, 4.0], [9.0, 9.0]]])
164+
output_context = OnnxOutputContext(model_output=model_output, attention_mask=None)
165+
166+
embedding._post_process_onnx_output(output_context)
167+
168+
mock_normalize.assert_called_once()
169+
# It should slice [:, 0]
170+
expected_slice = np.array([[1.0, 2.0], [3.0, 4.0]])
171+
assert np.array_equal(mock_normalize.call_args[0][0], expected_slice)
172+
173+
174+
@patch("qwen3_embed.text.onnx_embedding.OnnxTextEmbedding")
175+
def test_onnx_text_embedding_worker_init(
176+
mock_onnx_embedding: MagicMock,
177+
) -> None:
178+
worker = OnnxTextEmbeddingWorker.__new__(OnnxTextEmbeddingWorker)
179+
worker.__init__ = lambda *args, **kwargs: None
180+
worker.init_embedding(model_name=_MODEL_NAME, cache_dir="/tmp/cache", extra="arg")
181+
182+
mock_onnx_embedding.assert_called_once_with(
183+
model_name=_MODEL_NAME, cache_dir="/tmp/cache", threads=1, extra="arg"
184+
)

0 commit comments

Comments
 (0)