mirror of
https://github.com/modelscope/modelscope.git
synced 2026-09-03 04:32:02 +02:00
feat: sentence_embedding pipeline (#1435)
This commit is contained in:
@@ -1,3 +1,3 @@
|
||||
from .auto_class import *
|
||||
from .patcher import patch_context, patch_hub, unpatch_hub
|
||||
from .pipeline_builder import hf_pipeline
|
||||
from .pipeline_builder import hf_pipeline, sentence_transformers_pipeline
|
||||
|
||||
@@ -3,6 +3,9 @@ from typing import Optional, Union
|
||||
|
||||
from modelscope.hub import snapshot_download
|
||||
from modelscope.utils.hf_util.patcher import _patch_pretrained_class
|
||||
from modelscope.utils.logger import get_logger
|
||||
|
||||
logger = get_logger()
|
||||
|
||||
|
||||
def _get_hf_device(device):
|
||||
@@ -52,3 +55,38 @@ def hf_pipeline(
|
||||
device=device,
|
||||
pipeline_class=pipeline_class,
|
||||
**kwargs)
|
||||
|
||||
|
||||
def sentence_transformers_pipeline(model: str, **kwargs):
|
||||
try:
|
||||
from sentence_transformers import SentenceTransformer
|
||||
except ImportError:
|
||||
raise ImportError(
|
||||
'Could not import sentence_transformers, please upgrade to the latest version of sentence_transformers '
|
||||
"with: 'pip install -U sentence_transformers'") from None
|
||||
if isinstance(model, str):
|
||||
if not os.path.exists(model):
|
||||
model = snapshot_download(model)
|
||||
|
||||
from modelscope.pipelines import Pipeline
|
||||
|
||||
class SentenceTransformerPipeline(Pipeline):
|
||||
"""A wrapper for sentence_transformers.SentenceTransformer to make it compatible
|
||||
with the modelscope pipeline conventions."""
|
||||
|
||||
def __init__(self, model_path: str, **kwargs):
|
||||
self.model = SentenceTransformer(model_path, **kwargs)
|
||||
|
||||
def __call__(self,
|
||||
sentences: str | list[str] | None = None,
|
||||
prompt_name: str | None = None,
|
||||
**kwargs):
|
||||
input_data = kwargs.pop('input', None)
|
||||
if input_data is not None:
|
||||
sentences = input_data['source_sentence']
|
||||
res = self.model.encode(sentences, **kwargs)
|
||||
return {'text_embedding': res}
|
||||
return self.model.encode(
|
||||
sentences, prompt_name=prompt_name, **kwargs)
|
||||
|
||||
return SentenceTransformerPipeline(model, **kwargs)
|
||||
|
||||
@@ -82,6 +82,10 @@ def _inverted_index(forward_index):
|
||||
INVERTED_TASKS_LEVEL = _inverted_index(DEFAULT_TASKS_LEVEL)
|
||||
|
||||
|
||||
def is_embedding_task(task: str):
|
||||
return task == Tasks.sentence_embedding
|
||||
|
||||
|
||||
def get_task_by_subtask_name(group_key):
|
||||
if group_key in INVERTED_TASKS_LEVEL:
|
||||
return INVERTED_TASKS_LEVEL[group_key][
|
||||
|
||||
Reference in New Issue
Block a user