From bc989c803c8634286904c96631b98cdfd1222ba2 Mon Sep 17 00:00:00 2001 From: QuanshengGu Date: Thu, 13 Aug 2026 03:55:20 +0800 Subject: [PATCH 01/13] feat: add gemma4 text inference --- benchmark/fixed_token_microbench.py | 263 +++++++ src/sparsevllm/configs/model.py | 1 + src/sparsevllm/configs/sparse.py | 10 + .../engine/cache_manager/standard.py | 39 +- .../engine/cache_manager/storage/__init__.py | 14 + .../storage/heterogeneous_explicit_kv.py | 131 +++ src/sparsevllm/engine/model_runner.py | 5 + .../triton/gemma4_context_attention.py | 215 +++++ .../kernels/triton/gemma4_decode_attention.py | 319 ++++++++ .../kernels/triton/gemma4_fused_ops.py | 131 +++ .../kernels/triton/gemma4_gelu_and_mul.py | 49 ++ src/sparsevllm/kernels/triton/gemma4_moe.py | 125 +++ .../kernels/triton/gemma4_qkv_norm_rope.py | 227 ++++++ .../kernels/triton/gemma4_rmsnorm.py | 66 ++ .../kernels/triton/gemma4_router.py | 143 ++++ .../gemma4_single_block_decode_attention.py | 154 ++++ src/sparsevllm/layers/activation.py | 13 + src/sparsevllm/layers/gemma4_rmsnorm.py | 32 + src/sparsevllm/layers/linear.py | 93 +++ src/sparsevllm/layers/packed_moe.py | 2 + src/sparsevllm/method_registry.py | 11 +- src/sparsevllm/models/checkpoint.py | 23 + src/sparsevllm/models/gemma4.py | 744 ++++++++++++++++++ src/sparsevllm/models/layout.py | 157 +++- src/sparsevllm/models/spec.py | 32 +- src/sparsevllm/operators/activation.py | 107 +++ src/sparsevllm/operators/gemma4_attention.py | 104 +++ src/sparsevllm/operators/gemma4_moe.py | 139 ++++ src/sparsevllm/operators/moe.py | 6 + tests/test_activation.py | 29 +- tests/test_attention_cache_storage.py | 51 +- tests/test_gemma4_attention_kernels.py | 284 +++++++ tests/test_gemma4_model.py | 288 +++++++ tests/test_gemma4_rmsnorm.py | 34 + tests/test_triton_moe.py | 68 +- 35 files changed, 4067 insertions(+), 42 deletions(-) create mode 100644 benchmark/fixed_token_microbench.py create mode 100644 src/sparsevllm/engine/cache_manager/storage/heterogeneous_explicit_kv.py create mode 100644 src/sparsevllm/kernels/triton/gemma4_context_attention.py create mode 100644 src/sparsevllm/kernels/triton/gemma4_decode_attention.py create mode 100644 src/sparsevllm/kernels/triton/gemma4_fused_ops.py create mode 100644 src/sparsevllm/kernels/triton/gemma4_gelu_and_mul.py create mode 100644 src/sparsevllm/kernels/triton/gemma4_moe.py create mode 100644 src/sparsevllm/kernels/triton/gemma4_qkv_norm_rope.py create mode 100644 src/sparsevllm/kernels/triton/gemma4_rmsnorm.py create mode 100644 src/sparsevllm/kernels/triton/gemma4_router.py create mode 100644 src/sparsevllm/kernels/triton/gemma4_single_block_decode_attention.py create mode 100644 src/sparsevllm/layers/gemma4_rmsnorm.py create mode 100644 src/sparsevllm/models/gemma4.py create mode 100644 src/sparsevllm/operators/gemma4_attention.py create mode 100644 src/sparsevllm/operators/gemma4_moe.py create mode 100644 tests/test_gemma4_attention_kernels.py create mode 100644 tests/test_gemma4_model.py create mode 100644 tests/test_gemma4_rmsnorm.py diff --git a/benchmark/fixed_token_microbench.py b/benchmark/fixed_token_microbench.py new file mode 100644 index 00000000..9520292e --- /dev/null +++ b/benchmark/fixed_token_microbench.py @@ -0,0 +1,263 @@ +"""Reproducible fixed-token latency benchmark for Sparse-vLLM and vLLM.""" + +from __future__ import annotations + +import argparse +import json +import os +import shlex +import statistics +import subprocess +import sys +from datetime import datetime +from importlib.metadata import PackageNotFoundError, version +from pathlib import Path +from time import perf_counter +from typing import Any + + +def _parser() -> argparse.ArgumentParser: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--backend", choices=("sparsevllm", "vllm"), required=True) + parser.add_argument("--model-path", required=True) + parser.add_argument("--output-dir", required=True) + parser.add_argument("--input-len", type=int, default=128) + parser.add_argument("--output-len", type=int, default=32) + parser.add_argument("--batch-size", type=int, default=16) + parser.add_argument("--num-warmups", type=int, default=3) + parser.add_argument("--num-iters", type=int, default=5) + parser.add_argument("--tensor-parallel-size", type=int, default=1) + parser.add_argument("--expert-parallel-size", type=int, default=1) + parser.add_argument("--gpu-memory-utilization", type=float, default=0.8) + return parser + + +def _write_json(path: Path, payload: Any) -> None: + path.write_text( + json.dumps(payload, indent=2, sort_keys=True) + "\n", encoding="utf-8" + ) + + +def _write_jsonl(path: Path, rows: list[dict[str, Any]]) -> None: + with path.open("w", encoding="utf-8") as output: + for row in rows: + output.write(json.dumps(row, sort_keys=True) + "\n") + + +def _git(*args: str) -> str | None: + result = subprocess.run( + ["git", *args], + capture_output=True, + check=False, + cwd=Path(__file__).parents[1], + text=True, + ) + return result.stdout.strip() or None + + +def _package_version(name: str) -> str | None: + try: + return version(name) + except PackageNotFoundError: + return None + + +def _validate(args: argparse.Namespace) -> Path: + for name in ( + "input_len", + "output_len", + "batch_size", + "num_warmups", + "num_iters", + "tensor_parallel_size", + "expert_parallel_size", + ): + if int(getattr(args, name)) <= 0: + raise ValueError(f"{name} must be positive") + if not 0 < args.gpu_memory_utilization <= 1: + raise ValueError("gpu_memory_utilization must be in (0, 1]") + if args.backend == "vllm" and args.expert_parallel_size != 1: + raise ValueError( + "vLLM uses --expert-parallel-size=1 or --enable-expert-parallel" + ) + output_dir = Path(args.output_dir).expanduser().resolve() + if output_dir.exists() and any(output_dir.iterdir()): + raise FileExistsError(f"output directory is not empty: {output_dir}") + output_dir.mkdir(parents=True, exist_ok=True) + return output_dir + + +def _build_engine(args: argparse.Namespace): + common = { + "model": str(Path(args.model_path).expanduser().resolve()), + "max_model_len": args.input_len + args.output_len, + "gpu_memory_utilization": args.gpu_memory_utilization, + "tensor_parallel_size": args.tensor_parallel_size, + "enforce_eager": False, + } + if args.backend == "vllm": + from vllm import LLM + + return LLM( + **common, + max_num_seqs=args.batch_size, + max_num_batched_tokens=max(4096, args.batch_size * args.input_len), + enable_expert_parallel=args.expert_parallel_size > 1, + enable_flashinfer_autotune=False, + language_model_only=True, + trust_remote_code=True, + ) + from sparsevllm import LLM + + return LLM( + common.pop("model"), + **common, + expert_parallel_size=args.expert_parallel_size, + max_num_seqs_in_batch=args.batch_size, + max_decoding_seqs=args.batch_size, + max_num_seqs_in_gpu=args.batch_size, + max_num_batched_tokens=args.batch_size * args.input_len, + decode_cuda_graph=True, + ) + + +def _sampling_params(args: argparse.Namespace): + if args.backend == "vllm": + from vllm import SamplingParams + else: + from sparsevllm import SamplingParams + return SamplingParams(temperature=0.0, max_tokens=args.output_len, ignore_eos=True) + + +def _token_ids(args: argparse.Namespace, output: Any) -> list[int]: + if args.backend == "vllm": + return list(output.outputs[0].token_ids) + return list(output["token_ids"]) + + +def main() -> int: + args = _parser().parse_args() + output_dir = _validate(args) + run_info = { + "benchmark": "fixed_token_microbench", + "backend": args.backend, + "command": shlex.join(sys.argv), + "created_at": datetime.now().astimezone().isoformat(timespec="seconds"), + "git": { + "branch": _git("branch", "--show-current"), + "commit": _git("rev-parse", "HEAD"), + "dirty": bool(_git("status", "--porcelain")), + }, + "workload": { + "batch_size": args.batch_size, + "input_len": args.input_len, + "output_len": args.output_len, + "num_warmups": args.num_warmups, + "num_iters": args.num_iters, + "prompt": "[2] + [100 + (request + position) % 1000]", + "temperature": 0.0, + "ignore_eos": True, + }, + "topology": { + "tensor_parallel_size": args.tensor_parallel_size, + "expert_parallel_size": args.expert_parallel_size, + "cuda_graph": True, + }, + "environment": { + "cuda_visible_devices": os.getenv("CUDA_VISIBLE_DEVICES"), + "python": sys.version, + "torch": _package_version("torch"), + "transformers": _package_version("transformers"), + "flashinfer_python": _package_version("flashinfer-python"), + "triton": _package_version("triton"), + "vllm": _package_version("vllm"), + }, + } + _write_json(output_dir / "run_info.json", run_info) + prompts = [ + [ + 2, + *( + 100 + (request + position) % 1000 + for position in range(args.input_len - 1) + ), + ] + for request in range(args.batch_size) + ] + engine = None + raw_outputs: list[dict[str, Any]] = [] + sample_results: list[dict[str, Any]] = [] + performance: list[dict[str, Any]] = [] + try: + engine = _build_engine(args) + params = _sampling_params(args) + for _ in range(args.num_warmups): + outputs = engine.generate(prompts, params, use_tqdm=False) + if len(outputs) != args.batch_size: + raise RuntimeError(f"warmup returned {len(outputs)} requests") + for iteration in range(args.num_iters): + started = perf_counter() + outputs = engine.generate(prompts, params, use_tqdm=False) + elapsed = perf_counter() - started + if len(outputs) != args.batch_size: + raise RuntimeError( + f"iteration {iteration} returned {len(outputs)} requests" + ) + generated = 0 + for sample_index, output in enumerate(outputs): + token_ids = _token_ids(args, output) + status = ( + "success" if len(token_ids) == args.output_len else "model_failed" + ) + row = { + "iteration": iteration, + "sample_index": sample_index, + "status": status, + "input_tokens": args.input_len, + "output_tokens": len(token_ids), + } + sample_results.append(row) + raw_outputs.append({**row, "token_ids": token_ids}) + if status != "success": + raise RuntimeError( + f"iteration {iteration} sample {sample_index} produced {len(token_ids)} tokens" + ) + generated += len(token_ids) + performance.append( + { + "iteration": iteration, + "status": "success", + "elapsed_s": elapsed, + "output_tokens": generated, + "output_tok_s": generated / elapsed, + } + ) + rates = [row["output_tok_s"] for row in performance] + aggregate = { + "benchmark": "fixed_token_microbench", + "backend": args.backend, + "status": "success", + "output_tok_s_mean": statistics.fmean(rates), + "output_tok_s_median": statistics.median(rates), + "samples": len(rates), + } + except Exception as error: + aggregate = { + "benchmark": "fixed_token_microbench", + "backend": args.backend, + "status": "model_failed", + "error": repr(error), + } + raise + finally: + _write_jsonl(output_dir / "raw_outputs.jsonl", raw_outputs) + _write_jsonl(output_dir / "per_sample_results.jsonl", sample_results) + _write_jsonl(output_dir / "performance.jsonl", performance) + _write_json(output_dir / "aggregate_metrics.json", aggregate) + if args.backend == "sparsevllm" and engine is not None: + engine.exit() + return 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/src/sparsevllm/configs/model.py b/src/sparsevllm/configs/model.py index 6c335f1e..2adb63a4 100644 --- a/src/sparsevllm/configs/model.py +++ b/src/sparsevllm/configs/model.py @@ -97,6 +97,7 @@ def load_and_validate_model(config) -> None: config.tensor_parallel_size, config.expert_parallel_size, config.data_parallel_size, + config.hf_config, ) if config.tiny_random: from sparsevllm.debug.tiny_random import apply_tiny_random_overrides diff --git a/src/sparsevllm/configs/sparse.py b/src/sparsevllm/configs/sparse.py index af47fcd2..95a21168 100644 --- a/src/sparsevllm/configs/sparse.py +++ b/src/sparsevllm/configs/sparse.py @@ -12,6 +12,7 @@ ) from sparsevllm.utils.log import logger, log_once + def normalize_sparse_method_name(config) -> bool: raw_sparse_method = config.vllm_sparse_method raw_sparse_method_normalized = "" if raw_sparse_method is None else str(raw_sparse_method).strip().lower() @@ -175,6 +176,15 @@ def _normalize_skipkv(config) -> None: ) def normalize_sparse_methods(config) -> None: + if ( + getattr(config.hf_config, "model_type", "") == "gemma4_text" + and int(getattr(config.hf_config, "num_kv_shared_layers", 0) or 0) + and config.vllm_sparse_method == "streamingllm" + ): + raise NotImplementedError( + "Gemma 4 StreamingLLM requires independent per-layer KV caches; " + "KV-sharing variants support vanilla and OmniKV." + ) _normalize_quest(config) _normalize_h2o(config) _normalize_rkv(config) diff --git a/src/sparsevllm/engine/cache_manager/standard.py b/src/sparsevllm/engine/cache_manager/standard.py index 74a26780..da3e31d9 100644 --- a/src/sparsevllm/engine/cache_manager/standard.py +++ b/src/sparsevllm/engine/cache_manager/standard.py @@ -38,6 +38,7 @@ ) from .storage import ( ExplicitKVStorage, + HeterogeneousExplicitKVStorage, create_attention_cache_storage, ) @@ -139,7 +140,7 @@ def _init_prefix_offload(self) -> None: host_size_gb = getattr(self.config, "prefix_cache_host_size_gb", None) if host_size_gb is None: raise RuntimeError("Prefix cache offload requires prefix_cache_host_size_gb.") - storage = self._require_explicit_storage("Prefix cache offload") + storage = self._require_uniform_explicit_storage("Prefix cache offload") kv_cache = storage.cache bytes_per_block = int( self.prefix_cache_block_size @@ -186,7 +187,12 @@ def allocate_kv_cache(self): available_memory, slot_bytes_per_layer = self._get_available_slots_info() num_layers = self.num_kv_layers - slot_bytes = num_layers * slot_bytes_per_layer + storage = self.attention_cache_storage + slot_bytes = ( + storage.bytes_per_slot() + if isinstance(storage, HeterogeneousExplicitKVStorage) + else num_layers * slot_bytes_per_layer + ) self.config.num_kvcache_slots = available_memory // slot_bytes assert self.config.num_kvcache_slots > 0, "可用显存不足以分配 KV Cache" @@ -198,11 +204,7 @@ def allocate_kv_cache(self): num_slots=self.config.num_kvcache_slots, device=self.device, ) - self.kv_cache = ( - self.attention_cache_storage.kv_cache - if isinstance(self.attention_cache_storage, ExplicitKVStorage) - else None - ) + self.kv_cache = getattr(self.attention_cache_storage, "kv_cache", None) def attention_cache_bytes_per_slot_per_layer(self) -> int: storage = getattr(self, "attention_cache_storage", None) @@ -210,15 +212,31 @@ def attention_cache_bytes_per_slot_per_layer(self) -> int: return super().attention_cache_bytes_per_slot_per_layer() return int(storage.bytes_per_slot_per_layer()) - def _require_explicit_storage(self, operation: str) -> ExplicitKVStorage: + def _logical_live_kv_bytes(self) -> int: + storage = getattr(self, "attention_cache_storage", None) + if not isinstance(storage, HeterogeneousExplicitKVStorage): + return super()._logical_live_kv_bytes() + return int(self.row_seq_lens.sum()) * storage.bytes_per_slot() + + def _require_explicit_storage( + self, operation: str + ) -> ExplicitKVStorage | HeterogeneousExplicitKVStorage: storage = self.attention_cache_storage - if not isinstance(storage, ExplicitKVStorage): + if not isinstance(storage, (ExplicitKVStorage, HeterogeneousExplicitKVStorage)): raise TypeError( f"{operation} requires ExplicitKVStorage, got " f"{type(storage).__name__}." ) return storage + def _require_uniform_explicit_storage(self, operation: str) -> ExplicitKVStorage: + storage = self.attention_cache_storage + if not isinstance(storage, ExplicitKVStorage): + raise NotImplementedError( + f"{operation} does not support heterogeneous per-layer KV shapes." + ) + return storage + def get_layer_batch_states(self, layer_idx: int) -> LayerBatchStates: return self.layer_batch_state @@ -771,6 +789,9 @@ def prefix_kv_payload_nbytes(self, payload: object) -> int: raise RuntimeError("Standard mixed prefix KV payload is missing token slots.") if not isinstance(payload.token_slots, torch.Tensor): raise RuntimeError("Standard mixed prefix KV payload has no device slots.") + storage = self.attention_cache_storage + if isinstance(storage, HeterogeneousExplicitKVStorage): + return int(payload.token_slots.numel()) * storage.bytes_per_slot() dtype_size = self._cache_slot_dtype_size() return int( payload.token_slots.numel() diff --git a/src/sparsevllm/engine/cache_manager/storage/__init__.py b/src/sparsevllm/engine/cache_manager/storage/__init__.py index 2e564758..3f435769 100644 --- a/src/sparsevllm/engine/cache_manager/storage/__init__.py +++ b/src/sparsevllm/engine/cache_manager/storage/__init__.py @@ -4,6 +4,7 @@ from .base import AttentionCacheStorage, CacheLayout from .explicit_kv import ExplicitKVStorage +from .heterogeneous_explicit_kv import HeterogeneousExplicitKVStorage if TYPE_CHECKING: from .mla_latent import MlaLatentStorage @@ -23,6 +24,18 @@ def create_attention_cache_storage( ) dtype = config.hf_config.torch_dtype if layout is CacheLayout.EXPLICIT_KV: + runtime_layout = getattr(config, "runtime_layout", None) + parallel_topology = getattr(config, "parallel_topology", None) + layer_shapes = ( + runtime_layout.local_kv_shapes(parallel_topology.attention_tp_size) + if runtime_layout is not None and parallel_topology is not None + else () + ) + if len(set(layer_shapes)) > 1: + return HeterogeneousExplicitKVStorage( + layer_shapes=layer_shapes, + dtype=dtype, + ) return ExplicitKVStorage( num_kv_heads=num_kv_heads, head_dim=head_dim, @@ -51,6 +64,7 @@ def __getattr__(name: str): "AttentionCacheStorage", "CacheLayout", "ExplicitKVStorage", + "HeterogeneousExplicitKVStorage", "MlaLatentStorage", "create_attention_cache_storage", ] diff --git a/src/sparsevllm/engine/cache_manager/storage/heterogeneous_explicit_kv.py b/src/sparsevllm/engine/cache_manager/storage/heterogeneous_explicit_kv.py new file mode 100644 index 00000000..04a6dbb3 --- /dev/null +++ b/src/sparsevllm/engine/cache_manager/storage/heterogeneous_explicit_kv.py @@ -0,0 +1,131 @@ +from __future__ import annotations + +import torch + +from sparsevllm.kernels.triton.store_kvcache import store_kvcache + +from ..base import AttentionCacheWrite, ExplicitKVPayload, ExplicitKVWrite +from .base import CacheLayout + + +class HeterogeneousExplicitKVStorage: + """Explicit K/V tensors whose head layout may differ by layer.""" + + layout = CacheLayout.EXPLICIT_KV + + def __init__(self, *, layer_shapes: tuple[tuple[int, int], ...], dtype: torch.dtype) -> None: + self.layer_shapes = tuple((int(heads), int(dim)) for heads, dim in layer_shapes) + if not self.layer_shapes or any(heads <= 0 or dim <= 0 for heads, dim in self.layer_shapes): + raise ValueError(f"Heterogeneous KV layer shapes must be positive, got {self.layer_shapes}.") + self.dtype = dtype + self.kv_cache: list[torch.Tensor] = [] + + def allocate(self, *, num_layers: int, num_slots: int, device: torch.device) -> None: + if int(num_layers) != len(self.layer_shapes) or int(num_slots) <= 0: + raise ValueError( + "Heterogeneous KV allocation does not match its layout: " + f"layers={num_layers}/{len(self.layer_shapes)} slots={num_slots}." + ) + self.kv_cache = [ + torch.empty(2, int(num_slots), heads, dim, dtype=self.dtype, device=device) + for heads, dim in self.layer_shapes + ] + + def _layer_cache(self, layer_idx: int) -> torch.Tensor: + if not self.kv_cache: + raise RuntimeError("Heterogeneous KV storage has not been allocated.") + layer_idx = int(layer_idx) + if not 0 <= layer_idx < len(self.kv_cache): + raise IndexError(f"KV layer index {layer_idx} is outside [0, {len(self.kv_cache)}).") + return self.kv_cache[layer_idx] + + @property + def cache(self) -> list[torch.Tensor]: + if not self.kv_cache: + raise RuntimeError("Heterogeneous KV storage has not been allocated.") + return self.kv_cache + + def layer_payload(self, layer_idx: int) -> ExplicitKVPayload: + cache = self._layer_cache(layer_idx) + return ExplicitKVPayload(k_cache=cache[0], v_cache=cache[1]) + + def validate_slot_mapping(self, slot_mapping: torch.Tensor) -> None: + cache = self._layer_cache(0) + if slot_mapping.ndim != 1 or slot_mapping.dtype != torch.int32: + raise ValueError( + "Heterogeneous KV slot_mapping must be 1D int32, " + f"got shape={tuple(slot_mapping.shape)} dtype={slot_mapping.dtype}." + ) + if slot_mapping.device != cache.device: + raise ValueError( + f"KV slot_mapping device {slot_mapping.device} does not match cache {cache.device}." + ) + + def validate_slot_mappings(self, slot_mappings: tuple[torch.Tensor, ...]) -> None: + for slot_mapping in slot_mappings: + self.validate_slot_mapping(slot_mapping) + + def store(self, layer_idx: int, slot_mapping: torch.Tensor, payload: AttentionCacheWrite) -> None: + if not isinstance(payload, ExplicitKVWrite): + raise TypeError(f"Heterogeneous KV storage requires ExplicitKVWrite, got {type(payload).__name__}.") + destination = self.layer_payload(layer_idx) + expected = (int(payload.key.shape[0]), *self.layer_shapes[int(layer_idx)]) + if tuple(payload.key.shape) != expected or tuple(payload.value.shape) != expected: + raise ValueError( + f"KV payload for layer {layer_idx} must have shape {expected}, " + f"got K={tuple(payload.key.shape)} V={tuple(payload.value.shape)}." + ) + if payload.key.dtype != self.dtype or payload.value.dtype != self.dtype: + raise TypeError( + f"KV payload requires dtype={self.dtype}, got K={payload.key.dtype} V={payload.value.dtype}." + ) + if slot_mapping.shape != (int(payload.key.shape[0]),): + raise ValueError( + "Heterogeneous KV slot_mapping must match the token dimension, " + f"got slots={tuple(slot_mapping.shape)} tokens={payload.key.shape[0]}." + ) + if payload.key.device != destination.k_cache.device or payload.value.device != destination.v_cache.device: + raise ValueError( + "Heterogeneous KV payload must share the cache device, got " + f"K={payload.key.device} V={payload.value.device} cache={destination.k_cache.device}." + ) + self.validate_slot_mapping(slot_mapping) + if payload.key.is_cuda: + store_kvcache(payload.key, payload.value, destination.k_cache, destination.v_cache, slot_mapping) + else: + slots = slot_mapping.to(torch.long) + destination.k_cache.index_copy_(0, slots, payload.key) + destination.v_cache.index_copy_(0, slots, payload.value) + + def bytes_per_slot_per_layer(self) -> int: + total = self.bytes_per_slot() + return (total + len(self.layer_shapes) - 1) // len(self.layer_shapes) + + def bytes_per_slot(self) -> int: + element_size = torch.tensor([], dtype=self.dtype).element_size() + return sum(2 * heads * dim * element_size for heads, dim in self.layer_shapes) + + def slot_capacity(self) -> int: + return int(self._layer_cache(0).shape[1]) + + @torch.no_grad() + def copy_slots( + self, + layer_idx: int, + source_slots: torch.Tensor, + destination_slots: torch.Tensor, + ) -> None: + payload = self.layer_payload(layer_idx) + source = source_slots.to(device=payload.k_cache.device, dtype=torch.long).reshape(-1) + destination = destination_slots.to(device=payload.k_cache.device, dtype=torch.long).reshape(-1) + if source.shape != destination.shape: + raise ValueError( + f"KV slot copy requires equal shapes, got {tuple(source.shape)} and {tuple(destination.shape)}." + ) + if source.numel() == 0: + return + payload.k_cache.index_copy_(0, destination, payload.k_cache.index_select(0, source)) + payload.v_cache.index_copy_(0, destination, payload.v_cache.index_select(0, source)) + + def accounting_tensors(self) -> tuple[torch.Tensor, ...]: + return tuple(self.kv_cache) diff --git a/src/sparsevllm/engine/model_runner.py b/src/sparsevllm/engine/model_runner.py index c688334f..868d6928 100644 --- a/src/sparsevllm/engine/model_runner.py +++ b/src/sparsevllm/engine/model_runner.py @@ -65,6 +65,11 @@ except ImportError: Qwen35MoeForCausalLM = None +try: + from sparsevllm.models.gemma4 import Gemma4ForCausalLM +except ImportError: + Gemma4ForCausalLM = None + def _create_model(hf_config, model_spec: ModelSpec, **runtime_kwargs): class_name = model_spec.runtime_class_name diff --git a/src/sparsevllm/kernels/triton/gemma4_context_attention.py b/src/sparsevllm/kernels/triton/gemma4_context_attention.py new file mode 100644 index 00000000..ba1ac1f5 --- /dev/null +++ b/src/sparsevllm/kernels/triton/gemma4_context_attention.py @@ -0,0 +1,215 @@ +from __future__ import annotations + +import torch +import triton +import triton.language as tl + + +@triton.jit +def _gemma4_context_attention_kernel( + q, + k, + v, + output, + q_start, + context_lens, + cached_prefix_lens, + active_slots, + req_indices, + attn_score, + stride_qt, + stride_qh, + stride_kt, + stride_kh, + stride_vt, + stride_vh, + stride_ot, + stride_oh, + stride_sb, + stride_ss, + stride_asb, + stride_ash, + stride_asl, + group_size, + NUM_HEADS: tl.constexpr, + HEAD_DIM: tl.constexpr, + BLOCK_M: tl.constexpr, + BLOCK_N: tl.constexpr, + WINDOW: tl.constexpr, + SCORE_MODE: tl.constexpr, +): + query_block = tl.program_id(0) + batch_head = tl.program_id(1) + batch = batch_head // NUM_HEADS + query_head = batch_head % NUM_HEADS + kv_head = query_head // group_size + query_start = tl.load(q_start + batch) + prefix_len = tl.load(cached_prefix_lens + batch) + query_len = tl.load(context_lens + batch) - prefix_len + request = tl.load(req_indices + batch) + query_positions = query_block * BLOCK_M + tl.arange(0, BLOCK_M) + dims = tl.arange(0, HEAD_DIM) + query = tl.load( + q + + (query_start + query_positions[:, None]) * stride_qt + + query_head * stride_qh + + dims[None, :], + mask=query_positions[:, None] < query_len, + other=0.0, + ) + max_logit = tl.full((BLOCK_M,), -float("inf"), tl.float32) + denominator = tl.zeros((BLOCK_M,), tl.float32) + accumulator = tl.zeros((BLOCK_M, HEAD_DIM), tl.float32) + max_key = tl.minimum( + prefix_len + (query_block + 1) * BLOCK_M, prefix_len + query_len + ) + for key_start in range(0, max_key, BLOCK_N): + key_positions = key_start + tl.arange(0, BLOCK_N) + slots = tl.load( + active_slots + request * stride_sb + key_positions * stride_ss, + mask=key_positions < max_key, + other=0, + ) + key = tl.load( + k + slots[None, :] * stride_kt + kv_head * stride_kh + dims[:, None], + mask=key_positions[None, :] < max_key, + other=0.0, + ) + logits = tl.dot(query, key) * 1.4426950408889634 + absolute_queries = prefix_len + query_positions[:, None] + visible = key_positions[None, :] <= absolute_queries + if WINDOW > 0: + visible &= key_positions[None, :] > absolute_queries - WINDOW + if SCORE_MODE == 3: + score = tl.sum( + tl.where(visible, logits * 0.6931471805599453, 0.0), + axis=0, + ) + tl.atomic_add( + attn_score + + batch * stride_asb + + query_head * stride_ash + + key_positions * stride_asl, + score, + mask=key_positions < max_key, + ) + elif SCORE_MODE == 2: + score = ( + tl.sum( + tl.where(visible, logits * 0.6931471805599453, 0.0), + axis=0, + ) + / query_len + ) + tl.atomic_max( + attn_score + batch * stride_asb + key_positions * stride_asl, + score, + mask=key_positions < max_key, + ) + logits = tl.where(visible, logits, -float("inf")) + block_max = tl.max(logits, axis=1) + new_max = tl.maximum(max_logit, block_max) + probabilities = tl.exp2(logits - new_max[:, None]) + correction = tl.exp2(max_logit - new_max) + denominator = denominator * correction + tl.sum(probabilities, axis=1) + accumulator *= correction[:, None] + value = tl.load( + v + slots[:, None] * stride_vt + kv_head * stride_vh + dims[None, :], + mask=key_positions[:, None] < max_key, + other=0.0, + ) + accumulator = tl.dot(probabilities.to(value.dtype), value, accumulator) + max_logit = new_max + output_positions = query_start + query_positions + tl.store( + output + + output_positions[:, None] * stride_ot + + query_head * stride_oh + + dims[None, :], + accumulator / denominator[:, None], + mask=query_positions[:, None] < query_len, + ) + + +@torch.no_grad() +def gemma4_context_attention( + q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, + output: torch.Tensor, + req_indices: torch.Tensor, + q_start: torch.Tensor, + context_lens: torch.Tensor, + cached_prefix_lens: torch.Tensor, + max_query_len: int, + active_slots: torch.Tensor, + *, + sliding_window: int | None, + attn_score: torch.Tensor | None = None, +) -> None: + head_dim = int(q.shape[-1]) + if q.ndim != 3 or k.ndim != 3 or v.shape != k.shape or output.shape != q.shape: + raise ValueError("Gemma 4 attention requires matching rank-3 Q/K/V/output.") + if head_dim not in {256, 512} or k.shape[-1] != head_dim: + raise ValueError( + f"Gemma 4 attention requires head_dim 256 or 512, got {head_dim}." + ) + if not all(t.is_cuda for t in (q, k, v, output)): + raise TypeError("Gemma 4 attention requires CUDA tensors.") + if q.dtype not in {torch.float16, torch.bfloat16} or any( + t.dtype != q.dtype for t in (k, v, output) + ): + raise TypeError("Gemma 4 attention requires matching FP16 or BF16 tensors.") + if any(t.stride(-1) != 1 for t in (q, k, v, output)): + raise ValueError( + "Gemma 4 attention requires contiguous BF16/FP16 head dimensions." + ) + if int(q.shape[1]) % int(k.shape[1]): + raise ValueError("Gemma 4 attention requires divisible Q and KV heads.") + block_m = 32 if head_dim == 256 else 16 + block_n = block_m + if attn_score is not None and attn_score.dim() not in {2, 3}: + raise ValueError( + "Gemma 4 prefill attention scores must be [B, L] or [B, H, L], " + f"got {tuple(attn_score.shape)}." + ) + score = context_lens if attn_score is None else attn_score + score_head_stride = score.stride(1) if score.dim() == 3 else 0 + score_length_stride = score.stride(-1) + batch, num_heads = int(context_lens.numel()), int(q.shape[1]) + _gemma4_context_attention_kernel[ + (triton.cdiv(int(max_query_len), block_m), batch * num_heads) + ]( + q, + k, + v, + output, + q_start, + context_lens, + cached_prefix_lens, + active_slots, + req_indices, + score, + q.stride(0), + q.stride(1), + k.stride(0), + k.stride(1), + v.stride(0), + v.stride(1), + output.stride(0), + output.stride(1), + active_slots.stride(0), + active_slots.stride(1), + score.stride(0), + score_head_stride, + score_length_stride, + int(q.shape[1]) // int(k.shape[1]), + NUM_HEADS=num_heads, + HEAD_DIM=head_dim, + BLOCK_M=block_m, + BLOCK_N=block_n, + WINDOW=int(sliding_window or 0), + SCORE_MODE=0 if attn_score is None else attn_score.dim(), + num_warps=8, + num_stages=1, + ) diff --git a/src/sparsevllm/kernels/triton/gemma4_decode_attention.py b/src/sparsevllm/kernels/triton/gemma4_decode_attention.py new file mode 100644 index 00000000..2bcb5f1e --- /dev/null +++ b/src/sparsevllm/kernels/triton/gemma4_decode_attention.py @@ -0,0 +1,319 @@ +from __future__ import annotations + +import torch +import triton +import triton.language as tl + + +@triton.jit +def _gemma4_decode_stage1_kernel( + q, + k, + v, + active_slots, + req_indices, + context_lens, + mid_output, + mid_lse, + attn_score, + stride_qb, + stride_qh, + stride_kt, + stride_kh, + stride_vt, + stride_vh, + stride_sb, + stride_ss, + stride_mob, + stride_moh, + stride_mos, + stride_mlb, + stride_mlh, + stride_mls, + stride_asb, + stride_ash, + stride_asl, + group_size, + HEAD_DIM: tl.constexpr, + BLOCK_SEQ: tl.constexpr, + BLOCK_N: tl.constexpr, + WINDOW: tl.constexpr, + SCORE_MODE: tl.constexpr, +): + batch = tl.program_id(0) + query_head = tl.program_id(1) + sequence_block = tl.program_id(2) + kv_head = query_head // group_size + dims = tl.arange(0, HEAD_DIM) + sequence_len = tl.load(context_lens + batch) + request = tl.load(req_indices + batch) + block_start = sequence_block * BLOCK_SEQ + mid_offset = ( + batch * stride_mob + query_head * stride_moh + sequence_block * stride_mos + ) + if block_start >= sequence_len: + tl.store(mid_output + mid_offset + dims, 0.0) + tl.store( + mid_lse + + batch * stride_mlb + + query_head * stride_mlh + + sequence_block * stride_mls, + -float("inf"), + ) + return + if WINDOW > 0: + block_start = tl.maximum(block_start, sequence_len - WINDOW) + block_end = tl.minimum(sequence_len, (sequence_block + 1) * BLOCK_SEQ) + query = tl.load(q + batch * stride_qb + query_head * stride_qh + dims) + max_logit = tl.full((), -float("inf"), tl.float32) + denominator = tl.zeros((), tl.float32) + accumulator = tl.zeros((HEAD_DIM,), tl.float32) + for offset in range(0, BLOCK_SEQ, BLOCK_N): + positions = sequence_block * BLOCK_SEQ + offset + tl.arange(0, BLOCK_N) + visible = (positions >= block_start) & (positions < block_end) + slots = tl.load( + active_slots + request * stride_sb + positions * stride_ss, + mask=visible, + other=0, + ) + key = tl.load( + k + slots[None, :] * stride_kt + kv_head * stride_kh + dims[:, None], + mask=visible[None, :], + other=0.0, + ) + logits = ( + tl.reshape(tl.dot(query[None, :], key), (BLOCK_N,)) * 1.4426950408889634 + ) + if SCORE_MODE == 3: + tl.store( + attn_score + + batch * stride_asb + + query_head * stride_ash + + positions * stride_asl, + logits * 0.6931471805599453, + mask=visible, + ) + elif SCORE_MODE == 2: + tl.atomic_max( + attn_score + batch * stride_asb + positions * stride_asl, + logits * 0.6931471805599453, + mask=visible, + ) + logits = tl.where(visible, logits, -float("inf")) + block_max = tl.max(logits, axis=0) + new_max = tl.maximum(max_logit, block_max) + probabilities = tl.exp2(logits - new_max) + correction = tl.exp2(max_logit - new_max) + denominator = denominator * correction + tl.sum(probabilities, axis=0) + accumulator *= correction + value = tl.load( + v + slots[:, None] * stride_vt + kv_head * stride_vh + dims[None, :], + mask=visible[:, None], + other=0.0, + ) + accumulator += tl.reshape( + tl.dot(probabilities[None, :].to(value.dtype), value), (HEAD_DIM,) + ) + max_logit = new_max + valid_block = block_end > block_start + tl.store( + mid_output + mid_offset + dims, + tl.where(valid_block, accumulator / denominator, 0.0), + ) + tl.store( + mid_lse + + batch * stride_mlb + + query_head * stride_mlh + + sequence_block * stride_mls, + tl.where( + valid_block, + max_logit * 0.6931471805599453 + tl.log(denominator), + -float("inf"), + ), + ) + + +@torch.no_grad() +def gemma4_decode_stage1( + q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, + active_slots: torch.Tensor, + req_indices: torch.Tensor, + context_lens: torch.Tensor, + mid_output: torch.Tensor, + mid_lse: torch.Tensor, + *, + block_seq: int, + sliding_window: int | None, + attn_score: torch.Tensor | None = None, +) -> None: + head_dim = int(q.shape[-1]) + if q.ndim != 3 or k.ndim != 3 or v.shape != k.shape: + raise ValueError("Gemma 4 decode requires matching rank-3 Q/K/V.") + if head_dim not in {256, 512} or int(k.shape[-1]) != head_dim: + raise ValueError( + f"Gemma 4 decode requires head_dim 256 or 512, got {head_dim}." + ) + if not all(t.is_cuda for t in (q, k, v, mid_output, mid_lse)): + raise TypeError("Gemma 4 decode requires CUDA tensors.") + if q.dtype not in {torch.float16, torch.bfloat16} or any( + t.dtype != q.dtype for t in (k, v) + ): + raise TypeError("Gemma 4 decode requires matching FP16 or BF16 Q/K/V.") + if mid_output.dtype != torch.float32 or mid_lse.dtype != torch.float32: + raise TypeError( + "Gemma 4 decode workspace must use FP32 output and LSE tensors." + ) + if int(q.shape[1]) % int(k.shape[1]): + raise ValueError("Gemma 4 decode requires divisible Q and KV heads.") + if int(block_seq) <= 0: + raise ValueError(f"Gemma 4 decode requires block_seq > 0, got {block_seq}.") + block_n = 32 if head_dim == 256 else 16 + if attn_score is not None and attn_score.dim() not in {2, 3}: + raise ValueError( + "Gemma 4 decode attention scores must be [B, L] or [B, H, L], " + f"got {tuple(attn_score.shape)}." + ) + score = mid_lse if attn_score is None else attn_score + score_head_stride = score.stride(1) if score.dim() == 3 else 0 + score_length_stride = score.stride(-1) + _gemma4_decode_stage1_kernel[ + (int(q.shape[0]), int(q.shape[1]), int(mid_output.shape[2])) + ]( + q, + k, + v, + active_slots, + req_indices, + context_lens, + mid_output, + mid_lse, + score, + q.stride(0), + q.stride(1), + k.stride(0), + k.stride(1), + v.stride(0), + v.stride(1), + active_slots.stride(0), + active_slots.stride(1), + mid_output.stride(0), + mid_output.stride(1), + mid_output.stride(2), + mid_lse.stride(0), + mid_lse.stride(1), + mid_lse.stride(2), + score.stride(0), + score_head_stride, + score_length_stride, + int(q.shape[1]) // int(k.shape[1]), + HEAD_DIM=head_dim, + BLOCK_SEQ=int(block_seq), + BLOCK_N=block_n, + WINDOW=int(sliding_window or 0), + SCORE_MODE=0 if attn_score is None else attn_score.dim(), + num_warps=8, + num_stages=1, + ) + + +@triton.jit +def _gemma4_decode_stage2_kernel( + context_lens, + mid_output, + mid_lse, + output, + stride_mob, + stride_moh, + stride_mos, + stride_mlb, + stride_mlh, + stride_mls, + stride_ob, + stride_oh, + HEAD_DIM: tl.constexpr, + BLOCK_SEQ: tl.constexpr, + WINDOW: tl.constexpr, +): + batch = tl.program_id(0) + head = tl.program_id(1) + dims = tl.arange(0, HEAD_DIM) + sequence_len = tl.load(context_lens + batch) + first_block = 0 + if WINDOW > 0: + first_block = tl.maximum(0, sequence_len - WINDOW) // BLOCK_SEQ + block_count = (sequence_len + BLOCK_SEQ - 1) // BLOCK_SEQ + max_lse = tl.full((), -float("inf"), tl.float32) + denominator = tl.zeros((), tl.float32) + accumulator = tl.zeros((HEAD_DIM,), tl.float32) + for block in range(first_block, block_count): + lse = tl.load( + mid_lse + batch * stride_mlb + head * stride_mlh + block * stride_mls + ) + value = tl.load( + mid_output + + batch * stride_mob + + head * stride_moh + + block * stride_mos + + dims + ) + new_max = tl.maximum(max_lse, lse) + old_scale = tl.exp(max_lse - new_max) + new_scale = tl.exp(lse - new_max) + accumulator = accumulator * old_scale + value * new_scale + denominator = denominator * old_scale + new_scale + max_lse = new_max + tl.store( + output + batch * stride_ob + head * stride_oh + dims, accumulator / denominator + ) + + +@torch.no_grad() +def gemma4_decode_stage2( + mid_output: torch.Tensor, + mid_lse: torch.Tensor, + context_lens: torch.Tensor, + output: torch.Tensor, + *, + block_seq: int, + sliding_window: int | None, +) -> None: + head_dim = int(mid_output.shape[-1]) + if head_dim not in {256, 512}: + raise ValueError( + f"Gemma 4 decode stage 2 requires head_dim 256 or 512, got {head_dim}." + ) + if not all(t.is_cuda for t in (mid_output, mid_lse, output)): + raise TypeError("Gemma 4 decode stage 2 requires CUDA tensors.") + if mid_output.dtype != torch.float32 or mid_lse.dtype != torch.float32: + raise TypeError("Gemma 4 decode stage 2 workspace must use FP32 tensors.") + if output.dtype not in {torch.float16, torch.bfloat16}: + raise TypeError("Gemma 4 decode stage 2 output must use FP16 or BF16.") + if output.shape[:2] != mid_output.shape[:2] or output.shape[-1] != head_dim: + raise ValueError( + "Gemma 4 decode stage 2 requires matching batch/head/output shape." + ) + if int(block_seq) <= 0: + raise ValueError( + f"Gemma 4 decode stage 2 requires block_seq > 0, got {block_seq}." + ) + _gemma4_decode_stage2_kernel[(int(output.shape[0]), int(output.shape[1]))]( + context_lens, + mid_output, + mid_lse, + output, + mid_output.stride(0), + mid_output.stride(1), + mid_output.stride(2), + mid_lse.stride(0), + mid_lse.stride(1), + mid_lse.stride(2), + output.stride(0), + output.stride(1), + HEAD_DIM=head_dim, + BLOCK_SEQ=int(block_seq), + WINDOW=int(sliding_window or 0), + num_warps=8, + num_stages=2, + ) diff --git a/src/sparsevllm/kernels/triton/gemma4_fused_ops.py b/src/sparsevllm/kernels/triton/gemma4_fused_ops.py new file mode 100644 index 00000000..eda5838f --- /dev/null +++ b/src/sparsevllm/kernels/triton/gemma4_fused_ops.py @@ -0,0 +1,131 @@ +from __future__ import annotations + +import torch +import triton +import triton.language as tl +from triton.language.extra import libdevice + + +@triton.jit +def _rmsnorm_residual_kernel( + x_ptr, + weight_ptr, + residual_ptr, + scalar_ptr, + output_ptr, + stride, + hidden_size: tl.constexpr, + eps: tl.constexpr, + apply_scalar: tl.constexpr, + block: tl.constexpr, +): + row = tl.program_id(0) + cols = tl.arange(0, block) + mask = cols < hidden_size + offsets = row * stride + cols + x = tl.load(x_ptr + offsets, mask=mask, other=0.0).to(tl.float32) + variance = tl.sum(tl.where(mask, x * x, 0.0), axis=0) / hidden_size + x *= libdevice.pow(variance + eps, -0.5) + x *= tl.load(weight_ptr + cols, mask=mask, other=0.0).to(tl.float32) + x = x.to(x_ptr.dtype.element_ty).to(tl.float32) + output = x + tl.load(residual_ptr + offsets, mask=mask, other=0.0).to(tl.float32) + output = output.to(x_ptr.dtype.element_ty) + if apply_scalar: + output *= tl.load(scalar_ptr).to(tl.float32) + tl.store(output_ptr + offsets, output, mask=mask) + + +def gemma4_rmsnorm_residual( + x: torch.Tensor, + weight: torch.Tensor, + residual: torch.Tensor, + eps: float, + scalar: torch.Tensor | None = None, +) -> torch.Tensor: + if x.shape != residual.shape or x.stride(-1) != 1 or residual.stride(-1) != 1: + raise ValueError( + "Gemma 4 fused RMSNorm-residual requires matching contiguous features." + ) + if not x.is_cuda or not residual.is_cuda or not weight.is_cuda: + raise TypeError("Gemma 4 fused RMSNorm-residual requires CUDA tensors.") + if x.dtype not in {torch.float16, torch.bfloat16} or residual.dtype != x.dtype: + raise TypeError( + "Gemma 4 fused RMSNorm-residual requires matching FP16 or BF16 tensors." + ) + if weight.shape != (x.shape[-1],) or weight.device != x.device: + raise ValueError( + "Gemma 4 fused RMSNorm-residual requires a matching device-local weight." + ) + if scalar is not None and (scalar.numel() != 1 or scalar.device != x.device): + raise ValueError( + "Gemma 4 fused RMSNorm-residual scalar must be device-local and scalar." + ) + output = torch.empty_like(x) + rows = x.reshape(-1, x.shape[-1]) + hidden_size = int(x.shape[-1]) + block = triton.next_power_of_2(hidden_size) + _rmsnorm_residual_kernel[(rows.shape[0],)]( + x, + weight, + residual, + weight if scalar is None else scalar, + output, + rows.stride(0), + hidden_size=hidden_size, + eps=float(eps), + apply_scalar=scalar is not None, + block=block, + num_warps=min(max(block // 256, 1), 8), + ) + return output + + +@triton.jit +def _gelu_mul_kernel( + gate_ptr, + input_ptr, + rows, + cols, + gate_stride, + input_stride, + block: tl.constexpr, +): + row = tl.program_id(0) + offsets = tl.arange(0, block) + mask = offsets < cols + gate = tl.load(gate_ptr + row * gate_stride + offsets, mask=mask, other=0.0).to( + tl.float32 + ) + inner = 0.7978845608028654 * (gate + 0.044715 * gate * gate * gate) + gate = (gate * tl.sigmoid(2.0 * inner)).to(gate_ptr.dtype.element_ty) + value = tl.load(input_ptr + row * input_stride + offsets, mask=mask, other=0.0) + tl.store(gate_ptr + row * gate_stride + offsets, gate * value, mask=mask) + + +def gemma4_gelu_mul(gate: torch.Tensor, value: torch.Tensor) -> torch.Tensor: + if gate.shape != value.shape or gate.stride(-1) != 1 or value.stride(-1) != 1: + raise ValueError( + "Gemma 4 fused GELU-multiply requires matching contiguous features." + ) + if not gate.is_cuda or not value.is_cuda: + raise TypeError("Gemma 4 fused GELU-multiply requires CUDA tensors.") + if gate.dtype not in {torch.float16, torch.bfloat16} or value.dtype != gate.dtype: + raise TypeError( + "Gemma 4 fused GELU-multiply requires matching FP16 or BF16 tensors." + ) + rows, cols = gate.reshape(-1, gate.shape[-1]).shape + block = triton.next_power_of_2(cols) + _gelu_mul_kernel[(rows,)]( + gate, + value, + rows, + cols, + gate.stride(0), + value.stride(0), + block=block, + num_warps=min(max(block // 256, 1), 8), + ) + return gate + + +__all__ = ["gemma4_gelu_mul", "gemma4_rmsnorm_residual"] diff --git a/src/sparsevllm/kernels/triton/gemma4_gelu_and_mul.py b/src/sparsevllm/kernels/triton/gemma4_gelu_and_mul.py new file mode 100644 index 00000000..771dba33 --- /dev/null +++ b/src/sparsevllm/kernels/triton/gemma4_gelu_and_mul.py @@ -0,0 +1,49 @@ +import torch +import triton +import triton.language as tl + + +@triton.jit(do_not_specialize=["size_m"]) +def _gelu_tanh_and_mul_kernel( + input_ptr, + stride_m, + stride_n, + size_m, + size_n, + BLOCK_M: tl.constexpr, + BLOCK_N: tl.constexpr, +): + rows = tl.program_id(0) * BLOCK_M + tl.arange(0, BLOCK_M) + cols = tl.program_id(1) * BLOCK_N + tl.arange(0, BLOCK_N) + offsets = rows[:, None] * stride_m + cols[None, :] * stride_n + mask = (rows < size_m)[:, None] & (cols < size_n)[None, :] + gate = tl.load(input_ptr + offsets, mask=mask, other=0.0).to(tl.float32) + up = tl.load(input_ptr + offsets + size_n * stride_n, mask=mask, other=0.0) + inner = 0.7978845608028654 * (gate + 0.044715 * gate * gate * gate) + gate = gate * tl.sigmoid(2 * inner) + tl.store(input_ptr + offsets, gate.to(input_ptr.dtype.element_ty) * up, mask=mask) + + +def gelu_tanh_and_mul_fwd(input): + if not input.is_cuda or input.dtype not in {torch.float16, torch.bfloat16}: + raise TypeError("Gemma 4 GELU-and-multiply requires CUDA FP16 or BF16 input.") + if input.ndim != 2 or not input.is_contiguous() or input.shape[1] % 2: + raise ValueError( + "Gemma 4 GELU-and-multiply requires contiguous [tokens, 2 * hidden], " + f"got shape={tuple(input.shape)} contiguous={input.is_contiguous()}." + ) + size_m, size_n = input.shape[0], input.shape[1] // 2 + block_m, block_n = (32, 128) if size_m <= 256 else (128, 128) + _gelu_tanh_and_mul_kernel[ + (triton.cdiv(size_m, block_m), triton.cdiv(size_n, block_n)) + ]( + input, + input.stride(0), + input.stride(1), + size_m, + size_n, + BLOCK_M=block_m, + BLOCK_N=block_n, + num_warps=4 if size_m <= 256 else 8, + ) + return input[:, :size_n] diff --git a/src/sparsevllm/kernels/triton/gemma4_moe.py b/src/sparsevllm/kernels/triton/gemma4_moe.py new file mode 100644 index 00000000..3287dfd3 --- /dev/null +++ b/src/sparsevllm/kernels/triton/gemma4_moe.py @@ -0,0 +1,125 @@ +from __future__ import annotations + +import torch + +from sparsevllm.kernels.triton.gemma4_gelu_and_mul import gelu_tanh_and_mul_fwd +from sparsevllm.kernels.triton.moe import ( + _prepare_expert_assignment, + _routed_gemm, + _validate_fused_moe_inputs, + moe_sum, +) +from sparsevllm.kernels.triton.moe_config import device_info, resolve_moe_gemm_config + + +def _gemma4_moe_config(num_tokens: int) -> dict[str, int] | None: + if int(num_tokens) > 32: + return None + return { + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 64, + "BLOCK_SIZE_K": 128, + "GROUP_SIZE_M": 1, + "num_warps": 4, + "num_stages": 4, + } + + +def fused_gemma4_moe( + hidden_states: torch.Tensor, + w13_weight: torch.Tensor, + w2_weight: torch.Tensor, + topk_ids: torch.Tensor, + topk_weights: torch.Tensor, + *, + num_experts: int, + local_expert_start: int, +) -> torch.Tensor: + """Run Gemma 4 routed GEGLU experts without changing generic MoE kernels.""" + + num_experts, local_expert_start = int(num_experts), int(local_expert_start) + _validate_fused_moe_inputs( + hidden_states, + w13_weight, + w2_weight, + topk_ids, + topk_weights, + num_experts, + local_expert_start, + ) + num_tokens, top_k = int(hidden_states.shape[0]), int(topk_ids.shape[1]) + intermediate_size = int(w13_weight.shape[1]) // 2 + hidden_size = int(hidden_states.shape[1]) + local_expert_end = local_expert_start + int(w13_weight.shape[0]) + device_name, capability = device_info( + hidden_states.device.type, + int(hidden_states.device.index), + ) + w13_config = _gemma4_moe_config(num_tokens) or resolve_moe_gemm_config( + dtype=hidden_states.dtype, + num_tokens=num_tokens, + top_k=top_k, + num_local_experts=int(w13_weight.shape[0]), + hidden_size=hidden_size, + intermediate_size=intermediate_size, + stage="w13", + device_name=device_name, + device_capability=capability, + ).as_triton_kwargs() + alignment = _prepare_expert_assignment( + topk_ids, + block_size=w13_config["BLOCK_SIZE_M"], + num_experts=num_experts, + local_expert_start=local_expert_start, + local_expert_end=local_expert_end, + ) + w13_output = torch.empty( + (num_tokens * top_k, 2 * intermediate_size), + dtype=hidden_states.dtype, + device=hidden_states.device, + ) + _routed_gemm( + hidden_states, + w13_weight, + w13_output, + topk_weights, + alignment, + input_top_k=top_k, + multiply_routing_weight=False, + launch_config=w13_config, + ) + activated = gelu_tanh_and_mul_fwd(w13_output) + w2_config = _gemma4_moe_config(num_tokens) or resolve_moe_gemm_config( + dtype=hidden_states.dtype, + num_tokens=num_tokens, + top_k=top_k, + num_local_experts=int(w13_weight.shape[0]), + hidden_size=hidden_size, + intermediate_size=intermediate_size, + stage="w2", + device_name=device_name, + device_capability=capability, + ).as_triton_kwargs() + w2_config["BLOCK_SIZE_M"] = alignment.block_size + w2_output = torch.empty( + (num_tokens * top_k, hidden_size), + dtype=hidden_states.dtype, + device=hidden_states.device, + ) + _routed_gemm( + activated, + w2_weight, + w2_output, + topk_weights, + alignment, + input_top_k=1, + multiply_routing_weight=True, + launch_config=w2_config, + ) + return moe_sum( + w2_output.view(num_tokens, top_k, hidden_size), + topk_ids, + num_experts=num_experts, + local_expert_start=local_expert_start, + local_expert_end=local_expert_end, + ) diff --git a/src/sparsevllm/kernels/triton/gemma4_qkv_norm_rope.py b/src/sparsevllm/kernels/triton/gemma4_qkv_norm_rope.py new file mode 100644 index 00000000..83166d53 --- /dev/null +++ b/src/sparsevllm/kernels/triton/gemma4_qkv_norm_rope.py @@ -0,0 +1,227 @@ +from __future__ import annotations + +import torch +import triton +import triton.language as tl +from triton.language.extra import libdevice + + +@triton.jit +def _norm_rope( + x_ptr, + weight_ptr, + rope_ptr, + positions_ptr, + token, + head, + stride_token, + stride_head, + rope_stride, + eps: tl.constexpr, + head_dim: tl.constexpr, + block: tl.constexpr, +): + cols = tl.arange(0, block) + mask = cols < head_dim + offset = token * stride_token + head * stride_head + cols + x = tl.load(x_ptr + offset, mask=mask, other=0.0).to(tl.float32) + variance = tl.sum(tl.where(mask, x * x, 0.0), axis=0) / head_dim + x *= libdevice.pow(variance + eps, -0.5) + x *= tl.load(weight_ptr + cols, mask=mask, other=0.0).to(tl.float32) + x = x.to(x_ptr.dtype.element_ty).to(tl.float32) + half = head_dim // 2 + pair = (cols + half) % head_dim + other = tl.load(x_ptr + token * stride_token + head * stride_head + pair).to( + tl.float32 + ) + other *= libdevice.pow(variance + eps, -0.5) + other *= tl.load(weight_ptr + pair).to(tl.float32) + other = other.to(x_ptr.dtype.element_ty).to(tl.float32) + position = tl.load(positions_ptr + token) + cos = tl.load(rope_ptr + position * rope_stride + cols % half).to(tl.float32) + sin = tl.load(rope_ptr + position * rope_stride + half + cols % half).to(tl.float32) + rotated = tl.where(cols < half, x * cos - other * sin, x * cos + other * sin) + tl.store(x_ptr + offset, rotated, mask=mask) + + +@triton.jit +def _norm( + x_ptr, + weight_ptr, + token, + head, + stride_token, + stride_head, + eps: tl.constexpr, + head_dim: tl.constexpr, + has_weight: tl.constexpr, + block: tl.constexpr, +): + cols = tl.arange(0, block) + mask = cols < head_dim + offset = token * stride_token + head * stride_head + cols + x = tl.load(x_ptr + offset, mask=mask, other=0.0).to(tl.float32) + variance = tl.sum(tl.where(mask, x * x, 0.0), axis=0) / head_dim + x *= libdevice.pow(variance + eps, -0.5) + if has_weight: + x *= tl.load(weight_ptr + cols, mask=mask, other=0.0).to(tl.float32) + tl.store(x_ptr + offset, x, mask=mask) + + +@triton.jit +def _gemma4_qkv_norm_rope_kernel( + q_ptr, + k_ptr, + v_ptr, + q_weight_ptr, + k_weight_ptr, + rope_ptr, + positions_ptr, + q_stride_token, + q_stride_head, + k_stride_token, + k_stride_head, + v_stride_token, + v_stride_head, + rope_stride, + num_q_heads: tl.constexpr, + num_kv_heads: tl.constexpr, + head_dim: tl.constexpr, + eps: tl.constexpr, + has_kv: tl.constexpr, + block: tl.constexpr, +): + token = tl.program_id(0) + head = tl.program_id(1) + if head < num_q_heads: + _norm_rope( + q_ptr, + q_weight_ptr, + rope_ptr, + positions_ptr, + token, + head, + q_stride_token, + q_stride_head, + rope_stride, + eps, + head_dim, + block, + ) + elif has_kv and head < num_q_heads + num_kv_heads: + kv_head = head - num_q_heads + _norm_rope( + k_ptr, + k_weight_ptr, + rope_ptr, + positions_ptr, + token, + kv_head, + k_stride_token, + k_stride_head, + rope_stride, + eps, + head_dim, + block, + ) + elif has_kv: + kv_head = head - num_q_heads - num_kv_heads + _norm( + v_ptr, + q_weight_ptr, + token, + kv_head, + v_stride_token, + v_stride_head, + eps, + head_dim, + False, + block, + ) + + +def gemma4_qkv_norm_rope( + q: torch.Tensor, + k: torch.Tensor | None, + v: torch.Tensor | None, + q_weight: torch.Tensor, + k_weight: torch.Tensor | None, + rope_cache: torch.Tensor, + positions: torch.Tensor, + eps: float, +) -> None: + if q.dtype not in (torch.float16, torch.bfloat16) or not q.is_cuda: + raise TypeError( + "Gemma 4 fused QKV norm-RoPE requires CUDA FP16 or BF16 tensors." + ) + if q.ndim != 3 or q.stride(-1) != 1 or q.shape[-1] not in {256, 512}: + raise ValueError( + "Gemma 4 fused QKV norm-RoPE requires contiguous rank-3 heads of 256 or 512." + ) + has_kv = k is not None and v is not None + if (k is None) != (v is None): + raise ValueError("Gemma 4 fused QKV norm-RoPE requires K and V together.") + if has_kv != (k_weight is not None): + raise ValueError( + "Gemma 4 fused QKV norm-RoPE requires K weight with K/V tensors." + ) + head_dim = int(q.shape[-1]) + if ( + q_weight.shape != (head_dim,) + or q_weight.device != q.device + or q_weight.dtype != q.dtype + ): + raise ValueError( + "Gemma 4 Q norm weight must match Q head size, device, and dtype." + ) + if has_kv and ( + k.shape[0] != q.shape[0] + or k.shape[-1] != head_dim + or v.shape != k.shape + or any( + t.device != q.device or t.dtype != q.dtype or t.stride(-1) != 1 + for t in (k, v) + ) + or k_weight.shape != (head_dim,) + or k_weight.device != q.device + or k_weight.dtype != q.dtype + ): + raise ValueError( + "Gemma 4 K/V tensors and K norm weight must match Q layout and dtype." + ) + if ( + positions.shape != (q.shape[0],) + or positions.device != q.device + or rope_cache.device != q.device + or rope_cache.shape[-1] != head_dim + ): + raise ValueError( + "Gemma 4 positions and RoPE cache must match Q tokens, device, and head size." + ) + num_kv_heads = int(k.shape[1]) if k is not None else 0 + _gemma4_qkv_norm_rope_kernel[(int(q.shape[0]), int(q.shape[1]) + 2 * num_kv_heads)]( + q, + q if k is None else k, + q if v is None else v, + q_weight, + q_weight if k_weight is None else k_weight, + rope_cache, + positions, + q.stride(0), + q.stride(1), + 0 if k is None else k.stride(0), + 0 if k is None else k.stride(1), + 0 if v is None else v.stride(0), + 0 if v is None else v.stride(1), + rope_cache.stride(0), + num_q_heads=int(q.shape[1]), + num_kv_heads=num_kv_heads, + head_dim=head_dim, + eps=float(eps), + has_kv=has_kv, + block=triton.next_power_of_2(head_dim), + num_warps=4, + ) + + +__all__ = ["gemma4_qkv_norm_rope"] diff --git a/src/sparsevllm/kernels/triton/gemma4_rmsnorm.py b/src/sparsevllm/kernels/triton/gemma4_rmsnorm.py new file mode 100644 index 00000000..ebfcf46a --- /dev/null +++ b/src/sparsevllm/kernels/triton/gemma4_rmsnorm.py @@ -0,0 +1,66 @@ +from __future__ import annotations + +import torch +import triton +import triton.language as tl +from triton.language.extra import libdevice + + +@triton.jit +def _gemma4_rmsnorm_kernel( + x_ptr, + weight_ptr, + output_ptr, + row_stride, + hidden_size: tl.constexpr, + eps: tl.constexpr, + has_weight: tl.constexpr, + block_size: tl.constexpr, +): + row = tl.program_id(0) + cols = tl.arange(0, block_size) + mask = cols < hidden_size + x = tl.load(x_ptr + row * row_stride + cols, mask=mask, other=0.0).to(tl.float32) + variance = tl.sum(tl.where(mask, x * x, 0.0), axis=0) / hidden_size + output = x * libdevice.pow(variance + eps, -0.5) + if has_weight: + output *= tl.load(weight_ptr + cols, mask=mask, other=0.0).to(tl.float32) + tl.store(output_ptr + row * hidden_size + cols, output, mask=mask) + + +def gemma4_rmsnorm( + x: torch.Tensor, + weight: torch.Tensor | None, + eps: float, +) -> torch.Tensor: + if not x.is_cuda or x.dtype not in (torch.float16, torch.bfloat16): + raise TypeError("Gemma 4 Triton RMSNorm requires CUDA FP16 or BF16 input.") + if not x.is_contiguous(): + raise ValueError("Gemma 4 Triton RMSNorm requires contiguous input.") + hidden_size = int(x.shape[-1]) + if weight is not None and ( + weight.shape != (hidden_size,) + or weight.device != x.device + or weight.dtype != x.dtype + ): + raise ValueError( + "Gemma 4 RMSNorm weight must match the input feature dimension, device, and dtype." + ) + rows = x.reshape(-1, hidden_size) + output = torch.empty_like(x) + block_size = triton.next_power_of_2(hidden_size) + _gemma4_rmsnorm_kernel[(rows.shape[0],)]( + rows, + rows if weight is None else weight, + output, + rows.stride(0), + hidden_size=hidden_size, + eps=float(eps), + has_weight=weight is not None, + block_size=block_size, + num_warps=min(max(block_size // 256, 1), 8), + ) + return output + + +__all__ = ["gemma4_rmsnorm"] diff --git a/src/sparsevllm/kernels/triton/gemma4_router.py b/src/sparsevllm/kernels/triton/gemma4_router.py new file mode 100644 index 00000000..66b4c531 --- /dev/null +++ b/src/sparsevllm/kernels/triton/gemma4_router.py @@ -0,0 +1,143 @@ +from __future__ import annotations + +import torch +import triton +import triton.language as tl +from triton.language.extra import libdevice + + +@triton.jit +def _router_input_kernel( + x_ptr, + scale_ptr, + output_ptr, + stride, + root_size: tl.constexpr, + eps: tl.constexpr, + hidden_size: tl.constexpr, + block: tl.constexpr, +): + row = tl.program_id(0) + cols = tl.arange(0, block) + mask = cols < hidden_size + offsets = row * stride + cols + x = tl.load(x_ptr + offsets, mask=mask, other=0.0).to(tl.float32) + variance = tl.sum(tl.where(mask, x * x, 0.0), axis=0) / hidden_size + element_dtype = x_ptr.dtype.element_ty + x = (x * libdevice.pow(variance + eps, -0.5)).to(element_dtype).to(tl.float32) + x = ( + (x * tl.load(scale_ptr + cols, mask=mask, other=0.0)) + .to(element_dtype) + .to(tl.float32) + ) + x = (x * root_size).to(element_dtype) + tl.store(output_ptr + offsets, x, mask=mask) + + +@triton.jit +def _router_weights_kernel( + probabilities_ptr, + ids_ptr, + scale_ptr, + weights_ptr, + probabilities_stride, + top_k: tl.constexpr, + block: tl.constexpr, +): + row = tl.program_id(0) + routes = tl.arange(0, block) + mask = routes < top_k + experts = tl.load(ids_ptr + row * top_k + routes, mask=mask, other=0) + values = tl.load( + probabilities_ptr + row * probabilities_stride + experts, + mask=mask, + other=0.0, + ).to(tl.float32) + values /= tl.sum(values, axis=0) + values *= tl.load(scale_ptr + experts, mask=mask, other=0.0).to(tl.float32) + tl.store(weights_ptr + row * top_k + routes, values, mask=mask) + + +def gemma4_router_input( + hidden_states: torch.Tensor, + scale: torch.Tensor, + root_size: float, + eps: float, +) -> torch.Tensor: + if not hidden_states.is_cuda or hidden_states.dtype not in { + torch.float16, + torch.bfloat16, + }: + raise TypeError("Gemma 4 router input requires CUDA FP16 or BF16 tensors.") + if hidden_states.stride(-1) != 1 or scale.shape != (hidden_states.shape[-1],): + raise ValueError( + "Gemma 4 router input requires contiguous features and matching scale." + ) + if scale.device != hidden_states.device or scale.dtype != hidden_states.dtype: + raise TypeError( + "Gemma 4 router input scale must match activation dtype and device." + ) + output = torch.empty_like(hidden_states) + rows = hidden_states.reshape(-1, hidden_states.shape[-1]) + hidden_size = int(hidden_states.shape[-1]) + block = triton.next_power_of_2(hidden_size) + _router_input_kernel[(rows.shape[0],)]( + hidden_states, + scale, + output, + rows.stride(0), + root_size=float(root_size), + eps=float(eps), + hidden_size=hidden_size, + block=block, + num_warps=min(max(block // 256, 1), 8), + ) + return output + + +def gemma4_router_topk( + logits: torch.Tensor, + per_expert_scale: torch.Tensor, + top_k: int, +) -> tuple[torch.Tensor, torch.Tensor]: + num_experts = int(logits.shape[-1]) + if not logits.is_cuda or logits.dtype not in {torch.float16, torch.bfloat16}: + raise TypeError("Gemma 4 router top-k requires CUDA FP16 or BF16 logits.") + if ( + logits.ndim != 2 + or logits.stride(-1) != 1 + or per_expert_scale.shape != (num_experts,) + ): + raise ValueError( + "Gemma 4 router top-k requires contiguous 2D logits and matching scales." + ) + if ( + per_expert_scale.device != logits.device + or per_expert_scale.dtype != logits.dtype + ): + raise TypeError( + "Gemma 4 router expert scales must match logits dtype and device." + ) + if not 0 < int(top_k) <= num_experts: + raise ValueError( + f"Invalid Gemma 4 router top-k {top_k} for {num_experts} experts." + ) + probabilities = torch.softmax(logits, dim=-1, dtype=torch.float32) + ids = probabilities.topk(int(top_k), dim=-1).indices + weights = torch.empty( + (logits.shape[0], top_k), dtype=torch.float32, device=logits.device + ) + _router_weights_kernel[(int(logits.shape[0]),)]( + probabilities, + ids, + per_expert_scale, + weights, + probabilities.stride(0), + top_k=int(top_k), + block=triton.next_power_of_2(int(top_k)), + num_warps=1, + ) + return weights, ids + + +__all__ = ["gemma4_router_input", "gemma4_router_topk"] diff --git a/src/sparsevllm/kernels/triton/gemma4_single_block_decode_attention.py b/src/sparsevllm/kernels/triton/gemma4_single_block_decode_attention.py new file mode 100644 index 00000000..2106e7ae --- /dev/null +++ b/src/sparsevllm/kernels/triton/gemma4_single_block_decode_attention.py @@ -0,0 +1,154 @@ +from __future__ import annotations + +import torch +import triton +import triton.language as tl + + +@triton.jit +def _gemma4_single_block_decode_kernel( + q, + k, + v, + active_slots, + req_indices, + context_lens, + output, + stride_qb, + stride_qh, + stride_kt, + stride_kh, + stride_vt, + stride_vh, + stride_sb, + stride_ss, + stride_ob, + stride_oh, + GROUP_SIZE: tl.constexpr, + HEAD_DIM: tl.constexpr, + BLOCK_SEQ: tl.constexpr, + BLOCK_N: tl.constexpr, + WINDOW: tl.constexpr, +): + batch = tl.program_id(0) + kv_head = tl.program_id(1) + groups = tl.arange(0, GROUP_SIZE) + dims = tl.arange(0, HEAD_DIM) + sequence_len = tl.load(context_lens + batch) + start = tl.maximum(0, sequence_len - WINDOW) if WINDOW > 0 else 0 + query_head = kv_head * GROUP_SIZE + groups + query = tl.load( + q + batch * stride_qb + query_head[:, None] * stride_qh + dims[None, :] + ) + request = tl.load(req_indices + batch) + max_logit = tl.full((GROUP_SIZE,), -float("inf"), tl.float32) + denominator = tl.zeros((GROUP_SIZE,), tl.float32) + accumulator = tl.zeros((GROUP_SIZE, HEAD_DIM), tl.float32) + for offset in range(0, BLOCK_SEQ, BLOCK_N): + positions = offset + tl.arange(0, BLOCK_N) + visible = (positions >= start) & (positions < sequence_len) + slots = tl.load( + active_slots + request * stride_sb + positions * stride_ss, + mask=visible, + other=0, + ) + key = tl.load( + k + slots[None, :] * stride_kt + kv_head * stride_kh + dims[:, None], + mask=visible[None, :], + other=0.0, + ) + logits = tl.dot(query, key) * 1.4426950408889634 + logits = tl.where(visible[None, :], logits, -float("inf")) + block_max = tl.max(logits, axis=1) + new_max = tl.maximum(max_logit, block_max) + probabilities = tl.exp2(logits - new_max[:, None]) + correction = tl.exp2(max_logit - new_max) + denominator = denominator * correction + tl.sum(probabilities, axis=1) + accumulator *= correction[:, None] + value = tl.load( + v + slots[:, None] * stride_vt + kv_head * stride_vh + dims[None, :], + mask=visible[:, None], + other=0.0, + ) + accumulator += tl.dot(probabilities.to(value.dtype), value) + max_logit = new_max + offsets = batch * stride_ob + query_head[:, None] * stride_oh + dims[None, :] + tl.store(output + offsets, accumulator / denominator[:, None]) + + +def gemma4_single_block_decode( + q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, + active_slots: torch.Tensor, + req_indices: torch.Tensor, + context_lens: torch.Tensor, + output: torch.Tensor, + *, + block_seq: int, + sliding_window: int | None, +) -> None: + if q.ndim != 3 or k.ndim != 3 or v.shape != k.shape or output.shape != q.shape: + raise ValueError( + "Gemma 4 single-block decode requires matching rank-3 Q/K/V/output." + ) + if not all(t.is_cuda for t in (q, k, v, output)): + raise TypeError("Gemma 4 single-block decode requires CUDA tensors.") + if q.dtype not in {torch.float16, torch.bfloat16} or any( + t.dtype != q.dtype for t in (k, v, output) + ): + raise TypeError( + "Gemma 4 single-block decode requires matching FP16 or BF16 tensors." + ) + if any(t.stride(-1) != 1 for t in (q, k, v, output)): + raise ValueError( + "Gemma 4 single-block decode requires contiguous head dimensions." + ) + if int(q.shape[1]) % int(k.shape[1]): + raise ValueError( + "Gemma 4 single-block decode requires divisible Q and KV heads." + ) + group_size = int(q.shape[1]) // int(k.shape[1]) + if group_size not in {2, 4, 8}: + raise ValueError( + f"Gemma 4 single-block decode requires GQA group 2, 4, or 8, got {group_size}." + ) + head_dim = int(q.shape[-1]) + if head_dim not in {256, 512} or int(k.shape[-1]) != head_dim: + raise ValueError( + f"Gemma 4 single-block decode requires head_dim 256 or 512, got {head_dim}." + ) + if int(block_seq) <= 0: + raise ValueError( + f"Gemma 4 single-block decode requires block_seq > 0, got {block_seq}." + ) + block_n = 32 if head_dim == 256 else 16 + _gemma4_single_block_decode_kernel[(int(q.shape[0]), int(k.shape[1]))]( + q, + k, + v, + active_slots, + req_indices, + context_lens, + output, + q.stride(0), + q.stride(1), + k.stride(0), + k.stride(1), + v.stride(0), + v.stride(1), + active_slots.stride(0), + active_slots.stride(1), + output.stride(0), + output.stride(1), + GROUP_SIZE=group_size, + HEAD_DIM=head_dim, + BLOCK_SEQ=int(block_seq), + BLOCK_N=block_n, + WINDOW=int(sliding_window or 0), + num_warps=8, + num_stages=1, + ) + + +__all__ = ["gemma4_single_block_decode"] diff --git a/src/sparsevllm/layers/activation.py b/src/sparsevllm/layers/activation.py index 393c609f..456e6c6c 100755 --- a/src/sparsevllm/layers/activation.py +++ b/src/sparsevllm/layers/activation.py @@ -2,7 +2,9 @@ from torch import nn from sparsevllm.operators.activation import ( + GeluTanhAndMulProvider, SiluAndMulProvider, + TorchGeluTanhAndMulProvider, TorchSiluAndMulProvider, ) @@ -15,3 +17,14 @@ def __init__(self, provider: SiluAndMulProvider | None = None): def forward(self, x: torch.Tensor) -> torch.Tensor: return self.provider(x) + + +class GeluTanhAndMul(nn.Module): + def __init__(self, provider: GeluTanhAndMulProvider | None = None): + super().__init__() + self.provider = ( + provider if provider is not None else TorchGeluTanhAndMulProvider() + ) + + def forward(self, x: torch.Tensor) -> torch.Tensor: + return self.provider(x) diff --git a/src/sparsevllm/layers/gemma4_rmsnorm.py b/src/sparsevllm/layers/gemma4_rmsnorm.py new file mode 100644 index 00000000..3e3cc012 --- /dev/null +++ b/src/sparsevllm/layers/gemma4_rmsnorm.py @@ -0,0 +1,32 @@ +from __future__ import annotations + +import torch +from torch import nn + + +def _torch_rmsnorm(x: torch.Tensor, weight: torch.Tensor | None, eps: float) -> torch.Tensor: + output = x.float() + output *= torch.pow(output.square().mean(-1, keepdim=True) + eps, -0.5) + if weight is not None: + output *= weight.float() + return output.to(x.dtype) + + +class Gemma4RMSNorm(nn.Module): + def __init__(self, hidden_size: int, eps: float = 1e-6, *, with_scale: bool = True) -> None: + super().__init__() + self.eps = float(eps) + if with_scale: + self.weight = nn.Parameter(torch.ones(hidden_size)) + else: + self.register_parameter("weight", None) + + def forward(self, x: torch.Tensor) -> torch.Tensor: + if not x.is_cuda: + return _torch_rmsnorm(x, self.weight, self.eps) + from sparsevllm.kernels.triton.gemma4_rmsnorm import gemma4_rmsnorm + + return gemma4_rmsnorm(x, self.weight, self.eps) + + +__all__ = ["Gemma4RMSNorm"] diff --git a/src/sparsevllm/layers/linear.py b/src/sparsevllm/layers/linear.py index 995f0664..0a4bd4a2 100755 --- a/src/sparsevllm/layers/linear.py +++ b/src/sparsevllm/layers/linear.py @@ -432,6 +432,99 @@ def load_quantized_weight( ) +class ReplicatedKVQKVParallelLinear(QKVParallelLinear): + """QKV projection with replicated KV heads when TP exceeds KV heads.""" + + def __init__( + self, + hidden_size: int, + head_size: int, + total_num_heads: int, + total_num_kv_heads: int, + bias: bool = False, + quantization=None, + ): + tp_size = get_parallel_context().tp_size + if total_num_kv_heads >= tp_size: + raise ValueError( + "ReplicatedKVQKVParallelLinear requires total_num_kv_heads < TP size, " + f"got KV heads={total_num_kv_heads}, TP={tp_size}." + ) + self.head_size = int(head_size) + self.num_heads = divide(total_num_heads, tp_size) + self.total_num_kv_heads = int(total_num_kv_heads) + self.num_kv_head_replicas = divide(tp_size, self.total_num_kv_heads) + self.num_kv_heads = 1 + output_size = (self.num_heads + 2) * self.head_size + LinearBase.__init__(self, hidden_size, output_size, bias, 0, quantization=quantization) + + def _shard(self, loaded_shard_id: str) -> tuple[int, int, int, int]: + if loaded_shard_id == "q": + return self.num_heads * self.head_size, 0, self.tp_size, self.tp_rank + if loaded_shard_id not in ("k", "v"): + raise ValueError(f"Invalid QKV shard id {loaded_shard_id!r}.") + offset = self.num_heads * self.head_size + if loaded_shard_id == "v": + offset += self.head_size + return self.head_size, offset, self.total_num_kv_heads, self.tp_rank // self.num_kv_head_replicas + + def rank_local_weight_slice( + self, + source_shape: tuple[int, ...], + *, + loaded_shard_id=None, + is_scale: bool = False, + ) -> tuple[slice, ...] | None: + if self.tp_size == 1 or loaded_shard_id is None: + return None + _, _, shard_count, shard_rank = self._shard(str(loaded_shard_id)) + shard_size = divide(int(source_shape[self.tp_dim]), shard_count) + if is_scale and shard_size * shard_count != int(source_shape[self.tp_dim]): + raise ValueError("Replicated KV FP8 scale is not shardable.") + slices = [slice(None)] * len(source_shape) + slices[self.tp_dim] = slice(shard_rank * shard_size, (shard_rank + 1) * shard_size) + return tuple(slices) + + def weight_loader(self, param: nn.Parameter, loaded_weight: torch.Tensor, loaded_shard_id: str): + shard_size, shard_offset, shard_count, shard_rank = self._shard(loaded_shard_id) + target = param.data.narrow(self.tp_dim, shard_offset, shard_size) + if loaded_weight.size(self.tp_dim) != shard_size: + loaded_weight = loaded_weight.chunk(shard_count, self.tp_dim)[shard_rank] + target.copy_(loaded_weight) + + def load_quantized_weight( + self, + loaded_weight: torch.Tensor, + loaded_scale: torch.Tensor, + loaded_shard_id: str, + ) -> None: + self._ensure_quantized_loader() + shard_size, shard_offset, shard_count, shard_rank = self._shard(loaded_shard_id) + if shard_offset % 128 or shard_size % 128: + raise ValueError( + "Replicated KV QKV FP8 loading requires 128-aligned shards, " + f"got offset={shard_offset}, size={shard_size}." + ) + weight_target = self.weight.data.narrow(0, shard_offset, shard_size) + weight_shard = ( + loaded_weight + if loaded_weight.size(0) == shard_size + else loaded_weight.chunk(shard_count, 0)[shard_rank] + ) + scale_target = self.weight_scale_inv.narrow(0, shard_offset // 128, shard_size // 128) + scale_shard = ( + loaded_scale + if tuple(loaded_scale.shape) == tuple(scale_target.shape) + else loaded_scale.chunk(shard_count, 0)[shard_rank] + ) + self._copy_quantized_weight_and_scale( + weight_shard, + scale_shard, + weight_target=weight_target, + scale_target=scale_target, + ) + + class RowParallelLinear(LinearBase): def __init__( diff --git a/src/sparsevllm/layers/packed_moe.py b/src/sparsevllm/layers/packed_moe.py index 9a4d0f2e..55938a1b 100644 --- a/src/sparsevllm/layers/packed_moe.py +++ b/src/sparsevllm/layers/packed_moe.py @@ -40,6 +40,7 @@ def __init__( cuda_graph: bool, routing_method: str = "softmax", scale_dtype: torch.dtype | None = None, + activation: str = "silu", model_label: str = "PackedMoE", provider_resolver: Callable[[MoeOpSpec], MoeProvider] = resolve_moe_provider, parallel_context=None, @@ -121,6 +122,7 @@ def __init__( tp_size=int(self.tp_size), routing_method=str(routing_method), scale_dtype=scale_dtype, + activation=str(activation), ) self.provider = provider_resolver(self.op_spec) self.w13_weight = nn.Parameter( diff --git a/src/sparsevllm/method_registry.py b/src/sparsevllm/method_registry.py index f2447f1e..bd54f365 100644 --- a/src/sparsevllm/method_registry.py +++ b/src/sparsevllm/method_registry.py @@ -60,7 +60,7 @@ "skipkv", } -H2O_SUPPORTED_MODEL_TYPES = frozenset(MODEL_SPECS) +H2O_SUPPORTED_MODEL_TYPES = frozenset(MODEL_SPECS) - {"gemma4"} SKIPKV_ASSET_MODEL_NAMES = frozenset( { @@ -144,6 +144,13 @@ class ModelRuntimeCompatibility: ), ) +GEMMA4_COMPATIBILITY = ModelRuntimeCompatibility( + sparse_methods=frozenset({"", "streamingllm", "omnikv"}), + prefix_cache_methods=frozenset({"", "streamingllm", "omnikv"}), + requires_eager=False, + decode_cuda_graph_methods=frozenset({"", "streamingllm", "omnikv"}), +) + MODEL_RUNTIME_COMPATIBILITY = { **{ (model_type, ParallelMode.STANDARD): DENSE_MODEL_COMPATIBILITY @@ -157,6 +164,8 @@ class ModelRuntimeCompatibility: ("minimax_m2", ParallelMode.OUTER_TP_MOE): MINIMAX_M2_TP_EP_COMPATIBILITY, ("glm4_moe_lite", ParallelMode.STANDARD): GLM4_MOE_LITE_EP_COMPATIBILITY, ("glm4_moe_lite", ParallelMode.OUTER_TP_MOE): GLM4_MOE_LITE_EP_COMPATIBILITY, + ("gemma4", ParallelMode.STANDARD): GEMMA4_COMPATIBILITY, + ("gemma4", ParallelMode.OUTER_TP_MOE): GEMMA4_COMPATIBILITY, } # All shipped cache managers now expose a graph-stable decode preparation path. diff --git a/src/sparsevllm/models/checkpoint.py b/src/sparsevllm/models/checkpoint.py index 89ab4cf3..34c2cdeb 100644 --- a/src/sparsevllm/models/checkpoint.py +++ b/src/sparsevllm/models/checkpoint.py @@ -270,10 +270,33 @@ def _qwen3_moe_checkpoint(_outer, config, raw, quantization, topology) -> None: _validate_qwen3_moe(config, raw, quantization, topology) +def _gemma4_checkpoint(outer, config, _raw, quantization, topology) -> None: + _validate_architecture("Gemma 4", outer, "Gemma4ForConditionalGeneration") + _validate_bf16("Gemma 4", config, "BF16 weights") + _validate_fields( + "Gemma 4", + config, + { + "attention_bias": False, + "hidden_activation": "gelu_pytorch_tanh", + "rms_norm_eps": 1.0e-6, + "tie_word_embeddings": True, + }, + ) + if quantization.enabled: + raise NotImplementedError("Gemma 4 currently supports unquantized BF16 checkpoints only.") + enable_moe = bool(config_get(config, "enable_moe_block", False)) + if topology.expert_parallel_size > 1 and not enable_moe: + raise ValueError("Gemma 4 dense checkpoints require expert_parallel_size=1.") + if enable_moe and not int(config_get(config, "num_experts", 0) or 0): + raise ValueError("Gemma 4 MoE requires a positive num_experts.") + + CHECKPOINT_VALIDATORS = { "qwen3": _qwen3_checkpoint, "qwen3_moe": _qwen3_moe_checkpoint, "qwen3_5": _qwen35_checkpoint, "qwen3_5_moe": _qwen35_moe_checkpoint, "minimax_m2": _minimax_checkpoint, + "gemma4": _gemma4_checkpoint, } diff --git a/src/sparsevllm/models/gemma4.py b/src/sparsevllm/models/gemma4.py new file mode 100644 index 00000000..5f2da41a --- /dev/null +++ b/src/sparsevllm/models/gemma4.py @@ -0,0 +1,744 @@ +from __future__ import annotations + +import re +from typing import ClassVar + +import torch +import torch.nn.functional as F +from torch import nn +from transformers import Gemma4TextConfig + +from sparsevllm.distributed import get_parallel_context +from sparsevllm.layers.activation import GeluTanhAndMul +from sparsevllm.layers.attention import Attention +from sparsevllm.layers.embed_head import ParallelLMHead, VocabParallelEmbedding +from sparsevllm.layers.gemma4_rmsnorm import Gemma4RMSNorm +from sparsevllm.layers.linear import ( + ColumnParallelLinear, + MergedColumnParallelLinear, + QKVParallelLinear, + ReplicatedKVQKVParallelLinear, + ReplicatedLinear, + RowParallelLinear, +) +from sparsevllm.layers.packed_moe import PackedMoeExperts +from sparsevllm.layers.rotary_embedding import apply_rotary_emb +from sparsevllm.operators.activation import ( + GeluTanhAndMulProvider, + resolve_gelu_tanh_and_mul_provider, +) +from sparsevllm.operators.gemma4_attention import Gemma4AttentionBackend +from sparsevllm.operators.gemma4_moe import resolve_gemma4_moe_provider +from sparsevllm.operators.moe import model_activation_dtype +from sparsevllm.platforms import device_runtime +from sparsevllm.utils.context import get_context +from sparsevllm.utils.weight_target import WeightTarget + +_EXPERT_SOURCE_RE = re.compile( + r"^model\.language_model\.layers\.(\d+)\.experts\.(gate_up_proj|down_proj)$" +) +_EXPERT_TARGET_RE = re.compile( + r"^model\.layers\.(\d+)\.experts\.(gate_up_proj|down_proj)\.expert_weight$" +) + + +class Gemma4RotaryEmbedding(nn.Module): + def __init__( + self, config: Gemma4TextConfig, layer_type: str, head_dim: int + ) -> None: + super().__init__() + parameters = dict(config.rope_parameters[layer_type]) + rope_type = str(parameters.get("rope_type", "default")) + if rope_type not in {"default", "proportional"}: + raise NotImplementedError(f"Unsupported Gemma 4 RoPE type {rope_type!r}.") + head_dim = int(head_dim) + proportion = float(parameters.get("partial_rotary_factor", 1.0)) + rotated_pairs = int(proportion * head_dim // 2) + inv_freq = 1.0 / ( + float(parameters["rope_theta"]) + ** (torch.arange(0, 2 * rotated_pairs, 2, dtype=torch.float32) / head_dim) + ) + if rotated_pairs < head_dim // 2: + inv_freq = F.pad(inv_freq, (0, head_dim // 2 - rotated_pairs)) + inv_freq.div_(float(parameters.get("factor", 1.0))) + positions = torch.arange( + int(config.max_position_embeddings), dtype=torch.float32 + ) + freqs = torch.outer(positions, inv_freq) + self.register_buffer( + "cos_sin_cache", + torch.cat((freqs.cos(), freqs.sin()), -1).unsqueeze(1), + persistent=False, + ) + + def forward( + self, + positions: torch.Tensor, + query: torch.Tensor, + key: torch.Tensor, + ) -> tuple[torch.Tensor, torch.Tensor]: + cos, sin = self.cos_sin_cache[positions].chunk(2, -1) + return apply_rotary_emb(query, cos, sin), apply_rotary_emb(key, cos, sin) + + def forward_query( + self, + positions: torch.Tensor, + query: torch.Tensor, + ) -> torch.Tensor: + cos, sin = self.cos_sin_cache[positions].chunk(2, -1) + return apply_rotary_emb(query, cos, sin) + + +class _Gemma4QKVMixin: + use_k_eq_v: bool + + def _copy_k_to_v(self, param: nn.Parameter) -> None: + k_start = self.num_heads * self.head_size + k = param.data.narrow(0, k_start, self.num_kv_heads * self.head_size) + v = param.data.narrow( + 0, k_start + self.num_kv_heads * self.head_size, k.shape[0] + ) + v.copy_(k) + + def weight_loader( + self, param: nn.Parameter, loaded_weight: torch.Tensor, loaded_shard_id: str + ): + super().weight_loader(param, loaded_weight, loaded_shard_id) + if self.use_k_eq_v and loaded_shard_id == "k": + self._copy_k_to_v(param) + + +class Gemma4QKVParallelLinear(_Gemma4QKVMixin, QKVParallelLinear): + def __init__(self, *args, use_k_eq_v: bool = False, **kwargs) -> None: + self.use_k_eq_v = bool(use_k_eq_v) + super().__init__(*args, **kwargs) + + +class Gemma4ReplicatedKVQKVParallelLinear( + _Gemma4QKVMixin, ReplicatedKVQKVParallelLinear +): + def __init__(self, *args, use_k_eq_v: bool = False, **kwargs) -> None: + self.use_k_eq_v = bool(use_k_eq_v) + super().__init__(*args, **kwargs) + + +class Gemma4QueryParallelLinear(ColumnParallelLinear): + def weight_loader( + self, + param: nn.Parameter, + loaded_weight: torch.Tensor, + loaded_shard_id: str | None = None, + ) -> None: + if loaded_shard_id not in {None, "q"}: + raise ValueError( + f"Gemma 4 shared-KV query received shard {loaded_shard_id!r}." + ) + super().weight_loader(param, loaded_weight) + + +class Gemma4Attention(nn.Module): + def __init__( + self, + config: Gemma4TextConfig, + layer_idx: int, + rotary_emb: Gemma4RotaryEmbedding, + ) -> None: + super().__init__() + parallel_context = get_parallel_context() + tp_size = parallel_context.attention_tp_size + self.layer_type = str(config.layer_types[layer_idx]) + self.is_sliding = self.layer_type == "sliding_attention" + shared_start = int(config.num_hidden_layers) - int(config.num_kv_shared_layers) + self.is_kv_shared_layer = layer_idx >= shared_start > 0 + self.sliding_window = int(config.sliding_window) if self.is_sliding else None + self.head_dim = int( + config.head_dim if self.is_sliding else config.global_head_dim + ) + self.total_num_heads = int(config.num_attention_heads) + self.num_heads = self.total_num_heads // tp_size + self.use_k_eq_v = bool(config.attention_k_eq_v and not self.is_sliding) + self.total_num_kv_heads = int( + config.num_global_key_value_heads + if self.use_k_eq_v + else config.num_key_value_heads + ) + self.num_kv_heads = max(1, self.total_num_kv_heads // tp_size) + self.q_size = self.num_heads * self.head_dim + self.kv_size = self.num_kv_heads * self.head_dim + if self.is_kv_shared_layer: + self.qkv_proj = Gemma4QueryParallelLinear( + config.hidden_size, + self.total_num_heads * self.head_dim, + bias=config.attention_bias, + quantization=getattr(config, "quantization_config", None), + ) + else: + linear_cls = ( + Gemma4ReplicatedKVQKVParallelLinear + if self.total_num_kv_heads < tp_size + else Gemma4QKVParallelLinear + ) + self.qkv_proj = linear_cls( + config.hidden_size, + self.head_dim, + self.total_num_heads, + self.total_num_kv_heads, + bias=config.attention_bias, + quantization=getattr(config, "quantization_config", None), + use_k_eq_v=self.use_k_eq_v, + ) + self.o_proj = RowParallelLinear( + self.total_num_heads * self.head_dim, + config.hidden_size, + bias=config.attention_bias, + quantization=getattr(config, "quantization_config", None), + ) + self.q_norm = Gemma4RMSNorm(self.head_dim, eps=config.rms_norm_eps) + if not self.is_kv_shared_layer: + self.k_norm = Gemma4RMSNorm(self.head_dim, eps=config.rms_norm_eps) + self.v_norm = Gemma4RMSNorm( + self.head_dim, eps=config.rms_norm_eps, with_scale=False + ) + self.rotary_emb = rotary_emb + self.attn = Attention(self.num_heads, self.head_dim, 1.0, self.num_kv_heads) + self.attn.attention_backend = Gemma4AttentionBackend( + sliding_window=self.sliding_window + ) + + def forward( + self, positions: torch.Tensor, hidden_states: torch.Tensor + ) -> torch.Tensor: + if self.is_kv_shared_layer: + q = self.qkv_proj(hidden_states).view(-1, self.num_heads, self.head_dim) + if q.is_cuda: + from sparsevllm.kernels.triton.gemma4_qkv_norm_rope import ( + gemma4_qkv_norm_rope, + ) + + gemma4_qkv_norm_rope( + q, + None, + None, + self.q_norm.weight, + None, + self.rotary_emb.cos_sin_cache, + positions, + self.q_norm.eps, + ) + else: + q = self.rotary_emb.forward_query(positions, self.q_norm(q)) + empty = q.new_empty((0, self.num_kv_heads, self.head_dim)) + return self.o_proj(self.attn(q, empty, empty).flatten(1)) + q, k, v = self.qkv_proj(hidden_states).split( + (self.q_size, self.kv_size, self.kv_size), -1 + ) + q = q.view(-1, self.num_heads, self.head_dim) + k = k.view(-1, self.num_kv_heads, self.head_dim) + v = v.view(-1, self.num_kv_heads, self.head_dim) + if q.is_cuda: + from sparsevllm.kernels.triton.gemma4_qkv_norm_rope import ( + gemma4_qkv_norm_rope, + ) + + gemma4_qkv_norm_rope( + q, + k, + v, + self.q_norm.weight, + self.k_norm.weight, + self.rotary_emb.cos_sin_cache, + positions, + self.q_norm.eps, + ) + else: + q, k = self.rotary_emb(positions, self.q_norm(q), self.k_norm(k)) + v = self.v_norm(v) + context = get_context() + context.cache_manager.save_rope_kv_if_needed(context.now_layer_idx, k, v) + output = self.attn(q, k, v) + return self.o_proj(output.flatten(1)) + + +class Gemma4MLP(nn.Module): + def __init__( + self, + config: Gemma4TextConfig, + layer_idx: int, + activation_provider: GeluTanhAndMulProvider | None, + ) -> None: + super().__init__() + shared_start = int(config.num_hidden_layers) - int(config.num_kv_shared_layers) + width = int(config.intermediate_size) * ( + 2 + if bool(config.use_double_wide_mlp) and layer_idx >= shared_start > 0 + else 1 + ) + self.gate_up_proj = MergedColumnParallelLinear( + config.hidden_size, + [width, width], + quantization=getattr(config, "quantization_config", None), + ) + self.down_proj = RowParallelLinear( + width, + config.hidden_size, + quantization=getattr(config, "quantization_config", None), + ) + self.activation = GeluTanhAndMul(activation_provider) + + def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: + return self.down_proj(self.activation(self.gate_up_proj(hidden_states))) + + +class Gemma4Router(nn.Module): + def __init__(self, config: Gemma4TextConfig) -> None: + super().__init__() + self.top_k = int(config.top_k_experts) + self.root_size = float(config.hidden_size) ** -0.5 + self.norm = Gemma4RMSNorm( + config.hidden_size, eps=config.rms_norm_eps, with_scale=False + ) + self.scale = nn.Parameter(torch.ones(config.hidden_size)) + self.proj = ReplicatedLinear(config.hidden_size, config.num_experts) + self.per_expert_scale = nn.Parameter(torch.ones(config.num_experts)) + + def forward(self, hidden_states: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: + if hidden_states.is_cuda: + from sparsevllm.kernels.triton.gemma4_router import ( + gemma4_router_input, + gemma4_router_topk, + ) + + router_input = gemma4_router_input( + hidden_states, self.scale, self.root_size, self.norm.eps + ) + return gemma4_router_topk( + self.proj(router_input), self.per_expert_scale, self.top_k + ) + logits = self.proj(self.norm(hidden_states) * self.scale * self.root_size) + probabilities = F.softmax(logits, dim=-1, dtype=torch.float32) + weights, ids = probabilities.topk(self.top_k, dim=-1) + weights.div_(weights.sum(-1, keepdim=True)).mul_(self.per_expert_scale[ids]) + return weights, ids + + +class Gemma4PackedExperts(PackedMoeExperts): + def __init__(self, config: Gemma4TextConfig) -> None: + super().__init__( + num_experts=config.num_experts, + hidden_size=config.hidden_size, + intermediate_size=config.moe_intermediate_size, + top_k=config.top_k_experts, + activation_dtype=model_activation_dtype(config), + fp8_enabled=False, + cuda_graph=bool(getattr(config, "decode_cuda_graph", False)), + activation="gelu_tanh", + model_label="Gemma4MoE", + provider_resolver=resolve_gemma4_moe_provider, + parallel_context=get_parallel_context(), + ) + + def rank_local_weight_slice( + self, + source_shape: tuple[int, ...], + *, + loaded_shard_id: str, + is_scale: bool = False, + ) -> tuple[slice, ...] | None: + if is_scale: + raise ValueError("Gemma 4 BF16 experts do not use weight scales.") + if len(source_shape) != 3 or int(source_shape[0]) != self.num_experts: + raise ValueError(f"Invalid Gemma 4 packed expert shape {source_shape}.") + if self.ep_size == 1: + return None + return ( + slice(self.local_expert_start, self.local_expert_end), + slice(None), + slice(None), + ) + + def load_packed_weight(self, projection: str, loaded_weight: torch.Tensor) -> None: + if loaded_weight.shape[0] == self.num_experts: + loaded_weight = loaded_weight[ + self.local_expert_start : self.local_expert_end + ] + if projection == "gate_up_proj": + gate, up = loaded_weight.chunk(2, 1) + gate = gate.chunk(self.tp_size, 1)[self.tp_rank] + up = up.chunk(self.tp_size, 1)[self.tp_rank] + self.w13_weight.data.copy_(torch.cat((gate, up), 1)) + projections = ("gate_proj", "up_proj") + elif projection == "down_proj": + self.w2_weight.data.copy_( + loaded_weight.chunk(self.tp_size, 2)[self.tp_rank] + ) + projections = ("down_proj",) + else: + raise ValueError( + f"Unsupported Gemma 4 packed expert projection {projection!r}." + ) + self._loaded_expert_shards.update( + (expert_id, name) + for expert_id in range(self.local_expert_start, self.local_expert_end) + for name in projections + ) + + +class Gemma4DecoderLayer(nn.Module): + def __init__( + self, + config: Gemma4TextConfig, + layer_idx: int, + activation_provider: GeluTanhAndMulProvider | None, + rotary_embeddings: nn.ModuleDict, + ) -> None: + super().__init__() + layer_type = str(config.layer_types[layer_idx]) + self.self_attn = Gemma4Attention( + config, + layer_idx, + rotary_embeddings[layer_type], + ) + self.mlp = Gemma4MLP(config, layer_idx, activation_provider) + self.input_layernorm = Gemma4RMSNorm( + config.hidden_size, eps=config.rms_norm_eps + ) + self.post_attention_layernorm = Gemma4RMSNorm( + config.hidden_size, eps=config.rms_norm_eps + ) + self.pre_feedforward_layernorm = Gemma4RMSNorm( + config.hidden_size, eps=config.rms_norm_eps + ) + self.post_feedforward_layernorm = Gemma4RMSNorm( + config.hidden_size, eps=config.rms_norm_eps + ) + self.hidden_size_per_layer_input = int(config.hidden_size_per_layer_input) + if self.hidden_size_per_layer_input: + self.per_layer_input_gate = ReplicatedLinear( + config.hidden_size, + self.hidden_size_per_layer_input, + ) + self.per_layer_projection = ReplicatedLinear( + self.hidden_size_per_layer_input, + config.hidden_size, + ) + self.post_per_layer_input_norm = Gemma4RMSNorm( + config.hidden_size, + eps=config.rms_norm_eps, + ) + self.enable_moe_block = bool(config.enable_moe_block) + if self.enable_moe_block: + self.parallel_context = get_parallel_context() + self.router = Gemma4Router(config) + self.experts = Gemma4PackedExperts(config) + self.post_feedforward_layernorm_1 = Gemma4RMSNorm( + config.hidden_size, eps=config.rms_norm_eps + ) + self.pre_feedforward_layernorm_2 = Gemma4RMSNorm( + config.hidden_size, eps=config.rms_norm_eps + ) + self.post_feedforward_layernorm_2 = Gemma4RMSNorm( + config.hidden_size, eps=config.rms_norm_eps + ) + self.layer_scalar = nn.Parameter(torch.ones(1), requires_grad=False) + + def forward( + self, + positions: torch.Tensor, + hidden_states: torch.Tensor, + per_layer_input: torch.Tensor | None = None, + ) -> torch.Tensor: + residual = hidden_states + hidden_states = self.self_attn(positions, self.input_layernorm(hidden_states)) + if hidden_states.is_cuda: + from sparsevllm.kernels.triton.gemma4_fused_ops import ( + gemma4_rmsnorm_residual, + ) + + hidden_states = gemma4_rmsnorm_residual( + hidden_states, + self.post_attention_layernorm.weight, + residual, + self.post_attention_layernorm.eps, + ) + else: + hidden_states = self.post_attention_layernorm(hidden_states) + residual + residual = hidden_states + dense_input = self.pre_feedforward_layernorm(hidden_states) + hidden_states = self.mlp(dense_input) + if self.enable_moe_block: + weights, ids = self.router(residual) + expert_output = self.parallel_context.world_all_reduce( + self.experts(self.pre_feedforward_layernorm_2(residual), ids, weights) + ) + hidden_states = self.post_feedforward_layernorm_1( + hidden_states + ) + self.post_feedforward_layernorm_2(expert_output) + if hidden_states.is_cuda: + hidden_states = gemma4_rmsnorm_residual( + hidden_states, + self.post_feedforward_layernorm.weight, + residual, + self.post_feedforward_layernorm.eps, + None if self.hidden_size_per_layer_input else self.layer_scalar, + ) + else: + hidden_states = self.post_feedforward_layernorm(hidden_states) + residual + if self.hidden_size_per_layer_input: + if per_layer_input is None: + raise RuntimeError("Gemma 4 PLE layer requires per_layer_input.") + residual = hidden_states + hidden_states = self.per_layer_input_gate(hidden_states) + if hidden_states.is_cuda: + from sparsevllm.kernels.triton.gemma4_fused_ops import gemma4_gelu_mul + + hidden_states = gemma4_gelu_mul(hidden_states, per_layer_input) + else: + hidden_states = ( + F.gelu(hidden_states, approximate="tanh") * per_layer_input + ) + hidden_states = self.per_layer_projection(hidden_states) + if hidden_states.is_cuda: + hidden_states = gemma4_rmsnorm_residual( + hidden_states, + self.post_per_layer_input_norm.weight, + residual, + self.post_per_layer_input_norm.eps, + self.layer_scalar, + ) + else: + hidden_states = self.post_per_layer_input_norm(hidden_states) + residual + return ( + hidden_states + if hidden_states.is_cuda + else hidden_states * self.layer_scalar + ) + + +class Gemma4Model(nn.Module): + def __init__( + self, + config: Gemma4TextConfig, + activation_provider: GeluTanhAndMulProvider | None, + ) -> None: + super().__init__() + self.config = config + self.embed_tokens = VocabParallelEmbedding( + config.vocab_size, config.hidden_size + ) + self.hidden_size_per_layer_input = int(config.hidden_size_per_layer_input) + if self.hidden_size_per_layer_input: + packed_size = ( + int(config.num_hidden_layers) * self.hidden_size_per_layer_input + ) + self.embed_tokens_per_layer = VocabParallelEmbedding( + config.vocab_size_per_layer_input, + packed_size, + ) + self.per_layer_model_projection = ReplicatedLinear( + config.hidden_size, + packed_size, + ) + self.per_layer_projection_norm = Gemma4RMSNorm( + self.hidden_size_per_layer_input, + eps=config.rms_norm_eps, + ) + self.per_layer_model_projection_scale = float(config.hidden_size) ** -0.5 + self.per_layer_input_scale = 2.0**-0.5 + self.rotary_embeddings = nn.ModuleDict( + { + layer_type: Gemma4RotaryEmbedding( + config, + layer_type, + config.head_dim + if layer_type == "sliding_attention" + else config.global_head_dim, + ) + for layer_type in set(config.layer_types) + } + ) + self.layers = nn.ModuleList( + Gemma4DecoderLayer( + config, + layer_idx, + activation_provider, + self.rotary_embeddings, + ) + for layer_idx in range(config.num_hidden_layers) + ) + self.norm = Gemma4RMSNorm(config.hidden_size, eps=config.rms_norm_eps) + self.embedding_scale = float(config.hidden_size) ** 0.5 + self.sparse_controller = None + + def get_per_layer_inputs( + self, + input_ids: torch.Tensor, + hidden_states: torch.Tensor, + ) -> torch.Tensor | None: + if not self.hidden_size_per_layer_input: + return None + token_inputs = ( + self.embed_tokens_per_layer(input_ids).view( + -1, + self.config.num_hidden_layers, + self.hidden_size_per_layer_input, + ) + * float(self.hidden_size_per_layer_input) ** 0.5 + ) + model_inputs = self.per_layer_projection_norm( + ( + self.per_layer_model_projection(hidden_states) + * self.per_layer_model_projection_scale + ).view_as(token_inputs) + ) + return (model_inputs + token_inputs) * self.per_layer_input_scale + + def forward(self, input_ids: torch.Tensor, positions: torch.Tensor) -> torch.Tensor: + hidden_states = self.embed_tokens(input_ids) * self.embedding_scale + per_layer_inputs = self.get_per_layer_inputs(input_ids, hidden_states) + context = get_context() + for layer_idx, layer in enumerate(self.layers): + context.now_layer_idx = layer_idx + hidden_states = layer( + positions, + hidden_states, + None if per_layer_inputs is None else per_layer_inputs[:, layer_idx], + ) + if self.sparse_controller is not None: + hidden_states, _ = self.sparse_controller.apply_activation_hook( + layer_idx, hidden_states, None, context + ) + self.sparse_controller.on_layer_end(layer_idx, context) + return self.norm(hidden_states) + + +class Gemma4ForCausalLM(nn.Module): + special_weight_loaders = (".expert_weight",) + packed_modules_mapping: ClassVar = { + "q_proj": ("qkv_proj", "q"), + "k_proj": ("qkv_proj", "k"), + "v_proj": ("qkv_proj", "v"), + "gate_proj": ("gate_up_proj", 0), + "up_proj": ("gate_up_proj", 1), + } + + def __init__( + self, + config: Gemma4TextConfig, + activation_provider: GeluTanhAndMulProvider | None = None, + ) -> None: + super().__init__() + self.config = config + self.model = Gemma4Model(config, activation_provider) + self.lm_head = ParallelLMHead(config.vocab_size, config.hidden_size) + if config.tie_word_embeddings: + self.lm_head.weight.data = self.model.embed_tokens.weight.data + self.logit_softcap = float(config.final_logit_softcapping or 0.0) + + @classmethod + def build_runtime_kwargs(cls, config, *, device, **_): + return { + "activation_provider": resolve_gelu_tanh_and_mul_provider( + activation_dtype=model_activation_dtype(config), + device_index=device.index, + ) + } + + def map_weight_name(self, source_weight_name: str) -> str | None: + match = _EXPERT_SOURCE_RE.match(source_weight_name) + if match is not None: + layer_idx, projection = match.groups() + return f"model.layers.{layer_idx}.experts.{projection}.expert_weight" + prefix = "model.language_model." + if source_weight_name.startswith(prefix + "layers."): + parts = source_weight_name.split(".") + layer_idx = int(parts[3]) + shared_start = int(self.config.num_hidden_layers) - int( + self.config.num_kv_shared_layers + ) + if layer_idx >= shared_start > 0 and parts[-2] in { + "k_proj", + "v_proj", + "k_norm", + "v_norm", + }: + return None + return ( + "model." + source_weight_name[len(prefix) :] + if source_weight_name.startswith(prefix) + else None + ) + + def resolve_special_weight(self, target_weight_name: str) -> WeightTarget | None: + match = _EXPERT_TARGET_RE.match(target_weight_name) + if match is None: + return None + layer_idx, projection = match.groups() + return WeightTarget(self.model.layers[int(layer_idx)].experts, projection) + + def load_special_weight( + self, + target_weight_name: str, + loaded_weight: torch.Tensor, + loaded_scale: torch.Tensor | None, + ) -> int: + if loaded_scale is not None: + raise ValueError("Gemma 4 BF16 packed experts do not accept scales.") + target = self.resolve_special_weight(target_weight_name) + if target is None: + return 0 + target.module.load_packed_weight(str(target.shard_id), loaded_weight) + return 1 + + def validate_loaded_weights(self, loaded_parameter_names: set[str]) -> None: + packed_experts = { + name + for name, _ in self.named_parameters() + if name.endswith((".experts.w13_weight", ".experts.w2_weight")) + } + optional_tied_head = ( + {"lm_head.weight"} if self.config.tie_word_embeddings else set() + ) + missing = sorted( + {name for name, _ in self.named_parameters()} + - packed_experts + - optional_tied_head + - loaded_parameter_names + ) + if missing: + raise ValueError(f"Missing Gemma 4 weights: {missing[:8]}.") + if self.config.enable_moe_block: + for layer in self.model.layers: + layer.experts.validate_loaded_weights() + + @torch.inference_mode() + def warmup_moe(self, num_tokens: int = 1) -> None: + if not self.config.enable_moe_block: + return + experts = self.model.layers[0].experts + hidden = torch.zeros( + (int(num_tokens), experts.hidden_size), + dtype=model_activation_dtype(self.config), + device=experts.w13_weight.device, + ) + ids = ( + torch.arange( + int(num_tokens) * int(self.config.top_k_experts), + device=hidden.device, + ) + .remainder(experts.num_experts) + .view(int(num_tokens), -1) + ) + weights = torch.full_like( + ids, 1.0 / int(self.config.top_k_experts), dtype=hidden.dtype + ) + experts(hidden, ids, weights) + device_runtime.synchronize() + + def forward(self, input_ids: torch.Tensor, positions: torch.Tensor) -> torch.Tensor: + return self.model(input_ids, positions) + + def compute_logits(self, hidden_states: torch.Tensor) -> torch.Tensor | None: + logits = self.lm_head(hidden_states) + if logits is not None and self.logit_softcap: + logits = torch.tanh(logits / self.logit_softcap) * self.logit_softcap + return logits diff --git a/src/sparsevllm/models/layout.py b/src/sparsevllm/models/layout.py index 2e3f450a..acb65fac 100644 --- a/src/sparsevllm/models/layout.py +++ b/src/sparsevllm/models/layout.py @@ -1,6 +1,6 @@ from __future__ import annotations -from dataclasses import dataclass +from dataclasses import dataclass, replace from typing import Any from sparsevllm.utils.config import config_get @@ -25,7 +25,9 @@ def resolve_attention_qk_head_dim(hf_config: Any) -> int: ) head_dim = hidden_size // num_heads if head_dim <= 0: - raise ValueError(f"Attention QK head dimension must be positive, got {head_dim}.") + raise ValueError( + f"Attention QK head dimension must be positive, got {head_dim}." + ) return head_dim @@ -79,9 +81,41 @@ class RuntimeLayout: linear_attention_layer_indices: tuple[int, ...] layer_idx_to_kv_idx: tuple[int | None, ...] kv_idx_to_layer_idx: tuple[int, ...] + kv_num_heads: tuple[int, ...] = () + kv_head_dims: tuple[int, ...] = () + + @property + def heterogeneous_kv(self) -> bool: + return len(set(zip(self.kv_num_heads, self.kv_head_dims))) > 1 + + def local_kv_shapes(self, tp_size: int) -> tuple[tuple[int, int], ...]: + tp_size = int(tp_size) + if not self.kv_num_heads: + return () + shapes = [] + for num_heads, head_dim in zip(self.kv_num_heads, self.kv_head_dims): + if num_heads >= tp_size: + if num_heads % tp_size: + raise ValueError( + f"KV heads must be divisible by TP: heads={num_heads}, TP={tp_size}." + ) + local_heads = num_heads // tp_size + else: + if tp_size % num_heads: + raise ValueError( + f"TP must be divisible by replicated KV heads: heads={num_heads}, TP={tp_size}." + ) + local_heads = 1 + shapes.append((local_heads, head_dim)) + return tuple(shapes) + + def local_kv_shape(self, layer_idx: int, tp_size: int) -> tuple[int, int] | None: + if not self.kv_num_heads: + return None + return self.local_kv_shapes(tp_size)[self.kv_layer_index(layer_idx)] @classmethod - def dense(cls, num_layers: int) -> "RuntimeLayout": + def dense(cls, num_layers: int) -> RuntimeLayout: num_layers = int(num_layers) if num_layers <= 0: raise ValueError(f"num_hidden_layers must be positive, got {num_layers}.") @@ -101,7 +135,7 @@ def from_config( hf_config: Any, *, require_mixed: bool = False, - ) -> "RuntimeLayout": + ) -> RuntimeLayout: num_layers = int(config_get(hf_config, "num_hidden_layers")) if num_layers <= 0: raise ValueError(f"num_hidden_layers must be positive, got {num_layers}.") @@ -144,7 +178,7 @@ def from_config( "Mixed-attention models require layer_types or explicit " "full/linear attention layer indices." ) - return cls.dense(num_layers) + return cls._with_attention_shapes(cls.dense(num_layers), hf_config) if full_layers is None: linear_set = set(linear_layers or ()) full_layers = [idx for idx in range(num_layers) if idx not in linear_set] @@ -173,6 +207,35 @@ def from_config( layer_to_kv: list[int | None] = [None] * num_layers for kv_idx, layer_idx in enumerate(full_tuple): layer_to_kv[layer_idx] = kv_idx + num_shared = int(config_get(hf_config, "num_kv_shared_layers", 0) or 0) + if num_shared: + if str(config_get(hf_config, "model_type", "")) != "gemma4_text": + raise NotImplementedError( + "Automatic KV sharing is only defined for Gemma 4." + ) + shared_start = num_layers - num_shared + if shared_start <= 0: + raise ValueError( + "Gemma 4 num_kv_shared_layers must leave at least one " + f"physical KV layer, got {num_shared}/{num_layers}." + ) + layer_types = tuple(config_get(hf_config, "layer_types")) + sources = { + layer_type: max( + idx + for idx in range(shared_start) + if layer_types[idx] == layer_type + ) + for layer_type in set(layer_types[shared_start:]) + } + physical_to_kv = { + layer_idx: kv_idx + for kv_idx, layer_idx in enumerate(full_tuple[:shared_start]) + } + for layer_idx in range(shared_start, num_layers): + layer_to_kv[layer_idx] = physical_to_kv[ + sources[layer_types[layer_idx]] + ] else: if len(raw_layer_to_kv) != num_layers: raise ValueError( @@ -194,37 +257,77 @@ def from_config( f"{invalid_linear}." ) - kv_pairs = sorted( - (kv_idx, layer_idx) - for layer_idx, kv_idx in enumerate(layer_to_kv) - if kv_idx is not None - ) - if len(kv_pairs) != len(full_tuple): + assigned_layers = [ + layer_idx for layer_idx in full_tuple if layer_to_kv[layer_idx] is not None + ] + if len(assigned_layers) != len(full_tuple): raise ValueError( "RuntimeLayout must assign one KV index to each full-attention " - f"layer: full={len(full_tuple)}, assigned={len(kv_pairs)}." + f"layer: full={len(full_tuple)}, assigned={len(assigned_layers)}." ) - kv_indices = [kv_idx for kv_idx, _ in kv_pairs] - if kv_indices != list(range(len(kv_pairs))): + kv_indices = sorted({int(layer_to_kv[idx]) for idx in assigned_layers}) + if kv_indices != list(range(len(kv_indices))): raise ValueError(f"KV layer indices must be contiguous, got {kv_indices}.") - kv_tuple = tuple(layer_idx for _, layer_idx in kv_pairs) + kv_tuple = tuple( + next( + layer_idx + for layer_idx in assigned_layers + if layer_to_kv[layer_idx] == kv_idx + ) + for kv_idx in kv_indices + ) configured_num_kv_layers = config_get(hf_config, "num_kv_layers", None) - if ( - configured_num_kv_layers is not None - and int(configured_num_kv_layers) != len(kv_tuple) - ): + if configured_num_kv_layers is not None and int( + configured_num_kv_layers + ) != len(kv_tuple): raise ValueError( f"num_kv_layers={configured_num_kv_layers} does not match " - f"full-attention layers={len(kv_tuple)}." + f"physical KV layers={len(kv_tuple)}." ) - return cls( - num_layers=num_layers, - num_kv_layers=len(kv_tuple), - full_attention_layer_indices=full_tuple, - linear_attention_layer_indices=linear_tuple, - layer_idx_to_kv_idx=tuple(layer_to_kv), - kv_idx_to_layer_idx=kv_tuple, + return cls._with_attention_shapes( + cls( + num_layers=num_layers, + num_kv_layers=len(kv_tuple), + full_attention_layer_indices=full_tuple, + linear_attention_layer_indices=linear_tuple, + layer_idx_to_kv_idx=tuple(layer_to_kv), + kv_idx_to_layer_idx=kv_tuple, + ), + hf_config, + ) + + @classmethod + def _with_attention_shapes( + cls, layout: RuntimeLayout, hf_config: Any + ) -> RuntimeLayout: + if str(config_get(hf_config, "model_type", "")) != "gemma4_text": + return layout + layer_types = tuple(config_get(hf_config, "layer_types")) + if len(layer_types) != layout.num_layers: + raise ValueError( + "Gemma 4 layer_types must match num_hidden_layers, " + f"got {len(layer_types)} and {layout.num_layers}." + ) + invalid_types = sorted( + set(layer_types) - {"sliding_attention", "full_attention"} + ) + if invalid_types: + raise ValueError(f"Unsupported Gemma 4 layer types: {invalid_types}.") + sliding_heads = int(config_get(hf_config, "num_key_value_heads")) + sliding_dim = int(config_get(hf_config, "head_dim")) + global_dim = int(config_get(hf_config, "global_head_dim", sliding_dim)) + use_k_eq_v = bool(config_get(hf_config, "attention_k_eq_v", False)) + global_heads = int( + config_get(hf_config, "num_global_key_value_heads", sliding_heads) + if use_k_eq_v + else sliding_heads ) + heads, dims = [], [] + for layer_idx in layout.kv_idx_to_layer_idx: + is_full = str(layer_types[layer_idx]) == "full_attention" + heads.append(global_heads if is_full else sliding_heads) + dims.append(global_dim if is_full else sliding_dim) + return replace(layout, kv_num_heads=tuple(heads), kv_head_dims=tuple(dims)) def is_full_attention(self, layer_idx: int) -> bool: return self.layer_idx_to_kv_idx[int(layer_idx)] is not None diff --git a/src/sparsevllm/models/spec.py b/src/sparsevllm/models/spec.py index 0e1fae8a..981efc28 100644 --- a/src/sparsevllm/models/spec.py +++ b/src/sparsevllm/models/spec.py @@ -17,6 +17,7 @@ class ModelSpec: supports_tiny_random: bool = True supports_expert_parallel: bool = False supports_outer_tp_moe: bool = False + outer_tp_moe_config_field: str | None = None supports_data_parallel: bool = False prefix_cache_block_size_multiple: int | None = None deltakv_checkpoint_model_types: frozenset[str] = frozenset() @@ -27,14 +28,24 @@ class ModelSpec: moe_tp_fields: tuple[str, ...] = () top_k_field: str | None = None - def topology(self, tp_size: int, ep_size: int, dp_size: int) -> ParallelTopology: + def topology( + self, + tp_size: int, + ep_size: int, + dp_size: int, + hf_config: Any | None = None, + ) -> ParallelTopology: + use_outer_tp_moe = self.supports_outer_tp_moe and ( + self.outer_tp_moe_config_field is None + or bool(config_get(hf_config, self.outer_tp_moe_config_field, False)) + ) topology = ParallelTopology( int(tp_size), int(ep_size), int(dp_size), ( ParallelMode.OUTER_TP_MOE - if self.supports_outer_tp_moe and int(tp_size) > 1 + if use_outer_tp_moe and int(tp_size) > 1 else ParallelMode.STANDARD ), ) @@ -171,6 +182,23 @@ def validate_sharding(self, hf_config: Any, topology: ParallelTopology) -> None: moe_tp_fields=("intermediate_size", "moe_intermediate_size"), top_k_field="num_experts_per_tok", ), + "gemma4": ModelSpec( + "Gemma 4", + allow_raw_config=True, + supports_tiny_random=False, + supports_expert_parallel=True, + supports_outer_tp_moe=True, + outer_tp_moe_config_field="enable_moe_block", + runtime_class_name="Gemma4ForCausalLM", + attention_tp_fields=( + "num_attention_heads", + "vocab_size", + "intermediate_size", + ), + num_experts_field="num_experts", + moe_tp_fields=("moe_intermediate_size",), + top_k_field="top_k_experts", + ), } ) diff --git a/src/sparsevllm/operators/activation.py b/src/sparsevllm/operators/activation.py index 072c9c00..428b1543 100644 --- a/src/sparsevllm/operators/activation.py +++ b/src/sparsevllm/operators/activation.py @@ -21,6 +21,17 @@ def __post_init__(self) -> None: raise ValueError("SiluAndMul input_ndim must be positive.") +@dataclass(frozen=True) +class GeluTanhAndMulSpec: + activation_dtype: torch.dtype + input_ndim: int = 2 + contiguous: bool = True + + def __post_init__(self) -> None: + if int(self.input_ndim) <= 0: + raise ValueError("GeluTanhAndMul input_ndim must be positive.") + + def _validate_input(x: torch.Tensor) -> None: if int(x.shape[-1]) % 2: raise ValueError( @@ -53,9 +64,20 @@ def __call__(self, x: torch.Tensor) -> torch.Tensor: raise NotImplementedError +class GeluTanhAndMulProvider: + name = "" + priority = 0 + + def __call__(self, x: torch.Tensor) -> torch.Tensor: + raise NotImplementedError + + SILU_AND_MUL_REGISTRY: OpRegistry[SiluAndMulSpec, SiluAndMulProvider] = OpRegistry( "SiLU-and-multiply" ) +GELU_TANH_AND_MUL_REGISTRY: OpRegistry[ + GeluTanhAndMulSpec, GeluTanhAndMulProvider +] = OpRegistry("GELU-tanh-and-multiply") @SILU_AND_MUL_REGISTRY.register @@ -114,6 +136,62 @@ def __call__(self, x: torch.Tensor) -> torch.Tensor: return gate +@GELU_TANH_AND_MUL_REGISTRY.register +class TritonGeluTanhAndMulProvider(GeluTanhAndMulProvider): + name = "triton" + priority = 10 + + def __init__(self, *, op_spec: GeluTanhAndMulSpec) -> None: + self.spec = op_spec + + @classmethod + def supports(cls, spec: GeluTanhAndMulSpec, caps: DeviceCaps) -> SupportResult: + if caps.platform != PlatformEnum.CUDA: + return SupportResult.no(f"requires CUDA, got {caps.platform.name}") + if not caps.supports_triton: + return SupportResult.no("platform does not support Triton") + if spec.activation_dtype not in (torch.float16, torch.bfloat16): + return SupportResult.no( + "requires FP16 or BF16 activations, " + f"got {spec.activation_dtype}" + ) + if int(spec.input_ndim) != 2 or not spec.contiguous: + return SupportResult.no("requires contiguous rank-2 inputs") + return SupportResult.yes() + + def __call__(self, x: torch.Tensor) -> torch.Tensor: + _validate_bound_input(x, self.spec) + if not x.is_cuda: + raise ValueError("Triton GeluTanhAndMul provider requires a CUDA input.") + from sparsevllm.kernels.triton.gemma4_gelu_and_mul import ( + gelu_tanh_and_mul_fwd, + ) + + return gelu_tanh_and_mul_fwd(x) + + +@GELU_TANH_AND_MUL_REGISTRY.register +class TorchGeluTanhAndMulProvider(GeluTanhAndMulProvider): + name = "torch" + priority = 0 + + def __init__(self, *, op_spec: GeluTanhAndMulSpec | None = None) -> None: + self.spec = op_spec + + @classmethod + def supports(cls, spec: GeluTanhAndMulSpec, caps: DeviceCaps) -> SupportResult: + del spec, caps + return SupportResult.yes() + + def __call__(self, x: torch.Tensor) -> torch.Tensor: + if self.spec is None: + _validate_input(x) + else: + _validate_bound_input(x, self.spec) + gate, up = x.chunk(2, -1) + return F.gelu(gate, approximate="tanh") * up + + def resolve_silu_and_mul_provider( *, activation_dtype: torch.dtype, @@ -137,11 +215,40 @@ def resolve_silu_and_mul_provider( ).provider +def resolve_gelu_tanh_and_mul_provider( + *, + activation_dtype: torch.dtype, + input_ndim: int = 2, + contiguous: bool = True, + device_index: int | None = None, +) -> GeluTanhAndMulProvider: + platform = platforms.current_platform + if device_index is None: + device_index = torch.cuda.current_device() if platform.is_cuda_alike() else 0 + caps = platform.get_device_caps(int(device_index)) + spec = GeluTanhAndMulSpec( + activation_dtype=activation_dtype, + input_ndim=int(input_ndim), + contiguous=bool(contiguous), + ) + return OpResolver(GELU_TANH_AND_MUL_REGISTRY).resolve( + spec, + caps, + op_spec=spec, + ).provider + + __all__ = [ + "GELU_TANH_AND_MUL_REGISTRY", "SILU_AND_MUL_REGISTRY", + "GeluTanhAndMulProvider", + "GeluTanhAndMulSpec", "SiluAndMulProvider", "SiluAndMulSpec", + "TorchGeluTanhAndMulProvider", "TorchSiluAndMulProvider", + "TritonGeluTanhAndMulProvider", "TritonSiluAndMulProvider", + "resolve_gelu_tanh_and_mul_provider", "resolve_silu_and_mul_provider", ] diff --git a/src/sparsevllm/operators/gemma4_attention.py b/src/sparsevllm/operators/gemma4_attention.py new file mode 100644 index 00000000..39e628c2 --- /dev/null +++ b/src/sparsevllm/operators/gemma4_attention.py @@ -0,0 +1,104 @@ +from __future__ import annotations + +import torch + +from sparsevllm.layers.attention_backend import ( + TritonAttentionBackend, + _require_explicit_payload, +) + + +class Gemma4AttentionBackend(TritonAttentionBackend): + """Gemma 4 attention semantics isolated from the tuned generic kernels.""" + + name = "triton_gemma4" + + def __init__(self, *, sliding_window: int | None) -> None: + super().__init__() + self.sliding_window = None if sliding_window is None else int(sliding_window) + + def run_prefill( + self, + q: torch.Tensor, + view, + *, + b_start_loc: torch.Tensor, + chunk_lens: torch.Tensor, + max_input_len: int, + ) -> torch.Tensor: + payload = _require_explicit_payload(view, operation="Gemma 4 prefill") + output = torch.empty_like(q) + from sparsevllm.kernels.triton.gemma4_context_attention import ( + gemma4_context_attention, + ) + + gemma4_context_attention( + q, + payload.k_cache, + payload.v_cache, + output, + view.meta.req_indices, + b_start_loc, + view.meta.context_lens, + view.meta.context_lens - chunk_lens, + max_input_len, + view.meta.active_slots, + sliding_window=self.sliding_window, + attn_score=view.meta.attn_score, + ) + return output + + def run_decode( + self, + q: torch.Tensor, + view, + *, + mid_o: torch.Tensor, + mid_o_logexpsum: torch.Tensor, + max_len_in_batch: int, + block_seq: int, + num_heads: int, + num_kv_heads: int, + gqa_block_n: int = 16, + gqa_num_warps: int = 2, + ) -> torch.Tensor: + del max_len_in_batch, num_heads, num_kv_heads, gqa_block_n, gqa_num_warps + payload = _require_explicit_payload(view, operation="Gemma 4 decode") + from sparsevllm.kernels.triton.gemma4_decode_attention import ( + gemma4_decode_stage1, + gemma4_decode_stage2, + ) + + group_size = int(q.shape[1]) // int(payload.k_cache.shape[1]) + if mid_o.shape[2] == 1 and view.meta.attn_score is None and group_size in {2, 4, 8}: + from sparsevllm.kernels.triton.gemma4_single_block_decode_attention import ( + gemma4_single_block_decode, + ) + + output = torch.empty_like(q) + gemma4_single_block_decode( + q, payload.k_cache, payload.v_cache, view.meta.active_slots, + view.meta.req_indices, view.meta.context_lens, output, + block_seq=block_seq, sliding_window=self.sliding_window, + ) + return output + gemma4_decode_stage1( + q, payload.k_cache, payload.v_cache, view.meta.active_slots, + view.meta.req_indices, view.meta.context_lens, mid_o, + mid_o_logexpsum, block_seq=block_seq, + sliding_window=self.sliding_window, + attn_score=view.meta.attn_score, + ) + output = torch.empty_like(q) + gemma4_decode_stage2( + mid_o, + mid_o_logexpsum, + view.meta.context_lens, + output, + block_seq=block_seq, + sliding_window=self.sliding_window, + ) + return output + + +__all__ = ["Gemma4AttentionBackend"] diff --git a/src/sparsevllm/operators/gemma4_moe.py b/src/sparsevllm/operators/gemma4_moe.py new file mode 100644 index 00000000..3c340f86 --- /dev/null +++ b/src/sparsevllm/operators/gemma4_moe.py @@ -0,0 +1,139 @@ +from __future__ import annotations + +import torch +import torch.nn.functional as F + +from sparsevllm import platforms +from sparsevllm.operators.moe import MoeOpSpec, MoeProvider +from sparsevllm.operators.registry import OpRegistry, OpResolver, SupportResult +from sparsevllm.platforms.interface import DeviceCaps, PlatformEnum + +GEMMA4_MOE_REGISTRY: OpRegistry[MoeOpSpec, MoeProvider] = OpRegistry( + "Gemma 4 routed GEGLU MoE" +) + + +@GEMMA4_MOE_REGISTRY.register +class TritonGemma4MoeProvider(MoeProvider): + name = "triton_gemma4_geglu" + priority = 10 + gate_up_order = "gate_up" + + @classmethod + def supports(cls, spec: MoeOpSpec, caps: DeviceCaps) -> SupportResult: + if spec.activation != "gelu_tanh" or spec.routing_method != "softmax": + return SupportResult.no("requires Gemma 4 GELU-tanh and softmax routing") + if caps.platform != PlatformEnum.CUDA or not caps.supports_triton: + return SupportResult.no("requires CUDA with Triton") + if spec.cuda_graph and not caps.supports_graph_capture: + return SupportResult.no("device does not support CUDA Graph capture") + if spec.activation_dtype not in {torch.bfloat16, torch.float16}: + return SupportResult.no("requires BF16 or FP16 activations") + if spec.weight_dtype != spec.activation_dtype or spec.block_shape is not None: + return SupportResult.no( + "requires unquantized experts matching activation dtype" + ) + return SupportResult.yes() + + def run( + self, + spec, + hidden_states, + topk_ids, + topk_weights, + w13_weight, + w2_weight, + w13_scale_inv, + w2_scale_inv, + *, + local_expert_start, + ep_rank, + ): + del ep_rank + if w13_scale_inv is not None or w2_scale_inv is not None: + raise RuntimeError("Gemma 4 BF16 MoE does not accept expert scales.") + from sparsevllm.kernels.triton.gemma4_moe import fused_gemma4_moe + + return fused_gemma4_moe( + hidden_states, + w13_weight, + w2_weight, + topk_ids, + topk_weights, + num_experts=spec.num_experts, + local_expert_start=local_expert_start, + ) + + +@GEMMA4_MOE_REGISTRY.register +class TorchGemma4MoeProvider(MoeProvider): + name = "torch_gemma4_geglu" + priority = 0 + + @classmethod + def supports(cls, spec: MoeOpSpec, caps: DeviceCaps) -> SupportResult: + del caps + if spec.activation != "gelu_tanh" or spec.routing_method != "softmax": + return SupportResult.no("requires Gemma 4 GELU-tanh and softmax routing") + if spec.weight_dtype != spec.activation_dtype or spec.block_shape is not None: + return SupportResult.no("requires unquantized Gemma 4 GELU-tanh experts") + return SupportResult.yes() + + def run( + self, + spec, + hidden_states, + topk_ids, + topk_weights, + w13_weight, + w2_weight, + w13_scale_inv, + w2_scale_inv, + *, + local_expert_start, + ep_rank, + ): + del ep_rank + if w13_scale_inv is not None or w2_scale_inv is not None: + raise RuntimeError("Gemma 4 Torch MoE does not accept expert scales.") + output = torch.zeros_like(hidden_states) + for local_id in range(spec.num_local_experts): + global_id = int(local_expert_start) + local_id + token_ids, routes = torch.where(topk_ids == global_id) + if token_ids.numel() == 0: + continue + gate, up = F.linear(hidden_states[token_ids], w13_weight[local_id]).chunk( + 2, -1 + ) + routed = F.linear( + F.gelu(gate, approximate="tanh") * up, w2_weight[local_id] + ) + output.index_add_( + 0, token_ids, routed * topk_weights[token_ids, routes, None] + ) + return output + + +def resolve_gemma4_moe_provider( + spec: MoeOpSpec, + *, + device_index: int | None = None, +) -> MoeProvider: + if spec.activation != "gelu_tanh": + raise ValueError( + "Gemma 4 MoE resolver requires activation='gelu_tanh', " + f"got {spec.activation!r}." + ) + platform = platforms.current_platform + if device_index is None: + device_index = torch.cuda.current_device() if platform.is_cuda_alike() else 0 + caps = platform.get_device_caps(int(device_index)) + return OpResolver(GEMMA4_MOE_REGISTRY).resolve(spec, caps).provider + + +__all__ = [ + "GEMMA4_MOE_REGISTRY", + "TorchGemma4MoeProvider", + "TritonGemma4MoeProvider", + "resolve_gemma4_moe_provider", +] diff --git a/src/sparsevllm/operators/moe.py b/src/sparsevllm/operators/moe.py index 74eec8ba..613a813e 100644 --- a/src/sparsevllm/operators/moe.py +++ b/src/sparsevllm/operators/moe.py @@ -31,6 +31,7 @@ class MoeOpSpec: tp_size: int = 1 routing_method: str = "softmax" scale_dtype: torch.dtype | None = None + activation: str = "silu" def __post_init__(self) -> None: if self.num_experts <= 0 or self.num_local_experts <= 0: @@ -61,6 +62,11 @@ def __post_init__(self) -> None: "MoE routing_method must be 'softmax' or 'biased_sigmoid', " f"got {self.routing_method!r}." ) + if self.activation not in {"silu", "gelu_tanh"}: + raise ValueError( + "MoE activation must be 'silu' or 'gelu_tanh', " + f"got {self.activation!r}." + ) def model_activation_dtype(config) -> torch.dtype: diff --git a/tests/test_activation.py b/tests/test_activation.py index bfa6b055..df3a6975 100644 --- a/tests/test_activation.py +++ b/tests/test_activation.py @@ -2,10 +2,13 @@ import torch import torch.nn.functional as F -from sparsevllm.layers.activation import SiluAndMul +from sparsevllm.layers.activation import GeluTanhAndMul, SiluAndMul from sparsevllm.operators.activation import ( + GeluTanhAndMulSpec, SiluAndMulSpec, + TorchGeluTanhAndMulProvider, TorchSiluAndMulProvider, + TritonGeluTanhAndMulProvider, TritonSiluAndMulProvider, ) @@ -29,6 +32,30 @@ def test_silu_and_mul_rejects_odd_width(): SiluAndMul(provider=TorchSiluAndMulProvider())(torch.randn(2, 7)) +def test_gemma4_gelu_and_mul_cpu_matches_reference(): + x = torch.randn(3, 16) + gate, up = x.chunk(2, -1) + expected = F.gelu(gate, approximate="tanh") * up + actual = GeluTanhAndMul(provider=TorchGeluTanhAndMulProvider())(x) + torch.testing.assert_close(actual, expected) + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA is required") +@pytest.mark.parametrize("dtype", [torch.float16, torch.bfloat16]) +@pytest.mark.parametrize("rows", [1, 256, 257]) +def test_gemma4_gelu_and_mul_cuda_matches_reference(dtype, rows): + torch.manual_seed(20260813) + x = torch.randn(rows, 1408, dtype=dtype, device="cuda") + gate, up = x.chunk(2, -1) + expected = F.gelu(gate, approximate="tanh") * up + actual_input = x.clone() + actual = GeluTanhAndMul( + TritonGeluTanhAndMulProvider(op_spec=GeluTanhAndMulSpec(dtype)) + )(actual_input) + torch.testing.assert_close(actual, expected, rtol=3e-3, atol=3e-3) + assert actual.data_ptr() == actual_input.data_ptr() + + @pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA is required") @pytest.mark.parametrize("dtype", [torch.float16, torch.bfloat16]) @pytest.mark.parametrize("rows", [8, 256, 257]) diff --git a/tests/test_attention_cache_storage.py b/tests/test_attention_cache_storage.py index 9333c967..52eb0e7d 100644 --- a/tests/test_attention_cache_storage.py +++ b/tests/test_attention_cache_storage.py @@ -15,14 +15,63 @@ MlaLatentPayload, MlaLatentWrite, ) -from sparsevllm.engine.cache_manager.standard import StandardCacheManager from sparsevllm.engine.cache_manager.snapkv import SnapKVCacheManager +from sparsevllm.engine.cache_manager.standard import StandardCacheManager from sparsevllm.engine.cache_manager.storage import ( CacheLayout, ExplicitKVStorage, + HeterogeneousExplicitKVStorage, MlaLatentStorage, create_attention_cache_storage, ) + + +def test_heterogeneous_explicit_storage_preserves_per_layer_shapes(): + storage = HeterogeneousExplicitKVStorage( + layer_shapes=((2, 256), (1, 512)), + dtype=torch.bfloat16, + ) + storage.allocate(num_layers=2, num_slots=5, device=torch.device("cpu")) + + assert storage.layer_payload(0).k_cache.shape == (5, 2, 256) + assert storage.layer_payload(1).k_cache.shape == (5, 1, 512) + assert storage.bytes_per_slot_per_layer() == 2048 + assert storage.bytes_per_slot() == 4096 + + key = torch.randn(2, 1, 512, dtype=torch.bfloat16) + value = torch.randn_like(key) + slots = torch.tensor([1, 3], dtype=torch.int32) + storage.store(1, slots, ExplicitKVWrite(key=key, value=value)) + payload = storage.layer_payload(1) + torch.testing.assert_close(payload.k_cache[slots.long()], key) + torch.testing.assert_close(payload.v_cache[slots.long()], value) + + storage.copy_slots(1, slots, torch.tensor([0, 2], dtype=torch.int32)) + torch.testing.assert_close(payload.k_cache[[0, 2]], key) + torch.testing.assert_close(payload.v_cache[[0, 2]], value) + + +def test_standard_manager_uses_exact_heterogeneous_slot_size(): + storage = HeterogeneousExplicitKVStorage( + layer_shapes=((1, 256), (1, 512)), + dtype=torch.bfloat16, + ) + manager = object.__new__(StandardCacheManager) + manager.attention_cache_storage = storage + manager.num_kv_layers = 2 + manager.device = torch.device("cpu") + manager.config = SimpleNamespace(num_kvcache_slots=-1) + manager._get_available_slots_info = lambda: ( + 7 * storage.bytes_per_slot(), + storage.bytes_per_slot_per_layer(), + ) + + manager.allocate_kv_cache() + + assert manager.config.num_kvcache_slots == 7 + assert all(cache.shape[1] == 7 for cache in storage.cache) + + def test_explicit_storage_preserves_legacy_tensor_layout_and_size(): storage = ExplicitKVStorage( num_kv_heads=2, diff --git a/tests/test_gemma4_attention_kernels.py b/tests/test_gemma4_attention_kernels.py new file mode 100644 index 00000000..802e8365 --- /dev/null +++ b/tests/test_gemma4_attention_kernels.py @@ -0,0 +1,284 @@ +from __future__ import annotations + +import pytest +import torch + +from sparsevllm.kernels.triton.gemma4_context_attention import gemma4_context_attention +from sparsevllm.kernels.triton.gemma4_decode_attention import ( + gemma4_decode_stage1, + gemma4_decode_stage2, +) +from sparsevllm.kernels.triton.gemma4_single_block_decode_attention import ( + gemma4_single_block_decode, +) + +pytestmark = pytest.mark.skipif(not torch.cuda.is_available(), reason="requires CUDA") + + +def _slots_and_lengths(): + lengths = torch.tensor([21, 13], dtype=torch.int32, device="cuda") + slots = torch.zeros((2, 21), dtype=torch.int32, device="cuda") + slots[0, :21] = torch.arange(21, dtype=torch.int32, device="cuda") + slots[1, :13] = torch.arange(21, 34, dtype=torch.int32, device="cuda") + return slots, lengths + + +def _decode_reference(q, k, v, slots, lengths, window): + output = torch.empty_like(q) + for batch, length in enumerate(lengths.tolist()): + start = max(0, length - (window or length)) + indices = slots[batch, start:length].long() + for head in range(q.shape[1]): + logits = q[batch, head] @ k[indices, head // (q.shape[1] // k.shape[1])].T + output[batch, head] = logits.softmax(-1) @ v[ + indices, head // (q.shape[1] // v.shape[1]) + ] + return output + + +@pytest.mark.parametrize("head_dim", [256, 512]) +@pytest.mark.parametrize("sliding_window", [None, 4]) +def test_gemma4_prefill_matches_torch(head_dim, sliding_window): + torch.manual_seed(3) + prefix = torch.tensor([4, 2], dtype=torch.int32, device="cuda") + chunks = torch.tensor([5, 3], dtype=torch.int32, device="cuda") + lengths = prefix + chunks + starts = torch.tensor([0, 5], dtype=torch.int32, device="cuda") + slots = torch.zeros((2, 9), dtype=torch.int32, device="cuda") + slots[0, :9] = torch.arange(9, dtype=torch.int32, device="cuda") + slots[1, :5] = torch.arange(9, 14, dtype=torch.int32, device="cuda") + key = torch.randn(14, 2, head_dim, dtype=torch.bfloat16, device="cuda") + value = torch.randn_like(key) + query = torch.randn(8, 4, head_dim, dtype=torch.bfloat16, device="cuda") + output = torch.empty_like(query) + + gemma4_context_attention( + query, + key, + value, + output, + torch.tensor([0, 1], dtype=torch.int32, device="cuda"), + starts, + lengths, + prefix, + 5, + slots, + sliding_window=sliding_window, + ) + reference = torch.empty_like(output) + for batch, (prefix_len, chunk_len, start) in enumerate( + zip(prefix.tolist(), chunks.tolist(), starts.tolist()) + ): + for offset in range(chunk_len): + end = prefix_len + offset + 1 + begin = max(0, end - (sliding_window or end)) + indices = slots[batch, begin:end].long() + for head in range(query.shape[1]): + logits = query[start + offset, head] @ key[indices, head // 2].T + reference[start + offset, head] = logits.softmax(-1) @ value[indices, head // 2] + cosine = torch.nn.functional.cosine_similarity( + output.float().flatten(), reference.float().flatten(), dim=0 + ) + assert torch.isfinite(output).all() + assert cosine > 0.999 + + +@pytest.mark.parametrize("head_dim", [256, 512]) +@pytest.mark.parametrize("sliding_window", [None, 4]) +def test_gemma4_decode_matches_torch(head_dim, sliding_window): + torch.manual_seed(7) + slots, lengths = _slots_and_lengths() + key = torch.randn(34, 2, head_dim, dtype=torch.bfloat16, device="cuda") + value = torch.randn_like(key) + query = torch.randn(2, 4, head_dim, dtype=torch.bfloat16, device="cuda") + block_seq = 8 + blocks = (int(lengths.max()) + block_seq - 1) // block_seq + mid = torch.empty(2, 4, blocks, head_dim, dtype=torch.float32, device="cuda") + lse = torch.empty(2, 4, blocks, dtype=torch.float32, device="cuda") + output = torch.empty_like(query) + + gemma4_decode_stage1( + query, + key, + value, + slots, + torch.tensor([0, 1], dtype=torch.int32, device="cuda"), + lengths, + mid, + lse, + block_seq=block_seq, + sliding_window=sliding_window, + ) + gemma4_decode_stage2( + mid, + lse, + lengths, + output, + block_seq=block_seq, + sliding_window=sliding_window, + ) + + reference = _decode_reference(query, key, value, slots, lengths, sliding_window) + cosine = torch.nn.functional.cosine_similarity( + output.float().flatten(), reference.float().flatten(), dim=0 + ) + assert torch.isfinite(output).all() + assert cosine > 0.999 + + +@pytest.mark.parametrize("group_size", [2, 4, 8]) +@pytest.mark.parametrize("head_dim", [256, 512]) +def test_gemma4_single_block_decode_matches_torch(group_size, head_dim): + torch.manual_seed(11) + slots, lengths = _slots_and_lengths() + key = torch.randn(34, 2, head_dim, dtype=torch.bfloat16, device="cuda") + value = torch.randn_like(key) + query = torch.randn(2, 2 * group_size, head_dim, dtype=torch.bfloat16, device="cuda") + output = torch.empty_like(query) + gemma4_single_block_decode( + query, + key, + value, + slots, + torch.tensor([0, 1], dtype=torch.int32, device="cuda"), + lengths, + output, + block_seq=256, + sliding_window=None, + ) + reference = _decode_reference(query, key, value, slots, lengths, None) + cosine = torch.nn.functional.cosine_similarity( + output.float().flatten(), reference.float().flatten(), dim=0 + ) + assert torch.isfinite(output).all() + assert cosine > 0.999 + + +@pytest.mark.parametrize("head_dim", [256, 512]) +def test_gemma4_decode_supports_cuda_graph(head_dim): + slots, lengths = _slots_and_lengths() + query = torch.randn(2, 4, head_dim, dtype=torch.bfloat16, device="cuda") + key = torch.randn(34, 2, head_dim, dtype=torch.bfloat16, device="cuda") + value = torch.randn_like(key) + mid = torch.empty(2, 4, 1, head_dim, dtype=torch.float32, device="cuda") + lse = torch.empty(2, 4, 1, dtype=torch.float32, device="cuda") + output = torch.empty_like(query) + request_indices = torch.tensor([0, 1], dtype=torch.int32, device="cuda") + + def run(): + gemma4_decode_stage1( + query, + key, + value, + slots, + request_indices, + lengths, + mid, + lse, + block_seq=256, + sliding_window=1024, + ) + gemma4_decode_stage2( + mid, + lse, + lengths, + output, + block_seq=256, + sliding_window=1024, + ) + + for _ in range(3): + run() + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph): + run() + query.copy_(torch.randn_like(query)) + graph.replay() + first = output.clone() + graph.replay() + assert torch.equal(first, output) + + +@pytest.mark.parametrize("score_dims", [2, 3]) +def test_gemma4_decode_collects_raw_qk_scores(score_dims): + torch.manual_seed(13) + slots, lengths = _slots_and_lengths() + head_dim = 256 + query = torch.randn(2, 4, head_dim, dtype=torch.bfloat16, device="cuda") + key = torch.randn(34, 2, head_dim, dtype=torch.bfloat16, device="cuda") + value = torch.randn_like(key) + mid = torch.empty(2, 4, 3, head_dim, dtype=torch.float32, device="cuda") + lse = torch.empty(2, 4, 3, dtype=torch.float32, device="cuda") + score = torch.full( + (2, 4, 21) if score_dims == 3 else (2, 21), + -1e20, + dtype=torch.float32, + device="cuda", + ) + gemma4_decode_stage1( + query, + key, + value, + slots, + torch.tensor([0, 1], dtype=torch.int32, device="cuda"), + lengths, + mid, + lse, + block_seq=8, + sliding_window=None, + attn_score=score, + ) + expected = torch.empty(2, 4, 21, dtype=torch.float32, device="cuda") + expected.fill_(-1e20) + for batch, length in enumerate(lengths.tolist()): + for head in range(4): + indices = slots[batch, :length].long() + expected[batch, head, :length] = ( + query[batch, head].float() + @ key[indices, head // 2].float().T + ) + expected = expected if score_dims == 3 else expected.max(1).values + torch.testing.assert_close(score, expected, rtol=2e-2, atol=1.0) + + +@pytest.mark.parametrize("score_dims", [2, 3]) +def test_gemma4_prefill_collects_raw_qk_scores(score_dims): + torch.manual_seed(17) + head_dim = 256 + prefix = torch.tensor([0], dtype=torch.int32, device="cuda") + lengths = torch.tensor([4], dtype=torch.int32, device="cuda") + starts = torch.tensor([0], dtype=torch.int32, device="cuda") + slots = torch.arange(4, dtype=torch.int32, device="cuda").unsqueeze(0) + key = torch.randn(4, 2, head_dim, dtype=torch.bfloat16, device="cuda") + value = torch.randn_like(key) + query = torch.randn(4, 4, head_dim, dtype=torch.bfloat16, device="cuda") + output = torch.empty_like(query) + score = torch.zeros( + (1, 4, 4) if score_dims == 3 else (1, 4), + dtype=torch.float32, + device="cuda", + ) + gemma4_context_attention( + query, + key, + value, + output, + torch.tensor([0], dtype=torch.int32, device="cuda"), + starts, + lengths, + prefix, + 4, + slots, + sliding_window=None, + attn_score=score, + ) + expected = torch.zeros(1, 4, 4, dtype=torch.float32, device="cuda") + for head in range(4): + logits = query[:, head].float() @ key[:, head // 2].float().T + expected[0, head] = logits.tril().sum(0) + expected = ( + expected + if score_dims == 3 + else (expected / 4).max(1).values.clamp_min_(0) + ) + torch.testing.assert_close(score, expected, rtol=2e-2, atol=1.0) diff --git a/tests/test_gemma4_model.py b/tests/test_gemma4_model.py new file mode 100644 index 00000000..ef13d5e7 --- /dev/null +++ b/tests/test_gemma4_model.py @@ -0,0 +1,288 @@ +from __future__ import annotations + +from contextlib import ExitStack +from types import SimpleNamespace +from unittest.mock import patch + +import pytest +import torch +from transformers import Gemma4TextConfig +from transformers.models.gemma4.modeling_gemma4 import ( + Gemma4TextMLP as HFGemma4MLP, +) +from transformers.models.gemma4.modeling_gemma4 import ( + Gemma4TextModel as HFGemma4Model, +) +from transformers.models.gemma4.modeling_gemma4 import ( + Gemma4TextRotaryEmbedding as HFGemma4RotaryEmbedding, +) +from transformers.models.gemma4.modeling_gemma4 import ( + Gemma4TextRouter as HFGemma4Router, +) + +from sparsevllm.configs.sparse import normalize_sparse_methods +from sparsevllm.distributed import ParallelContext, ParallelGroup, ParallelMode +from sparsevllm.method_registry import MODEL_RUNTIME_COMPATIBILITY +from sparsevllm.models.gemma4 import ( + Gemma4Attention, + Gemma4MLP, + Gemma4Model, + Gemma4RotaryEmbedding, + Gemma4Router, +) +from sparsevllm.models.layout import RuntimeLayout +from sparsevllm.operators.activation import TorchGeluTanhAndMulProvider + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="requires CUDA") +def test_gemma4_router_kernels_match_torch(): + from sparsevllm.kernels.triton.gemma4_router import ( + gemma4_router_input, + gemma4_router_topk, + ) + + torch.manual_seed(19) + hidden = torch.randn(7, 2816, dtype=torch.bfloat16, device="cuda") + scale = torch.randn(2816, dtype=torch.bfloat16, device="cuda") + actual_input = gemma4_router_input(hidden, scale, 2816**-0.5, 1e-6) + variance = hidden.float().square().mean(-1, keepdim=True) + expected_input = (hidden.float() * torch.rsqrt(variance + 1e-6)).to(torch.bfloat16) + expected_input = (expected_input * scale * 2816**-0.5).to(torch.bfloat16) + torch.testing.assert_close(actual_input, expected_input, rtol=0, atol=0) + + logits = torch.randn(7, 128, dtype=torch.bfloat16, device="cuda") + expert_scale = torch.randn(128, dtype=torch.bfloat16, device="cuda") + actual_weights, actual_ids = gemma4_router_topk(logits, expert_scale, 8) + probabilities = logits.float().softmax(-1) + expected_weights, expected_ids = probabilities.topk(8, dim=-1) + expected_weights.div_(expected_weights.sum(-1, keepdim=True)).mul_( + expert_scale[expected_ids] + ) + assert torch.equal(actual_ids, expected_ids) + torch.testing.assert_close(actual_weights, expected_weights, rtol=1e-5, atol=1e-6) + + +def _parallel_context() -> ParallelContext: + group = ParallelGroup(process_group=None, ranks=(0,), rank=0, size=1) + return ParallelContext(world=group, tensor=group, expert=group, data=group) + + +def _patch_parallel_context(): + stack = ExitStack() + context = _parallel_context() + for target in ( + "sparsevllm.models.gemma4.get_parallel_context", + "sparsevllm.layers.linear.get_parallel_context", + "sparsevllm.layers.embed_head.get_parallel_context", + ): + stack.enter_context(patch(target, return_value=context)) + return stack + + +def _config(**overrides) -> Gemma4TextConfig: + values = { + "vocab_size": 32, + "hidden_size": 8, + "intermediate_size": 16, + "num_hidden_layers": 2, + "num_attention_heads": 2, + "num_key_value_heads": 2, + "head_dim": 4, + "global_head_dim": 8, + "num_global_key_value_heads": 1, + "max_position_embeddings": 32, + "layer_types": ["sliding_attention", "full_attention"], + "rope_parameters": { + "sliding_attention": {"rope_type": "default", "rope_theta": 10000.0}, + "full_attention": { + "rope_type": "proportional", + "rope_theta": 1000000.0, + "partial_rotary_factor": 0.25, + }, + }, + "sliding_window": 4, + "hidden_size_per_layer_input": 0, + "final_logit_softcapping": 30.0, + } + values.update(overrides) + return Gemma4TextConfig(**values) + + +def test_gemma4_rope_matches_transformers_for_both_layer_types(): + config = _config() + positions = torch.arange(9) + for layer_type, head_dim in (("sliding_attention", 4), ("full_attention", 8)): + actual = Gemma4RotaryEmbedding(config, layer_type, head_dim) + reference = HFGemma4RotaryEmbedding(config, layer_type=layer_type) + cos, sin = reference(torch.zeros(1), positions.unsqueeze(0), layer_type) + torch.testing.assert_close( + actual.cos_sin_cache[positions, 0, : head_dim // 2], + cos[0, :, : head_dim // 2], + ) + torch.testing.assert_close( + actual.cos_sin_cache[positions, 0, head_dim // 2 :], + sin[0, :, : head_dim // 2], + ) + + +def test_gemma4_dense_mlp_matches_transformers(): + config = _config() + with _patch_parallel_context(): + actual = Gemma4MLP(config, 0, TorchGeluTanhAndMulProvider()) + reference = HFGemma4MLP(config, 0) + torch.manual_seed(3) + for parameter in reference.parameters(): + parameter.data.normal_(0, 0.1) + actual.gate_up_proj.weight_loader( + actual.gate_up_proj.weight, reference.gate_proj.weight, 0 + ) + actual.gate_up_proj.weight_loader( + actual.gate_up_proj.weight, reference.up_proj.weight, 1 + ) + actual.down_proj.weight_loader(actual.down_proj.weight, reference.down_proj.weight) + hidden_states = torch.randn(7, config.hidden_size) + torch.testing.assert_close( + actual(hidden_states), reference(hidden_states), atol=1e-6, rtol=1e-5 + ) + + +def test_gemma4_router_matches_transformers(): + config = _config( + enable_moe_block=True, num_experts=4, top_k_experts=2, moe_intermediate_size=4 + ) + with _patch_parallel_context(): + actual = Gemma4Router(config) + reference = HFGemma4Router(config) + torch.manual_seed(5) + reference.proj.weight.data.normal_(0, 0.2) + reference.scale.data.normal_(1, 0.1) + reference.per_expert_scale.data.normal_(1, 0.1) + actual.load_state_dict(reference.state_dict()) + hidden_states = torch.randn(11, config.hidden_size) + _, expected_weights, expected_ids = reference(hidden_states) + weights, ids = actual(hidden_states) + torch.testing.assert_close(weights, expected_weights) + assert torch.equal(ids, expected_ids) + + +def test_gemma4_k_eq_v_loader_duplicates_normalized_projection_slot(): + config = _config(attention_k_eq_v=True) + with _patch_parallel_context(): + attention = Gemma4Attention( + config, + 1, + Gemma4RotaryEmbedding(config, "full_attention", config.global_head_dim), + ) + loaded_key = torch.randn( + config.num_global_key_value_heads * config.global_head_dim, config.hidden_size + ) + attention.qkv_proj.weight_loader(attention.qkv_proj.weight, loaded_key, "k") + q_end = attention.q_size + key = attention.qkv_proj.weight[q_end : q_end + attention.kv_size] + value = attention.qkv_proj.weight[q_end + attention.kv_size :] + torch.testing.assert_close(key, loaded_key) + torch.testing.assert_close(value, loaded_key) + + +def test_gemma4_ple_matches_transformers(): + config = _config( + hidden_size_per_layer_input=2, + vocab_size_per_layer_input=32, + ) + with _patch_parallel_context(): + actual = Gemma4Model(config, TorchGeluTanhAndMulProvider()) + reference = HFGemma4Model(config) + torch.manual_seed(11) + for parameter in reference.parameters(): + parameter.data.normal_(0, 0.1) + actual.embed_tokens.weight.data.copy_(reference.embed_tokens.weight) + actual.embed_tokens_per_layer.weight.data.copy_( + reference.embed_tokens_per_layer.weight + ) + actual.per_layer_model_projection.weight.data.copy_( + reference.per_layer_model_projection.weight + ) + actual.per_layer_projection_norm.load_state_dict( + reference.per_layer_projection_norm.state_dict() + ) + actual.per_layer_projection_norm._ops = SimpleNamespace( + rmsnorm=lambda x, weight, eps: ( + x.float() + * torch.rsqrt(x.float().square().mean(-1, keepdim=True) + eps) + * weight.float() + ).to(x.dtype) + ) + input_ids = torch.tensor([1, 7, 3, 9]) + actual_hidden = actual.embed_tokens(input_ids) * actual.embedding_scale + reference_hidden = reference.embed_tokens(input_ids) + expected = reference.project_per_layer_inputs( + reference_hidden, + reference.get_per_layer_inputs(input_ids, reference_hidden), + ) + torch.testing.assert_close( + actual.get_per_layer_inputs(input_ids, actual_hidden), + expected, + ) + + +def test_gemma4_shared_kv_layout_aliases_last_source_by_type(): + config = _config( + num_hidden_layers=4, + num_kv_shared_layers=2, + layer_types=[ + "sliding_attention", + "full_attention", + "sliding_attention", + "full_attention", + ], + ) + layout = RuntimeLayout.from_config(config) + assert layout.num_kv_layers == 2 + assert layout.kv_idx_to_layer_idx == (0, 1) + assert layout.layer_idx_to_kv_idx == (0, 1, 0, 1) + assert layout.kv_num_heads == (2, 2) + assert layout.kv_head_dims == (4, 8) + + +def test_gemma4_shared_kv_attention_only_allocates_query_projection(): + config = _config( + num_hidden_layers=4, + num_kv_shared_layers=2, + layer_types=[ + "sliding_attention", + "full_attention", + "sliding_attention", + "full_attention", + ], + ) + with _patch_parallel_context(): + attention = Gemma4Attention( + config, + 2, + Gemma4RotaryEmbedding(config, "sliding_attention", config.head_dim), + ) + assert attention.is_kv_shared_layer + assert tuple(attention.qkv_proj.weight.shape) == ( + config.num_attention_heads * config.head_dim, + config.hidden_size, + ) + assert not hasattr(attention, "k_norm") + assert not hasattr(attention, "v_norm") + + +def test_gemma4_sparse_registry_keeps_dedicated_validated_methods(): + compatibility = MODEL_RUNTIME_COMPATIBILITY[("gemma4", ParallelMode.STANDARD)] + assert compatibility.sparse_methods == {"", "streamingllm", "omnikv"} + assert compatibility.decode_cuda_graph_methods == compatibility.sparse_methods + + +def test_gemma4_shared_kv_rejects_per_layer_streaming_eviction(): + config = SimpleNamespace( + hf_config=SimpleNamespace( + model_type="gemma4_text", + num_kv_shared_layers=18, + ), + vllm_sparse_method="streamingllm", + ) + with pytest.raises(NotImplementedError, match="KV-sharing"): + normalize_sparse_methods(config) diff --git a/tests/test_gemma4_rmsnorm.py b/tests/test_gemma4_rmsnorm.py new file mode 100644 index 00000000..19b8d734 --- /dev/null +++ b/tests/test_gemma4_rmsnorm.py @@ -0,0 +1,34 @@ +import pytest +import torch + +from sparsevllm.layers.gemma4_rmsnorm import Gemma4RMSNorm + + +def _reference(x: torch.Tensor, weight: torch.Tensor | None, eps: float) -> torch.Tensor: + output = x.float() + output *= torch.pow(output.square().mean(-1, keepdim=True) + eps, -0.5) + if weight is not None: + output *= weight.float() + return output.to(x.dtype) + + +@pytest.mark.parametrize("with_scale", [False, True]) +def test_gemma4_rmsnorm_matches_torch(with_scale): + torch.manual_seed(7) + layer = Gemma4RMSNorm(32, with_scale=with_scale) + x = torch.randn(5, 32) + torch.testing.assert_close(layer(x), _reference(x, layer.weight, layer.eps)) + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="requires CUDA") +def test_gemma4_rmsnorm_cuda_graph_matches_torch(): + torch.manual_seed(11) + layer = Gemma4RMSNorm(2816).cuda().to(torch.bfloat16) + x = torch.randn(4, 2816, device="cuda", dtype=torch.bfloat16) + for _ in range(2): + layer(x) + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph): + output = layer(x) + graph.replay() + torch.testing.assert_close(output, _reference(x, layer.weight, layer.eps), rtol=0, atol=0) diff --git a/tests/test_triton_moe.py b/tests/test_triton_moe.py index 840b1c40..754ac75e 100644 --- a/tests/test_triton_moe.py +++ b/tests/test_triton_moe.py @@ -4,8 +4,8 @@ import torch import torch.nn.functional as F -from sparsevllm.operators.gated_shared_add import gated_shared_add from sparsevllm.kernels.triton.gate_up_swiglu import h20_gate_up_swiglu +from sparsevllm.kernels.triton.gemma4_moe import fused_gemma4_moe from sparsevllm.kernels.triton.moe import ( _prepare_expert_assignment, append_shared_expert_route, @@ -15,6 +15,7 @@ ) from sparsevllm.kernels.triton.moe_topk import topk_softmax from sparsevllm.kernels.triton.silu_and_mul import _resolve_silu_launch_config +from sparsevllm.operators.gated_shared_add import gated_shared_add def test_silu_launch_config_uses_decode_tile_only_for_small_rows(): @@ -96,6 +97,71 @@ def _oracle_local_moe( return output +def _oracle_gemma4_moe( + hidden_states, + w13_weight, + w2_weight, + topk_ids, + topk_weights, + local_expert_start, +): + output = torch.zeros_like(hidden_states) + for local_id in range(w13_weight.shape[0]): + global_id = local_expert_start + local_id + token_ids, routes = torch.where(topk_ids == global_id) + if token_ids.numel() == 0: + continue + gate, up = F.linear(hidden_states[token_ids], w13_weight[local_id]).chunk(2, -1) + routed = F.linear(F.gelu(gate, approximate="tanh") * up, w2_weight[local_id]) + output.index_add_(0, token_ids, routed * topk_weights[token_ids, routes, None]) + return output + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA is required") +@pytest.mark.parametrize("num_tokens", [1, 19]) +def test_gemma4_moe_matches_torch(num_tokens): + torch.manual_seed(71 + num_tokens) + hidden_size, intermediate_size, num_experts, top_k = 64, 32, 7, 3 + hidden_states = torch.randn( + num_tokens, hidden_size, dtype=torch.bfloat16, device="cuda" + ) + w13_weight = torch.randn( + num_experts, + 2 * intermediate_size, + hidden_size, + dtype=torch.bfloat16, + device="cuda", + ) * 0.1 + w2_weight = torch.randn( + num_experts, + hidden_size, + intermediate_size, + dtype=torch.bfloat16, + device="cuda", + ) * 0.1 + ids = torch.randint( + num_experts, + (num_tokens, top_k), + dtype=torch.int64, + device="cuda", + ) + weights = torch.rand(num_tokens, top_k, dtype=torch.bfloat16, device="cuda") + weights /= weights.sum(-1, keepdim=True) + expected = _oracle_gemma4_moe( + hidden_states, w13_weight, w2_weight, ids, weights, 0 + ) + actual = fused_gemma4_moe( + hidden_states, + w13_weight, + w2_weight, + ids, + weights, + num_experts=num_experts, + local_expert_start=0, + ) + torch.testing.assert_close(actual, expected, atol=0.04, rtol=0.04) + + @unittest.skipUnless(torch.cuda.is_available(), "CUDA is required for Triton MoE tests.") def test_append_shared_expert_route_matches_cat_and_replays_graph(): ids = torch.tensor( From 12fcdfecae0f931f7a1d21589ae3a4574d42e785 Mon Sep 17 00:00:00 2001 From: QuanshengGu Date: Thu, 13 Aug 2026 05:41:37 +0800 Subject: [PATCH 02/13] refactor: isolate gemma4 operators --- src/sparsevllm/layers/activation.py | 13 - src/sparsevllm/layers/gemma4_rmsnorm.py | 24 +- src/sparsevllm/models/gemma4.py | 313 +++++++++--------------- src/sparsevllm/operators/activation.py | 107 -------- src/sparsevllm/operators/gemma4.py | 224 +++++++++++++++++ src/sparsevllm/operators/gemma4_moe.py | 100 +++++++- tests/test_activation.py | 29 +-- tests/test_gemma4_model.py | 38 ++- tests/test_gemma4_rmsnorm.py | 37 ++- 9 files changed, 520 insertions(+), 365 deletions(-) create mode 100644 src/sparsevllm/operators/gemma4.py diff --git a/src/sparsevllm/layers/activation.py b/src/sparsevllm/layers/activation.py index 456e6c6c..393c609f 100755 --- a/src/sparsevllm/layers/activation.py +++ b/src/sparsevllm/layers/activation.py @@ -2,9 +2,7 @@ from torch import nn from sparsevllm.operators.activation import ( - GeluTanhAndMulProvider, SiluAndMulProvider, - TorchGeluTanhAndMulProvider, TorchSiluAndMulProvider, ) @@ -17,14 +15,3 @@ def __init__(self, provider: SiluAndMulProvider | None = None): def forward(self, x: torch.Tensor) -> torch.Tensor: return self.provider(x) - - -class GeluTanhAndMul(nn.Module): - def __init__(self, provider: GeluTanhAndMulProvider | None = None): - super().__init__() - self.provider = ( - provider if provider is not None else TorchGeluTanhAndMulProvider() - ) - - def forward(self, x: torch.Tensor) -> torch.Tensor: - return self.provider(x) diff --git a/src/sparsevllm/layers/gemma4_rmsnorm.py b/src/sparsevllm/layers/gemma4_rmsnorm.py index 3e3cc012..99d3297e 100644 --- a/src/sparsevllm/layers/gemma4_rmsnorm.py +++ b/src/sparsevllm/layers/gemma4_rmsnorm.py @@ -3,30 +3,28 @@ import torch from torch import nn - -def _torch_rmsnorm(x: torch.Tensor, weight: torch.Tensor | None, eps: float) -> torch.Tensor: - output = x.float() - output *= torch.pow(output.square().mean(-1, keepdim=True) + eps, -0.5) - if weight is not None: - output *= weight.float() - return output.to(x.dtype) +from sparsevllm.operators.gemma4 import Gemma4OperatorProvider class Gemma4RMSNorm(nn.Module): - def __init__(self, hidden_size: int, eps: float = 1e-6, *, with_scale: bool = True) -> None: + def __init__( + self, + hidden_size: int, + eps: float = 1e-6, + *, + with_scale: bool = True, + provider: Gemma4OperatorProvider, + ) -> None: super().__init__() self.eps = float(eps) + self._ops = provider if with_scale: self.weight = nn.Parameter(torch.ones(hidden_size)) else: self.register_parameter("weight", None) def forward(self, x: torch.Tensor) -> torch.Tensor: - if not x.is_cuda: - return _torch_rmsnorm(x, self.weight, self.eps) - from sparsevllm.kernels.triton.gemma4_rmsnorm import gemma4_rmsnorm - - return gemma4_rmsnorm(x, self.weight, self.eps) + return self._ops.rmsnorm(x, self.weight, self.eps) __all__ = ["Gemma4RMSNorm"] diff --git a/src/sparsevllm/models/gemma4.py b/src/sparsevllm/models/gemma4.py index 5f2da41a..53176db0 100644 --- a/src/sparsevllm/models/gemma4.py +++ b/src/sparsevllm/models/gemma4.py @@ -9,7 +9,6 @@ from transformers import Gemma4TextConfig from sparsevllm.distributed import get_parallel_context -from sparsevllm.layers.activation import GeluTanhAndMul from sparsevllm.layers.attention import Attention from sparsevllm.layers.embed_head import ParallelLMHead, VocabParallelEmbedding from sparsevllm.layers.gemma4_rmsnorm import Gemma4RMSNorm @@ -21,14 +20,13 @@ ReplicatedLinear, RowParallelLinear, ) -from sparsevllm.layers.packed_moe import PackedMoeExperts from sparsevllm.layers.rotary_embedding import apply_rotary_emb -from sparsevllm.operators.activation import ( - GeluTanhAndMulProvider, - resolve_gelu_tanh_and_mul_provider, +from sparsevllm.operators.gemma4 import ( + Gemma4OperatorProvider, + Gemma4OpSpec, + resolve_gemma4_provider, ) -from sparsevllm.operators.gemma4_attention import Gemma4AttentionBackend -from sparsevllm.operators.gemma4_moe import resolve_gemma4_moe_provider +from sparsevllm.operators.gemma4_moe import Gemma4PackedExperts from sparsevllm.operators.moe import model_activation_dtype from sparsevllm.platforms import device_runtime from sparsevllm.utils.context import get_context @@ -142,6 +140,7 @@ def __init__( config: Gemma4TextConfig, layer_idx: int, rotary_emb: Gemma4RotaryEmbedding, + operator_provider: Gemma4OperatorProvider, ) -> None: super().__init__() parallel_context = get_parallel_context() @@ -193,15 +192,23 @@ def __init__( bias=config.attention_bias, quantization=getattr(config, "quantization_config", None), ) - self.q_norm = Gemma4RMSNorm(self.head_dim, eps=config.rms_norm_eps) + self.q_norm = Gemma4RMSNorm( + self.head_dim, eps=config.rms_norm_eps, provider=operator_provider + ) if not self.is_kv_shared_layer: - self.k_norm = Gemma4RMSNorm(self.head_dim, eps=config.rms_norm_eps) + self.k_norm = Gemma4RMSNorm( + self.head_dim, eps=config.rms_norm_eps, provider=operator_provider + ) self.v_norm = Gemma4RMSNorm( - self.head_dim, eps=config.rms_norm_eps, with_scale=False + self.head_dim, + eps=config.rms_norm_eps, + with_scale=False, + provider=operator_provider, ) self.rotary_emb = rotary_emb + self._ops = operator_provider self.attn = Attention(self.num_heads, self.head_dim, 1.0, self.num_kv_heads) - self.attn.attention_backend = Gemma4AttentionBackend( + self.attn.attention_backend = operator_provider.attention_backend( sliding_window=self.sliding_window ) @@ -210,23 +217,16 @@ def forward( ) -> torch.Tensor: if self.is_kv_shared_layer: q = self.qkv_proj(hidden_states).view(-1, self.num_heads, self.head_dim) - if q.is_cuda: - from sparsevllm.kernels.triton.gemma4_qkv_norm_rope import ( - gemma4_qkv_norm_rope, - ) - - gemma4_qkv_norm_rope( - q, - None, - None, - self.q_norm.weight, - None, - self.rotary_emb.cos_sin_cache, - positions, - self.q_norm.eps, - ) - else: - q = self.rotary_emb.forward_query(positions, self.q_norm(q)) + q, _, _ = self._ops.qkv_norm_rope( + q, + None, + None, + self.q_norm.weight, + None, + self.rotary_emb.cos_sin_cache, + positions, + self.q_norm.eps, + ) empty = q.new_empty((0, self.num_kv_heads, self.head_dim)) return self.o_proj(self.attn(q, empty, empty).flatten(1)) q, k, v = self.qkv_proj(hidden_states).split( @@ -235,24 +235,16 @@ def forward( q = q.view(-1, self.num_heads, self.head_dim) k = k.view(-1, self.num_kv_heads, self.head_dim) v = v.view(-1, self.num_kv_heads, self.head_dim) - if q.is_cuda: - from sparsevllm.kernels.triton.gemma4_qkv_norm_rope import ( - gemma4_qkv_norm_rope, - ) - - gemma4_qkv_norm_rope( - q, - k, - v, - self.q_norm.weight, - self.k_norm.weight, - self.rotary_emb.cos_sin_cache, - positions, - self.q_norm.eps, - ) - else: - q, k = self.rotary_emb(positions, self.q_norm(q), self.k_norm(k)) - v = self.v_norm(v) + q, k, v = self._ops.qkv_norm_rope( + q, + k, + v, + self.q_norm.weight, + self.k_norm.weight, + self.rotary_emb.cos_sin_cache, + positions, + self.q_norm.eps, + ) context = get_context() context.cache_manager.save_rope_kv_if_needed(context.now_layer_idx, k, v) output = self.attn(q, k, v) @@ -264,7 +256,7 @@ def __init__( self, config: Gemma4TextConfig, layer_idx: int, - activation_provider: GeluTanhAndMulProvider | None, + operator_provider: Gemma4OperatorProvider, ) -> None: super().__init__() shared_start = int(config.num_hidden_layers) - int(config.num_kv_shared_layers) @@ -283,103 +275,40 @@ def __init__( config.hidden_size, quantization=getattr(config, "quantization_config", None), ) - self.activation = GeluTanhAndMul(activation_provider) + self._ops = operator_provider def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - return self.down_proj(self.activation(self.gate_up_proj(hidden_states))) + return self.down_proj( + self._ops.gelu_tanh_and_mul(self.gate_up_proj(hidden_states)) + ) class Gemma4Router(nn.Module): - def __init__(self, config: Gemma4TextConfig) -> None: + def __init__( + self, + config: Gemma4TextConfig, + operator_provider: Gemma4OperatorProvider, + ) -> None: super().__init__() self.top_k = int(config.top_k_experts) self.root_size = float(config.hidden_size) ** -0.5 self.norm = Gemma4RMSNorm( - config.hidden_size, eps=config.rms_norm_eps, with_scale=False + config.hidden_size, + eps=config.rms_norm_eps, + with_scale=False, + provider=operator_provider, ) + self._ops = operator_provider self.scale = nn.Parameter(torch.ones(config.hidden_size)) self.proj = ReplicatedLinear(config.hidden_size, config.num_experts) self.per_expert_scale = nn.Parameter(torch.ones(config.num_experts)) def forward(self, hidden_states: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: - if hidden_states.is_cuda: - from sparsevllm.kernels.triton.gemma4_router import ( - gemma4_router_input, - gemma4_router_topk, - ) - - router_input = gemma4_router_input( - hidden_states, self.scale, self.root_size, self.norm.eps - ) - return gemma4_router_topk( - self.proj(router_input), self.per_expert_scale, self.top_k - ) - logits = self.proj(self.norm(hidden_states) * self.scale * self.root_size) - probabilities = F.softmax(logits, dim=-1, dtype=torch.float32) - weights, ids = probabilities.topk(self.top_k, dim=-1) - weights.div_(weights.sum(-1, keepdim=True)).mul_(self.per_expert_scale[ids]) - return weights, ids - - -class Gemma4PackedExperts(PackedMoeExperts): - def __init__(self, config: Gemma4TextConfig) -> None: - super().__init__( - num_experts=config.num_experts, - hidden_size=config.hidden_size, - intermediate_size=config.moe_intermediate_size, - top_k=config.top_k_experts, - activation_dtype=model_activation_dtype(config), - fp8_enabled=False, - cuda_graph=bool(getattr(config, "decode_cuda_graph", False)), - activation="gelu_tanh", - model_label="Gemma4MoE", - provider_resolver=resolve_gemma4_moe_provider, - parallel_context=get_parallel_context(), - ) - - def rank_local_weight_slice( - self, - source_shape: tuple[int, ...], - *, - loaded_shard_id: str, - is_scale: bool = False, - ) -> tuple[slice, ...] | None: - if is_scale: - raise ValueError("Gemma 4 BF16 experts do not use weight scales.") - if len(source_shape) != 3 or int(source_shape[0]) != self.num_experts: - raise ValueError(f"Invalid Gemma 4 packed expert shape {source_shape}.") - if self.ep_size == 1: - return None - return ( - slice(self.local_expert_start, self.local_expert_end), - slice(None), - slice(None), - ) - - def load_packed_weight(self, projection: str, loaded_weight: torch.Tensor) -> None: - if loaded_weight.shape[0] == self.num_experts: - loaded_weight = loaded_weight[ - self.local_expert_start : self.local_expert_end - ] - if projection == "gate_up_proj": - gate, up = loaded_weight.chunk(2, 1) - gate = gate.chunk(self.tp_size, 1)[self.tp_rank] - up = up.chunk(self.tp_size, 1)[self.tp_rank] - self.w13_weight.data.copy_(torch.cat((gate, up), 1)) - projections = ("gate_proj", "up_proj") - elif projection == "down_proj": - self.w2_weight.data.copy_( - loaded_weight.chunk(self.tp_size, 2)[self.tp_rank] - ) - projections = ("down_proj",) - else: - raise ValueError( - f"Unsupported Gemma 4 packed expert projection {projection!r}." - ) - self._loaded_expert_shards.update( - (expert_id, name) - for expert_id in range(self.local_expert_start, self.local_expert_end) - for name in projections + router_input = self._ops.router_input( + hidden_states, self.scale, self.root_size, self.norm.eps + ) + return self._ops.router_topk( + self.proj(router_input), self.per_expert_scale, self.top_k ) @@ -388,7 +317,7 @@ def __init__( self, config: Gemma4TextConfig, layer_idx: int, - activation_provider: GeluTanhAndMulProvider | None, + operator_provider: Gemma4OperatorProvider, rotary_embeddings: nn.ModuleDict, ) -> None: super().__init__() @@ -397,19 +326,21 @@ def __init__( config, layer_idx, rotary_embeddings[layer_type], + operator_provider, ) - self.mlp = Gemma4MLP(config, layer_idx, activation_provider) + self._ops = operator_provider + self.mlp = Gemma4MLP(config, layer_idx, operator_provider) self.input_layernorm = Gemma4RMSNorm( - config.hidden_size, eps=config.rms_norm_eps + config.hidden_size, eps=config.rms_norm_eps, provider=operator_provider ) self.post_attention_layernorm = Gemma4RMSNorm( - config.hidden_size, eps=config.rms_norm_eps + config.hidden_size, eps=config.rms_norm_eps, provider=operator_provider ) self.pre_feedforward_layernorm = Gemma4RMSNorm( - config.hidden_size, eps=config.rms_norm_eps + config.hidden_size, eps=config.rms_norm_eps, provider=operator_provider ) self.post_feedforward_layernorm = Gemma4RMSNorm( - config.hidden_size, eps=config.rms_norm_eps + config.hidden_size, eps=config.rms_norm_eps, provider=operator_provider ) self.hidden_size_per_layer_input = int(config.hidden_size_per_layer_input) if self.hidden_size_per_layer_input: @@ -424,20 +355,27 @@ def __init__( self.post_per_layer_input_norm = Gemma4RMSNorm( config.hidden_size, eps=config.rms_norm_eps, + provider=operator_provider, ) self.enable_moe_block = bool(config.enable_moe_block) if self.enable_moe_block: self.parallel_context = get_parallel_context() - self.router = Gemma4Router(config) + self.router = Gemma4Router(config, operator_provider) self.experts = Gemma4PackedExperts(config) self.post_feedforward_layernorm_1 = Gemma4RMSNorm( - config.hidden_size, eps=config.rms_norm_eps + config.hidden_size, + eps=config.rms_norm_eps, + provider=operator_provider, ) self.pre_feedforward_layernorm_2 = Gemma4RMSNorm( - config.hidden_size, eps=config.rms_norm_eps + config.hidden_size, + eps=config.rms_norm_eps, + provider=operator_provider, ) self.post_feedforward_layernorm_2 = Gemma4RMSNorm( - config.hidden_size, eps=config.rms_norm_eps + config.hidden_size, + eps=config.rms_norm_eps, + provider=operator_provider, ) self.layer_scalar = nn.Parameter(torch.ones(1), requires_grad=False) @@ -449,19 +387,12 @@ def forward( ) -> torch.Tensor: residual = hidden_states hidden_states = self.self_attn(positions, self.input_layernorm(hidden_states)) - if hidden_states.is_cuda: - from sparsevllm.kernels.triton.gemma4_fused_ops import ( - gemma4_rmsnorm_residual, - ) - - hidden_states = gemma4_rmsnorm_residual( - hidden_states, - self.post_attention_layernorm.weight, - residual, - self.post_attention_layernorm.eps, - ) - else: - hidden_states = self.post_attention_layernorm(hidden_states) + residual + hidden_states = self._ops.rmsnorm_residual( + hidden_states, + self.post_attention_layernorm.weight, + residual, + self.post_attention_layernorm.eps, + ) residual = hidden_states dense_input = self.pre_feedforward_layernorm(hidden_states) hidden_states = self.mlp(dense_input) @@ -473,52 +404,35 @@ def forward( hidden_states = self.post_feedforward_layernorm_1( hidden_states ) + self.post_feedforward_layernorm_2(expert_output) - if hidden_states.is_cuda: - hidden_states = gemma4_rmsnorm_residual( - hidden_states, - self.post_feedforward_layernorm.weight, - residual, - self.post_feedforward_layernorm.eps, - None if self.hidden_size_per_layer_input else self.layer_scalar, - ) - else: - hidden_states = self.post_feedforward_layernorm(hidden_states) + residual + hidden_states = self._ops.rmsnorm_residual( + hidden_states, + self.post_feedforward_layernorm.weight, + residual, + self.post_feedforward_layernorm.eps, + None if self.hidden_size_per_layer_input else self.layer_scalar, + ) if self.hidden_size_per_layer_input: if per_layer_input is None: raise RuntimeError("Gemma 4 PLE layer requires per_layer_input.") residual = hidden_states hidden_states = self.per_layer_input_gate(hidden_states) - if hidden_states.is_cuda: - from sparsevllm.kernels.triton.gemma4_fused_ops import gemma4_gelu_mul - - hidden_states = gemma4_gelu_mul(hidden_states, per_layer_input) - else: - hidden_states = ( - F.gelu(hidden_states, approximate="tanh") * per_layer_input - ) + hidden_states = self._ops.gelu_mul(hidden_states, per_layer_input) hidden_states = self.per_layer_projection(hidden_states) - if hidden_states.is_cuda: - hidden_states = gemma4_rmsnorm_residual( - hidden_states, - self.post_per_layer_input_norm.weight, - residual, - self.post_per_layer_input_norm.eps, - self.layer_scalar, - ) - else: - hidden_states = self.post_per_layer_input_norm(hidden_states) + residual - return ( - hidden_states - if hidden_states.is_cuda - else hidden_states * self.layer_scalar - ) + hidden_states = self._ops.rmsnorm_residual( + hidden_states, + self.post_per_layer_input_norm.weight, + residual, + self.post_per_layer_input_norm.eps, + self.layer_scalar, + ) + return hidden_states class Gemma4Model(nn.Module): def __init__( self, config: Gemma4TextConfig, - activation_provider: GeluTanhAndMulProvider | None, + operator_provider: Gemma4OperatorProvider, ) -> None: super().__init__() self.config = config @@ -541,6 +455,7 @@ def __init__( self.per_layer_projection_norm = Gemma4RMSNorm( self.hidden_size_per_layer_input, eps=config.rms_norm_eps, + provider=operator_provider, ) self.per_layer_model_projection_scale = float(config.hidden_size) ** -0.5 self.per_layer_input_scale = 2.0**-0.5 @@ -560,12 +475,14 @@ def __init__( Gemma4DecoderLayer( config, layer_idx, - activation_provider, + operator_provider, self.rotary_embeddings, ) for layer_idx in range(config.num_hidden_layers) ) - self.norm = Gemma4RMSNorm(config.hidden_size, eps=config.rms_norm_eps) + self.norm = Gemma4RMSNorm( + config.hidden_size, eps=config.rms_norm_eps, provider=operator_provider + ) self.embedding_scale = float(config.hidden_size) ** 0.5 self.sparse_controller = None @@ -624,11 +541,11 @@ class Gemma4ForCausalLM(nn.Module): def __init__( self, config: Gemma4TextConfig, - activation_provider: GeluTanhAndMulProvider | None = None, + operator_provider: Gemma4OperatorProvider, ) -> None: super().__init__() self.config = config - self.model = Gemma4Model(config, activation_provider) + self.model = Gemma4Model(config, operator_provider) self.lm_head = ParallelLMHead(config.vocab_size, config.hidden_size) if config.tie_word_embeddings: self.lm_head.weight.data = self.model.embed_tokens.weight.data @@ -636,9 +553,21 @@ def __init__( @classmethod def build_runtime_kwargs(cls, config, *, device, **_): + head_dims = tuple( + sorted( + { + int(config.head_dim), + int(getattr(config, "global_head_dim", config.head_dim)), + } + ) + ) return { - "activation_provider": resolve_gelu_tanh_and_mul_provider( - activation_dtype=model_activation_dtype(config), + "operator_provider": resolve_gemma4_provider( + Gemma4OpSpec( + activation_dtype=model_activation_dtype(config), + head_dims=head_dims, + cuda_graph=bool(getattr(config, "decode_cuda_graph", False)), + ), device_index=device.index, ) } diff --git a/src/sparsevllm/operators/activation.py b/src/sparsevllm/operators/activation.py index 428b1543..072c9c00 100644 --- a/src/sparsevllm/operators/activation.py +++ b/src/sparsevllm/operators/activation.py @@ -21,17 +21,6 @@ def __post_init__(self) -> None: raise ValueError("SiluAndMul input_ndim must be positive.") -@dataclass(frozen=True) -class GeluTanhAndMulSpec: - activation_dtype: torch.dtype - input_ndim: int = 2 - contiguous: bool = True - - def __post_init__(self) -> None: - if int(self.input_ndim) <= 0: - raise ValueError("GeluTanhAndMul input_ndim must be positive.") - - def _validate_input(x: torch.Tensor) -> None: if int(x.shape[-1]) % 2: raise ValueError( @@ -64,20 +53,9 @@ def __call__(self, x: torch.Tensor) -> torch.Tensor: raise NotImplementedError -class GeluTanhAndMulProvider: - name = "" - priority = 0 - - def __call__(self, x: torch.Tensor) -> torch.Tensor: - raise NotImplementedError - - SILU_AND_MUL_REGISTRY: OpRegistry[SiluAndMulSpec, SiluAndMulProvider] = OpRegistry( "SiLU-and-multiply" ) -GELU_TANH_AND_MUL_REGISTRY: OpRegistry[ - GeluTanhAndMulSpec, GeluTanhAndMulProvider -] = OpRegistry("GELU-tanh-and-multiply") @SILU_AND_MUL_REGISTRY.register @@ -136,62 +114,6 @@ def __call__(self, x: torch.Tensor) -> torch.Tensor: return gate -@GELU_TANH_AND_MUL_REGISTRY.register -class TritonGeluTanhAndMulProvider(GeluTanhAndMulProvider): - name = "triton" - priority = 10 - - def __init__(self, *, op_spec: GeluTanhAndMulSpec) -> None: - self.spec = op_spec - - @classmethod - def supports(cls, spec: GeluTanhAndMulSpec, caps: DeviceCaps) -> SupportResult: - if caps.platform != PlatformEnum.CUDA: - return SupportResult.no(f"requires CUDA, got {caps.platform.name}") - if not caps.supports_triton: - return SupportResult.no("platform does not support Triton") - if spec.activation_dtype not in (torch.float16, torch.bfloat16): - return SupportResult.no( - "requires FP16 or BF16 activations, " - f"got {spec.activation_dtype}" - ) - if int(spec.input_ndim) != 2 or not spec.contiguous: - return SupportResult.no("requires contiguous rank-2 inputs") - return SupportResult.yes() - - def __call__(self, x: torch.Tensor) -> torch.Tensor: - _validate_bound_input(x, self.spec) - if not x.is_cuda: - raise ValueError("Triton GeluTanhAndMul provider requires a CUDA input.") - from sparsevllm.kernels.triton.gemma4_gelu_and_mul import ( - gelu_tanh_and_mul_fwd, - ) - - return gelu_tanh_and_mul_fwd(x) - - -@GELU_TANH_AND_MUL_REGISTRY.register -class TorchGeluTanhAndMulProvider(GeluTanhAndMulProvider): - name = "torch" - priority = 0 - - def __init__(self, *, op_spec: GeluTanhAndMulSpec | None = None) -> None: - self.spec = op_spec - - @classmethod - def supports(cls, spec: GeluTanhAndMulSpec, caps: DeviceCaps) -> SupportResult: - del spec, caps - return SupportResult.yes() - - def __call__(self, x: torch.Tensor) -> torch.Tensor: - if self.spec is None: - _validate_input(x) - else: - _validate_bound_input(x, self.spec) - gate, up = x.chunk(2, -1) - return F.gelu(gate, approximate="tanh") * up - - def resolve_silu_and_mul_provider( *, activation_dtype: torch.dtype, @@ -215,40 +137,11 @@ def resolve_silu_and_mul_provider( ).provider -def resolve_gelu_tanh_and_mul_provider( - *, - activation_dtype: torch.dtype, - input_ndim: int = 2, - contiguous: bool = True, - device_index: int | None = None, -) -> GeluTanhAndMulProvider: - platform = platforms.current_platform - if device_index is None: - device_index = torch.cuda.current_device() if platform.is_cuda_alike() else 0 - caps = platform.get_device_caps(int(device_index)) - spec = GeluTanhAndMulSpec( - activation_dtype=activation_dtype, - input_ndim=int(input_ndim), - contiguous=bool(contiguous), - ) - return OpResolver(GELU_TANH_AND_MUL_REGISTRY).resolve( - spec, - caps, - op_spec=spec, - ).provider - - __all__ = [ - "GELU_TANH_AND_MUL_REGISTRY", "SILU_AND_MUL_REGISTRY", - "GeluTanhAndMulProvider", - "GeluTanhAndMulSpec", "SiluAndMulProvider", "SiluAndMulSpec", - "TorchGeluTanhAndMulProvider", "TorchSiluAndMulProvider", - "TritonGeluTanhAndMulProvider", "TritonSiluAndMulProvider", - "resolve_gelu_tanh_and_mul_provider", "resolve_silu_and_mul_provider", ] diff --git a/src/sparsevllm/operators/gemma4.py b/src/sparsevllm/operators/gemma4.py new file mode 100644 index 00000000..cde36cf7 --- /dev/null +++ b/src/sparsevllm/operators/gemma4.py @@ -0,0 +1,224 @@ +from __future__ import annotations + +from dataclasses import dataclass + +import torch +import torch.nn.functional as F + +import sparsevllm.platforms as platforms +from sparsevllm.layers.rotary_embedding import apply_rotary_emb +from sparsevllm.operators.registry import OpRegistry, OpResolver, SupportResult +from sparsevllm.platforms.interface import DeviceCaps, PlatformEnum + + +@dataclass(frozen=True) +class Gemma4OpSpec: + activation_dtype: torch.dtype + head_dims: tuple[int, ...] + cuda_graph: bool + + def __post_init__(self) -> None: + if not self.head_dims or any(int(value) <= 0 for value in self.head_dims): + raise ValueError("Gemma 4 head dimensions must be positive.") + + +class Gemma4OperatorProvider: + name = "" + priority = 0 + + def attention_backend(self, *, sliding_window: int | None): + raise NotImplementedError + + def rmsnorm( + self, x: torch.Tensor, weight: torch.Tensor | None, eps: float + ) -> torch.Tensor: + raise NotImplementedError + + def qkv_norm_rope( + self, + q: torch.Tensor, + k: torch.Tensor | None, + v: torch.Tensor | None, + q_weight: torch.Tensor, + k_weight: torch.Tensor | None, + rope_cache: torch.Tensor, + positions: torch.Tensor, + eps: float, + ) -> tuple[torch.Tensor, torch.Tensor | None, torch.Tensor | None]: + raise NotImplementedError + + def router_input( + self, + hidden_states: torch.Tensor, + scale: torch.Tensor, + root_size: float, + eps: float, + ) -> torch.Tensor: + raise NotImplementedError + + def router_topk( + self, + logits: torch.Tensor, + per_expert_scale: torch.Tensor, + top_k: int, + ) -> tuple[torch.Tensor, torch.Tensor]: + raise NotImplementedError + + def gelu_tanh_and_mul(self, x: torch.Tensor) -> torch.Tensor: + raise NotImplementedError + + def gelu_mul(self, gate: torch.Tensor, value: torch.Tensor) -> torch.Tensor: + raise NotImplementedError + + def rmsnorm_residual( + self, + x: torch.Tensor, + weight: torch.Tensor, + residual: torch.Tensor, + eps: float, + scalar: torch.Tensor | None = None, + ) -> torch.Tensor: + raise NotImplementedError + + +GEMMA4_REGISTRY: OpRegistry[Gemma4OpSpec, Gemma4OperatorProvider] = OpRegistry( + "Gemma 4 model operations" +) + + +@GEMMA4_REGISTRY.register +class TritonGemma4OperatorProvider(Gemma4OperatorProvider): + name = "triton" + priority = 10 + + @classmethod + def supports(cls, spec: Gemma4OpSpec, caps: DeviceCaps) -> SupportResult: + if caps.platform != PlatformEnum.CUDA or not caps.supports_triton: + return SupportResult.no("requires CUDA with Triton") + if spec.cuda_graph and not caps.supports_graph_capture: + return SupportResult.no("device does not support CUDA Graph capture") + if spec.activation_dtype not in {torch.bfloat16, torch.float16}: + return SupportResult.no("requires BF16 or FP16 activations") + if any(head_dim not in {256, 512} for head_dim in spec.head_dims): + return SupportResult.no("requires attention head dimensions 256 or 512") + return SupportResult.yes() + + def attention_backend(self, *, sliding_window: int | None): + from sparsevllm.operators.gemma4_attention import Gemma4AttentionBackend + + return Gemma4AttentionBackend(sliding_window=sliding_window) + + def rmsnorm(self, x, weight, eps): + from sparsevllm.kernels.triton.gemma4_rmsnorm import gemma4_rmsnorm + + return gemma4_rmsnorm(x, weight, eps) + + def qkv_norm_rope( + self, q, k, v, q_weight, k_weight, rope_cache, positions, eps + ): + from sparsevllm.kernels.triton.gemma4_qkv_norm_rope import ( + gemma4_qkv_norm_rope, + ) + + gemma4_qkv_norm_rope( + q, k, v, q_weight, k_weight, rope_cache, positions, eps + ) + return q, k, v + + def router_input(self, hidden_states, scale, root_size, eps): + from sparsevllm.kernels.triton.gemma4_router import gemma4_router_input + + return gemma4_router_input(hidden_states, scale, root_size, eps) + + def router_topk(self, logits, per_expert_scale, top_k): + from sparsevllm.kernels.triton.gemma4_router import gemma4_router_topk + + return gemma4_router_topk(logits, per_expert_scale, top_k) + + def gelu_tanh_and_mul(self, x): + from sparsevllm.kernels.triton.gemma4_gelu_and_mul import ( + gelu_tanh_and_mul_fwd, + ) + + return gelu_tanh_and_mul_fwd(x) + + def gelu_mul(self, gate, value): + from sparsevllm.kernels.triton.gemma4_fused_ops import gemma4_gelu_mul + + return gemma4_gelu_mul(gate, value) + + def rmsnorm_residual(self, x, weight, residual, eps, scalar=None): + from sparsevllm.kernels.triton.gemma4_fused_ops import ( + gemma4_rmsnorm_residual, + ) + + return gemma4_rmsnorm_residual(x, weight, residual, eps, scalar) + + +class TorchGemma4OperatorProvider(Gemma4OperatorProvider): + """Explicit correctness oracle; never selected for production inference.""" + + name = "torch_oracle" + + def attention_backend(self, *, sliding_window: int | None): + from sparsevllm.operators.gemma4_attention import Gemma4AttentionBackend + + return Gemma4AttentionBackend(sliding_window=sliding_window) + + def rmsnorm(self, x, weight, eps): + output = x.float() + output *= torch.rsqrt(output.square().mean(-1, keepdim=True) + eps) + if weight is not None: + output *= weight.float() + return output.to(x.dtype) + + def qkv_norm_rope( + self, q, k, v, q_weight, k_weight, rope_cache, positions, eps + ): + q = self.rmsnorm(q, q_weight, eps) + cos, sin = rope_cache[positions].chunk(2, -1) + q = apply_rotary_emb(q, cos, sin) + if k is not None: + k = apply_rotary_emb(self.rmsnorm(k, k_weight, eps), cos, sin) + v = self.rmsnorm(v, None, eps) + return q, k, v + + def router_input(self, hidden_states, scale, root_size, eps): + return self.rmsnorm(hidden_states, None, eps) * scale * root_size + + def router_topk(self, logits, per_expert_scale, top_k): + probabilities = F.softmax(logits, dim=-1, dtype=torch.float32) + weights, ids = probabilities.topk(top_k, dim=-1) + weights.div_(weights.sum(-1, keepdim=True)).mul_(per_expert_scale[ids]) + return weights, ids + + def gelu_tanh_and_mul(self, x): + gate, up = x.chunk(2, -1) + return F.gelu(gate, approximate="tanh") * up + + def gelu_mul(self, gate, value): + return F.gelu(gate, approximate="tanh") * value + + def rmsnorm_residual(self, x, weight, residual, eps, scalar=None): + output = self.rmsnorm(x, weight, eps) + residual + return output if scalar is None else output * scalar + + +def resolve_gemma4_provider( + spec: Gemma4OpSpec, *, device_index: int | None = None +) -> Gemma4OperatorProvider: + platform = platforms.current_platform + if device_index is None: + device_index = torch.cuda.current_device() if platform.is_cuda_alike() else 0 + caps = platform.get_device_caps(int(device_index)) + return OpResolver(GEMMA4_REGISTRY).resolve(spec, caps).provider + + +__all__ = [ + "GEMMA4_REGISTRY", + "Gemma4OperatorProvider", + "Gemma4OpSpec", + "TorchGemma4OperatorProvider", + "TritonGemma4OperatorProvider", + "resolve_gemma4_provider", +] diff --git a/src/sparsevllm/operators/gemma4_moe.py b/src/sparsevllm/operators/gemma4_moe.py index 3c340f86..59e08929 100644 --- a/src/sparsevllm/operators/gemma4_moe.py +++ b/src/sparsevllm/operators/gemma4_moe.py @@ -4,17 +4,48 @@ import torch.nn.functional as F from sparsevllm import platforms -from sparsevllm.operators.moe import MoeOpSpec, MoeProvider +from sparsevllm.distributed import get_parallel_context +from sparsevllm.layers.packed_moe import PackedMoeExperts +from sparsevllm.operators.moe import MoeOpSpec, MoeProvider, model_activation_dtype from sparsevllm.operators.registry import OpRegistry, OpResolver, SupportResult from sparsevllm.platforms.interface import DeviceCaps, PlatformEnum -GEMMA4_MOE_REGISTRY: OpRegistry[MoeOpSpec, MoeProvider] = OpRegistry( + +class Gemma4MoeProvider(MoeProvider): + def load_packed_projection( + self, + spec: MoeOpSpec, + *, + projection: str, + loaded_weight: torch.Tensor, + tp_rank: int, + tp_size: int, + local_expert_start: int, + local_expert_end: int, + w13_weight: torch.Tensor, + w2_weight: torch.Tensor, + ) -> tuple[str, ...]: + if loaded_weight.shape[0] == spec.num_experts: + loaded_weight = loaded_weight[local_expert_start:local_expert_end] + if projection == "gate_up_proj": + gate, up = loaded_weight.chunk(2, 1) + gate = gate.chunk(tp_size, 1)[tp_rank] + up = up.chunk(tp_size, 1)[tp_rank] + w13_weight.copy_(torch.cat((gate, up), 1)) + return "gate_proj", "up_proj" + if projection == "down_proj": + w2_weight.copy_(loaded_weight.chunk(tp_size, 2)[tp_rank]) + return ("down_proj",) + raise ValueError(f"Unsupported Gemma 4 projection {projection!r}.") + + +GEMMA4_MOE_REGISTRY: OpRegistry[MoeOpSpec, Gemma4MoeProvider] = OpRegistry( "Gemma 4 routed GEGLU MoE" ) @GEMMA4_MOE_REGISTRY.register -class TritonGemma4MoeProvider(MoeProvider): +class TritonGemma4MoeProvider(Gemma4MoeProvider): name = "triton_gemma4_geglu" priority = 10 gate_up_order = "gate_up" @@ -65,8 +96,9 @@ def run( ) -@GEMMA4_MOE_REGISTRY.register -class TorchGemma4MoeProvider(MoeProvider): +class TorchGemma4MoeProvider(Gemma4MoeProvider): + """Explicit correctness oracle; never selected for production inference.""" + name = "torch_gemma4_geglu" priority = 0 @@ -118,7 +150,7 @@ def resolve_gemma4_moe_provider( spec: MoeOpSpec, *, device_index: int | None = None, -) -> MoeProvider: +) -> Gemma4MoeProvider: if spec.activation != "gelu_tanh": raise ValueError( "Gemma 4 MoE resolver requires activation='gelu_tanh', " @@ -131,8 +163,64 @@ def resolve_gemma4_moe_provider( return OpResolver(GEMMA4_MOE_REGISTRY).resolve(spec, caps).provider +class Gemma4PackedExperts(PackedMoeExperts): + def __init__(self, config) -> None: + super().__init__( + num_experts=config.num_experts, + hidden_size=config.hidden_size, + intermediate_size=config.moe_intermediate_size, + top_k=config.top_k_experts, + activation_dtype=model_activation_dtype(config), + fp8_enabled=False, + cuda_graph=bool(getattr(config, "decode_cuda_graph", False)), + activation="gelu_tanh", + model_label="Gemma4MoE", + provider_resolver=resolve_gemma4_moe_provider, + parallel_context=get_parallel_context(), + ) + + def rank_local_weight_slice( + self, + source_shape: tuple[int, ...], + *, + loaded_shard_id: str, + is_scale: bool = False, + ) -> tuple[slice, ...] | None: + if is_scale: + raise ValueError("Gemma 4 BF16 experts do not use weight scales.") + if len(source_shape) != 3 or int(source_shape[0]) != self.num_experts: + raise ValueError(f"Invalid Gemma 4 packed expert shape {source_shape}.") + if self.ep_size == 1: + return None + return ( + slice(self.local_expert_start, self.local_expert_end), + slice(None), + slice(None), + ) + + def load_packed_weight(self, projection: str, loaded_weight: torch.Tensor) -> None: + projections = self.provider.load_packed_projection( + self.op_spec, + projection=projection, + loaded_weight=loaded_weight, + tp_rank=self.tp_rank, + tp_size=self.tp_size, + local_expert_start=self.local_expert_start, + local_expert_end=self.local_expert_end, + w13_weight=self.w13_weight.data, + w2_weight=self.w2_weight.data, + ) + self._loaded_expert_shards.update( + (expert_id, name) + for expert_id in range(self.local_expert_start, self.local_expert_end) + for name in projections + ) + + __all__ = [ "GEMMA4_MOE_REGISTRY", + "Gemma4MoeProvider", + "Gemma4PackedExperts", "TorchGemma4MoeProvider", "TritonGemma4MoeProvider", "resolve_gemma4_moe_provider", diff --git a/tests/test_activation.py b/tests/test_activation.py index df3a6975..bfa6b055 100644 --- a/tests/test_activation.py +++ b/tests/test_activation.py @@ -2,13 +2,10 @@ import torch import torch.nn.functional as F -from sparsevllm.layers.activation import GeluTanhAndMul, SiluAndMul +from sparsevllm.layers.activation import SiluAndMul from sparsevllm.operators.activation import ( - GeluTanhAndMulSpec, SiluAndMulSpec, - TorchGeluTanhAndMulProvider, TorchSiluAndMulProvider, - TritonGeluTanhAndMulProvider, TritonSiluAndMulProvider, ) @@ -32,30 +29,6 @@ def test_silu_and_mul_rejects_odd_width(): SiluAndMul(provider=TorchSiluAndMulProvider())(torch.randn(2, 7)) -def test_gemma4_gelu_and_mul_cpu_matches_reference(): - x = torch.randn(3, 16) - gate, up = x.chunk(2, -1) - expected = F.gelu(gate, approximate="tanh") * up - actual = GeluTanhAndMul(provider=TorchGeluTanhAndMulProvider())(x) - torch.testing.assert_close(actual, expected) - - -@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA is required") -@pytest.mark.parametrize("dtype", [torch.float16, torch.bfloat16]) -@pytest.mark.parametrize("rows", [1, 256, 257]) -def test_gemma4_gelu_and_mul_cuda_matches_reference(dtype, rows): - torch.manual_seed(20260813) - x = torch.randn(rows, 1408, dtype=dtype, device="cuda") - gate, up = x.chunk(2, -1) - expected = F.gelu(gate, approximate="tanh") * up - actual_input = x.clone() - actual = GeluTanhAndMul( - TritonGeluTanhAndMulProvider(op_spec=GeluTanhAndMulSpec(dtype)) - )(actual_input) - torch.testing.assert_close(actual, expected, rtol=3e-3, atol=3e-3) - assert actual.data_ptr() == actual_input.data_ptr() - - @pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA is required") @pytest.mark.parametrize("dtype", [torch.float16, torch.bfloat16]) @pytest.mark.parametrize("rows", [8, 256, 257]) diff --git a/tests/test_gemma4_model.py b/tests/test_gemma4_model.py index ef13d5e7..52f64566 100644 --- a/tests/test_gemma4_model.py +++ b/tests/test_gemma4_model.py @@ -6,6 +6,7 @@ import pytest import torch +import torch.nn.functional as F from transformers import Gemma4TextConfig from transformers.models.gemma4.modeling_gemma4 import ( Gemma4TextMLP as HFGemma4MLP, @@ -31,7 +32,15 @@ Gemma4Router, ) from sparsevllm.models.layout import RuntimeLayout -from sparsevllm.operators.activation import TorchGeluTanhAndMulProvider +from sparsevllm.operators.gemma4 import ( + TorchGemma4OperatorProvider, + TritonGemma4OperatorProvider, +) +from sparsevllm.operators.gemma4_moe import ( + GEMMA4_MOE_REGISTRY, + TorchGemma4MoeProvider, + TritonGemma4MoeProvider, +) @pytest.mark.skipif(not torch.cuda.is_available(), reason="requires CUDA") @@ -62,6 +71,25 @@ def test_gemma4_router_kernels_match_torch(): torch.testing.assert_close(actual_weights, expected_weights, rtol=1e-5, atol=1e-6) +@pytest.mark.skipif(not torch.cuda.is_available(), reason="requires CUDA") +@pytest.mark.parametrize("dtype", [torch.float16, torch.bfloat16]) +@pytest.mark.parametrize("rows", [1, 256, 257]) +def test_gemma4_provider_gelu_tanh_and_mul_matches_torch(dtype, rows): + torch.manual_seed(20260813) + x = torch.randn(rows, 1408, dtype=dtype, device="cuda") + gate, up = x.chunk(2, -1) + expected = F.gelu(gate, approximate="tanh") * up + actual_input = x.clone() + actual = TritonGemma4OperatorProvider().gelu_tanh_and_mul(actual_input) + torch.testing.assert_close(actual, expected, rtol=3e-3, atol=3e-3) + assert actual.data_ptr() == actual_input.data_ptr() + + +def test_gemma4_moe_torch_oracle_is_not_a_production_fallback(): + assert GEMMA4_MOE_REGISTRY.providers == (TritonGemma4MoeProvider,) + assert TorchGemma4MoeProvider not in GEMMA4_MOE_REGISTRY.providers + + def _parallel_context() -> ParallelContext: group = ParallelGroup(process_group=None, ranks=(0,), rank=0, size=1) return ParallelContext(world=group, tensor=group, expert=group, data=group) @@ -128,7 +156,7 @@ def test_gemma4_rope_matches_transformers_for_both_layer_types(): def test_gemma4_dense_mlp_matches_transformers(): config = _config() with _patch_parallel_context(): - actual = Gemma4MLP(config, 0, TorchGeluTanhAndMulProvider()) + actual = Gemma4MLP(config, 0, TorchGemma4OperatorProvider()) reference = HFGemma4MLP(config, 0) torch.manual_seed(3) for parameter in reference.parameters(): @@ -151,7 +179,7 @@ def test_gemma4_router_matches_transformers(): enable_moe_block=True, num_experts=4, top_k_experts=2, moe_intermediate_size=4 ) with _patch_parallel_context(): - actual = Gemma4Router(config) + actual = Gemma4Router(config, TorchGemma4OperatorProvider()) reference = HFGemma4Router(config) torch.manual_seed(5) reference.proj.weight.data.normal_(0, 0.2) @@ -172,6 +200,7 @@ def test_gemma4_k_eq_v_loader_duplicates_normalized_projection_slot(): config, 1, Gemma4RotaryEmbedding(config, "full_attention", config.global_head_dim), + TorchGemma4OperatorProvider(), ) loaded_key = torch.randn( config.num_global_key_value_heads * config.global_head_dim, config.hidden_size @@ -190,7 +219,7 @@ def test_gemma4_ple_matches_transformers(): vocab_size_per_layer_input=32, ) with _patch_parallel_context(): - actual = Gemma4Model(config, TorchGeluTanhAndMulProvider()) + actual = Gemma4Model(config, TorchGemma4OperatorProvider()) reference = HFGemma4Model(config) torch.manual_seed(11) for parameter in reference.parameters(): @@ -260,6 +289,7 @@ def test_gemma4_shared_kv_attention_only_allocates_query_projection(): config, 2, Gemma4RotaryEmbedding(config, "sliding_attention", config.head_dim), + TorchGemma4OperatorProvider(), ) assert attention.is_kv_shared_layer assert tuple(attention.qkv_proj.weight.shape) == ( diff --git a/tests/test_gemma4_rmsnorm.py b/tests/test_gemma4_rmsnorm.py index 19b8d734..aaf329ef 100644 --- a/tests/test_gemma4_rmsnorm.py +++ b/tests/test_gemma4_rmsnorm.py @@ -2,6 +2,14 @@ import torch from sparsevllm.layers.gemma4_rmsnorm import Gemma4RMSNorm +from sparsevllm.operators.gemma4 import ( + GEMMA4_REGISTRY, + Gemma4OpSpec, + TorchGemma4OperatorProvider, + TritonGemma4OperatorProvider, +) +from sparsevllm.operators.registry import OpResolver +from sparsevllm.platforms.interface import DeviceCaps, PlatformEnum def _reference(x: torch.Tensor, weight: torch.Tensor | None, eps: float) -> torch.Tensor: @@ -15,7 +23,9 @@ def _reference(x: torch.Tensor, weight: torch.Tensor | None, eps: float) -> torc @pytest.mark.parametrize("with_scale", [False, True]) def test_gemma4_rmsnorm_matches_torch(with_scale): torch.manual_seed(7) - layer = Gemma4RMSNorm(32, with_scale=with_scale) + layer = Gemma4RMSNorm( + 32, with_scale=with_scale, provider=TorchGemma4OperatorProvider() + ) x = torch.randn(5, 32) torch.testing.assert_close(layer(x), _reference(x, layer.weight, layer.eps)) @@ -23,7 +33,9 @@ def test_gemma4_rmsnorm_matches_torch(with_scale): @pytest.mark.skipif(not torch.cuda.is_available(), reason="requires CUDA") def test_gemma4_rmsnorm_cuda_graph_matches_torch(): torch.manual_seed(11) - layer = Gemma4RMSNorm(2816).cuda().to(torch.bfloat16) + layer = Gemma4RMSNorm( + 2816, provider=TritonGemma4OperatorProvider() + ).cuda().to(torch.bfloat16) x = torch.randn(4, 2816, device="cuda", dtype=torch.bfloat16) for _ in range(2): layer(x) @@ -32,3 +44,24 @@ def test_gemma4_rmsnorm_cuda_graph_matches_torch(): output = layer(x) graph.replay() torch.testing.assert_close(output, _reference(x, layer.weight, layer.eps), rtol=0, atol=0) + + +def test_gemma4_provider_requires_supported_cuda_profile(): + spec = Gemma4OpSpec(torch.bfloat16, (256, 512), cuda_graph=True) + caps = DeviceCaps( + platform=PlatformEnum.CUDA, + device_type="cuda", + device_index=0, + device_name="test", + supports_graph_capture=True, + supports_triton=True, + ) + assert isinstance( + OpResolver(GEMMA4_REGISTRY).resolve(spec, caps).provider, + TritonGemma4OperatorProvider, + ) + + with pytest.raises(RuntimeError, match="requires attention head dimensions"): + OpResolver(GEMMA4_REGISTRY).resolve( + Gemma4OpSpec(torch.bfloat16, (128,), cuda_graph=True), caps + ) From ea57c14870318bbf1d450eae7aecc6eab5fb0d92 Mon Sep 17 00:00:00 2001 From: QuanshengGu Date: Thu, 13 Aug 2026 05:14:48 +0800 Subject: [PATCH 03/13] feat: add native multimodal inference --- src/sparsevllm/__init__.py | 6 +- src/sparsevllm/configs/runtime.py | 1 + .../cache_manager/prefix_cache_mixin.py | 2 + src/sparsevllm/engine/chain_cache.py | 1 + src/sparsevllm/engine/llm_engine.py | 79 ++++- src/sparsevllm/engine/model_runner.py | 57 +++- .../engine/prefix_cache_coordinator.py | 2 + src/sparsevllm/engine/runtime_state.py | 5 + src/sparsevllm/engine/sequence.py | 16 +- .../entrypoints/openai/dispatcher.py | 7 +- .../entrypoints/openai/protocol/chat.py | 14 +- src/sparsevllm/entrypoints/openai/render.py | 57 +++- .../entrypoints/openai/serving/chat.py | 4 +- .../entrypoints/openai/serving/responses.py | 4 +- .../gemma4_multimodal_context_attention.py | 244 ++++++++++++++++ src/sparsevllm/models/gemma4.py | 70 ++++- src/sparsevllm/models/gemma4_multimodal.py | 90 ++++++ src/sparsevllm/models/qwen3_5.py | 89 +++++- src/sparsevllm/models/qwen3_5_moe.py | 29 +- src/sparsevllm/models/qwen3_5_multimodal.py | 110 +++++++ src/sparsevllm/multimodal/__init__.py | 3 + src/sparsevllm/multimodal/inputs.py | 171 +++++++++++ src/sparsevllm/multimodal/runtime.py | 158 ++++++++++ src/sparsevllm/operators/gemma4_attention.py | 24 ++ src/sparsevllm/operators/qwen35_mrope.py | 77 +++++ src/sparsevllm/utils/context.py | 2 + src/sparsevllm/utils/loader.py | 35 ++- tests/test_gemma4_model.py | 9 + tests/test_multimodal.py | 275 ++++++++++++++++++ tests/test_openai_api_server.py | 45 +++ tests/test_qwen35_mixed_runtime.py | 22 +- tests/test_weight_loading.py | 19 ++ 32 files changed, 1666 insertions(+), 61 deletions(-) create mode 100644 src/sparsevllm/kernels/triton/gemma4_multimodal_context_attention.py create mode 100644 src/sparsevllm/models/gemma4_multimodal.py create mode 100644 src/sparsevllm/models/qwen3_5_multimodal.py create mode 100644 src/sparsevllm/multimodal/__init__.py create mode 100644 src/sparsevllm/multimodal/inputs.py create mode 100644 src/sparsevllm/multimodal/runtime.py create mode 100644 src/sparsevllm/operators/qwen35_mrope.py create mode 100644 tests/test_multimodal.py diff --git a/src/sparsevllm/__init__.py b/src/sparsevllm/__init__.py index b2354b45..26243b43 100644 --- a/src/sparsevllm/__init__.py +++ b/src/sparsevllm/__init__.py @@ -1,6 +1,6 @@ from __future__ import annotations -__all__ = ["LLM", "SamplingParams"] +__all__ = ["LLM", "MultiModalPrompt", "SamplingParams"] def __getattr__(name: str): @@ -12,4 +12,8 @@ def __getattr__(name: str): from sparsevllm.sampling_params import SamplingParams return SamplingParams + if name == "MultiModalPrompt": + from sparsevllm.multimodal import MultiModalPrompt + + return MultiModalPrompt raise AttributeError(f"module 'sparsevllm' has no attribute {name!r}") diff --git a/src/sparsevllm/configs/runtime.py b/src/sparsevllm/configs/runtime.py index c4d1722c..f863222f 100644 --- a/src/sparsevllm/configs/runtime.py +++ b/src/sparsevllm/configs/runtime.py @@ -72,6 +72,7 @@ class Config( # least one synchronous loading path when the budget is smaller. weight_loading_workers: int = 1 enforce_eager: bool = True + enable_multimodal: bool = True hf_config: AutoConfig | None = None outer_hf_config: Any | None = None runtime_layout: RuntimeLayout | None = None diff --git a/src/sparsevllm/engine/cache_manager/prefix_cache_mixin.py b/src/sparsevllm/engine/cache_manager/prefix_cache_mixin.py index 2b8483b8..c6e8a523 100644 --- a/src/sparsevllm/engine/cache_manager/prefix_cache_mixin.py +++ b/src/sparsevllm/engine/cache_manager/prefix_cache_mixin.py @@ -296,6 +296,8 @@ def _record_prefix_materialization( token_ids: list[int], slots: torch.Tensor, ) -> None: + if getattr(seq, "multimodal_digest", None) is not None: + return if not self.enable_prefix_caching or self.prefix_cache is None: return if len(token_ids) != int(slots.numel()): diff --git a/src/sparsevllm/engine/chain_cache.py b/src/sparsevllm/engine/chain_cache.py index 02e07ad7..250ab1f7 100644 --- a/src/sparsevllm/engine/chain_cache.py +++ b/src/sparsevllm/engine/chain_cache.py @@ -80,6 +80,7 @@ class RequestAdmission: chain_status: str reused_tokens: int prefilled_tokens: int = 0 + prompt_token_ids: list[int] | None = None @dataclass(slots=True) diff --git a/src/sparsevllm/engine/llm_engine.py b/src/sparsevllm/engine/llm_engine.py index a18b3465..abd22b95 100644 --- a/src/sparsevllm/engine/llm_engine.py +++ b/src/sparsevllm/engine/llm_engine.py @@ -1,7 +1,9 @@ import atexit import gc import os +import pickle from dataclasses import fields +from multiprocessing.shared_memory import SharedMemory from time import perf_counter import threading from tqdm.auto import tqdm @@ -20,6 +22,11 @@ from sparsevllm.engine.scheduler import Scheduler from sparsevllm.engine.model_runner import ModelRunner, make_tp_shm_name from sparsevllm.engine.input_processor import tokenize_text_prompt +from sparsevllm.multimodal.inputs import ( + MultiModalInputProcessor, + MultiModalPrompt, + is_multimodal_prompt, +) from sparsevllm.engine.prefix_cache import PrefixCacheRoutingSnapshot from sparsevllm.engine.chain_cache import ( ChainCacheIndex, @@ -246,6 +253,12 @@ def __init__(self, model, **kwargs): # 加载分词器 self.tokenizer: Qwen2Tokenizer = AutoTokenizer.from_pretrained(config.model, use_fast=True) + self.multimodal_processor = ( + MultiModalInputProcessor(config.model) + if config.enable_multimodal + and callable(getattr(self.model_runner.model, "encode_multimodal", None)) + else None + ) generation_config = GenerationConfig.from_pretrained(config.model) eos_values = generation_config.eos_token_id if eos_values is None: @@ -580,7 +593,7 @@ def _tokenize_prompt(self, prompt: str | list[int]) -> list[int]: def admit_request( self, - prompt: str | list[int], + prompt: str | list[int] | MultiModalPrompt | dict, sampling_params: SamplingParams, chain_id: str | None = None, chain_append_only: bool = False, @@ -590,6 +603,16 @@ def admit_request( In chain mode the returned seq_id is the resident sequence identity and remains stable across turns. The caller's request identity is separate. """ + multimodal = None + if is_multimodal_prompt(prompt): + if self.multimodal_processor is None: + raise NotImplementedError( + "Multimodal input is disabled or unsupported by this model." + ) + if chain_id or chain_append_only: + raise ChainModeError("Multimodal requests do not support chain mode.") + multimodal = self.multimodal_processor.process(prompt) + prompt = multimodal.token_ids mode = str( getattr(self.config, "resolved_prefix_cache_mode", "disabled") ) @@ -644,6 +667,31 @@ def admit_request( ) logger.debug(f'add prompt with {len(prompt)} tokens.') seq = Sequence(prompt, sampling_params) + if multimodal is not None: + seq.multimodal_digest = multimodal.digest + seq.multimodal_full_prefill = ( + getattr(self.config.hf_config, "use_bidirectional_attention", None) + == "vision" + ) + payload = pickle.dumps(multimodal.tensors, protocol=pickle.HIGHEST_PROTOCOL) + payload_shm = SharedMemory(create=True, size=len(payload)) + try: + payload_shm.buf[: len(payload)] = payload + seq.multimodal_position_delta = int( + self.model_runner.call( + "register_multimodal_shared", + int(seq.seq_id), + list(prompt), + payload_shm.name, + len(payload), + ) + ) + except Exception: + self.model_runner.call("free_multimodal", int(seq.seq_id)) + raise + finally: + payload_shm.close() + payload_shm.unlink() if mode != "chain": if normalized_chain_id: raise ChainModeError( @@ -651,13 +699,19 @@ def admit_request( "prefix_cache_mode='chain'.", chain_id=normalized_chain_id, ) - self.scheduler.add(seq) + try: + self.scheduler.add(seq) + except Exception: + if multimodal is not None: + self.model_runner.call("free_multimodal", int(seq.seq_id)) + raise return RequestAdmission( seq_id=int(seq.seq_id), chain_id=None, chain_status="disabled", reused_tokens=0, prefilled_tokens=prompt_len, + prompt_token_ids=list(prompt), ) coordinator = self.model_runner.runtime_state.chain_cache_coordinator @@ -743,9 +797,14 @@ def admit_request( chain_status=chain_status, reused_tokens=int(plan.reused_tokens), prefilled_tokens=prompt_len - int(plan.reused_tokens), + prompt_token_ids=list(prompt), ) - def add_request(self, prompt: str | list[int], sampling_params: SamplingParams): + def add_request( + self, + prompt: str | list[int] | MultiModalPrompt | dict, + sampling_params: SamplingParams, + ): """Backward-compatible request API returning only seq_id.""" return self.admit_request(prompt, sampling_params).seq_id @@ -762,6 +821,14 @@ def abort_request(self, seq_id: int, disposition: str = "invalidate"): f"{disposition!r}." ) chain_seq = self._active_chain_sequences.get(int(seq_id)) + multimodal = any( + seq.seq_id == seq_id and seq.multimodal_digest is not None + for queue in ( + getattr(self.scheduler, "waiting", ()), + getattr(self.scheduler, "decoding", ()), + ) + for seq in queue + ) should_free = self.scheduler.abort(seq_id) if chain_seq is not None: self._active_chain_sequences.pop(int(seq_id), None) @@ -787,6 +854,8 @@ def abort_request(self, seq_id: int, disposition: str = "invalidate"): return if should_free: self.model_runner.call("free_slots", seq_id) + elif multimodal: + self.model_runner.call("free_multimodal", seq_id) def chain_cache_routing_match(self, chain_id: str) -> dict[str, object]: return self.model_runner.runtime_state.chain_routing_match( @@ -1258,7 +1327,7 @@ def step(self): ) ) if finished_seq_ids: - self.model_runner.call("free_slots_batch", finished_seq_ids) + self.model_runner.call("finish_slots_batch", finished_seq_ids) # 计算吞吐量统计数据 (正数表示 Prefill,负数表示 Decode) num_tokens = sum(seq.current_chunk_size for seq in seqs) if is_prefill else -len(seqs) @@ -1294,7 +1363,7 @@ def is_finished(self): def generate( self, - prompts: list[str] | list[list[int]], + prompts: list[str] | list[list[int]] | list[MultiModalPrompt] | list[dict], sampling_params: SamplingParams | list[SamplingParams], use_tqdm: bool = True, ) -> list[dict]: diff --git a/src/sparsevllm/engine/model_runner.py b/src/sparsevllm/engine/model_runner.py index 868d6928..1db4b039 100644 --- a/src/sparsevllm/engine/model_runner.py +++ b/src/sparsevllm/engine/model_runner.py @@ -30,6 +30,7 @@ from sparsevllm.engine.chain_cache import ChainAdmissionPlan, ChainCacheCoordinator from sparsevllm.engine.recurrent_state_manager import RecurrentStateManager, RecurrentStateSpec from sparsevllm.engine.runtime_state import RuntimeState +from sparsevllm.multimodal.runtime import MultiModalRuntime from sparsevllm.engine.sparse_controller import SparseController from sparsevllm.models.spec import ModelSpec import sparsevllm.platforms as platforms @@ -77,10 +78,19 @@ def _create_model(hf_config, model_spec: ModelSpec, **runtime_kwargs): if model_class is None: raise ImportError(f"{class_name} is unavailable for {model_spec.name}.") builder = getattr(model_class, "build_runtime_kwargs", None) - return model_class( + model = model_class( hf_config, **(builder(hf_config, **runtime_kwargs) if callable(builder) else {}), ) + engine_config = runtime_kwargs.get("engine_config") + configure_multimodal = getattr(model, "configure_multimodal", None) + if ( + callable(configure_multimodal) + and bool(getattr(engine_config, "enable_multimodal", True)) + and getattr(engine_config, "outer_hf_config", hf_config) is not hf_config + ): + configure_multimodal(engine_config.outer_hf_config) + return model TP_SHM_NAME_PREFIX = "sparsevllm_" @@ -104,10 +114,13 @@ def _create_model(hf_config, model_spec: ModelSpec, **runtime_kwargs): "debug_moe_states_cpu", "free_slots", "free_slots_batch", + "finish_slots_batch", + "free_multimodal", "log_operator_implementations", "refresh_prefix_cache_hit", "reset_after_warmup", "run", + "register_multimodal_shared", "set_warmup_fake_prefill_attention", "warmup_moe_workspace", } @@ -208,6 +221,10 @@ def __init__( show_progress=self.parallel_context.world_rank == 0, progress_rank=0 if self.parallel_context.world_rank == 0 else None, ) + self.model.eval() + self.multimodal_runtime = MultiModalRuntime(self.model, self.device) + self._prefill_inputs_embeds = None + self._prefill_multimodal_mask = None warmup_moe = getattr(self.model, "warmup_moe", None) if callable(warmup_moe): warmup_moe() @@ -648,6 +665,7 @@ def free_slots(self, seq_id: int): before = self.cache_manager.free_slot_stats() logger.info("model_runner.free_slots seq_id={} before={}", seq_id, before) self.runtime_state.free_seq(seq_id) + self.multimodal_runtime.free(seq_id) if os.getenv("SPARSEVLLM_DEBUG_SLOTS", "0") == "1": after = self.cache_manager.free_slot_stats() logger.info("model_runner.free_slots seq_id={} after={}", seq_id, after) @@ -667,6 +685,27 @@ def free_slots_batch(self, seq_ids: list[int]): after = self.cache_manager.free_slot_stats() logger.info("model_runner.free_slots_batch seq_ids={} after={}", seq_ids, after) + def finish_slots_batch(self, seq_ids: list[int]): + self.free_slots_batch(seq_ids) + self.multimodal_runtime.free_batch(seq_ids) + + def free_multimodal(self, seq_id: int): + self.multimodal_runtime.free(seq_id) + + def register_multimodal_shared( + self, + seq_id: int, + input_ids: list[int], + shared_name: str, + payload_size: int, + ) -> int: + payload_shm = SharedMemory(name=shared_name) + try: + tensors = pickle.loads(bytes(payload_shm.buf[: int(payload_size)])) + finally: + payload_shm.close() + return self.multimodal_runtime.register(seq_id, input_ids, tensors) + def chain_admission_plan( self, chain_id: str, @@ -1114,6 +1153,11 @@ def prepare_step(self, seqs: list[Sequence], is_prefill: bool): seqs=seqs, recurrent_state_manager=self.recurrent_state_manager, ) + ( + self._prefill_inputs_embeds, + positions, + self._prefill_multimodal_mask, + ) = self.multimodal_runtime.prepare(seqs, input_ids, positions, is_prefill) return input_ids, positions def prepare_sample(self, seqs: list[Sequence]): @@ -1274,7 +1318,16 @@ def run_model(self, input_ids: torch.Tensor, positions: torch.Tensor, is_prefill """物理执行逻辑:统一使用 Eager 模式""" _stage = 'prefill' if is_prefill else 'decode' with profiler.record(f"model_run_model_{_stage}"): - logits = self.model.compute_logits(self.model(input_ids, positions)) + if is_prefill and self._prefill_inputs_embeds is not None: + hidden_states = self.multimodal_runtime.forward( + input_ids, + positions, + self._prefill_inputs_embeds, + self._prefill_multimodal_mask, + ) + else: + hidden_states = self.model(input_ids, positions) + logits = self.model.compute_logits(hidden_states) self._record_debug_logits(logits) return logits diff --git a/src/sparsevllm/engine/prefix_cache_coordinator.py b/src/sparsevllm/engine/prefix_cache_coordinator.py index 15aab5bd..b003c2aa 100644 --- a/src/sparsevllm/engine/prefix_cache_coordinator.py +++ b/src/sparsevllm/engine/prefix_cache_coordinator.py @@ -574,6 +574,8 @@ def finish_step(self) -> None: self._step_h2d_operations.clear() def _record_tokens(self, seq: Sequence, token_ids: list[int]) -> None: + if getattr(seq, "multimodal_digest", None) is not None: + return if not token_ids: return state = self.runtime_states.get(int(seq.seq_id)) diff --git a/src/sparsevllm/engine/runtime_state.py b/src/sparsevllm/engine/runtime_state.py index 4caa7078..b56f30cc 100644 --- a/src/sparsevllm/engine/runtime_state.py +++ b/src/sparsevllm/engine/runtime_state.py @@ -308,6 +308,9 @@ def reset_after_warmup(self) -> None: self._resident_seq_ids.clear() def refresh_prefix_cache_hit(self, seq: Sequence) -> None: + if getattr(seq, "multimodal_digest", None) is not None: + self.clear_prefix_cache_hit(seq) + return if self.prefix_cache_coordinator is not None: self.prefix_cache_coordinator.refresh_prefix_cache_hit(seq) return @@ -337,6 +340,8 @@ def remaining_prefill_tokens(self, seq: Sequence) -> int: return int(self.cache_manager.remaining_prefill_tokens(seq)) def prefill_execution_mode(self, seq: Sequence) -> str: + if getattr(seq, "multimodal_full_prefill", False): + return "full" return str(self.cache_manager.prefill_execution_mode(seq)) def prefill_batch_compatibility_key(self, seq: Sequence) -> object: diff --git a/src/sparsevllm/engine/sequence.py b/src/sparsevllm/engine/sequence.py index c7da4239..87a6c7f0 100644 --- a/src/sparsevllm/engine/sequence.py +++ b/src/sparsevllm/engine/sequence.py @@ -108,6 +108,9 @@ def __init__(self, token_ids: list[int], sampling_params = SamplingParams()): self.chain_id: str | None = None self.chain_status = "disabled" self.chain_reused_tokens = 0 + self.multimodal_digest: str | None = None + self.multimodal_position_delta = 0 + self.multimodal_full_prefill = False # None means normal generation. During recompute replay the original # prompt is prefetched again, then accepted completion tokens before the # current last token are replayed through decode without being sampled @@ -294,8 +297,10 @@ def decode_input_token(self) -> int: @property def decode_input_position(self) -> int: if self.is_recompute_decode: - return int(self.num_prompt_tokens + int(self.recompute_replay_cursor or 0)) - return int(self.num_tokens - 1) + position = self.num_prompt_tokens + int(self.recompute_replay_cursor or 0) + else: + position = self.num_tokens - 1 + return int(position + self.multimodal_position_delta) @property def should_publish_sample(self) -> bool: @@ -392,6 +397,9 @@ def __getstate__(self): self.chain_reused_tokens, self.recompute_replay_cursor, self.decode_progress_checkpoint, + self.multimodal_digest, + self.multimodal_position_delta, + self.multimodal_full_prefill, ) def __setstate__(self, state): @@ -403,7 +411,9 @@ def __setstate__(self, state): self.prefix_cache_enabled, self.prefix_cache_hit_len, self.prefix_cache_hit_block_count, self.prefix_cache_hit_last_block_id, self.prefix_cache_block_size, self.prefix_cache_method, self.chain_id, self.chain_status, self.chain_reused_tokens, - self.recompute_replay_cursor, self.decode_progress_checkpoint) = state + self.recompute_replay_cursor, self.decode_progress_checkpoint, + self.multimodal_digest, self.multimodal_position_delta, + self.multimodal_full_prefill) = state self.completion_token_logprobs = [] self.completion_top_logprobs = [] # TP workers intentionally receive only the active prompt chunk or one diff --git a/src/sparsevllm/entrypoints/openai/dispatcher.py b/src/sparsevllm/entrypoints/openai/dispatcher.py index 35d2b166..2443bdc4 100644 --- a/src/sparsevllm/entrypoints/openai/dispatcher.py +++ b/src/sparsevllm/entrypoints/openai/dispatcher.py @@ -581,8 +581,13 @@ def _admit(self, item: _QueuedRequest, active: dict[int, _ActiveRequest]): item.handle.terminal.set() self._resolve_admission(item, asyncio.CancelledError()) return + admitted_prompt_token_ids = ( + getattr(admission, "prompt_token_ids", None) if callable(admit) else None + ) prompt_token_ids = ( - list(item.prompt) + list(admitted_prompt_token_ids) + if admitted_prompt_token_ids is not None + else list(item.prompt) if isinstance(item.prompt, list) else self.engine.tokenizer.encode(item.prompt) ) diff --git a/src/sparsevllm/entrypoints/openai/protocol/chat.py b/src/sparsevllm/entrypoints/openai/protocol/chat.py index cd3bbdab..8124b963 100644 --- a/src/sparsevllm/entrypoints/openai/protocol/chat.py +++ b/src/sparsevllm/entrypoints/openai/protocol/chat.py @@ -11,8 +11,18 @@ class ChatContentPart(BaseModel): model_config = ConfigDict(extra="forbid") - type: Literal["text"] - text: str + type: Literal["text", "image_url", "video_url", "input_audio"] + text: str | None = None + image_url: str | dict[str, Any] | None = None + video_url: str | dict[str, Any] | None = None + input_audio: dict[str, Any] | None = None + + @model_validator(mode="after") + def validate_content(self): + value = getattr(self, self.type) + if value is None: + raise ValueError(f"{self.type} content requires its matching field.") + return self class ChatMessage(BaseModel): diff --git a/src/sparsevllm/entrypoints/openai/render.py b/src/sparsevllm/entrypoints/openai/render.py index 2293bd5a..feae654c 100644 --- a/src/sparsevllm/entrypoints/openai/render.py +++ b/src/sparsevllm/entrypoints/openai/render.py @@ -3,6 +3,7 @@ from typing import Any from sparsevllm.entrypoints.openai.protocol.chat import ChatContentPart +from sparsevllm.multimodal import MultiModalPrompt from sparsevllm.entrypoints.openai.protocol.chat import ChatCompletionRequest from sparsevllm.entrypoints.openai.protocol.chat import ChatMessage from sparsevllm.entrypoints.openai.protocol.responses import ResponseRequest @@ -98,7 +99,15 @@ def _chat_content_text(content: str | list[ChatContentPart] | None) -> str: return "" if isinstance(content, str): return content - return "\n".join(part.text for part in content) + return "\n".join(part.text or "" for part in content) + + +def _has_multimodal_content(messages) -> bool: + return any( + isinstance(message.content, list) + and any(part.type != "text" for part in message.content) + for message in messages + ) def validate_chat_template_kwargs(value: Any) -> dict[str, Any] | None: @@ -155,6 +164,8 @@ def _chat_request_append_prompt( tokenizer: Any, request: ChatCompletionRequest, ) -> str: + if _has_multimodal_content(request.messages): + raise ValueError("Multimodal chat does not support chain append rendering.") append_start = request.chain_append_start if append_start is None: raise ValueError("chain_append_start is required for append rendering.") @@ -217,12 +228,18 @@ def _chat_prompt( tools: list[dict[str, Any]] | None = None, *, add_generation_prompt: bool = True, -) -> str: +) -> str | MultiModalPrompt: chat = [] for message in messages: rendered_message = { "role": _chat_template_role(message.role), - "content": None if message.content is None else _chat_content_text(message.content), + "content": ( + None + if message.content is None + else [part.model_dump(exclude_none=True) for part in message.content] + if isinstance(message.content, list) and _has_multimodal_content(messages) + else _chat_content_text(message.content) + ), } if message.reasoning_content is not None: rendered_message["reasoning_content"] = message.reasoning_content @@ -231,6 +248,13 @@ def _chat_prompt( if message.tool_call_id is not None: rendered_message["tool_call_id"] = message.tool_call_id chat.append(rendered_message) + if _has_multimodal_content(messages): + return MultiModalPrompt( + chat, + chat_template_kwargs=chat_template_kwargs, + tools=_tools_for_chat_template(tokenizer, tools) if tools else None, + add_generation_prompt=add_generation_prompt, + ) if getattr(tokenizer, "chat_template", None) and hasattr(tokenizer, "apply_chat_template"): kwargs = { "tokenize": False, @@ -280,10 +304,20 @@ def _chat_template_tool_calls(tool_calls: list[dict[str, Any]]) -> list[dict[str return rendered -def _response_prompt(tokenizer: Any, request: ResponseRequest) -> str: +def _response_prompt(tokenizer: Any, request: ResponseRequest) -> str | MultiModalPrompt: chat_template_kwargs = resolve_response_chat_template_kwargs(request) tools = normalize_tools(request.tools) if request.tools else None messages = _response_messages(request) + if any( + isinstance(message.get("content"), list) + and any(part.get("type") in {"input_image", "input_audio", "input_video"} for part in message["content"]) + for message in messages + ): + return MultiModalPrompt( + messages, + chat_template_kwargs=chat_template_kwargs, + tools=_tools_for_chat_template(tokenizer, tools) if tools else None, + ) has_template = bool(getattr(tokenizer, "chat_template", None)) and hasattr(tokenizer, "apply_chat_template") if has_template: @@ -357,7 +391,7 @@ def _response_message_item(item: dict[str, Any]) -> dict[str, Any]: raise ValueError("message.role must be one of developer, system, user, assistant.") return { "role": _chat_template_role(str(role)), - "content": _response_content_text(item.get("content")), + "content": _response_content(item.get("content")), } @@ -394,23 +428,26 @@ def _minimax_response_messages(messages: list[dict[str, Any]]) -> list[dict[str, return adapted -def _response_content_text(content: Any) -> str: +def _response_content(content: Any) -> str | list[dict[str, Any]]: if isinstance(content, str): return content if isinstance(content, list): - texts = [] + normalized = [] for part in content: if not isinstance(part, dict): raise ValueError("message content parts must be JSON objects.") part_type = part.get("type") + if part_type in {"input_image", "input_audio", "input_video"}: + normalized.append(dict(part)) + continue if part_type not in {"text", "input_text", "output_text"}: raise ValueError(f"Unsupported message content part type: {part_type!r}.") text = part.get("text") if not isinstance(text, str): raise ValueError("message content text parts require a string text field.") - texts.append(text) - return "\n".join(texts) - raise ValueError("message.content must be a string or a text-only content part list.") + normalized.append({"type": "text", "text": text}) + return normalized if any(part["type"] != "text" for part in normalized) else "\n".join(part["text"] for part in normalized) + raise ValueError("message.content must be a string or a content part list.") def _messages_require_chat_template(messages: list[dict[str, Any]]) -> bool: diff --git a/src/sparsevllm/entrypoints/openai/serving/chat.py b/src/sparsevllm/entrypoints/openai/serving/chat.py index d10b165e..aa53d834 100644 --- a/src/sparsevllm/entrypoints/openai/serving/chat.py +++ b/src/sparsevllm/entrypoints/openai/serving/chat.py @@ -127,7 +127,7 @@ async def serve_chat_completion( started, tokenizer, _stream_include_usage(request.stream_options), - prompt=prompt, + prompt=prompt if isinstance(prompt, str) else "", reasoning_parser_name=reasoning_parser_name, parse_tools=bool(chat_tools), response_parser=response_parser, @@ -144,7 +144,7 @@ async def serve_chat_completion( request.model, handles, tokenizer, - prompt=prompt, + prompt=prompt if isinstance(prompt, str) else "", reasoning_parser_name=reasoning_parser_name, parse_tools=bool(chat_tools), response_parser=response_parser, diff --git a/src/sparsevllm/entrypoints/openai/serving/responses.py b/src/sparsevllm/entrypoints/openai/serving/responses.py index 8674b5ac..a04efff1 100644 --- a/src/sparsevllm/entrypoints/openai/serving/responses.py +++ b/src/sparsevllm/entrypoints/openai/serving/responses.py @@ -111,7 +111,7 @@ async def serve_response( started, request_log_path, request, - prompt=prompt, + prompt=prompt if isinstance(prompt, str) else "", reasoning_parser_name=reasoning_parser_name, response_parser=response_parser, is_disconnected=is_disconnected, @@ -126,7 +126,7 @@ async def serve_response( created_at, request.model, handle, - prompt=prompt, + prompt=prompt if isinstance(prompt, str) else "", reasoning_parser_name=reasoning_parser_name, parse_tools=bool(request.tools), response_parser=response_parser, diff --git a/src/sparsevllm/kernels/triton/gemma4_multimodal_context_attention.py b/src/sparsevllm/kernels/triton/gemma4_multimodal_context_attention.py new file mode 100644 index 00000000..97479d44 --- /dev/null +++ b/src/sparsevllm/kernels/triton/gemma4_multimodal_context_attention.py @@ -0,0 +1,244 @@ +from __future__ import annotations + +import torch +import triton +import triton.language as tl + + +@triton.jit +def _gemma4_multimodal_context_attention_kernel( + q, + k, + v, + output, + q_start, + context_lens, + cached_prefix_lens, + active_slots, + req_indices, + image_groups, + attn_score, + stride_qt, + stride_qh, + stride_kt, + stride_kh, + stride_vt, + stride_vh, + stride_ot, + stride_oh, + stride_sb, + stride_ss, + stride_asb, + stride_ash, + stride_asl, + group_size, + NUM_HEADS: tl.constexpr, + HEAD_DIM: tl.constexpr, + BLOCK_M: tl.constexpr, + BLOCK_N: tl.constexpr, + WINDOW: tl.constexpr, + SCORE_MODE: tl.constexpr, +): + query_block = tl.program_id(0) + batch_head = tl.program_id(1) + batch = batch_head // NUM_HEADS + query_head = batch_head % NUM_HEADS + kv_head = query_head // group_size + query_start = tl.load(q_start + batch) + prefix_len = tl.load(cached_prefix_lens + batch) + query_len = tl.load(context_lens + batch) - prefix_len + request = tl.load(req_indices + batch) + query_positions = query_block * BLOCK_M + tl.arange(0, BLOCK_M) + dims = tl.arange(0, HEAD_DIM) + query = tl.load( + q + + (query_start + query_positions[:, None]) * stride_qt + + query_head * stride_qh + + dims[None, :], + mask=query_positions[:, None] < query_len, + other=0.0, + ) + query_groups = tl.load( + image_groups + query_start + query_positions, + mask=query_positions < query_len, + other=0, + ) + max_logit = tl.full((BLOCK_M,), -float("inf"), tl.float32) + denominator = tl.zeros((BLOCK_M,), tl.float32) + accumulator = tl.zeros((BLOCK_M, HEAD_DIM), tl.float32) + max_key = prefix_len + query_len + for key_start in range(0, max_key, BLOCK_N): + key_positions = key_start + tl.arange(0, BLOCK_N) + slots = tl.load( + active_slots + request * stride_sb + key_positions * stride_ss, + mask=key_positions < max_key, + other=0, + ) + key = tl.load( + k + slots[None, :] * stride_kt + kv_head * stride_kh + dims[:, None], + mask=key_positions[None, :] < max_key, + other=0.0, + ) + logits = tl.dot(query, key) * 1.4426950408889634 + absolute_queries = prefix_len + query_positions[:, None] + key_relative = key_positions - prefix_len + key_groups = tl.load( + image_groups + query_start + key_relative, + mask=(key_relative >= 0) & (key_relative < query_len), + other=0, + ) + same_image = (query_groups[:, None] > 0) & ( + query_groups[:, None] == key_groups[None, :] + ) + visible = (key_positions[None, :] <= absolute_queries) | same_image + if WINDOW > 0: + visible &= key_positions[None, :] > absolute_queries - WINDOW + visible &= key_positions[None, :] < max_key + if SCORE_MODE == 3: + score = tl.sum( + tl.where(visible, logits * 0.6931471805599453, 0.0), axis=0 + ) + tl.atomic_add( + attn_score + + batch * stride_asb + + query_head * stride_ash + + key_positions * stride_asl, + score, + mask=key_positions < max_key, + ) + elif SCORE_MODE == 2: + score = ( + tl.sum( + tl.where(visible, logits * 0.6931471805599453, 0.0), axis=0 + ) + / query_len + ) + tl.atomic_max( + attn_score + batch * stride_asb + key_positions * stride_asl, + score, + mask=key_positions < max_key, + ) + logits = tl.where(visible, logits, -float("inf")) + block_max = tl.max(logits, axis=1) + new_max = tl.maximum(max_logit, block_max) + probabilities = tl.exp2(logits - new_max[:, None]) + correction = tl.exp2(max_logit - new_max) + denominator = denominator * correction + tl.sum(probabilities, axis=1) + accumulator *= correction[:, None] + value = tl.load( + v + slots[:, None] * stride_vt + kv_head * stride_vh + dims[None, :], + mask=key_positions[:, None] < max_key, + other=0.0, + ) + accumulator = tl.dot(probabilities.to(value.dtype), value, accumulator) + max_logit = new_max + output_positions = query_start + query_positions + tl.store( + output + + output_positions[:, None] * stride_ot + + query_head * stride_oh + + dims[None, :], + accumulator / denominator[:, None], + mask=query_positions[:, None] < query_len, + ) + + +@torch.no_grad() +def gemma4_multimodal_context_attention( + q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, + output: torch.Tensor, + req_indices: torch.Tensor, + q_start: torch.Tensor, + context_lens: torch.Tensor, + cached_prefix_lens: torch.Tensor, + max_query_len: int, + active_slots: torch.Tensor, + image_groups: torch.Tensor, + *, + sliding_window: int, + attn_score: torch.Tensor | None = None, +) -> None: + head_dim = int(q.shape[-1]) + if q.ndim != 3 or k.ndim != 3 or v.shape != k.shape or output.shape != q.shape: + raise ValueError( + "Gemma 4 multimodal attention requires matching rank-3 Q/K/V/output." + ) + if head_dim not in {256, 512} or k.shape[-1] != head_dim: + raise ValueError( + f"Gemma 4 multimodal attention requires head_dim 256 or 512, got {head_dim}." + ) + if not all(tensor.is_cuda for tensor in (q, k, v, output, image_groups)): + raise TypeError("Gemma 4 multimodal attention requires CUDA tensors.") + if q.dtype not in {torch.float16, torch.bfloat16} or any( + tensor.dtype != q.dtype for tensor in (k, v, output) + ): + raise TypeError( + "Gemma 4 multimodal attention requires matching FP16 or BF16 tensors." + ) + if any(tensor.stride(-1) != 1 for tensor in (q, k, v, output)): + raise ValueError( + "Gemma 4 multimodal attention requires contiguous head dimensions." + ) + if int(q.shape[1]) % int(k.shape[1]): + raise ValueError("Gemma 4 multimodal attention requires divisible Q and KV heads.") + if image_groups.shape != (q.shape[0],): + raise ValueError( + "Gemma 4 image groups must align with flattened queries: " + f"groups={tuple(image_groups.shape)} q={tuple(q.shape)}." + ) + if image_groups.dtype != torch.int32: + raise TypeError("Gemma 4 image groups must use int32.") + if int(sliding_window) <= 0: + raise ValueError(f"Gemma 4 multimodal sliding window must be positive, got {sliding_window}.") + if attn_score is not None and attn_score.dim() not in {2, 3}: + raise ValueError( + "Gemma 4 multimodal prefill scores must be [B, L] or [B, H, L], " + f"got {tuple(attn_score.shape)}." + ) + block_m = 32 if head_dim == 256 else 16 + block_n = block_m + batch, num_heads = int(context_lens.numel()), int(q.shape[1]) + score = context_lens if attn_score is None else attn_score + score_head_stride = score.stride(1) if score.dim() == 3 else 0 + _gemma4_multimodal_context_attention_kernel[ + (triton.cdiv(int(max_query_len), block_m), batch * num_heads) + ]( + q, + k, + v, + output, + q_start, + context_lens, + cached_prefix_lens, + active_slots, + req_indices, + image_groups, + score, + q.stride(0), + q.stride(1), + k.stride(0), + k.stride(1), + v.stride(0), + v.stride(1), + output.stride(0), + output.stride(1), + active_slots.stride(0), + active_slots.stride(1), + score.stride(0), + score_head_stride, + score.stride(-1), + int(q.shape[1]) // int(k.shape[1]), + NUM_HEADS=num_heads, + HEAD_DIM=head_dim, + BLOCK_M=block_m, + BLOCK_N=block_n, + WINDOW=int(sliding_window), + SCORE_MODE=0 if attn_score is None else attn_score.dim(), + num_warps=8, + num_stages=1, + ) + + +__all__ = ["gemma4_multimodal_context_attention"] diff --git a/src/sparsevllm/models/gemma4.py b/src/sparsevllm/models/gemma4.py index 53176db0..230ee7ed 100644 --- a/src/sparsevllm/models/gemma4.py +++ b/src/sparsevllm/models/gemma4.py @@ -509,9 +509,22 @@ def get_per_layer_inputs( ) return (model_inputs + token_inputs) * self.per_layer_input_scale - def forward(self, input_ids: torch.Tensor, positions: torch.Tensor) -> torch.Tensor: - hidden_states = self.embed_tokens(input_ids) * self.embedding_scale - per_layer_inputs = self.get_per_layer_inputs(input_ids, hidden_states) + def forward( + self, + input_ids: torch.Tensor, + positions: torch.Tensor, + inputs_embeds: torch.Tensor | None = None, + multimodal_mask: torch.Tensor | None = None, + ) -> torch.Tensor: + hidden_states = ( + self.embed_tokens(input_ids) * self.embedding_scale + if inputs_embeds is None + else inputs_embeds + ) + per_layer_ids = input_ids + if self.hidden_size_per_layer_input and multimodal_mask is not None: + per_layer_ids = input_ids.masked_fill(multimodal_mask, int(self.config.pad_token_id)) + per_layer_inputs = self.get_per_layer_inputs(per_layer_ids, hidden_states) context = get_context() for layer_idx, layer in enumerate(self.layers): context.now_layer_idx = layer_idx @@ -530,6 +543,7 @@ def forward(self, input_ids: torch.Tensor, positions: torch.Tensor) -> torch.Ten class Gemma4ForCausalLM(nn.Module): special_weight_loaders = (".expert_weight",) + packed_modules_excluded_prefixes = ("multimodal_encoder.",) packed_modules_mapping: ClassVar = { "q_proj": ("qkv_proj", "q"), "k_proj": ("qkv_proj", "k"), @@ -550,6 +564,33 @@ def __init__( if config.tie_word_embeddings: self.lm_head.weight.data = self.model.embed_tokens.weight.data self.logit_softcap = float(config.final_logit_softcapping or 0.0) + self.multimodal_encoder = None + + def configure_multimodal(self, outer_config) -> None: + from sparsevllm.models.gemma4_multimodal import Gemma4MultimodalEncoder + + self.multimodal_encoder = Gemma4MultimodalEncoder(outer_config) + self.multimodal_bidirectional = ( + getattr(outer_config.text_config, "use_bidirectional_attention", None) + == "vision" + ) + + def encode_multimodal(self, input_ids, tensors): + if self.multimodal_encoder is None: + raise RuntimeError("Gemma 4 multimodal encoder is disabled.") + return self.multimodal_encoder.encode(input_ids, tensors) + + def embed_input_ids(self, input_ids: torch.Tensor) -> torch.Tensor: + return self.model.embed_tokens(input_ids) * self.model.embedding_scale + + def forward_multimodal( + self, + input_ids: torch.Tensor, + positions: torch.Tensor, + inputs_embeds: torch.Tensor, + multimodal_mask: torch.Tensor, + ) -> torch.Tensor: + return self.model(input_ids, positions, inputs_embeds, multimodal_mask) @classmethod def build_runtime_kwargs(cls, config, *, device, **_): @@ -591,6 +632,19 @@ def map_weight_name(self, source_weight_name: str) -> str | None: "v_norm", }: return None + multimodal_prefixes = ( + ("model.vision_tower.", "multimodal_encoder.vision_tower."), + ("model.audio_tower.", "multimodal_encoder.audio_tower."), + ("model.embed_vision.", "multimodal_encoder.embed_vision."), + ("model.embed_audio.", "multimodal_encoder.embed_audio."), + ) + if self.multimodal_encoder is None and source_weight_name.startswith( + tuple(prefix for prefix, _ in multimodal_prefixes) + ): + return None + for source_prefix, target_prefix in multimodal_prefixes: + if source_weight_name.startswith(source_prefix): + return target_prefix + source_weight_name[len(source_prefix) :] return ( "model." + source_weight_name[len(prefix) :] if source_weight_name.startswith(prefix) @@ -663,8 +717,14 @@ def warmup_moe(self, num_tokens: int = 1) -> None: experts(hidden, ids, weights) device_runtime.synchronize() - def forward(self, input_ids: torch.Tensor, positions: torch.Tensor) -> torch.Tensor: - return self.model(input_ids, positions) + def forward( + self, + input_ids: torch.Tensor, + positions: torch.Tensor, + inputs_embeds: torch.Tensor | None = None, + multimodal_mask: torch.Tensor | None = None, + ) -> torch.Tensor: + return self.model(input_ids, positions, inputs_embeds, multimodal_mask) def compute_logits(self, hidden_states: torch.Tensor) -> torch.Tensor | None: logits = self.lm_head(hidden_states) diff --git a/src/sparsevllm/models/gemma4_multimodal.py b/src/sparsevllm/models/gemma4_multimodal.py new file mode 100644 index 00000000..31ee2758 --- /dev/null +++ b/src/sparsevllm/models/gemma4_multimodal.py @@ -0,0 +1,90 @@ +from __future__ import annotations + +import torch +from torch import nn +from transformers.models.gemma4.modeling_gemma4 import ( + Gemma4AudioModel, + Gemma4MultimodalEmbedder, + Gemma4VisionModel, +) + +from sparsevllm.multimodal.runtime import MultiModalState + + +class Gemma4MultimodalEncoder(nn.Module): + def __init__(self, config) -> None: + super().__init__() + text_config = config.text_config + self.vision_tower = ( + Gemma4VisionModel(config.vision_config) + if config.vision_config is not None + else None + ) + self.embed_vision = ( + Gemma4MultimodalEmbedder(config.vision_config, text_config) + if config.vision_config is not None + else None + ) + self.audio_tower = ( + Gemma4AudioModel(config.audio_config) + if config.audio_config is not None + else None + ) + self.embed_audio = ( + Gemma4MultimodalEmbedder(config.audio_config, text_config) + if config.audio_config is not None + else None + ) + + def _vision_features( + self, + pixels: torch.Tensor, + position_ids: torch.Tensor, + *, + video: bool, + ) -> torch.Tensor: + if self.vision_tower is None or self.embed_vision is None: + raise ValueError("This Gemma 4 checkpoint has no vision tower.") + device = next(self.vision_tower.parameters()).device + if video: + pixels = pixels.flatten(0, 1) + position_ids = position_ids.flatten(0, 1) + output = self.vision_tower( + pixel_values=pixels.to(device=device, dtype=self.vision_tower.dtype), + pixel_position_ids=position_ids.to(device), + return_dict=True, + ) + return self.embed_vision(output.last_hidden_state) + + @torch.inference_mode() + def encode( + self, + input_ids: list[int], + tensors: dict[str, torch.Tensor], + ) -> MultiModalState: + del input_ids + device = next(self.parameters()).device + type_ids = tensors["mm_token_type_ids"].squeeze(0).to(device) + embeddings = {} + if "pixel_values" in tensors: + embeddings[1] = self._vision_features( + tensors["pixel_values"], tensors["image_position_ids"], video=False + ) + if "pixel_values_videos" in tensors: + embeddings[2] = self._vision_features( + tensors["pixel_values_videos"], tensors["video_position_ids"], video=True + ) + if "input_features" in tensors: + if self.audio_tower is None or self.embed_audio is None: + raise ValueError("This Gemma 4 checkpoint has no audio tower.") + output = self.audio_tower( + tensors["input_features"].to(device=device, dtype=self.audio_tower.dtype), + tensors["input_features_mask"].to(device), + return_dict=True, + ) + features = self.embed_audio(output.last_hidden_state) + embeddings[3] = features[output.attention_mask.to(device=features.device)] + return MultiModalState(type_ids, embeddings, None, 0) + + +__all__ = ["Gemma4MultimodalEncoder"] diff --git a/src/sparsevllm/models/qwen3_5.py b/src/sparsevllm/models/qwen3_5.py index 71a23c4d..a43a7307 100644 --- a/src/sparsevllm/models/qwen3_5.py +++ b/src/sparsevllm/models/qwen3_5.py @@ -20,6 +20,7 @@ resolve_gate_up_swiglu_provider, ) from sparsevllm.layers.rotary_embedding import apply_partial_rotary_emb, get_rope +from sparsevllm.operators.qwen35_mrope import Qwen35MRotaryEmbedding from sparsevllm.layers.embed_head import VocabParallelEmbedding, ParallelLMHead from sparsevllm.utils.context import get_context from sparsevllm.utils.weight_target import WeightTarget @@ -245,12 +246,28 @@ def __init__(self, config) -> None: bias=False, quantization=quantization, ) - self.rotary_emb = get_rope( - self.rotary_dim, - rotary_dim=self.rotary_dim, - max_position=int(config.max_position_embeddings), - base=_get_rope_theta(config), - rope_scaling=_get_rope_scaling(config), + rope_parameters = getattr(config, "rope_parameters", None) + mrope_sections = ( + rope_parameters.get("mrope_section") + if isinstance(rope_parameters, dict) + else None + ) + self.rotary_emb = ( + Qwen35MRotaryEmbedding( + self.head_dim, + self.rotary_dim, + int(config.max_position_embeddings), + _get_rope_theta(config), + mrope_sections, + ) + if mrope_sections + else get_rope( + self.rotary_dim, + rotary_dim=self.rotary_dim, + max_position=int(config.max_position_embeddings), + base=_get_rope_theta(config), + rope_scaling=_get_rope_scaling(config), + ) ) self.attn = Attention(self.num_heads, self.head_dim, self.scaling, self.num_kv_heads) self.q_norm = Qwen35RMSNorm(self.head_dim, eps=float(getattr(config, "rms_norm_eps", 1.0e-6))) @@ -280,12 +297,12 @@ def forward(self, positions: torch.Tensor, hidden_states: torch.Tensor) -> torch cache_manager.save_raw_kv_if_needed(layer_idx, pre_rope_k, v) q = self.q_norm(q) k = self.k_norm(k) - q, k = apply_partial_rotary_emb( - self.rotary_emb, - positions, - q, - k, - self.rotary_dim, + q, k = ( + self.rotary_emb(positions, q, k) + if isinstance(self.rotary_emb, Qwen35MRotaryEmbedding) + else apply_partial_rotary_emb( + self.rotary_emb, positions, q, k, self.rotary_dim + ) ) cache_manager.save_rope_kv_if_needed(layer_idx, k, v) o = self.attn(q, k, v) @@ -928,8 +945,13 @@ def __init__(self, config, layer_cls=Qwen35DecoderLayer) -> None: self.sparse_controller = None self.recurrent_state_manager = None - def forward(self, input_ids: torch.Tensor, positions: torch.Tensor) -> torch.Tensor: - hidden_states = self.embed_tokens(input_ids) + def forward( + self, + input_ids: torch.Tensor, + positions: torch.Tensor, + inputs_embeds: torch.Tensor | None = None, + ) -> torch.Tensor: + hidden_states = self.embed_tokens(input_ids) if inputs_embeds is None else inputs_embeds residual = None context = get_context() debug_layers_env = os.getenv("SPARSEVLLM_DEBUG_HIDDEN_LAYERS") @@ -962,6 +984,7 @@ def forward(self, input_ids: torch.Tensor, positions: torch.Tensor) -> torch.Ten class Qwen35ForCausalLM(nn.Module): ignored_weight_prefixes = ("model.visual.", "visual.", "mtp.") + packed_modules_excluded_prefixes = ("multimodal_encoder.",) special_weight_loaders = { ".linear_attn.in_proj_qkv.weight": "load_packed_in_proj_qkv", ".linear_attn.in_proj_qkvz.weight": "load_packed_in_proj_qkvz", @@ -986,6 +1009,31 @@ def __init__(self, config) -> None: self.lm_head = ParallelLMHead(int(config.vocab_size), int(config.hidden_size)) if bool(getattr(config, "tie_word_embeddings", False)): self.lm_head.weight.data = self.model.embed_tokens.weight.data + self.multimodal_encoder = None + + def configure_multimodal(self, outer_config) -> None: + from sparsevllm.models.qwen3_5_multimodal import Qwen35MultimodalEncoder + + self.multimodal_encoder = Qwen35MultimodalEncoder(outer_config) + self.ignored_weight_prefixes = ("mtp.",) + + def encode_multimodal(self, input_ids, tensors): + if self.multimodal_encoder is None: + raise RuntimeError("Qwen3.5 multimodal encoder is disabled.") + return self.multimodal_encoder.encode(input_ids, tensors) + + def embed_input_ids(self, input_ids: torch.Tensor) -> torch.Tensor: + return self.model.embed_tokens(input_ids) + + def forward_multimodal( + self, + input_ids: torch.Tensor, + positions: torch.Tensor, + inputs_embeds: torch.Tensor, + multimodal_mask: torch.Tensor, + ) -> torch.Tensor: + del multimodal_mask + return self.model(input_ids, positions, inputs_embeds) @staticmethod def recurrent_state_spec(config, world_size: int) -> RecurrentStateSpec: @@ -1015,6 +1063,10 @@ def map_weight_name(self, source_weight_name: str) -> str: prefix = "model.language_model." if source_weight_name.startswith(prefix): return "model." + source_weight_name[len(prefix) :] + if source_weight_name.startswith("model.visual."): + return "multimodal_encoder.visual." + source_weight_name[len("model.visual.") :] + if source_weight_name.startswith("visual."): + return "multimodal_encoder.visual." + source_weight_name[len("visual.") :] return source_weight_name def resolve_special_weight( @@ -1046,8 +1098,13 @@ def load_special_weight( ) return int(loader(loaded_weight, loaded_scale)) - def forward(self, input_ids: torch.Tensor, positions: torch.Tensor) -> torch.Tensor: - return self.model(input_ids, positions) + def forward( + self, + input_ids: torch.Tensor, + positions: torch.Tensor, + inputs_embeds: torch.Tensor | None = None, + ) -> torch.Tensor: + return self.model(input_ids, positions, inputs_embeds) def compute_logits(self, hidden_states: torch.Tensor) -> torch.Tensor: return self.lm_head(hidden_states) diff --git a/src/sparsevllm/models/qwen3_5_moe.py b/src/sparsevllm/models/qwen3_5_moe.py index f0edae0e..530e66d9 100644 --- a/src/sparsevllm/models/qwen3_5_moe.py +++ b/src/sparsevllm/models/qwen3_5_moe.py @@ -228,15 +228,24 @@ def load_packed_expert_weight( self._loaded_packed_projections.add(projection) def validate_loaded_weights(self) -> None: - if self.fp8_enabled: - super().validate_loaded_weights() - return - missing = {"gate_up_proj", "down_proj"} - self._loaded_packed_projections + if not self.fp8_enabled: + missing = {"gate_up_proj", "down_proj"} - self._loaded_packed_projections + if missing: + raise ValueError( + f"Missing Qwen3.6 packed expert projections: {sorted(missing)}." + ) + expected = { + (expert_id, projection) + for expert_id in range(self.local_expert_start, self.local_expert_end) + for projection in ("gate_proj", "up_proj", "down_proj") + } + missing = sorted(expected - self._loaded_expert_shards) if missing: raise ValueError( - f"Missing Qwen3.6 packed expert projections: {sorted(missing)}." + "Missing local Qwen3.6 expert weights: " + f"local_range=[{self.local_expert_start}, " + f"{self.local_expert_end}), missing={missing[:8]}." ) - super().validate_loaded_weights() class Qwen35MoeSparseMoeBlock(nn.Module): @@ -311,6 +320,7 @@ def __init__(self, config) -> None: ) if bool(getattr(config, "tie_word_embeddings", False)): self.lm_head.weight.data = self.model.embed_tokens.weight.data + self.multimodal_encoder = None self._loaded_linear_special_weights: set[str] = set() self._intentionally_skipped_weights: set[str] = set() self._intentionally_skipped_expert_weights: set[str] = set() @@ -647,7 +657,12 @@ def validate_loaded_weights(self, loaded_parameter_names: set[str]) -> None: for name in self._intentionally_skipped_weights ), } - missing_skip_groups = [name for name, seen in skip_groups.items() if not seen] + required_skip_groups = ( + {"mtp"} if self.multimodal_encoder is not None else {"visual", "mtp"} + ) + missing_skip_groups = [ + name for name in required_skip_groups if not skip_groups[name] + ] if missing_skip_groups: raise ValueError( "Qwen3.6 MoE checkpoint is missing expected intentional-skip " diff --git a/src/sparsevllm/models/qwen3_5_multimodal.py b/src/sparsevllm/models/qwen3_5_multimodal.py new file mode 100644 index 00000000..57f72616 --- /dev/null +++ b/src/sparsevllm/models/qwen3_5_multimodal.py @@ -0,0 +1,110 @@ +from __future__ import annotations + +import itertools + +import torch +from torch import nn +from transformers.models.qwen3_5.modeling_qwen3_5 import Qwen3_5VisionModel + +from sparsevllm.multimodal.runtime import MultiModalState + + +def _vision_positions( + start: int, + grid: torch.Tensor, + spatial_merge_size: int, +) -> torch.Tensor: + t, h, w = [int(value) for value in grid.tolist()] + h //= spatial_merge_size + w //= spatial_merge_size + temporal = torch.arange(t, device=grid.device) + start + height = torch.arange(h, device=grid.device) + start + width = torch.arange(w, device=grid.device) + start + values = torch.meshgrid(temporal, height, width, indexing="ij") + return torch.stack(values).reshape(3, -1) + + +def qwen35_mrope_positions( + type_ids: torch.Tensor, + image_grid_thw: torch.Tensor | None, + video_grid_thw: torch.Tensor | None, + spatial_merge_size: int, +) -> tuple[torch.Tensor, int]: + if video_grid_thw is not None: + video_grid_thw = torch.repeat_interleave( + video_grid_thw, video_grid_thw[:, 0], dim=0 + ).clone() + video_grid_thw[:, 0] = 1 + grids = { + 1: iter(image_grid_thw) if image_grid_thw is not None else None, + 2: iter(video_grid_thw) if video_grid_thw is not None else None, + } + current = 0 + segments = [] + for modality, group in itertools.groupby( + enumerate(type_ids.tolist()), lambda item: item[1] + ): + group = list(group) + length = group[-1][0] - group[0][0] + 1 + if modality == 0: + segment = torch.arange(length, device=type_ids.device).expand(3, -1) + current + current += length + else: + grid_iter = grids.get(int(modality)) + if grid_iter is None: + raise ValueError(f"Missing Qwen3.5 grid for modality={modality}.") + grid = next(grid_iter) + segment = _vision_positions(current, grid, spatial_merge_size) + current += max(int(grid[1]), int(grid[2])) // spatial_merge_size + if int(segment.shape[1]) != length: + raise ValueError( + "Qwen3.5 M-RoPE span/grid mismatch: " + f"modality={modality} span={length} positions={segment.shape[1]}." + ) + segments.append(segment) + positions = torch.cat(segments, dim=1) + return positions, int(positions.max().item() + 1 - type_ids.numel()) + + +class Qwen35MultimodalEncoder(nn.Module): + def __init__(self, config) -> None: + super().__init__() + self.visual = Qwen3_5VisionModel(config.vision_config) + self.spatial_merge_size = int(config.vision_config.spatial_merge_size) + + @torch.inference_mode() + def encode( + self, + input_ids: list[int], + tensors: dict[str, torch.Tensor], + ) -> MultiModalState: + device = next(self.parameters()).device + type_ids = tensors["mm_token_type_ids"].squeeze(0).to(device) + embeddings = {} + inputs = ( + (1, "pixel_values", "image_grid_thw"), + (2, "pixel_values_videos", "video_grid_thw"), + ) + for modality, pixels_name, grid_name in inputs: + pixels = tensors.get(pixels_name) + if pixels is None: + continue + grid = tensors.get(grid_name) + if grid is None: + raise ValueError(f"{pixels_name} requires {grid_name}.") + output = self.visual( + pixels.to(device=device, dtype=self.visual.dtype), + grid_thw=grid.to(device), + return_dict=True, + ) + embeddings[modality] = output.pooler_output + positions, delta = qwen35_mrope_positions( + type_ids, + None if tensors.get("image_grid_thw") is None else tensors["image_grid_thw"].to(device), + None if tensors.get("video_grid_thw") is None else tensors["video_grid_thw"].to(device), + self.spatial_merge_size, + ) + return MultiModalState(type_ids, embeddings, positions, delta) + + +__all__ = ["Qwen35MultimodalEncoder", "qwen35_mrope_positions"] diff --git a/src/sparsevllm/multimodal/__init__.py b/src/sparsevllm/multimodal/__init__.py new file mode 100644 index 00000000..71b645bd --- /dev/null +++ b/src/sparsevllm/multimodal/__init__.py @@ -0,0 +1,3 @@ +from sparsevllm.multimodal.inputs import MultiModalPrompt + +__all__ = ["MultiModalPrompt"] diff --git a/src/sparsevllm/multimodal/inputs.py b/src/sparsevllm/multimodal/inputs.py new file mode 100644 index 00000000..5b5351da --- /dev/null +++ b/src/sparsevllm/multimodal/inputs.py @@ -0,0 +1,171 @@ +from __future__ import annotations + +import base64 +import hashlib +import io +import wave +from dataclasses import dataclass +from typing import Any + +import numpy as np +import torch +from transformers import AutoProcessor + + +@dataclass(frozen=True) +class MultiModalPrompt: + messages: list[dict[str, Any]] + chat_template_kwargs: dict[str, Any] | None = None + tools: list[dict[str, Any]] | None = None + add_generation_prompt: bool = True + + +@dataclass(frozen=True) +class ProcessedMultiModalPrompt: + token_ids: list[int] + tensors: dict[str, torch.Tensor] + digest: str + + +def is_multimodal_prompt(prompt: object) -> bool: + return isinstance(prompt, MultiModalPrompt) or ( + isinstance(prompt, dict) and "messages" in prompt + ) + + +def _audio_part(part: dict[str, Any]) -> dict[str, Any]: + audio = part.get("input_audio") + if not isinstance(audio, dict) or not isinstance(audio.get("data"), str): + raise TypeError("input_audio requires a base64 data string.") + if str(audio.get("format", "wav")).lower() != "wav": + raise ValueError("Only WAV input_audio is supported.") + with wave.open(io.BytesIO(base64.b64decode(audio["data"])), "rb") as wav: + if wav.getcomptype() != "NONE": + raise ValueError("Compressed WAV input_audio is unsupported.") + channels, sample_width, sampling_rate = ( + wav.getnchannels(), + wav.getsampwidth(), + wav.getframerate(), + ) + raw = wav.readframes(wav.getnframes()) + if sample_width == 1: + waveform = (np.frombuffer(raw, np.uint8).astype(np.float32) - 128) / 128 + elif sample_width in {2, 4}: + dtype = np.dtype(f" 1: + waveform = waveform.reshape(-1, channels).mean(axis=1) + return { + "type": "audio", + "audio": np.asarray(waveform, dtype=np.float32), + "sampling_rate": int(sampling_rate), + } + + +def normalize_messages(messages: list[dict[str, Any]]) -> list[dict[str, Any]]: + normalized = [] + for message in messages: + if not isinstance(message, dict): + raise TypeError("Multimodal messages must be dictionaries.") + content = message.get("content") + if not isinstance(content, list): + normalized.append(dict(message)) + continue + parts = [] + for raw_part in content: + part = dict(raw_part) + part_type = part.get("type") + if part_type in {"image_url", "input_image"}: + value = part.get("image_url") + url = value.get("url") if isinstance(value, dict) else value + if not isinstance(url, str): + raise TypeError("image_url requires a URL string.") + part = {"type": "image", "image": url} + elif part_type in {"video_url", "input_video"}: + value = part.get("video_url") + url = value.get("url") if isinstance(value, dict) else value + if not isinstance(url, str): + raise TypeError("video_url requires a URL string.") + part = {"type": "video", "video": url} + elif part_type == "input_audio": + part = _audio_part(part) + parts.append(part) + normalized.append({**message, "content": parts}) + return normalized + + +def _tensor_digest(tensors: dict[str, torch.Tensor]) -> str: + digest = hashlib.sha256() + for name, tensor in sorted(tensors.items()): + value = tensor.detach().cpu().contiguous() + digest.update(name.encode()) + digest.update(str(value.dtype).encode()) + digest.update(np.asarray(value.shape, dtype=np.int64).tobytes()) + digest.update(value.view(torch.uint8).numpy().tobytes()) + return digest.hexdigest() + + +class MultiModalInputProcessor: + def __init__(self, model_path: str) -> None: + self.processor = AutoProcessor.from_pretrained( + model_path, trust_remote_code=True + ) + + def process( + self, prompt: MultiModalPrompt | dict[str, Any] + ) -> ProcessedMultiModalPrompt: + messages = prompt.messages if isinstance(prompt, MultiModalPrompt) else prompt.get("messages") + if not isinstance(messages, list) or not messages: + raise ValueError("A multimodal prompt requires non-empty messages.") + outputs = self.processor.apply_chat_template( + normalize_messages(messages), + tokenize=True, + add_generation_prompt=( + prompt.add_generation_prompt + if isinstance(prompt, MultiModalPrompt) + else True + ), + return_dict=True, + return_tensors="pt", + **( + { + **(prompt.chat_template_kwargs or {}), + **({"tools": prompt.tools} if prompt.tools else {}), + } + if isinstance(prompt, MultiModalPrompt) + else {} + ), + ) + input_ids = outputs.pop("input_ids") + outputs.pop("attention_mask", None) + if input_ids.ndim != 2 or input_ids.shape[0] != 1: + raise ValueError( + f"Multimodal processor must return one input sequence, got {tuple(input_ids.shape)}." + ) + tensors = { + str(name): value.detach().cpu().contiguous() + for name, value in outputs.items() + if isinstance(value, torch.Tensor) + } + if not tensors or "mm_token_type_ids" not in tensors: + raise ValueError("The processor did not return multimodal tensors.") + return ProcessedMultiModalPrompt( + token_ids=[int(token_id) for token_id in input_ids[0].tolist()], + tensors=tensors, + digest=_tensor_digest(tensors), + ) + + +__all__ = [ + "MultiModalInputProcessor", + "MultiModalPrompt", + "ProcessedMultiModalPrompt", + "is_multimodal_prompt", + "normalize_messages", +] diff --git a/src/sparsevllm/multimodal/runtime.py b/src/sparsevllm/multimodal/runtime.py new file mode 100644 index 00000000..a9a16821 --- /dev/null +++ b/src/sparsevllm/multimodal/runtime.py @@ -0,0 +1,158 @@ +from __future__ import annotations + +from dataclasses import dataclass + +import torch + + +@dataclass +class MultiModalState: + type_ids: torch.Tensor + embeddings: dict[int, torch.Tensor] + position_ids: torch.Tensor | None + position_delta: int + + +class MultiModalRuntime: + """Rank-local encoder feature cache keyed by the engine sequence id.""" + + def __init__(self, model, device: torch.device) -> None: + self.model = model + self.device = device + self.states: dict[int, MultiModalState] = {} + + @property + def enabled(self) -> bool: + return callable(getattr(self.model, "encode_multimodal", None)) + + def register( + self, + seq_id: int, + input_ids: list[int], + tensors: dict[str, torch.Tensor], + ) -> int: + if not self.enabled: + raise NotImplementedError( + f"{type(self.model).__name__} has no native multimodal encoder." + ) + seq_id = int(seq_id) + if seq_id in self.states: + raise RuntimeError(f"Multimodal state already exists for seq_id={seq_id}.") + encoded = self.model.encode_multimodal(input_ids, tensors) + if not isinstance(encoded, MultiModalState): + raise TypeError( + "encode_multimodal() must return MultiModalState, " + f"got {type(encoded).__name__}." + ) + if encoded.type_ids.ndim != 1 or encoded.type_ids.numel() != len(input_ids): + raise ValueError( + "Multimodal token types must align with the prompt: " + f"types={tuple(encoded.type_ids.shape)} tokens={len(input_ids)}." + ) + for modality, features in encoded.embeddings.items(): + expected = int(encoded.type_ids.eq(int(modality)).sum().item()) + if features.ndim != 2 or int(features.shape[0]) != expected: + raise ValueError( + "Multimodal feature/token mismatch: " + f"modality={modality} features={tuple(features.shape)} tokens={expected}." + ) + self.states[seq_id] = encoded + return int(encoded.position_delta) + + def free(self, seq_id: int) -> None: + self.states.pop(int(seq_id), None) + + def free_batch(self, seq_ids: list[int]) -> None: + for seq_id in seq_ids: + self.free(seq_id) + + def forward( + self, + input_ids: torch.Tensor, + positions: torch.Tensor, + inputs_embeds: torch.Tensor, + multimodal_mask: torch.Tensor, + ) -> torch.Tensor: + forward = getattr(self.model, "forward_multimodal", None) + if not callable(forward): + raise NotImplementedError( + f"{type(self.model).__name__} has no multimodal forward path." + ) + return forward(input_ids, positions, inputs_embeds, multimodal_mask) + + def prepare( + self, + seqs, + input_ids: torch.Tensor, + positions: torch.Tensor, + is_prefill: bool, + ) -> tuple[torch.Tensor | None, torch.Tensor, torch.Tensor | None]: + if not is_prefill or not any(int(seq.seq_id) in self.states for seq in seqs): + return None, positions, None + + inputs_embeds = self.model.embed_input_ids(input_ids) + multimodal_mask = torch.zeros( + input_ids.shape, dtype=torch.bool, device=input_ids.device + ) + position_rows = [] + batch_offset = 0 + image_groups = torch.zeros_like(input_ids, dtype=torch.int32) + next_image_group = 1 + for seq in seqs: + chunk_size = int(seq.current_chunk_size) + start = int(seq.num_prefilled_tokens) + end = start + chunk_size + state = self.states.get(int(seq.seq_id)) + if state is None: + position_rows.append( + positions[batch_offset : batch_offset + chunk_size].expand(3, -1) + ) + batch_offset += chunk_size + continue + + chunk_types = state.type_ids[start:end] + position_rows.append( + state.position_ids[:, start:end] + if state.position_ids is not None + else positions[batch_offset : batch_offset + chunk_size].expand(3, -1) + ) + for modality, features in state.embeddings.items(): + prompt_indices = state.type_ids.eq(int(modality)).nonzero().flatten() + selected = (prompt_indices >= start) & (prompt_indices < end) + if not selected.any(): + continue + chunk_indices = prompt_indices[selected] - start + batch_offset + feature_start = int((prompt_indices < start).sum().item()) + feature_end = feature_start + int(selected.sum().item()) + inputs_embeds[chunk_indices] = features[feature_start:feature_end] + multimodal_mask[chunk_indices] = True + + image_positions = ((chunk_types == 1) | (chunk_types == 2)).nonzero().flatten() + if image_positions.numel(): + split = torch.where(image_positions[1:] != image_positions[:-1] + 1)[0] + 1 + for group in torch.tensor_split(image_positions, split.cpu().tolist()): + image_groups[batch_offset + group] = next_image_group + next_image_group += 1 + batch_offset += chunk_size + + from sparsevllm.utils.context import get_context + + get_context().multimodal_image_groups = ( + image_groups + if next_image_group > 1 + and bool(getattr(self.model, "multimodal_bidirectional", False)) + else None + ) + use_mrope = any( + state.position_ids is not None + for seq in seqs + if (state := self.states.get(int(seq.seq_id))) is not None + ) + return ( + inputs_embeds, + torch.cat(position_rows, dim=1) if use_mrope else positions, + multimodal_mask, + ) + + +__all__ = ["MultiModalRuntime", "MultiModalState"] diff --git a/src/sparsevllm/operators/gemma4_attention.py b/src/sparsevllm/operators/gemma4_attention.py index 39e628c2..8736c067 100644 --- a/src/sparsevllm/operators/gemma4_attention.py +++ b/src/sparsevllm/operators/gemma4_attention.py @@ -28,6 +28,30 @@ def run_prefill( ) -> torch.Tensor: payload = _require_explicit_payload(view, operation="Gemma 4 prefill") output = torch.empty_like(q) + from sparsevllm.utils.context import get_context + + image_groups = getattr(get_context(), "multimodal_image_groups", None) + if self.sliding_window is not None and isinstance(image_groups, torch.Tensor): + from sparsevllm.kernels.triton.gemma4_multimodal_context_attention import ( + gemma4_multimodal_context_attention, + ) + + gemma4_multimodal_context_attention( + q, + payload.k_cache, + payload.v_cache, + output, + view.meta.req_indices, + b_start_loc, + view.meta.context_lens, + view.meta.context_lens - chunk_lens, + max_input_len, + view.meta.active_slots, + image_groups, + sliding_window=self.sliding_window, + attn_score=view.meta.attn_score, + ) + return output from sparsevllm.kernels.triton.gemma4_context_attention import ( gemma4_context_attention, ) diff --git a/src/sparsevllm/operators/qwen35_mrope.py b/src/sparsevllm/operators/qwen35_mrope.py new file mode 100644 index 00000000..0b5dde9a --- /dev/null +++ b/src/sparsevllm/operators/qwen35_mrope.py @@ -0,0 +1,77 @@ +from __future__ import annotations + +import torch +from torch import nn + +from sparsevllm.layers.rotary_embedding import ( + apply_rotary_emb, + get_rope, +) + + +class Qwen35MRotaryEmbedding(nn.Module): + """Dedicated Qwen3.5 M-RoPE path; 1-D text decode keeps FlashInfer RoPE.""" + + def __init__( + self, + head_dim: int, + rotary_dim: int, + max_position: int, + base: float, + sections: list[int], + ) -> None: + super().__init__() + if sum(sections) != rotary_dim // 2: + raise ValueError( + f"Qwen3.5 M-RoPE sections must sum to {rotary_dim // 2}, got {sections}." + ) + self.rotary_dim = int(rotary_dim) + self.sections = tuple(int(section) for section in sections) + self.text_rope = get_rope( + rotary_dim, + rotary_dim=rotary_dim, + max_position=max_position, + base=base, + ) + + def _multimodal_cos_sin( + self, positions: torch.Tensor + ) -> tuple[torch.Tensor, torch.Tensor]: + cache = self.text_rope.cos_sin_cache[positions] + cos, sin = cache.chunk(2, dim=-1) + merged_cos = cos[0].clone() + merged_sin = sin[0].clone() + h_end = self.sections[1] * 3 + w_end = self.sections[2] * 3 + merged_cos[..., 1:h_end:3] = cos[1, ..., 1:h_end:3] + merged_sin[..., 1:h_end:3] = sin[1, ..., 1:h_end:3] + merged_cos[..., 2:w_end:3] = cos[2, ..., 2:w_end:3] + merged_sin[..., 2:w_end:3] = sin[2, ..., 2:w_end:3] + return merged_cos, merged_sin + + def forward( + self, + positions: torch.Tensor, + query: torch.Tensor, + key: torch.Tensor, + ) -> tuple[torch.Tensor, torch.Tensor]: + if positions.ndim == 1: + if self.rotary_dim == query.shape[-1] == key.shape[-1]: + return self.text_rope(positions, query, key) + cos_sin = self.text_rope.cos_sin_cache[positions] + cos, sin = cos_sin.chunk(2, dim=-1) + elif positions.ndim == 2 and positions.shape[0] == 3: + cos, sin = self._multimodal_cos_sin(positions) + else: + raise ValueError( + f"Qwen3.5 positions must be [tokens] or [3, tokens], got {tuple(positions.shape)}." + ) + query_rot = apply_rotary_emb(query[..., : self.rotary_dim], cos, sin) + key_rot = apply_rotary_emb(key[..., : self.rotary_dim], cos, sin) + return ( + torch.cat((query_rot, query[..., self.rotary_dim :]), dim=-1), + torch.cat((key_rot, key[..., self.rotary_dim :]), dim=-1), + ) + + +__all__ = ["Qwen35MRotaryEmbedding"] diff --git a/src/sparsevllm/utils/context.py b/src/sparsevllm/utils/context.py index a53d850c..6385b23d 100644 --- a/src/sparsevllm/utils/context.py +++ b/src/sparsevllm/utils/context.py @@ -12,6 +12,7 @@ def __init__(self): self.seqs = None self.decode_mid_o = None self.decode_mid_o_logexpsum = None + self.multimodal_image_groups = None _CONTEXT = Context() @@ -38,6 +39,7 @@ def set_context( _CONTEXT.cache_manager = cache_manager _CONTEXT.recurrent_state_manager = recurrent_state_manager _CONTEXT.seqs = seqs + _CONTEXT.multimodal_image_groups = None def reset_context(): global _CONTEXT diff --git a/src/sparsevllm/utils/loader.py b/src/sparsevllm/utils/loader.py index 67ae5ef8..3de3fca0 100644 --- a/src/sparsevllm/utils/loader.py +++ b/src/sparsevllm/utils/loader.py @@ -19,6 +19,15 @@ def default_weight_loader(param: nn.Parameter, loaded_weight: torch.Tensor): param.data.copy_(loaded_weight) +def _packed_mapping_for(model: nn.Module, target_parameter_name: str) -> dict: + excluded = tuple(getattr(model, "packed_modules_excluded_prefixes", ())) + return ( + {} + if target_parameter_name.startswith(excluded) + else getattr(model, "packed_modules_mapping", {}) + ) + + @dataclass(frozen=True) class TensorMetadata: shape: tuple[int, ...] @@ -52,7 +61,7 @@ def _resolve_weight_target( return target loaded_shard_id = None - packed_modules_mapping = getattr(model, "packed_modules_mapping", {}) + packed_modules_mapping = _packed_mapping_for(model, target_parameter_name) for source_fragment, (target_fragment, shard_id) in packed_modules_mapping.items(): if source_fragment in target_parameter_name: target_parameter_name = target_parameter_name.replace( @@ -626,7 +635,6 @@ def load_model( num_threads = int(num_threads) if num_threads <= 0: raise ValueError(f"num_threads must be positive, got {num_threads}.") - packed_modules_mapping = getattr(model, "packed_modules_mapping", {}) files = sorted(glob(os.path.join(path, "*.safetensors"))) assert len(files) > 0, f"No safetensors found in {path}" checkpoint_is_rank_local = False @@ -752,9 +760,9 @@ def load_model( consumed_scale_keys.add(scale_key) loaded_count += special_count continue - for k in packed_modules_mapping: + packed_modules_mapping = _packed_mapping_for(model, param_name) + for k, (v, shard_id) in packed_modules_mapping.items(): if k in param_name: - v, shard_id = packed_modules_mapping[k] packed_param_name = param_name.replace(k, v) module = _module_for_parameter(model, packed_param_name) if _load_grouped_quantized_weight( @@ -775,6 +783,25 @@ def load_model( loaded_count += 1 break else: + try: + buffer = model.get_buffer(param_name) + except AttributeError: + buffer = None + if buffer is not None: + loaded_weight = tensors[source_weight_name] + if loaded_scale is not None: + raise ValueError( + f"Buffer {param_name!r} cannot load a quantization scale." + ) + if buffer.shape != loaded_weight.shape: + raise ValueError( + f"Buffer {param_name!r} shape mismatch: " + f"expected {tuple(buffer.shape)}, got {tuple(loaded_weight.shape)}." + ) + buffer.copy_(loaded_weight) + loaded_parameter_names.add(param_name) + loaded_count += 1 + continue module = _module_for_parameter(model, param_name) loaded_scale = None if source_weight_name.endswith(".weight"): diff --git a/tests/test_gemma4_model.py b/tests/test_gemma4_model.py index 52f64566..624b7c5f 100644 --- a/tests/test_gemma4_model.py +++ b/tests/test_gemma4_model.py @@ -26,6 +26,7 @@ from sparsevllm.method_registry import MODEL_RUNTIME_COMPATIBILITY from sparsevllm.models.gemma4 import ( Gemma4Attention, + Gemma4ForCausalLM, Gemma4MLP, Gemma4Model, Gemma4RotaryEmbedding, @@ -107,6 +108,14 @@ def _patch_parallel_context(): return stack +def test_gemma4_skips_multimodal_weights_when_disabled(): + model = SimpleNamespace(multimodal_encoder=None) + + assert Gemma4ForCausalLM.map_weight_name( + model, "model.vision_tower.encoder.layers.0.weight" + ) is None + + def _config(**overrides) -> Gemma4TextConfig: values = { "vocab_size": 32, diff --git a/tests/test_multimodal.py b/tests/test_multimodal.py new file mode 100644 index 00000000..f4832f80 --- /dev/null +++ b/tests/test_multimodal.py @@ -0,0 +1,275 @@ +import base64 +import io +import pickle +import wave +from types import SimpleNamespace + +import pytest +import torch + +from sparsevllm.engine.sequence import Sequence +from sparsevllm.engine.llm_engine import LLMEngine +from sparsevllm.multimodal.inputs import MultiModalInputProcessor, normalize_messages +from sparsevllm.multimodal.runtime import MultiModalRuntime, MultiModalState +from sparsevllm.models.qwen3_5_multimodal import qwen35_mrope_positions +from sparsevllm.operators.qwen35_mrope import Qwen35MRotaryEmbedding +from sparsevllm.sampling_params import SamplingParams +from sparsevllm.utils.context import get_context, set_context + + +def test_normalize_openai_multimodal_parts(): + messages = [ + { + "role": "user", + "content": [ + {"type": "text", "text": "describe"}, + {"type": "image_url", "image_url": {"url": "https://x/image.png"}}, + {"type": "video_url", "video_url": "https://x/video.mp4"}, + ], + } + ] + assert normalize_messages(messages)[0]["content"] == [ + {"type": "text", "text": "describe"}, + {"type": "image", "image": "https://x/image.png"}, + {"type": "video", "video": "https://x/video.mp4"}, + ] + + +def test_normalize_openai_wav_audio_without_optional_dependencies(): + buffer = io.BytesIO() + with wave.open(buffer, "wb") as wav: + wav.setnchannels(1) + wav.setsampwidth(2) + wav.setframerate(16_000) + wav.writeframes(torch.tensor([-32768, 0, 32767], dtype=torch.int16).numpy().tobytes()) + + part = normalize_messages( + [ + { + "role": "user", + "content": [ + { + "type": "input_audio", + "input_audio": { + "format": "wav", + "data": base64.b64encode(buffer.getvalue()).decode(), + }, + } + ], + } + ] + )[0]["content"][0] + + assert part["type"] == "audio" and part["sampling_rate"] == 16_000 + torch.testing.assert_close( + torch.from_numpy(part["audio"]), + torch.tensor([-1.0, 0.0, 32767 / 32768]), + ) + +def test_multimodal_processor_returns_stable_cpu_payload(): + class Processor: + def apply_chat_template(self, messages, **kwargs): + assert messages[0]["content"][1]["type"] == "image" + assert kwargs["tokenize"] and kwargs["return_dict"] + return { + "input_ids": torch.tensor([[7, 8, 9]]), + "attention_mask": torch.ones(1, 3), + "mm_token_type_ids": torch.tensor([[0, 1, 0]]), + "pixel_values": torch.arange(6, dtype=torch.float32).reshape(1, 2, 3), + } + + processor = object.__new__(MultiModalInputProcessor) + processor.processor = Processor() + prompt = { + "messages": [ + { + "role": "user", + "content": [ + {"type": "text", "text": "x"}, + {"type": "image_url", "image_url": "https://x/image.png"}, + ], + } + ] + } + + first = processor.process(prompt) + second = processor.process(prompt) + + assert first.token_ids == [7, 8, 9] + assert first.digest == second.digest + assert set(first.tensors) == {"mm_token_type_ids", "pixel_values"} + assert all(tensor.device.type == "cpu" and tensor.is_contiguous() for tensor in first.tensors.values()) + + +def test_qwen35_mrope_positions_match_image_grid(): + positions, delta = qwen35_mrope_positions( + torch.tensor([0, 0, 1, 1, 1, 1, 0]), + torch.tensor([[1, 4, 4]]), + None, + spatial_merge_size=2, + ) + + assert positions.tolist() == [ + [0, 1, 2, 2, 2, 2, 4], + [0, 1, 2, 2, 3, 3, 4], + [0, 1, 2, 3, 2, 3, 4], + ] + assert delta == -2 + + +def test_qwen35_multimodal_rope_matches_transformers(): + from transformers.models.qwen3_5.modeling_qwen3_5 import apply_rotary_pos_emb + + torch.manual_seed(7) + positions = torch.tensor( + [[0, 1, 2, 3], [0, 1, 4, 5], [0, 1, 6, 7]], dtype=torch.long + ) + query = torch.randn(4, 2, 128) + key = torch.randn(4, 1, 128) + rope = Qwen35MRotaryEmbedding(128, 64, 32, 10_000, [11, 11, 10]) + + actual_q, actual_k = rope(positions, query, key) + inv_freq = 1 / 10_000 ** (torch.arange(0, 64, 2).float() / 64) + freqs = positions[:, :, None].float() * inv_freq + merged = freqs[0].clone() + merged[:, 1:33:3] = freqs[1, :, 1:33:3] + merged[:, 2:30:3] = freqs[2, :, 2:30:3] + cos, sin = torch.cat((merged, merged), -1).cos(), torch.cat((merged, merged), -1).sin() + expected_q, expected_k = apply_rotary_pos_emb( + query.unsqueeze(0), key.unsqueeze(0), cos.unsqueeze(0), sin.unsqueeze(0), unsqueeze_dim=2 + ) + + torch.testing.assert_close(actual_q, expected_q[0]) + torch.testing.assert_close(actual_k, expected_k[0]) + + +def test_multimodal_runtime_replaces_chunk_features_and_tracks_vision_groups(): + type_ids = torch.tensor([0, 1, 1, 2, 2, 0]) + state = MultiModalState( + type_ids=type_ids, + embeddings={ + 1: torch.tensor([[10.0, 11.0], [12.0, 13.0]]), + 2: torch.tensor([[20.0, 21.0], [22.0, 23.0]]), + }, + position_ids=torch.arange(18).reshape(3, 6), + position_delta=-1, + ) + + class Model: + multimodal_bidirectional = True + + def encode_multimodal(self, input_ids, tensors): + assert input_ids == list(range(6)) and tensors == {} + return state + + def embed_input_ids(self, input_ids): + return input_ids.float().unsqueeze(1).expand(-1, 2).clone() + + runtime = MultiModalRuntime(Model(), torch.device("cpu")) + assert runtime.register(3, list(range(6)), {}) == -1 + seq = Sequence(list(range(6)), SamplingParams(max_tokens=1)) + seq.seq_id = 3 + seq.current_chunk_size = 6 + set_context(True) + + embeds, positions, mask = runtime.prepare( + [seq], torch.arange(6), torch.arange(6), is_prefill=True + ) + + assert embeds.tolist() == [ + [0.0, 0.0], + [10.0, 11.0], + [12.0, 13.0], + [20.0, 21.0], + [22.0, 23.0], + [5.0, 5.0], + ] + assert positions.equal(state.position_ids) + assert mask.tolist() == [False, True, True, True, True, False] + assert get_context().multimodal_image_groups.tolist() == [0, 1, 1, 1, 1, 0] + runtime.free(3) + assert runtime.states == {} + + +def test_multimodal_sequence_state_preserves_decode_delta(): + seq = Sequence([1, 2, 3], SamplingParams(max_tokens=2)) + seq.multimodal_digest = "digest" + seq.multimodal_position_delta = -4 + seq.multimodal_full_prefill = True + restored = pickle.loads(pickle.dumps(seq)) + + assert restored.multimodal_digest == "digest" + assert restored.multimodal_full_prefill + assert restored.decode_input_position == -2 + + +def test_abort_queued_multimodal_request_releases_encoder_state_only(): + seq = Sequence([1], SamplingParams(max_tokens=1)) + seq.multimodal_digest = "digest" + calls = [] + + class Scheduler: + waiting = [seq] + decoding = [] + + def abort(self, seq_id): + self.waiting.clear() + return False + + engine = object.__new__(LLMEngine) + engine.scheduler = Scheduler() + engine._active_chain_sequences = {} + engine.model_runner = SimpleNamespace( + runtime_state=SimpleNamespace(chain_cache_coordinator=None), + call=lambda method, *args: calls.append((method, args)), + ) + + engine.abort_request(seq.seq_id) + + assert calls == [("free_multimodal", (seq.seq_id,))] + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="requires CUDA") +def test_gemma4_multimodal_context_attention_matches_reference(): + from sparsevllm.kernels.triton.gemma4_multimodal_context_attention import ( + gemma4_multimodal_context_attention, + ) + + device = torch.device("cuda") + torch.manual_seed(0) + length, num_heads, head_dim, window = 9, 2, 256, 4 + q = (torch.randn(length, num_heads, head_dim, device=device) / head_dim**0.5).bfloat16() + k = torch.randn(length, 1, head_dim, device=device).bfloat16() + v = torch.randn_like(k) + output = torch.empty_like(q) + attention_score = torch.zeros( + 1, num_heads, length, device=device, dtype=torch.float32 + ) + groups = torch.tensor([0, 0, 1, 1, 1, 1, 0, 0, 0], device=device, dtype=torch.int32) + gemma4_multimodal_context_attention( + q, + k, + v, + output, + torch.tensor([0], device=device, dtype=torch.int32), + torch.tensor([0], device=device, dtype=torch.int32), + torch.tensor([length], device=device, dtype=torch.int32), + torch.tensor([0], device=device, dtype=torch.int32), + length, + torch.arange(length, device=device, dtype=torch.int32).unsqueeze(0), + groups, + sliding_window=window, + attn_score=attention_score, + ) + + scores = torch.einsum("qhd,khd->hqk", q.float(), k.expand(-1, num_heads, -1).float()) + query = torch.arange(length, device=device)[:, None] + key = torch.arange(length, device=device)[None, :] + same_group = (groups[:, None] == groups[None, :]) & (groups[:, None] > 0) + visible = ((key <= query) | same_group) & (key > query - window) + expected_score = torch.where(visible.unsqueeze(0), scores, 0).sum(1) + scores.masked_fill_(~visible.unsqueeze(0), float("-inf")) + reference = torch.einsum("hqk,khd->qhd", scores.softmax(-1), v.expand(-1, num_heads, -1).float()) + + torch.testing.assert_close(output.float(), reference, atol=2e-2, rtol=2e-2) + torch.testing.assert_close(attention_score[0], expected_score, atol=2e-2, rtol=2e-2) diff --git a/tests/test_openai_api_server.py b/tests/test_openai_api_server.py index 72b65563..26b7bcd2 100644 --- a/tests/test_openai_api_server.py +++ b/tests/test_openai_api_server.py @@ -1237,6 +1237,27 @@ def apply_chat_template(self, chat, **_kwargs): self.assertEqual(prompt, "rendered") self.assertEqual(tokenizer.chat, [{"role": "system", "content": "policy\ndetails"}]) + def test_chat_prompt_preserves_multimodal_parts_for_model_processor(self): + from sparsevllm.entrypoints.openai.api_server import ChatMessage, _chat_prompt + from sparsevllm.multimodal import MultiModalPrompt + + prompt = _chat_prompt( + object(), + [ + ChatMessage( + role="user", + content=[ + {"type": "image_url", "image_url": {"url": "https://x/image.png"}}, + {"type": "text", "text": "describe"}, + ], + ) + ], + ) + + self.assertIsInstance(prompt, MultiModalPrompt) + self.assertEqual(prompt.messages[0]["content"][0]["type"], "image_url") + self.assertEqual(prompt.messages[0]["content"][1]["text"], "describe") + def test_chat_prompt_preserves_reasoning_content_for_templates(self): from sparsevllm.entrypoints.openai.api_server import ChatMessage, _chat_prompt @@ -4948,6 +4969,30 @@ class Tokenizer: ResponseRequest(model="model", input=[{"type": "image", "image_url": "x"}]), ) + def test_response_prompt_preserves_multimodal_parts_for_model_processor(self): + from sparsevllm.entrypoints.openai.api_server import ResponseRequest, _response_prompt + from sparsevllm.multimodal import MultiModalPrompt + + request = ResponseRequest( + model="model", + input=[ + { + "type": "message", + "role": "user", + "content": [ + {"type": "input_image", "image_url": "https://x/image.png"}, + {"type": "input_text", "text": "describe"}, + ], + } + ], + ) + + prompt = _response_prompt(object(), request) + + self.assertIsInstance(prompt, MultiModalPrompt) + self.assertEqual(prompt.messages[0]["content"][0]["type"], "input_image") + self.assertEqual(prompt.messages[0]["content"][1], {"type": "text", "text": "describe"}) + def test_response_prompt_passes_tools_and_tool_outputs(self): from sparsevllm.entrypoints.openai.api_server import ResponseRequest, _response_prompt diff --git a/tests/test_qwen35_mixed_runtime.py b/tests/test_qwen35_mixed_runtime.py index 27cc1fc8..2ed149cf 100644 --- a/tests/test_qwen35_mixed_runtime.py +++ b/tests/test_qwen35_mixed_runtime.py @@ -45,7 +45,11 @@ Qwen35RMSNorm, _get_rotary_dim, ) -from sparsevllm.models.qwen3_5_moe import Qwen35MoeRouter, Qwen35MoeSparseMoeBlock +from sparsevllm.models.qwen3_5_moe import ( + Qwen35MoePackedExperts, + Qwen35MoeRouter, + Qwen35MoeSparseMoeBlock, +) from sparsevllm.models.checkpoint import validate_checkpoint from sparsevllm.models.spec import resolve_model_spec from sparsevllm.platforms.cpu import CpuPlatform @@ -58,6 +62,22 @@ def _single_process_parallel_context() -> ParallelContext: return ParallelContext(world=group, tensor=group, expert=group, data=group) +def test_qwen35_fp8_expert_validation_uses_per_projection_weights(): + experts = SimpleNamespace( + fp8_enabled=True, + local_expert_start=0, + local_expert_end=2, + _loaded_packed_projections=set(), + _loaded_expert_shards={ + (expert_id, projection) + for expert_id in range(2) + for projection in ("gate_proj", "up_proj", "down_proj") + }, + ) + + Qwen35MoePackedExperts.validate_loaded_weights(experts) + + def _qwen35_outer_config(*, num_layers: int = 64, full_layers: tuple[int, ...] | None = None): if full_layers is None: full_layers = tuple(range(0, num_layers, 4)) diff --git a/tests/test_weight_loading.py b/tests/test_weight_loading.py index da142adc..30882df8 100644 --- a/tests/test_weight_loading.py +++ b/tests/test_weight_loading.py @@ -16,6 +16,13 @@ def __init__(self): self.right = nn.Linear(2, 2, bias=False) +class _BufferModel(nn.Module): + def __init__(self): + super().__init__() + self.weight = nn.Parameter(torch.empty(2)) + self.register_buffer("clip_max", torch.tensor(float("inf"))) + + class _RankLocalWeight(nn.Module): def __init__(self): super().__init__() @@ -147,6 +154,18 @@ def test_load_model_can_disable_progress(tmp_path, capsys): assert "loading shards" not in capsys.readouterr().err.lower() +def test_load_model_restores_checkpoint_buffers(tmp_path): + save_file( + {"weight": torch.ones(2), "clip_max": torch.tensor(3.5)}, + tmp_path / "model.safetensors", + ) + model = _BufferModel() + + loader.load_model(model, str(tmp_path), show_progress=False) + + torch.testing.assert_close(model.clip_max, torch.tensor(3.5)) + + def test_load_model_labels_progress_rank(tmp_path, capsys): _write_two_shards(tmp_path) From c61aa0bdb68793a74d7d50819424632e3ebecf91 Mon Sep 17 00:00:00 2001 From: QuanshengGu Date: Thu, 13 Aug 2026 05:41:43 +0800 Subject: [PATCH 04/13] fix: harden multimodal admission --- src/sparsevllm/engine/llm_engine.py | 13 ++- src/sparsevllm/engine/model_runner.py | 16 +++- src/sparsevllm/entrypoints/openai/render.py | 23 ++++- .../entrypoints/openai/serving/chat.py | 2 + .../entrypoints/openai/serving/responses.py | 2 + src/sparsevllm/multimodal/inputs.py | 23 +++-- tests/test_multimodal.py | 81 ++++++++++++++++-- tests/test_openai_api_server.py | 84 +++++++++++++++++++ tests/test_tp_rpc.py | 25 ++++++ 9 files changed, 248 insertions(+), 21 deletions(-) diff --git a/src/sparsevllm/engine/llm_engine.py b/src/sparsevllm/engine/llm_engine.py index abd22b95..78bb2e48 100644 --- a/src/sparsevllm/engine/llm_engine.py +++ b/src/sparsevllm/engine/llm_engine.py @@ -686,8 +686,17 @@ def admit_request( len(payload), ) ) - except Exception: - self.model_runner.call("free_multimodal", int(seq.seq_id)) + except Exception as register_error: + try: + self.model_runner.call("free_multimodal", int(seq.seq_id)) + except Exception as cleanup_error: + logger.error( + "Failed to roll back multimodal seq_id={} after registration " + "error {}: {}", + seq.seq_id, + type(register_error).__name__, + cleanup_error, + ) raise finally: payload_shm.close() diff --git a/src/sparsevllm/engine/model_runner.py b/src/sparsevllm/engine/model_runner.py index 1db4b039..ba288e48 100644 --- a/src/sparsevllm/engine/model_runner.py +++ b/src/sparsevllm/engine/model_runner.py @@ -104,6 +104,13 @@ def _create_model(hf_config, model_spec: ModelSpec, **runtime_kwargs): "prefix_cache_delete_subtree", "prefix_cache_set_eviction_priority", } +RECOVERABLE_TP_CONTROL_RPC_METHODS = PREFIX_CACHE_CONTROL_RPC_METHODS | { + "finish_slots_batch", + "free_multimodal", + "free_slots", + "free_slots_batch", + "register_multimodal_shared", +} TP_RPC_STATUS_SYNC_METHODS = PREFIX_CACHE_CONTROL_RPC_METHODS | { "chain_admission_plan", "chain_apply_admission", @@ -373,8 +380,13 @@ def loop(self): try: self.call(method_name, *args) except Exception as exc: - if method_name in PREFIX_CACHE_CONTROL_RPC_METHODS: - logger.error("TP worker prefix-cache control RPC failed: {}: {}", type(exc).__name__, exc) + if method_name in RECOVERABLE_TP_CONTROL_RPC_METHODS: + logger.error( + "TP worker recoverable control RPC {} failed: {}: {}", + method_name, + type(exc).__name__, + exc, + ) else: raise if method_name == "exit": diff --git a/src/sparsevllm/entrypoints/openai/render.py b/src/sparsevllm/entrypoints/openai/render.py index feae654c..c25e1c0a 100644 --- a/src/sparsevllm/entrypoints/openai/render.py +++ b/src/sparsevllm/entrypoints/openai/render.py @@ -437,8 +437,27 @@ def _response_content(content: Any) -> str | list[dict[str, Any]]: if not isinstance(part, dict): raise ValueError("message content parts must be JSON objects.") part_type = part.get("type") - if part_type in {"input_image", "input_audio", "input_video"}: - normalized.append(dict(part)) + if part_type in {"input_image", "input_video"}: + field = "image_url" if part_type == "input_image" else "video_url" + if set(part) - {"type", field}: + raise ValueError(f"{part_type} contains unsupported fields.") + value = part.get(field) + url = value.get("url") if isinstance(value, dict) else value + if not isinstance(url, str) or not url: + raise ValueError(f"{part_type} requires a non-empty {field}.") + normalized.append({"type": part_type, field: value}) + continue + if part_type == "input_audio": + if set(part) - {"type", "input_audio"}: + raise ValueError("input_audio contains unsupported fields.") + audio = part.get("input_audio") + if not isinstance(audio, dict) or set(audio) - {"data", "format"}: + raise ValueError("input_audio requires data and optional format fields.") + if not isinstance(audio.get("data"), str) or not audio["data"]: + raise ValueError("input_audio requires a non-empty base64 data string.") + if "format" in audio and str(audio["format"]).lower() != "wav": + raise ValueError("Only WAV input_audio is supported.") + normalized.append({"type": "input_audio", "input_audio": dict(audio)}) continue if part_type not in {"text", "input_text", "output_text"}: raise ValueError(f"Unsupported message content part type: {part_type!r}.") diff --git a/src/sparsevllm/entrypoints/openai/serving/chat.py b/src/sparsevllm/entrypoints/openai/serving/chat.py index aa53d834..42bb24f9 100644 --- a/src/sparsevllm/entrypoints/openai/serving/chat.py +++ b/src/sparsevllm/entrypoints/openai/serving/chat.py @@ -109,6 +109,8 @@ async def serve_chat_completion( ) except ChainCacheError as exc: raise _chain_http_exception(exc) from exc + except (ValueError, TypeError, NotImplementedError) as exc: + raise HTTPException(status_code=400, detail=str(exc)) from exc handles = [handle] headers = ( {"X-SparseVLLM-Chain-ID": getattr(handle, "chain_id", None)} diff --git a/src/sparsevllm/entrypoints/openai/serving/responses.py b/src/sparsevllm/entrypoints/openai/serving/responses.py index a04efff1..e767dc68 100644 --- a/src/sparsevllm/entrypoints/openai/serving/responses.py +++ b/src/sparsevllm/entrypoints/openai/serving/responses.py @@ -86,6 +86,8 @@ async def serve_response( ) except ChainCacheError as exc: raise _chain_http_exception(exc) from exc + except (ValueError, TypeError, NotImplementedError) as exc: + raise HTTPException(status_code=400, detail=str(exc)) from exc headers = ( {"X-SparseVLLM-Chain-ID": getattr(handle, "chain_id", None)} if getattr(handle, "chain_id", None) is not None diff --git a/src/sparsevllm/multimodal/inputs.py b/src/sparsevllm/multimodal/inputs.py index 5b5351da..a9de2673 100644 --- a/src/sparsevllm/multimodal/inputs.py +++ b/src/sparsevllm/multimodal/inputs.py @@ -1,6 +1,7 @@ from __future__ import annotations import base64 +import binascii import hashlib import io import wave @@ -39,15 +40,19 @@ def _audio_part(part: dict[str, Any]) -> dict[str, Any]: raise TypeError("input_audio requires a base64 data string.") if str(audio.get("format", "wav")).lower() != "wav": raise ValueError("Only WAV input_audio is supported.") - with wave.open(io.BytesIO(base64.b64decode(audio["data"])), "rb") as wav: - if wav.getcomptype() != "NONE": - raise ValueError("Compressed WAV input_audio is unsupported.") - channels, sample_width, sampling_rate = ( - wav.getnchannels(), - wav.getsampwidth(), - wav.getframerate(), - ) - raw = wav.readframes(wav.getnframes()) + try: + encoded = base64.b64decode(audio["data"], validate=True) + with wave.open(io.BytesIO(encoded), "rb") as wav: + if wav.getcomptype() != "NONE": + raise ValueError("Compressed WAV input_audio is unsupported.") + channels, sample_width, sampling_rate = ( + wav.getnchannels(), + wav.getsampwidth(), + wav.getframerate(), + ) + raw = wav.readframes(wav.getnframes()) + except (binascii.Error, wave.Error) as exc: + raise ValueError("input_audio data must contain a valid base64 WAV file.") from exc if sample_width == 1: waveform = (np.frombuffer(raw, np.uint8).astype(np.float32) - 128) / 128 elif sample_width in {2, 4}: diff --git a/tests/test_multimodal.py b/tests/test_multimodal.py index f4832f80..555ed94b 100644 --- a/tests/test_multimodal.py +++ b/tests/test_multimodal.py @@ -9,7 +9,11 @@ from sparsevllm.engine.sequence import Sequence from sparsevllm.engine.llm_engine import LLMEngine -from sparsevllm.multimodal.inputs import MultiModalInputProcessor, normalize_messages +from sparsevllm.multimodal.inputs import ( + MultiModalInputProcessor, + ProcessedMultiModalPrompt, + normalize_messages, +) from sparsevllm.multimodal.runtime import MultiModalRuntime, MultiModalState from sparsevllm.models.qwen3_5_multimodal import qwen35_mrope_positions from sparsevllm.operators.qwen35_mrope import Qwen35MRotaryEmbedding @@ -66,6 +70,24 @@ def test_normalize_openai_wav_audio_without_optional_dependencies(): torch.tensor([-1.0, 0.0, 32767 / 32768]), ) + +def test_normalize_openai_audio_rejects_invalid_base64_wav(): + with pytest.raises(ValueError, match="valid base64 WAV"): + normalize_messages( + [ + { + "role": "user", + "content": [ + { + "type": "input_audio", + "input_audio": {"format": "wav", "data": "not-base64"}, + } + ], + } + ] + ) + + def test_multimodal_processor_returns_stable_cpu_payload(): class Processor: def apply_chat_template(self, messages, **kwargs): @@ -229,23 +251,61 @@ def abort(self, seq_id): assert calls == [("free_multimodal", (seq.seq_id,))] +def test_multimodal_registration_error_survives_failed_rollback(): + calls = [] + + class Runner: + def call(self, method, *args): + calls.append(method) + if method == "register_multimodal_shared": + raise ValueError("rank-local encoder failure") + raise TimeoutError("rollback timeout") + + engine = object.__new__(LLMEngine) + engine.config = SimpleNamespace( + hf_config=SimpleNamespace(use_bidirectional_attention=None), + max_model_len=16, + resolved_prefix_cache_mode="disabled", + ) + engine.multimodal_processor = SimpleNamespace( + process=lambda _prompt: ProcessedMultiModalPrompt( + token_ids=[1, 2], + tensors={"mm_token_type_ids": torch.tensor([[0, 1]])}, + digest="digest", + ) + ) + engine.model_runner = Runner() + + with pytest.raises(ValueError, match="rank-local encoder failure"): + engine.admit_request( + {"messages": [{"role": "user", "content": []}]}, + SamplingParams(max_tokens=1), + ) + + assert calls == ["register_multimodal_shared", "free_multimodal"] + + @pytest.mark.skipif(not torch.cuda.is_available(), reason="requires CUDA") -def test_gemma4_multimodal_context_attention_matches_reference(): +@pytest.mark.parametrize("score_ndim", [2, 3]) +def test_gemma4_multimodal_context_attention_matches_reference(score_ndim): from sparsevllm.kernels.triton.gemma4_multimodal_context_attention import ( gemma4_multimodal_context_attention, ) device = torch.device("cuda") torch.manual_seed(0) - length, num_heads, head_dim, window = 9, 2, 256, 4 + length, num_heads, head_dim, window = 73, 2, 256, 64 q = (torch.randn(length, num_heads, head_dim, device=device) / head_dim**0.5).bfloat16() k = torch.randn(length, 1, head_dim, device=device).bfloat16() v = torch.randn_like(k) output = torch.empty_like(q) attention_score = torch.zeros( - 1, num_heads, length, device=device, dtype=torch.float32 + (1, num_heads, length) if score_ndim == 3 else (1, length), + device=device, + dtype=torch.float32, ) - groups = torch.tensor([0, 0, 1, 1, 1, 1, 0, 0, 0], device=device, dtype=torch.int32) + groups = torch.zeros(length, device=device, dtype=torch.int32) + groups[20:51] = 1 gemma4_multimodal_context_attention( q, k, @@ -267,7 +327,16 @@ def test_gemma4_multimodal_context_attention_matches_reference(): key = torch.arange(length, device=device)[None, :] same_group = (groups[:, None] == groups[None, :]) & (groups[:, None] > 0) visible = ((key <= query) | same_group) & (key > query - window) - expected_score = torch.where(visible.unsqueeze(0), scores, 0).sum(1) + visible_scores = torch.where(visible.unsqueeze(0), scores, 0) + if score_ndim == 3: + expected_score = visible_scores.sum(1) + else: + expected_score = torch.stack( + [ + visible_scores[:, start : start + 32].sum(1) / length + for start in range(0, length, 32) + ] + ).amax((0, 1)).clamp_min_(0) scores.masked_fill_(~visible.unsqueeze(0), float("-inf")) reference = torch.einsum("hqk,khd->qhd", scores.softmax(-1), v.expand(-1, num_heads, -1).float()) diff --git a/tests/test_openai_api_server.py b/tests/test_openai_api_server.py index 26b7bcd2..07e6d694 100644 --- a/tests/test_openai_api_server.py +++ b/tests/test_openai_api_server.py @@ -4993,6 +4993,90 @@ def test_response_prompt_preserves_multimodal_parts_for_model_processor(self): self.assertEqual(prompt.messages[0]["content"][0]["type"], "input_image") self.assertEqual(prompt.messages[0]["content"][1], {"type": "text", "text": "describe"}) + def test_response_prompt_validates_multimodal_parts(self): + from sparsevllm.entrypoints.openai.api_server import ResponseRequest, _response_prompt + + cases = [ + ({"type": "input_image"}, "requires a non-empty image_url"), + ({"type": "input_video", "video_url": ""}, "requires a non-empty video_url"), + ({"type": "input_audio", "input_audio": {}}, "base64 data string"), + ( + { + "type": "input_audio", + "input_audio": {"data": "AAAA", "format": "mp3"}, + }, + "Only WAV", + ), + ( + {"type": "input_image", "image_url": "https://x/image.png", "extra": 1}, + "unsupported fields", + ), + ] + for part, error in cases: + with self.subTest(part=part), self.assertRaisesRegex(ValueError, error): + _response_prompt( + object(), + ResponseRequest( + model="model", + input=[ + { + "type": "message", + "role": "user", + "content": [part], + } + ], + ), + ) + + async def test_multimodal_admission_errors_return_bad_request(self): + from fastapi import HTTPException + + from sparsevllm.entrypoints.openai.api_server import ( + ChatCompletionRequest, + ResponseRequest, + ) + from sparsevllm.entrypoints.openai.serving.chat import serve_chat_completion + from sparsevllm.entrypoints.openai.serving.responses import serve_response + + class Dispatcher: + admission_ack_enabled = True + + async def submit(self, *_args, **_kwargs): + raise AssertionError("submit_admitted must be used") + + async def submit_admitted(self, *_args, **_kwargs): + raise NotImplementedError("checkpoint does not support audio") + + class Tokenizer: + chat_template = "template" + + def apply_chat_template(self, *_args, **_kwargs): + return "rendered" + + requests = [ + serve_chat_completion( + ChatCompletionRequest( + model="model", messages=[{"role": "user", "content": "hello"}] + ), + Dispatcher(), + Tokenizer(), + "model", + None, + ), + serve_response( + ResponseRequest(model="model", input="hello"), + Dispatcher(), + Tokenizer(), + "model", + None, + None, + ), + ] + for request in requests: + with self.assertRaises(HTTPException) as ctx: + await request + self.assertEqual(ctx.exception.status_code, 400) + def test_response_prompt_passes_tools_and_tool_outputs(self): from sparsevllm.entrypoints.openai.api_server import ResponseRequest, _response_prompt diff --git a/tests/test_tp_rpc.py b/tests/test_tp_rpc.py index 8c686274..b8bf5e15 100644 --- a/tests/test_tp_rpc.py +++ b/tests/test_tp_rpc.py @@ -13,6 +13,7 @@ from sparsevllm.engine.model_runner import ( ModelRunner, PREFIX_CACHE_CONTROL_RPC_METHODS, + RECOVERABLE_TP_CONTROL_RPC_METHODS, TP_RUN_STATUS_FAILED, TP_RUN_STATUS_SUCCESS, TP_RPC_STATUS_SYNC_METHODS, @@ -433,6 +434,30 @@ def test_hidden_state_debug_uses_failure_synchronized_world_rpc(): assert "debug_moe_states_cpu" in TP_RPC_STATUS_SYNC_METHODS +def test_tp_worker_continues_after_multimodal_registration_failure(): + assert "register_multimodal_shared" in RECOVERABLE_TP_CONTROL_RPC_METHODS + runner = object.__new__(ModelRunner) + commands = iter( + [ + ("register_multimodal_shared", []), + ("free_multimodal", []), + ("exit", []), + ] + ) + calls = [] + runner.read_shm = lambda: next(commands) + + def call(method_name, *_args): + calls.append(method_name) + if method_name == "register_multimodal_shared": + raise ValueError("rank-local encoder failure") + + runner.call = call + ModelRunner.loop(runner) + + assert calls == ["register_multimodal_shared", "free_multimodal", "exit"] + + def test_model_runner_reset_after_warmup_resets_local_runtime_state(): calls = [] runner = object.__new__(ModelRunner) From c5da2a13176adb4d6181cc3fc1383b58a72ebd82 Mon Sep 17 00:00:00 2001 From: QuanshengGu Date: Thu, 13 Aug 2026 05:41:47 +0800 Subject: [PATCH 05/13] docs: document multimodal support --- README.md | 4 ++++ docs/en/features/supported-models.md | 17 +++++++++++++++++ docs/zh/features/supported-models.md | 17 +++++++++++++++++ 3 files changed, 38 insertions(+) diff --git a/README.md b/README.md index 2b87ad4b..7f41c198 100644 --- a/README.md +++ b/README.md @@ -71,12 +71,16 @@ Read the method overview and integration rules in | Qwen3.5 / Qwen3.6 | ✅ | | Qwen3.5 / Qwen3.6 MoE | ✅ | | GLM-4.7-Flash | ✅ | +| Gemma 4 Dense / MoE | ✅ | | Llama 3 / 3.1 | ✅ | | MiniMax M2.7 | ✅ | See [Supported Models](docs/en/features/supported-models.md) for the precision, parallelism, and sparse-method compatibility matrices. +Native image, video, and audio inputs are enabled per checkpoint with +`enable_multimodal=True`; see the supported-model matrix for media coverage. + ## Documentation | Topic | Link | diff --git a/docs/en/features/supported-models.md b/docs/en/features/supported-models.md index fd374fe3..bb5326f2 100644 --- a/docs/en/features/supported-models.md +++ b/docs/en/features/supported-models.md @@ -18,6 +18,7 @@ parallel size must use that value. | Qwen3.5 / Qwen3.6 | `qwen3_5` | BF16 / block FP8 | ✅ | 1 only | 1 only | | Qwen3.6 MoE | `qwen3_5_moe` | BF16 / block FP8 | ✅ | 1 only | ✅ | | GLM-4.7-Flash | `glm4_moe_lite` | BF16 | 1 / 2 / 4 (H100 only)⁵ | 1 only | 1 / 2 / 4⁵ | +| Gemma 4 Dense / MoE | `gemma4` | BF16 / FP16 | ✅ | 1 only | ✅ (MoE only) | | Llama 3 / 3.1 | `llama` | BF16 / FP16 | ✅ | 1 only | 1 only | | MiniMax M2.7 | `minimax_m2` | block FP8 with BF16 non-quantized weights | ✅ | 1 only | ✅ | @@ -59,6 +60,7 @@ layer. | Qwen3.5 / Qwen3.6 | ✅ | ✅ | ✅ | Experimental⁴ | ✅ | ✅ | ✅ | ✅ | — | Matched checkpoint³ | | Qwen3.6 MoE | ✅ | ✅ | ✅ | Experimental⁴ | ✅ | ✅ | ✅ | ✅ | — | — | | GLM-4.7-Flash | ✅⁵ | ✅⁵ | ✅⁵ | Experimental⁴⁵ | — | ✅⁵ | — | ✅⁵ | — | — | +| Gemma 4 Dense / MoE | ✅ | ✅⁶ | — | — | — | ✅ | — | — | — | — | | Llama 3 / 3.1 | ✅ | ✅ | ✅ | Experimental⁴ | ✅ | ✅ | ✅ | ✅ | Selected checkpoint¹ | Compressor required² | | MiniMax M2.7 | ✅ | ✅ | ✅ | Experimental⁴ | ✅ | ✅ | ✅ | ✅ | — | — | @@ -81,4 +83,19 @@ global-head selection. Model-specific TP, EP, and DP restrictions still apply. cross-rank sparse-index aggregation, so their selection semantics are not guaranteed to match `TP=1`. +⁶ Gemma 4 checkpoints with shared KV layers reject per-layer StreamingLLM +eviction. Vanilla and OmniKV remain supported. + +## Native Multimodal Support + +Set `enable_multimodal=True` to use a checkpoint's native media towers. +Sparse-vLLM accepts OpenAI-compatible Chat and Responses content parts and +uses the checkpoint processor and chat template. Unsupported media fail +explicitly during admission. + +| Model family | Image | Video | Audio | +| --- | :---: | :---: | :---: | +| Qwen3.5 / Qwen3.6 Dense and MoE | ✅ | ✅ | — | +| Gemma 4 Dense and MoE | ✅ | ✅ | Checkpoint dependent | + `—` means that the combination is not currently supported. diff --git a/docs/zh/features/supported-models.md b/docs/zh/features/supported-models.md index af43594a..bac80a78 100644 --- a/docs/zh/features/supported-models.md +++ b/docs/zh/features/supported-models.md @@ -14,6 +14,7 @@ | Qwen3.5 / Qwen3.6 | `qwen3_5` | BF16 / 块级 FP8 | ✅ | 仅支持 1 | 仅支持 1 | | Qwen3.6 MoE | `qwen3_5_moe` | BF16 / 块级 FP8 | ✅ | 仅支持 1 | ✅ | | GLM-4.7-Flash | `glm4_moe_lite` | BF16 | 1 / 2 / 4(仅 H100)⁵ | 仅支持 1 | 1 / 2 / 4⁵ | +| Gemma 4 Dense / MoE | `gemma4` | BF16 / FP16 | ✅ | 仅支持 1 | ✅(仅 MoE) | | Llama 3 / 3.1 | `llama` | BF16 / FP16 | ✅ | 仅支持 1 | 仅支持 1 | | MiniMax M2.7 | `minimax_m2` | 块级 FP8,非量化权重使用 BF16 | ✅ | 仅支持 1 | ✅ | @@ -50,6 +51,7 @@ vanilla 和 OmniKV 使用 radix 模式,对 StreamingLLM、SnapKV、H2O 和 R-K | Qwen3.5 / Qwen3.6 | ✅ | ✅ | ✅ | 实验性⁴ | ✅ | ✅ | ✅ | ✅ | — | 匹配的 checkpoint³ | | Qwen3.6 MoE | ✅ | ✅ | ✅ | 实验性⁴ | ✅ | ✅ | ✅ | ✅ | — | — | | GLM-4.7-Flash | ✅⁵ | ✅⁵ | ✅⁵ | 实验性⁴⁵ | — | ✅⁵ | — | ✅⁵ | — | — | +| Gemma 4 Dense / MoE | ✅ | ✅⁶ | — | — | — | ✅ | — | — | — | — | | Llama 3 / 3.1 | ✅ | ✅ | ✅ | 实验性⁴ | ✅ | ✅ | ✅ | ✅ | 指定 checkpoint¹ | 需要 compressor² | | MiniMax M2.7 | ✅ | ✅ | ✅ | 实验性⁴ | ✅ | ✅ | ✅ | ✅ | — | — | @@ -70,4 +72,19 @@ TP、EP、DP 限制仍然适用。 `TP>1` 时,基于 head 评分的稀疏方法使用 TP-local selection,不跨 rank 聚合 sparse index,因此其选择语义不保证与 `TP=1` 相同。 +⁶ 带共享 KV 层的 Gemma 4 checkpoint 不支持逐层 StreamingLLM eviction; +Vanilla 和 OmniKV 仍受支持。 + +## 原生多模态支持 + +设置 `enable_multimodal=True` 后,可使用 checkpoint 自带的媒体塔。 +Sparse-vLLM 接受 OpenAI 兼容的 Chat 与 Responses content part,并使用 +checkpoint 自身的 processor 和 chat template;不受支持的媒体会在接纳阶段 +明确报错。 + +| 模型家族 | 图片 | 视频 | 音频 | +| --- | :---: | :---: | :---: | +| Qwen3.5 / Qwen3.6 Dense 与 MoE | ✅ | ✅ | — | +| Gemma 4 Dense 与 MoE | ✅ | ✅ | 取决于 checkpoint | + `—` 表示当前不支持该组合。 From 8fc5c8aa60902ac4b535258e21c1d20a080b605e Mon Sep 17 00:00:00 2001 From: QuanshengGu Date: Thu, 13 Aug 2026 13:10:28 +0800 Subject: [PATCH 06/13] chore: add nsys benchmark capture --- benchmark/fixed_token_microbench.py | 27 +++++++++++++++++++++++++-- 1 file changed, 25 insertions(+), 2 deletions(-) diff --git a/benchmark/fixed_token_microbench.py b/benchmark/fixed_token_microbench.py index 9520292e..a27ecade 100644 --- a/benchmark/fixed_token_microbench.py +++ b/benchmark/fixed_token_microbench.py @@ -29,6 +29,12 @@ def _parser() -> argparse.ArgumentParser: parser.add_argument("--tensor-parallel-size", type=int, default=1) parser.add_argument("--expert-parallel-size", type=int, default=1) parser.add_argument("--gpu-memory-utilization", type=float, default=0.8) + parser.add_argument( + "--nsys-iteration", + type=int, + default=-1, + help="Wrap one timed iteration in an NVTX range for nsys capture.", + ) return parser @@ -76,6 +82,8 @@ def _validate(args: argparse.Namespace) -> Path: raise ValueError(f"{name} must be positive") if not 0 < args.gpu_memory_utilization <= 1: raise ValueError("gpu_memory_utilization must be in (0, 1]") + if not -1 <= args.nsys_iteration < args.num_iters: + raise ValueError("nsys_iteration must be -1 or a timed iteration index") if args.backend == "vllm" and args.expert_parallel_size != 1: raise ValueError( "vLLM uses --expert-parallel-size=1 or --enable-expert-parallel" @@ -94,6 +102,7 @@ def _build_engine(args: argparse.Namespace): "gpu_memory_utilization": args.gpu_memory_utilization, "tensor_parallel_size": args.tensor_parallel_size, "enforce_eager": False, + "enable_prefix_caching": False, } if args.backend == "vllm": from vllm import LLM @@ -162,6 +171,11 @@ def main() -> int: "tensor_parallel_size": args.tensor_parallel_size, "expert_parallel_size": args.expert_parallel_size, "cuda_graph": True, + "prefix_cache": False, + }, + "profiler": { + "kind": "nsys_nvtx" if args.nsys_iteration >= 0 else None, + "iteration": args.nsys_iteration if args.nsys_iteration >= 0 else None, }, "environment": { "cuda_visible_devices": os.getenv("CUDA_VISIBLE_DEVICES"), @@ -196,9 +210,18 @@ def main() -> int: if len(outputs) != args.batch_size: raise RuntimeError(f"warmup returned {len(outputs)} requests") for iteration in range(args.num_iters): + profiling = iteration == args.nsys_iteration + if profiling: + import torch + + torch.cuda.nvtx.range_push(f"fixed_token_iteration_{iteration}") started = perf_counter() - outputs = engine.generate(prompts, params, use_tqdm=False) - elapsed = perf_counter() - started + try: + outputs = engine.generate(prompts, params, use_tqdm=False) + elapsed = perf_counter() - started + finally: + if profiling: + torch.cuda.nvtx.range_pop() if len(outputs) != args.batch_size: raise RuntimeError( f"iteration {iteration} returned {len(outputs)} requests" From 58cc2f318c492686a30f4226d530d0e43ae81e21 Mon Sep 17 00:00:00 2001 From: QuanshengGu Date: Thu, 13 Aug 2026 18:27:20 +0800 Subject: [PATCH 07/13] perf: optimize gemma inference --- benchmark/fixed_token_microbench.py | 27 +- src/sparsevllm/engine/cache_manager/base.py | 14 +- src/sparsevllm/engine/sparse_controller.py | 9 +- .../triton/gemma4_context_attention.py | 13 +- .../kernels/triton/gemma4_fused_router.py | 108 +++++++ .../triton/gemma4_global_decode_attention.py | 185 ++++++++++++ src/sparsevllm/kernels/triton/gemma4_moe.py | 33 ++- .../triton/gemma4_window_decode_attention.py | 266 +++++++++++++++++ src/sparsevllm/models/gemma4.py | 34 ++- src/sparsevllm/models/layout.py | 36 ++- src/sparsevllm/operators/gemma4.py | 72 ++++- src/sparsevllm/operators/gemma4_attention.py | 193 +++++++++++- src/sparsevllm/operators/gemma4_moe.py | 34 ++- src/sparsevllm/utils/config.py | 32 ++ tests/test_gemma4_attention_kernels.py | 276 ++++++++++++++++++ tests/test_gemma4_model.py | 123 +++++++- 16 files changed, 1397 insertions(+), 58 deletions(-) create mode 100644 src/sparsevllm/kernels/triton/gemma4_fused_router.py create mode 100644 src/sparsevllm/kernels/triton/gemma4_global_decode_attention.py create mode 100644 src/sparsevllm/kernels/triton/gemma4_window_decode_attention.py diff --git a/benchmark/fixed_token_microbench.py b/benchmark/fixed_token_microbench.py index a27ecade..f4ac19b8 100644 --- a/benchmark/fixed_token_microbench.py +++ b/benchmark/fixed_token_microbench.py @@ -84,9 +84,13 @@ def _validate(args: argparse.Namespace) -> Path: raise ValueError("gpu_memory_utilization must be in (0, 1]") if not -1 <= args.nsys_iteration < args.num_iters: raise ValueError("nsys_iteration must be -1 or a timed iteration index") - if args.backend == "vllm" and args.expert_parallel_size != 1: + if args.backend == "vllm" and args.expert_parallel_size not in { + 1, + args.tensor_parallel_size, + }: raise ValueError( - "vLLM uses --expert-parallel-size=1 or --enable-expert-parallel" + "vLLM expert parallelism is disabled with EP=1 or spans TP ranks " + "with EP=TP" ) output_dir = Path(args.output_dir).expanduser().resolve() if output_dir.exists() and any(output_dir.iterdir()): @@ -105,8 +109,24 @@ def _build_engine(args: argparse.Namespace): "enable_prefix_caching": False, } if args.backend == "vllm": + import vllm.config.model as vllm_model_config from vllm import LLM + get_config = vllm_model_config.get_config + + def get_compatible_config(*config_args, **config_kwargs): + config = get_config(*config_args, **config_kwargs) + text = getattr(config, "text_config", config) + layers = getattr(text, "per_layer_config", None) + if str(getattr(config, "model_type", "")) == "gemma4" and layers: + text.allow_global_per_layer_attribute_access = True + full_idx = text.layer_types.index("full_attention") + text.global_head_dim = layers[full_idx].head_dim + text.num_global_key_value_heads = layers[full_idx].num_key_value_heads + return config + + vllm_model_config.get_config = get_compatible_config + return LLM( **common, max_num_seqs=args.batch_size, @@ -127,6 +147,7 @@ def _build_engine(args: argparse.Namespace): max_num_seqs_in_gpu=args.batch_size, max_num_batched_tokens=args.batch_size * args.input_len, decode_cuda_graph=True, + enable_multimodal=False, ) @@ -215,12 +236,14 @@ def main() -> int: import torch torch.cuda.nvtx.range_push(f"fixed_token_iteration_{iteration}") + torch.cuda.cudart().cudaProfilerStart() started = perf_counter() try: outputs = engine.generate(prompts, params, use_tqdm=False) elapsed = perf_counter() - started finally: if profiling: + torch.cuda.cudart().cudaProfilerStop() torch.cuda.nvtx.range_pop() if len(outputs) != args.batch_size: raise RuntimeError( diff --git a/src/sparsevllm/engine/cache_manager/base.py b/src/sparsevllm/engine/cache_manager/base.py index 65bcab09..517bd048 100644 --- a/src/sparsevllm/engine/cache_manager/base.py +++ b/src/sparsevllm/engine/cache_manager/base.py @@ -244,8 +244,18 @@ def __init__(self, config: Config, parallel_context: ParallelContext): raise ValueError("CacheManager requires config.runtime_layout.") self.num_kv_layers = int(self.runtime_layout.num_kv_layers) - self.num_kv_heads = self.hf_config.num_key_value_heads // self.tp_size - self.head_dim = resolve_attention_qk_head_dim(self.hf_config) + layout_heads = tuple(getattr(self.runtime_layout, "kv_num_heads", ())) + layout_dims = tuple(getattr(self.runtime_layout, "kv_head_dims", ())) + self.num_kv_heads = ( + int(layout_heads[0]) // self.tp_size + if layout_heads + else int(self.hf_config.num_key_value_heads) // self.tp_size + ) + self.head_dim = ( + int(layout_dims[0]) + if layout_dims + else resolve_attention_qk_head_dim(self.hf_config) + ) self.max_model_len = config.max_model_len resident_buffer_rows = int(config.max_num_seqs_in_gpu) diff --git a/src/sparsevllm/engine/sparse_controller.py b/src/sparsevllm/engine/sparse_controller.py index 1f9d7214..08b5269c 100644 --- a/src/sparsevllm/engine/sparse_controller.py +++ b/src/sparsevllm/engine/sparse_controller.py @@ -84,7 +84,14 @@ def __init__(self, config: Config, cache_manager: CacheManager): self.num_sink = self.config.num_sink_tokens self.num_recent = self.config.num_recent_tokens self.decode_keep_tokens = self.config.decode_keep_tokens - head_dim = resolve_attention_qk_head_dim(self.config.hf_config) + layout_dims = tuple( + getattr(getattr(self.config, "runtime_layout", None), "kv_head_dims", ()) + ) + head_dim = ( + int(layout_dims[0]) + if layout_dims + else resolve_attention_qk_head_dim(self.config.hf_config) + ) self.attn_softmax_scale = float(head_dim) ** -0.5 score_dtype_name = str(getattr(self.config, "sparse_attn_score_dtype", "float32") or "float32").lower() self.attn_score_dtype = { diff --git a/src/sparsevllm/kernels/triton/gemma4_context_attention.py b/src/sparsevllm/kernels/triton/gemma4_context_attention.py index ba1ac1f5..7e005142 100644 --- a/src/sparsevllm/kernels/triton/gemma4_context_attention.py +++ b/src/sparsevllm/kernels/triton/gemma4_context_attention.py @@ -107,10 +107,17 @@ def _gemma4_context_attention_kernel( mask=key_positions < max_key, ) logits = tl.where(visible, logits, -float("inf")) + has_visible_key = tl.max(visible.to(tl.int32), axis=1) > 0 block_max = tl.max(logits, axis=1) - new_max = tl.maximum(max_logit, block_max) - probabilities = tl.exp2(logits - new_max[:, None]) - correction = tl.exp2(max_logit - new_max) + new_max = tl.where( + has_visible_key, tl.maximum(max_logit, block_max), max_logit + ) + probabilities = tl.where( + visible, tl.exp2(logits - new_max[:, None]), 0.0 + ) + correction = tl.where( + has_visible_key, tl.exp2(max_logit - new_max), 1.0 + ) denominator = denominator * correction + tl.sum(probabilities, axis=1) accumulator *= correction[:, None] value = tl.load( diff --git a/src/sparsevllm/kernels/triton/gemma4_fused_router.py b/src/sparsevllm/kernels/triton/gemma4_fused_router.py new file mode 100644 index 00000000..cbdaaab4 --- /dev/null +++ b/src/sparsevllm/kernels/triton/gemma4_fused_router.py @@ -0,0 +1,108 @@ +from __future__ import annotations + +import torch +import triton +import triton.language as tl + + +@triton.jit +def _gemma4_fused_router_kernel( + logits_ptr, + scale_ptr, + weights_ptr, + ids_ptr, + stride_logits, + NUM_EXPERTS: tl.constexpr, + TOP_K: tl.constexpr, + BLOCK_EXPERTS: tl.constexpr, +): + row = tl.program_id(0) + experts = tl.arange(0, BLOCK_EXPERTS) + valid = experts < NUM_EXPERTS + logits = tl.load( + logits_ptr + row * stride_logits + experts, + mask=valid, + other=-float("inf"), + ).to(tl.float32) + + # Pack a descending-float sort key and the expert id into one int64 value. + min_int32 = -2147483648 + bits = logits.to(tl.int32, bitcast=True) + keys = tl.where(bits >> 31 == 0, bits ^ -1, bits ^ min_int32) + keys = tl.where(valid, keys, 0x7FFFFFFF) + packed = ((keys.to(tl.int64) & 0xFFFFFFFF) << 32) | experts.to(tl.int64) + sorted_packed = tl.sort(packed, descending=False) + sorted_keys = ((sorted_packed >> 32) & 0xFFFFFFFF).to(tl.int32) + sorted_ids = (sorted_packed & 0xFFFFFFFF).to(tl.int32) + sorted_bits = tl.where( + sorted_keys >> 31 < 0, + sorted_keys ^ -1, + sorted_keys ^ min_int32, + ) + sorted_logits = sorted_bits.to(tl.float32, bitcast=True) + + selected = experts < TOP_K + selected_logits = tl.where(selected, sorted_logits, -float("inf")) + selected_max = tl.max(selected_logits, axis=0) + probabilities = tl.where( + selected, + tl.exp2((sorted_logits - selected_max) * 1.4426950408889634), + 0.0, + ) + probabilities /= tl.sum(probabilities, axis=0) + probabilities *= tl.load( + scale_ptr + sorted_ids, + mask=selected, + other=1.0, + ).to(tl.float32) + output_offsets = row * TOP_K + experts + tl.store(weights_ptr + output_offsets, probabilities, mask=selected) + tl.store(ids_ptr + output_offsets, sorted_ids, mask=selected) + + +def gemma4_fused_router_topk( + logits: torch.Tensor, + per_expert_scale: torch.Tensor, + top_k: int, +) -> tuple[torch.Tensor, torch.Tensor]: + if not logits.is_cuda or logits.dtype not in {torch.float16, torch.bfloat16}: + raise TypeError("Gemma 4 fused router requires CUDA FP16 or BF16 logits.") + num_experts = int(logits.shape[-1]) + if ( + logits.ndim != 2 + or logits.stride(-1) != 1 + or per_expert_scale.shape != (num_experts,) + or per_expert_scale.stride(0) != 1 + ): + raise ValueError( + "Gemma 4 fused router requires contiguous [tokens, experts] logits " + "and matching contiguous expert scales." + ) + if per_expert_scale.device != logits.device or per_expert_scale.dtype != logits.dtype: + raise TypeError("Gemma 4 fused router scales must match logits dtype and device.") + if not 0 < int(top_k) <= num_experts or num_experts > 1024: + raise ValueError( + f"Gemma 4 fused router requires 0 < top_k <= experts <= 1024, got " + f"top_k={top_k}, experts={num_experts}." + ) + weights = torch.empty( + (logits.shape[0], int(top_k)), dtype=torch.float32, device=logits.device + ) + ids = torch.empty( + (logits.shape[0], int(top_k)), dtype=torch.int32, device=logits.device + ) + _gemma4_fused_router_kernel[(int(logits.shape[0]),)]( + logits, + per_expert_scale, + weights, + ids, + logits.stride(0), + NUM_EXPERTS=num_experts, + TOP_K=int(top_k), + BLOCK_EXPERTS=triton.next_power_of_2(num_experts), + num_warps=1, + ) + return weights, ids + + +__all__ = ["gemma4_fused_router_topk"] diff --git a/src/sparsevllm/kernels/triton/gemma4_global_decode_attention.py b/src/sparsevllm/kernels/triton/gemma4_global_decode_attention.py new file mode 100644 index 00000000..a5b33a88 --- /dev/null +++ b/src/sparsevllm/kernels/triton/gemma4_global_decode_attention.py @@ -0,0 +1,185 @@ +from __future__ import annotations + +import torch +import triton +import triton.language as tl + + +@triton.jit +def _gemma4_global_decode_stage1_kernel( + q, + k, + v, + active_slots, + req_indices, + context_lens, + mid_output, + mid_lse, + stride_qb, + stride_qh, + stride_kt, + stride_kh, + stride_vt, + stride_vh, + stride_sb, + stride_ss, + stride_mob, + stride_moh, + stride_mos, + stride_mlb, + stride_mlh, + stride_mls, + GROUP_SIZE: tl.constexpr, + HEADS_PER_PROGRAM: tl.constexpr, + HEAD_DIM: tl.constexpr, + BLOCK_SEQ: tl.constexpr, + BLOCK_N: tl.constexpr, +): + batch = tl.program_id(0) + head_group = tl.program_id(1) + sequence_block = tl.program_id(2) + heads = head_group * HEADS_PER_PROGRAM + tl.arange(0, HEADS_PER_PROGRAM) + kv_head = head_group * HEADS_PER_PROGRAM // GROUP_SIZE + dims = tl.arange(0, HEAD_DIM) + sequence_len = tl.load(context_lens + batch) + block_start = sequence_block * BLOCK_SEQ + mid_offset = ( + batch * stride_mob + + heads[:, None] * stride_moh + + sequence_block * stride_mos + ) + lse_offset = ( + batch * stride_mlb + heads * stride_mlh + sequence_block * stride_mls + ) + if block_start >= sequence_len: + tl.store(mid_output + mid_offset + dims[None, :], 0.0) + tl.store(mid_lse + lse_offset, -float("inf")) + return + + query = tl.load(q + batch * stride_qb + heads[:, None] * stride_qh + dims) + max_logit = tl.full((HEADS_PER_PROGRAM,), -float("inf"), tl.float32) + denominator = tl.zeros((HEADS_PER_PROGRAM,), tl.float32) + accumulator = tl.zeros((HEADS_PER_PROGRAM, HEAD_DIM), tl.float32) + block_end = tl.minimum(sequence_len, block_start + BLOCK_SEQ) + request = tl.load(req_indices + batch) + for offset in range(0, BLOCK_SEQ, BLOCK_N): + positions = block_start + offset + tl.arange(0, BLOCK_N) + visible = positions < block_end + slots = tl.load( + active_slots + request * stride_sb + positions * stride_ss, + mask=visible, + other=0, + ) + key = tl.load( + k + slots[None, :] * stride_kt + kv_head * stride_kh + dims[:, None], + mask=visible[None, :], + other=0.0, + ) + logits = tl.dot(query, key) * 1.4426950408889634 + logits = tl.where(visible[None, :], logits, -float("inf")) + block_max = tl.max(logits, axis=1) + new_max = tl.maximum(max_logit, block_max) + probabilities = tl.exp2(logits - new_max[:, None]) + correction = tl.exp2(max_logit - new_max) + denominator = denominator * correction + tl.sum(probabilities, axis=1) + accumulator *= correction[:, None] + value = tl.load( + v + slots[:, None] * stride_vt + kv_head * stride_vh + dims[None, :], + mask=visible[:, None], + other=0.0, + ) + accumulator += tl.dot(probabilities.to(value.dtype), value) + max_logit = new_max + tl.store(mid_output + mid_offset + dims[None, :], accumulator / denominator[:, None]) + tl.store( + mid_lse + lse_offset, + max_logit * 0.6931471805599453 + tl.log(denominator), + ) + + +def gemma4_global_decode_stage1( + q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, + active_slots: torch.Tensor, + req_indices: torch.Tensor, + context_lens: torch.Tensor, + mid_output: torch.Tensor, + mid_lse: torch.Tensor, + *, + block_seq: int, + heads_per_program: int = 4, +) -> None: + if q.ndim != 3 or k.ndim != 3 or v.shape != k.shape: + raise ValueError("Gemma 4 global decode requires matching rank-3 Q/K/V.") + if not all(t.is_cuda for t in (q, k, v, mid_output, mid_lse)): + raise TypeError("Gemma 4 global decode requires CUDA tensors.") + if q.dtype not in {torch.float16, torch.bfloat16} or any( + tensor.dtype != q.dtype for tensor in (k, v) + ): + raise TypeError("Gemma 4 global decode requires matching FP16 or BF16 Q/K/V.") + if int(q.shape[1]) % int(k.shape[1]): + raise ValueError("Gemma 4 global decode requires divisible Q and KV heads.") + head_dim = int(q.shape[-1]) + group_size = int(q.shape[1]) // int(k.shape[1]) + heads_per_program = int(heads_per_program) + if ( + head_dim != 512 + or int(k.shape[-1]) != head_dim + or group_size % heads_per_program + or heads_per_program not in {2, 4} + ): + raise ValueError( + "Gemma 4 global decode requires head_dim=512 and GQA groups divisible " + f"by 2 or 4, got head_dim={head_dim}, group_size={group_size}, " + f"heads_per_program={heads_per_program}." + ) + if mid_output.dtype != torch.float32 or mid_lse.dtype != torch.float32: + raise TypeError("Gemma 4 global decode workspace must use FP32 tensors.") + expected_mid = (q.shape[0], q.shape[1], mid_output.shape[2], head_dim) + expected_lse = expected_mid[:-1] + if mid_output.shape != expected_mid or mid_lse.shape != expected_lse: + raise ValueError( + f"Gemma 4 global decode workspace must have shapes {expected_mid} and " + f"{expected_lse}, got {tuple(mid_output.shape)} and {tuple(mid_lse.shape)}." + ) + _gemma4_global_decode_stage1_kernel[ + ( + int(q.shape[0]), + int(q.shape[1]) // heads_per_program, + int(mid_output.shape[2]), + ) + ]( + q, + k, + v, + active_slots, + req_indices, + context_lens, + mid_output, + mid_lse, + q.stride(0), + q.stride(1), + k.stride(0), + k.stride(1), + v.stride(0), + v.stride(1), + active_slots.stride(0), + active_slots.stride(1), + mid_output.stride(0), + mid_output.stride(1), + mid_output.stride(2), + mid_lse.stride(0), + mid_lse.stride(1), + mid_lse.stride(2), + GROUP_SIZE=group_size, + HEADS_PER_PROGRAM=heads_per_program, + HEAD_DIM=head_dim, + BLOCK_SEQ=int(block_seq), + BLOCK_N=16, + num_warps=8, + num_stages=1, + ) + + +__all__ = ["gemma4_global_decode_stage1"] diff --git a/src/sparsevllm/kernels/triton/gemma4_moe.py b/src/sparsevllm/kernels/triton/gemma4_moe.py index 3287dfd3..6c11cf5c 100644 --- a/src/sparsevllm/kernels/triton/gemma4_moe.py +++ b/src/sparsevllm/kernels/triton/gemma4_moe.py @@ -12,17 +12,21 @@ from sparsevllm.kernels.triton.moe_config import device_info, resolve_moe_gemm_config -def _gemma4_moe_config(num_tokens: int) -> dict[str, int] | None: - if int(num_tokens) > 32: +def _gemma4_moe_config( + num_tokens: int, large_token_config: dict[str, int] | None +) -> dict[str, int] | None: + if int(num_tokens) <= 32: + return { + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 64, + "BLOCK_SIZE_K": 128, + "GROUP_SIZE_M": 1, + "num_warps": 4, + "num_stages": 4, + } + if int(num_tokens) < 512: return None - return { - "BLOCK_SIZE_M": 16, - "BLOCK_SIZE_N": 64, - "BLOCK_SIZE_K": 128, - "GROUP_SIZE_M": 1, - "num_warps": 4, - "num_stages": 4, - } + return None if large_token_config is None else dict(large_token_config) def fused_gemma4_moe( @@ -34,6 +38,7 @@ def fused_gemma4_moe( *, num_experts: int, local_expert_start: int, + large_token_config: dict[str, int] | None = None, ) -> torch.Tensor: """Run Gemma 4 routed GEGLU experts without changing generic MoE kernels.""" @@ -55,7 +60,9 @@ def fused_gemma4_moe( hidden_states.device.type, int(hidden_states.device.index), ) - w13_config = _gemma4_moe_config(num_tokens) or resolve_moe_gemm_config( + w13_config = _gemma4_moe_config( + num_tokens, large_token_config + ) or resolve_moe_gemm_config( dtype=hidden_states.dtype, num_tokens=num_tokens, top_k=top_k, @@ -89,7 +96,9 @@ def fused_gemma4_moe( launch_config=w13_config, ) activated = gelu_tanh_and_mul_fwd(w13_output) - w2_config = _gemma4_moe_config(num_tokens) or resolve_moe_gemm_config( + w2_config = _gemma4_moe_config( + num_tokens, large_token_config + ) or resolve_moe_gemm_config( dtype=hidden_states.dtype, num_tokens=num_tokens, top_k=top_k, diff --git a/src/sparsevllm/kernels/triton/gemma4_window_decode_attention.py b/src/sparsevllm/kernels/triton/gemma4_window_decode_attention.py new file mode 100644 index 00000000..118fb2d5 --- /dev/null +++ b/src/sparsevllm/kernels/triton/gemma4_window_decode_attention.py @@ -0,0 +1,266 @@ +from __future__ import annotations + +import torch +import triton +import triton.language as tl + + +@triton.jit +def _gemma4_window_decode_stage1_kernel( + q, + k, + v, + active_slots, + req_indices, + context_lens, + mid_output, + mid_lse, + stride_qb, + stride_qh, + stride_kt, + stride_kh, + stride_vt, + stride_vh, + stride_sb, + stride_ss, + stride_mob, + stride_moh, + stride_mos, + stride_mlb, + stride_mlh, + stride_mls, + GROUP_SIZE: tl.constexpr, + HEAD_DIM: tl.constexpr, + BLOCK_SEQ: tl.constexpr, + BLOCK_N: tl.constexpr, + WINDOW: tl.constexpr, +): + batch = tl.program_id(0) + kv_head = tl.program_id(1) + sequence_block = tl.program_id(2) + groups = tl.arange(0, GROUP_SIZE) + dims = tl.arange(0, HEAD_DIM) + sequence_len = tl.load(context_lens + batch) + window_start = tl.maximum(0, sequence_len - WINDOW) + block_start = window_start + sequence_block * BLOCK_SEQ + query_head = kv_head * GROUP_SIZE + groups + mid_offset = ( + batch * stride_mob + + query_head[:, None] * stride_moh + + sequence_block * stride_mos + ) + lse_offset = ( + batch * stride_mlb + + query_head * stride_mlh + + sequence_block * stride_mls + ) + if block_start >= sequence_len: + tl.store(mid_output + mid_offset + dims[None, :], 0.0) + tl.store(mid_lse + lse_offset, -float("inf")) + return + + block_end = tl.minimum(sequence_len, block_start + BLOCK_SEQ) + query = tl.load( + q + batch * stride_qb + query_head[:, None] * stride_qh + dims[None, :] + ) + max_logit = tl.full((GROUP_SIZE,), -float("inf"), tl.float32) + denominator = tl.zeros((GROUP_SIZE,), tl.float32) + accumulator = tl.zeros((GROUP_SIZE, HEAD_DIM), tl.float32) + for offset in range(0, BLOCK_SEQ, BLOCK_N): + positions = block_start + offset + tl.arange(0, BLOCK_N) + visible = positions < block_end + request = tl.load(req_indices + batch) + slots = tl.load( + active_slots + request * stride_sb + positions * stride_ss, + mask=visible, + other=0, + ) + key = tl.load( + k + slots[None, :] * stride_kt + kv_head * stride_kh + dims[:, None], + mask=visible[None, :], + other=0.0, + ) + logits = tl.dot(query, key) * 1.4426950408889634 + logits = tl.where(visible[None, :], logits, -float("inf")) + block_max = tl.max(logits, axis=1) + new_max = tl.maximum(max_logit, block_max) + probabilities = tl.exp2(logits - new_max[:, None]) + correction = tl.exp2(max_logit - new_max) + denominator = denominator * correction + tl.sum(probabilities, axis=1) + accumulator *= correction[:, None] + value = tl.load( + v + slots[:, None] * stride_vt + kv_head * stride_vh + dims[None, :], + mask=visible[:, None], + other=0.0, + ) + accumulator += tl.dot(probabilities.to(value.dtype), value) + max_logit = new_max + tl.store(mid_output + mid_offset + dims[None, :], accumulator / denominator[:, None]) + tl.store( + mid_lse + lse_offset, + max_logit * 0.6931471805599453 + tl.log(denominator), + ) + + +@triton.jit +def _gemma4_window_decode_stage2_kernel( + context_lens, + mid_output, + mid_lse, + output, + stride_mob, + stride_moh, + stride_mos, + stride_mlb, + stride_mlh, + stride_mls, + stride_ob, + stride_oh, + GROUP_SIZE: tl.constexpr, + HEAD_DIM: tl.constexpr, + BLOCK_SEQ: tl.constexpr, + NUM_BLOCKS: tl.constexpr, + WINDOW: tl.constexpr, +): + batch = tl.program_id(0) + kv_head = tl.program_id(1) + groups = tl.arange(0, GROUP_SIZE) + dims = tl.arange(0, HEAD_DIM) + query_head = kv_head * GROUP_SIZE + groups + sequence_len = tl.load(context_lens + batch) + block_count = (tl.minimum(sequence_len, WINDOW) + BLOCK_SEQ - 1) // BLOCK_SEQ + max_lse = tl.full((GROUP_SIZE,), -float("inf"), tl.float32) + denominator = tl.zeros((GROUP_SIZE,), tl.float32) + accumulator = tl.zeros((GROUP_SIZE, HEAD_DIM), tl.float32) + for block in range(0, NUM_BLOCKS): + valid = block < block_count + lse = tl.load( + mid_lse + + batch * stride_mlb + + query_head * stride_mlh + + block * stride_mls + ) + lse = tl.where(valid, lse, -float("inf")) + value = tl.load( + mid_output + + batch * stride_mob + + query_head[:, None] * stride_moh + + block * stride_mos + + dims[None, :] + ) + new_max = tl.maximum(max_lse, lse) + old_scale = tl.exp(max_lse - new_max) + new_scale = tl.exp(lse - new_max) + accumulator = accumulator * old_scale[:, None] + value * new_scale[:, None] + denominator = denominator * old_scale + new_scale + max_lse = new_max + tl.store( + output + + batch * stride_ob + + query_head[:, None] * stride_oh + + dims[None, :], + accumulator / denominator[:, None], + ) + + +def gemma4_window_decode( + q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, + active_slots: torch.Tensor, + req_indices: torch.Tensor, + context_lens: torch.Tensor, + mid_output: torch.Tensor, + mid_lse: torch.Tensor, + output: torch.Tensor, + *, + block_seq: int, + sliding_window: int, +) -> None: + if q.ndim != 3 or k.ndim != 3 or v.shape != k.shape or output.shape != q.shape: + raise ValueError("Gemma 4 window decode requires matching rank-3 Q/K/V/output.") + if not all(t.is_cuda for t in (q, k, v, mid_output, mid_lse, output)): + raise TypeError("Gemma 4 window decode requires CUDA tensors.") + if q.dtype not in {torch.float16, torch.bfloat16} or any( + tensor.dtype != q.dtype for tensor in (k, v, output) + ): + raise TypeError("Gemma 4 window decode requires matching FP16 or BF16 Q/K/V.") + if mid_output.dtype != torch.float32 or mid_lse.dtype != torch.float32: + raise TypeError("Gemma 4 window decode workspace must use FP32 tensors.") + if int(q.shape[1]) % int(k.shape[1]): + raise ValueError("Gemma 4 window decode requires divisible Q and KV heads.") + group_size = int(q.shape[1]) // int(k.shape[1]) + head_dim = int(q.shape[-1]) + if group_size not in {2, 4} or head_dim != 256 or int(k.shape[-1]) != head_dim: + raise ValueError( + "Gemma 4 window decode requires head_dim=256 and GQA group 2 or 4, " + f"got head_dim={head_dim}, group_size={group_size}." + ) + block_seq, sliding_window = int(block_seq), int(sliding_window) + if block_seq <= 0 or sliding_window <= 0: + raise ValueError("Gemma 4 window decode requires positive block and window sizes.") + num_blocks = triton.cdiv(sliding_window, block_seq) + if mid_output.shape[2] < num_blocks or mid_lse.shape[2] < num_blocks: + raise ValueError( + f"Gemma 4 window workspace needs {num_blocks} blocks, got " + f"{mid_output.shape[2]}/{mid_lse.shape[2]}." + ) + mid_output = mid_output[:, :, :num_blocks] + mid_lse = mid_lse[:, :, :num_blocks] + _gemma4_window_decode_stage1_kernel[ + (int(q.shape[0]), int(k.shape[1]), num_blocks) + ]( + q, + k, + v, + active_slots, + req_indices, + context_lens, + mid_output, + mid_lse, + q.stride(0), + q.stride(1), + k.stride(0), + k.stride(1), + v.stride(0), + v.stride(1), + active_slots.stride(0), + active_slots.stride(1), + mid_output.stride(0), + mid_output.stride(1), + mid_output.stride(2), + mid_lse.stride(0), + mid_lse.stride(1), + mid_lse.stride(2), + GROUP_SIZE=group_size, + HEAD_DIM=head_dim, + BLOCK_SEQ=block_seq, + BLOCK_N=32, + WINDOW=sliding_window, + num_warps=8, + num_stages=1, + ) + _gemma4_window_decode_stage2_kernel[(int(q.shape[0]), int(k.shape[1]))]( + context_lens, + mid_output, + mid_lse, + output, + mid_output.stride(0), + mid_output.stride(1), + mid_output.stride(2), + mid_lse.stride(0), + mid_lse.stride(1), + mid_lse.stride(2), + output.stride(0), + output.stride(1), + GROUP_SIZE=group_size, + HEAD_DIM=head_dim, + BLOCK_SEQ=block_seq, + NUM_BLOCKS=num_blocks, + WINDOW=sliding_window, + num_warps=8, + num_stages=2, + ) + + +__all__ = ["gemma4_window_decode"] diff --git a/src/sparsevllm/models/gemma4.py b/src/sparsevllm/models/gemma4.py index 230ee7ed..9e1da5ec 100644 --- a/src/sparsevllm/models/gemma4.py +++ b/src/sparsevllm/models/gemma4.py @@ -30,6 +30,7 @@ from sparsevllm.operators.moe import model_activation_dtype from sparsevllm.platforms import device_runtime from sparsevllm.utils.context import get_context +from sparsevllm.utils.config import config_get, config_layer_get from sparsevllm.utils.weight_target import WeightTarget _EXPERT_SOURCE_RE = re.compile( @@ -151,15 +152,25 @@ def __init__( self.is_kv_shared_layer = layer_idx >= shared_start > 0 self.sliding_window = int(config.sliding_window) if self.is_sliding else None self.head_dim = int( - config.head_dim if self.is_sliding else config.global_head_dim + config_layer_get( + config, + layer_idx, + "head_dim", + "head_dim" if self.is_sliding else "global_head_dim", + ) ) self.total_num_heads = int(config.num_attention_heads) self.num_heads = self.total_num_heads // tp_size self.use_k_eq_v = bool(config.attention_k_eq_v and not self.is_sliding) self.total_num_kv_heads = int( - config.num_global_key_value_heads - if self.use_k_eq_v - else config.num_key_value_heads + config_layer_get( + config, + layer_idx, + "num_key_value_heads", + "num_global_key_value_heads" + if self.use_k_eq_v + else "num_key_value_heads", + ) ) self.num_kv_heads = max(1, self.total_num_kv_heads // tp_size) self.q_size = self.num_heads * self.head_dim @@ -464,9 +475,14 @@ def __init__( layer_type: Gemma4RotaryEmbedding( config, layer_type, - config.head_dim - if layer_type == "sliding_attention" - else config.global_head_dim, + config_layer_get( + config, + config.layer_types.index(layer_type), + "head_dim", + "head_dim" + if layer_type == "sliding_attention" + else "global_head_dim", + ), ) for layer_type in set(config.layer_types) } @@ -597,8 +613,8 @@ def build_runtime_kwargs(cls, config, *, device, **_): head_dims = tuple( sorted( { - int(config.head_dim), - int(getattr(config, "global_head_dim", config.head_dim)), + int(config_layer_get(config, layer_idx, "head_dim")) + for layer_idx in range(int(config.num_hidden_layers)) } ) ) diff --git a/src/sparsevllm/models/layout.py b/src/sparsevllm/models/layout.py index acb65fac..1b798b12 100644 --- a/src/sparsevllm/models/layout.py +++ b/src/sparsevllm/models/layout.py @@ -3,7 +3,7 @@ from dataclasses import dataclass, replace from typing import Any -from sparsevllm.utils.config import config_get +from sparsevllm.utils.config import config_get, config_layer_get def resolve_attention_qk_head_dim(hf_config: Any) -> int: @@ -313,20 +313,32 @@ def _with_attention_shapes( ) if invalid_types: raise ValueError(f"Unsupported Gemma 4 layer types: {invalid_types}.") - sliding_heads = int(config_get(hf_config, "num_key_value_heads")) - sliding_dim = int(config_get(hf_config, "head_dim")) - global_dim = int(config_get(hf_config, "global_head_dim", sliding_dim)) - use_k_eq_v = bool(config_get(hf_config, "attention_k_eq_v", False)) - global_heads = int( - config_get(hf_config, "num_global_key_value_heads", sliding_heads) - if use_k_eq_v - else sliding_heads - ) heads, dims = [], [] for layer_idx in layout.kv_idx_to_layer_idx: is_full = str(layer_types[layer_idx]) == "full_attention" - heads.append(global_heads if is_full else sliding_heads) - dims.append(global_dim if is_full else sliding_dim) + dims.append( + int( + config_layer_get( + hf_config, + layer_idx, + "head_dim", + "global_head_dim" if is_full else "head_dim", + ) + ) + ) + heads.append( + int( + config_layer_get( + hf_config, + layer_idx, + "num_key_value_heads", + "num_global_key_value_heads" + if is_full + and config_get(hf_config, "attention_k_eq_v", False) + else "num_key_value_heads", + ) + ) + ) return replace(layout, kv_num_heads=tuple(heads), kv_head_dims=tuple(dims)) def is_full_attention(self, layer_idx: int) -> bool: diff --git a/src/sparsevllm/operators/gemma4.py b/src/sparsevllm/operators/gemma4.py index cde36cf7..a93d383f 100644 --- a/src/sparsevllm/operators/gemma4.py +++ b/src/sparsevllm/operators/gemma4.py @@ -1,13 +1,20 @@ from __future__ import annotations from dataclasses import dataclass +from importlib.metadata import PackageNotFoundError, version +from importlib.util import find_spec import torch import torch.nn.functional as F import sparsevllm.platforms as platforms from sparsevllm.layers.rotary_embedding import apply_rotary_emb -from sparsevllm.operators.registry import OpRegistry, OpResolver, SupportResult +from sparsevllm.operators.registry import ( + OpRegistry, + OpResolver, + SupportResult, + runtime_version_at_least, +) from sparsevllm.platforms.interface import DeviceCaps, PlatformEnum @@ -155,6 +162,68 @@ def rmsnorm_residual(self, x, weight, residual, eps, scalar=None): return gemma4_rmsnorm_residual(x, weight, residual, eps, scalar) +@GEMMA4_REGISTRY.register +class H20Gemma4OperatorProvider(TritonGemma4OperatorProvider): + """Profiled H20 provider; generic Gemma kernels remain unchanged.""" + + name = "gemma4_h20" + priority = 100 + + @classmethod + def supports(cls, spec: Gemma4OpSpec, caps: DeviceCaps) -> SupportResult: + triton = super().supports(spec, caps) + if not triton.supported: + return triton + if caps.compute_capability != (9, 0) or caps.device_name != "NVIDIA H20": + return SupportResult.no( + "requires profiled NVIDIA H20 SM90 hardware, " + f"got {caps.device_name} {caps.compute_capability}" + ) + if not runtime_version_at_least(caps.runtime_version, (12, 8)): + return SupportResult.no( + f"requires CUDA runtime >= 12.8, got {caps.runtime_version or 'unknown'}" + ) + if find_spec("flashinfer") is None: + return SupportResult.no("flashinfer is not installed") + try: + installed = version("flashinfer-python") + except PackageNotFoundError: + return SupportResult.no("flashinfer-python package metadata is unavailable") + try: + numeric = tuple(int(part) for part in installed.split(".")[:3]) + except ValueError: + return SupportResult.no( + f"cannot parse flashinfer-python version {installed!r}" + ) + if numeric < (0, 6, 15): + return SupportResult.no( + f"requires flashinfer-python >= 0.6.15, got {installed}" + ) + return SupportResult.yes() + + def __init__(self) -> None: + from sparsevllm.operators.gemma4_attention import Gemma4FlashInferPrefill + + self._prefill = Gemma4FlashInferPrefill() + + def router_topk(self, logits, per_expert_scale, top_k): + from sparsevllm.kernels.triton.gemma4_fused_router import ( + gemma4_fused_router_topk, + ) + + return gemma4_fused_router_topk(logits, per_expert_scale, top_k) + + def attention_backend(self, *, sliding_window: int | None): + from sparsevllm.operators.gemma4_attention import Gemma4AttentionBackend + + return Gemma4AttentionBackend( + sliding_window=sliding_window, + flashinfer_prefill=self._prefill, + use_window_decode=True, + global_decode_heads_per_program=4, + ) + + class TorchGemma4OperatorProvider(Gemma4OperatorProvider): """Explicit correctness oracle; never selected for production inference.""" @@ -216,6 +285,7 @@ def resolve_gemma4_provider( __all__ = [ "GEMMA4_REGISTRY", + "H20Gemma4OperatorProvider", "Gemma4OperatorProvider", "Gemma4OpSpec", "TorchGemma4OperatorProvider", diff --git a/src/sparsevllm/operators/gemma4_attention.py b/src/sparsevllm/operators/gemma4_attention.py index 8736c067..8ead4b4e 100644 --- a/src/sparsevllm/operators/gemma4_attention.py +++ b/src/sparsevllm/operators/gemma4_attention.py @@ -1,5 +1,7 @@ from __future__ import annotations +from dataclasses import dataclass + import torch from sparsevllm.layers.attention_backend import ( @@ -8,14 +10,127 @@ ) +@dataclass +class _FlashInferState: + wrapper: object + plan_key: tuple[object, int, int, int] | None = None + + +class Gemma4FlashInferPrefill: + """Shared FlashInfer plans for Gemma 4 text-prefill head shapes.""" + + def __init__(self) -> None: + self._states: dict[tuple[int, int, int, int], _FlashInferState] = {} + + @staticmethod + def _page_metadata(view, max_context_len: int): + meta = view.meta + rows = meta.active_slots.index_select(0, meta.req_indices.to(torch.long))[ + :, :max_context_len + ] + positions = torch.arange( + max_context_len, + device=meta.context_lens.device, + dtype=meta.context_lens.dtype, + ) + indices = rows.masked_select( + positions.unsqueeze(0) < meta.context_lens.unsqueeze(1) + ).to(torch.int32).contiguous() + indptr = torch.cat( + ( + torch.zeros(1, device=indices.device, dtype=torch.int32), + meta.context_lens.to(torch.int32).cumsum(0, dtype=torch.int32), + ) + ) + return indices, indptr, torch.ones_like(meta.context_lens, dtype=torch.int32) + + def run( + self, + q: torch.Tensor, + view, + *, + q_start: torch.Tensor, + chunk_lens: torch.Tensor, + max_context_len: int, + sliding_window: int | None, + ) -> torch.Tensor: + from flashinfer.prefill import BatchPrefillWithPagedKVCacheWrapper + from sparsevllm.utils.context import get_context + + payload = _require_explicit_payload(view, operation="Gemma 4 prefill") + meta = view.meta + if meta.active_slots.dtype != torch.int32 or meta.active_slots.ndim != 2: + raise TypeError("Gemma 4 FlashInfer prefill requires an int32 page table.") + q_heads, kv_heads, head_dim = map( + int, (q.shape[1], payload.k_cache.shape[1], q.shape[2]) + ) + window_left = -1 if sliding_window is None else int(sliding_window) - 1 + key = q_heads, kv_heads, head_dim, window_left + state = self._states.get(key) + if state is None: + workspace = torch.empty(128 * 1024 * 1024, dtype=torch.uint8, device=q.device) + state = _FlashInferState( + BatchPrefillWithPagedKVCacheWrapper( + workspace, kv_layout="NHD", backend="auto" + ) + ) + self._states[key] = state + context = get_context() + plan_key = ( + context.attention_validation_scope, + meta.active_slots.data_ptr(), + meta.req_indices.data_ptr(), + meta.context_lens.data_ptr(), + ) + if state.plan_key != plan_key: + indices, kv_indptr, last_page_len = self._page_metadata( + view, int(max_context_len) + ) + qo_indptr = torch.cat((q_start, q_start[-1:] + chunk_lens[-1:])) + state.wrapper.plan( + qo_indptr, + kv_indptr, + indices, + last_page_len, + q_heads, + kv_heads, + head_dim, + 1, + causal=True, + sm_scale=1.0, + window_left=window_left, + q_data_type=q.dtype, + kv_data_type=payload.k_cache.dtype, + non_blocking=True, + ) + state.plan_key = plan_key + output = torch.empty_like(q) + state.wrapper.run( + q, + (payload.k_cache.unsqueeze(1), payload.v_cache.unsqueeze(1)), + out=output, + ) + return output + + class Gemma4AttentionBackend(TritonAttentionBackend): """Gemma 4 attention semantics isolated from the tuned generic kernels.""" name = "triton_gemma4" - def __init__(self, *, sliding_window: int | None) -> None: + def __init__( + self, + *, + sliding_window: int | None, + flashinfer_prefill: Gemma4FlashInferPrefill | None = None, + use_window_decode: bool = False, + global_decode_heads_per_program: int | None = None, + ) -> None: super().__init__() self.sliding_window = None if sliding_window is None else int(sliding_window) + self.flashinfer_prefill = flashinfer_prefill + self.use_window_decode = bool(use_window_decode) + self.global_decode_heads_per_program = global_decode_heads_per_program def run_prefill( self, @@ -27,7 +142,6 @@ def run_prefill( max_input_len: int, ) -> torch.Tensor: payload = _require_explicit_payload(view, operation="Gemma 4 prefill") - output = torch.empty_like(q) from sparsevllm.utils.context import get_context image_groups = getattr(get_context(), "multimodal_image_groups", None) @@ -36,6 +150,7 @@ def run_prefill( gemma4_multimodal_context_attention, ) + output = torch.empty_like(q) gemma4_multimodal_context_attention( q, payload.k_cache, @@ -52,6 +167,16 @@ def run_prefill( attn_score=view.meta.attn_score, ) return output + if self.flashinfer_prefill is not None and view.meta.attn_score is None: + return self.flashinfer_prefill.run( + q, + view, + q_start=b_start_loc, + chunk_lens=chunk_lens, + max_context_len=max_input_len, + sliding_window=self.sliding_window, + ) + output = torch.empty_like(q) from sparsevllm.kernels.triton.gemma4_context_attention import ( gemma4_context_attention, ) @@ -94,6 +219,35 @@ def run_decode( ) group_size = int(q.shape[1]) // int(payload.k_cache.shape[1]) + if ( + self.use_window_decode + and self.sliding_window is not None + and view.meta.attn_score is None + and int(q.shape[-1]) == 256 + and group_size in {2, 4} + and mid_o.shape[2] + >= (self.sliding_window + block_seq - 1) // block_seq + ): + from sparsevllm.kernels.triton.gemma4_window_decode_attention import ( + gemma4_window_decode, + ) + + output = torch.empty_like(q) + window_blocks = (self.sliding_window + block_seq - 1) // block_seq + gemma4_window_decode( + q, + payload.k_cache, + payload.v_cache, + view.meta.active_slots, + view.meta.req_indices, + view.meta.context_lens, + mid_o[:, :, :window_blocks], + mid_o_logexpsum[:, :, :window_blocks], + output, + block_seq=block_seq, + sliding_window=self.sliding_window, + ) + return output if mid_o.shape[2] == 1 and view.meta.attn_score is None and group_size in {2, 4, 8}: from sparsevllm.kernels.triton.gemma4_single_block_decode_attention import ( gemma4_single_block_decode, @@ -106,6 +260,39 @@ def run_decode( block_seq=block_seq, sliding_window=self.sliding_window, ) return output + if ( + self.sliding_window is None + and view.meta.attn_score is None + and int(q.shape[-1]) == 512 + and self.global_decode_heads_per_program is not None + and group_size % self.global_decode_heads_per_program == 0 + ): + from sparsevllm.kernels.triton.gemma4_global_decode_attention import ( + gemma4_global_decode_stage1, + ) + + gemma4_global_decode_stage1( + q, + payload.k_cache, + payload.v_cache, + view.meta.active_slots, + view.meta.req_indices, + view.meta.context_lens, + mid_o, + mid_o_logexpsum, + block_seq=block_seq, + heads_per_program=self.global_decode_heads_per_program, + ) + output = torch.empty_like(q) + gemma4_decode_stage2( + mid_o, + mid_o_logexpsum, + view.meta.context_lens, + output, + block_seq=block_seq, + sliding_window=None, + ) + return output gemma4_decode_stage1( q, payload.k_cache, payload.v_cache, view.meta.active_slots, view.meta.req_indices, view.meta.context_lens, mid_o, @@ -125,4 +312,4 @@ def run_decode( return output -__all__ = ["Gemma4AttentionBackend"] +__all__ = ["Gemma4AttentionBackend", "Gemma4FlashInferPrefill"] diff --git a/src/sparsevllm/operators/gemma4_moe.py b/src/sparsevllm/operators/gemma4_moe.py index 59e08929..06168c34 100644 --- a/src/sparsevllm/operators/gemma4_moe.py +++ b/src/sparsevllm/operators/gemma4_moe.py @@ -7,7 +7,12 @@ from sparsevllm.distributed import get_parallel_context from sparsevllm.layers.packed_moe import PackedMoeExperts from sparsevllm.operators.moe import MoeOpSpec, MoeProvider, model_activation_dtype -from sparsevllm.operators.registry import OpRegistry, OpResolver, SupportResult +from sparsevllm.operators.registry import ( + OpRegistry, + OpResolver, + SupportResult, + runtime_version_at_least, +) from sparsevllm.platforms.interface import DeviceCaps, PlatformEnum @@ -49,6 +54,7 @@ class TritonGemma4MoeProvider(Gemma4MoeProvider): name = "triton_gemma4_geglu" priority = 10 gate_up_order = "gate_up" + _large_token_config = None @classmethod def supports(cls, spec: MoeOpSpec, caps: DeviceCaps) -> SupportResult: @@ -93,9 +99,34 @@ def run( topk_weights, num_experts=spec.num_experts, local_expert_start=local_expert_start, + large_token_config=self._large_token_config, ) +@GEMMA4_MOE_REGISTRY.register +class H20Gemma4MoeProvider(TritonGemma4MoeProvider): + name = "triton_gemma4_geglu_h20" + priority = 100 + _large_token_config = { + "BLOCK_SIZE_M": 64, + "BLOCK_SIZE_N": 64, + "BLOCK_SIZE_K": 64, + "GROUP_SIZE_M": 1, + "num_warps": 4, + "num_stages": 3, + } + + @classmethod + def supports(cls, spec: MoeOpSpec, caps: DeviceCaps) -> SupportResult: + triton = super().supports(spec, caps) + if not triton.supported: + return triton + if caps.compute_capability != (9, 0) or caps.device_name != "NVIDIA H20": + return SupportResult.no("requires profiled NVIDIA H20 SM90 hardware") + if not runtime_version_at_least(caps.runtime_version, (12, 8)): + return SupportResult.no("requires CUDA runtime >= 12.8") + return SupportResult.yes() + class TorchGemma4MoeProvider(Gemma4MoeProvider): """Explicit correctness oracle; never selected for production inference.""" @@ -221,6 +252,7 @@ def load_packed_weight(self, projection: str, loaded_weight: torch.Tensor) -> No "GEMMA4_MOE_REGISTRY", "Gemma4MoeProvider", "Gemma4PackedExperts", + "H20Gemma4MoeProvider", "TorchGemma4MoeProvider", "TritonGemma4MoeProvider", "resolve_gemma4_moe_provider", diff --git a/src/sparsevllm/utils/config.py b/src/sparsevllm/utils/config.py index 665270c9..386975c8 100644 --- a/src/sparsevllm/utils/config.py +++ b/src/sparsevllm/utils/config.py @@ -5,3 +5,35 @@ def config_get(config: Any, name: str, default: Any = None) -> Any: if config is None: return default return config.get(name, default) if isinstance(config, dict) else getattr(config, name, default) + + +def config_layer(config: Any, layer_idx: int) -> Any: + layers = config_get(config, "per_layer_config", None) + if not layers: + return config + if isinstance(layers, dict): + index = int(layer_idx) + return layers.get( + index, layers.get(str(index), layers.get(f"{index:02d}", config)) + ) + return layers[layer_idx] + + +def config_layer_get( + config: Any, layer_idx: int, name: str, legacy_name: str | None = None +) -> Any: + layer = config_layer(config, layer_idx) + missing = object() + if layer is not config: + value = config_get(layer, name, missing) + if value is not missing: + return value + fallback_name = legacy_name or name + if isinstance(config, dict): + return config.get(fallback_name) + values = vars(config) + return ( + values[fallback_name] + if fallback_name in values + else getattr(config, fallback_name, None) + ) diff --git a/tests/test_gemma4_attention_kernels.py b/tests/test_gemma4_attention_kernels.py index 802e8365..421f59af 100644 --- a/tests/test_gemma4_attention_kernels.py +++ b/tests/test_gemma4_attention_kernels.py @@ -1,20 +1,122 @@ from __future__ import annotations +from types import SimpleNamespace + import pytest import torch +from sparsevllm.engine.cache_manager.base import ExplicitKVPayload from sparsevllm.kernels.triton.gemma4_context_attention import gemma4_context_attention from sparsevllm.kernels.triton.gemma4_decode_attention import ( gemma4_decode_stage1, gemma4_decode_stage2, ) +from sparsevllm.kernels.triton.gemma4_global_decode_attention import ( + gemma4_global_decode_stage1, +) from sparsevllm.kernels.triton.gemma4_single_block_decode_attention import ( gemma4_single_block_decode, ) +from sparsevllm.kernels.triton.gemma4_window_decode_attention import ( + gemma4_window_decode, +) +from sparsevllm.operators.gemma4_attention import Gemma4FlashInferPrefill +from sparsevllm.utils.context import reset_context, set_context pytestmark = pytest.mark.skipif(not torch.cuda.is_available(), reason="requires CUDA") +@pytest.mark.parametrize( + ("head_dim", "q_heads", "kv_heads", "sliding_window"), + [(256, 4, 2, 32), (512, 4, 1, None)], +) +def test_gemma4_flashinfer_prefill_matches_torch( + head_dim, q_heads, kv_heads, sliding_window +): + pytest.importorskip("flashinfer") + torch.manual_seed(20260813) + prefix, chunk, length = 11, 54, 65 + slots = torch.randperm(length, device="cuda", dtype=torch.int64).to(torch.int32) + key = torch.randn(length, kv_heads, head_dim, device="cuda", dtype=torch.bfloat16) + value = torch.randn_like(key) + query = torch.randn(chunk, q_heads, head_dim, device="cuda", dtype=torch.bfloat16) + view = SimpleNamespace( + payload=ExplicitKVPayload(key, value), + meta=SimpleNamespace( + active_slots=slots.view(1, -1), + req_indices=torch.zeros(1, device="cuda", dtype=torch.int32), + context_lens=torch.tensor([length], device="cuda", dtype=torch.int32), + attn_score=None, + ), + ) + reset_context() + set_context( + True, + cu_seqlens_q=torch.tensor([0, chunk], device="cuda", dtype=torch.int32), + ) + prefill = Gemma4FlashInferPrefill() + try: + output = prefill.run( + query, + view, + q_start=torch.zeros(1, device="cuda", dtype=torch.int32), + chunk_lens=torch.tensor([chunk], device="cuda", dtype=torch.int32), + max_context_len=length, + sliding_window=sliding_window, + ) + second_slots = slots.flip(0).contiguous() + view.meta.active_slots = second_slots.view(1, -1) + second_output = prefill.run( + query, + view, + q_start=torch.zeros(1, device="cuda", dtype=torch.int32), + chunk_lens=torch.tensor([chunk], device="cuda", dtype=torch.int32), + max_context_len=length, + sliding_window=sliding_window, + ) + set_context( + True, + cu_seqlens_q=torch.tensor([0, chunk], device="cuda", dtype=torch.int32), + ) + second_slots.copy_(slots) + reused_output = prefill.run( + query, + view, + q_start=torch.zeros(1, device="cuda", dtype=torch.int32), + chunk_lens=torch.tensor([chunk], device="cuda", dtype=torch.int32), + max_context_len=length, + sliding_window=sliding_window, + ) + finally: + reset_context() + kv_head_ids = torch.arange(q_heads, device="cuda") // (q_heads // kv_heads) + query_positions = prefix + torch.arange(chunk, device="cuda") + key_positions = torch.arange(length, device="cuda") + visible = key_positions[None] <= query_positions[:, None] + if sliding_window is not None: + visible &= key_positions[None] > query_positions[:, None] - sliding_window + for actual, slot_ids in ( + (output, slots), + (second_output, slots.flip(0)), + (reused_output, slots), + ): + logical_key, logical_value = key[slot_ids.long()], value[slot_ids.long()] + logits = torch.einsum( + "qhd,khd->hqk", query, logical_key[:, kv_head_ids] + ).float() + probabilities = logits.masked_fill(~visible[None], -torch.inf).softmax(-1) + reference = torch.einsum( + "hqk,khd->qhd", + probabilities.to(value.dtype), + logical_value[:, kv_head_ids], + ) + cosine = torch.nn.functional.cosine_similarity( + actual.float().flatten(), reference.float().flatten(), dim=0 + ) + assert torch.isfinite(actual).all() + assert cosine > 0.999 + + def _slots_and_lengths(): lengths = torch.tensor([21, 13], dtype=torch.int32, device="cuda") slots = torch.zeros((2, 21), dtype=torch.int32, device="cuda") @@ -36,6 +138,53 @@ def _decode_reference(q, k, v, slots, lengths, window): return output +@pytest.mark.parametrize("group_size", [8, 16]) +@pytest.mark.parametrize("length", [513, 8193]) +def test_gemma4_global_decode_matches_torch_and_graph(group_size, length): + torch.manual_seed(20260813) + block_seq = 256 + slots = torch.arange(length, device="cuda", dtype=torch.int32).view(1, -1) + lengths = torch.tensor([length], device="cuda", dtype=torch.int32) + key = torch.randn(length, 1, 512, device="cuda", dtype=torch.bfloat16) + value = torch.randn_like(key) + query = torch.randn(1, group_size, 512, device="cuda", dtype=torch.bfloat16) + blocks = (length + block_seq - 1) // block_seq + mid = torch.empty(1, group_size, blocks, 512, device="cuda", dtype=torch.float32) + lse = torch.empty(1, group_size, blocks, device="cuda", dtype=torch.float32) + output = torch.empty_like(query) + + def run(): + gemma4_global_decode_stage1( + query, + key, + value, + slots, + torch.zeros(1, device="cuda", dtype=torch.int32), + lengths, + mid, + lse, + block_seq=block_seq, + ) + gemma4_decode_stage2( + mid, lse, lengths, output, block_seq=block_seq, sliding_window=None + ) + + run() + reference = _decode_reference(query, key, value, slots, lengths, None) + cosine = torch.nn.functional.cosine_similarity( + output.float().flatten(), reference.float().flatten(), dim=0 + ) + assert cosine > 0.999 + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph): + run() + query.copy_(torch.randn_like(query)) + graph.replay() + replay = output.clone() + graph.replay() + torch.testing.assert_close(output, replay, rtol=0, atol=0) + + @pytest.mark.parametrize("head_dim", [256, 512]) @pytest.mark.parametrize("sliding_window", [None, 4]) def test_gemma4_prefill_matches_torch(head_dim, sliding_window): @@ -83,6 +232,50 @@ def test_gemma4_prefill_matches_torch(head_dim, sliding_window): assert cosine > 0.999 +def test_gemma4_long_window_prefill_matches_torch(): + torch.manual_seed(5) + prefix = torch.tensor([256], dtype=torch.int32, device="cuda") + chunks = torch.tensor([1088], dtype=torch.int32, device="cuda") + lengths = prefix + chunks + slots = torch.arange(1344, dtype=torch.int32, device="cuda").view(1, -1) + key = torch.randn(1344, 2, 256, dtype=torch.bfloat16, device="cuda") + value = torch.randn_like(key) + query = torch.randn(1088, 4, 256, dtype=torch.bfloat16, device="cuda") + output = torch.empty_like(query) + gemma4_context_attention( + query, + key, + value, + output, + torch.tensor([0], dtype=torch.int32, device="cuda"), + torch.tensor([0], dtype=torch.int32, device="cuda"), + lengths, + prefix, + 1088, + slots, + sliding_window=1024, + ) + kv_heads = torch.arange(4, device="cuda") // 2 + logits = torch.bmm( + query.permute(1, 0, 2), + key[:, kv_heads].permute(1, 2, 0), + ).float() + query_positions = 256 + torch.arange(1088, device="cuda") + key_positions = torch.arange(1344, device="cuda") + visible = (key_positions[None, :] <= query_positions[:, None]) & ( + key_positions[None, :] > query_positions[:, None] - 1024 + ) + probabilities = logits.masked_fill(~visible, -float("inf")).softmax(-1) + reference = torch.bmm( + probabilities.to(value.dtype), value[:, kv_heads].permute(1, 0, 2) + ).permute(1, 0, 2) + cosine = torch.nn.functional.cosine_similarity( + output.float().flatten(), reference.float().flatten(), dim=0 + ) + assert torch.isfinite(output).all() + assert cosine > 0.999 + + @pytest.mark.parametrize("head_dim", [256, 512]) @pytest.mark.parametrize("sliding_window", [None, 4]) def test_gemma4_decode_matches_torch(head_dim, sliding_window): @@ -154,6 +347,89 @@ def test_gemma4_single_block_decode_matches_torch(group_size, head_dim): assert cosine > 0.999 +@pytest.mark.parametrize("group_size", [2, 4]) +@pytest.mark.parametrize("block_seq", [250, 256]) +def test_gemma4_window_decode_matches_torch(group_size, block_seq): + torch.manual_seed(17) + lengths = torch.tensor([1301, 1177], dtype=torch.int32, device="cuda") + slots = torch.zeros((2, 1301), dtype=torch.int32, device="cuda") + slots[0, :1301] = torch.arange(1301, dtype=torch.int32, device="cuda") + slots[1, :1177] = torch.arange(1301, 2478, dtype=torch.int32, device="cuda") + key = torch.randn(2478, 2, 256, dtype=torch.bfloat16, device="cuda") + value = torch.randn_like(key) + query = torch.randn(2, 2 * group_size, 256, dtype=torch.bfloat16, device="cuda") + blocks = (1024 + block_seq - 1) // block_seq + mid = torch.empty( + 2, 2 * group_size, blocks, 256, dtype=torch.float32, device="cuda" + ) + lse = torch.empty( + 2, 2 * group_size, blocks, dtype=torch.float32, device="cuda" + ) + output = torch.empty_like(query) + gemma4_window_decode( + query, + key, + value, + slots, + torch.tensor([0, 1], dtype=torch.int32, device="cuda"), + lengths, + mid, + lse, + output, + block_seq=block_seq, + sliding_window=1024, + ) + reference = _decode_reference(query, key, value, slots, lengths, 1024) + cosine = torch.nn.functional.cosine_similarity( + output.float().flatten(), reference.float().flatten(), dim=0 + ) + assert torch.isfinite(output).all() + assert cosine > 0.999 + + +def test_gemma4_window_decode_supports_cuda_graph(): + torch.manual_seed(23) + slots, lengths = _slots_and_lengths() + key = torch.randn(34, 2, 256, dtype=torch.bfloat16, device="cuda") + value = torch.randn_like(key) + query = torch.randn(2, 4, 256, dtype=torch.bfloat16, device="cuda") + mid = torch.empty(2, 4, 2, 256, dtype=torch.float32, device="cuda") + lse = torch.empty(2, 4, 2, dtype=torch.float32, device="cuda") + output = torch.empty_like(query) + request_indices = torch.tensor([0, 1], dtype=torch.int32, device="cuda") + + def run(): + gemma4_window_decode( + query, + key, + value, + slots, + request_indices, + lengths, + mid, + lse, + output, + block_seq=8, + sliding_window=16, + ) + + for _ in range(3): + run() + reference = _decode_reference(query, key, value, slots, lengths, 16) + cosine = torch.nn.functional.cosine_similarity( + output.float().flatten(), reference.float().flatten(), dim=0 + ) + assert cosine > 0.999 + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph): + run() + query.copy_(torch.randn_like(query)) + graph.replay() + first = output.clone() + graph.replay() + assert torch.equal(first, output) + + @pytest.mark.parametrize("head_dim", [256, 512]) def test_gemma4_decode_supports_cuda_graph(head_dim): slots, lengths = _slots_and_lengths() diff --git a/tests/test_gemma4_model.py b/tests/test_gemma4_model.py index 624b7c5f..cdd36011 100644 --- a/tests/test_gemma4_model.py +++ b/tests/test_gemma4_model.py @@ -1,5 +1,6 @@ from __future__ import annotations +import inspect from contextlib import ExitStack from types import SimpleNamespace from unittest.mock import patch @@ -34,21 +35,31 @@ ) from sparsevllm.models.layout import RuntimeLayout from sparsevllm.operators.gemma4 import ( + GEMMA4_REGISTRY, + Gemma4OpSpec, + H20Gemma4OperatorProvider, TorchGemma4OperatorProvider, TritonGemma4OperatorProvider, ) from sparsevllm.operators.gemma4_moe import ( GEMMA4_MOE_REGISTRY, + H20Gemma4MoeProvider, TorchGemma4MoeProvider, TritonGemma4MoeProvider, ) +from sparsevllm.operators.moe import MoeOpSpec +from sparsevllm.operators.registry import OpResolver +from sparsevllm.platforms import DeviceCaps, PlatformEnum +from sparsevllm.utils.config import config_layer_get @pytest.mark.skipif(not torch.cuda.is_available(), reason="requires CUDA") def test_gemma4_router_kernels_match_torch(): + from sparsevllm.kernels.triton.gemma4_fused_router import ( + gemma4_fused_router_topk, + ) from sparsevllm.kernels.triton.gemma4_router import ( gemma4_router_input, - gemma4_router_topk, ) torch.manual_seed(19) @@ -62,14 +73,32 @@ def test_gemma4_router_kernels_match_torch(): logits = torch.randn(7, 128, dtype=torch.bfloat16, device="cuda") expert_scale = torch.randn(128, dtype=torch.bfloat16, device="cuda") - actual_weights, actual_ids = gemma4_router_topk(logits, expert_scale, 8) + actual_weights, actual_ids = gemma4_fused_router_topk(logits, expert_scale, 8) probabilities = logits.float().softmax(-1) expected_weights, expected_ids = probabilities.topk(8, dim=-1) expected_weights.div_(expected_weights.sum(-1, keepdim=True)).mul_( expert_scale[expected_ids] ) - assert torch.equal(actual_ids, expected_ids) - torch.testing.assert_close(actual_weights, expected_weights, rtol=1e-5, atol=1e-6) + actual_routes = torch.zeros_like(probabilities).scatter_( + 1, actual_ids.long(), actual_weights + ) + expected_routes = torch.zeros_like(probabilities).scatter_( + 1, expected_ids, expected_weights + ) + torch.testing.assert_close(actual_routes, expected_routes, rtol=1e-5, atol=1e-6) + assert actual_ids.dtype == torch.int32 + + for _ in range(3): + gemma4_fused_router_topk(logits, expert_scale, 8) + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph): + graph_weights, graph_ids = gemma4_fused_router_topk(logits, expert_scale, 8) + logits.copy_(torch.randn_like(logits)) + graph.replay() + replay_weights, replay_ids = graph_weights.clone(), graph_ids.clone() + graph.replay() + assert torch.equal(replay_ids, graph_ids) + assert torch.equal(replay_weights, graph_weights) @pytest.mark.skipif(not torch.cuda.is_available(), reason="requires CUDA") @@ -87,10 +116,61 @@ def test_gemma4_provider_gelu_tanh_and_mul_matches_torch(dtype, rows): def test_gemma4_moe_torch_oracle_is_not_a_production_fallback(): - assert GEMMA4_MOE_REGISTRY.providers == (TritonGemma4MoeProvider,) + assert GEMMA4_MOE_REGISTRY.providers == ( + TritonGemma4MoeProvider, + H20Gemma4MoeProvider, + ) assert TorchGemma4MoeProvider not in GEMMA4_MOE_REGISTRY.providers +@patch("sparsevllm.operators.gemma4.version", return_value="0.6.15") +@patch("sparsevllm.operators.gemma4.find_spec", return_value=object()) +def test_gemma4_h20_prefers_dedicated_flashinfer_prefill(_find, _version): + caps = DeviceCaps( + platform=PlatformEnum.CUDA, + device_type="cuda", + device_index=0, + device_name="NVIDIA H20", + compute_capability=(9, 0), + runtime_version="13.0", + supports_graph_capture=True, + supports_triton=True, + supports_bfloat16=True, + supports_native_fp8=True, + ) + resolved = OpResolver(GEMMA4_REGISTRY).resolve( + Gemma4OpSpec(torch.bfloat16, (256, 512), True), caps + ) + assert resolved.provider.name == "gemma4_h20" + assert isinstance(resolved.provider, H20Gemma4OperatorProvider) + h20_backend = resolved.provider.attention_backend(sliding_window=1024) + generic_backend = TritonGemma4OperatorProvider().attention_backend( + sliding_window=1024 + ) + assert h20_backend.use_window_decode + assert h20_backend.global_decode_heads_per_program == 4 + assert not generic_backend.use_window_decode + assert generic_backend.global_decode_heads_per_program is None + assert TorchGemma4OperatorProvider not in GEMMA4_REGISTRY.providers + moe = OpResolver(GEMMA4_MOE_REGISTRY).resolve( + MoeOpSpec( + num_experts=128, + num_local_experts=128, + hidden_size=2816, + intermediate_size=1408, + top_k=8, + activation_dtype=torch.bfloat16, + weight_dtype=torch.bfloat16, + block_shape=None, + ep_size=1, + cuda_graph=True, + activation="gelu_tanh", + ), + caps, + ) + assert isinstance(moe.provider, H20Gemma4MoeProvider) + + def _parallel_context() -> ParallelContext: group = ParallelGroup(process_group=None, ranks=(0,), rank=0, size=1) return ParallelContext(world=group, tensor=group, expert=group, data=group) @@ -145,13 +225,32 @@ def _config(**overrides) -> Gemma4TextConfig: return Gemma4TextConfig(**values) +def test_gemma4_serialized_layer_config_uses_global_defaults(): + config = { + "head_dim": 256, + "num_key_value_heads": 8, + "per_layer_config": {"05": {"head_dim": 512}}, + } + assert config_layer_get(config, 0, "head_dim") == 256 + assert config_layer_get(config, 5, "head_dim") == 512 + assert config_layer_get(config, 5, "num_key_value_heads") == 8 + + def test_gemma4_rope_matches_transformers_for_both_layer_types(): config = _config() positions = torch.arange(9) - for layer_type, head_dim in (("sliding_attention", 4), ("full_attention", 8)): + for layer_idx, (layer_type, head_dim) in enumerate( + (("sliding_attention", 4), ("full_attention", 8)) + ): actual = Gemma4RotaryEmbedding(config, layer_type, head_dim) - reference = HFGemma4RotaryEmbedding(config, layer_type=layer_type) - cos, sin = reference(torch.zeros(1), positions.unsqueeze(0), layer_type) + if "layer_type" in inspect.signature(HFGemma4RotaryEmbedding).parameters: + reference = HFGemma4RotaryEmbedding(config, layer_type=layer_type) + cos, sin = reference(torch.zeros(1), positions.unsqueeze(0), layer_type) + else: + reference = HFGemma4RotaryEmbedding(config.per_layer_config[layer_idx]) + cos, sin = reference( + torch.zeros(1), positions.unsqueeze(0), layer_type + ) torch.testing.assert_close( actual.cos_sin_cache[positions, 0, : head_dim // 2], cos[0, :, : head_dim // 2], @@ -208,11 +307,11 @@ def test_gemma4_k_eq_v_loader_duplicates_normalized_projection_slot(): attention = Gemma4Attention( config, 1, - Gemma4RotaryEmbedding(config, "full_attention", config.global_head_dim), + Gemma4RotaryEmbedding(config, "full_attention", 8), TorchGemma4OperatorProvider(), ) loaded_key = torch.randn( - config.num_global_key_value_heads * config.global_head_dim, config.hidden_size + 8, config.hidden_size ) attention.qkv_proj.weight_loader(attention.qkv_proj.weight, loaded_key, "k") q_end = attention.q_size @@ -297,12 +396,12 @@ def test_gemma4_shared_kv_attention_only_allocates_query_projection(): attention = Gemma4Attention( config, 2, - Gemma4RotaryEmbedding(config, "sliding_attention", config.head_dim), + Gemma4RotaryEmbedding(config, "sliding_attention", 4), TorchGemma4OperatorProvider(), ) assert attention.is_kv_shared_layer assert tuple(attention.qkv_proj.weight.shape) == ( - config.num_attention_heads * config.head_dim, + config.num_attention_heads * 4, config.hidden_size, ) assert not hasattr(attention, "k_norm") From c65c34670303d3e5770d3b77f1a566e3373ce0d1 Mon Sep 17 00:00:00 2001 From: QuanshengGu Date: Thu, 13 Aug 2026 18:27:25 +0800 Subject: [PATCH 08/13] docs: correct multimodal support --- docs/en/features/supported-models.md | 2 +- docs/zh/features/supported-models.md | 2 +- 2 files changed, 2 insertions(+), 2 deletions(-) diff --git a/docs/en/features/supported-models.md b/docs/en/features/supported-models.md index bb5326f2..7ed25b7c 100644 --- a/docs/en/features/supported-models.md +++ b/docs/en/features/supported-models.md @@ -96,6 +96,6 @@ explicitly during admission. | Model family | Image | Video | Audio | | --- | :---: | :---: | :---: | | Qwen3.5 / Qwen3.6 Dense and MoE | ✅ | ✅ | — | -| Gemma 4 Dense and MoE | ✅ | ✅ | Checkpoint dependent | +| Gemma 4 Dense and MoE | ✅ | ✅ | — | `—` means that the combination is not currently supported. diff --git a/docs/zh/features/supported-models.md b/docs/zh/features/supported-models.md index bac80a78..0f9853f0 100644 --- a/docs/zh/features/supported-models.md +++ b/docs/zh/features/supported-models.md @@ -85,6 +85,6 @@ checkpoint 自身的 processor 和 chat template;不受支持的媒体会在 | 模型家族 | 图片 | 视频 | 音频 | | --- | :---: | :---: | :---: | | Qwen3.5 / Qwen3.6 Dense 与 MoE | ✅ | ✅ | — | -| Gemma 4 Dense 与 MoE | ✅ | ✅ | 取决于 checkpoint | +| Gemma 4 Dense 与 MoE | ✅ | ✅ | — | `—` 表示当前不支持该组合。 From 11102e77ba5c670aa759c29af847a81f80561d91 Mon Sep 17 00:00:00 2001 From: QuanshengGu Date: Thu, 13 Aug 2026 18:44:17 +0800 Subject: [PATCH 09/13] fix: harden gemma operator selection --- src/sparsevllm/models/gemma4.py | 85 ++++++++++---- src/sparsevllm/operators/gemma4.py | 26 ----- src/sparsevllm/operators/gemma4_router.py | 132 ++++++++++++++++++++++ tests/test_gemma4_model.py | 75 +++++++++++- 4 files changed, 268 insertions(+), 50 deletions(-) create mode 100644 src/sparsevllm/operators/gemma4_router.py diff --git a/src/sparsevllm/models/gemma4.py b/src/sparsevllm/models/gemma4.py index 9e1da5ec..2c096958 100644 --- a/src/sparsevllm/models/gemma4.py +++ b/src/sparsevllm/models/gemma4.py @@ -26,6 +26,11 @@ Gemma4OpSpec, resolve_gemma4_provider, ) +from sparsevllm.operators.gemma4_router import ( + Gemma4RouterOpSpec, + Gemma4RouterProvider, + resolve_gemma4_router_provider, +) from sparsevllm.operators.gemma4_moe import Gemma4PackedExperts from sparsevllm.operators.moe import model_activation_dtype from sparsevllm.platforms import device_runtime @@ -299,6 +304,7 @@ def __init__( self, config: Gemma4TextConfig, operator_provider: Gemma4OperatorProvider, + router_provider: Gemma4RouterProvider, ) -> None: super().__init__() self.top_k = int(config.top_k_experts) @@ -310,6 +316,7 @@ def __init__( provider=operator_provider, ) self._ops = operator_provider + self._router_ops = router_provider self.scale = nn.Parameter(torch.ones(config.hidden_size)) self.proj = ReplicatedLinear(config.hidden_size, config.num_experts) self.per_expert_scale = nn.Parameter(torch.ones(config.num_experts)) @@ -318,7 +325,7 @@ def forward(self, hidden_states: torch.Tensor) -> tuple[torch.Tensor, torch.Tens router_input = self._ops.router_input( hidden_states, self.scale, self.root_size, self.norm.eps ) - return self._ops.router_topk( + return self._router_ops.topk( self.proj(router_input), self.per_expert_scale, self.top_k ) @@ -329,14 +336,14 @@ def __init__( config: Gemma4TextConfig, layer_idx: int, operator_provider: Gemma4OperatorProvider, - rotary_embeddings: nn.ModuleDict, + router_provider: Gemma4RouterProvider | None, + rotary_emb: Gemma4RotaryEmbedding, ) -> None: super().__init__() - layer_type = str(config.layer_types[layer_idx]) self.self_attn = Gemma4Attention( config, layer_idx, - rotary_embeddings[layer_type], + rotary_emb, operator_provider, ) self._ops = operator_provider @@ -371,7 +378,9 @@ def __init__( self.enable_moe_block = bool(config.enable_moe_block) if self.enable_moe_block: self.parallel_context = get_parallel_context() - self.router = Gemma4Router(config, operator_provider) + if router_provider is None: + raise RuntimeError("Gemma 4 MoE requires a router provider.") + self.router = Gemma4Router(config, operator_provider, router_provider) self.experts = Gemma4PackedExperts(config) self.post_feedforward_layernorm_1 = Gemma4RMSNorm( config.hidden_size, @@ -444,6 +453,7 @@ def __init__( self, config: Gemma4TextConfig, operator_provider: Gemma4OperatorProvider, + router_provider: Gemma4RouterProvider | None = None, ) -> None: super().__init__() self.config = config @@ -470,29 +480,46 @@ def __init__( ) self.per_layer_model_projection_scale = float(config.hidden_size) ** -0.5 self.per_layer_input_scale = 2.0**-0.5 - self.rotary_embeddings = nn.ModuleDict( - { - layer_type: Gemma4RotaryEmbedding( + rotary_keys = [] + rotary_embeddings = {} + rotary_signatures = {} + for layer_idx, layer_type in enumerate(config.layer_types): + head_dim = int( + config_layer_get( config, - layer_type, - config_layer_get( - config, - config.layer_types.index(layer_type), - "head_dim", - "head_dim" - if layer_type == "sliding_attention" - else "global_head_dim", - ), + layer_idx, + "head_dim", + "head_dim" + if layer_type == "sliding_attention" + else "global_head_dim", ) - for layer_type in set(config.layer_types) - } - ) + ) + signature = ( + str(layer_type), + head_dim, + tuple( + sorted( + (key, repr(value)) + for key, value in config.rope_parameters[layer_type].items() + ) + ), + ) + key = rotary_signatures.setdefault( + signature, f"rope_{len(rotary_signatures)}" + ) + rotary_keys.append(key) + if key not in rotary_embeddings: + rotary_embeddings[key] = Gemma4RotaryEmbedding( + config, str(layer_type), head_dim + ) + self.rotary_embeddings = nn.ModuleDict(rotary_embeddings) self.layers = nn.ModuleList( Gemma4DecoderLayer( config, layer_idx, operator_provider, - self.rotary_embeddings, + router_provider, + self.rotary_embeddings[rotary_keys[layer_idx]], ) for layer_idx in range(config.num_hidden_layers) ) @@ -572,10 +599,11 @@ def __init__( self, config: Gemma4TextConfig, operator_provider: Gemma4OperatorProvider, + router_provider: Gemma4RouterProvider | None = None, ) -> None: super().__init__() self.config = config - self.model = Gemma4Model(config, operator_provider) + self.model = Gemma4Model(config, operator_provider, router_provider) self.lm_head = ParallelLMHead(config.vocab_size, config.hidden_size) if config.tie_word_embeddings: self.lm_head.weight.data = self.model.embed_tokens.weight.data @@ -618,7 +646,7 @@ def build_runtime_kwargs(cls, config, *, device, **_): } ) ) - return { + runtime_kwargs = { "operator_provider": resolve_gemma4_provider( Gemma4OpSpec( activation_dtype=model_activation_dtype(config), @@ -628,6 +656,17 @@ def build_runtime_kwargs(cls, config, *, device, **_): device_index=device.index, ) } + if bool(config.enable_moe_block): + runtime_kwargs["router_provider"] = resolve_gemma4_router_provider( + Gemma4RouterOpSpec( + activation_dtype=model_activation_dtype(config), + num_experts=int(config.num_experts), + top_k=int(config.top_k_experts), + cuda_graph=bool(getattr(config, "decode_cuda_graph", False)), + ), + device_index=device.index, + ) + return runtime_kwargs def map_weight_name(self, source_weight_name: str) -> str | None: match = _EXPERT_SOURCE_RE.match(source_weight_name) diff --git a/src/sparsevllm/operators/gemma4.py b/src/sparsevllm/operators/gemma4.py index a93d383f..8a1a84a2 100644 --- a/src/sparsevllm/operators/gemma4.py +++ b/src/sparsevllm/operators/gemma4.py @@ -63,14 +63,6 @@ def router_input( ) -> torch.Tensor: raise NotImplementedError - def router_topk( - self, - logits: torch.Tensor, - per_expert_scale: torch.Tensor, - top_k: int, - ) -> tuple[torch.Tensor, torch.Tensor]: - raise NotImplementedError - def gelu_tanh_and_mul(self, x: torch.Tensor) -> torch.Tensor: raise NotImplementedError @@ -137,11 +129,6 @@ def router_input(self, hidden_states, scale, root_size, eps): return gemma4_router_input(hidden_states, scale, root_size, eps) - def router_topk(self, logits, per_expert_scale, top_k): - from sparsevllm.kernels.triton.gemma4_router import gemma4_router_topk - - return gemma4_router_topk(logits, per_expert_scale, top_k) - def gelu_tanh_and_mul(self, x): from sparsevllm.kernels.triton.gemma4_gelu_and_mul import ( gelu_tanh_and_mul_fwd, @@ -206,13 +193,6 @@ def __init__(self) -> None: self._prefill = Gemma4FlashInferPrefill() - def router_topk(self, logits, per_expert_scale, top_k): - from sparsevllm.kernels.triton.gemma4_fused_router import ( - gemma4_fused_router_topk, - ) - - return gemma4_fused_router_topk(logits, per_expert_scale, top_k) - def attention_backend(self, *, sliding_window: int | None): from sparsevllm.operators.gemma4_attention import Gemma4AttentionBackend @@ -255,12 +235,6 @@ def qkv_norm_rope( def router_input(self, hidden_states, scale, root_size, eps): return self.rmsnorm(hidden_states, None, eps) * scale * root_size - def router_topk(self, logits, per_expert_scale, top_k): - probabilities = F.softmax(logits, dim=-1, dtype=torch.float32) - weights, ids = probabilities.topk(top_k, dim=-1) - weights.div_(weights.sum(-1, keepdim=True)).mul_(per_expert_scale[ids]) - return weights, ids - def gelu_tanh_and_mul(self, x): gate, up = x.chunk(2, -1) return F.gelu(gate, approximate="tanh") * up diff --git a/src/sparsevllm/operators/gemma4_router.py b/src/sparsevllm/operators/gemma4_router.py new file mode 100644 index 00000000..9a6b17c8 --- /dev/null +++ b/src/sparsevllm/operators/gemma4_router.py @@ -0,0 +1,132 @@ +from __future__ import annotations + +from dataclasses import dataclass + +import torch + +import sparsevllm.platforms as platforms +from sparsevllm.operators.registry import ( + OpRegistry, + OpResolver, + SupportResult, + runtime_version_at_least, +) +from sparsevllm.platforms.interface import DeviceCaps, PlatformEnum + + +@dataclass(frozen=True) +class Gemma4RouterOpSpec: + activation_dtype: torch.dtype + num_experts: int + top_k: int + cuda_graph: bool + + def __post_init__(self) -> None: + if self.num_experts <= 0 or not 0 < self.top_k <= self.num_experts: + raise ValueError( + "Gemma 4 router requires 0 < top_k <= num_experts, got " + f"top_k={self.top_k}, num_experts={self.num_experts}." + ) + + +class Gemma4RouterProvider: + name = "" + priority = 0 + + def topk( + self, + logits: torch.Tensor, + per_expert_scale: torch.Tensor, + top_k: int, + ) -> tuple[torch.Tensor, torch.Tensor]: + raise NotImplementedError + + +GEMMA4_ROUTER_REGISTRY: OpRegistry[Gemma4RouterOpSpec, Gemma4RouterProvider] = ( + OpRegistry("Gemma 4 router") +) + + +@GEMMA4_ROUTER_REGISTRY.register +class TritonGemma4RouterProvider(Gemma4RouterProvider): + name = "triton" + priority = 10 + + @classmethod + def supports(cls, spec: Gemma4RouterOpSpec, caps: DeviceCaps) -> SupportResult: + if caps.platform != PlatformEnum.CUDA or not caps.supports_triton: + return SupportResult.no("requires CUDA with Triton") + if spec.cuda_graph and not caps.supports_graph_capture: + return SupportResult.no("device does not support CUDA Graph capture") + if spec.activation_dtype not in {torch.bfloat16, torch.float16}: + return SupportResult.no("requires BF16 or FP16 activations") + return SupportResult.yes() + + def topk(self, logits, per_expert_scale, top_k): + from sparsevllm.kernels.triton.gemma4_router import gemma4_router_topk + + return gemma4_router_topk(logits, per_expert_scale, top_k) + + +@GEMMA4_ROUTER_REGISTRY.register +class H20Gemma4RouterProvider(TritonGemma4RouterProvider): + name = "gemma4_h20" + priority = 100 + + @classmethod + def supports(cls, spec: Gemma4RouterOpSpec, caps: DeviceCaps) -> SupportResult: + generic = super().supports(spec, caps) + if not generic.supported: + return generic + if caps.compute_capability != (9, 0) or caps.device_name != "NVIDIA H20": + return SupportResult.no( + "requires profiled NVIDIA H20 SM90 hardware, " + f"got {caps.device_name} {caps.compute_capability}" + ) + if not runtime_version_at_least(caps.runtime_version, (12, 8)): + return SupportResult.no( + f"requires CUDA runtime >= 12.8, got {caps.runtime_version or 'unknown'}" + ) + if spec.num_experts > 1024: + return SupportResult.no("fused router requires at most 1024 experts") + return SupportResult.yes() + + def topk(self, logits, per_expert_scale, top_k): + from sparsevllm.kernels.triton.gemma4_fused_router import ( + gemma4_fused_router_topk, + ) + + return gemma4_fused_router_topk(logits, per_expert_scale, top_k) + + +class TorchGemma4RouterProvider(Gemma4RouterProvider): + """Explicit correctness oracle; never selected for production inference.""" + + name = "torch_oracle" + + def topk(self, logits, per_expert_scale, top_k): + probabilities = torch.softmax(logits, dim=-1, dtype=torch.float32) + weights, ids = probabilities.topk(top_k, dim=-1) + weights.div_(weights.sum(-1, keepdim=True)).mul_(per_expert_scale[ids]) + return weights, ids + + +def resolve_gemma4_router_provider( + spec: Gemma4RouterOpSpec, *, device_index: int | None = None +) -> Gemma4RouterProvider: + platform = platforms.current_platform + if device_index is None: + device_index = torch.cuda.current_device() if platform.is_cuda_alike() else 0 + caps = platform.get_device_caps(int(device_index)) + return OpResolver(GEMMA4_ROUTER_REGISTRY).resolve(spec, caps).provider + + +__all__ = [ + "GEMMA4_ROUTER_REGISTRY", + "Gemma4RouterOpSpec", + "Gemma4RouterProvider", + "H20Gemma4RouterProvider", + "TorchGemma4RouterProvider", + "TritonGemma4RouterProvider", + "resolve_gemma4_router_provider", +] diff --git a/tests/test_gemma4_model.py b/tests/test_gemma4_model.py index cdd36011..6f1585e8 100644 --- a/tests/test_gemma4_model.py +++ b/tests/test_gemma4_model.py @@ -47,6 +47,13 @@ TorchGemma4MoeProvider, TritonGemma4MoeProvider, ) +from sparsevllm.operators.gemma4_router import ( + GEMMA4_ROUTER_REGISTRY, + Gemma4RouterOpSpec, + H20Gemma4RouterProvider, + TorchGemma4RouterProvider, + TritonGemma4RouterProvider, +) from sparsevllm.operators.moe import MoeOpSpec from sparsevllm.operators.registry import OpResolver from sparsevllm.platforms import DeviceCaps, PlatformEnum @@ -171,6 +178,53 @@ def test_gemma4_h20_prefers_dedicated_flashinfer_prefill(_find, _version): assert isinstance(moe.provider, H20Gemma4MoeProvider) +@pytest.mark.parametrize( + ("num_experts", "provider_type"), + ((1024, H20Gemma4RouterProvider), (1025, TritonGemma4RouterProvider)), +) +def test_gemma4_h20_router_resolves_expert_boundary(num_experts, provider_type): + caps = DeviceCaps( + platform=PlatformEnum.CUDA, + device_type="cuda", + device_index=0, + device_name="NVIDIA H20", + compute_capability=(9, 0), + runtime_version="13.0", + supports_graph_capture=True, + supports_triton=True, + supports_bfloat16=True, + supports_native_fp8=True, + ) + resolved = OpResolver(GEMMA4_ROUTER_REGISTRY).resolve( + Gemma4RouterOpSpec(torch.bfloat16, num_experts, 8, True), caps + ) + assert isinstance(resolved.provider, provider_type) + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="requires CUDA") +@pytest.mark.parametrize("num_experts", [1024, 1025]) +def test_gemma4_h20_router_expert_boundary_matches_torch(num_experts): + provider_type = ( + H20Gemma4RouterProvider + if num_experts == 1024 + else TritonGemma4RouterProvider + ) + torch.manual_seed(41) + logits = torch.randn(3, num_experts, dtype=torch.bfloat16, device="cuda") + scales = torch.randn(num_experts, dtype=torch.bfloat16, device="cuda") + weights, ids = provider_type().topk(logits, scales, 8) + probabilities = logits.float().softmax(-1) + expected_weights, expected_ids = probabilities.topk(8, dim=-1) + expected_weights.div_(expected_weights.sum(-1, keepdim=True)).mul_( + scales[expected_ids] + ) + actual = torch.zeros_like(probabilities).scatter_(1, ids.long(), weights) + expected = torch.zeros_like(probabilities).scatter_( + 1, expected_ids, expected_weights + ) + torch.testing.assert_close(actual, expected, rtol=1e-5, atol=1e-6) + + def _parallel_context() -> ParallelContext: group = ParallelGroup(process_group=None, ranks=(0,), rank=0, size=1) return ParallelContext(world=group, tensor=group, expert=group, data=group) @@ -261,6 +315,23 @@ def test_gemma4_rope_matches_transformers_for_both_layer_types(): ) +def test_gemma4_rope_cache_separates_same_type_head_dims(): + config = _config( + num_hidden_layers=3, + layer_types=["sliding_attention", "sliding_attention", "full_attention"], + per_layer_config={ + "0": {"head_dim": 4}, + "1": {"head_dim": 8}, + "2": {"head_dim": 8}, + }, + ) + with _patch_parallel_context(): + model = Gemma4Model(config, TorchGemma4OperatorProvider()) + assert len(model.rotary_embeddings) == 3 + assert model.layers[0].self_attn.rotary_emb.cos_sin_cache.shape[-1] == 4 + assert model.layers[1].self_attn.rotary_emb.cos_sin_cache.shape[-1] == 8 + + def test_gemma4_dense_mlp_matches_transformers(): config = _config() with _patch_parallel_context(): @@ -287,7 +358,9 @@ def test_gemma4_router_matches_transformers(): enable_moe_block=True, num_experts=4, top_k_experts=2, moe_intermediate_size=4 ) with _patch_parallel_context(): - actual = Gemma4Router(config, TorchGemma4OperatorProvider()) + actual = Gemma4Router( + config, TorchGemma4OperatorProvider(), TorchGemma4RouterProvider() + ) reference = HFGemma4Router(config) torch.manual_seed(5) reference.proj.weight.data.normal_(0, 0.2) From cc8d1548616f0c5206022af92fd254d5c9830a82 Mon Sep 17 00:00:00 2001 From: QuanshengGu Date: Thu, 13 Aug 2026 18:49:19 +0800 Subject: [PATCH 10/13] fix: honor per-layer rope settings --- src/sparsevllm/models/gemma4.py | 19 +++++++++------ tests/test_gemma4_model.py | 42 +++++++++++++++++++++++++++++++++ 2 files changed, 54 insertions(+), 7 deletions(-) diff --git a/src/sparsevllm/models/gemma4.py b/src/sparsevllm/models/gemma4.py index 2c096958..afcaa6f9 100644 --- a/src/sparsevllm/models/gemma4.py +++ b/src/sparsevllm/models/gemma4.py @@ -48,10 +48,16 @@ class Gemma4RotaryEmbedding(nn.Module): def __init__( - self, config: Gemma4TextConfig, layer_type: str, head_dim: int + self, + config: Gemma4TextConfig, + layer_type: str, + head_dim: int, + parameters: dict | None = None, ) -> None: super().__init__() - parameters = dict(config.rope_parameters[layer_type]) + if parameters is None: + parameters = config_layer_get(config, 0, "rope_parameters") + parameters = dict(parameters.get(layer_type, parameters)) rope_type = str(parameters.get("rope_type", "default")) if rope_type not in {"default", "proportional"}: raise NotImplementedError(f"Unsupported Gemma 4 RoPE type {rope_type!r}.") @@ -494,14 +500,13 @@ def __init__( else "global_head_dim", ) ) + rope_parameters = config_layer_get(config, layer_idx, "rope_parameters") + parameters = dict(rope_parameters.get(layer_type, rope_parameters)) signature = ( str(layer_type), head_dim, tuple( - sorted( - (key, repr(value)) - for key, value in config.rope_parameters[layer_type].items() - ) + sorted((key, repr(value)) for key, value in parameters.items()) ), ) key = rotary_signatures.setdefault( @@ -510,7 +515,7 @@ def __init__( rotary_keys.append(key) if key not in rotary_embeddings: rotary_embeddings[key] = Gemma4RotaryEmbedding( - config, str(layer_type), head_dim + config, str(layer_type), head_dim, parameters ) self.rotary_embeddings = nn.ModuleDict(rotary_embeddings) self.layers = nn.ModuleList( diff --git a/tests/test_gemma4_model.py b/tests/test_gemma4_model.py index 6f1585e8..4985030c 100644 --- a/tests/test_gemma4_model.py +++ b/tests/test_gemma4_model.py @@ -332,6 +332,48 @@ def test_gemma4_rope_cache_separates_same_type_head_dims(): assert model.layers[1].self_attn.rotary_emb.cos_sin_cache.shape[-1] == 8 +def test_gemma4_rope_cache_separates_per_layer_parameters(): + parameters = { + "sliding_attention": {"rope_type": "default", "rope_theta": 10000.0}, + "full_attention": { + "rope_type": "proportional", + "rope_theta": 1000000.0, + "partial_rotary_factor": 0.25, + }, + } + second_parameters = { + **parameters, + "sliding_attention": {"rope_type": "default", "rope_theta": 20000.0}, + } + config = _config( + num_hidden_layers=3, + layer_types=["sliding_attention", "sliding_attention", "full_attention"], + allow_global_per_layer_attribute_access=True, + per_layer_config={ + "0": {"head_dim": 4, "rope_parameters": parameters}, + "1": { + "head_dim": 4, + "rope_parameters": second_parameters, + }, + "2": {"head_dim": 8, "rope_parameters": parameters}, + }, + ) + config.allow_global_per_layer_attribute_access = False + with _patch_parallel_context(): + model = Gemma4Model(config, TorchGemma4OperatorProvider()) + first = model.layers[0].self_attn.rotary_emb.cos_sin_cache + second = model.layers[1].self_attn.rotary_emb.cos_sin_cache + assert len(model.rotary_embeddings) == 3 + assert not torch.equal(first, second) + expected = Gemma4RotaryEmbedding( + config, + "sliding_attention", + 4, + second_parameters, + ) + torch.testing.assert_close(second, expected.cos_sin_cache) + + def test_gemma4_dense_mlp_matches_transformers(): config = _config() with _patch_parallel_context(): From 972541d2ace532ccf65e811554fbe29c4bbbd2d1 Mon Sep 17 00:00:00 2001 From: QuanshengGu Date: Thu, 13 Aug 2026 22:35:49 +0800 Subject: [PATCH 11/13] fix: declare multimodal vision runtime --- pyproject.toml | 6 ++++++ tests/test_dependency_constraints.py | 6 ++++++ 2 files changed, 12 insertions(+) diff --git a/pyproject.toml b/pyproject.toml index 3c111c58..fd1b09a2 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -46,11 +46,13 @@ dependencies = [ [project.optional-dependencies] cu129 = [ "torch==2.11.0", + "torchvision==0.26.0", "flashinfer-python[cu12]>=0.6.15,<0.7", "flashinfer-jit-cache>=0.6.15,<0.7", ] cu130 = [ "torch==2.11.0", + "torchvision==0.26.0", "flashinfer-python[cu13]>=0.6.15,<0.7", "flashinfer-jit-cache>=0.6.15,<0.7", ] @@ -80,6 +82,10 @@ torch = [ { index = "pytorch-cu129", extra = "cu129" }, { index = "pytorch-cu130", extra = "cu130" }, ] +torchvision = [ + { index = "pytorch-cu129", extra = "cu129" }, + { index = "pytorch-cu130", extra = "cu130" }, +] flashinfer-jit-cache = [ { index = "flashinfer-cu129", extra = "cu129" }, { index = "flashinfer-cu130", extra = "cu130" }, diff --git a/tests/test_dependency_constraints.py b/tests/test_dependency_constraints.py index 1e3187e0..80656cb4 100644 --- a/tests/test_dependency_constraints.py +++ b/tests/test_dependency_constraints.py @@ -52,11 +52,13 @@ def test_uv_routes_cuda_packages_to_explicit_indexes(): assert config["project"]["optional-dependencies"] == { "cu129": [ "torch==2.11.0", + "torchvision==0.26.0", "flashinfer-python[cu12]>=0.6.15,<0.7", "flashinfer-jit-cache>=0.6.15,<0.7", ], "cu130": [ "torch==2.11.0", + "torchvision==0.26.0", "flashinfer-python[cu13]>=0.6.15,<0.7", "flashinfer-jit-cache>=0.6.15,<0.7", ], @@ -66,6 +68,10 @@ def test_uv_routes_cuda_packages_to_explicit_indexes(): {"index": "pytorch-cu129", "extra": "cu129"}, {"index": "pytorch-cu130", "extra": "cu130"}, ], + "torchvision": [ + {"index": "pytorch-cu129", "extra": "cu129"}, + {"index": "pytorch-cu130", "extra": "cu130"}, + ], "flashinfer-jit-cache": [ {"index": "flashinfer-cu129", "extra": "cu129"}, {"index": "flashinfer-cu130", "extra": "cu130"}, From f32325cd5025bddc761d3c043a8998433f108a82 Mon Sep 17 00:00:00 2001 From: kuma_Gu2006 Date: Thu, 13 Aug 2026 22:41:44 +0800 Subject: [PATCH 12/13] fix: use package-managed torchvision --- pyproject.toml | 7 +------ tests/test_dependency_constraints.py | 18 ++++++++---------- 2 files changed, 9 insertions(+), 16 deletions(-) diff --git a/pyproject.toml b/pyproject.toml index fd1b09a2..8d3cbec4 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -21,6 +21,7 @@ dependencies = [ "apache-tvm-ffi==0.1.10", "nvidia-cutlass-dsl>=4.6,<5", "pillow", + "torchvision", "einops", "sglang-kernel>=0.4.5,<0.4.6", "tqdm", @@ -46,13 +47,11 @@ dependencies = [ [project.optional-dependencies] cu129 = [ "torch==2.11.0", - "torchvision==0.26.0", "flashinfer-python[cu12]>=0.6.15,<0.7", "flashinfer-jit-cache>=0.6.15,<0.7", ] cu130 = [ "torch==2.11.0", - "torchvision==0.26.0", "flashinfer-python[cu13]>=0.6.15,<0.7", "flashinfer-jit-cache>=0.6.15,<0.7", ] @@ -82,10 +81,6 @@ torch = [ { index = "pytorch-cu129", extra = "cu129" }, { index = "pytorch-cu130", extra = "cu130" }, ] -torchvision = [ - { index = "pytorch-cu129", extra = "cu129" }, - { index = "pytorch-cu130", extra = "cu130" }, -] flashinfer-jit-cache = [ { index = "flashinfer-cu129", extra = "cu129" }, { index = "flashinfer-cu130", extra = "cu130" }, diff --git a/tests/test_dependency_constraints.py b/tests/test_dependency_constraints.py index 80656cb4..2fb781ca 100644 --- a/tests/test_dependency_constraints.py +++ b/tests/test_dependency_constraints.py @@ -17,10 +17,14 @@ def test_runtime_compatibility_bounds_cover_canonical_lock(): assert "transformers>=5.13,<6" in dependencies assert "nvidia-cutlass-dsl>=4.6,<5" in dependencies assert "sglang-kernel>=0.4.5,<0.4.6" in dependencies - assert {"fire", "pillow", "einops", "tqdm", "loguru"} <= dependencies - assert not any( - dependency.startswith("torchvision") for dependency in dependencies - ) + assert { + "fire", + "pillow", + "torchvision", + "einops", + "tqdm", + "loguru", + } <= dependencies assert not any( dependency.startswith("flashinfer-cubin") for dependency in dependencies @@ -52,13 +56,11 @@ def test_uv_routes_cuda_packages_to_explicit_indexes(): assert config["project"]["optional-dependencies"] == { "cu129": [ "torch==2.11.0", - "torchvision==0.26.0", "flashinfer-python[cu12]>=0.6.15,<0.7", "flashinfer-jit-cache>=0.6.15,<0.7", ], "cu130": [ "torch==2.11.0", - "torchvision==0.26.0", "flashinfer-python[cu13]>=0.6.15,<0.7", "flashinfer-jit-cache>=0.6.15,<0.7", ], @@ -68,10 +70,6 @@ def test_uv_routes_cuda_packages_to_explicit_indexes(): {"index": "pytorch-cu129", "extra": "cu129"}, {"index": "pytorch-cu130", "extra": "cu130"}, ], - "torchvision": [ - {"index": "pytorch-cu129", "extra": "cu129"}, - {"index": "pytorch-cu130", "extra": "cu130"}, - ], "flashinfer-jit-cache": [ {"index": "flashinfer-cu129", "extra": "cu129"}, {"index": "flashinfer-cu130", "extra": "cu130"}, From b5a8f0bd647cf2349470f01ff046c19f0a664708 Mon Sep 17 00:00:00 2001 From: QuanshengGu Date: Thu, 13 Aug 2026 22:47:41 +0800 Subject: [PATCH 13/13] chore: remove fixed token benchmark --- benchmark/fixed_token_microbench.py | 309 ---------------------------- 1 file changed, 309 deletions(-) delete mode 100644 benchmark/fixed_token_microbench.py diff --git a/benchmark/fixed_token_microbench.py b/benchmark/fixed_token_microbench.py deleted file mode 100644 index f4ac19b8..00000000 --- a/benchmark/fixed_token_microbench.py +++ /dev/null @@ -1,309 +0,0 @@ -"""Reproducible fixed-token latency benchmark for Sparse-vLLM and vLLM.""" - -from __future__ import annotations - -import argparse -import json -import os -import shlex -import statistics -import subprocess -import sys -from datetime import datetime -from importlib.metadata import PackageNotFoundError, version -from pathlib import Path -from time import perf_counter -from typing import Any - - -def _parser() -> argparse.ArgumentParser: - parser = argparse.ArgumentParser(description=__doc__) - parser.add_argument("--backend", choices=("sparsevllm", "vllm"), required=True) - parser.add_argument("--model-path", required=True) - parser.add_argument("--output-dir", required=True) - parser.add_argument("--input-len", type=int, default=128) - parser.add_argument("--output-len", type=int, default=32) - parser.add_argument("--batch-size", type=int, default=16) - parser.add_argument("--num-warmups", type=int, default=3) - parser.add_argument("--num-iters", type=int, default=5) - parser.add_argument("--tensor-parallel-size", type=int, default=1) - parser.add_argument("--expert-parallel-size", type=int, default=1) - parser.add_argument("--gpu-memory-utilization", type=float, default=0.8) - parser.add_argument( - "--nsys-iteration", - type=int, - default=-1, - help="Wrap one timed iteration in an NVTX range for nsys capture.", - ) - return parser - - -def _write_json(path: Path, payload: Any) -> None: - path.write_text( - json.dumps(payload, indent=2, sort_keys=True) + "\n", encoding="utf-8" - ) - - -def _write_jsonl(path: Path, rows: list[dict[str, Any]]) -> None: - with path.open("w", encoding="utf-8") as output: - for row in rows: - output.write(json.dumps(row, sort_keys=True) + "\n") - - -def _git(*args: str) -> str | None: - result = subprocess.run( - ["git", *args], - capture_output=True, - check=False, - cwd=Path(__file__).parents[1], - text=True, - ) - return result.stdout.strip() or None - - -def _package_version(name: str) -> str | None: - try: - return version(name) - except PackageNotFoundError: - return None - - -def _validate(args: argparse.Namespace) -> Path: - for name in ( - "input_len", - "output_len", - "batch_size", - "num_warmups", - "num_iters", - "tensor_parallel_size", - "expert_parallel_size", - ): - if int(getattr(args, name)) <= 0: - raise ValueError(f"{name} must be positive") - if not 0 < args.gpu_memory_utilization <= 1: - raise ValueError("gpu_memory_utilization must be in (0, 1]") - if not -1 <= args.nsys_iteration < args.num_iters: - raise ValueError("nsys_iteration must be -1 or a timed iteration index") - if args.backend == "vllm" and args.expert_parallel_size not in { - 1, - args.tensor_parallel_size, - }: - raise ValueError( - "vLLM expert parallelism is disabled with EP=1 or spans TP ranks " - "with EP=TP" - ) - output_dir = Path(args.output_dir).expanduser().resolve() - if output_dir.exists() and any(output_dir.iterdir()): - raise FileExistsError(f"output directory is not empty: {output_dir}") - output_dir.mkdir(parents=True, exist_ok=True) - return output_dir - - -def _build_engine(args: argparse.Namespace): - common = { - "model": str(Path(args.model_path).expanduser().resolve()), - "max_model_len": args.input_len + args.output_len, - "gpu_memory_utilization": args.gpu_memory_utilization, - "tensor_parallel_size": args.tensor_parallel_size, - "enforce_eager": False, - "enable_prefix_caching": False, - } - if args.backend == "vllm": - import vllm.config.model as vllm_model_config - from vllm import LLM - - get_config = vllm_model_config.get_config - - def get_compatible_config(*config_args, **config_kwargs): - config = get_config(*config_args, **config_kwargs) - text = getattr(config, "text_config", config) - layers = getattr(text, "per_layer_config", None) - if str(getattr(config, "model_type", "")) == "gemma4" and layers: - text.allow_global_per_layer_attribute_access = True - full_idx = text.layer_types.index("full_attention") - text.global_head_dim = layers[full_idx].head_dim - text.num_global_key_value_heads = layers[full_idx].num_key_value_heads - return config - - vllm_model_config.get_config = get_compatible_config - - return LLM( - **common, - max_num_seqs=args.batch_size, - max_num_batched_tokens=max(4096, args.batch_size * args.input_len), - enable_expert_parallel=args.expert_parallel_size > 1, - enable_flashinfer_autotune=False, - language_model_only=True, - trust_remote_code=True, - ) - from sparsevllm import LLM - - return LLM( - common.pop("model"), - **common, - expert_parallel_size=args.expert_parallel_size, - max_num_seqs_in_batch=args.batch_size, - max_decoding_seqs=args.batch_size, - max_num_seqs_in_gpu=args.batch_size, - max_num_batched_tokens=args.batch_size * args.input_len, - decode_cuda_graph=True, - enable_multimodal=False, - ) - - -def _sampling_params(args: argparse.Namespace): - if args.backend == "vllm": - from vllm import SamplingParams - else: - from sparsevllm import SamplingParams - return SamplingParams(temperature=0.0, max_tokens=args.output_len, ignore_eos=True) - - -def _token_ids(args: argparse.Namespace, output: Any) -> list[int]: - if args.backend == "vllm": - return list(output.outputs[0].token_ids) - return list(output["token_ids"]) - - -def main() -> int: - args = _parser().parse_args() - output_dir = _validate(args) - run_info = { - "benchmark": "fixed_token_microbench", - "backend": args.backend, - "command": shlex.join(sys.argv), - "created_at": datetime.now().astimezone().isoformat(timespec="seconds"), - "git": { - "branch": _git("branch", "--show-current"), - "commit": _git("rev-parse", "HEAD"), - "dirty": bool(_git("status", "--porcelain")), - }, - "workload": { - "batch_size": args.batch_size, - "input_len": args.input_len, - "output_len": args.output_len, - "num_warmups": args.num_warmups, - "num_iters": args.num_iters, - "prompt": "[2] + [100 + (request + position) % 1000]", - "temperature": 0.0, - "ignore_eos": True, - }, - "topology": { - "tensor_parallel_size": args.tensor_parallel_size, - "expert_parallel_size": args.expert_parallel_size, - "cuda_graph": True, - "prefix_cache": False, - }, - "profiler": { - "kind": "nsys_nvtx" if args.nsys_iteration >= 0 else None, - "iteration": args.nsys_iteration if args.nsys_iteration >= 0 else None, - }, - "environment": { - "cuda_visible_devices": os.getenv("CUDA_VISIBLE_DEVICES"), - "python": sys.version, - "torch": _package_version("torch"), - "transformers": _package_version("transformers"), - "flashinfer_python": _package_version("flashinfer-python"), - "triton": _package_version("triton"), - "vllm": _package_version("vllm"), - }, - } - _write_json(output_dir / "run_info.json", run_info) - prompts = [ - [ - 2, - *( - 100 + (request + position) % 1000 - for position in range(args.input_len - 1) - ), - ] - for request in range(args.batch_size) - ] - engine = None - raw_outputs: list[dict[str, Any]] = [] - sample_results: list[dict[str, Any]] = [] - performance: list[dict[str, Any]] = [] - try: - engine = _build_engine(args) - params = _sampling_params(args) - for _ in range(args.num_warmups): - outputs = engine.generate(prompts, params, use_tqdm=False) - if len(outputs) != args.batch_size: - raise RuntimeError(f"warmup returned {len(outputs)} requests") - for iteration in range(args.num_iters): - profiling = iteration == args.nsys_iteration - if profiling: - import torch - - torch.cuda.nvtx.range_push(f"fixed_token_iteration_{iteration}") - torch.cuda.cudart().cudaProfilerStart() - started = perf_counter() - try: - outputs = engine.generate(prompts, params, use_tqdm=False) - elapsed = perf_counter() - started - finally: - if profiling: - torch.cuda.cudart().cudaProfilerStop() - torch.cuda.nvtx.range_pop() - if len(outputs) != args.batch_size: - raise RuntimeError( - f"iteration {iteration} returned {len(outputs)} requests" - ) - generated = 0 - for sample_index, output in enumerate(outputs): - token_ids = _token_ids(args, output) - status = ( - "success" if len(token_ids) == args.output_len else "model_failed" - ) - row = { - "iteration": iteration, - "sample_index": sample_index, - "status": status, - "input_tokens": args.input_len, - "output_tokens": len(token_ids), - } - sample_results.append(row) - raw_outputs.append({**row, "token_ids": token_ids}) - if status != "success": - raise RuntimeError( - f"iteration {iteration} sample {sample_index} produced {len(token_ids)} tokens" - ) - generated += len(token_ids) - performance.append( - { - "iteration": iteration, - "status": "success", - "elapsed_s": elapsed, - "output_tokens": generated, - "output_tok_s": generated / elapsed, - } - ) - rates = [row["output_tok_s"] for row in performance] - aggregate = { - "benchmark": "fixed_token_microbench", - "backend": args.backend, - "status": "success", - "output_tok_s_mean": statistics.fmean(rates), - "output_tok_s_median": statistics.median(rates), - "samples": len(rates), - } - except Exception as error: - aggregate = { - "benchmark": "fixed_token_microbench", - "backend": args.backend, - "status": "model_failed", - "error": repr(error), - } - raise - finally: - _write_jsonl(output_dir / "raw_outputs.jsonl", raw_outputs) - _write_jsonl(output_dir / "per_sample_results.jsonl", sample_results) - _write_jsonl(output_dir / "performance.jsonl", performance) - _write_json(output_dir / "aggregate_metrics.json", aggregate) - if args.backend == "sparsevllm" and engine is not None: - engine.exit() - return 0 - - -if __name__ == "__main__": - raise SystemExit(main())