# 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)