Compare commits
1
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
4ac64ec057 |
@@ -644,11 +644,6 @@ steps:
|
||||
- pytest -v -s plugins_tests/test_terratorch_io_processor_plugins.py
|
||||
- pip uninstall prithvi_io_processor_plugin -y
|
||||
# END: `io_processor` plugins test
|
||||
# BEGIN: `bge_m3_sparse io_processor` test
|
||||
- pip install -e ./plugins/bge_m3_sparse_plugin
|
||||
- pytest -v -s plugins_tests/test_bge_m3_sparse_io_processor_plugins.py
|
||||
- pip uninstall bge_m3_sparse_plugin -y
|
||||
# END: `bge_m3_sparse io_processor` test
|
||||
# BEGIN: `stat_logger` plugins test
|
||||
- pip install -e ./plugins/vllm_add_dummy_stat_logger
|
||||
- pytest -v -s plugins_tests/test_stats_logger_plugins.py
|
||||
|
||||
@@ -23,10 +23,6 @@ steps:
|
||||
- pip install -e ./plugins/prithvi_io_processor_plugin
|
||||
- pytest -v -s plugins_tests/test_terratorch_io_processor_plugins.py
|
||||
- pip uninstall prithvi_io_processor_plugin -y
|
||||
# test bge_m3_sparse io_processor plugin
|
||||
- pip install -e ./plugins/bge_m3_sparse_plugin
|
||||
- pytest -v -s plugins_tests/test_bge_m3_sparse_io_processor_plugins.py
|
||||
- pip uninstall bge_m3_sparse_plugin -y
|
||||
# end io_processor plugins test
|
||||
# begin stat_logger plugins test
|
||||
- pip install -e ./plugins/vllm_add_dummy_stat_logger
|
||||
|
||||
@@ -1,6 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
|
||||
|
||||
def register_bge_m3_sparse_embeddings_processor():
|
||||
return "bge_m3_sparse_processor.sparse_embeddings_processor.BgeM3SparseEmbeddingsProcessor" # noqa: E501
|
||||
-206
@@ -1,206 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
|
||||
from collections.abc import Sequence
|
||||
|
||||
from vllm.config import ModelConfig, PoolerConfig, VllmConfig
|
||||
from vllm.entrypoints.openai.engine.protocol import UsageInfo
|
||||
from vllm.entrypoints.pooling.base.protocol import EmbedRequestMixin
|
||||
from vllm.inputs import PromptType
|
||||
from vllm.outputs import PoolingRequestOutput
|
||||
from vllm.plugins.io_processors.interface import IOProcessor
|
||||
from vllm.pooling_params import PoolingParams
|
||||
from vllm.renderers import BaseRenderer
|
||||
from vllm.tokenizers.detokenizer_utils import convert_ids_list_to_tokens
|
||||
|
||||
from .types import (
|
||||
EMBED_TASKS,
|
||||
SparseEmbeddingCompletionRequestMixin,
|
||||
SparseEmbeddingResponse,
|
||||
SparseEmbeddingResponseData,
|
||||
SparseEmbeddingTokenWeight,
|
||||
)
|
||||
|
||||
|
||||
class BgeM3SparseEmbeddingsProcessor(
|
||||
IOProcessor[SparseEmbeddingCompletionRequestMixin, SparseEmbeddingResponse]
|
||||
):
|
||||
def __init__(self, vllm_config: VllmConfig, renderer: BaseRenderer):
|
||||
super().__init__(vllm_config, renderer)
|
||||
self.offline_requests: list[SparseEmbeddingCompletionRequestMixin] = []
|
||||
self.online_requests: dict[str, SparseEmbeddingCompletionRequestMixin] = {}
|
||||
self.renderer: BaseRenderer = renderer
|
||||
self.default_pooling_params = {}
|
||||
pooler_config: PoolerConfig = vllm_config.model_config.pooler_config
|
||||
if pooler_config is not None:
|
||||
for param in ["use_activation", "dimensions"]:
|
||||
if getattr(pooler_config, param, None) is None:
|
||||
continue
|
||||
self.default_pooling_params[param] = getattr(pooler_config, param)
|
||||
self.embed_dimensions = vllm_config.model_config.embedding_size
|
||||
self.embed_request_queue: list[EmbedRequestMixin] = []
|
||||
|
||||
def __repr__(self) -> str:
|
||||
return (
|
||||
f"BgeM3SparseEmbeddingsProcessor("
|
||||
f"embed_dimensions={self.embed_dimensions}, "
|
||||
f"default_pooling_params={self.default_pooling_params})"
|
||||
)
|
||||
|
||||
def merge_pooling_params(
|
||||
self,
|
||||
params: PoolingParams | None = None,
|
||||
) -> PoolingParams:
|
||||
if params is None:
|
||||
params = PoolingParams()
|
||||
# refer to PoolingCompletionRequest.to_pooling_params
|
||||
# set and verify pooling params
|
||||
params.skip_reading_prefix_cache = True
|
||||
|
||||
raw_embed_request = self.embed_request_queue.pop(0)
|
||||
if raw_embed_request.embed_task not in EMBED_TASKS:
|
||||
raise ValueError(
|
||||
f"Unsupported task {raw_embed_request}, "
|
||||
f"Supported tasks are {EMBED_TASKS}"
|
||||
)
|
||||
params.task = "embed&token_classify"
|
||||
params.use_activation = raw_embed_request.use_activation
|
||||
if params.use_activation is None:
|
||||
params.use_activation = True
|
||||
|
||||
params.dimensions = raw_embed_request.dimensions
|
||||
|
||||
model_config: ModelConfig = self.vllm_config.model_config
|
||||
for param in self.default_pooling_params:
|
||||
if getattr(params, param, None) is None:
|
||||
setattr(params, param, self.default_pooling_params[param])
|
||||
|
||||
if params.dimensions is not None:
|
||||
if not model_config.is_matryoshka:
|
||||
raise ValueError(
|
||||
f'Model "{model_config.served_model_name}" does not '
|
||||
f"support matryoshka representation, "
|
||||
f"changing output dimensions will lead to poor results."
|
||||
)
|
||||
|
||||
mds = model_config.matryoshka_dimensions
|
||||
if mds is not None:
|
||||
if params.dimensions not in mds:
|
||||
raise ValueError(
|
||||
f"Model {model_config.served_model_name!r} "
|
||||
f"only supports {str(mds)} matryoshka dimensions, "
|
||||
f"use other output dimensions will "
|
||||
f"lead to poor results."
|
||||
)
|
||||
elif params.dimensions < 1:
|
||||
raise ValueError("Dimensions must be greater than 0")
|
||||
return params
|
||||
|
||||
def parse_request(
|
||||
self, request_data: object
|
||||
) -> SparseEmbeddingCompletionRequestMixin:
|
||||
# for vllm.entrypoints.llm.LLM, offline mode, calls `encode` directly.
|
||||
if isinstance(request_data, dict):
|
||||
return SparseEmbeddingCompletionRequestMixin(**request_data)
|
||||
raise TypeError("request_data should be a dictionary")
|
||||
|
||||
def pre_process(
|
||||
self,
|
||||
prompt: SparseEmbeddingCompletionRequestMixin,
|
||||
request_id: str | None = None,
|
||||
**kwargs,
|
||||
) -> PromptType | Sequence[PromptType]:
|
||||
if request_id is not None:
|
||||
assert request_id not in self.online_requests, "request_id duplicated"
|
||||
self.online_requests[request_id] = prompt
|
||||
self.embed_request_queue.extend(prompt.to_embed_requests_online())
|
||||
else:
|
||||
self.offline_requests.append(prompt)
|
||||
self.embed_request_queue.extend(prompt.to_embed_requests_offline())
|
||||
return prompt.input
|
||||
|
||||
def _get_sparse_embedding_request(self, request_id: str | None = None):
|
||||
if request_id:
|
||||
return self.online_requests.pop(request_id, None)
|
||||
return self.offline_requests.pop(0)
|
||||
|
||||
def _build_sparse_embedding_token_weights(
|
||||
self,
|
||||
sparse_embedding: dict[int, float],
|
||||
return_tokens: bool = False,
|
||||
) -> list[SparseEmbeddingTokenWeight]:
|
||||
token_ids = sparse_embedding.keys()
|
||||
token_weights = sparse_embedding.values()
|
||||
tokens = [None] * len(token_ids)
|
||||
|
||||
if return_tokens and self.renderer is not None:
|
||||
tokens = convert_ids_list_to_tokens(
|
||||
self.renderer.get_tokenizer(), token_ids
|
||||
)
|
||||
sparse_embedding_output: list[SparseEmbeddingTokenWeight] = []
|
||||
for token_id, weight, token in zip(token_ids, token_weights, tokens):
|
||||
sparse_embedding_output.append(
|
||||
SparseEmbeddingTokenWeight(
|
||||
token_id=token_id, weight=weight, token=token
|
||||
)
|
||||
)
|
||||
return sparse_embedding_output
|
||||
|
||||
def post_process(
|
||||
self,
|
||||
model_output: Sequence[PoolingRequestOutput],
|
||||
request_id: str | None = None,
|
||||
**kwargs,
|
||||
) -> SparseEmbeddingResponse:
|
||||
num_prompt_tokens = 0
|
||||
response_data = []
|
||||
raw_request = self._get_sparse_embedding_request(request_id)
|
||||
has_dense_embed = raw_request.embed_task in ["dense", "dense&sparse"]
|
||||
has_sparse_embed = raw_request.embed_task in ["sparse", "dense&sparse"]
|
||||
embed_dimensions = (
|
||||
self.embed_dimensions
|
||||
if raw_request.dimensions is None
|
||||
else raw_request.dimensions
|
||||
)
|
||||
for idx in range(len(model_output)):
|
||||
mo = model_output[idx]
|
||||
sparse_embedding_dict: dict[int, float] = {}
|
||||
num_prompt_tokens += len(mo.prompt_token_ids)
|
||||
dense_embedding: list[float] | None = None
|
||||
sparse_embedding: list[SparseEmbeddingTokenWeight] | None = None
|
||||
if has_dense_embed:
|
||||
dense_embedding = mo.outputs.data[:embed_dimensions].tolist()
|
||||
if has_sparse_embed:
|
||||
sparse_weights = mo.outputs.data[embed_dimensions:].tolist()
|
||||
if len(mo.prompt_token_ids) != len(sparse_weights):
|
||||
# this is the case that add_special_tokens is True,
|
||||
# which means first token and last token are special tokens
|
||||
mo.prompt_token_ids = mo.prompt_token_ids[1:]
|
||||
for token_id, weight in zip(mo.prompt_token_ids, sparse_weights):
|
||||
sparse_embedding_dict[token_id] = max(
|
||||
weight, sparse_embedding_dict.get(token_id, 0.0)
|
||||
)
|
||||
sparse_embedding = self._build_sparse_embedding_token_weights(
|
||||
sparse_embedding_dict,
|
||||
raw_request.return_tokens,
|
||||
)
|
||||
|
||||
response_data.append(
|
||||
SparseEmbeddingResponseData(
|
||||
index=idx,
|
||||
object=raw_request.embed_task,
|
||||
sparse_embedding=sparse_embedding,
|
||||
dense_embedding=dense_embedding,
|
||||
)
|
||||
)
|
||||
|
||||
usage = UsageInfo(
|
||||
prompt_tokens=num_prompt_tokens,
|
||||
total_tokens=num_prompt_tokens,
|
||||
)
|
||||
resp = SparseEmbeddingResponse(
|
||||
data=response_data,
|
||||
usage=usage,
|
||||
)
|
||||
|
||||
return resp
|
||||
@@ -1,59 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
|
||||
from typing import Literal, get_args
|
||||
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
from vllm.entrypoints.openai.engine.protocol import UsageInfo
|
||||
from vllm.entrypoints.pooling.base.protocol import (
|
||||
CompletionRequestMixin,
|
||||
EmbedRequestMixin,
|
||||
)
|
||||
|
||||
EmbedTask = Literal[
|
||||
"sparse",
|
||||
"dense",
|
||||
"dense&sparse",
|
||||
]
|
||||
|
||||
EMBED_TASKS: tuple[EmbedTask, ...] = get_args(EmbedTask)
|
||||
|
||||
|
||||
class SparseEmbeddingCompletionRequestMixin(CompletionRequestMixin, EmbedRequestMixin):
|
||||
return_tokens: bool | None = Field(
|
||||
default=None,
|
||||
description="Whether to return dict shows the mapping of token_id to text."
|
||||
"`None` or False means not return.",
|
||||
)
|
||||
embed_task: EmbedTask = Field(
|
||||
default="dense&sparse",
|
||||
description="embed task, can be one of 'sparse', 'dense' , 'dense&sparse', "
|
||||
"default to 'dense&sparse'",
|
||||
)
|
||||
|
||||
def to_embed_requests_offline(self) -> list[EmbedRequestMixin]:
|
||||
if isinstance(self.input, list):
|
||||
return [self] * len(self.input)
|
||||
return [self]
|
||||
|
||||
def to_embed_requests_online(self) -> list[EmbedRequestMixin]:
|
||||
return [self]
|
||||
|
||||
|
||||
class SparseEmbeddingTokenWeight(BaseModel):
|
||||
token_id: int
|
||||
weight: float
|
||||
token: str | None
|
||||
|
||||
|
||||
class SparseEmbeddingResponseData(BaseModel):
|
||||
index: int
|
||||
object: str = "dense&sparse"
|
||||
sparse_embedding: list[SparseEmbeddingTokenWeight] | None
|
||||
dense_embedding: list[float] | None
|
||||
|
||||
|
||||
class SparseEmbeddingResponse(BaseModel):
|
||||
data: list[SparseEmbeddingResponseData]
|
||||
usage: UsageInfo
|
||||
@@ -1,15 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
|
||||
from setuptools import setup
|
||||
|
||||
setup(
|
||||
name="bge-m3-sparse-plugin",
|
||||
version="0.1",
|
||||
packages=["bge_m3_sparse_processor"],
|
||||
entry_points={
|
||||
"vllm.io_processor_plugins": [
|
||||
"bge_m3_sparse_plugin = bge_m3_sparse_processor:register_bge_m3_sparse_embeddings_processor", # noqa: E501
|
||||
]
|
||||
},
|
||||
)
|
||||
@@ -1,235 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
|
||||
import json
|
||||
|
||||
import pytest
|
||||
import requests
|
||||
|
||||
# Test configuration for BGE-M3 sparse plugin
|
||||
from tests.utils import RemoteOpenAIServer
|
||||
from vllm.entrypoints.pooling.pooling.protocol import IOProcessorResponse
|
||||
|
||||
model_config = {
|
||||
"model_name": "BAAI/bge-m3",
|
||||
"plugin": "bge_m3_sparse_plugin",
|
||||
"test_input": "What is the capital of France?",
|
||||
"hf_overrides": json.dumps(
|
||||
{"architectures": ["BgeM3EmbeddingModel"], "head_dtype": "float16"}
|
||||
),
|
||||
}
|
||||
|
||||
dense_embedding_sum = [
|
||||
-0.7214539647102356, # "What is the capital of France?"
|
||||
-0.6926871538162231, # "What is the capital of Germany?"
|
||||
-0.7129564881324768, # "What is the capital of Spain?"
|
||||
]
|
||||
|
||||
|
||||
def _float_close(expected: object, result: object):
|
||||
assert isinstance(expected, float) and isinstance(result, float), (
|
||||
f"{expected=} or {result=} is not float"
|
||||
)
|
||||
return (expected - result) < 1e-3 or abs(expected / result - 1) < 1e-3
|
||||
|
||||
|
||||
def _get_attr_or_val(obj: object | dict, key: str):
|
||||
if isinstance(obj, dict) and key in obj:
|
||||
return obj[key]
|
||||
return getattr(obj, key, None)
|
||||
|
||||
|
||||
def _check_dense_embedding(data, index=0):
|
||||
assert _float_close(sum(data), dense_embedding_sum[index]), (
|
||||
"dense-embedding result not match"
|
||||
)
|
||||
|
||||
|
||||
def _check_sparse_embedding(data, check_tokens=False):
|
||||
expected_weights = [
|
||||
{"token_id": 32, "weight": 0.0552978515625, "token": "?"},
|
||||
{"token_id": 70, "weight": 0.09808349609375, "token": "the"},
|
||||
{"token_id": 83, "weight": 0.08154296875, "token": "is"},
|
||||
{"token_id": 111, "weight": 0.11810302734375, "token": "of"},
|
||||
{"token_id": 4865, "weight": 0.1171875, "token": "What"},
|
||||
{"token_id": 9942, "weight": 0.292236328125, "token": "France"},
|
||||
{"token_id": 10323, "weight": 0.2802734375, "token": "capital"},
|
||||
]
|
||||
expected_embed = {x["token_id"]: x for x in expected_weights}
|
||||
|
||||
assert len(data) == len(expected_embed)
|
||||
for entry in data:
|
||||
expected_val = expected_embed[_get_attr_or_val(entry, "token_id")]
|
||||
assert _float_close(
|
||||
expected_val["weight"], _get_attr_or_val(entry, "weight")
|
||||
), f"actual embed {entry} not equal to {expected_val}"
|
||||
if check_tokens:
|
||||
assert expected_val["token"] == _get_attr_or_val(entry, "token"), (
|
||||
f"actual embed {entry} not equal to {expected_val}"
|
||||
)
|
||||
else:
|
||||
assert _get_attr_or_val(entry, "token") is None, (
|
||||
f"{entry} should not return token"
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture(scope="function")
|
||||
def server():
|
||||
args = [
|
||||
"--runner",
|
||||
"pooling",
|
||||
"--enforce-eager",
|
||||
"--max-num-seqs",
|
||||
"32",
|
||||
"--hf_overrides",
|
||||
model_config["hf_overrides"],
|
||||
"--io-processor-plugin",
|
||||
model_config["plugin"],
|
||||
]
|
||||
|
||||
with RemoteOpenAIServer(model_config["model_name"], args) as remote_server:
|
||||
yield remote_server
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"return_tokens",
|
||||
[True, False],
|
||||
)
|
||||
async def test_bge_m3_sparse_plugin_online(
|
||||
server: RemoteOpenAIServer, return_tokens: bool
|
||||
):
|
||||
"""Test BGE-M3 sparse plugin in online mode via API."""
|
||||
request_payload = {
|
||||
"model": model_config["model_name"],
|
||||
"task": "plugin",
|
||||
"data": {"input": model_config["test_input"], "return_tokens": return_tokens},
|
||||
}
|
||||
|
||||
ret = requests.post(
|
||||
server.url_for("pooling"),
|
||||
json=request_payload,
|
||||
)
|
||||
|
||||
response = ret.json()
|
||||
|
||||
# Verify the request response is in the correct format
|
||||
assert (parsed_response := IOProcessorResponse(**response).data)
|
||||
|
||||
# Verify the output is formatted as expected for this plugin
|
||||
assert _get_attr_or_val(parsed_response, "data")
|
||||
assert len(_get_attr_or_val(parsed_response, "data")) > 0
|
||||
|
||||
data_entry = _get_attr_or_val(parsed_response, "data")[0]
|
||||
assert _get_attr_or_val(data_entry, "object") == "dense&sparse"
|
||||
assert _get_attr_or_val(data_entry, "sparse_embedding")
|
||||
|
||||
# Verify sparse embedding format
|
||||
sparse_embedding = _get_attr_or_val(data_entry, "sparse_embedding")
|
||||
assert isinstance(sparse_embedding, list)
|
||||
_check_sparse_embedding(sparse_embedding, return_tokens)
|
||||
|
||||
# Verify dense embedding format
|
||||
dense_embedding = _get_attr_or_val(data_entry, "dense_embedding")
|
||||
assert isinstance(dense_embedding, list)
|
||||
_check_dense_embedding(dense_embedding)
|
||||
|
||||
# Verify usage information
|
||||
usage = _get_attr_or_val(parsed_response, "usage")
|
||||
assert usage, f"usage not found for {parsed_response}"
|
||||
assert _get_attr_or_val(usage, "prompt_tokens") > 0
|
||||
assert _get_attr_or_val(usage, "total_tokens") == _get_attr_or_val(
|
||||
usage, "prompt_tokens"
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"return_tokens",
|
||||
[True, False],
|
||||
)
|
||||
def test_bge_m3_sparse_plugin_offline(vllm_runner, return_tokens: bool):
|
||||
"""Test BGE-M3 sparse plugin in offline mode."""
|
||||
prompt = {
|
||||
"data": {
|
||||
"input": model_config["test_input"],
|
||||
"return_tokens": return_tokens,
|
||||
}
|
||||
}
|
||||
|
||||
with vllm_runner(
|
||||
model_config["model_name"],
|
||||
runner="pooling",
|
||||
enforce_eager=True,
|
||||
max_num_seqs=32,
|
||||
io_processor_plugin=model_config["plugin"],
|
||||
hf_overrides=json.loads(model_config["hf_overrides"]),
|
||||
default_torch_num_threads=1,
|
||||
) as llm_runner:
|
||||
llm = llm_runner.get_llm()
|
||||
pooler_output = llm.encode(prompt, pooling_task="plugin")
|
||||
|
||||
outputs = pooler_output[0]
|
||||
|
||||
# Verify output structure
|
||||
assert hasattr(outputs, "outputs")
|
||||
response = outputs.outputs
|
||||
assert hasattr(response, "data")
|
||||
assert len(response.data) == 1
|
||||
# Verify response data
|
||||
for i, output in enumerate(response.data):
|
||||
# Each output should have sparse embeddings
|
||||
sparse_embedding = output.sparse_embedding
|
||||
assert isinstance(sparse_embedding, list)
|
||||
_check_sparse_embedding(sparse_embedding, return_tokens)
|
||||
dense_embedding = output.dense_embedding
|
||||
assert isinstance(dense_embedding, list)
|
||||
_check_dense_embedding(dense_embedding)
|
||||
|
||||
# Verify usage
|
||||
assert response.usage.prompt_tokens > 0
|
||||
assert response.usage.total_tokens == response.usage.prompt_tokens
|
||||
|
||||
|
||||
def test_bge_m3_sparse_plugin_offline_multiple_inputs(vllm_runner):
|
||||
"""Test BGE-M3 sparse plugin with multiple inputs in offline mode."""
|
||||
prompts = {
|
||||
"data": {
|
||||
"input": [
|
||||
"What is the capital of France?",
|
||||
"What is the capital of Germany?",
|
||||
"What is the capital of Spain?",
|
||||
],
|
||||
"return_tokens": True,
|
||||
}
|
||||
}
|
||||
|
||||
with vllm_runner(
|
||||
model_config["model_name"],
|
||||
runner="pooling",
|
||||
enforce_eager=True,
|
||||
max_num_seqs=32,
|
||||
io_processor_plugin=model_config["plugin"],
|
||||
hf_overrides=json.loads(model_config["hf_overrides"]),
|
||||
default_torch_num_threads=1,
|
||||
) as llm_runner:
|
||||
llm = llm_runner.get_llm()
|
||||
pooler_output = llm.encode(prompts, pooling_task="plugin")
|
||||
|
||||
outputs = pooler_output[0]
|
||||
|
||||
# Verify output structure
|
||||
assert hasattr(outputs, "outputs")
|
||||
response = outputs.outputs
|
||||
assert hasattr(response, "data")
|
||||
assert len(response.data) == 3
|
||||
for i, output in enumerate(response.data):
|
||||
# Each output should have sparse embeddings
|
||||
sparse_embedding = output.sparse_embedding
|
||||
assert isinstance(sparse_embedding, list)
|
||||
dense_embedding = output.dense_embedding
|
||||
assert isinstance(dense_embedding, list)
|
||||
_check_dense_embedding(dense_embedding, i)
|
||||
|
||||
# Verify usage
|
||||
assert response.usage.prompt_tokens > 0
|
||||
assert response.usage.total_tokens == response.usage.prompt_tokens
|
||||
@@ -1540,7 +1540,6 @@ class ModelConfig:
|
||||
return "token_classify"
|
||||
|
||||
priority: list[PoolingTask] = [
|
||||
"embed&token_classify",
|
||||
"embed",
|
||||
"classify",
|
||||
"token_embed",
|
||||
|
||||
@@ -196,42 +196,4 @@ class BOSEOSFilter(Pooler):
|
||||
return pooled_outputs
|
||||
|
||||
|
||||
class BgeM3Pooler(Pooler):
|
||||
def __init__(self, token_classify_pooler: Pooler, embed_pooler: Pooler) -> None:
|
||||
super().__init__()
|
||||
self.token_classify_pooler = token_classify_pooler
|
||||
self.embed_pooler = embed_pooler
|
||||
|
||||
def forward(
|
||||
self, hidden_states: torch.Tensor, pooling_metadata: PoolingMetadata
|
||||
) -> PoolerOutput:
|
||||
embed_outputs = self.embed_pooler(hidden_states, pooling_metadata)
|
||||
token_classify_outputs = self.token_classify_pooler(
|
||||
hidden_states, pooling_metadata
|
||||
)
|
||||
pooler_outputs: list[torch.Tensor] = []
|
||||
for embed_output, token_classify_output in zip(
|
||||
embed_outputs, token_classify_outputs
|
||||
):
|
||||
pooler_outputs.append(
|
||||
torch.cat(
|
||||
[embed_output.view(-1), token_classify_output.view(-1)], dim=-1
|
||||
)
|
||||
)
|
||||
|
||||
return pooler_outputs
|
||||
|
||||
def get_supported_tasks(self) -> Set[PoolingTask]:
|
||||
return {"embed&token_classify"}
|
||||
|
||||
def get_pooling_updates(self, task: PoolingTask) -> PoolingParamsUpdate:
|
||||
return self.embed_pooler.get_pooling_updates(
|
||||
"embed"
|
||||
) | self.token_classify_pooler.get_pooling_updates("token_classify")
|
||||
|
||||
def extra_repr(self) -> str:
|
||||
s = f"supported_task={self.get_supported_tasks()}"
|
||||
return s
|
||||
|
||||
|
||||
__all__ = ["BOSEOSFilter", "DispatchPooler", "IdentityPooler", "BgeM3Pooler"]
|
||||
__all__ = ["BOSEOSFilter", "DispatchPooler", "IdentityPooler"]
|
||||
|
||||
@@ -10,7 +10,6 @@ from transformers import RobertaConfig
|
||||
|
||||
from vllm.config import ModelConfig, PoolerConfig, VllmConfig
|
||||
from vllm.model_executor.layers.pooler import (
|
||||
BgeM3Pooler,
|
||||
BOSEOSFilter,
|
||||
DispatchPooler,
|
||||
Pooler,
|
||||
@@ -238,9 +237,6 @@ class BgeM3EmbeddingModel(RobertaEmbeddingModel):
|
||||
# for some reason m3 only filters the bos for colbert vectors
|
||||
),
|
||||
"token_classify": token_classify_pooler,
|
||||
"embed&token_classify": BgeM3Pooler(
|
||||
token_classify_pooler, embed_pooler
|
||||
),
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
@@ -11,7 +11,6 @@ PoolingTask = Literal[
|
||||
"token_embed",
|
||||
"token_classify",
|
||||
"plugin",
|
||||
"embed&token_classify",
|
||||
]
|
||||
POOLING_TASKS: tuple[PoolingTask, ...] = get_args(PoolingTask)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user