Compare commits

...
2 changed files with 112 additions and 3 deletions
+56 -3
View File
@@ -3,6 +3,7 @@
from typing import List from typing import List
from urllib.parse import unquote, urlparse
import numpy as np import numpy as np
@@ -125,9 +126,20 @@ class InstructorEmbeddingFunction(TextEmbeddingFunction):
@weak_lru(maxsize=1) @weak_lru(maxsize=1)
def get_model(self): def get_model(self):
instructor_embedding = attempt_import_or_raise( huggingface_hub = attempt_import_or_raise("huggingface_hub", "huggingface-hub")
"InstructorEmbedding", "InstructorEmbedding" 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)
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") torch = attempt_import_or_raise("torch", "torch")
model = instructor_embedding.INSTRUCTOR(self.name) model = instructor_embedding.INSTRUCTOR(self.name)
@@ -140,3 +152,44 @@ class InstructorEmbeddingFunction(TextEmbeddingFunction):
model, {torch.nn.Linear}, dtype=torch.qint8 model, {torch.nn.Linear}, dtype=torch.qint8
) )
return model return model
def _cached_download(huggingface_hub):
"""Provide the legacy download API used by sentence-transformers 2.2.x."""
def cached_download(
*,
url,
cache_dir=None,
force_filename=None,
library_name=None,
library_version=None,
user_agent=None,
use_auth_token=None,
**_,
):
path = urlparse(url).path.lstrip("/")
try:
repo_id, resolved_path = path.split("/resolve/", maxsplit=1)
revision, filename = resolved_path.split("/", maxsplit=1)
except ValueError as err:
raise ValueError(f"Unsupported Hugging Face Hub URL: {url}") from err
repo_id = unquote(repo_id)
revision = unquote(revision)
filename = unquote(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,
revision=revision,
local_dir=cache_dir,
library_name=library_name,
library_version=library_version,
user_agent=user_agent,
token=use_auth_token,
)
return cached_download
+56
View File
@@ -1,8 +1,11 @@
# SPDX-License-Identifier: Apache-2.0 # SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright The LanceDB Authors # SPDX-FileCopyrightText: Copyright The LanceDB Authors
import ntpath
import os import os
import pickle import pickle
import sys
from types import ModuleType
from typing import List, Optional, Union from typing import List, Optional, Union
from unittest.mock import MagicMock, patch from unittest.mock import MagicMock, patch
@@ -522,6 +525,59 @@ def test_embedding_function_safe_model_dump(embedding_type):
) )
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 = 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"
)
embedding = InstructorEmbeddingFunction.create(show_progress_bar=False)
instructor_model = embedding.get_model()
assert instructor_model.name == "hkunlp/instructor-base"
assert not hasattr(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=ntpath.join("1_Pooling", "config.json"),
library_name="sentence-transformers",
library_version="2.2.2",
use_auth_token="token",
)
assert path == "/cache/1_Pooling/config.json"
hub_download.assert_called_once_with(
repo_id="hkunlp/instructor-base",
filename="1_Pooling/config.json",
revision="abc123",
local_dir="/cache",
library_name="sentence-transformers",
library_version="2.2.2",
user_agent=None,
token="token",
)
@patch("time.sleep") @patch("time.sleep")
def test_retry(mock_sleep): def test_retry(mock_sleep):
test_function = MagicMock(side_effect=[Exception] * 9 + ["result"]) test_function = MagicMock(side_effect=[Exception] * 9 + ["result"])