|
| 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