fix(python): scope Instructor compatibility shim

This commit is contained in:
Gatefixer
2026-08-06 01:30:35 +00:00
parent bd779bb7d5
commit 88d8a69a99
2 changed files with 40 additions and 30 deletions
+14 -7
View File
@@ -127,12 +127,19 @@ class InstructorEmbeddingFunction(TextEmbeddingFunction):
@weak_lru(maxsize=1)
def get_model(self):
huggingface_hub = attempt_import_or_raise("huggingface_hub", "huggingface-hub")
if not hasattr(huggingface_hub, "cached_download"):
missing = object()
original_cached_download = getattr(huggingface_hub, "cached_download", missing)
if original_cached_download is missing:
huggingface_hub.cached_download = _cached_download(huggingface_hub)
instructor_embedding = attempt_import_or_raise(
"InstructorEmbedding", "InstructorEmbedding"
)
try:
instructor_embedding = attempt_import_or_raise(
"InstructorEmbedding", "InstructorEmbedding"
)
finally:
if original_cached_download is missing:
del huggingface_hub.cached_download
torch = attempt_import_or_raise("torch", "torch")
model = instructor_embedding.INSTRUCTOR(self.name)
@@ -171,9 +178,9 @@ def _cached_download(huggingface_hub):
repo_id = unquote(repo_id)
revision = unquote(revision)
filename = unquote(filename)
if force_filename is not None:
filename = force_filename
# sentence-transformers derives force_filename from this Hub path with
# os.path.join. Using the URL path beneath local_dir produces the same
# local destination without sending Windows separators to the Hub.
return huggingface_hub.hf_hub_download(
repo_id=repo_id,
filename=filename,
+26 -23
View File
@@ -1,9 +1,11 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright The LanceDB Authors
import ntpath
import os
import pickle
from types import SimpleNamespace
import sys
from types import ModuleType
from typing import List, Optional, Union
from unittest.mock import MagicMock, patch
@@ -523,40 +525,41 @@ def test_embedding_function_safe_model_dump(embedding_type):
)
def test_instructor_embedding_supports_huggingface_hub_without_cached_download():
def test_instructor_embedding_supports_huggingface_hub_without_cached_download(
tmp_path, monkeypatch
):
from lancedb.embeddings.instructor import InstructorEmbeddingFunction
hub_download = MagicMock(return_value="/cache/1_Pooling/config.json")
huggingface_hub = SimpleNamespace(hf_hub_download=hub_download)
instructor_model = MagicMock()
instructor_embedding = SimpleNamespace(
INSTRUCTOR=MagicMock(return_value=instructor_model)
huggingface_hub = ModuleType("huggingface_hub")
huggingface_hub.hf_hub_download = hub_download
torch = ModuleType("torch")
monkeypatch.setitem(sys.modules, "huggingface_hub", huggingface_hub)
monkeypatch.setitem(sys.modules, "torch", torch)
monkeypatch.delitem(sys.modules, "InstructorEmbedding", raising=False)
monkeypatch.syspath_prepend(str(tmp_path))
(tmp_path / "InstructorEmbedding.py").write_text(
"from huggingface_hub import cached_download\n\n"
"class INSTRUCTOR:\n"
" def __init__(self, name):\n"
" self.name = name\n"
)
def import_dependency(module, _mitigation):
if module == "huggingface_hub":
return huggingface_hub
if module == "InstructorEmbedding":
assert hasattr(huggingface_hub, "cached_download")
return instructor_embedding
if module == "torch":
return SimpleNamespace()
raise AssertionError(f"Unexpected import: {module}")
embedding = InstructorEmbeddingFunction.create(show_progress_bar=False)
instructor_model = embedding.get_model()
with patch(
"lancedb.embeddings.instructor.attempt_import_or_raise",
side_effect=import_dependency,
):
embedding = InstructorEmbeddingFunction.create(show_progress_bar=False)
assert embedding.get_model() is instructor_model
assert instructor_model.name == "hkunlp/instructor-base"
assert not hasattr(huggingface_hub, "cached_download")
path = huggingface_hub.cached_download(
instructor_embedding = sys.modules["InstructorEmbedding"]
path = instructor_embedding.cached_download(
url=(
"https://huggingface.co/hkunlp/instructor-base/resolve/abc123/"
"1_Pooling/config.json"
),
cache_dir="/cache",
force_filename="1_Pooling/config.json",
force_filename=ntpath.join("1_Pooling", "config.json"),
library_name="sentence-transformers",
library_version="2.2.2",
use_auth_token="token",