From 16347276b87a574719f60bd2f20035b0c9c4e1a6 Mon Sep 17 00:00:00 2001 From: svonava Date: Wed, 7 Oct 2026 22:01:09 -0700 Subject: [PATCH] fix(server): apply default_instruction in the remaining flash embedders RoPEFlashAdapter, BertFlashAdapter, ModernBERTFlashAdapter, NomicFlashAdapter, GTESparseFlashAdapter and SPLADEFlashAdapter never read the `default_instruction` runtime option, so a query without an explicit instruction was formatted with an empty instruction slot. For shipped configs this affected NovaSearch/stella_en_400M_v5 on CUDA, whose SentenceTransformer CPU fallback already applied the instruction. Add `_utils.resolve_query_instruction` and use it in these adapters and in Qwen2FlashAdapter (#621): an explicit instruction (including "") wins, queries otherwise get `default_instruction`, documents get none. --- .../src/sie_server/adapters/_utils.py | 18 ++ .../adapters/bert_flash/__init__.py | 6 +- .../adapters/gte_sparse_flash/__init__.py | 12 +- .../adapters/modernbert_flash/__init__.py | 6 +- .../adapters/nomic_flash/__init__.py | 12 +- .../adapters/qwen2_flash/__init__.py | 14 +- .../adapters/rope_flash/__init__.py | 12 +- .../adapters/splade_flash/adapter.py | 4 +- .../test_flash_default_instruction.py | 195 ++++++++++++++++++ 9 files changed, 257 insertions(+), 22 deletions(-) create mode 100644 packages/sie_server/tests/adapters/test_flash_default_instruction.py diff --git a/packages/sie_server/src/sie_server/adapters/_utils.py b/packages/sie_server/src/sie_server/adapters/_utils.py index fd871d83d..c72d13c5f 100644 --- a/packages/sie_server/src/sie_server/adapters/_utils.py +++ b/packages/sie_server/src/sie_server/adapters/_utils.py @@ -163,6 +163,24 @@ def resolve_embedding_options( ) +def resolve_query_instruction( + instruction: str | None, + options: dict[str, Any] | None, + *, + is_query: bool, +) -> str | None: + """Return the instruction to format into the texts for ``extract_texts``. + + A request instruction always wins, including an explicit ``""``. When the + request gives none, queries fall back to the ``default_instruction`` + runtime option (the profile default, or a request override merged into + ``options``). Documents never receive the default. + """ + if instruction is not None or not is_query: + return instruction + return (options or {}).get("default_instruction") + + # --------------------------------------------------------------------------- # Score-pair grouping (shared by ColBERT-family adapters) # --------------------------------------------------------------------------- diff --git a/packages/sie_server/src/sie_server/adapters/bert_flash/__init__.py b/packages/sie_server/src/sie_server/adapters/bert_flash/__init__.py index 42b5c105a..7008d5e7a 100644 --- a/packages/sie_server/src/sie_server/adapters/bert_flash/__init__.py +++ b/packages/sie_server/src/sie_server/adapters/bert_flash/__init__.py @@ -30,6 +30,7 @@ from sie_server.adapters._utils import ( extract_texts, resolve_embedding_options, + resolve_query_instruction, validate_output_types, ) from sie_server.core.inference_output import EncodeOutput @@ -212,7 +213,8 @@ def encode( Args: items: List of items to encode. output_types: Which outputs to compute (only "dense" supported). - instruction: Optional instruction prefix. + instruction: Optional instruction prefix. For queries, ``None`` + falls back to the ``default_instruction`` runtime option. is_query: Whether items are queries (affects template selection). prepared_items: Not used by this adapter. @@ -236,7 +238,7 @@ def encode( texts = extract_texts( items, - instruction, + resolve_query_instruction(instruction, options, is_query=is_query), is_query=is_query, query_template=query_template, doc_template=doc_template, diff --git a/packages/sie_server/src/sie_server/adapters/gte_sparse_flash/__init__.py b/packages/sie_server/src/sie_server/adapters/gte_sparse_flash/__init__.py index 5d271e52b..ff19e1a2d 100644 --- a/packages/sie_server/src/sie_server/adapters/gte_sparse_flash/__init__.py +++ b/packages/sie_server/src/sie_server/adapters/gte_sparse_flash/__init__.py @@ -10,7 +10,12 @@ from sie_server.adapters._flash_base import FlashBaseAdapter from sie_server.adapters._spec import AdapterSpec from sie_server.adapters._types import ERR_NOT_LOADED, ERR_REQUIRES_TEXT, ComputePrecision -from sie_server.adapters._utils import apply_rotary_pos_emb, extract_texts, validate_output_types +from sie_server.adapters._utils import ( + apply_rotary_pos_emb, + extract_texts, + resolve_query_instruction, + validate_output_types, +) from sie_server.adapters.peft_lora_mixin import PEFTLoRAMixin from sie_server.core.inference_output import EncodeOutput, SparseVector from sie_server.types.inputs import Item @@ -201,7 +206,8 @@ def encode( Args: items: List of items to encode. output_types: Which outputs to compute (only "sparse" supported). - instruction: Optional instruction prefix. + instruction: Optional instruction prefix. For queries, ``None`` + falls back to the ``default_instruction`` runtime option. is_query: Whether items are queries (affects template selection). prepared_items: Not used by this adapter. @@ -221,7 +227,7 @@ def encode( texts = extract_texts( items, - instruction, + resolve_query_instruction(instruction, options, is_query=is_query), is_query=is_query, query_template=query_template, doc_template=doc_template, diff --git a/packages/sie_server/src/sie_server/adapters/modernbert_flash/__init__.py b/packages/sie_server/src/sie_server/adapters/modernbert_flash/__init__.py index 9debb25af..c311d6366 100644 --- a/packages/sie_server/src/sie_server/adapters/modernbert_flash/__init__.py +++ b/packages/sie_server/src/sie_server/adapters/modernbert_flash/__init__.py @@ -25,6 +25,7 @@ from sie_server.adapters._utils import ( extract_texts, resolve_embedding_options, + resolve_query_instruction, validate_output_types, ) from sie_server.adapters.peft_lora_mixin import PEFTLoRAMixin @@ -219,7 +220,8 @@ def encode( Args: items: List of items to encode. output_types: Which outputs to compute (only "dense" supported). - instruction: Optional instruction prefix. + instruction: Optional instruction prefix. For queries, ``None`` + falls back to the ``default_instruction`` runtime option. is_query: Whether items are queries (affects template selection). prepared_items: Not used by this adapter. options: Runtime options for normalize, pooling, templates. @@ -244,7 +246,7 @@ def encode( texts = extract_texts( items, - instruction, + resolve_query_instruction(instruction, options, is_query=is_query), is_query=is_query, query_template=query_template, doc_template=doc_template, diff --git a/packages/sie_server/src/sie_server/adapters/nomic_flash/__init__.py b/packages/sie_server/src/sie_server/adapters/nomic_flash/__init__.py index 658ef7953..583b7f17d 100644 --- a/packages/sie_server/src/sie_server/adapters/nomic_flash/__init__.py +++ b/packages/sie_server/src/sie_server/adapters/nomic_flash/__init__.py @@ -12,7 +12,12 @@ from sie_server.adapters._flash_pack import build_position_ids, mean_pool_packed from sie_server.adapters._spec import AdapterSpec from sie_server.adapters._types import ERR_NOT_LOADED, ComputePrecision, PoolingStrategy -from sie_server.adapters._utils import apply_rotary_pos_emb, extract_texts, validate_output_types +from sie_server.adapters._utils import ( + apply_rotary_pos_emb, + extract_texts, + resolve_query_instruction, + validate_output_types, +) from sie_server.core.inference_output import EncodeOutput from sie_server.types.inputs import Item @@ -236,7 +241,8 @@ def encode( Args: items: List of items to encode. output_types: Which outputs to compute (only "dense" supported). - instruction: Optional instruction prefix (unused, template-based). + instruction: Optional instruction prefix. For queries, ``None`` + falls back to the ``default_instruction`` runtime option. is_query: Whether items are queries (affects template selection). prepared_items: Not used by this adapter. @@ -256,7 +262,7 @@ def encode( texts = extract_texts( items, - instruction, + resolve_query_instruction(instruction, options, is_query=is_query), is_query=is_query, query_template=query_template, doc_template=doc_template, diff --git a/packages/sie_server/src/sie_server/adapters/qwen2_flash/__init__.py b/packages/sie_server/src/sie_server/adapters/qwen2_flash/__init__.py index eebfe9c2f..8fe1f0dac 100644 --- a/packages/sie_server/src/sie_server/adapters/qwen2_flash/__init__.py +++ b/packages/sie_server/src/sie_server/adapters/qwen2_flash/__init__.py @@ -12,7 +12,12 @@ from sie_server.adapters._flash_pack import build_position_ids, mean_pool_packed from sie_server.adapters._spec import AdapterSpec from sie_server.adapters._types import ERR_NOT_LOADED, ComputePrecision, PoolingStrategy -from sie_server.adapters._utils import apply_rotary_pos_emb, extract_texts, validate_output_types +from sie_server.adapters._utils import ( + apply_rotary_pos_emb, + extract_texts, + resolve_query_instruction, + validate_output_types, +) from sie_server.adapters.peft_lora_mixin import PEFTLoRAMixin from sie_server.core.inference_output import EncodeOutput from sie_server.types.inputs import Item @@ -270,15 +275,10 @@ def encode( doc_template = opts.get("doc_template", self._doc_template) normalize = opts.get("normalize", self._normalize) pooling = opts.get("pooling", self._pooling) - # The profile's default_instruction is query-only and fills in only - # when the request gave no instruction; an explicit "" is kept. - effective_instruction = instruction - if effective_instruction is None and is_query: - effective_instruction = opts.get("default_instruction") texts = extract_texts( items, - effective_instruction, + resolve_query_instruction(instruction, opts, is_query=is_query), is_query=is_query, query_template=query_template, doc_template=doc_template, diff --git a/packages/sie_server/src/sie_server/adapters/rope_flash/__init__.py b/packages/sie_server/src/sie_server/adapters/rope_flash/__init__.py index 77735b600..ecc7a95e8 100644 --- a/packages/sie_server/src/sie_server/adapters/rope_flash/__init__.py +++ b/packages/sie_server/src/sie_server/adapters/rope_flash/__init__.py @@ -11,7 +11,12 @@ from sie_server.adapters._flash_pack import mean_pool_packed from sie_server.adapters._spec import AdapterSpec from sie_server.adapters._types import ERR_NOT_LOADED, ComputePrecision, PoolingStrategy -from sie_server.adapters._utils import apply_rotary_pos_emb, extract_texts, validate_output_types +from sie_server.adapters._utils import ( + apply_rotary_pos_emb, + extract_texts, + resolve_query_instruction, + validate_output_types, +) from sie_server.adapters.peft_lora_mixin import PEFTLoRAMixin from sie_server.core.inference_output import EncodeOutput from sie_server.types.inputs import Item @@ -172,7 +177,8 @@ def encode( Args: items: List of items to encode. output_types: Which outputs to compute (only "dense" supported). - instruction: Optional instruction prefix. + instruction: Optional instruction prefix. For queries, ``None`` + falls back to the ``default_instruction`` runtime option. is_query: Whether items are queries (affects template selection). prepared_items: Not used by this adapter. @@ -194,7 +200,7 @@ def encode( texts = extract_texts( items, - instruction, + resolve_query_instruction(instruction, options, is_query=is_query), is_query=is_query, query_template=query_template, doc_template=doc_template, diff --git a/packages/sie_server/src/sie_server/adapters/splade_flash/adapter.py b/packages/sie_server/src/sie_server/adapters/splade_flash/adapter.py index abe7268dd..22bac8991 100644 --- a/packages/sie_server/src/sie_server/adapters/splade_flash/adapter.py +++ b/packages/sie_server/src/sie_server/adapters/splade_flash/adapter.py @@ -14,7 +14,7 @@ from sie_server.adapters._flash_pack import build_position_ids from sie_server.adapters._spec import AdapterSpec from sie_server.adapters._types import ERR_NOT_LOADED, ComputePrecision -from sie_server.adapters._utils import extract_texts, validate_output_types +from sie_server.adapters._utils import extract_texts, resolve_query_instruction, validate_output_types from sie_server.adapters.base import ModelAdapter from sie_server.adapters.peft_lora_mixin import PEFTLoRAMixin from sie_server.core.inference_output import EncodeOutput, SparseVector @@ -230,7 +230,7 @@ def encode( texts = extract_texts( items, - instruction, + resolve_query_instruction(instruction, options, is_query=is_query), is_query=is_query, query_template=query_template, doc_template=doc_template, diff --git a/packages/sie_server/tests/adapters/test_flash_default_instruction.py b/packages/sie_server/tests/adapters/test_flash_default_instruction.py new file mode 100644 index 000000000..429b176b5 --- /dev/null +++ b/packages/sie_server/tests/adapters/test_flash_default_instruction.py @@ -0,0 +1,195 @@ +"""CPU unit tests for query instructions in the extract_texts-based flash adapters. + +The flash kernels need CUDA, so these tests stop each adapter right after it +formats its texts and check the exact strings it would tokenize. +""" + +from __future__ import annotations + +import importlib +from pathlib import Path +from typing import Any + +import pytest +from sie_server.adapters import _utils +from sie_server.adapters._utils import resolve_query_instruction +from sie_server.core.loader import load_model_config +from sie_server.core.runtime_options import merge_runtime_options +from sie_server.types.inputs import Item + +_MODELS_DIR = Path(__file__).resolve().parents[2] / "models" +_TEMPLATE = "Instruct: {instruction}\nQuery: {text}" +_DEFAULT = "Given a web search query, retrieve relevant passages that answer the query." + +# (module, class, output type) for every flash adapter that formats texts with +# the shared extract_texts helper and resolves its instruction through +# resolve_query_instruction. +_ADAPTERS = [ + ("bert_flash", "BertFlashAdapter", "dense"), + ("modernbert_flash", "ModernBERTFlashAdapter", "dense"), + ("nomic_flash", "NomicFlashAdapter", "dense"), + ("qwen2_flash", "Qwen2FlashAdapter", "dense"), + ("rope_flash", "RoPEFlashAdapter", "dense"), + ("gte_sparse_flash", "GTESparseFlashAdapter", "sparse"), + ("splade_flash.adapter", "SPLADEFlashAdapter", "sparse"), +] + + +class _FormattedTexts(Exception): # noqa: N818 - control-flow signal, not an error + def __init__(self, texts: list[str]) -> None: + super().__init__(texts) + self.texts = texts + + +def _formatted_texts( + monkeypatch: pytest.MonkeyPatch, + adapter_spec: tuple[str, str, str], + texts: list[str], + *, + is_query: bool, + instruction: str | None = None, + options: dict[str, Any] | None = None, +) -> list[str]: + module_name, class_name, output_type = adapter_spec + module = importlib.import_module(f"sie_server.adapters.{module_name}") + + def stop_after_formatting(*args: Any, **kwargs: Any) -> list[str]: + raise _FormattedTexts(_utils.extract_texts(*args, **kwargs)) + + monkeypatch.setattr(module, "extract_texts", stop_after_formatting) + adapter = getattr(module, class_name)("unused") + # Satisfy each adapter's loaded check; NomicFlashAdapter keeps _layers, not _model. + for loaded_attr in ("_model", "_layers", "_tokenizer"): + if hasattr(adapter, loaded_attr): + monkeypatch.setattr(adapter, loaded_attr, object()) + + with pytest.raises(_FormattedTexts) as formatted: + adapter.encode( + [Item(text=text) for text in texts], + [output_type], + instruction=instruction, + is_query=is_query, + options=options, + ) + return formatted.value.texts + + +_ADAPTER_IDS = [class_name for _, class_name, _ in _ADAPTERS] + + +@pytest.mark.parametrize("adapter_spec", _ADAPTERS, ids=_ADAPTER_IDS) +def test_query_uses_default_instruction(monkeypatch: pytest.MonkeyPatch, adapter_spec: tuple[str, str, str]) -> None: + texts = _formatted_texts( + monkeypatch, + adapter_spec, + ["what is sie?", "flash attention"], + is_query=True, + options={"query_template": _TEMPLATE, "default_instruction": _DEFAULT}, + ) + + assert texts == [ + f"Instruct: {_DEFAULT}\nQuery: what is sie?", + f"Instruct: {_DEFAULT}\nQuery: flash attention", + ] + + +@pytest.mark.parametrize("adapter_spec", _ADAPTERS, ids=_ADAPTER_IDS) +def test_explicit_instruction_wins(monkeypatch: pytest.MonkeyPatch, adapter_spec: tuple[str, str, str]) -> None: + texts = _formatted_texts( + monkeypatch, + adapter_spec, + ["what is sie?"], + is_query=True, + instruction="Find related questions", + options={"query_template": _TEMPLATE, "default_instruction": _DEFAULT}, + ) + + assert texts == ["Instruct: Find related questions\nQuery: what is sie?"] + + +@pytest.mark.parametrize("adapter_spec", _ADAPTERS, ids=_ADAPTER_IDS) +def test_explicit_empty_instruction_is_kept( + monkeypatch: pytest.MonkeyPatch, adapter_spec: tuple[str, str, str] +) -> None: + texts = _formatted_texts( + monkeypatch, + adapter_spec, + ["what is sie?"], + is_query=True, + instruction="", + options={"query_template": _TEMPLATE, "default_instruction": _DEFAULT}, + ) + + assert texts == ["Instruct: \nQuery: what is sie?"] + + +@pytest.mark.parametrize("adapter_spec", _ADAPTERS, ids=_ADAPTER_IDS) +def test_documents_get_no_default_instruction( + monkeypatch: pytest.MonkeyPatch, + adapter_spec: tuple[str, str, str], +) -> None: + texts = _formatted_texts( + monkeypatch, + adapter_spec, + ["SIE serves embedding models."], + is_query=False, + options={ + "query_template": _TEMPLATE, + "doc_template": "Instruct: {instruction}\nDocument: {text}", + "default_instruction": _DEFAULT, + }, + ) + + assert texts == ["Instruct: \nDocument: SIE serves embedding models."] + + +@pytest.mark.parametrize("adapter_spec", _ADAPTERS, ids=_ADAPTER_IDS) +def test_query_without_default_instruction_keeps_empty_slot( + monkeypatch: pytest.MonkeyPatch, + adapter_spec: tuple[str, str, str], +) -> None: + texts = _formatted_texts( + monkeypatch, + adapter_spec, + ["what is sie?"], + is_query=True, + options={"query_template": _TEMPLATE}, + ) + + assert texts == ["Instruct: \nQuery: what is sie?"] + + +def test_stella_400m_query_uses_profile_default_instruction(monkeypatch: pytest.MonkeyPatch) -> None: + config = load_model_config(_MODELS_DIR / "NovaSearch__stella_en_400M_v5.yaml") + assert config.resolve_profile("default").adapter_path == "sie_server.adapters.rope_flash:RoPEFlashAdapter" + + texts = _formatted_texts( + monkeypatch, + ("rope_flash", "RoPEFlashAdapter", "dense"), + ["what is sie?"], + is_query=True, + options=merge_runtime_options(config, {"is_query": True}), + ) + + assert texts == [f"Instruct: {_DEFAULT}\nQuery: what is sie?"] + + +@pytest.mark.parametrize( + ("instruction", "options", "is_query", "expected"), + [ + (None, {"default_instruction": "default"}, True, "default"), + ("explicit", {"default_instruction": "default"}, True, "explicit"), + ("", {"default_instruction": "default"}, True, ""), + (None, {"default_instruction": "default"}, False, None), + ("explicit", {"default_instruction": "default"}, False, "explicit"), + (None, {}, True, None), + (None, None, True, None), + ], +) +def test_resolve_query_instruction( + instruction: str | None, + options: dict[str, Any] | None, + is_query: bool, + expected: str | None, +) -> None: + assert resolve_query_instruction(instruction, options, is_query=is_query) == expected