forked from Karylab-cklius/vllm
169 lines
6.1 KiB
Python
169 lines
6.1 KiB
Python
# SPDX-License-Identifier: Apache-2.0
|
|
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
|
|
|
from unittest.mock import MagicMock, patch
|
|
|
|
import numpy as np
|
|
import pytest
|
|
import torch
|
|
import torch.nn as nn
|
|
|
|
from vllm.model_executor.models.bert import (
|
|
BertMLMHead,
|
|
SPLADESparsePooler,
|
|
)
|
|
from vllm.platforms import current_platform
|
|
from vllm.pooling_params import PoolingParams
|
|
from vllm.utils.torch_utils import PIN_MEMORY
|
|
from vllm.v1.pool.metadata import PoolingMetadata, PoolingStates
|
|
from vllm.v1.worker.gpu.input_batch import InputBatch
|
|
from vllm.v1.worker.gpu.pool.pooling_runner import PoolingRunner
|
|
from vllm.v1.worker.gpu.states import RequestState
|
|
|
|
# ---------------------------------------------------------------------
|
|
# Functional test: SPLADE formula correctness (no HF download needed)
|
|
# ---------------------------------------------------------------------
|
|
|
|
|
|
@pytest.mark.parametrize("B,T,H,V", [(2, 3, 5, 7)])
|
|
@torch.inference_mode
|
|
def test_splade_pooler_matches_reference_formula(B, T, H, V):
|
|
"""Ensure SPLADESparsePooler forward() matches the mathematical formula:
|
|
log1p(relu(logits)) -> max over sequence length (after masking)."""
|
|
torch.manual_seed(0)
|
|
|
|
# Prepare [B] sequences of shape [T, H]
|
|
hs_list = [torch.randn(T, H) for _ in range(B)]
|
|
hs_tenser = torch.cat(hs_list)
|
|
|
|
# Simulate PoolingMetadata (only required fields)
|
|
prompt_lens = [T, T - 1]
|
|
prompt_lens_tenser = torch.tensor(prompt_lens, dtype=torch.int32)
|
|
token_ids = torch.tensor(
|
|
[
|
|
[101, 5, 102], # Batch 0: [CLS], token, [SEP]
|
|
[101, 6, 6], # Batch 1: [CLS], token, token (last token ignored)
|
|
],
|
|
dtype=torch.long,
|
|
)
|
|
meta = PoolingMetadata(
|
|
prompt_lens=prompt_lens_tenser,
|
|
prompt_token_ids=token_ids,
|
|
prompt_token_ids_cpu=token_ids,
|
|
pooling_params=[PoolingParams(task="embed")] * B,
|
|
pooling_states=[PoolingStates() for _ in range(B)],
|
|
)
|
|
|
|
# MLM head (prefer BertMLMHead, fallback to Linear if unavailable)
|
|
try:
|
|
mlm_head = BertMLMHead(hidden_size=H, vocab_size=V, layer_norm_eps=1e-12)
|
|
except Exception:
|
|
mlm_head = nn.Linear(H, V, bias=True)
|
|
|
|
# Forward pass through SPLADE pooler
|
|
pooler = SPLADESparsePooler(mlm_head=mlm_head, pooling="max", remove_cls_sep=True)
|
|
pooled = pooler(hidden_states=hs_tenser, pooling_metadata=meta) # list of [V]
|
|
|
|
# Basic output checks
|
|
assert isinstance(pooled, torch.Tensor) and len(pooled) == B
|
|
for vec in pooled:
|
|
assert vec.shape == (V,)
|
|
assert torch.isfinite(vec).all()
|
|
assert (vec >= 0).all(), "SPLADE outputs must be non-negative."
|
|
|
|
# Reference implementation for comparison
|
|
def ref_one(hs: torch.Tensor, L: int, tid_row: torch.Tensor) -> torch.Tensor:
|
|
keep = torch.ones(L, dtype=torch.bool)
|
|
if L > 0 and tid_row[0].item() == 101: # remove CLS
|
|
keep[0] = False
|
|
if L > 0 and tid_row[L - 1].item() == 102: # remove SEP
|
|
keep[L - 1] = False
|
|
|
|
valid = hs[:L][keep[:L]]
|
|
if valid.numel() == 0:
|
|
return torch.zeros(V, dtype=torch.float32)
|
|
|
|
logits = mlm_head(valid) # [L', V]
|
|
scores = torch.log1p(torch.relu(logits)) # [L', V]
|
|
return scores.max(dim=0).values.to(torch.float32)
|
|
|
|
torch.testing.assert_close(
|
|
pooled[0],
|
|
ref_one(hs_list[0], prompt_lens[0], token_ids[0]),
|
|
rtol=1e-4,
|
|
atol=1e-4,
|
|
)
|
|
torch.testing.assert_close(
|
|
pooled[1],
|
|
ref_one(hs_list[1], prompt_lens[1], token_ids[1]),
|
|
rtol=1e-4,
|
|
atol=1e-4,
|
|
)
|
|
|
|
|
|
def test_pooling_runner_gathers_required_token_ids() -> None:
|
|
runner = PoolingRunner.__new__(PoolingRunner)
|
|
pooling_params = PoolingParams(task="embed", requires_token_ids=True)
|
|
runner.pooling_params = {1: pooling_params, 3: pooling_params}
|
|
runner.pooling_states = {1: PoolingStates(), 3: PoolingStates()}
|
|
runner.prompt_token_ids = {
|
|
1: torch.tensor([101, 102]),
|
|
3: torch.tensor([101, 11, 102]),
|
|
}
|
|
|
|
input_batch = MagicMock(spec=InputBatch)
|
|
input_batch.idx_mapping_np = np.array([3, 1], dtype=np.int32)
|
|
input_batch.num_reqs = 2
|
|
req_states = MagicMock(spec=RequestState)
|
|
req_states.prompt_len = MagicMock(np=np.array([0, 2, 0, 3], dtype=np.int32))
|
|
metadata = runner._get_pooling_metadata(
|
|
input_batch, req_states, torch.device(current_platform.device_type)
|
|
)
|
|
|
|
expected = torch.tensor([[101, 11, 102], [101, 102, 0]])
|
|
assert metadata.prompt_token_ids_cpu is not None
|
|
assert metadata.prompt_token_ids is not None
|
|
assert metadata.prompt_token_ids_cpu.is_pinned() == PIN_MEMORY
|
|
torch.testing.assert_close(
|
|
metadata.prompt_lens, torch.tensor([3, 2], dtype=torch.int32)
|
|
)
|
|
torch.testing.assert_close(metadata.prompt_token_ids_cpu, expected)
|
|
torch.testing.assert_close(metadata.prompt_token_ids.cpu(), expected)
|
|
|
|
|
|
def test_pooling_runner_stores_only_required_token_ids() -> None:
|
|
runner = PoolingRunner.__new__(PoolingRunner)
|
|
runner.model = MagicMock()
|
|
runner.supported_tasks = frozenset({"embed"})
|
|
runner.pooling_params = {}
|
|
runner.pooling_states = {}
|
|
runner.prompt_token_ids = {}
|
|
|
|
runner.add_request(1, PoolingParams(task="embed"), [101, 102])
|
|
runner.add_request(
|
|
2,
|
|
PoolingParams(task="embed", requires_token_ids=True),
|
|
[101, 11, 102],
|
|
)
|
|
|
|
assert 1 not in runner.prompt_token_ids
|
|
torch.testing.assert_close(runner.prompt_token_ids[2], torch.tensor([101, 11, 102]))
|
|
|
|
|
|
def test_pooling_runner_rejects_unsupported_selected_task() -> None:
|
|
model = MagicMock()
|
|
model.pooler.get_supported_tasks.return_value = {
|
|
"embed",
|
|
"embed&token_classify",
|
|
"token_classify",
|
|
}
|
|
vllm_config = MagicMock()
|
|
vllm_config.scheduler_config.max_num_seqs = 2
|
|
vllm_config.model_config.get_pooling_task.return_value = "embed&token_classify"
|
|
|
|
with (
|
|
patch.object(PoolingRunner, "get_supported_tasks", return_value=["embed"]),
|
|
pytest.raises(ValueError, match="selects 'embed&token_classify'"),
|
|
):
|
|
PoolingRunner(model, vllm_config)
|