Overlay: build the n-gram table parameter on the meta device when VLLM_PLE_QUANT_DIR is set, and swap the stub Parameter instead of set_data — removes the 102 GB virtual reservation that the kernel's overcommit heuristic refuses on hosts with less RAM+swap than the table (field report, 64 GB host); validated under an emulated 67/99 GiB commit limit, sanity PASS, tool-calling 77.0 (n=3), 81 tok/s c1
c75abdf verified | # SPDX-License-Identifier: Apache-2.0 | |
| # SPDX-FileCopyrightText: Copyright contributors to the vLLM project | |
| """GPU-resident Qwen3.8-Flash-Next position-learning enhancement layers.""" | |
| import math | |
| import os | |
| from contextlib import nullcontext | |
| from collections.abc import Iterable, Sequence | |
| import torch | |
| import torch.nn.functional as F | |
| from torch import nn | |
| import vllm.envs as envs | |
| from vllm.config import CacheConfig, ModelConfig, VllmConfig, get_current_vllm_config | |
| from vllm.forward_context import get_forward_context | |
| from vllm.model_executor.layers.linear import ReplicatedLinear | |
| from vllm.model_executor.layers.mamba.abstract import MambaBase | |
| from vllm.model_executor.layers.mamba.mamba_utils import ( | |
| MambaStateDtypeCalculator, | |
| MambaStateShapeCalculator, | |
| is_conv_state_dim_first, | |
| ) | |
| from vllm.model_executor.layers.ple_offload_layer import ( | |
| PleOffloadLayer, | |
| is_offload_process, | |
| ) | |
| from vllm.model_executor.layers.quantization.base_config import ( | |
| QuantizationConfig, | |
| QuantizeMethodBase, | |
| ) | |
| from vllm.model_executor.layers.quantization.fp8 import Fp8Config | |
| from vllm.model_executor.layers.quantization.utils.fp8_utils import ( | |
| create_fp8_scale_parameter, | |
| create_fp8_weight_parameter, | |
| is_fp8, | |
| ) | |
| from vllm.model_executor.layers.quantization.utils.quant_utils import ( | |
| is_layer_skipped, | |
| ) | |
| from vllm.model_executor.layers.vocab_parallel_embedding import ( | |
| VocabParallelEmbedding, | |
| ) | |
| from vllm.model_executor.models.utils import AutoWeightsLoader | |
| from vllm.model_executor.parameter import PerTensorScaleParameter | |
| from vllm.transformers_utils.configs.qwen3_8_flash_next import ( | |
| Qwen3_8FlashNextTextConfig, | |
| ) | |
| from vllm.utils.torch_utils import direct_register_custom_op | |
| from vllm.v1.attention.backends.registry import MambaAttentionBackendEnum | |
| from vllm.v1.attention.backends.short_conv_attn import ( | |
| PleShortConvAttentionBackend, | |
| PleShortConvAttentionMetadata, | |
| ) | |
| from vllm.v1.attention.backends.utils import NULL_BLOCK_ID | |
| from ..common.ple import copy_ple_embedding_shard_ | |
| _MASK64 = (1 << 64) - 1 | |
| _SPLITMIX_GAMMA = 0x9E3779B97F4A7C15 | |
| _SPLITMIX_M1 = 0xBF58476D1CE4E5B9 | |
| _SPLITMIX_M2 = 0x94D049BB133111EB | |
| _PLE_LAYER_PRIME = 10007 | |
| def _splitmix64(value: int) -> int: | |
| value = (value + _SPLITMIX_GAMMA) & _MASK64 | |
| value = ((value ^ (value >> 30)) * _SPLITMIX_M1) & _MASK64 | |
| value = ((value ^ (value >> 27)) * _SPLITMIX_M2) & _MASK64 | |
| return (value ^ (value >> 31)) & _MASK64 | |
| def _is_prime_64(value: int) -> bool: | |
| if value < 2: | |
| return False | |
| for prime in (2, 3, 5, 7, 11, 13, 17, 19, 23, 29, 31, 37): | |
| if value % prime == 0: | |
| return value == prime | |
| exponent = value - 1 | |
| shifts = 0 | |
| while exponent % 2 == 0: | |
| exponent //= 2 | |
| shifts += 1 | |
| for base in (2, 325, 9375, 28178, 450775, 9780504, 1795265022): | |
| if base % value == 0: | |
| continue | |
| witness = pow(base, exponent, value) | |
| if witness in (1, value - 1): | |
| continue | |
| for _ in range(shifts - 1): | |
| witness = pow(witness, 2, value) | |
| if witness == value - 1: | |
| break | |
| else: | |
| return False | |
| return True | |
| def _nth_prime_after(start: int, count: int) -> int: | |
| prime = int(start) | |
| for _ in range(count): | |
| candidate = prime + 1 | |
| if candidate <= 2: | |
| prime = 2 | |
| continue | |
| if candidate % 2 == 0: | |
| candidate += 1 | |
| while not _is_prime_64(candidate): | |
| candidate += 2 | |
| prime = candidate | |
| return prime | |
| class Qwen3_8FlashNextPLEGroupedNorm(nn.Module): | |
| def __init__( | |
| self, | |
| hidden_size: int, | |
| eps: float, | |
| group_size: int | None, | |
| dtype: torch.dtype | None, | |
| ) -> None: | |
| super().__init__() | |
| if group_size is not None and hidden_size % group_size: | |
| raise ValueError( | |
| f"hidden_size ({hidden_size}) must be divisible by " | |
| f"group_size ({group_size})" | |
| ) | |
| self.eps = eps | |
| self.group_size = group_size | |
| self.weight = nn.Parameter(torch.zeros(hidden_size, dtype=dtype)) | |
| def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: | |
| input_dtype = hidden_states.dtype | |
| hidden_states = hidden_states.float() | |
| if self.group_size is None: | |
| variance = hidden_states.square().mean(dim=-1, keepdim=True) | |
| normalized = hidden_states * torch.rsqrt(variance + self.eps) | |
| else: | |
| grouped = hidden_states.unflatten( | |
| -1, (hidden_states.shape[-1] // self.group_size, self.group_size) | |
| ) | |
| variance = grouped.square().mean(dim=-1, keepdim=True) | |
| normalized = (grouped * torch.rsqrt(variance + self.eps)).flatten(-2) | |
| return (normalized * (1.0 + self.weight.float())).to(input_dtype) | |
| class Qwen3_8FlashNextPLEFp8EmbeddingMethod(QuantizeMethodBase): | |
| """FP8 PLE embedding with one global checkpoint scale.""" | |
| def create_weights( | |
| self, | |
| layer: nn.Module, | |
| input_size_per_partition: int, | |
| output_partition_sizes: list[int], | |
| input_size: int, | |
| output_size: int, | |
| params_dtype: torch.dtype, | |
| **extra_weight_attrs, | |
| ) -> None: | |
| del input_size, output_size, params_dtype | |
| weight_loader = extra_weight_attrs.get("weight_loader") | |
| weight = create_fp8_weight_parameter( | |
| sum(output_partition_sizes), input_size_per_partition, weight_loader | |
| ) | |
| layer.register_parameter("weight", weight) | |
| weight_scale = create_fp8_scale_parameter( | |
| PerTensorScaleParameter, | |
| output_partition_sizes, | |
| input_size_per_partition, | |
| None, | |
| weight_loader, | |
| scale_dtype=torch.bfloat16, | |
| ) | |
| layer.register_parameter("weight_scale", weight_scale) | |
| def apply( | |
| self, | |
| layer: nn.Module, | |
| x: torch.Tensor, | |
| bias: torch.Tensor | None = None, | |
| ) -> torch.Tensor: | |
| raise NotImplementedError("PLE FP8 weights only support embedding lookup") | |
| def embedding(self, layer: nn.Module, input_: torch.Tensor) -> torch.Tensor: | |
| return F.embedding(input_, layer.weight) | |
| def _get_ple_embedding_quant_method( | |
| quant_config: QuantizationConfig | None, | |
| prefix: str, | |
| ) -> QuantizeMethodBase | None: | |
| """Select global-scale FP8 only for quantized PLE checkpoint shards.""" | |
| if not isinstance(quant_config, Fp8Config): | |
| return None | |
| if not quant_config.is_checkpoint_fp8_serialized: | |
| return None | |
| ignored_layers = quant_config.ignored_layers | |
| if is_layer_skipped( | |
| prefix, | |
| ignored_layers, | |
| quant_config.packed_modules_mapping, | |
| match_mode=quant_config.ignored_layers_match_mode, | |
| ): | |
| return None | |
| # PLE checkpoint shards form one runtime embedding parameter. | |
| shard_prefix = f"{prefix}.shard_" | |
| if any(name.startswith(shard_prefix) for name in ignored_layers): | |
| return None | |
| return Qwen3_8FlashNextPLEFp8EmbeddingMethod() | |
| class Qwen3_8FlashNextNGramEmbedding(PleOffloadLayer): | |
| def __init__( | |
| self, | |
| config: Qwen3_8FlashNextTextConfig, | |
| embedding_dim: int, | |
| ple_dense_layer_id: int, | |
| max_total_tokens: int, | |
| max_num_reqs: int, | |
| prefix: str, | |
| quant_config: QuantizationConfig | None = None, | |
| params_dtype: torch.dtype | None = None, | |
| ) -> None: | |
| super().__init__() | |
| self.embedding_dim = embedding_dim | |
| self.ngram_size = int(config.ngram_size) | |
| self.heads_per_ngram = int(config.heads_per_ngram) | |
| self.ngram_heads = (self.ngram_size - 1) * self.heads_per_ngram | |
| if self.ngram_size < 2: | |
| raise ValueError(f"ngram_size must be >= 2, got {self.ngram_size}") | |
| if self.heads_per_ngram <= 0: | |
| raise ValueError(f"heads_per_ngram must be > 0, got {self.heads_per_ngram}") | |
| if embedding_dim % self.ngram_heads: | |
| raise ValueError( | |
| "ple_embed_dim must be divisible by total ngram heads: " | |
| f"{embedding_dim} % {self.ngram_heads} != 0" | |
| ) | |
| self.head_dim = embedding_dim // self.ngram_heads | |
| self.eos_token_id = int(config.eos_token_id) | |
| self.unigram_vocab_size = int(config.vocab_size) | |
| self.split_ngram_parts = int(getattr(config, "split_ngram_parts", 512)) | |
| if self.split_ngram_parts <= 0: | |
| raise ValueError("split_ngram_parts must be positive") | |
| max_multiplier = ((1 << 63) - 1) // self.unigram_vocab_size | |
| half_bound = max(1, max_multiplier // 2) | |
| seed = int(getattr(config, "seed", 1234)) | |
| base_seed = seed + _PLE_LAYER_PRIME * ple_dense_layer_id | |
| multipliers = [] | |
| for index in range(self.ngram_size): | |
| value = base_seed + _SPLITMIX_GAMMA * (index + 1) | |
| multipliers.append(2 * (_splitmix64(value) % half_bound) + 1) | |
| self.register_buffer( | |
| "layer_multipliers", | |
| torch.tensor(multipliers, dtype=torch.long), | |
| persistent=True, | |
| ) | |
| ngram_vocab_size_base = int(config.ngram_vocab_size_base) | |
| sizes: list[int] = [] | |
| offsets: list[int] = [] | |
| offset = 0 | |
| for local_head in range(self.ngram_heads): | |
| global_head = ple_dense_layer_id * self.ngram_heads + local_head | |
| size = _nth_prime_after(ngram_vocab_size_base - 1, global_head + 1) | |
| sizes.append(size) | |
| offsets.append(offset) | |
| offset += size | |
| self.register_buffer( | |
| "ngram_heads_vocab_sizes", | |
| torch.tensor(sizes, dtype=torch.long), | |
| persistent=True, | |
| ) | |
| self.register_buffer( | |
| "ngram_heads_offsets", | |
| torch.tensor(offsets, dtype=torch.long), | |
| persistent=True, | |
| ) | |
| divisor = int(config.make_ngram_vocab_size_divisible_by) | |
| padded_vocab_size = ((offset + divisor - 1) // divisor) * divisor | |
| # With a quantized sidecar (VLLM_PLE_QUANT_DIR) the BF16 table is never | |
| # materialised, so build it on the meta device: otherwise the offload worker | |
| # reserves 95 GB of virtual memory for it, and the kernel's default | |
| # overcommit heuristic (vm.overcommit_memory=0) refuses any single allocation | |
| # larger than RAM + swap on hosts smaller than the table. The worker replaces | |
| # the parameter's storage with an empty stub before anything loads into it. | |
| table_device = ( | |
| torch.device("meta") if os.environ.get("VLLM_PLE_QUANT_DIR") else nullcontext() | |
| ) | |
| with table_device: | |
| self.ngram_embedding = VocabParallelEmbedding( | |
| padded_vocab_size, | |
| self.head_dim, | |
| params_dtype=params_dtype, | |
| padding_size=divisor, | |
| prefix=f"{prefix}.ngram_embedding", | |
| quant_method=_get_ple_embedding_quant_method( | |
| quant_config, f"{prefix}.ngram_embedding" | |
| ), | |
| ) | |
| self.register_buffer( | |
| "positions_buffer", | |
| torch.arange(max_total_tokens, dtype=torch.int64), | |
| persistent=False, | |
| ) | |
| self.register_buffer( | |
| "padded_buffer", | |
| torch.full( | |
| (max_num_reqs, max_total_tokens), | |
| self.eos_token_id, | |
| dtype=torch.int64, | |
| ), | |
| persistent=False, | |
| ) | |
| def _shift_precompute( | |
| tokens: torch.Tensor, eos_token_id: int | |
| ) -> tuple[torch.Tensor, torch.Tensor]: | |
| if tokens.dim() != 2: | |
| raise ValueError("tokens must be a 2D tensor") | |
| batch_size, seq_len = tokens.shape | |
| positions = torch.arange(seq_len, device=tokens.device, dtype=torch.int64) | |
| eos_positions = torch.where(tokens == eos_token_id, positions, -1) | |
| previous_eos_inclusive = torch.cummax(eos_positions, dim=1).values | |
| previous_eos = torch.cat( | |
| [ | |
| eos_positions.new_full((batch_size, 1), -1), | |
| previous_eos_inclusive[:, :-1], | |
| ], | |
| dim=1, | |
| ) | |
| return positions, positions.unsqueeze(0) - previous_eos - 1 | |
| def _shift_apply( | |
| tokens: torch.Tensor, | |
| positions: torch.Tensor, | |
| position_in_segment: torch.Tensor, | |
| shift: int, | |
| eos_token_id: int, | |
| ) -> torch.Tensor: | |
| if shift == 0: | |
| return tokens | |
| source = positions - shift | |
| gather_indices = source.clamp_min(0).unsqueeze(0).expand(tokens.shape[0], -1) | |
| shifted = tokens.gather(1, gather_indices) | |
| valid = (source.unsqueeze(0) >= 0) & (position_in_segment >= shift) | |
| return torch.where(valid, shifted, tokens.new_full((), eos_token_id)) | |
| def forward_impl( # type: ignore[override] | |
| self, | |
| hidden_states: torch.Tensor, | |
| input_ids: torch.Tensor, | |
| query_start_loc: torch.Tensor, | |
| ngram_context: torch.Tensor, | |
| output_buffer: torch.Tensor | None = None, | |
| ) -> torch.Tensor: | |
| del hidden_states | |
| input_ids = input_ids.reshape(-1).long() | |
| query_start_loc = query_start_loc.long() | |
| num_reqs = query_start_loc.numel() - 1 | |
| num_tokens = input_ids.shape[0] | |
| if num_tokens > self.positions_buffer.numel(): | |
| raise ValueError( | |
| f"PLE received {num_tokens} tokens, but its workspace supports " | |
| f"at most {self.positions_buffer.numel()}" | |
| ) | |
| if num_reqs > self.padded_buffer.shape[0]: | |
| raise ValueError( | |
| f"PLE received {num_reqs} requests, but its workspace supports " | |
| f"at most {self.padded_buffer.shape[0]}" | |
| ) | |
| # The CPU-offload subprocess is never captured by a CUDA Graph, so its | |
| # pack workspace can narrow to the actual maximum sequence length. The | |
| # regular GPU path retains the static maximum-width buffer for capture. | |
| if is_offload_process(): | |
| if num_reqs <= 0: | |
| raise ValueError("PLE CPU offload requires at least one request") | |
| max_seq_len = max( | |
| 1, | |
| int((query_start_loc[1:] - query_start_loc[:-1]).max().item()), | |
| ) | |
| # The model runner sends the CUDA-graph padded token count together | |
| # with an unpadded query_start_loc. Stale padding must not enter the | |
| # scatter: its clamped indices would overwrite the last real token. | |
| num_valid_tokens = min(int(query_start_loc[-1].item()), num_tokens) | |
| else: | |
| max_seq_len = self.padded_buffer.shape[1] | |
| num_valid_tokens = num_tokens | |
| positions = self.positions_buffer[:num_tokens] | |
| packed = self.padded_buffer[:num_reqs, :max_seq_len] | |
| packed.fill_(self.eos_token_id) | |
| request_indices = torch.searchsorted(query_start_loc, positions, right=True) - 1 | |
| request_indices.clamp_(max=num_reqs - 1) | |
| columns = (positions - query_start_loc[request_indices]).clamp( | |
| 0, packed.shape[1] - 1 | |
| ) | |
| packed[request_indices[:num_valid_tokens], columns[:num_valid_tokens]] = ( | |
| input_ids[:num_valid_tokens] | |
| ) | |
| ngram_context = ngram_context[:num_reqs].to( | |
| device=input_ids.device, dtype=torch.long | |
| ) | |
| context = torch.cat([ngram_context, packed], dim=-1) | |
| positions_2d, position_in_segment = self._shift_precompute( | |
| context, self.eos_token_id | |
| ) | |
| shifted = [context] | |
| for shift in range(1, self.ngram_size): | |
| shifted.append( | |
| self._shift_apply( | |
| context, | |
| positions_2d, | |
| position_in_segment, | |
| shift, | |
| self.eos_token_id, | |
| ) | |
| ) | |
| adjusted_columns = columns + self.ngram_size - 1 | |
| id_blocks = [] | |
| for ngram in range(2, self.ngram_size + 1): | |
| start = (ngram - 2) * self.heads_per_ngram | |
| end = start + self.heads_per_ngram | |
| mixed = shifted[0] * self.layer_multipliers[0] | |
| for index in range(1, ngram): | |
| mixed = torch.bitwise_xor( | |
| mixed, shifted[index] * self.layer_multipliers[index] | |
| ) | |
| sizes = self.ngram_heads_vocab_sizes[start:end] | |
| offsets = self.ngram_heads_offsets[start:end] | |
| ids = torch.remainder(mixed.unsqueeze(-1), sizes) + offsets | |
| id_blocks.append(ids[request_indices, adjusted_columns]) | |
| ngram_ids = torch.cat(id_blocks, dim=-1) | |
| quant = getattr(self.ngram_embedding, "_ple_quant", None) | |
| if output_buffer is not None: | |
| output = output_buffer[:num_tokens, : self.embedding_dim] | |
| if quant is not None: | |
| quant.gather_into( | |
| ngram_ids.reshape(-1), output.reshape(-1, self.head_dim) | |
| ) | |
| else: | |
| torch.index_select( | |
| self.ngram_embedding.weight, | |
| 0, | |
| ngram_ids.reshape(-1), | |
| out=output.reshape(-1, self.head_dim), | |
| ) | |
| return output | |
| if quant is not None: | |
| flat = torch.empty( | |
| ngram_ids.numel(), | |
| self.head_dim, | |
| dtype=torch.bfloat16, | |
| device=ngram_ids.device, | |
| ) | |
| quant.gather_into(ngram_ids.reshape(-1), flat) | |
| return flat.view(*ngram_ids.shape, self.head_dim).flatten(-2) | |
| return self.ngram_embedding(ngram_ids).flatten(-2) | |
| def get_offload_output_dtype(self, default_dtype: torch.dtype) -> torch.dtype: | |
| """Keep quantized lookup results in their embedding storage dtype.""" | |
| embedding = getattr(self, "ngram_embedding", None) | |
| weight = getattr(embedding, "weight", None) | |
| if weight is not None: | |
| return weight.dtype | |
| if hasattr(self, "_offload_weight_scale"): | |
| return torch.float8_e4m3fn | |
| return default_dtype | |
| def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]: | |
| """Load hash buffers and checkpoint-split embedding rows.""" | |
| # GPU workers retain only the global FP8 scale. The CPU process owns the | |
| # embedding weight and returns its quantized lookup output unchanged. | |
| if envs.VLLM_PLE_CPU_OFFLOAD and not is_offload_process(): | |
| retained: set[str] = set() | |
| for name, loaded_weight in weights: | |
| if name != "ngram_embedding.weight_scale": | |
| continue | |
| self.register_buffer( | |
| "_offload_weight_scale", | |
| loaded_weight.to(device=torch.accelerator.current_accelerator()), | |
| persistent=False, | |
| ) | |
| retained.add(name) | |
| return retained | |
| persistent_buffers = { | |
| "layer_multipliers": self.layer_multipliers, | |
| "ngram_heads_offsets": self.ngram_heads_offsets, | |
| "ngram_heads_vocab_sizes": self.ngram_heads_vocab_sizes, | |
| } | |
| loaded: set[str] = set() | |
| regular_weights: list[tuple[str, torch.Tensor]] = [] | |
| shard_prefix = "ngram_embedding.shard_" | |
| for name, loaded_weight in weights: | |
| leaf_name = name.rsplit(".", 1)[-1] | |
| if leaf_name.startswith("hashstats_") or leaf_name == "token_lookup": | |
| continue | |
| if name in persistent_buffers: | |
| buffer = persistent_buffers[name] | |
| if buffer.shape != loaded_weight.shape: | |
| raise ValueError( | |
| f"Shape mismatch for {name}: expected " | |
| f"{tuple(buffer.shape)}, got {tuple(loaded_weight.shape)}" | |
| ) | |
| buffer.copy_(loaded_weight.to(device=buffer.device, dtype=buffer.dtype)) | |
| loaded.add(name) | |
| continue | |
| if name.startswith(shard_prefix) and name.endswith(".weight"): | |
| shard_text = name[len(shard_prefix) : -len(".weight")] | |
| if not shard_text.isdigit(): | |
| regular_weights.append((name, loaded_weight)) | |
| continue | |
| shard_index = int(shard_text) | |
| if shard_index >= self.split_ngram_parts: | |
| raise ValueError( | |
| f"PLE embedding shard index {shard_index} exceeds " | |
| f"split_ngram_parts={self.split_ngram_parts}" | |
| ) | |
| embedding = self.ngram_embedding | |
| shard_size = ( | |
| embedding.org_vocab_size + self.split_ngram_parts - 1 | |
| ) // self.split_ngram_parts | |
| checkpoint_start = shard_index * shard_size | |
| expected_rows = max( | |
| 0, | |
| min(shard_size, embedding.org_vocab_size - checkpoint_start), | |
| ) | |
| expected_shape = (expected_rows, embedding.embedding_dim) | |
| if tuple(loaded_weight.shape) != expected_shape: | |
| raise ValueError( | |
| f"Shape mismatch for PLE embedding shard {shard_index}: " | |
| f"expected {expected_shape}, got " | |
| f"{tuple(loaded_weight.shape)}" | |
| ) | |
| copy_ple_embedding_shard_( | |
| embedding.weight.data, | |
| loaded_weight, | |
| checkpoint_start=checkpoint_start, | |
| tp_start=embedding.shard_indices.org_vocab_start_index, | |
| tp_end=embedding.shard_indices.org_vocab_end_index, | |
| ) | |
| loaded.add("ngram_embedding.weight") | |
| continue | |
| regular_weights.append((name, loaded_weight)) | |
| if regular_weights: | |
| loaded.update(AutoWeightsLoader(self).load_weights(regular_weights)) | |
| return loaded | |
| class Qwen3_8FlashNextPLELayer(nn.Module, MambaBase): | |
| def __init__( | |
| self, | |
| config: Qwen3_8FlashNextTextConfig, | |
| vllm_config: VllmConfig, | |
| layer_idx: int = 0, | |
| ple_dense_layer_id: int | None = None, | |
| prefix: str = "", | |
| ) -> None: | |
| super().__init__() | |
| model_config = vllm_config.model_config | |
| cache_config = vllm_config.cache_config | |
| quant_config = vllm_config.quant_config | |
| self.model_config: ModelConfig = model_config | |
| self.cache_config: CacheConfig = cache_config | |
| self.layer_idx = layer_idx | |
| self.ple_dense_layer_id = ( | |
| int(ple_dense_layer_id) | |
| if ple_dense_layer_id is not None | |
| else int(layer_idx) | |
| ) | |
| self.prefix = prefix | |
| self.hidden_size = int(config.hidden_size) | |
| self.hc_count = config.hc_count | |
| self.hc_hidden_size = self.hidden_size * self.hc_count | |
| self.conv_kernel_size = int(config.ple_conv_kernel_size) | |
| self.short_conv_dilation = int(config.ngram_size) | |
| self.conv_state_len = (self.conv_kernel_size - 1) * self.short_conv_dilation | |
| self.num_spec_tokens = vllm_config.num_speculative_tokens | |
| self.activation = "silu" | |
| # The offload process builds the surrounding model on meta while | |
| # this subtree must own real CPU storage. GPU workers skip the | |
| # subclass constructor and retain only an empty IPC placeholder. | |
| with torch.device(PleOffloadLayer.get_target_device()): | |
| self.ple_embedding: nn.Module = Qwen3_8FlashNextNGramEmbedding( | |
| config, | |
| int(config.ple_embed_dim), | |
| self.ple_dense_layer_id, | |
| vllm_config.scheduler_config.max_num_batched_tokens, | |
| vllm_config.scheduler_config.max_num_seqs, | |
| f"{prefix}.ple_embedding", | |
| quant_config=quant_config, | |
| params_dtype=model_config.dtype, | |
| ) | |
| self.key_proj = ReplicatedLinear( | |
| int(config.ple_embed_dim), | |
| self.hc_hidden_size, | |
| bias=False, | |
| quant_config=quant_config, | |
| prefix=f"{prefix}.key_proj", | |
| ) | |
| self.value_proj = ReplicatedLinear( | |
| int(config.ple_embed_dim), | |
| self.hidden_size, | |
| bias=False, | |
| quant_config=quant_config, | |
| prefix=f"{prefix}.value_proj", | |
| ) | |
| norm_args = ( | |
| self.hc_hidden_size, | |
| config.rms_norm_eps, | |
| self.hidden_size, | |
| model_config.dtype, | |
| ) | |
| self.norm_key = Qwen3_8FlashNextPLEGroupedNorm(*norm_args) | |
| self.norm_query = Qwen3_8FlashNextPLEGroupedNorm(*norm_args) | |
| self.norm_conv = Qwen3_8FlashNextPLEGroupedNorm(*norm_args) | |
| self.conv1d = nn.Conv1d( | |
| self.hc_hidden_size, | |
| self.hc_hidden_size, | |
| self.conv_kernel_size, | |
| groups=self.hc_hidden_size, | |
| padding=self.conv_state_len, | |
| dilation=self.short_conv_dilation, | |
| bias=False, | |
| dtype=model_config.dtype, | |
| ) | |
| nn.init.zeros_(self.conv1d.weight) | |
| self.conv1d.weight._no_reinit = True | |
| self.kv_cache = (torch.tensor([]),) | |
| compilation_config = get_current_vllm_config().compilation_config | |
| if prefix in compilation_config.static_forward_context: | |
| raise ValueError(f"Duplicate layer name: {prefix}") | |
| compilation_config.static_forward_context[prefix] = self | |
| def _get_embedding_weight_scale(self) -> torch.Tensor | None: | |
| embedding = getattr(self.ple_embedding, "ngram_embedding", None) | |
| weight_scale = getattr(embedding, "weight_scale", None) | |
| if weight_scale is not None: | |
| return weight_scale | |
| return getattr(self.ple_embedding, "_offload_weight_scale", None) | |
| def _dequantize_embeddings( | |
| self, | |
| embeddings: torch.Tensor, | |
| output_dtype: torch.dtype, | |
| ) -> torch.Tensor: | |
| """Dequantize PLE lookup output.""" | |
| if not is_fp8(embeddings): | |
| return embeddings | |
| weight_scale = self._get_embedding_weight_scale() | |
| if weight_scale is None: | |
| raise RuntimeError("FP8 PLE embedding is missing its global scale") | |
| if weight_scale.device != embeddings.device: | |
| raise RuntimeError("FP8 PLE embedding scale must be on the output device") | |
| return embeddings.to(output_dtype) * weight_scale.to(output_dtype) | |
| def mamba_type(self) -> MambaAttentionBackendEnum: | |
| return MambaAttentionBackendEnum.SHORT_CONV | |
| def is_kv_cache_tp_replicated(self) -> bool: | |
| return True | |
| def get_attn_backend(self) -> type[PleShortConvAttentionBackend]: | |
| return PleShortConvAttentionBackend | |
| def get_state_dtype(self) -> tuple[torch.dtype, ...]: | |
| return MambaStateDtypeCalculator.short_conv_state_dtype( | |
| self.model_config.dtype, self.cache_config.mamba_cache_dtype | |
| ) | |
| def get_state_shape(self) -> Sequence[tuple[int, ...]]: | |
| return MambaStateShapeCalculator.short_conv_state_shape( | |
| tp_world_size=1, | |
| intermediate_size=self.hc_hidden_size, | |
| conv_kernel=self.conv_state_len + 1, | |
| num_spec=self.num_spec_tokens, | |
| ) | |
| def _apply_norm( | |
| self, norm: Qwen3_8FlashNextPLEGroupedNorm, hidden_states: torch.Tensor | |
| ) -> torch.Tensor: | |
| shape = hidden_states.shape | |
| return norm(hidden_states.flatten(-2)).reshape(shape) | |
| def _short_conv_fallback(self, inputs: torch.Tensor) -> torch.Tensor: | |
| # Profiling / CUDA graph capture only; conv state is not updated. | |
| inputs_t = inputs.transpose(0, 1).unsqueeze(0) | |
| output = self.conv1d(inputs_t)[..., : inputs_t.size(-1)] | |
| return F.silu(output).squeeze(0).transpose(0, 1) | |
| def _short_conv_dilated_decode_batched( | |
| self, | |
| x_d: torch.Tensor, | |
| conv_state: torch.Tensor, | |
| conv_weights: torch.Tensor, | |
| state_indices_tensor_d: torch.Tensor, | |
| has_initial_states_d: torch.Tensor | None, | |
| ) -> torch.Tensor: | |
| state_indices = state_indices_tensor_d.to( | |
| device=conv_state.device, dtype=torch.int64 | |
| ) | |
| # TODO: need double-check | |
| # FULL cudagraph padded decode rows use NULL_BLOCK_ID. Remap them to | |
| # slot 0 for a safe gather, then zero output and skip write-back. | |
| valid_state = state_indices != NULL_BLOCK_ID | |
| state_indices = torch.where( | |
| valid_state, state_indices, torch.zeros_like(state_indices) | |
| ) | |
| if has_initial_states_d is None: | |
| has_initial_state = valid_state | |
| else: | |
| if has_initial_states_d.numel() < state_indices_tensor_d.numel(): | |
| raise ValueError( | |
| "has_initial_states_d size mismatch: " | |
| f"got {has_initial_states_d.numel()}, " | |
| f"need >= {state_indices_tensor_d.numel()}." | |
| ) | |
| has_initial_state = has_initial_states_d[ | |
| : state_indices_tensor_d.numel() | |
| ].to(device=conv_state.device, dtype=torch.bool) | |
| has_initial_state = has_initial_state & valid_state | |
| cached_state = conv_state.index_select(0, state_indices) | |
| state = cached_state[..., : self.conv_state_len].to(x_d.dtype) | |
| if self.conv_state_len > 0: | |
| initial_state = torch.where( | |
| has_initial_state.view(-1, 1, 1), | |
| state, | |
| torch.zeros_like(state), | |
| ) | |
| history = torch.cat((initial_state, x_d.unsqueeze(-1)), dim=-1) | |
| else: | |
| history = x_d.unsqueeze(-1) | |
| conv_output = F.conv1d( | |
| history, | |
| conv_weights.unsqueeze(1).contiguous(), | |
| groups=history.size(1), | |
| dilation=self.short_conv_dilation, | |
| ).squeeze(-1) | |
| output = F.silu(conv_output) | |
| output = output * valid_state.view(-1, 1).to(output.dtype) | |
| if self.conv_state_len > 0: | |
| next_state = history[..., -self.conv_state_len :] | |
| # Padded rows are remapped to the reserved null slot. Preserve its | |
| # existing value while writing the new states for valid rows. | |
| existing_base_state = cached_state[..., : self.conv_state_len] | |
| safe_next_state = torch.where( | |
| valid_state.view(-1, 1, 1), | |
| next_state.to(conv_state.dtype), | |
| existing_base_state, | |
| ) | |
| cached_state[..., : self.conv_state_len] = safe_next_state | |
| conv_state.index_copy_(0, state_indices, cached_state) | |
| return output | |
| def _short_conv_dilated_prefill_batched( | |
| self, | |
| x_p: torch.Tensor, | |
| metadata: PleShortConvAttentionMetadata, | |
| conv_state: torch.Tensor, | |
| conv_weights: torch.Tensor, | |
| state_indices_tensor_p: torch.Tensor, | |
| num_prefills: int, | |
| num_decode_tokens: int, | |
| num_prefill_tokens: int, | |
| ) -> torch.Tensor: | |
| # ``non_spec_query_start_loc`` covers the non-spec (decode + prefill) | |
| # requests and equals ``query_start_loc`` when spec-decode is inactive. | |
| non_spec_query_start_loc = metadata.non_spec_query_start_loc | |
| if non_spec_query_start_loc is None: | |
| raise ValueError("query_start_loc is required for prefill short-conv") | |
| query_start_loc_p = ( | |
| non_spec_query_start_loc[-num_prefills - 1 :] - num_decode_tokens | |
| ) | |
| # The metadata builder guarantees that the prefill query offsets start | |
| # at 0 and end at num_prefill_tokens. Avoid reading those values here, | |
| # since doing so would force a device-to-host synchronization. | |
| has_initial_states_p = metadata.has_initial_states_p | |
| if has_initial_states_p is None: | |
| raise ValueError("has_initial_states_p is required for prefill short-conv") | |
| output = torch.empty_like(x_p) | |
| q_starts = query_start_loc_p.to(torch.int64) | |
| if state_indices_tensor_p.numel() < num_prefills: | |
| raise ValueError( | |
| "state_indices_tensor_p size mismatch: " | |
| f"got {state_indices_tensor_p.numel()}, " | |
| f"need >= {num_prefills}." | |
| ) | |
| if has_initial_states_p.numel() < num_prefills: | |
| raise ValueError( | |
| "has_initial_states_p size mismatch: " | |
| f"got {has_initial_states_p.numel()}, " | |
| f"need >= {num_prefills}." | |
| ) | |
| if num_prefills == 0 or x_p.numel() == 0: | |
| return output | |
| lengths = q_starts[1:] - q_starts[:-1] | |
| # Use the CPU-computed packing width from the metadata builder instead | |
| # of synchronizing on lengths.max(). | |
| max_len = metadata.max_prefill_query_len | |
| if max_len <= 0: | |
| return output | |
| hidden_size = x_p.shape[1] | |
| positions = torch.arange( | |
| num_prefill_tokens, device=x_p.device, dtype=torch.int64 | |
| ) | |
| req_indices = torch.searchsorted(q_starts[1:], positions, right=True) | |
| col_indices = positions - q_starts[req_indices] | |
| packed_tokens = x_p.new_zeros((num_prefills, max_len, hidden_size)) | |
| packed_tokens[req_indices, col_indices] = x_p | |
| packed_tokens = packed_tokens.transpose(1, 2).contiguous() | |
| state_indices = state_indices_tensor_p[:num_prefills].to( | |
| device=conv_state.device, dtype=torch.int64 | |
| ) | |
| valid_state = state_indices != NULL_BLOCK_ID | |
| state_indices = torch.where( | |
| valid_state, state_indices, torch.zeros_like(state_indices) | |
| ) | |
| has_initial = has_initial_states_p[:num_prefills].to( | |
| device=conv_state.device, dtype=torch.bool | |
| ) | |
| if self.conv_state_len > 0: | |
| if conv_state.shape[0] == 0: | |
| state = conv_state.new_zeros( | |
| (num_prefills, hidden_size, self.conv_state_len), | |
| dtype=x_p.dtype, | |
| ) | |
| else: | |
| state = conv_state.index_select(0, state_indices)[ | |
| ..., : self.conv_state_len | |
| ].to(x_p.dtype) | |
| use_initial_mask = (valid_state & has_initial).view(num_prefills, 1, 1) | |
| initial_state = torch.where( | |
| use_initial_mask, | |
| state, | |
| torch.zeros_like(state), | |
| ) | |
| history = torch.cat((initial_state, packed_tokens), dim=-1) | |
| else: | |
| history = packed_tokens | |
| conv_output = F.conv1d( | |
| history, | |
| conv_weights.unsqueeze(1).contiguous(), | |
| groups=history.size(1), | |
| dilation=self.short_conv_dilation, | |
| ) | |
| conv_output = F.silu(conv_output).transpose(1, 2).contiguous() | |
| token_positions = torch.arange(max_len, device=x_p.device, dtype=torch.int64) | |
| valid_tokens = token_positions.view(1, max_len) < lengths.view(num_prefills, 1) | |
| valid_output_mask = valid_tokens & valid_state.to(device=x_p.device).view( | |
| num_prefills, 1 | |
| ) | |
| conv_output.masked_fill_(~valid_output_mask.unsqueeze(-1), 0) | |
| output.copy_(conv_output[req_indices, col_indices]) | |
| if self.conv_state_len > 0 and conv_state.shape[0] > 0: | |
| state_starts = lengths.to(device=history.device, dtype=torch.int64).view( | |
| num_prefills, 1, 1 | |
| ) | |
| state_offsets = torch.arange( | |
| self.conv_state_len, device=history.device, dtype=torch.int64 | |
| ).view(1, 1, self.conv_state_len) | |
| next_state = history.gather( | |
| dim=2, | |
| index=(state_starts + state_offsets).expand(-1, history.size(1), -1), | |
| ) | |
| # Write back without a host synchronization. Valid, non-empty rows | |
| # receive their new state; padding and zero-length rows keep the | |
| # current cache value. | |
| existing_state = conv_state.index_select(0, state_indices) | |
| existing_base_state = existing_state[..., : self.conv_state_len] | |
| update_mask = valid_state & (lengths.to(device=conv_state.device) > 0) | |
| safe_next_state = torch.where( | |
| update_mask.view(num_prefills, 1, 1), | |
| next_state.to(conv_state.dtype), | |
| existing_base_state, | |
| ) | |
| existing_state[..., : self.conv_state_len] = safe_next_state | |
| conv_state.index_copy_(0, state_indices, existing_state) | |
| return output | |
| def _short_conv_dilated_spec_batched( | |
| self, | |
| x_spec: torch.Tensor, | |
| conv_state: torch.Tensor, | |
| conv_weights: torch.Tensor, | |
| spec_state_indices_tensor: torch.Tensor, | |
| spec_query_start_loc: torch.Tensor, | |
| num_accepted_tokens: torch.Tensor, | |
| spec_query_len: int, | |
| ) -> torch.Tensor: | |
| """Dilated short-conv for speculative-decode (MTP) requests. | |
| Each spec request feeds multiple (draft + 1) query tokens. The conv | |
| outputs are computed causally after rolling back the previous draft | |
| state by ``num_accepted_tokens - 1``. The current candidate inputs stay | |
| in the extended cache for the next forward, matching | |
| ``causal_conv1d_update``. | |
| ``spec_query_len`` (== num_speculative_tokens + 1) is the maximum query | |
| length and is a Python int, so no host synchronization is needed; this | |
| keeps the path safe for full CUDA-graph capture/replay where the buffers | |
| are padded at the request level. | |
| """ | |
| num_reqs = spec_state_indices_tensor.numel() | |
| hidden_size = x_spec.size(-1) | |
| # Use a fixed packing width instead of synchronizing on lengths.max(). | |
| max_len = spec_query_len | |
| # Full CUDA graphs can pad these buffers. Only the first num_reqs | |
| # accepted-token counts belong to actual speculative requests. | |
| num_accepted_tokens = num_accepted_tokens[:num_reqs] | |
| q_starts = spec_query_start_loc[: num_reqs + 1].to(torch.int64) | |
| # Keep the number of real speculative tokens on the device. | |
| total_real_tokens = q_starts[num_reqs] | |
| state_indices = spec_state_indices_tensor.to( | |
| device=conv_state.device, dtype=torch.int64 | |
| ) | |
| valid_state = state_indices != NULL_BLOCK_ID | |
| state_indices = torch.where( | |
| valid_state, state_indices, torch.zeros_like(state_indices) | |
| ) | |
| positions = torch.arange( | |
| x_spec.size(0), device=x_spec.device, dtype=torch.int64 | |
| ) | |
| # Route graph-padded token rows to the discarded dummy request so that | |
| # they cannot overwrite real packed data. | |
| req_indices = torch.searchsorted(q_starts[1:], positions, right=True) | |
| valid_tokens = (positions < total_real_tokens) & (req_indices < num_reqs) | |
| clamped_req_indices = req_indices.clamp_max(max(num_reqs - 1, 0)) | |
| col_indices = (positions - q_starts[clamped_req_indices]).clamp_(0, max_len - 1) | |
| pack_req_indices = torch.where( | |
| valid_tokens, | |
| clamped_req_indices, | |
| torch.full_like(req_indices, num_reqs), | |
| ) | |
| pack_col_indices = torch.where( | |
| valid_tokens, col_indices, torch.zeros_like(col_indices) | |
| ) | |
| # The last request row is the dummy sink for graph padding. | |
| packed = x_spec.new_zeros((num_reqs + 1, max_len, hidden_size)) | |
| packed[pack_req_indices, pack_col_indices] = x_spec | |
| packed = packed.transpose(1, 2).contiguous() | |
| if self.conv_state_len > 0: | |
| cached_state = conv_state.index_select(0, state_indices) | |
| rollback_offsets = num_accepted_tokens.to( | |
| device=conv_state.device, dtype=torch.int64 | |
| ).sub(1) | |
| rollback_offsets = torch.where( | |
| valid_state, | |
| rollback_offsets.clamp_(0, max_len - 1), | |
| torch.zeros_like(rollback_offsets), | |
| ) | |
| state_offsets = torch.arange( | |
| self.conv_state_len, device=conv_state.device, dtype=torch.int64 | |
| ).view(1, 1, self.conv_state_len) | |
| rollback_indices = rollback_offsets.view(-1, 1, 1) + state_offsets | |
| state = cached_state.gather( | |
| 2, rollback_indices.expand(-1, hidden_size, -1) | |
| ).to(x_spec.dtype) | |
| state = torch.where( | |
| valid_state.view(num_reqs, 1, 1), | |
| state, | |
| torch.zeros_like(state), | |
| ) | |
| # Append a zeroed dummy-row state to match the [num_reqs + 1] pack. | |
| dummy_state = state.new_zeros((1, hidden_size, self.conv_state_len)) | |
| state_full = torch.cat((state, dummy_state), dim=0) | |
| history = torch.cat((state_full, packed), dim=-1) | |
| else: | |
| history = packed | |
| conv_output = F.conv1d( | |
| history, | |
| conv_weights.unsqueeze(1).contiguous(), | |
| groups=history.size(1), | |
| dilation=self.short_conv_dilation, | |
| ) | |
| conv_output = F.silu(conv_output).transpose(1, 2).contiguous() | |
| output = conv_output[pack_req_indices, pack_col_indices] | |
| output = output * valid_tokens.view(-1, 1).to(output.dtype) | |
| # Keep all current candidate inputs in the extended state. On the next | |
| # target forward, ``num_accepted_tokens - 1`` selects the rollback | |
| # window before processing the newly scheduled tokens. | |
| if self.conv_state_len > 0: | |
| state_capacity = self.conv_state_len + max_len - 1 | |
| if conv_state.size(-1) < state_capacity: | |
| raise RuntimeError( | |
| "PLE short-conv cache cannot retain speculative tokens: " | |
| f"got {conv_state.size(-1)}, need {state_capacity}." | |
| ) | |
| candidate_state = history[:num_reqs, :, 1 : state_capacity + 1] | |
| query_lengths = q_starts[1:] - q_starts[:-1] | |
| state_positions = torch.arange( | |
| state_capacity, device=history.device, dtype=torch.int64 | |
| ).view(1, 1, state_capacity) | |
| update_lengths = (self.conv_state_len + query_lengths - 1).view( | |
| num_reqs, 1, 1 | |
| ) | |
| update_mask = valid_state.view(num_reqs, 1, 1) & ( | |
| state_positions < update_lengths | |
| ) | |
| existing_state = cached_state[..., :state_capacity] | |
| next_state = torch.where( | |
| update_mask, | |
| candidate_state.to(conv_state.dtype), | |
| existing_state, | |
| ) | |
| cached_state[..., :state_capacity] = next_state | |
| conv_state.index_copy_(0, state_indices, cached_state) | |
| return output | |
| def _short_conv_dilated_dispatch( | |
| self, | |
| inputs: torch.Tensor, | |
| metadata: PleShortConvAttentionMetadata, | |
| conv_state: torch.Tensor, | |
| conv_weights: torch.Tensor, | |
| ) -> torch.Tensor: | |
| num_prefills = metadata.num_prefills | |
| num_decodes = metadata.num_decodes | |
| num_decode_tokens = metadata.num_decode_tokens | |
| num_prefill_tokens = metadata.num_prefill_tokens | |
| has_prefill = num_prefills > 0 | |
| has_decode = num_decodes > 0 | |
| has_spec = metadata.spec_sequence_masks is not None | |
| x = inputs[: metadata.num_actual_tokens] | |
| # Split spec / non-spec tokens. | |
| if has_spec: | |
| if has_prefill or has_decode: | |
| assert metadata.spec_token_indx is not None | |
| assert metadata.non_spec_token_indx is not None | |
| x_spec = x.index_select(0, metadata.spec_token_indx.long()) | |
| x_non_spec = x.index_select(0, metadata.non_spec_token_indx.long()) | |
| else: | |
| x_spec = x | |
| x_non_spec = None | |
| else: | |
| x_spec = None | |
| x_non_spec = x | |
| spec_output = None | |
| # 1. Run the multi-query speculative-decode part. | |
| if has_spec: | |
| assert metadata.spec_state_indices_tensor is not None | |
| assert metadata.spec_query_start_loc is not None | |
| assert metadata.num_accepted_tokens is not None | |
| spec_output = self._short_conv_dilated_spec_batched( | |
| x_spec=x_spec, | |
| conv_state=conv_state, | |
| conv_weights=conv_weights, | |
| spec_state_indices_tensor=metadata.spec_state_indices_tensor[ | |
| : metadata.num_spec_decodes | |
| ], | |
| spec_query_start_loc=metadata.spec_query_start_loc, | |
| num_accepted_tokens=metadata.num_accepted_tokens, | |
| spec_query_len=metadata.spec_query_len, | |
| ) | |
| # 2. Run regular decode and prefill requests. | |
| conv_out_non_spec = None | |
| state_indices_tensor = metadata.state_indices_tensor | |
| if x_non_spec is not None: | |
| assert state_indices_tensor is not None | |
| if has_prefill: | |
| state_indices_tensor_d, state_indices_tensor_p = torch.split( | |
| state_indices_tensor, | |
| [num_decodes, num_prefills], | |
| dim=0, | |
| ) | |
| x_d, x_p = torch.split( | |
| x_non_spec, | |
| [num_decode_tokens, num_prefill_tokens], | |
| dim=0, | |
| ) | |
| non_spec_parts: list[torch.Tensor] = [] | |
| if has_decode: | |
| non_spec_parts.append( | |
| self._short_conv_dilated_decode_batched( | |
| x_d=x_d, | |
| conv_state=conv_state, | |
| conv_weights=conv_weights, | |
| state_indices_tensor_d=state_indices_tensor_d, | |
| has_initial_states_d=metadata.has_initial_states_d, | |
| ) | |
| ) | |
| non_spec_parts.append( | |
| self._short_conv_dilated_prefill_batched( | |
| x_p=x_p, | |
| metadata=metadata, | |
| conv_state=conv_state, | |
| conv_weights=conv_weights, | |
| state_indices_tensor_p=state_indices_tensor_p, | |
| num_prefills=num_prefills, | |
| num_decode_tokens=num_decode_tokens, | |
| num_prefill_tokens=num_prefill_tokens, | |
| ) | |
| ) | |
| conv_out_non_spec = torch.vstack(non_spec_parts) | |
| else: | |
| conv_out_non_spec = self._short_conv_dilated_decode_batched( | |
| x_d=x_non_spec, | |
| conv_state=conv_state, | |
| conv_weights=conv_weights, | |
| state_indices_tensor_d=state_indices_tensor[: x_non_spec.size(0)], | |
| has_initial_states_d=metadata.has_initial_states_d, | |
| ) | |
| # 3. Merge both parts back into the original token order. | |
| if has_spec and conv_out_non_spec is not None: | |
| assert metadata.spec_token_indx is not None | |
| assert metadata.non_spec_token_indx is not None | |
| assert spec_output is not None | |
| output = x.new_empty((metadata.num_actual_tokens, x.size(-1))) | |
| output.index_copy_(0, metadata.spec_token_indx, spec_output) | |
| output.index_copy_(0, metadata.non_spec_token_indx, conv_out_non_spec) | |
| return output | |
| elif has_spec: | |
| assert spec_output is not None | |
| return spec_output | |
| if conv_out_non_spec is None: | |
| return x | |
| return conv_out_non_spec | |
| def _short_conv(self, inputs: torch.Tensor) -> torch.Tensor: | |
| forward_context = get_forward_context() | |
| attn_metadata = forward_context.attn_metadata | |
| if attn_metadata is None: | |
| return self._short_conv_fallback(inputs) | |
| if not isinstance(attn_metadata, dict): | |
| raise RuntimeError( | |
| "PLE short-conv expects per-layer attention metadata dict " | |
| f"during inference, got {type(attn_metadata).__name__}." | |
| ) | |
| layer_attn_metadata = attn_metadata.get(self.prefix) | |
| if layer_attn_metadata is None: | |
| raise RuntimeError( | |
| f"Missing short-conv metadata for layer '{self.prefix}'. " | |
| "This would bypass conv-state updates and is not allowed." | |
| ) | |
| if not isinstance(layer_attn_metadata, PleShortConvAttentionMetadata): | |
| raise TypeError( | |
| "Expected PleShortConvAttentionMetadata for layer " | |
| f"'{self.prefix}', got " | |
| f"{type(layer_attn_metadata).__name__}." | |
| ) | |
| conv_state = self.kv_cache[0] | |
| if not is_conv_state_dim_first(): | |
| conv_state = conv_state.transpose(-1, -2) | |
| conv_weights = self.conv1d.weight.squeeze(1) | |
| state_capacity = self.conv_state_len + self.num_spec_tokens | |
| if state_capacity > 0: | |
| if conv_state.size(-1) < state_capacity: | |
| raise RuntimeError( | |
| "PLE short-conv cache is smaller than expected for " | |
| f"dilated convolution: got {conv_state.size(-1)}, " | |
| f"expect at least {state_capacity}." | |
| ) | |
| conv_state = conv_state[..., -state_capacity:] | |
| return self._short_conv_dilated_dispatch( | |
| inputs, | |
| layer_attn_metadata, | |
| conv_state, | |
| conv_weights.to(dtype=inputs.dtype), | |
| ) | |
| def forward( | |
| self, | |
| hidden_states: torch.Tensor, | |
| input_ids: torch.Tensor, | |
| query_start_loc: torch.Tensor, | |
| ngram_context: torch.Tensor, | |
| ) -> torch.Tensor: | |
| input_ids = input_ids.reshape(-1) | |
| if input_ids.shape[0] != hidden_states.shape[0]: | |
| raise ValueError( | |
| "PLE expects input_ids and hidden_states to have the same " | |
| f"token length, got {input_ids.shape[0]} and " | |
| f"{hidden_states.shape[0]}" | |
| ) | |
| embeddings = self.ple_embedding( | |
| hidden_states, | |
| input_ids, | |
| query_start_loc, | |
| ngram_context, | |
| ) | |
| embeddings = self._dequantize_embeddings(embeddings, hidden_states.dtype) | |
| key, _ = self.key_proj(embeddings) | |
| value, _ = self.value_proj(embeddings) | |
| token_count = hidden_states.shape[0] | |
| key = key.reshape(token_count, self.hc_count, self.hidden_size) | |
| query = hidden_states.reshape(token_count, self.hc_count, self.hidden_size) | |
| key = self._apply_norm(self.norm_key, key) | |
| query = self._apply_norm(self.norm_query, query) | |
| gate = (key * query).sum(dim=-1, keepdim=True) / math.sqrt(self.hidden_size) | |
| gate = torch.sigmoid(gate.sign() * gate.abs().clamp_min(1e-6).sqrt()) | |
| gated_value = gate * value.unsqueeze(-2) | |
| normalized = self._apply_norm(self.norm_conv, gated_value).flatten(-2) | |
| conv_output = torch.zeros_like(normalized) | |
| torch.ops.vllm.qwen3_8_flash_next_ple_short_conv( | |
| normalized, | |
| conv_output, | |
| self.prefix, | |
| ) | |
| return gated_value.flatten(-2) + conv_output | |
| def qwen3_8_flash_next_ple_short_conv( | |
| inputs: torch.Tensor, | |
| output: torch.Tensor, | |
| layer_name: str, | |
| ) -> None: | |
| layer = get_forward_context().no_compile_layers[layer_name] | |
| result = layer._short_conv(inputs) | |
| output[: result.shape[0]].copy_(result) | |
| def qwen3_8_flash_next_ple_short_conv_fake( | |
| inputs: torch.Tensor, | |
| output: torch.Tensor, | |
| layer_name: str, | |
| ) -> None: | |
| return | |
| direct_register_custom_op( | |
| op_name="qwen3_8_flash_next_ple_short_conv", | |
| op_func=qwen3_8_flash_next_ple_short_conv, | |
| mutates_args=["output"], | |
| fake_impl=qwen3_8_flash_next_ple_short_conv_fake, | |
| ) | |
| __all__ = [ | |
| "Qwen3_8FlashNextNGramEmbedding", | |
| "Qwen3_8FlashNextPLEGroupedNorm", | |
| "Qwen3_8FlashNextPLELayer", | |
| ] | |