Skip to content

Commit 8367ad7

Browse files
n24q02mn24q02m
andauthored
fix: propagate custom reranker registry to spawned workers (#799)
TextCrossEncoder(custom).rerank_pairs(..., parallel=N) with input >= batch_size spawned worker processes whose fresh interpreter had an empty registry, so the custom model could not be resolved and the worker raised "Model not supported". The embedding side already handled this; the reranker did not. Mirror the embedding fix: OnnxCrossEncoderModel gains _extra_worker_params() and the parallel pool path now does params.update(self._extra_worker_params()); CustomTextCrossEncoder gains _export_registry/_import_registry (idempotent) + _extra_worker_params (carries custom_registry) + a CustomTextCrossEncoderWorker that re-imports the registry before constructing the model in the worker. Co-authored-by: n24q02m <n24q02m@outlook.com>
1 parent 77c95f2 commit 8367ad7

3 files changed

Lines changed: 158 additions & 0 deletions

File tree

qwen3_embed/rerank/cross_encoder/custom_text_cross_encoder.py

Lines changed: 48 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1,5 +1,8 @@
1+
from typing import Any
2+
13
from qwen3_embed.common.model_description import BaseModelDescription
24
from qwen3_embed.rerank.cross_encoder.onnx_text_cross_encoder import OnnxTextCrossEncoder
5+
from qwen3_embed.rerank.cross_encoder.onnx_text_model import TextRerankerWorker
36

47

58
class CustomTextCrossEncoder(OnnxTextCrossEncoder):
@@ -16,3 +19,48 @@ def add_model(
1619
) -> None:
1720
cls._clear_model_cache()
1821
cls.SUPPORTED_MODELS.append(model_description)
22+
23+
@classmethod
24+
def _export_registry(cls) -> list[BaseModelDescription]:
25+
# Snapshot the runtime registry so it can be pickled into spawned
26+
# worker processes, which start with a fresh interpreter and an empty
27+
# SUPPORTED_MODELS. BaseModelDescription is a frozen dataclass (picklable).
28+
return list(cls.SUPPORTED_MODELS)
29+
30+
@classmethod
31+
def _import_registry(cls, payload: list[BaseModelDescription]) -> None:
32+
# Re-register custom models in a worker process, idempotently (same id
33+
# imported twice must not create a duplicate entry).
34+
existing = {m.model.lower() for m in cls.SUPPORTED_MODELS}
35+
for desc in payload:
36+
if desc.model.lower() not in existing:
37+
cls.SUPPORTED_MODELS.append(desc)
38+
existing.add(desc.model.lower())
39+
40+
@classmethod
41+
def _get_worker_class(cls) -> type["CustomTextCrossEncoderWorker"]:
42+
return CustomTextCrossEncoderWorker
43+
44+
def _extra_worker_params(self) -> dict[str, Any]:
45+
# Propagate the runtime registry so spawned workers (fresh interpreters
46+
# with an empty SUPPORTED_MODELS) can resolve + re-register this custom
47+
# model instead of raising "Model ... not supported".
48+
return {"custom_registry": self._export_registry()}
49+
50+
51+
class CustomTextCrossEncoderWorker(TextRerankerWorker):
52+
def init_embedding(
53+
self,
54+
model_name: str,
55+
cache_dir: str,
56+
**kwargs: Any,
57+
) -> CustomTextCrossEncoder:
58+
registry = kwargs.pop("custom_registry", None)
59+
if registry is not None:
60+
CustomTextCrossEncoder._import_registry(registry)
61+
return CustomTextCrossEncoder(
62+
model_name=model_name,
63+
cache_dir=cache_dir,
64+
threads=1,
65+
**kwargs,
66+
)

qwen3_embed/rerank/cross_encoder/onnx_text_model.py

Lines changed: 13 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -24,6 +24,14 @@ class OnnxCrossEncoderModel(OnnxModel[float]):
2424
def _get_worker_class(cls) -> type["TextRerankerWorker"]:
2525
raise NotImplementedError("Subclasses must implement this method")
2626

27+
def _extra_worker_params(self) -> dict[str, Any]:
28+
"""Extra kwargs injected into each spawned worker's init.
29+
30+
Subclasses override this to carry state (e.g. a runtime-registered custom
31+
model registry) into worker processes that start with a fresh interpreter.
32+
"""
33+
return {}
34+
2735
def tokenize(self, pairs: list[tuple[str, str]], **_: Any) -> list[Encoding]:
2836
assert self.tokenizer is not None
2937
return self.tokenizer.encode_batch(pairs)
@@ -110,6 +118,11 @@ def _rerank_pairs(
110118
if extra_session_options is not None:
111119
params.update(extra_session_options)
112120

121+
# Carry per-instance worker state (e.g. a runtime custom-model
122+
# registry) into the spawned workers, which start with a fresh
123+
# interpreter + empty registry. Mirrors the embedding pool path.
124+
params.update(self._extra_worker_params())
125+
113126
pool = ParallelWorkerPool(
114127
worker=self._get_worker_class(),
115128
config=PoolConfig(
Lines changed: 97 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,97 @@
1+
"""F5: custom reranker registry must propagate to spawned multiprocessing workers.
2+
3+
The embedding side already does this (CustomTextEmbedding._export_registry /
4+
_import_registry / _extra_worker_params + the parallel pool params.update). The
5+
reranker side did not, so TextCrossEncoder(custom).rerank_pairs(..., parallel=N)
6+
with input >= batch_size raised "Model ... not supported" in the spawned worker
7+
(fresh interpreter with an empty registry). These tests guard the mirror fix.
8+
"""
9+
10+
import sys
11+
12+
import pytest
13+
14+
from qwen3_embed.common.model_description import BaseModelDescription, ModelSource
15+
from qwen3_embed.rerank.cross_encoder.custom_text_cross_encoder import (
16+
CustomTextCrossEncoder,
17+
CustomTextCrossEncoderWorker,
18+
)
19+
20+
21+
def _clear() -> None:
22+
CustomTextCrossEncoder.SUPPORTED_MODELS.clear()
23+
24+
25+
def test_custom_reranker_registry_survives_serialization():
26+
_clear()
27+
CustomTextCrossEncoder.add_model(
28+
BaseModelDescription(model="Org/My-Reranker", sources=ModelSource(hf="Org/My-Reranker"))
29+
)
30+
payload = CustomTextCrossEncoder._export_registry()
31+
_clear()
32+
assert CustomTextCrossEncoder.SUPPORTED_MODELS == [] # fresh-worker simulation
33+
CustomTextCrossEncoder._import_registry(payload)
34+
try:
35+
models = [m.model for m in CustomTextCrossEncoder._list_supported_models()]
36+
assert "Org/My-Reranker" in models
37+
finally:
38+
_clear()
39+
40+
41+
def test_custom_reranker_import_registry_is_idempotent():
42+
_clear()
43+
desc = BaseModelDescription(model="Org/Dup", sources=ModelSource(hf="Org/Dup"))
44+
CustomTextCrossEncoder.add_model(desc)
45+
payload = CustomTextCrossEncoder._export_registry()
46+
CustomTextCrossEncoder._import_registry(payload) # re-import same payload
47+
try:
48+
ids = [m.model for m in CustomTextCrossEncoder._list_supported_models()]
49+
assert ids.count("Org/Dup") == 1 # no duplicate entry
50+
finally:
51+
_clear()
52+
53+
54+
def test_custom_reranker_extra_worker_params_carries_registry():
55+
_clear()
56+
CustomTextCrossEncoder.add_model(
57+
BaseModelDescription(model="Org/R", sources=ModelSource(hf="Org/R"))
58+
)
59+
try:
60+
enc = CustomTextCrossEncoder.__new__(CustomTextCrossEncoder) # skip ONNX load
61+
params = enc._extra_worker_params()
62+
assert "custom_registry" in params
63+
assert any(d.model == "Org/R" for d in params["custom_registry"])
64+
finally:
65+
_clear()
66+
67+
68+
def test_custom_reranker_uses_custom_worker_class():
69+
# The base OnnxTextCrossEncoder worker would construct a plain (non-custom)
70+
# cross-encoder + never import the registry. The custom subclass must use a
71+
# worker that re-registers the custom model in the spawned process.
72+
assert CustomTextCrossEncoder._get_worker_class() is CustomTextCrossEncoderWorker
73+
74+
75+
@pytest.mark.integration
76+
@pytest.mark.skipif(sys.platform == "win32", reason="multiprocessing spawn deadlock on Windows")
77+
def test_custom_reranker_parallel_resolves_in_workers():
78+
"""Register a custom reranker, rerank_pairs with parallel=2 + batch_size=1
79+
(forces the worker-pool path), assert no 'not supported' error + scores."""
80+
_clear()
81+
model_id = "Org/Custom-Reranker"
82+
CustomTextCrossEncoder.add_model(
83+
BaseModelDescription(
84+
model=model_id,
85+
sources=ModelSource(hf="n24q02m/Qwen3-Reranker-0.6B-ONNX"),
86+
model_file="onnx/model_quantized.onnx",
87+
)
88+
)
89+
try:
90+
from qwen3_embed.rerank.cross_encoder.text_cross_encoder import TextCrossEncoder
91+
92+
ce = TextCrossEncoder(model_name=model_id)
93+
pairs = [("q", "d1"), ("q", "d2"), ("q", "d3")]
94+
scores = list(ce.rerank_pairs(pairs, batch_size=1, parallel=2))
95+
assert len(scores) == 3
96+
finally:
97+
_clear()

0 commit comments

Comments
 (0)