mirror of
https://github.com/lancedb/lancedb.git
synced 2026-08-18 12:08:35 +00:00
fix(python): instruct dimension probe for instructor embeddings (#3874)
## Summary - pass an Instructor-compatible `[instruction, text]` pair when detecting embedding dimensions - add a regression test that verifies the dimension probe uses the configured source instruction ## Root cause `InstructorEmbeddingFunction.ndims()` encoded a bare string even though Instructor models require instruction/text pairs. With affected `sentence-transformers` versions, the bare input omitted `instruction_mask` and raised `KeyError` while defining the LanceDB schema. ## Validation - `uv run --extra tests pytest python/tests/test_embeddings.py -q` (`14 passed, 9 skipped`) - `uv run --project python --extra tests --extra dev ruff format --check python/python/lancedb/embeddings/instructor.py python/python/tests/test_embeddings.py` - `uv run --project python --extra tests --extra dev ruff check .` Fixes #2041 <!-- lance-gatekeeper-fix:v1 agent=4b05e0d9f3eef17bccfb446e788294f4 generation=1 --> Co-authored-by: Gatefixer <313497061+lancedb-gatefixer[bot]@users.noreply.github.com>
This commit is contained in:
committed by
GitHub
parent
123c921c4f
commit
f1f34dfdd3
@@ -101,8 +101,7 @@ class InstructorEmbeddingFunction(TextEmbeddingFunction):
|
||||
|
||||
@weak_lru(maxsize=1)
|
||||
def ndims(self):
|
||||
model = self.get_model()
|
||||
return model.encode("foo").shape[0]
|
||||
return len(self.generate_embeddings([[self.source_instruction, "foo"]])[0])
|
||||
|
||||
def compute_query_embeddings(self, query: str, *args, **kwargs) -> List[np.array]:
|
||||
return self.generate_embeddings([[self.query_instruction, query]])
|
||||
|
||||
@@ -64,6 +64,23 @@ def test_embedding_function(tmp_path):
|
||||
assert np.allclose(actual, expected)
|
||||
|
||||
|
||||
def test_instructor_ndims_uses_instruction():
|
||||
instructor = get_registry().get("instructor").create()
|
||||
model = MagicMock()
|
||||
model.encode.return_value = np.zeros((1, 384))
|
||||
|
||||
with patch.object(type(instructor), "get_model", return_value=model):
|
||||
assert instructor.ndims() == 384
|
||||
|
||||
model.encode.assert_called_once_with(
|
||||
[[instructor.source_instruction, "foo"]],
|
||||
batch_size=instructor.batch_size,
|
||||
show_progress_bar=instructor.show_progress_bar,
|
||||
normalize_embeddings=instructor.normalize_embeddings,
|
||||
device=instructor.device,
|
||||
)
|
||||
|
||||
|
||||
def test_embedding_function_variables():
|
||||
@register("variable-testing")
|
||||
class VariableTestingFunction(TextEmbeddingFunction):
|
||||
|
||||
Reference in New Issue
Block a user