jagat-primitive-org commited on
Commit
10713f4
·
verified ·
1 Parent(s): c7550ff

Upload ple_layer_quant.py with huggingface_hub

Browse files
Files changed (1) hide show
  1. ple_layer_quant.py +1259 -0
ple_layer_quant.py ADDED
@@ -0,0 +1,1259 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # SPDX-License-Identifier: Apache-2.0
2
+ # SPDX-FileCopyrightText: Copyright contributors to the vLLM project
3
+ """GPU-resident Qwen3.8-Flash-Next position-learning enhancement layers."""
4
+
5
+ import math
6
+ from collections.abc import Iterable, Sequence
7
+
8
+ import torch
9
+ import torch.nn.functional as F
10
+ from torch import nn
11
+
12
+ import vllm.envs as envs
13
+ from vllm.config import CacheConfig, ModelConfig, VllmConfig, get_current_vllm_config
14
+ from vllm.forward_context import get_forward_context
15
+ from vllm.model_executor.layers.linear import ReplicatedLinear
16
+ from vllm.model_executor.layers.mamba.abstract import MambaBase
17
+ from vllm.model_executor.layers.mamba.mamba_utils import (
18
+ MambaStateDtypeCalculator,
19
+ MambaStateShapeCalculator,
20
+ is_conv_state_dim_first,
21
+ )
22
+ from vllm.model_executor.layers.ple_offload_layer import (
23
+ PleOffloadLayer,
24
+ is_offload_process,
25
+ )
26
+ from vllm.model_executor.layers.quantization.base_config import (
27
+ QuantizationConfig,
28
+ QuantizeMethodBase,
29
+ )
30
+ from vllm.model_executor.layers.quantization.fp8 import Fp8Config
31
+ from vllm.model_executor.layers.quantization.utils.fp8_utils import (
32
+ create_fp8_scale_parameter,
33
+ create_fp8_weight_parameter,
34
+ is_fp8,
35
+ )
36
+ from vllm.model_executor.layers.quantization.utils.quant_utils import (
37
+ is_layer_skipped,
38
+ )
39
+ from vllm.model_executor.layers.vocab_parallel_embedding import (
40
+ VocabParallelEmbedding,
41
+ )
42
+ from vllm.model_executor.models.utils import AutoWeightsLoader
43
+ from vllm.model_executor.parameter import PerTensorScaleParameter
44
+ from vllm.transformers_utils.configs.qwen3_8_flash_next import (
45
+ Qwen3_8FlashNextTextConfig,
46
+ )
47
+ from vllm.utils.torch_utils import direct_register_custom_op
48
+ from vllm.v1.attention.backends.registry import MambaAttentionBackendEnum
49
+ from vllm.v1.attention.backends.short_conv_attn import (
50
+ PleShortConvAttentionBackend,
51
+ PleShortConvAttentionMetadata,
52
+ )
53
+ from vllm.v1.attention.backends.utils import NULL_BLOCK_ID
54
+
55
+ from ..common.ple import copy_ple_embedding_shard_
56
+
57
+ _MASK64 = (1 << 64) - 1
58
+ _SPLITMIX_GAMMA = 0x9E3779B97F4A7C15
59
+ _SPLITMIX_M1 = 0xBF58476D1CE4E5B9
60
+ _SPLITMIX_M2 = 0x94D049BB133111EB
61
+ _PLE_LAYER_PRIME = 10007
62
+
63
+
64
+ def _splitmix64(value: int) -> int:
65
+ value = (value + _SPLITMIX_GAMMA) & _MASK64
66
+ value = ((value ^ (value >> 30)) * _SPLITMIX_M1) & _MASK64
67
+ value = ((value ^ (value >> 27)) * _SPLITMIX_M2) & _MASK64
68
+ return (value ^ (value >> 31)) & _MASK64
69
+
70
+
71
+ def _is_prime_64(value: int) -> bool:
72
+ if value < 2:
73
+ return False
74
+ for prime in (2, 3, 5, 7, 11, 13, 17, 19, 23, 29, 31, 37):
75
+ if value % prime == 0:
76
+ return value == prime
77
+ exponent = value - 1
78
+ shifts = 0
79
+ while exponent % 2 == 0:
80
+ exponent //= 2
81
+ shifts += 1
82
+ for base in (2, 325, 9375, 28178, 450775, 9780504, 1795265022):
83
+ if base % value == 0:
84
+ continue
85
+ witness = pow(base, exponent, value)
86
+ if witness in (1, value - 1):
87
+ continue
88
+ for _ in range(shifts - 1):
89
+ witness = pow(witness, 2, value)
90
+ if witness == value - 1:
91
+ break
92
+ else:
93
+ return False
94
+ return True
95
+
96
+
97
+ def _nth_prime_after(start: int, count: int) -> int:
98
+ prime = int(start)
99
+ for _ in range(count):
100
+ candidate = prime + 1
101
+ if candidate <= 2:
102
+ prime = 2
103
+ continue
104
+ if candidate % 2 == 0:
105
+ candidate += 1
106
+ while not _is_prime_64(candidate):
107
+ candidate += 2
108
+ prime = candidate
109
+ return prime
110
+
111
+
112
+ class Qwen3_8FlashNextPLEGroupedNorm(nn.Module):
113
+ def __init__(
114
+ self,
115
+ hidden_size: int,
116
+ eps: float,
117
+ group_size: int | None,
118
+ dtype: torch.dtype | None,
119
+ ) -> None:
120
+ super().__init__()
121
+ if group_size is not None and hidden_size % group_size:
122
+ raise ValueError(
123
+ f"hidden_size ({hidden_size}) must be divisible by "
124
+ f"group_size ({group_size})"
125
+ )
126
+ self.eps = eps
127
+ self.group_size = group_size
128
+ self.weight = nn.Parameter(torch.zeros(hidden_size, dtype=dtype))
129
+
130
+ def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
131
+ input_dtype = hidden_states.dtype
132
+ hidden_states = hidden_states.float()
133
+ if self.group_size is None:
134
+ variance = hidden_states.square().mean(dim=-1, keepdim=True)
135
+ normalized = hidden_states * torch.rsqrt(variance + self.eps)
136
+ else:
137
+ grouped = hidden_states.unflatten(
138
+ -1, (hidden_states.shape[-1] // self.group_size, self.group_size)
139
+ )
140
+ variance = grouped.square().mean(dim=-1, keepdim=True)
141
+ normalized = (grouped * torch.rsqrt(variance + self.eps)).flatten(-2)
142
+ return (normalized * (1.0 + self.weight.float())).to(input_dtype)
143
+
144
+
145
+ class Qwen3_8FlashNextPLEFp8EmbeddingMethod(QuantizeMethodBase):
146
+ """FP8 PLE embedding with one global checkpoint scale."""
147
+
148
+ def create_weights(
149
+ self,
150
+ layer: nn.Module,
151
+ input_size_per_partition: int,
152
+ output_partition_sizes: list[int],
153
+ input_size: int,
154
+ output_size: int,
155
+ params_dtype: torch.dtype,
156
+ **extra_weight_attrs,
157
+ ) -> None:
158
+ del input_size, output_size, params_dtype
159
+ weight_loader = extra_weight_attrs.get("weight_loader")
160
+ weight = create_fp8_weight_parameter(
161
+ sum(output_partition_sizes), input_size_per_partition, weight_loader
162
+ )
163
+ layer.register_parameter("weight", weight)
164
+
165
+ weight_scale = create_fp8_scale_parameter(
166
+ PerTensorScaleParameter,
167
+ output_partition_sizes,
168
+ input_size_per_partition,
169
+ None,
170
+ weight_loader,
171
+ scale_dtype=torch.bfloat16,
172
+ )
173
+ layer.register_parameter("weight_scale", weight_scale)
174
+
175
+ def apply(
176
+ self,
177
+ layer: nn.Module,
178
+ x: torch.Tensor,
179
+ bias: torch.Tensor | None = None,
180
+ ) -> torch.Tensor:
181
+ raise NotImplementedError("PLE FP8 weights only support embedding lookup")
182
+
183
+ def embedding(self, layer: nn.Module, input_: torch.Tensor) -> torch.Tensor:
184
+ return F.embedding(input_, layer.weight)
185
+
186
+
187
+ def _get_ple_embedding_quant_method(
188
+ quant_config: QuantizationConfig | None,
189
+ prefix: str,
190
+ ) -> QuantizeMethodBase | None:
191
+ """Select global-scale FP8 only for quantized PLE checkpoint shards."""
192
+
193
+ if not isinstance(quant_config, Fp8Config):
194
+ return None
195
+ if not quant_config.is_checkpoint_fp8_serialized:
196
+ return None
197
+
198
+ ignored_layers = quant_config.ignored_layers
199
+ if is_layer_skipped(
200
+ prefix,
201
+ ignored_layers,
202
+ quant_config.packed_modules_mapping,
203
+ match_mode=quant_config.ignored_layers_match_mode,
204
+ ):
205
+ return None
206
+ # PLE checkpoint shards form one runtime embedding parameter.
207
+ shard_prefix = f"{prefix}.shard_"
208
+ if any(name.startswith(shard_prefix) for name in ignored_layers):
209
+ return None
210
+ return Qwen3_8FlashNextPLEFp8EmbeddingMethod()
211
+
212
+
213
+ class Qwen3_8FlashNextNGramEmbedding(PleOffloadLayer):
214
+ def __init__(
215
+ self,
216
+ config: Qwen3_8FlashNextTextConfig,
217
+ embedding_dim: int,
218
+ ple_dense_layer_id: int,
219
+ max_total_tokens: int,
220
+ max_num_reqs: int,
221
+ prefix: str,
222
+ quant_config: QuantizationConfig | None = None,
223
+ params_dtype: torch.dtype | None = None,
224
+ ) -> None:
225
+ super().__init__()
226
+ self.embedding_dim = embedding_dim
227
+ self.ngram_size = int(config.ngram_size)
228
+ self.heads_per_ngram = int(config.heads_per_ngram)
229
+ self.ngram_heads = (self.ngram_size - 1) * self.heads_per_ngram
230
+ if self.ngram_size < 2:
231
+ raise ValueError(f"ngram_size must be >= 2, got {self.ngram_size}")
232
+ if self.heads_per_ngram <= 0:
233
+ raise ValueError(f"heads_per_ngram must be > 0, got {self.heads_per_ngram}")
234
+ if embedding_dim % self.ngram_heads:
235
+ raise ValueError(
236
+ "ple_embed_dim must be divisible by total ngram heads: "
237
+ f"{embedding_dim} % {self.ngram_heads} != 0"
238
+ )
239
+ self.head_dim = embedding_dim // self.ngram_heads
240
+ self.eos_token_id = int(config.eos_token_id)
241
+ self.unigram_vocab_size = int(config.vocab_size)
242
+ self.split_ngram_parts = int(getattr(config, "split_ngram_parts", 512))
243
+ if self.split_ngram_parts <= 0:
244
+ raise ValueError("split_ngram_parts must be positive")
245
+
246
+ max_multiplier = ((1 << 63) - 1) // self.unigram_vocab_size
247
+ half_bound = max(1, max_multiplier // 2)
248
+ seed = int(getattr(config, "seed", 1234))
249
+ base_seed = seed + _PLE_LAYER_PRIME * ple_dense_layer_id
250
+ multipliers = []
251
+ for index in range(self.ngram_size):
252
+ value = base_seed + _SPLITMIX_GAMMA * (index + 1)
253
+ multipliers.append(2 * (_splitmix64(value) % half_bound) + 1)
254
+ self.register_buffer(
255
+ "layer_multipliers",
256
+ torch.tensor(multipliers, dtype=torch.long),
257
+ persistent=True,
258
+ )
259
+
260
+ ngram_vocab_size_base = int(config.ngram_vocab_size_base)
261
+ sizes: list[int] = []
262
+ offsets: list[int] = []
263
+ offset = 0
264
+ for local_head in range(self.ngram_heads):
265
+ global_head = ple_dense_layer_id * self.ngram_heads + local_head
266
+ size = _nth_prime_after(ngram_vocab_size_base - 1, global_head + 1)
267
+ sizes.append(size)
268
+ offsets.append(offset)
269
+ offset += size
270
+ self.register_buffer(
271
+ "ngram_heads_vocab_sizes",
272
+ torch.tensor(sizes, dtype=torch.long),
273
+ persistent=True,
274
+ )
275
+ self.register_buffer(
276
+ "ngram_heads_offsets",
277
+ torch.tensor(offsets, dtype=torch.long),
278
+ persistent=True,
279
+ )
280
+ divisor = int(config.make_ngram_vocab_size_divisible_by)
281
+ padded_vocab_size = ((offset + divisor - 1) // divisor) * divisor
282
+ self.ngram_embedding = VocabParallelEmbedding(
283
+ padded_vocab_size,
284
+ self.head_dim,
285
+ params_dtype=params_dtype,
286
+ padding_size=divisor,
287
+ prefix=f"{prefix}.ngram_embedding",
288
+ quant_method=_get_ple_embedding_quant_method(
289
+ quant_config, f"{prefix}.ngram_embedding"
290
+ ),
291
+ )
292
+ self.register_buffer(
293
+ "positions_buffer",
294
+ torch.arange(max_total_tokens, dtype=torch.int64),
295
+ persistent=False,
296
+ )
297
+ self.register_buffer(
298
+ "padded_buffer",
299
+ torch.full(
300
+ (max_num_reqs, max_total_tokens),
301
+ self.eos_token_id,
302
+ dtype=torch.int64,
303
+ ),
304
+ persistent=False,
305
+ )
306
+
307
+ @staticmethod
308
+ def _shift_precompute(
309
+ tokens: torch.Tensor, eos_token_id: int
310
+ ) -> tuple[torch.Tensor, torch.Tensor]:
311
+ if tokens.dim() != 2:
312
+ raise ValueError("tokens must be a 2D tensor")
313
+ batch_size, seq_len = tokens.shape
314
+ positions = torch.arange(seq_len, device=tokens.device, dtype=torch.int64)
315
+ eos_positions = torch.where(tokens == eos_token_id, positions, -1)
316
+ previous_eos_inclusive = torch.cummax(eos_positions, dim=1).values
317
+ previous_eos = torch.cat(
318
+ [
319
+ eos_positions.new_full((batch_size, 1), -1),
320
+ previous_eos_inclusive[:, :-1],
321
+ ],
322
+ dim=1,
323
+ )
324
+ return positions, positions.unsqueeze(0) - previous_eos - 1
325
+
326
+ @staticmethod
327
+ def _shift_apply(
328
+ tokens: torch.Tensor,
329
+ positions: torch.Tensor,
330
+ position_in_segment: torch.Tensor,
331
+ shift: int,
332
+ eos_token_id: int,
333
+ ) -> torch.Tensor:
334
+ if shift == 0:
335
+ return tokens
336
+ source = positions - shift
337
+ gather_indices = source.clamp_min(0).unsqueeze(0).expand(tokens.shape[0], -1)
338
+ shifted = tokens.gather(1, gather_indices)
339
+ valid = (source.unsqueeze(0) >= 0) & (position_in_segment >= shift)
340
+ return torch.where(valid, shifted, tokens.new_full((), eos_token_id))
341
+
342
+ def forward_impl( # type: ignore[override]
343
+ self,
344
+ hidden_states: torch.Tensor,
345
+ input_ids: torch.Tensor,
346
+ query_start_loc: torch.Tensor,
347
+ ngram_context: torch.Tensor,
348
+ output_buffer: torch.Tensor | None = None,
349
+ ) -> torch.Tensor:
350
+ del hidden_states
351
+ input_ids = input_ids.reshape(-1).long()
352
+ query_start_loc = query_start_loc.long()
353
+ num_reqs = query_start_loc.numel() - 1
354
+ num_tokens = input_ids.shape[0]
355
+ if num_tokens > self.positions_buffer.numel():
356
+ raise ValueError(
357
+ f"PLE received {num_tokens} tokens, but its workspace supports "
358
+ f"at most {self.positions_buffer.numel()}"
359
+ )
360
+ if num_reqs > self.padded_buffer.shape[0]:
361
+ raise ValueError(
362
+ f"PLE received {num_reqs} requests, but its workspace supports "
363
+ f"at most {self.padded_buffer.shape[0]}"
364
+ )
365
+
366
+ # The CPU-offload subprocess is never captured by a CUDA Graph, so its
367
+ # pack workspace can narrow to the actual maximum sequence length. The
368
+ # regular GPU path retains the static maximum-width buffer for capture.
369
+ if is_offload_process():
370
+ if num_reqs <= 0:
371
+ raise ValueError("PLE CPU offload requires at least one request")
372
+ max_seq_len = max(
373
+ 1,
374
+ int((query_start_loc[1:] - query_start_loc[:-1]).max().item()),
375
+ )
376
+ # The model runner sends the CUDA-graph padded token count together
377
+ # with an unpadded query_start_loc. Stale padding must not enter the
378
+ # scatter: its clamped indices would overwrite the last real token.
379
+ num_valid_tokens = min(int(query_start_loc[-1].item()), num_tokens)
380
+ else:
381
+ max_seq_len = self.padded_buffer.shape[1]
382
+ num_valid_tokens = num_tokens
383
+
384
+ positions = self.positions_buffer[:num_tokens]
385
+ packed = self.padded_buffer[:num_reqs, :max_seq_len]
386
+ packed.fill_(self.eos_token_id)
387
+ request_indices = torch.searchsorted(query_start_loc, positions, right=True) - 1
388
+ request_indices.clamp_(max=num_reqs - 1)
389
+ columns = (positions - query_start_loc[request_indices]).clamp(
390
+ 0, packed.shape[1] - 1
391
+ )
392
+ packed[request_indices[:num_valid_tokens], columns[:num_valid_tokens]] = (
393
+ input_ids[:num_valid_tokens]
394
+ )
395
+ ngram_context = ngram_context[:num_reqs].to(
396
+ device=input_ids.device, dtype=torch.long
397
+ )
398
+
399
+ context = torch.cat([ngram_context, packed], dim=-1)
400
+ positions_2d, position_in_segment = self._shift_precompute(
401
+ context, self.eos_token_id
402
+ )
403
+ shifted = [context]
404
+ for shift in range(1, self.ngram_size):
405
+ shifted.append(
406
+ self._shift_apply(
407
+ context,
408
+ positions_2d,
409
+ position_in_segment,
410
+ shift,
411
+ self.eos_token_id,
412
+ )
413
+ )
414
+ adjusted_columns = columns + self.ngram_size - 1
415
+ id_blocks = []
416
+ for ngram in range(2, self.ngram_size + 1):
417
+ start = (ngram - 2) * self.heads_per_ngram
418
+ end = start + self.heads_per_ngram
419
+ mixed = shifted[0] * self.layer_multipliers[0]
420
+ for index in range(1, ngram):
421
+ mixed = torch.bitwise_xor(
422
+ mixed, shifted[index] * self.layer_multipliers[index]
423
+ )
424
+ sizes = self.ngram_heads_vocab_sizes[start:end]
425
+ offsets = self.ngram_heads_offsets[start:end]
426
+ ids = torch.remainder(mixed.unsqueeze(-1), sizes) + offsets
427
+ id_blocks.append(ids[request_indices, adjusted_columns])
428
+ ngram_ids = torch.cat(id_blocks, dim=-1)
429
+ quant = getattr(self.ngram_embedding, "_ple_quant", None)
430
+ if output_buffer is not None:
431
+ output = output_buffer[:num_tokens, : self.embedding_dim]
432
+ if quant is not None:
433
+ quant.gather_into(
434
+ ngram_ids.reshape(-1), output.reshape(-1, self.head_dim)
435
+ )
436
+ else:
437
+ torch.index_select(
438
+ self.ngram_embedding.weight,
439
+ 0,
440
+ ngram_ids.reshape(-1),
441
+ out=output.reshape(-1, self.head_dim),
442
+ )
443
+ return output
444
+ if quant is not None:
445
+ flat = torch.empty(
446
+ ngram_ids.numel(),
447
+ self.head_dim,
448
+ dtype=torch.bfloat16,
449
+ device=ngram_ids.device,
450
+ )
451
+ quant.gather_into(ngram_ids.reshape(-1), flat)
452
+ return flat.view(*ngram_ids.shape, self.head_dim).flatten(-2)
453
+ return self.ngram_embedding(ngram_ids).flatten(-2)
454
+
455
+ def get_offload_output_dtype(self, default_dtype: torch.dtype) -> torch.dtype:
456
+ """Keep quantized lookup results in their embedding storage dtype."""
457
+ embedding = getattr(self, "ngram_embedding", None)
458
+ weight = getattr(embedding, "weight", None)
459
+ if weight is not None:
460
+ return weight.dtype
461
+ if hasattr(self, "_offload_weight_scale"):
462
+ return torch.float8_e4m3fn
463
+ return default_dtype
464
+
465
+ def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]:
466
+ """Load hash buffers and checkpoint-split embedding rows."""
467
+
468
+ # GPU workers retain only the global FP8 scale. The CPU process owns the
469
+ # embedding weight and returns its quantized lookup output unchanged.
470
+ if envs.VLLM_PLE_CPU_OFFLOAD and not is_offload_process():
471
+ retained: set[str] = set()
472
+ for name, loaded_weight in weights:
473
+ if name != "ngram_embedding.weight_scale":
474
+ continue
475
+ self.register_buffer(
476
+ "_offload_weight_scale",
477
+ loaded_weight.to(device=torch.accelerator.current_accelerator()),
478
+ persistent=False,
479
+ )
480
+ retained.add(name)
481
+ return retained
482
+
483
+ persistent_buffers = {
484
+ "layer_multipliers": self.layer_multipliers,
485
+ "ngram_heads_offsets": self.ngram_heads_offsets,
486
+ "ngram_heads_vocab_sizes": self.ngram_heads_vocab_sizes,
487
+ }
488
+ loaded: set[str] = set()
489
+ regular_weights: list[tuple[str, torch.Tensor]] = []
490
+ shard_prefix = "ngram_embedding.shard_"
491
+
492
+ for name, loaded_weight in weights:
493
+ leaf_name = name.rsplit(".", 1)[-1]
494
+ if leaf_name.startswith("hashstats_") or leaf_name == "token_lookup":
495
+ continue
496
+ if name in persistent_buffers:
497
+ buffer = persistent_buffers[name]
498
+ if buffer.shape != loaded_weight.shape:
499
+ raise ValueError(
500
+ f"Shape mismatch for {name}: expected "
501
+ f"{tuple(buffer.shape)}, got {tuple(loaded_weight.shape)}"
502
+ )
503
+ buffer.copy_(loaded_weight.to(device=buffer.device, dtype=buffer.dtype))
504
+ loaded.add(name)
505
+ continue
506
+ if name.startswith(shard_prefix) and name.endswith(".weight"):
507
+ shard_text = name[len(shard_prefix) : -len(".weight")]
508
+ if not shard_text.isdigit():
509
+ regular_weights.append((name, loaded_weight))
510
+ continue
511
+ shard_index = int(shard_text)
512
+ if shard_index >= self.split_ngram_parts:
513
+ raise ValueError(
514
+ f"PLE embedding shard index {shard_index} exceeds "
515
+ f"split_ngram_parts={self.split_ngram_parts}"
516
+ )
517
+ embedding = self.ngram_embedding
518
+ shard_size = (
519
+ embedding.org_vocab_size + self.split_ngram_parts - 1
520
+ ) // self.split_ngram_parts
521
+ checkpoint_start = shard_index * shard_size
522
+ expected_rows = max(
523
+ 0,
524
+ min(shard_size, embedding.org_vocab_size - checkpoint_start),
525
+ )
526
+ expected_shape = (expected_rows, embedding.embedding_dim)
527
+ if tuple(loaded_weight.shape) != expected_shape:
528
+ raise ValueError(
529
+ f"Shape mismatch for PLE embedding shard {shard_index}: "
530
+ f"expected {expected_shape}, got "
531
+ f"{tuple(loaded_weight.shape)}"
532
+ )
533
+ copy_ple_embedding_shard_(
534
+ embedding.weight.data,
535
+ loaded_weight,
536
+ checkpoint_start=checkpoint_start,
537
+ tp_start=embedding.shard_indices.org_vocab_start_index,
538
+ tp_end=embedding.shard_indices.org_vocab_end_index,
539
+ )
540
+ loaded.add("ngram_embedding.weight")
541
+ continue
542
+ regular_weights.append((name, loaded_weight))
543
+
544
+ if regular_weights:
545
+ loaded.update(AutoWeightsLoader(self).load_weights(regular_weights))
546
+ return loaded
547
+
548
+
549
+ class Qwen3_8FlashNextPLELayer(nn.Module, MambaBase):
550
+ def __init__(
551
+ self,
552
+ config: Qwen3_8FlashNextTextConfig,
553
+ vllm_config: VllmConfig,
554
+ layer_idx: int = 0,
555
+ ple_dense_layer_id: int | None = None,
556
+ prefix: str = "",
557
+ ) -> None:
558
+ super().__init__()
559
+ model_config = vllm_config.model_config
560
+ cache_config = vllm_config.cache_config
561
+ quant_config = vllm_config.quant_config
562
+ self.model_config: ModelConfig = model_config
563
+ self.cache_config: CacheConfig = cache_config
564
+ self.layer_idx = layer_idx
565
+ self.ple_dense_layer_id = (
566
+ int(ple_dense_layer_id)
567
+ if ple_dense_layer_id is not None
568
+ else int(layer_idx)
569
+ )
570
+ self.prefix = prefix
571
+ self.hidden_size = int(config.hidden_size)
572
+ self.hc_count = config.hc_count
573
+ self.hc_hidden_size = self.hidden_size * self.hc_count
574
+ self.conv_kernel_size = int(config.ple_conv_kernel_size)
575
+ self.short_conv_dilation = int(config.ngram_size)
576
+ self.conv_state_len = (self.conv_kernel_size - 1) * self.short_conv_dilation
577
+ self.num_spec_tokens = vllm_config.num_speculative_tokens
578
+ self.activation = "silu"
579
+ # The offload process builds the surrounding model on meta while
580
+ # this subtree must own real CPU storage. GPU workers skip the
581
+ # subclass constructor and retain only an empty IPC placeholder.
582
+ with torch.device(PleOffloadLayer.get_target_device()):
583
+ self.ple_embedding: nn.Module = Qwen3_8FlashNextNGramEmbedding(
584
+ config,
585
+ int(config.ple_embed_dim),
586
+ self.ple_dense_layer_id,
587
+ vllm_config.scheduler_config.max_num_batched_tokens,
588
+ vllm_config.scheduler_config.max_num_seqs,
589
+ f"{prefix}.ple_embedding",
590
+ quant_config=quant_config,
591
+ params_dtype=model_config.dtype,
592
+ )
593
+ self.key_proj = ReplicatedLinear(
594
+ int(config.ple_embed_dim),
595
+ self.hc_hidden_size,
596
+ bias=False,
597
+ quant_config=quant_config,
598
+ prefix=f"{prefix}.key_proj",
599
+ )
600
+ self.value_proj = ReplicatedLinear(
601
+ int(config.ple_embed_dim),
602
+ self.hidden_size,
603
+ bias=False,
604
+ quant_config=quant_config,
605
+ prefix=f"{prefix}.value_proj",
606
+ )
607
+ norm_args = (
608
+ self.hc_hidden_size,
609
+ config.rms_norm_eps,
610
+ self.hidden_size,
611
+ model_config.dtype,
612
+ )
613
+ self.norm_key = Qwen3_8FlashNextPLEGroupedNorm(*norm_args)
614
+ self.norm_query = Qwen3_8FlashNextPLEGroupedNorm(*norm_args)
615
+ self.norm_conv = Qwen3_8FlashNextPLEGroupedNorm(*norm_args)
616
+ self.conv1d = nn.Conv1d(
617
+ self.hc_hidden_size,
618
+ self.hc_hidden_size,
619
+ self.conv_kernel_size,
620
+ groups=self.hc_hidden_size,
621
+ padding=self.conv_state_len,
622
+ dilation=self.short_conv_dilation,
623
+ bias=False,
624
+ dtype=model_config.dtype,
625
+ )
626
+ nn.init.zeros_(self.conv1d.weight)
627
+ self.conv1d.weight._no_reinit = True
628
+ self.kv_cache = (torch.tensor([]),)
629
+ compilation_config = get_current_vllm_config().compilation_config
630
+ if prefix in compilation_config.static_forward_context:
631
+ raise ValueError(f"Duplicate layer name: {prefix}")
632
+ compilation_config.static_forward_context[prefix] = self
633
+
634
+ def _get_embedding_weight_scale(self) -> torch.Tensor | None:
635
+ embedding = getattr(self.ple_embedding, "ngram_embedding", None)
636
+ weight_scale = getattr(embedding, "weight_scale", None)
637
+ if weight_scale is not None:
638
+ return weight_scale
639
+ return getattr(self.ple_embedding, "_offload_weight_scale", None)
640
+
641
+ def _dequantize_embeddings(
642
+ self,
643
+ embeddings: torch.Tensor,
644
+ output_dtype: torch.dtype,
645
+ ) -> torch.Tensor:
646
+ """Dequantize PLE lookup output."""
647
+
648
+ if not is_fp8(embeddings):
649
+ return embeddings
650
+ weight_scale = self._get_embedding_weight_scale()
651
+ if weight_scale is None:
652
+ raise RuntimeError("FP8 PLE embedding is missing its global scale")
653
+ if weight_scale.device != embeddings.device:
654
+ raise RuntimeError("FP8 PLE embedding scale must be on the output device")
655
+ return embeddings.to(output_dtype) * weight_scale.to(output_dtype)
656
+
657
+ @property
658
+ def mamba_type(self) -> MambaAttentionBackendEnum:
659
+ return MambaAttentionBackendEnum.SHORT_CONV
660
+
661
+ @property
662
+ def is_kv_cache_tp_replicated(self) -> bool:
663
+ return True
664
+
665
+ def get_attn_backend(self) -> type[PleShortConvAttentionBackend]:
666
+ return PleShortConvAttentionBackend
667
+
668
+ def get_state_dtype(self) -> tuple[torch.dtype, ...]:
669
+ return MambaStateDtypeCalculator.short_conv_state_dtype(
670
+ self.model_config.dtype, self.cache_config.mamba_cache_dtype
671
+ )
672
+
673
+ def get_state_shape(self) -> Sequence[tuple[int, ...]]:
674
+ return MambaStateShapeCalculator.short_conv_state_shape(
675
+ tp_world_size=1,
676
+ intermediate_size=self.hc_hidden_size,
677
+ conv_kernel=self.conv_state_len + 1,
678
+ num_spec=self.num_spec_tokens,
679
+ )
680
+
681
+ def _apply_norm(
682
+ self, norm: Qwen3_8FlashNextPLEGroupedNorm, hidden_states: torch.Tensor
683
+ ) -> torch.Tensor:
684
+ shape = hidden_states.shape
685
+ return norm(hidden_states.flatten(-2)).reshape(shape)
686
+
687
+ def _short_conv_fallback(self, inputs: torch.Tensor) -> torch.Tensor:
688
+ # Profiling / CUDA graph capture only; conv state is not updated.
689
+ inputs_t = inputs.transpose(0, 1).unsqueeze(0)
690
+ output = self.conv1d(inputs_t)[..., : inputs_t.size(-1)]
691
+ return F.silu(output).squeeze(0).transpose(0, 1)
692
+
693
+ def _short_conv_dilated_decode_batched(
694
+ self,
695
+ x_d: torch.Tensor,
696
+ conv_state: torch.Tensor,
697
+ conv_weights: torch.Tensor,
698
+ state_indices_tensor_d: torch.Tensor,
699
+ has_initial_states_d: torch.Tensor | None,
700
+ ) -> torch.Tensor:
701
+ state_indices = state_indices_tensor_d.to(
702
+ device=conv_state.device, dtype=torch.int64
703
+ )
704
+ # TODO: need double-check
705
+ # FULL cudagraph padded decode rows use NULL_BLOCK_ID. Remap them to
706
+ # slot 0 for a safe gather, then zero output and skip write-back.
707
+ valid_state = state_indices != NULL_BLOCK_ID
708
+ state_indices = torch.where(
709
+ valid_state, state_indices, torch.zeros_like(state_indices)
710
+ )
711
+ if has_initial_states_d is None:
712
+ has_initial_state = valid_state
713
+ else:
714
+ if has_initial_states_d.numel() < state_indices_tensor_d.numel():
715
+ raise ValueError(
716
+ "has_initial_states_d size mismatch: "
717
+ f"got {has_initial_states_d.numel()}, "
718
+ f"need >= {state_indices_tensor_d.numel()}."
719
+ )
720
+ has_initial_state = has_initial_states_d[
721
+ : state_indices_tensor_d.numel()
722
+ ].to(device=conv_state.device, dtype=torch.bool)
723
+ has_initial_state = has_initial_state & valid_state
724
+
725
+ cached_state = conv_state.index_select(0, state_indices)
726
+ state = cached_state[..., : self.conv_state_len].to(x_d.dtype)
727
+ if self.conv_state_len > 0:
728
+ initial_state = torch.where(
729
+ has_initial_state.view(-1, 1, 1),
730
+ state,
731
+ torch.zeros_like(state),
732
+ )
733
+ history = torch.cat((initial_state, x_d.unsqueeze(-1)), dim=-1)
734
+ else:
735
+ history = x_d.unsqueeze(-1)
736
+
737
+ conv_output = F.conv1d(
738
+ history,
739
+ conv_weights.unsqueeze(1).contiguous(),
740
+ groups=history.size(1),
741
+ dilation=self.short_conv_dilation,
742
+ ).squeeze(-1)
743
+ output = F.silu(conv_output)
744
+ output = output * valid_state.view(-1, 1).to(output.dtype)
745
+
746
+ if self.conv_state_len > 0:
747
+ next_state = history[..., -self.conv_state_len :]
748
+ # Padded rows are remapped to the reserved null slot. Preserve its
749
+ # existing value while writing the new states for valid rows.
750
+ existing_base_state = cached_state[..., : self.conv_state_len]
751
+ safe_next_state = torch.where(
752
+ valid_state.view(-1, 1, 1),
753
+ next_state.to(conv_state.dtype),
754
+ existing_base_state,
755
+ )
756
+ cached_state[..., : self.conv_state_len] = safe_next_state
757
+ conv_state.index_copy_(0, state_indices, cached_state)
758
+
759
+ return output
760
+
761
+ def _short_conv_dilated_prefill_batched(
762
+ self,
763
+ x_p: torch.Tensor,
764
+ metadata: PleShortConvAttentionMetadata,
765
+ conv_state: torch.Tensor,
766
+ conv_weights: torch.Tensor,
767
+ state_indices_tensor_p: torch.Tensor,
768
+ num_prefills: int,
769
+ num_decode_tokens: int,
770
+ num_prefill_tokens: int,
771
+ ) -> torch.Tensor:
772
+ # ``non_spec_query_start_loc`` covers the non-spec (decode + prefill)
773
+ # requests and equals ``query_start_loc`` when spec-decode is inactive.
774
+ non_spec_query_start_loc = metadata.non_spec_query_start_loc
775
+ if non_spec_query_start_loc is None:
776
+ raise ValueError("query_start_loc is required for prefill short-conv")
777
+ query_start_loc_p = (
778
+ non_spec_query_start_loc[-num_prefills - 1 :] - num_decode_tokens
779
+ )
780
+ # The metadata builder guarantees that the prefill query offsets start
781
+ # at 0 and end at num_prefill_tokens. Avoid reading those values here,
782
+ # since doing so would force a device-to-host synchronization.
783
+ has_initial_states_p = metadata.has_initial_states_p
784
+ if has_initial_states_p is None:
785
+ raise ValueError("has_initial_states_p is required for prefill short-conv")
786
+
787
+ output = torch.empty_like(x_p)
788
+ q_starts = query_start_loc_p.to(torch.int64)
789
+ if state_indices_tensor_p.numel() < num_prefills:
790
+ raise ValueError(
791
+ "state_indices_tensor_p size mismatch: "
792
+ f"got {state_indices_tensor_p.numel()}, "
793
+ f"need >= {num_prefills}."
794
+ )
795
+ if has_initial_states_p.numel() < num_prefills:
796
+ raise ValueError(
797
+ "has_initial_states_p size mismatch: "
798
+ f"got {has_initial_states_p.numel()}, "
799
+ f"need >= {num_prefills}."
800
+ )
801
+ if num_prefills == 0 or x_p.numel() == 0:
802
+ return output
803
+ lengths = q_starts[1:] - q_starts[:-1]
804
+ # Use the CPU-computed packing width from the metadata builder instead
805
+ # of synchronizing on lengths.max().
806
+ max_len = metadata.max_prefill_query_len
807
+ if max_len <= 0:
808
+ return output
809
+
810
+ hidden_size = x_p.shape[1]
811
+ positions = torch.arange(
812
+ num_prefill_tokens, device=x_p.device, dtype=torch.int64
813
+ )
814
+ req_indices = torch.searchsorted(q_starts[1:], positions, right=True)
815
+ col_indices = positions - q_starts[req_indices]
816
+
817
+ packed_tokens = x_p.new_zeros((num_prefills, max_len, hidden_size))
818
+ packed_tokens[req_indices, col_indices] = x_p
819
+ packed_tokens = packed_tokens.transpose(1, 2).contiguous()
820
+
821
+ state_indices = state_indices_tensor_p[:num_prefills].to(
822
+ device=conv_state.device, dtype=torch.int64
823
+ )
824
+ valid_state = state_indices != NULL_BLOCK_ID
825
+ state_indices = torch.where(
826
+ valid_state, state_indices, torch.zeros_like(state_indices)
827
+ )
828
+ has_initial = has_initial_states_p[:num_prefills].to(
829
+ device=conv_state.device, dtype=torch.bool
830
+ )
831
+ if self.conv_state_len > 0:
832
+ if conv_state.shape[0] == 0:
833
+ state = conv_state.new_zeros(
834
+ (num_prefills, hidden_size, self.conv_state_len),
835
+ dtype=x_p.dtype,
836
+ )
837
+ else:
838
+ state = conv_state.index_select(0, state_indices)[
839
+ ..., : self.conv_state_len
840
+ ].to(x_p.dtype)
841
+ use_initial_mask = (valid_state & has_initial).view(num_prefills, 1, 1)
842
+ initial_state = torch.where(
843
+ use_initial_mask,
844
+ state,
845
+ torch.zeros_like(state),
846
+ )
847
+ history = torch.cat((initial_state, packed_tokens), dim=-1)
848
+ else:
849
+ history = packed_tokens
850
+
851
+ conv_output = F.conv1d(
852
+ history,
853
+ conv_weights.unsqueeze(1).contiguous(),
854
+ groups=history.size(1),
855
+ dilation=self.short_conv_dilation,
856
+ )
857
+ conv_output = F.silu(conv_output).transpose(1, 2).contiguous()
858
+
859
+ token_positions = torch.arange(max_len, device=x_p.device, dtype=torch.int64)
860
+ valid_tokens = token_positions.view(1, max_len) < lengths.view(num_prefills, 1)
861
+ valid_output_mask = valid_tokens & valid_state.to(device=x_p.device).view(
862
+ num_prefills, 1
863
+ )
864
+ conv_output.masked_fill_(~valid_output_mask.unsqueeze(-1), 0)
865
+ output.copy_(conv_output[req_indices, col_indices])
866
+
867
+ if self.conv_state_len > 0 and conv_state.shape[0] > 0:
868
+ state_starts = lengths.to(device=history.device, dtype=torch.int64).view(
869
+ num_prefills, 1, 1
870
+ )
871
+ state_offsets = torch.arange(
872
+ self.conv_state_len, device=history.device, dtype=torch.int64
873
+ ).view(1, 1, self.conv_state_len)
874
+ next_state = history.gather(
875
+ dim=2,
876
+ index=(state_starts + state_offsets).expand(-1, history.size(1), -1),
877
+ )
878
+ # Write back without a host synchronization. Valid, non-empty rows
879
+ # receive their new state; padding and zero-length rows keep the
880
+ # current cache value.
881
+ existing_state = conv_state.index_select(0, state_indices)
882
+ existing_base_state = existing_state[..., : self.conv_state_len]
883
+ update_mask = valid_state & (lengths.to(device=conv_state.device) > 0)
884
+ safe_next_state = torch.where(
885
+ update_mask.view(num_prefills, 1, 1),
886
+ next_state.to(conv_state.dtype),
887
+ existing_base_state,
888
+ )
889
+ existing_state[..., : self.conv_state_len] = safe_next_state
890
+ conv_state.index_copy_(0, state_indices, existing_state)
891
+ return output
892
+
893
+ def _short_conv_dilated_spec_batched(
894
+ self,
895
+ x_spec: torch.Tensor,
896
+ conv_state: torch.Tensor,
897
+ conv_weights: torch.Tensor,
898
+ spec_state_indices_tensor: torch.Tensor,
899
+ spec_query_start_loc: torch.Tensor,
900
+ num_accepted_tokens: torch.Tensor,
901
+ spec_query_len: int,
902
+ ) -> torch.Tensor:
903
+ """Dilated short-conv for speculative-decode (MTP) requests.
904
+
905
+ Each spec request feeds multiple (draft + 1) query tokens. The conv
906
+ outputs are computed causally after rolling back the previous draft
907
+ state by ``num_accepted_tokens - 1``. The current candidate inputs stay
908
+ in the extended cache for the next forward, matching
909
+ ``causal_conv1d_update``.
910
+
911
+ ``spec_query_len`` (== num_speculative_tokens + 1) is the maximum query
912
+ length and is a Python int, so no host synchronization is needed; this
913
+ keeps the path safe for full CUDA-graph capture/replay where the buffers
914
+ are padded at the request level.
915
+ """
916
+ num_reqs = spec_state_indices_tensor.numel()
917
+ hidden_size = x_spec.size(-1)
918
+ # Use a fixed packing width instead of synchronizing on lengths.max().
919
+ max_len = spec_query_len
920
+ # Full CUDA graphs can pad these buffers. Only the first num_reqs
921
+ # accepted-token counts belong to actual speculative requests.
922
+ num_accepted_tokens = num_accepted_tokens[:num_reqs]
923
+ q_starts = spec_query_start_loc[: num_reqs + 1].to(torch.int64)
924
+ # Keep the number of real speculative tokens on the device.
925
+ total_real_tokens = q_starts[num_reqs]
926
+
927
+ state_indices = spec_state_indices_tensor.to(
928
+ device=conv_state.device, dtype=torch.int64
929
+ )
930
+ valid_state = state_indices != NULL_BLOCK_ID
931
+ state_indices = torch.where(
932
+ valid_state, state_indices, torch.zeros_like(state_indices)
933
+ )
934
+ positions = torch.arange(
935
+ x_spec.size(0), device=x_spec.device, dtype=torch.int64
936
+ )
937
+ # Route graph-padded token rows to the discarded dummy request so that
938
+ # they cannot overwrite real packed data.
939
+ req_indices = torch.searchsorted(q_starts[1:], positions, right=True)
940
+ valid_tokens = (positions < total_real_tokens) & (req_indices < num_reqs)
941
+ clamped_req_indices = req_indices.clamp_max(max(num_reqs - 1, 0))
942
+ col_indices = (positions - q_starts[clamped_req_indices]).clamp_(0, max_len - 1)
943
+ pack_req_indices = torch.where(
944
+ valid_tokens,
945
+ clamped_req_indices,
946
+ torch.full_like(req_indices, num_reqs),
947
+ )
948
+ pack_col_indices = torch.where(
949
+ valid_tokens, col_indices, torch.zeros_like(col_indices)
950
+ )
951
+
952
+ # The last request row is the dummy sink for graph padding.
953
+ packed = x_spec.new_zeros((num_reqs + 1, max_len, hidden_size))
954
+ packed[pack_req_indices, pack_col_indices] = x_spec
955
+ packed = packed.transpose(1, 2).contiguous()
956
+
957
+ if self.conv_state_len > 0:
958
+ cached_state = conv_state.index_select(0, state_indices)
959
+ rollback_offsets = num_accepted_tokens.to(
960
+ device=conv_state.device, dtype=torch.int64
961
+ ).sub(1)
962
+ rollback_offsets = torch.where(
963
+ valid_state,
964
+ rollback_offsets.clamp_(0, max_len - 1),
965
+ torch.zeros_like(rollback_offsets),
966
+ )
967
+ state_offsets = torch.arange(
968
+ self.conv_state_len, device=conv_state.device, dtype=torch.int64
969
+ ).view(1, 1, self.conv_state_len)
970
+ rollback_indices = rollback_offsets.view(-1, 1, 1) + state_offsets
971
+ state = cached_state.gather(
972
+ 2, rollback_indices.expand(-1, hidden_size, -1)
973
+ ).to(x_spec.dtype)
974
+ state = torch.where(
975
+ valid_state.view(num_reqs, 1, 1),
976
+ state,
977
+ torch.zeros_like(state),
978
+ )
979
+ # Append a zeroed dummy-row state to match the [num_reqs + 1] pack.
980
+ dummy_state = state.new_zeros((1, hidden_size, self.conv_state_len))
981
+ state_full = torch.cat((state, dummy_state), dim=0)
982
+ history = torch.cat((state_full, packed), dim=-1)
983
+ else:
984
+ history = packed
985
+
986
+ conv_output = F.conv1d(
987
+ history,
988
+ conv_weights.unsqueeze(1).contiguous(),
989
+ groups=history.size(1),
990
+ dilation=self.short_conv_dilation,
991
+ )
992
+ conv_output = F.silu(conv_output).transpose(1, 2).contiguous()
993
+
994
+ output = conv_output[pack_req_indices, pack_col_indices]
995
+ output = output * valid_tokens.view(-1, 1).to(output.dtype)
996
+
997
+ # Keep all current candidate inputs in the extended state. On the next
998
+ # target forward, ``num_accepted_tokens - 1`` selects the rollback
999
+ # window before processing the newly scheduled tokens.
1000
+ if self.conv_state_len > 0:
1001
+ state_capacity = self.conv_state_len + max_len - 1
1002
+ if conv_state.size(-1) < state_capacity:
1003
+ raise RuntimeError(
1004
+ "PLE short-conv cache cannot retain speculative tokens: "
1005
+ f"got {conv_state.size(-1)}, need {state_capacity}."
1006
+ )
1007
+ candidate_state = history[:num_reqs, :, 1 : state_capacity + 1]
1008
+ query_lengths = q_starts[1:] - q_starts[:-1]
1009
+ state_positions = torch.arange(
1010
+ state_capacity, device=history.device, dtype=torch.int64
1011
+ ).view(1, 1, state_capacity)
1012
+ update_lengths = (self.conv_state_len + query_lengths - 1).view(
1013
+ num_reqs, 1, 1
1014
+ )
1015
+ update_mask = valid_state.view(num_reqs, 1, 1) & (
1016
+ state_positions < update_lengths
1017
+ )
1018
+ existing_state = cached_state[..., :state_capacity]
1019
+ next_state = torch.where(
1020
+ update_mask,
1021
+ candidate_state.to(conv_state.dtype),
1022
+ existing_state,
1023
+ )
1024
+ cached_state[..., :state_capacity] = next_state
1025
+ conv_state.index_copy_(0, state_indices, cached_state)
1026
+
1027
+ return output
1028
+
1029
+ def _short_conv_dilated_dispatch(
1030
+ self,
1031
+ inputs: torch.Tensor,
1032
+ metadata: PleShortConvAttentionMetadata,
1033
+ conv_state: torch.Tensor,
1034
+ conv_weights: torch.Tensor,
1035
+ ) -> torch.Tensor:
1036
+ num_prefills = metadata.num_prefills
1037
+ num_decodes = metadata.num_decodes
1038
+ num_decode_tokens = metadata.num_decode_tokens
1039
+ num_prefill_tokens = metadata.num_prefill_tokens
1040
+ has_prefill = num_prefills > 0
1041
+ has_decode = num_decodes > 0
1042
+ has_spec = metadata.spec_sequence_masks is not None
1043
+ x = inputs[: metadata.num_actual_tokens]
1044
+
1045
+ # Split spec / non-spec tokens.
1046
+ if has_spec:
1047
+ if has_prefill or has_decode:
1048
+ assert metadata.spec_token_indx is not None
1049
+ assert metadata.non_spec_token_indx is not None
1050
+ x_spec = x.index_select(0, metadata.spec_token_indx.long())
1051
+ x_non_spec = x.index_select(0, metadata.non_spec_token_indx.long())
1052
+ else:
1053
+ x_spec = x
1054
+ x_non_spec = None
1055
+ else:
1056
+ x_spec = None
1057
+ x_non_spec = x
1058
+
1059
+ spec_output = None
1060
+ # 1. Run the multi-query speculative-decode part.
1061
+ if has_spec:
1062
+ assert metadata.spec_state_indices_tensor is not None
1063
+ assert metadata.spec_query_start_loc is not None
1064
+ assert metadata.num_accepted_tokens is not None
1065
+ spec_output = self._short_conv_dilated_spec_batched(
1066
+ x_spec=x_spec,
1067
+ conv_state=conv_state,
1068
+ conv_weights=conv_weights,
1069
+ spec_state_indices_tensor=metadata.spec_state_indices_tensor[
1070
+ : metadata.num_spec_decodes
1071
+ ],
1072
+ spec_query_start_loc=metadata.spec_query_start_loc,
1073
+ num_accepted_tokens=metadata.num_accepted_tokens,
1074
+ spec_query_len=metadata.spec_query_len,
1075
+ )
1076
+
1077
+ # 2. Run regular decode and prefill requests.
1078
+ conv_out_non_spec = None
1079
+ state_indices_tensor = metadata.state_indices_tensor
1080
+ if x_non_spec is not None:
1081
+ assert state_indices_tensor is not None
1082
+ if has_prefill:
1083
+ state_indices_tensor_d, state_indices_tensor_p = torch.split(
1084
+ state_indices_tensor,
1085
+ [num_decodes, num_prefills],
1086
+ dim=0,
1087
+ )
1088
+ x_d, x_p = torch.split(
1089
+ x_non_spec,
1090
+ [num_decode_tokens, num_prefill_tokens],
1091
+ dim=0,
1092
+ )
1093
+ non_spec_parts: list[torch.Tensor] = []
1094
+ if has_decode:
1095
+ non_spec_parts.append(
1096
+ self._short_conv_dilated_decode_batched(
1097
+ x_d=x_d,
1098
+ conv_state=conv_state,
1099
+ conv_weights=conv_weights,
1100
+ state_indices_tensor_d=state_indices_tensor_d,
1101
+ has_initial_states_d=metadata.has_initial_states_d,
1102
+ )
1103
+ )
1104
+ non_spec_parts.append(
1105
+ self._short_conv_dilated_prefill_batched(
1106
+ x_p=x_p,
1107
+ metadata=metadata,
1108
+ conv_state=conv_state,
1109
+ conv_weights=conv_weights,
1110
+ state_indices_tensor_p=state_indices_tensor_p,
1111
+ num_prefills=num_prefills,
1112
+ num_decode_tokens=num_decode_tokens,
1113
+ num_prefill_tokens=num_prefill_tokens,
1114
+ )
1115
+ )
1116
+ conv_out_non_spec = torch.vstack(non_spec_parts)
1117
+ else:
1118
+ conv_out_non_spec = self._short_conv_dilated_decode_batched(
1119
+ x_d=x_non_spec,
1120
+ conv_state=conv_state,
1121
+ conv_weights=conv_weights,
1122
+ state_indices_tensor_d=state_indices_tensor[: x_non_spec.size(0)],
1123
+ has_initial_states_d=metadata.has_initial_states_d,
1124
+ )
1125
+
1126
+ # 3. Merge both parts back into the original token order.
1127
+ if has_spec and conv_out_non_spec is not None:
1128
+ assert metadata.spec_token_indx is not None
1129
+ assert metadata.non_spec_token_indx is not None
1130
+ assert spec_output is not None
1131
+ output = x.new_empty((metadata.num_actual_tokens, x.size(-1)))
1132
+ output.index_copy_(0, metadata.spec_token_indx, spec_output)
1133
+ output.index_copy_(0, metadata.non_spec_token_indx, conv_out_non_spec)
1134
+ return output
1135
+ elif has_spec:
1136
+ assert spec_output is not None
1137
+ return spec_output
1138
+ if conv_out_non_spec is None:
1139
+ return x
1140
+ return conv_out_non_spec
1141
+
1142
+ def _short_conv(self, inputs: torch.Tensor) -> torch.Tensor:
1143
+ forward_context = get_forward_context()
1144
+ attn_metadata = forward_context.attn_metadata
1145
+ if attn_metadata is None:
1146
+ return self._short_conv_fallback(inputs)
1147
+
1148
+ if not isinstance(attn_metadata, dict):
1149
+ raise RuntimeError(
1150
+ "PLE short-conv expects per-layer attention metadata dict "
1151
+ f"during inference, got {type(attn_metadata).__name__}."
1152
+ )
1153
+
1154
+ layer_attn_metadata = attn_metadata.get(self.prefix)
1155
+ if layer_attn_metadata is None:
1156
+ raise RuntimeError(
1157
+ f"Missing short-conv metadata for layer '{self.prefix}'. "
1158
+ "This would bypass conv-state updates and is not allowed."
1159
+ )
1160
+ if not isinstance(layer_attn_metadata, PleShortConvAttentionMetadata):
1161
+ raise TypeError(
1162
+ "Expected PleShortConvAttentionMetadata for layer "
1163
+ f"'{self.prefix}', got "
1164
+ f"{type(layer_attn_metadata).__name__}."
1165
+ )
1166
+
1167
+ conv_state = self.kv_cache[0]
1168
+ if not is_conv_state_dim_first():
1169
+ conv_state = conv_state.transpose(-1, -2)
1170
+ conv_weights = self.conv1d.weight.squeeze(1)
1171
+
1172
+ state_capacity = self.conv_state_len + self.num_spec_tokens
1173
+ if state_capacity > 0:
1174
+ if conv_state.size(-1) < state_capacity:
1175
+ raise RuntimeError(
1176
+ "PLE short-conv cache is smaller than expected for "
1177
+ f"dilated convolution: got {conv_state.size(-1)}, "
1178
+ f"expect at least {state_capacity}."
1179
+ )
1180
+ conv_state = conv_state[..., -state_capacity:]
1181
+ return self._short_conv_dilated_dispatch(
1182
+ inputs,
1183
+ layer_attn_metadata,
1184
+ conv_state,
1185
+ conv_weights.to(dtype=inputs.dtype),
1186
+ )
1187
+
1188
+ def forward(
1189
+ self,
1190
+ hidden_states: torch.Tensor,
1191
+ input_ids: torch.Tensor,
1192
+ query_start_loc: torch.Tensor,
1193
+ ngram_context: torch.Tensor,
1194
+ ) -> torch.Tensor:
1195
+ input_ids = input_ids.reshape(-1)
1196
+ if input_ids.shape[0] != hidden_states.shape[0]:
1197
+ raise ValueError(
1198
+ "PLE expects input_ids and hidden_states to have the same "
1199
+ f"token length, got {input_ids.shape[0]} and "
1200
+ f"{hidden_states.shape[0]}"
1201
+ )
1202
+ embeddings = self.ple_embedding(
1203
+ hidden_states,
1204
+ input_ids,
1205
+ query_start_loc,
1206
+ ngram_context,
1207
+ )
1208
+ embeddings = self._dequantize_embeddings(embeddings, hidden_states.dtype)
1209
+ key, _ = self.key_proj(embeddings)
1210
+ value, _ = self.value_proj(embeddings)
1211
+ token_count = hidden_states.shape[0]
1212
+ key = key.reshape(token_count, self.hc_count, self.hidden_size)
1213
+ query = hidden_states.reshape(token_count, self.hc_count, self.hidden_size)
1214
+ key = self._apply_norm(self.norm_key, key)
1215
+ query = self._apply_norm(self.norm_query, query)
1216
+ gate = (key * query).sum(dim=-1, keepdim=True) / math.sqrt(self.hidden_size)
1217
+ gate = torch.sigmoid(gate.sign() * gate.abs().clamp_min(1e-6).sqrt())
1218
+ gated_value = gate * value.unsqueeze(-2)
1219
+ normalized = self._apply_norm(self.norm_conv, gated_value).flatten(-2)
1220
+ conv_output = torch.zeros_like(normalized)
1221
+ torch.ops.vllm.qwen3_8_flash_next_ple_short_conv(
1222
+ normalized,
1223
+ conv_output,
1224
+ self.prefix,
1225
+ )
1226
+ return gated_value.flatten(-2) + conv_output
1227
+
1228
+
1229
+ def qwen3_8_flash_next_ple_short_conv(
1230
+ inputs: torch.Tensor,
1231
+ output: torch.Tensor,
1232
+ layer_name: str,
1233
+ ) -> None:
1234
+ layer = get_forward_context().no_compile_layers[layer_name]
1235
+ result = layer._short_conv(inputs)
1236
+ output[: result.shape[0]].copy_(result)
1237
+
1238
+
1239
+ def qwen3_8_flash_next_ple_short_conv_fake(
1240
+ inputs: torch.Tensor,
1241
+ output: torch.Tensor,
1242
+ layer_name: str,
1243
+ ) -> None:
1244
+ return
1245
+
1246
+
1247
+ direct_register_custom_op(
1248
+ op_name="qwen3_8_flash_next_ple_short_conv",
1249
+ op_func=qwen3_8_flash_next_ple_short_conv,
1250
+ mutates_args=["output"],
1251
+ fake_impl=qwen3_8_flash_next_ple_short_conv_fake,
1252
+ )
1253
+
1254
+
1255
+ __all__ = [
1256
+ "Qwen3_8FlashNextNGramEmbedding",
1257
+ "Qwen3_8FlashNextPLEGroupedNorm",
1258
+ "Qwen3_8FlashNextPLELayer",
1259
+ ]