diff --git a/FlagEmbedding/inference/_model_mapping.py b/FlagEmbedding/inference/_model_mapping.py new file mode 100644 index 000000000..3a9ae524c --- /dev/null +++ b/FlagEmbedding/inference/_model_mapping.py @@ -0,0 +1,22 @@ +import os +from typing import Collection + + +def resolve_model_name(model_name_or_path: str, model_names: Collection[str]) -> str: + """Resolve a model mapping key from a Hub ID or local model path.""" + model_path = os.path.normpath(model_name_or_path) + if os.path.basename(model_path).startswith("checkpoint-"): + model_path = os.path.dirname(model_path) + + for candidate in (model_path, os.path.basename(model_path)): + if candidate in model_names: + return candidate + + local_name = os.path.basename(model_path) + matching_names = [ + name for name in model_names if os.path.basename(name) == local_name + ] + if len(matching_names) == 1: + return matching_names[0] + + return local_name diff --git a/FlagEmbedding/inference/auto_embedder.py b/FlagEmbedding/inference/auto_embedder.py index 6fad0330f..e7d782b1c 100644 --- a/FlagEmbedding/inference/auto_embedder.py +++ b/FlagEmbedding/inference/auto_embedder.py @@ -1,10 +1,11 @@ -import os import logging -from typing import List, Union, Optional +from typing import List, Optional, Union +from FlagEmbedding.inference._model_mapping import resolve_model_name from FlagEmbedding.inference.embedder.model_mapping import ( + AUTO_EMBEDDER_MAPPING, + EMBEDDER_CLASS_MAPPING, EmbedderModelClass, - AUTO_EMBEDDER_MAPPING, EMBEDDER_CLASS_MAPPING ) logger = logging.getLogger(__name__) @@ -59,10 +60,6 @@ def from_finetuned( Returns: AbsEmbedder: The model class to load model, which is child class of :class:`AbsEmbedder`. """ - model_name = os.path.basename(model_name_or_path) - if model_name.startswith("checkpoint-"): - model_name = os.path.basename(os.path.dirname(model_name_or_path)) - if model_class is not None: _model_class = EMBEDDER_CLASS_MAPPING[EmbedderModelClass(model_class)] if pooling_method is None: @@ -81,6 +78,7 @@ def from_finetuned( f"`query_instruction_format` is not specified, set to default value '{query_instruction_format}'." ) else: + model_name = resolve_model_name(model_name_or_path, AUTO_EMBEDDER_MAPPING) if model_name not in AUTO_EMBEDDER_MAPPING: raise ValueError( f"Model name '{model_name}' not found in the model mapping. You can pull request to add the model to " diff --git a/FlagEmbedding/inference/auto_reranker.py b/FlagEmbedding/inference/auto_reranker.py index a8b2d699e..98e567032 100644 --- a/FlagEmbedding/inference/auto_reranker.py +++ b/FlagEmbedding/inference/auto_reranker.py @@ -1,11 +1,11 @@ -import os import logging -from typing import Union, Optional +from typing import Optional, Union +from FlagEmbedding.inference._model_mapping import resolve_model_name from FlagEmbedding.inference.reranker.model_mapping import ( - RerankerModelClass, + AUTO_RERANKER_MAPPING, RERANKER_CLASS_MAPPING, - AUTO_RERANKER_MAPPING + RerankerModelClass, ) logger = logging.getLogger(__name__) @@ -46,10 +46,6 @@ def from_finetuned( Returns: AbsReranker: The reranker class to load model, which is child class of :class:`AbsReranker`. """ - model_name = os.path.basename(model_name_or_path) - if model_name.startswith("checkpoint-"): - model_name = os.path.basename(os.path.dirname(model_name_or_path)) - if model_class is not None: _model_class = RERANKER_CLASS_MAPPING[RerankerModelClass(model_class)] if trust_remote_code is None: @@ -58,6 +54,7 @@ def from_finetuned( f"`trust_remote_code` is not specified, set to default value '{trust_remote_code}'." ) else: + model_name = resolve_model_name(model_name_or_path, AUTO_RERANKER_MAPPING) if model_name not in AUTO_RERANKER_MAPPING: raise ValueError( f"Model name '{model_name}' not found in the model mapping. You can pull request to add the model to " diff --git a/tests/test_auto_model_resolution.py b/tests/test_auto_model_resolution.py new file mode 100644 index 000000000..76ce19321 --- /dev/null +++ b/tests/test_auto_model_resolution.py @@ -0,0 +1,70 @@ +import os + +import pytest + +from FlagEmbedding.inference.auto_embedder import FlagAutoModel +from FlagEmbedding.inference.auto_reranker import FlagAutoReranker +from FlagEmbedding.inference.embedder.model_mapping import AUTO_EMBEDDER_MAPPING +from FlagEmbedding.inference.reranker.model_mapping import AUTO_RERANKER_MAPPING + + +class RecordingModel: + def __init__(self, model_name_or_path, **kwargs): + self.model_name_or_path = model_name_or_path + self.kwargs = kwargs + + +@pytest.mark.parametrize( + "model_name", + [ + "jinaai/jina-reranker-v2-base-multilingual", + "Alibaba-NLP/gte-multilingual-reranker-base", + "maidalun1020/bce-reranker-base_v1", + "jinaai/jina-reranker-v1-turbo-en", + ], +) +def test_auto_reranker_resolves_namespaced_model_ids(monkeypatch, model_name): + monkeypatch.setattr( + AUTO_RERANKER_MAPPING[model_name], "model_class", RecordingModel + ) + + model = FlagAutoReranker.from_finetuned(model_name) + + assert model.model_name_or_path == model_name + + +def test_auto_reranker_resolves_local_model_directory(monkeypatch, tmp_path): + model_name = "jinaai/jina-reranker-v2-base-multilingual" + monkeypatch.setattr( + AUTO_RERANKER_MAPPING[model_name], "model_class", RecordingModel + ) + model_path = tmp_path / "jina-reranker-v2-base-multilingual" + + model = FlagAutoReranker.from_finetuned(str(model_path)) + + assert model.model_name_or_path == str(model_path) + + +def test_auto_reranker_resolves_local_checkpoint_directory(monkeypatch, tmp_path): + model_name = "jinaai/jina-reranker-v2-base-multilingual" + monkeypatch.setattr( + AUTO_RERANKER_MAPPING[model_name], "model_class", RecordingModel + ) + checkpoint_path = tmp_path / "jina-reranker-v2-base-multilingual" / "checkpoint-100" + + model = FlagAutoReranker.from_finetuned(str(checkpoint_path)) + + assert model.model_name_or_path == str(checkpoint_path) + + +def test_auto_embedder_resolves_path_with_trailing_separator(monkeypatch, tmp_path): + model_name = "bge-base-en-v1.5" + monkeypatch.setattr( + AUTO_EMBEDDER_MAPPING[model_name], "model_class", RecordingModel + ) + model_path = tmp_path / model_name + model_path_with_separator = f"{model_path}{os.sep}" + + model = FlagAutoModel.from_finetuned(model_path_with_separator) + + assert model.model_name_or_path == model_path_with_separator