Skip to content

Commit 683af0d

Browse files
Add tests for CustomTextCrossEncoder (#238)
Adds comprehensive testing for the custom model registration and initialization for the CustomTextCrossEncoder class. Co-authored-by: google-labs-jules[bot] <161369871+google-labs-jules[bot]@users.noreply.github.com>
1 parent 5d1e462 commit 683af0d

1 file changed

Lines changed: 95 additions & 0 deletions

File tree

Lines changed: 95 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,95 @@
1+
"""Tests for CustomTextCrossEncoder."""
2+
3+
from qwen3_embed.common.model_description import BaseModelDescription, ModelSource
4+
from qwen3_embed.rerank.cross_encoder.custom_text_cross_encoder import CustomTextCrossEncoder
5+
6+
7+
class TestCustomTextCrossEncoder:
8+
"""Tests for CustomTextCrossEncoder custom model registry methods."""
9+
10+
def setup_method(self) -> None:
11+
"""Clear custom model registry between tests."""
12+
CustomTextCrossEncoder.SUPPORTED_MODELS.clear()
13+
14+
def test_list_supported_models_empty_initially(self) -> None:
15+
"""Verify the model list is empty initially."""
16+
models = CustomTextCrossEncoder._list_supported_models()
17+
assert isinstance(models, list)
18+
assert len(models) == 0
19+
20+
def test_add_model(self) -> None:
21+
"""Verify add_model correctly adds a model description."""
22+
desc = BaseModelDescription(
23+
model="org/test-model",
24+
sources=ModelSource(hf="org/test-model"),
25+
model_file="onnx/model.onnx",
26+
description="",
27+
license="",
28+
size_in_GB=0.0,
29+
)
30+
CustomTextCrossEncoder.add_model(desc)
31+
32+
models = CustomTextCrossEncoder._list_supported_models()
33+
assert len(models) == 1
34+
assert models[0].model == "org/test-model"
35+
36+
def test_add_multiple_models(self) -> None:
37+
"""Verify multiple models can be added and listed."""
38+
desc1 = BaseModelDescription(
39+
model="org/test-model-1",
40+
sources=ModelSource(hf="org/test-model-1"),
41+
model_file="onnx/model.onnx",
42+
description="",
43+
license="",
44+
size_in_GB=0.0,
45+
)
46+
desc2 = BaseModelDescription(
47+
model="org/test-model-2",
48+
sources=ModelSource(hf="org/test-model-2"),
49+
model_file="onnx/model.onnx",
50+
description="",
51+
license="",
52+
size_in_GB=0.0,
53+
)
54+
55+
CustomTextCrossEncoder.add_model(desc1)
56+
CustomTextCrossEncoder.add_model(desc2)
57+
58+
models = CustomTextCrossEncoder._list_supported_models()
59+
assert len(models) == 2
60+
assert models[0].model == "org/test-model-1"
61+
assert models[1].model == "org/test-model-2"
62+
63+
def test_custom_text_cross_encoder_init_passes_args(self, tmp_path) -> None:
64+
"""Verify that CustomTextCrossEncoder passes arguments correctly during initialization."""
65+
# Create a dummy model directory to avoid actual downloading
66+
model_dir = tmp_path / "dummy_model"
67+
model_dir.mkdir()
68+
69+
# Add a test model to the registry
70+
desc = BaseModelDescription(
71+
model="org/init-test-model",
72+
sources=ModelSource(hf="org/init-test-model"),
73+
model_file="onnx/model.onnx",
74+
description="",
75+
license="",
76+
size_in_GB=0.0,
77+
)
78+
CustomTextCrossEncoder.add_model(desc)
79+
80+
# We patch load_onnx_model so we don't actually try to load a real ONNX file
81+
# and download_model so we don't try to download from HF
82+
from unittest.mock import patch
83+
84+
with (
85+
patch.object(CustomTextCrossEncoder, "download_model", return_value=model_dir),
86+
patch.object(CustomTextCrossEncoder, "load_onnx_model") as mock_load,
87+
):
88+
encoder = CustomTextCrossEncoder(
89+
model_name="org/init-test-model",
90+
threads=4,
91+
)
92+
93+
assert encoder.model_name == "org/init-test-model"
94+
assert encoder.threads == 4
95+
mock_load.assert_called_once()

0 commit comments

Comments
 (0)