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..7ed25b7c 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 | ✅ | ✅ | — | + `—` 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..0f9853f0 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 | ✅ | ✅ | — | + `—` 表示当前不支持该组合。 diff --git a/pyproject.toml b/pyproject.toml index 3c111c58..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", 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/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/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/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/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/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/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/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..78bb2e48 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,40 @@ 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 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() + payload_shm.unlink() if mode != "chain": if normalized_chain_id: raise ChainModeError( @@ -651,13 +708,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 +806,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 +830,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 +863,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 +1336,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 +1372,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 c688334f..ba288e48 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 @@ -65,6 +66,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 @@ -72,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_" @@ -89,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", @@ -99,10 +121,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", } @@ -203,6 +228,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() @@ -351,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": @@ -643,6 +677,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) @@ -662,6 +697,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, @@ -1109,6 +1165,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]): @@ -1269,7 +1330,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/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/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..c25e1c0a 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,45 @@ 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_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}.") 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..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)} @@ -127,7 +129,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 +146,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..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 @@ -111,7 +113,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 +128,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_context_attention.py b/src/sparsevllm/kernels/triton/gemma4_context_attention.py new file mode 100644 index 00000000..7e005142 --- /dev/null +++ b/src/sparsevllm/kernels/triton/gemma4_context_attention.py @@ -0,0 +1,222 @@ +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")) + has_visible_key = tl.max(visible.to(tl.int32), axis=1) > 0 + block_max = tl.max(logits, axis=1) + 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( + 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_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_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_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 new file mode 100644 index 00000000..6c11cf5c --- /dev/null +++ b/src/sparsevllm/kernels/triton/gemma4_moe.py @@ -0,0 +1,134 @@ +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, 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 None if large_token_config is None else dict(large_token_config) + + +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, + large_token_config: dict[str, int] | None = None, +) -> 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, large_token_config + ) 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, large_token_config + ) 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_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/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/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/layers/gemma4_rmsnorm.py b/src/sparsevllm/layers/gemma4_rmsnorm.py new file mode 100644 index 00000000..99d3297e --- /dev/null +++ b/src/sparsevllm/layers/gemma4_rmsnorm.py @@ -0,0 +1,30 @@ +from __future__ import annotations + +import torch +from torch import nn + +from sparsevllm.operators.gemma4 import Gemma4OperatorProvider + + +class Gemma4RMSNorm(nn.Module): + 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: + return self._ops.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..afcaa6f9 --- /dev/null +++ b/src/sparsevllm/models/gemma4.py @@ -0,0 +1,793 @@ +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.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.rotary_embedding import apply_rotary_emb +from sparsevllm.operators.gemma4 import ( + Gemma4OperatorProvider, + 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 +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( + 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, + parameters: dict | None = None, + ) -> None: + super().__init__() + 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}.") + 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, + operator_provider: Gemma4OperatorProvider, + ) -> 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_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_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 + 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, provider=operator_provider + ) + if not self.is_kv_shared_layer: + 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, + 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 = operator_provider.attention_backend( + 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) + 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( + (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) + 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) + return self.o_proj(output.flatten(1)) + + +class Gemma4MLP(nn.Module): + def __init__( + self, + config: Gemma4TextConfig, + layer_idx: int, + operator_provider: Gemma4OperatorProvider, + ) -> 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._ops = operator_provider + + def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: + 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, + operator_provider: Gemma4OperatorProvider, + router_provider: Gemma4RouterProvider, + ) -> 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, + 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)) + + def forward(self, hidden_states: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: + router_input = self._ops.router_input( + hidden_states, self.scale, self.root_size, self.norm.eps + ) + return self._router_ops.topk( + self.proj(router_input), self.per_expert_scale, self.top_k + ) + + +class Gemma4DecoderLayer(nn.Module): + def __init__( + self, + config: Gemma4TextConfig, + layer_idx: int, + operator_provider: Gemma4OperatorProvider, + router_provider: Gemma4RouterProvider | None, + rotary_emb: Gemma4RotaryEmbedding, + ) -> None: + super().__init__() + self.self_attn = Gemma4Attention( + config, + layer_idx, + rotary_emb, + operator_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, provider=operator_provider + ) + self.post_attention_layernorm = Gemma4RMSNorm( + config.hidden_size, eps=config.rms_norm_eps, provider=operator_provider + ) + self.pre_feedforward_layernorm = Gemma4RMSNorm( + config.hidden_size, eps=config.rms_norm_eps, provider=operator_provider + ) + self.post_feedforward_layernorm = Gemma4RMSNorm( + 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: + 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, + provider=operator_provider, + ) + self.enable_moe_block = bool(config.enable_moe_block) + if self.enable_moe_block: + self.parallel_context = get_parallel_context() + 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, + eps=config.rms_norm_eps, + provider=operator_provider, + ) + self.pre_feedforward_layernorm_2 = Gemma4RMSNorm( + 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, + provider=operator_provider, + ) + 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)) + 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) + 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) + 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) + hidden_states = self._ops.gelu_mul(hidden_states, per_layer_input) + hidden_states = self.per_layer_projection(hidden_states) + 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, + operator_provider: Gemma4OperatorProvider, + router_provider: Gemma4RouterProvider | None = 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, + provider=operator_provider, + ) + self.per_layer_model_projection_scale = float(config.hidden_size) ** -0.5 + self.per_layer_input_scale = 2.0**-0.5 + rotary_keys = [] + rotary_embeddings = {} + rotary_signatures = {} + for layer_idx, layer_type in enumerate(config.layer_types): + head_dim = int( + config_layer_get( + config, + layer_idx, + "head_dim", + "head_dim" + if layer_type == "sliding_attention" + 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 parameters.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, parameters + ) + self.rotary_embeddings = nn.ModuleDict(rotary_embeddings) + self.layers = nn.ModuleList( + Gemma4DecoderLayer( + config, + layer_idx, + operator_provider, + router_provider, + self.rotary_embeddings[rotary_keys[layer_idx]], + ) + for layer_idx in range(config.num_hidden_layers) + ) + 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 + + 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, + 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 + 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_excluded_prefixes = ("multimodal_encoder.",) + 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, + operator_provider: Gemma4OperatorProvider, + router_provider: Gemma4RouterProvider | None = None, + ) -> None: + super().__init__() + self.config = config + 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 + 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, **_): + head_dims = tuple( + sorted( + { + int(config_layer_get(config, layer_idx, "head_dim")) + for layer_idx in range(int(config.num_hidden_layers)) + } + ) + ) + runtime_kwargs = { + "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, + ) + } + 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) + 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 + 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) + 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, + 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) + 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/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/layout.py b/src/sparsevllm/models/layout.py index 2e3f450a..1b798b12 100644 --- a/src/sparsevllm/models/layout.py +++ b/src/sparsevllm/models/layout.py @@ -1,9 +1,9 @@ from __future__ import annotations -from dataclasses import dataclass +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: @@ -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,89 @@ 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}.") + heads, dims = [], [] + for layer_idx in layout.kv_idx_to_layer_idx: + is_full = str(layer_types[layer_idx]) == "full_attention" + 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: return self.layer_idx_to_kv_idx[int(layer_idx)] is not None 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/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/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..a9de2673 --- /dev/null +++ b/src/sparsevllm/multimodal/inputs.py @@ -0,0 +1,176 @@ +from __future__ import annotations + +import base64 +import binascii +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.") + 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}: + 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.py b/src/sparsevllm/operators/gemma4.py new file mode 100644 index 00000000..8a1a84a2 --- /dev/null +++ b/src/sparsevllm/operators/gemma4.py @@ -0,0 +1,268 @@ +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, + runtime_version_at_least, +) +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 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 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) + + +@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 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.""" + + 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 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", + "H20Gemma4OperatorProvider", + "Gemma4OperatorProvider", + "Gemma4OpSpec", + "TorchGemma4OperatorProvider", + "TritonGemma4OperatorProvider", + "resolve_gemma4_provider", +] diff --git a/src/sparsevllm/operators/gemma4_attention.py b/src/sparsevllm/operators/gemma4_attention.py new file mode 100644 index 00000000..8ead4b4e --- /dev/null +++ b/src/sparsevllm/operators/gemma4_attention.py @@ -0,0 +1,315 @@ +from __future__ import annotations + +from dataclasses import dataclass + +import torch + +from sparsevllm.layers.attention_backend import ( + TritonAttentionBackend, + _require_explicit_payload, +) + + +@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, + 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, + 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") + 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, + ) + + output = torch.empty_like(q) + 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 + 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, + ) + + 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 ( + 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, + ) + + 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 + 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, + 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", "Gemma4FlashInferPrefill"] diff --git a/src/sparsevllm/operators/gemma4_moe.py b/src/sparsevllm/operators/gemma4_moe.py new file mode 100644 index 00000000..06168c34 --- /dev/null +++ b/src/sparsevllm/operators/gemma4_moe.py @@ -0,0 +1,259 @@ +from __future__ import annotations + +import torch +import torch.nn.functional as F + +from sparsevllm import platforms +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, + runtime_version_at_least, +) +from sparsevllm.platforms.interface import DeviceCaps, PlatformEnum + + +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(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: + 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, + 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.""" + + 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, +) -> Gemma4MoeProvider: + 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 + + +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", + "H20Gemma4MoeProvider", + "TorchGemma4MoeProvider", + "TritonGemma4MoeProvider", + "resolve_gemma4_moe_provider", +] 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/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/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/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/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_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_dependency_constraints.py b/tests/test_dependency_constraints.py index 1e3187e0..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 diff --git a/tests/test_gemma4_attention_kernels.py b/tests/test_gemma4_attention_kernels.py new file mode 100644 index 00000000..421f59af --- /dev/null +++ b/tests/test_gemma4_attention_kernels.py @@ -0,0 +1,560 @@ +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") + 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("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): + 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 + + +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): + 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("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() + 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..4985030c --- /dev/null +++ b/tests/test_gemma4_model.py @@ -0,0 +1,541 @@ +from __future__ import annotations + +import inspect +from contextlib import ExitStack +from types import SimpleNamespace +from unittest.mock import patch + +import pytest +import torch +import torch.nn.functional as F +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, + Gemma4ForCausalLM, + Gemma4MLP, + Gemma4Model, + Gemma4RotaryEmbedding, + Gemma4Router, +) +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.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 +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, + ) + + 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_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] + ) + 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") +@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, + 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) + + +@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) + + +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 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, + "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_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_idx, (layer_type, head_dim) in enumerate( + (("sliding_attention", 4), ("full_attention", 8)) + ): + actual = Gemma4RotaryEmbedding(config, layer_type, head_dim) + 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], + ) + torch.testing.assert_close( + actual.cos_sin_cache[positions, 0, head_dim // 2 :], + sin[0, :, : head_dim // 2], + ) + + +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_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(): + actual = Gemma4MLP(config, 0, TorchGemma4OperatorProvider()) + 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, TorchGemma4OperatorProvider(), TorchGemma4RouterProvider() + ) + 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", 8), + TorchGemma4OperatorProvider(), + ) + loaded_key = torch.randn( + 8, 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, TorchGemma4OperatorProvider()) + 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", 4), + TorchGemma4OperatorProvider(), + ) + assert attention.is_kv_shared_layer + assert tuple(attention.qkv_proj.weight.shape) == ( + config.num_attention_heads * 4, + 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..aaf329ef --- /dev/null +++ b/tests/test_gemma4_rmsnorm.py @@ -0,0 +1,67 @@ +import pytest +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: + 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, provider=TorchGemma4OperatorProvider() + ) + 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, provider=TritonGemma4OperatorProvider() + ).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) + + +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 + ) diff --git a/tests/test_multimodal.py b/tests/test_multimodal.py new file mode 100644 index 00000000..555ed94b --- /dev/null +++ b/tests/test_multimodal.py @@ -0,0 +1,344 @@ +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, + 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 +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_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): + 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,))] + + +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") +@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 = 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) if score_ndim == 3 else (1, length), + device=device, + dtype=torch.float32, + ) + groups = torch.zeros(length, device=device, dtype=torch.int32) + groups[20:51] = 1 + 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) + 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()) + + 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..07e6d694 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,114 @@ 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_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_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_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) 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( 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)