diff --git a/.github/workflows/check.yml b/.github/workflows/check.yml index 125617b..7f3b278 100644 --- a/.github/workflows/check.yml +++ b/.github/workflows/check.yml @@ -45,7 +45,23 @@ jobs: python-version: ${{ matrix.python-version }} - name: Install package run: pip install -e ".[test]" - - name: Run smoke tests + - name: Run tests + run: pytest tests/ -v + + compatibility-vllm-023: + name: Compatibility (vLLM 0.23.0) + runs-on: ubuntu-latest + needs: lint + steps: + - uses: actions/checkout@v7 + - uses: actions/setup-python@v7 + with: + python-version: "3.12" + - name: Install vLLM 0.23 and test dependencies + run: pip install "vllm==0.23.0" "pytest>=8.0" + - name: Install plugin without changing vLLM + run: pip install -e . --no-deps + - name: Run compatibility tests run: pytest tests/ -v # ============ Build ============ diff --git a/CHANGELOG.md b/CHANGELOG.md index a816294..7659c96 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -7,6 +7,18 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 ## [Unreleased] +### Fixed + +- Restore weight loading on vLLM 0.23 while preserving the gate/up shard mapping + and the checkpoint's already-fused QKV projection. +- Keep the `spark25` tool parser importable on vLLM releases that predate the + shared `find_tool_name` helper. + +### Changed + +- Run and document the complete test suite, including packed-weight regression + coverage for the legacy loader path. + ## [0.1.0] - 2026-08-28 ### Added diff --git a/README.md b/README.md index dcfd807..4190f4a 100644 --- a/README.md +++ b/README.md @@ -194,6 +194,13 @@ concurrency of 500. Decode latency had a mean TPOT of 11.58 ms and a P99 of ## Compatibility and maintenance +The compatibility path is tested against vLLM `0.23.0` and the recorded vLLM +commit `81efe7883f30582696b69f9b9ea93c4819a8c608`. It detects the relevant loader +capabilities instead of branching on a version string. On vLLM 0.23, the +plugin preserves the checkpoint's split `gate_proj` / `up_proj` tensors when +loading vLLM's packed `gate_up_proj`; it also supplies the parser-name lookup +helper that this release lacks. + The model implementation depends on vLLM's internal runtime-layer APIs. The `TESTED_VLLM` value in `src/vllm_spark2_5_plugin/__init__.py` records the vLLM revision from which the implementation was extracted. When upgrading vLLM, @@ -204,15 +211,16 @@ registers, making version drift visible in the server log. ## Tests -The plugin smoke tests do not require a GPU or model weights. After installing -the plugin in the test environment, run: +The plugin tests do not require a GPU or model weights. After installing the +plugin in the test environment, run the complete suite: ```bash -.venv/bin/python -m pytest Spark-plugin/tests/test_smoke.py -q +.venv/bin/python -m pytest Spark-plugin/tests/ -q ``` -The tests verify both the public registration entry point and a complete -Spark2_5 XML tool-call parse. +The tests verify the public registration entry point, a complete Spark2_5 +XML tool-call parse, and the legacy packed-weight loading path (including +gate/up shard IDs, the already-fused QKV projection, and tied embeddings). ## Plugin loading and allowlists @@ -267,6 +275,14 @@ vLLM. Check the complete server log for the original traceback, then compare the running vLLM version with the vendored reference commit recorded in `src/vllm_spark2_5_plugin/__init__.py`. +### `WeightsMapper` rejects `orig_to_new_stacked` on vLLM 0.23 + +Use a plugin revision containing the vLLM 0.23 compatibility path and run the +complete test suite above. Removing only the unsupported constructor argument +is insufficient: Spark checkpoints store separate gate/up tensors that must be +loaded into distinct shards of vLLM's packed projection. If the reported serve +command uses `--tool-call-parser spark25`, verify that parser path as well. + ## Project layout ```text @@ -278,6 +294,7 @@ Spark-plugin/ | +-- spark2_5_tool_parser.py # Spark2_5 XML/KV parser +-- tests/ | +-- test_smoke.py # registration and parser smoke tests +| +-- test_weight_loading.py # packed-weight compatibility regression +-- pyproject.toml ``` diff --git a/src/vllm_spark2_5_plugin/spark2_5.py b/src/vllm_spark2_5_plugin/spark2_5.py index 71d6073..f084a67 100644 --- a/src/vllm_spark2_5_plugin/spark2_5.py +++ b/src/vllm_spark2_5_plugin/spark2_5.py @@ -8,6 +8,7 @@ """ from collections.abc import Iterable +from inspect import signature from itertools import islice import torch @@ -32,6 +33,10 @@ ParallelLMHead, VocabParallelEmbedding, ) +from vllm.model_executor.model_loader.weight_utils import ( + default_weight_loader, + maybe_remap_kv_scale_name, +) from vllm.sequence import IntermediateTensors from vllm.model_executor.models.interfaces import SupportsPP @@ -40,12 +45,21 @@ PPMissingLayer, WeightsMapper, extract_layer_index, + is_pp_missing_parameter, make_empty_intermediate_tensors_factory, make_layers, maybe_prefix, ) +_SUPPORTS_STACKED_WEIGHTS = ( + "orig_to_new_stacked" in signature(WeightsMapper).parameters +) +_SUPPORTS_SKIP_PREFIXES = "skip_prefixes" in signature( + AutoWeightsLoader +).parameters + + class Spark2_5MLP(nn.Module): def __init__( self, @@ -326,11 +340,15 @@ class Spark2_5ForCausalLM(nn.Module, SupportsPP): "gate_up_proj": ["gate_proj", "up_proj"], } - hf_to_vllm_mapper = WeightsMapper( - orig_to_new_stacked={ - ".gate_proj": (".gate_up_proj", 0), - ".up_proj": (".gate_up_proj", 1), - } + hf_to_vllm_mapper = ( + WeightsMapper( + orig_to_new_stacked={ + ".gate_proj": (".gate_up_proj", 0), + ".up_proj": (".gate_up_proj", 1), + } + ) + if _SUPPORTS_STACKED_WEIGHTS + else None ) def __init__(self, *, vllm_config: VllmConfig, prefix: str = "") -> None: @@ -384,8 +402,79 @@ def compute_logits( return logits def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]: - loader = AutoWeightsLoader( - self, - skip_prefixes=(["lm_head."] if self.config.tie_word_embeddings else None), + mapper = Spark2_5ForCausalLM.hf_to_vllm_mapper + if mapper is None: + return Spark2_5ForCausalLM._load_weights_without_stacked_mapper( + self, weights + ) + + loader_kwargs = {} + if self.config.tie_word_embeddings: + if _SUPPORTS_SKIP_PREFIXES: + loader_kwargs["skip_prefixes"] = ["lm_head."] + else: + weights = ( + (name, weight) + for name, weight in weights + if not name.startswith("lm_head.") + ) + loader = AutoWeightsLoader(self, **loader_kwargs) + return loader.load_weights(weights, mapper=mapper) + + def _load_weights_without_stacked_mapper( + self, weights: Iterable[tuple[str, torch.Tensor]] + ) -> set[str]: + """Load split MLP weights on vLLM releases predating stacked mappers.""" + stacked_params_mapping = ( + (".gate_up_proj.", ".gate_proj.", 0), + (".gate_up_proj.", ".up_proj.", 1), ) - return loader.load_weights(weights, mapper=self.hf_to_vllm_mapper) + params_dict = dict(self.named_parameters(remove_duplicate=False)) + loaded_params: set[str] = set() + + for name, loaded_weight in weights: + if self.config.tie_word_embeddings and name.startswith("lm_head."): + continue + if "rotary_emb.inv_freq" in name: + continue + + for param_name, weight_name, shard_id in stacked_params_mapping: + if weight_name not in name: + continue + name = name.replace(weight_name, param_name, 1) + if name.endswith(".bias") and name not in params_dict: + break + if name.endswith("scale"): + name = maybe_remap_kv_scale_name(name, params_dict) + if name is None: + break + if is_pp_missing_parameter(name, self) or name not in params_dict: + break + + param = params_dict[name] + weight_loader = getattr( + param, "weight_loader", default_weight_loader + ) + if weight_loader is default_weight_loader: + weight_loader(param, loaded_weight) + else: + weight_loader(param, loaded_weight, shard_id) + loaded_params.add(name) + break + else: + if name.endswith(".bias") and name not in params_dict: + continue + name = maybe_remap_kv_scale_name(name, params_dict) + if name is None: + continue + if is_pp_missing_parameter(name, self) or name not in params_dict: + continue + + param = params_dict[name] + weight_loader = getattr( + param, "weight_loader", default_weight_loader + ) + weight_loader(param, loaded_weight) + loaded_params.add(name) + + return loaded_params diff --git a/src/vllm_spark2_5_plugin/spark2_5_tool_parser.py b/src/vllm_spark2_5_plugin/spark2_5_tool_parser.py index b7c1f93..b35fb8c 100644 --- a/src/vllm_spark2_5_plugin/spark2_5_tool_parser.py +++ b/src/vllm_spark2_5_plugin/spark2_5_tool_parser.py @@ -25,13 +25,26 @@ from vllm.entrypoints.openai.responses.protocol import ResponsesRequest from vllm.logger import init_logger from vllm.tokenizers import TokenizerLike -from vllm.tool_parsers.abstract_tool_parser import Tool, ToolParser +from vllm.tool_parsers.abstract_tool_parser import ToolParser from vllm.tool_parsers.utils import ( - find_tool_name, + Tool, find_tool_properties, partial_tag_overlap, ) +try: + from vllm.tool_parsers.utils import find_tool_name +except ImportError: + + def find_tool_name(tools: list[Tool] | None, tool_name: str) -> bool: + """Compatibility helper for vLLM versions without this utility.""" + if not tools: + return False + return any( + getattr(getattr(tool, "function", tool), "name", None) == tool_name + for tool in tools + ) + logger = init_logger(__name__) TOOL_CALL_BEGIN = "" diff --git a/tests/test_weight_loading.py b/tests/test_weight_loading.py new file mode 100644 index 0000000..bdee812 --- /dev/null +++ b/tests/test_weight_loading.py @@ -0,0 +1,87 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project + +from types import SimpleNamespace + +import torch +from torch import nn + +from vllm_spark2_5_plugin.spark2_5 import Spark2_5ForCausalLM + + +def _parameter(values): + return nn.Parameter(torch.tensor(values), requires_grad=False) + + +class _TinySparkModel(nn.Module): + def __init__(self, calls): + super().__init__() + self.config = SimpleNamespace(tie_word_embeddings=True) + self.model = nn.Module() + self.model.layers = nn.ModuleList([nn.Module()]) + layer = self.model.layers[0] + layer.mlp = nn.Module() + layer.mlp.gate_up_proj = nn.Module() + layer.self_attn = nn.Module() + layer.self_attn.q_k_v_proj = nn.Module() + self.model.embedding = nn.Module() + + gate_up = _parameter([0.0, 0.0, 0.0, 0.0]) + q_k_v = _parameter([0.0, 0.0, 0.0]) + embedding = _parameter([0.0, 0.0]) + + def packed_loader(param, weight, shard_id=None): + if shard_id is None: + shard_id = getattr(weight, "shard_id", None) + calls.append(("gate_up", shard_id)) + assert shard_id in (0, 1) + start = shard_id * weight.numel() + param.data[start : start + weight.numel()].copy_(weight) + + def direct_qkv_loader(param, weight): + calls.append(("q_k_v", None)) + param.data.copy_(weight) + + gate_up.weight_loader = packed_loader + q_k_v.weight_loader = direct_qkv_loader + layer.mlp.gate_up_proj.register_parameter("weight", gate_up) + layer.self_attn.q_k_v_proj.register_parameter("weight", q_k_v) + self.model.embedding.register_parameter("weight", embedding) + self.lm_head = self.model.embedding + + +def test_loader_preserves_spark_checkpoint_layout(): + calls = [] + model = _TinySparkModel(calls) + weights = [ + ("model.layers.0.mlp.gate_proj.weight", torch.tensor([1.0, 2.0])), + ("model.layers.0.mlp.up_proj.weight", torch.tensor([3.0, 4.0])), + ( + "model.layers.0.self_attn.q_k_v_proj.weight", + torch.tensor([5.0, 6.0, 7.0]), + ), + ("model.embedding.weight", torch.tensor([8.0, 9.0])), + ("lm_head.weight", torch.tensor([98.0, 99.0])), + ] + + loaded = Spark2_5ForCausalLM.load_weights(model, iter(weights)) + + assert loaded == { + "model.layers.0.mlp.gate_up_proj.weight", + "model.layers.0.self_attn.q_k_v_proj.weight", + "model.embedding.weight", + } + assert calls == [("gate_up", 0), ("gate_up", 1), ("q_k_v", None)] + params = dict(model.named_parameters(remove_duplicate=False)) + assert params["model.layers.0.mlp.gate_up_proj.weight"].tolist() == [ + 1.0, + 2.0, + 3.0, + 4.0, + ] + assert params["model.layers.0.self_attn.q_k_v_proj.weight"].tolist() == [ + 5.0, + 6.0, + 7.0, + ] + assert params["lm_head.weight"].tolist() == [8.0, 9.0]