Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
22 changes: 22 additions & 0 deletions FlagEmbedding/inference/_model_mapping.py
Original file line number Diff line number Diff line change
@@ -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
12 changes: 5 additions & 7 deletions FlagEmbedding/inference/auto_embedder.py
Original file line number Diff line number Diff line change
@@ -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__)
Expand Down Expand Up @@ -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:
Expand All @@ -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 "
Expand Down
13 changes: 5 additions & 8 deletions FlagEmbedding/inference/auto_reranker.py
Original file line number Diff line number Diff line change
@@ -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__)
Expand Down Expand Up @@ -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:
Expand All @@ -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 "
Expand Down
70 changes: 70 additions & 0 deletions tests/test_auto_model_resolution.py
Original file line number Diff line number Diff line change
@@ -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