diff --git a/.buildkite/test_areas/kernels.yaml b/.buildkite/test_areas/kernels.yaml index 35cd72bff356..640a7f5b3990 100644 --- a/.buildkite/test_areas/kernels.yaml +++ b/.buildkite/test_areas/kernels.yaml @@ -147,7 +147,7 @@ steps: commands: - pytest -v -s kernels/attention/test_triton_unified_attention_diffkv.py -- label: ":nvidia: (H100) FlashMLA Kernels" +- label: ":nvidia: (H100) MLA Kernel Test" key: kernels-flashmla-test-h100 timeout_in_minutes: 25 device: h100 @@ -157,12 +157,16 @@ steps: - vllm/v1/attention/ops/flashmla.py - vllm/v1/attention/backends/mla/flashmla.py - vllm/v1/attention/backends/mla/flashmla_sparse.py + - vllm/v1/attention/backends/mla/flashinfer_mla_sparse_sm90.py + - vllm/utils/flashinfer.py + - tests/kernels/attention/test_flashinfer_mla_decode.py - tests/kernels/attention/test_flashmla.py - tests/kernels/attention/test_flashmla_sparse.py - tests/kernels/attention/test_mla_cross_layer_kernel_equivalence.py commands: - pytest -v -s kernels/attention/test_flashmla.py - pytest -v -s kernels/attention/test_flashmla_sparse.py + - pytest -v -s kernels/attention/test_flashinfer_mla_decode.py - pytest -v -s kernels/attention/test_mla_cross_layer_kernel_equivalence.py - label: ":nvidia: (L4) Quantization Kernels Shard %N" @@ -318,6 +322,7 @@ steps: - csrc/libtorch_stable/fused_minimax_m3_qknorm_rope_kv_insert_kernel.cu - csrc/libtorch_stable/ops.h - csrc/libtorch_stable/torch_bindings.cpp + - tests/kernels/attention/test_flashinfer_mla_decode.py - tests/kernels/attention/test_minimax_m3_msa_cutlass_sparse_decode.py - tests/kernels/test_fused_minimax_m3_qknorm_rope_kv_insert.py - tests/kernels/mamba/test_gdn_prefill_cutedsl.py diff --git a/csrc/libtorch_stable/cache_kernels.cu b/csrc/libtorch_stable/cache_kernels.cu index 1807a81f4c0a..79aa63a5b83a 100644 --- a/csrc/libtorch_stable/cache_kernels.cu +++ b/csrc/libtorch_stable/cache_kernels.cu @@ -1268,6 +1268,9 @@ __global__ void gather_and_maybe_dequant_cache_page( #define CALL_GATHER_CACHE_576(SCALAR_T, CACHE_T, KV_DTYPE) \ CALL_GATHER_CACHE(SCALAR_T, CACHE_T, KV_DTYPE, 576) +#define CALL_GATHER_CACHE_512(SCALAR_T, CACHE_T, KV_DTYPE) \ + CALL_GATHER_CACHE(SCALAR_T, CACHE_T, KV_DTYPE, 512) + #define CALL_GATHER_CACHE_320(SCALAR_T, CACHE_T, KV_DTYPE) \ CALL_GATHER_CACHE(SCALAR_T, CACHE_T, KV_DTYPE, 320) @@ -1305,10 +1308,9 @@ void gather_and_maybe_dequant_cache( seq_starts.value().scalar_type() == torch::headeronly::ScalarType::Int, "seq_starts must be int32"); } - STD_TORCH_CHECK( - head_dim == 320 || head_dim == 576, - "gather_and_maybe_dequant_cache only support the head_dim to 320 or 576 " - "for better performance") + STD_TORCH_CHECK(head_dim == 320 || head_dim == 512 || head_dim == 576, + "gather_and_maybe_dequant_cache only support the head_dim to " + "320 or 512 or 576 for better performance") STD_TORCH_CHECK(src_cache.device() == dst.device(), "src_cache and dst must be on the same device"); @@ -1344,6 +1346,9 @@ void gather_and_maybe_dequant_cache( if (head_dim == 576) { DISPATCH_BY_KV_CACHE_DTYPE(dst.scalar_type(), kv_cache_dtype, CALL_GATHER_CACHE_576); + } else if (head_dim == 512) { + DISPATCH_BY_KV_CACHE_DTYPE(dst.scalar_type(), kv_cache_dtype, + CALL_GATHER_CACHE_512); } else { DISPATCH_BY_KV_CACHE_DTYPE(dst.scalar_type(), kv_cache_dtype, CALL_GATHER_CACHE_320); diff --git a/tests/kernels/attention/test_flashinfer_mla_decode.py b/tests/kernels/attention/test_flashinfer_mla_decode.py index d1bd55eebedd..a49ae3ae1ad3 100644 --- a/tests/kernels/attention/test_flashinfer_mla_decode.py +++ b/tests/kernels/attention/test_flashinfer_mla_decode.py @@ -9,13 +9,23 @@ FLASHINFER_WORKSPACE_BUFFER_SIZE = 128 * 1024 * 1024 -if not current_platform.has_device_capability(100): +if not current_platform.is_cuda() or not current_platform.has_device_capability(90): pytest.skip( - reason="FlashInfer MLA Requires compute capability of 10 or above.", + reason="FlashInfer MLA requires CUDA compute capability 9.0 or above.", allow_module_level=True, ) else: from flashinfer.decode import trtllm_batch_decode_with_kv_cache_mla + from flashinfer.mla import BatchMLAPagedAttentionWrapper + +requires_sm90 = pytest.mark.skipif( + not current_platform.is_device_capability_family(90), + reason="This test requires an SM90 GPU.", +) +requires_sm10x = pytest.mark.skipif( + not current_platform.is_device_capability_family(100), + reason="This test requires an SM10x GPU.", +) # Deepseek R1 MLA config. NUM_HEADS = 128 @@ -82,6 +92,7 @@ def ref_mla( @pytest.mark.parametrize("dtype", [torch.bfloat16]) @pytest.mark.parametrize("bs", [1, 2, 4, 16]) @pytest.mark.parametrize("block_size", [32, 64]) +@requires_sm10x def test_flashinfer_mla_decode(dtype: torch.dtype, bs: int, block_size: int): torch.set_default_device("cuda") torch.manual_seed(42) @@ -119,6 +130,161 @@ def test_flashinfer_mla_decode(dtype: torch.dtype, bs: int, block_size: int): torch.testing.assert_close(out_ans, out_ref, atol=1e-2, rtol=1e-2) +@requires_sm10x +def test_flashinfer_trtllm_sparse_mla_decode_without_rope(): + """The native sparse MLA path supports a zero-width rotary tail.""" + torch.set_default_device("cuda") + torch.manual_seed(42) + + batch_size = 2 + block_size = 64 + num_blocks = 4 + sparse_topk = 128 + valid_lens = torch.tensor([17, 73], dtype=torch.int32) + + query = torch.randn( + batch_size, + 1, + NUM_HEADS, + KV_LORA_RANK, + dtype=torch.bfloat16, + ) + kv_cache = torch.randn( + num_blocks, + block_size, + KV_LORA_RANK, + dtype=torch.bfloat16, + ) + + num_slots = num_blocks * block_size + slot_tables = torch.stack( + [torch.randperm(num_slots)[:sparse_topk] for _ in range(batch_size)] + ).to(torch.int32) + for row, valid_len in zip(slot_tables, valid_lens.tolist()): + row[valid_len:] = -1 + + workspace_buffer = torch.empty( + FLASHINFER_WORKSPACE_BUFFER_SIZE, + dtype=torch.int8, + ) + out = trtllm_batch_decode_with_kv_cache_mla( + query=query, + kv_cache=kv_cache.unsqueeze(1), + workspace_buffer=workspace_buffer, + qk_nope_head_dim=QK_NOPE_HEAD_DIM, + kv_lora_rank=KV_LORA_RANK, + qk_rope_head_dim=0, + block_tables=slot_tables.unsqueeze(1), + seq_lens=valid_lens, + max_seq_len=sparse_topk, + sparse_mla_top_k=sparse_topk, + sparse_mla_top_k_lens=valid_lens, + bmm1_scale=QK_NOPE_HEAD_DIM**-0.5, + bmm2_scale=1.0, + ).squeeze(1) + + flat_cache = kv_cache.view(num_slots, KV_LORA_RANK).float() + refs = [] + for batch_idx, valid_len in enumerate(valid_lens.tolist()): + selected_kv = flat_cache[slot_tables[batch_idx, :valid_len].long()] + scores = torch.einsum("hd,kd->hk", query[batch_idx, 0].float(), selected_kv) + probs = torch.softmax(scores * QK_NOPE_HEAD_DIM**-0.5, dim=-1) + refs.append(torch.einsum("hk,kd->hd", probs, selected_kv)) + ref = torch.stack(refs).to(torch.bfloat16) + + torch.testing.assert_close(out, ref, atol=2e-2, rtol=2e-2) + + +@requires_sm90 +def test_flashinfer_sm90_fp8_mla_decode_without_rope(): + """Hopper FA3 supports BF16 queries over an FP8 cache without KPE.""" + torch.manual_seed(42) + device = torch.device("cuda") + batch_size = 2 + num_heads = 16 + page_size = 16 + num_pages = 6 + + q_nope = torch.randn( + batch_size, + num_heads, + KV_LORA_RANK, + dtype=torch.bfloat16, + device=device, + ) + q_pe = torch.empty( + batch_size, + num_heads, + 0, + dtype=torch.bfloat16, + device=device, + ) + + ckv = torch.randn( + num_pages, + page_size, + KV_LORA_RANK, + device=device, + ) + fp8_max = torch.finfo(torch.float8_e4m3fn).max + ckv_scale = ckv.abs().max().item() / fp8_max + ckv_fp8 = (ckv / ckv_scale).clamp(-fp8_max, fp8_max).to(torch.float8_e4m3fn) + scale_bf16 = torch.tensor(ckv_scale, dtype=torch.bfloat16, device=device) + ckv_ref = ckv_fp8.to(torch.bfloat16) * scale_bf16 + kpe_fp8 = torch.empty( + num_pages, + page_size, + 0, + dtype=torch.float8_e4m3fn, + device=device, + ) + kpe_ref = torch.empty( + num_pages, + page_size, + 0, + dtype=torch.bfloat16, + device=device, + ) + + qo_indptr = torch.tensor([0, 1, 2], dtype=torch.int32, device=device) + kv_indptr = torch.tensor([0, 3, 5], dtype=torch.int32, device=device) + kv_indices = torch.tensor([4, 1, 3, 0, 5], dtype=torch.int32, device=device) + kv_lens = torch.tensor([45, 29], dtype=torch.int32, device=device) + sm_scale = QK_NOPE_HEAD_DIM**-0.5 + + def run( + ckv_cache: torch.Tensor, + kpe_cache: torch.Tensor, + **kwargs, + ) -> torch.Tensor: + workspace = torch.empty( + FLASHINFER_WORKSPACE_BUFFER_SIZE, + dtype=torch.uint8, + device=device, + ) + wrapper = BatchMLAPagedAttentionWrapper(workspace, backend="fa3") + wrapper.plan( + qo_indptr, + kv_indptr, + kv_indices, + kv_lens, + num_heads, + KV_LORA_RANK, + 0, + page_size, + False, + sm_scale, + q_data_type=torch.bfloat16, + kv_data_type=ckv_cache.dtype, + ) + return wrapper.run(q_nope, q_pe, ckv_cache, kpe_cache, **kwargs) + + out_ref = run(ckv_ref, kpe_ref) + out = run(ckv_fp8, kpe_fp8, ckv_scale=ckv_scale, kpe_scale=1.0) + torch.testing.assert_close(out, out_ref, atol=2e-2, rtol=2e-2) + + +@requires_sm10x def test_flashinfer_mla_decode_workspace_supports_autotune(): """vLLM's FlashInfer MLA decode workspace must be int8 for autotuning. diff --git a/tests/kernels/attention/test_rocm_aiter_mla_sparse_metadata_sync.py b/tests/kernels/attention/test_rocm_aiter_mla_sparse_metadata_sync.py index aa603626414c..d2f5251c62de 100644 --- a/tests/kernels/attention/test_rocm_aiter_mla_sparse_metadata_sync.py +++ b/tests/kernels/attention/test_rocm_aiter_mla_sparse_metadata_sync.py @@ -41,6 +41,8 @@ def _make_builder(): builder.device = torch.device("cpu") builder.kv_cache_spec = SimpleNamespace(block_size=1) builder.model_dtype = torch.bfloat16 + builder.kv_cache_dtype = "fp8" + builder.mla_dims = SimpleNamespace(kv_lora_rank=512, qk_rope_head_dim=64) builder.topk_tokens = topk_tokens builder.req_id_per_token_buffer = torch.zeros( max_num_batched_tokens, dtype=torch.int32, device="cpu" diff --git a/tests/kernels/mamba/test_gdn_prefill_flashinfer.py b/tests/kernels/mamba/test_gdn_prefill_flashinfer.py new file mode 100644 index 000000000000..157e589ecaae --- /dev/null +++ b/tests/kernels/mamba/test_gdn_prefill_flashinfer.py @@ -0,0 +1,53 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project + +import pytest +import torch + +from vllm.platforms import current_platform + +if current_platform.is_rocm(): + pytest.skip( + reason="FlashInfer GDN prefill is not supported on ROCm.", + allow_module_level=True, + ) + +import flashinfer.gdn_prefill # noqa: E402 + +from vllm.model_executor.layers.mamba.gdn.qwen_gdn_linear_attn import ( + fi_chunk_gated_delta_rule, +) # noqa: E402 + + +def test_flashinfer_gdn_prefill_uses_int64_cu_seqlens(monkeypatch): + captured_cu_seqlens = None + + def fake_chunk_gated_delta_rule(**kwargs): + nonlocal captured_cu_seqlens + captured_cu_seqlens = kwargs["cu_seqlens"] + return kwargs["q"] + + monkeypatch.setattr( + flashinfer.gdn_prefill, + "chunk_gated_delta_rule", + fake_chunk_gated_delta_rule, + ) + q = torch.zeros(1, 2, 1, 2) + cu_seqlens = torch.tensor([0, 2], dtype=torch.int32) + + output, final_state = fi_chunk_gated_delta_rule( + q=q, + k=q, + v=q, + g=torch.zeros(1, 2, 1), + beta=torch.zeros(1, 2, 1), + initial_state=torch.zeros(1, 1, 2, 2), + output_final_state=False, + cu_seqlens=cu_seqlens, + use_qk_l2norm_in_kernel=False, + ) + + assert captured_cu_seqlens is not None + assert captured_cu_seqlens.dtype == torch.int64 + assert output.shape == q.shape + assert final_state is None diff --git a/tests/kernels/test_kpool_decode_update_batched.py b/tests/kernels/test_kpool_decode_update_batched.py new file mode 100644 index 000000000000..afc0c1f836b7 --- /dev/null +++ b/tests/kernels/test_kpool_decode_update_batched.py @@ -0,0 +1,549 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +"""Unit test for the batched kpool decode-update kernel. + +Validates ``kpool_decode_update_and_maybe_write_cache_batched`` against an +independent pure-torch reference that replicates the per-request, in-position +order semantics: stash each token into a paged tail ring; on pool completion +(``pos % pool_size == pool_size-1``) softmax(gate+ape)-weighted sum + Hadamard-128 ++ per-vector fp8 absmax quant + write to the indexer K cache. Covers +no-completion, completion-at-end, completion-mid-batch, non-uniform padding, +plain decode, plus a randomized fuzz pass. + +The kernel iterates each request's ``next_n`` tokens in position order inside +one program (grid = num_requests) to preserve the pool-completion +read-after-stash dependency; the reference mirrors that ordering. + +``test_decode_writer_matches_prefill_writer`` is deliberately NOT +reference-based: it checks the decode writer against the *prefill* writer +(``kpool_compress_and_write_cache``), the invariant that actually matters in +production. A hand-written reference can drift to match a buggy kernel -- that +is exactly how the stash-gating bug (intra-pool tokens never entering the tail +ring, because the stash was gated on the pool-granular ``slot_mapping``) stayed +green here. +""" + +import math + +import pytest +import torch + +from vllm.platforms import current_platform + +if current_platform.is_rocm(): + from vllm.models.glm5next.amd.ops.kpool_compress import ( + kpool_compress_and_write_cache, + kpool_decode_update_and_maybe_write_cache_batched, + kpool_seed_tail_cache, + ) +else: + from vllm.models.glm5next.nvidia.ops.kpool_compress import ( + kpool_compress_and_write_cache, + kpool_decode_update_and_maybe_write_cache_batched, + kpool_seed_tail_cache, + ) + +HEAD_DIM = 128 +POOL_SIZE = 16 +PAGE_SIZE = 64 +NUM_BLOCKS = 32 +ROUND_SCALE = True +FP8_DTYPE = current_platform.fp8_dtype() +FP8_MAX = torch.finfo(FP8_DTYPE).max + + +def _make_caches(): + kv = torch.zeros( + NUM_BLOCKS, PAGE_SIZE, HEAD_DIM + 4, dtype=torch.uint8, device="cuda" + ) + tail = torch.zeros( + NUM_BLOCKS, 2, POOL_SIZE, HEAD_DIM, dtype=torch.bfloat16, device="cuda" + ) + return kv, tail + + +def _tail_slot_for(blocks, pos): + """tail_slot = block*POOL + pos%POOL; each request owns a distinct tail block.""" + blk = torch.tensor(blocks, device=pos.device, dtype=torch.int32).unsqueeze(1) + return (blk * POOL_SIZE + pos % POOL_SIZE).to(torch.int32) + + +def _seed_prior(tail, blocks, n_prior, seed=42): + if n_prior <= 0: + return + g = torch.Generator(device=tail.device).manual_seed(seed) + prior_k = torch.randn( + len(blocks), + n_prior, + HEAD_DIM, + dtype=torch.bfloat16, + device=tail.device, + generator=g, + ) + prior_s = torch.randn( + len(blocks), + n_prior, + HEAD_DIM, + dtype=torch.bfloat16, + device=tail.device, + generator=g, + ) + for i, blk in enumerate(blocks): + tail[blk, 0, :n_prior, :] = prior_k[i] + tail[blk, 1, :n_prior, :] = prior_s[i] + + +def _hadamard128_torch(x: torch.Tensor) -> torch.Tensor: + """Reference Hadamard-128 on the last dim (must be 128).""" + n = x.shape[-1] + assert n == 128 + h = torch.tensor([[1.0, 1.0], [1.0, -1.0]], dtype=torch.float32, device=x.device) + while h.shape[0] < n: + h = torch.cat([torch.cat([h, h], dim=1), torch.cat([h, -h], dim=1)], dim=0) + h = h / math.sqrt(n) + return x @ h + + +def _torch_reference( + kv: torch.Tensor, + tail: torch.Tensor, + tail_slot: torch.Tensor, + key: torch.Tensor, + score: torch.Tensor, + ape: torch.Tensor, + slot_map: torch.Tensor, + pos: torch.Tensor, +) -> tuple[torch.Tensor, torch.Tensor]: + """Independent reference for the batched decode-update kernel. + + For each request, iterate its next_n tokens in order. On each token: + - if pos%POOL == POOL-1 and pos_valid: compress the pool (slots + [pool_start..pool_start+POOL-1], current token via is_current) and write + fp8 K + fp32 scale to kv_cache at cache_loc. + - always stash the current token's K/score into tail[block, pos%POOL]. + """ + kv = kv.clone() + tail = tail.clone() + B, next_n = pos.shape + # The indexer K cache is [num_blocks, PAGE_SIZE, HEAD_DIM+4] uint8 but the + # kernels interpret each page as [HEAD_DIM*PAGE_SIZE bytes of K (token-major) + # | 4*PAGE_SIZE bytes of fp32 scale (token-major)]. Operate on a flat byte + # view so the reference writes K and scale at the exact offsets the kernel + # uses (page_base + tok*HEAD_DIM for K; page_base + HEAD_DIM*PAGE_SIZE + + # tok*4 for scale). + page_bytes = PAGE_SIZE * (HEAD_DIM + 4) + k_region = HEAD_DIM * PAGE_SIZE + tail_slot_cpu = tail_slot.cpu().tolist() + slot_map_cpu = slot_map.cpu().tolist() + pos_cpu = pos.cpu().tolist() + key_cpu = key.float().cpu() + score_cpu = score.float().cpu() + ape_cpu = ape.cpu() + tail_cpu = tail.float().cpu() + kv_flat = kv.view(torch.uint8).reshape(-1).cpu() + + for b in range(B): + for t in range(next_n): + cache_loc = slot_map_cpu[b][t] + p = pos_cpu[b][t] + pos_valid = cache_loc >= 0 and p >= 0 + safe_pos = max(p, 0) + slot = safe_pos % POOL_SIZE + phys_slot = safe_pos % POOL_SIZE + # Per-token block derivation (a leading invalid sentinel must not + # poison the base for the rest of the request); clamped like the + # kernel so an invalid entry can't form a negative base. + block = max(tail_slot_cpu[b][t], 0) // POOL_SIZE + + cur_key = key_cpu[b, t] + cur_score = score_cpu[b, t] + + if pos_valid and slot == POOL_SIZE - 1: + pool_logical_start = safe_pos - slot + pool_scores = [] + pool_ks = [] + for ps in range(POOL_SIZE): + is_current = ps == slot + phys = (pool_logical_start + ps) % POOL_SIZE + if is_current: + s = cur_score + k = cur_key + else: + s = tail_cpu[block, 1, phys] + k = tail_cpu[block, 0, phys] + s = s + ape_cpu[ps] + pool_scores.append(s) + pool_ks.append(k) + pool_scores = torch.stack(pool_scores) # [POOL, D] + pool_ks = torch.stack(pool_ks) # [POOL, D] + max_score = pool_scores.max(dim=0).values + prob = torch.exp(pool_scores - max_score) + denom = prob.sum(dim=0) + acc = (pool_ks * prob).sum(dim=0) + x = (acc / denom).to(torch.bfloat16).to(torch.float32) + x = _hadamard128_torch(x).to(torch.bfloat16).to(torch.float32) + absmax = torch.clamp(x.abs().max(), min=1e-4) + if ROUND_SCALE: + scale = torch.exp2(torch.ceil(torch.log2(absmax / FP8_MAX))) + else: + scale = absmax / FP8_MAX + quantized = torch.clamp(x / scale, -FP8_MAX, FP8_MAX).to(FP8_DTYPE) + # write K and scale at the separated-layout offsets + loc = cache_loc + loc_page_index = loc // PAGE_SIZE + loc_tok = loc % PAGE_SIZE + page_base = loc_page_index * page_bytes + if current_platform.is_rocm(): + dims = torch.arange(HEAD_DIM) + k_off = ( + page_base + + (loc_tok // 16) * 16 * HEAD_DIM + + (dims // 16) * 16 * 16 + + (loc_tok % 16) * 16 + + dims % 16 + ) + else: + k_off = page_base + loc_tok * HEAD_DIM + torch.arange(HEAD_DIM) + s_off = page_base + k_region + loc_tok * 4 + kv_flat[k_off] = quantized.view(torch.uint8) + kv_flat[s_off : s_off + 4] = scale.detach().reshape(1).view(torch.uint8) + + # stash -- gated on the TOKEN-granular tail slot, not on pos_valid. + # pos_valid keys off the pool-granular cache_loc, which is -1 for + # every token that is not the pool's last, so gating the stash on it + # would drop all intra-pool tokens. + if p >= 0 and tail_slot_cpu[b][t] >= 0: + tail_cpu[block, 0, phys_slot] = cur_key + tail_cpu[block, 1, phys_slot] = cur_score + + kv_out = kv_flat.view(NUM_BLOCKS, PAGE_SIZE, HEAD_DIM + 4).to(device="cuda") + return kv_out, tail_cpu.to(torch.bfloat16).to(device="cuda") + + +def _assert_eq(r_ref, r_kern): + kv_ref, tail_ref = r_ref + kv_kern, tail_kern = r_kern + assert torch.equal(kv_ref, kv_kern), ( + "kv_cache differs: max diff " + f"{(kv_ref.int() - kv_kern.int()).abs().max().item()}" + ) + assert torch.equal(tail_ref, tail_kern), ( + "tail_kv_cache differs: max diff " + f"{(tail_ref.float() - tail_kern.float()).abs().max().item()}" + ) + + +@pytest.mark.skipif(not current_platform.is_rocm(), reason="ROCm required") +def test_amd_prefill_writer_uses_preshuffled_cache_layout(): + from vllm.models.glm5next.amd.ops.kpool_compress import ( + kpool_compress_and_write_cache as amd_kpool_compress, + ) + + torch.manual_seed(0) + token_offset = 17 + kv = torch.zeros(1, PAGE_SIZE, HEAD_DIM + 4, dtype=torch.uint8, device="cuda") + key = torch.randn(1, POOL_SIZE, HEAD_DIM, dtype=torch.bfloat16, device="cuda") + score = torch.randn_like(key) + ape = torch.randn(POOL_SIZE, HEAD_DIM, dtype=torch.float32, device="cuda") + compressed_k, compressed_scale = amd_kpool_compress( + kv, + key, + score, + ape, + torch.tensor([token_offset], dtype=torch.int64, device="cuda"), + pool_size=POOL_SIZE, + head_dim=HEAD_DIM, + round_scale=ROUND_SCALE, + return_compressed=True, + ) + + dim = torch.arange(HEAD_DIM, device="cuda") + offsets = ( + (token_offset // 16) * 16 * HEAD_DIM + + (dim // 16) * 16 * 16 + + (token_offset % 16) * 16 + + dim % 16 + ) + flat = kv.view(torch.uint8).reshape(-1) + stored_k = flat[offsets].view(compressed_k.dtype) + scale_offset = PAGE_SIZE * HEAD_DIM + token_offset * 4 + stored_scale = flat[scale_offset : scale_offset + 4].view(torch.float32) + + assert torch.equal(stored_k, compressed_k[0]) + assert torch.equal(stored_scale, compressed_scale) + + +def _run_kernel(kv, tail, tail_slot, key, score, ape, slot_map, pos): + kv = kv.clone() + tail = tail.clone() + kpool_decode_update_and_maybe_write_cache_batched( + kv, + tail, + tail_slot, + key, + score, + ape, + slot_map, + pos, + POOL_SIZE, + HEAD_DIM, + round_scale=ROUND_SCALE, + ) + return kv, tail + + +@pytest.mark.parametrize("pool_size", [4, 16]) +def test_decode_writer_matches_prefill_writer(pool_size): + """Compare production decode and prefill writers for pool sizes 4 and 16.""" + n_pools, page, nblk = 8, 64, 4 + n_tok = n_pools * pool_size + dev = "cuda" + torch.manual_seed(0) + k = torch.randn(n_tok, HEAD_DIM, dtype=torch.bfloat16, device=dev) + score = torch.randn(n_tok, HEAD_DIM, dtype=torch.bfloat16, device=dev) + ape = torch.randn(pool_size, HEAD_DIM, dtype=torch.float32, device=dev) + + kv_prefill = torch.zeros(nblk, page, HEAD_DIM + 4, dtype=torch.uint8, device=dev) + kpool_compress_and_write_cache( + kv_prefill, + k.view(n_pools, pool_size, HEAD_DIM), + score.view(n_pools, pool_size, HEAD_DIM), + ape, + torch.arange(n_pools, dtype=torch.int64, device=dev), + pool_size=pool_size, + head_dim=HEAD_DIM, + round_scale=ROUND_SCALE, + ) + + # One request owning tail block 0, fed one token per decode step. + kv_decode = torch.zeros_like(kv_prefill) + tail = torch.zeros(nblk, 2, pool_size, HEAD_DIM, dtype=torch.bfloat16, device=dev) + for t in range(n_tok): + completes = t % pool_size == pool_size - 1 + kpool_decode_update_and_maybe_write_cache_batched( + kv_decode, + tail, + # token-granular: every token has a valid tail slot + torch.tensor([[t % pool_size]], dtype=torch.int32, device=dev), + k[t].view(1, 1, HEAD_DIM), + score[t].view(1, 1, HEAD_DIM), + ape, + # pool-granular: only the pool's last token carries a cache slot + torch.tensor( + [[t // pool_size if completes else -1]], dtype=torch.int32, device=dev + ), + torch.tensor([[t]], dtype=torch.int32, device=dev), + pool_size, + HEAD_DIM, + round_scale=ROUND_SCALE, + ) + + differing = [ + p + for p in range(n_pools) + if not torch.equal( + kv_prefill[p // page, p % page], kv_decode[p // page, p % page] + ) + ] + assert not differing, ( + f"decode-written pools differ from prefill-written pools: " + f"{len(differing)}/{n_pools} (pool_size={pool_size}, first={differing[:5]})" + ) + + +def test_leading_invalid_tail_slot(): + """A request whose FIRST token carries an invalid (-1) tail slot while a + later token is a real pool completion. + + The tail block must be derived per token, not from token 0: a leading + invalid sentinel would otherwise poison the base address for the whole + request (out-of-bounds tail reads on the completion). + """ + torch.manual_seed(0) + B, next_n, blocks = 2, 4, [3, 5] + # req 0: token 0 invalid (pos -1), tokens 1..3 valid, completion at pos 15 + # req 1: all valid, no completion + pos = torch.tensor( + [[-1, 13, 14, 15], [4, 5, 6, 7]], dtype=torch.int32, device="cuda" + ) + safe_pos = torch.where(pos >= 0, pos, 0) + tail_slot = _tail_slot_for(blocks, safe_pos) + # leading invalid entry carries the -1 sentinel, as the scatter path emits + tail_slot[0, 0] = -1 + slot_map = torch.full((B, next_n), -1, dtype=torch.int32, device="cuda") + slot_map[0, 3] = 15 # req 0 completes its pool on the last verify token + + key = torch.randn(B, next_n, HEAD_DIM, dtype=torch.bfloat16, device="cuda") + score = torch.randn(B, next_n, HEAD_DIM, dtype=torch.bfloat16, device="cuda") + ape = torch.randn(POOL_SIZE, HEAD_DIM, dtype=torch.float32, device="cuda") + + kv, tail = _make_caches() + _seed_prior(tail, blocks, 13) + r_ref = _torch_reference(kv, tail, tail_slot, key, score, ape, slot_map, pos) + r_kern = _run_kernel(kv, tail, tail_slot, key, score, ape, slot_map, pos) + _assert_eq(r_ref, r_kern) + + +@pytest.mark.skipif(not current_platform.is_rocm(), reason="ROCm required") +def test_amd_prefill_seed_honors_padded_tail_block_stride(): + """The tail shares a padded indexer allocation in production.""" + kpool = 4 + num_blocks = 6 + logical_block_elems = 2 * kpool * HEAD_DIM + padded_block_elems = logical_block_elems + 256 + sentinel = -123.0 + backing = torch.full( + (num_blocks * padded_block_elems,), + sentinel, + dtype=torch.bfloat16, + device="cuda", + ) + tail = torch.as_strided( + backing, + size=(num_blocks, 2, kpool, HEAD_DIM), + stride=(padded_block_elems, kpool * HEAD_DIM, HEAD_DIM, 1), + ) + + block = 3 + ring_slot = 2 + key = torch.arange(HEAD_DIM, dtype=torch.bfloat16, device="cuda").unsqueeze(0) + score = (key + 256).to(torch.bfloat16) + tail_slot = torch.tensor( + [block * kpool + ring_slot], dtype=torch.int32, device="cuda" + ) + + kpool_seed_tail_cache(tail, key, score, tail_slot, kpool, HEAD_DIM) + torch.accelerator.synchronize() + + assert torch.equal(tail[block, 0, ring_slot], key[0]) + assert torch.equal(tail[block, 1, ring_slot], score[0]) + + compact_offset = (block * 2 * kpool + ring_slot) * HEAD_DIM + assert torch.all(backing[compact_offset : compact_offset + HEAD_DIM] == sentinel) + + +@pytest.mark.parametrize( + "case_id", + [ + "no_completion", + "completion_at_end", + "completion_mid_batch", + "non_uniform_padding", + "plain_decode", + ], +) +def test_batched_matches_reference(case_id): + torch.manual_seed(0) + if case_id == "no_completion": + B, next_n, blocks = 3, 4, [0, 1, 2] + pos = ( + torch.arange(next_n, device="cuda", dtype=torch.int32) + .unsqueeze(0) + .expand(B, -1) + .contiguous() + ) + slot_map = torch.full((B, next_n), -1, dtype=torch.int32, device="cuda") + n_prior = 0 + elif case_id == "completion_at_end": + B, next_n, blocks = 2, 4, [0, 1] + pos = torch.tensor( + [[12, 13, 14, 15], [12, 13, 14, 15]], dtype=torch.int32, device="cuda" + ) + slot_map = torch.full((B, next_n), -1, dtype=torch.int32, device="cuda") + slot_map[:, 3] = torch.tensor( + [15, PAGE_SIZE + 15], dtype=torch.int32, device="cuda" + ) + n_prior = POOL_SIZE - next_n + elif case_id == "completion_mid_batch": + B, next_n, blocks = 3, 4, [0, 1, 2] + pos = torch.tensor([[13, 14, 15, 16]] * B, dtype=torch.int32, device="cuda") + slot_map = torch.full((B, next_n), -1, dtype=torch.int32, device="cuda") + slot_map[:, 2] = torch.tensor( + [15, PAGE_SIZE + 15, 2 * PAGE_SIZE + 15], dtype=torch.int32, device="cuda" + ) + n_prior = 13 + elif case_id == "non_uniform_padding": + B, next_n, blocks = 2, 4, [0, 1] + pos = torch.tensor( + [[12, 13, 14, 15], [12, 13, -1, -1]], dtype=torch.int32, device="cuda" + ) + slot_map = torch.full((B, next_n), -1, dtype=torch.int32, device="cuda") + slot_map[0, 3] = 15 + n_prior = POOL_SIZE - 4 + else: # plain_decode + B, next_n, blocks = 4, 1, [0, 1, 2, 3] + pos = torch.tensor([[5], [6], [7], [8]], dtype=torch.int32, device="cuda") + slot_map = torch.full((B, next_n), -1, dtype=torch.int32, device="cuda") + n_prior = 0 + + if case_id == "non_uniform_padding": + safe_pos = torch.where(pos >= 0, pos, 0) + tail_slot = torch.where(pos >= 0, _tail_slot_for(blocks, safe_pos), 0) + else: + tail_slot = _tail_slot_for(blocks, pos) + + key = torch.randn(B, next_n, HEAD_DIM, dtype=torch.bfloat16, device="cuda") + score = torch.randn(B, next_n, HEAD_DIM, dtype=torch.bfloat16, device="cuda") + ape = torch.randn(POOL_SIZE, HEAD_DIM, dtype=torch.float32, device="cuda") + + kv, tail = _make_caches() + _seed_prior(tail, blocks, n_prior) + r_ref = _torch_reference(kv, tail, tail_slot, key, score, ape, slot_map, pos) + r_kern = _run_kernel(kv, tail, tail_slot, key, score, ape, slot_map, pos) + _assert_eq(r_ref, r_kern) + + +@pytest.mark.parametrize("seed", list(range(20))) +def test_batched_matches_reference_fuzz(seed): + """Random B / next_n / start positions; covers 0, 1, and multi completion.""" + g = torch.Generator(device="cuda").manual_seed(seed) + B = int(torch.randint(1, 6, (1,), generator=g, device="cuda").item()) + next_n = int(torch.randint(1, 8, (1,), generator=g, device="cuda").item()) + blocks = list(range(B)) + + starts = torch.randint(0, 33, (B,), generator=g, device="cuda", dtype=torch.int32) + pos = starts.unsqueeze(1) + torch.arange( + next_n, device="cuda", dtype=torch.int32 + ).unsqueeze(0) + tail_slot = _tail_slot_for(blocks, pos) + + is_completion = pos % POOL_SIZE == POOL_SIZE - 1 + blk = torch.tensor(blocks, device="cuda", dtype=torch.int32).unsqueeze(1) + pool_slot = blk * PAGE_SIZE + (POOL_SIZE - 1) + slot_map = torch.where(is_completion, pool_slot, torch.full_like(pos, -1)) + + key = torch.randn( + B, next_n, HEAD_DIM, dtype=torch.bfloat16, device="cuda", generator=g + ) + score = torch.randn( + B, next_n, HEAD_DIM, dtype=torch.bfloat16, device="cuda", generator=g + ) + ape = torch.randn( + POOL_SIZE, HEAD_DIM, dtype=torch.float32, device="cuda", generator=g + ) + + kv, tail = _make_caches() + prior_g = torch.Generator(device="cuda").manual_seed(seed + 1000) + for b in range(B): + n_prior = int(starts[b].item()) % POOL_SIZE + if n_prior > 0: + pk = torch.randn( + n_prior, + HEAD_DIM, + dtype=torch.bfloat16, + device="cuda", + generator=prior_g, + ) + ps = torch.randn( + n_prior, + HEAD_DIM, + dtype=torch.bfloat16, + device="cuda", + generator=prior_g, + ) + tail[blocks[b], 0, :n_prior, :] = pk + tail[blocks[b], 1, :n_prior, :] = ps + + r_ref = _torch_reference(kv, tail, tail_slot, key, score, ape, slot_map, pos) + r_kern = _run_kernel(kv, tail, tail_slot, key, score, ape, slot_map, pos) + _assert_eq(r_ref, r_kern) diff --git a/tests/kernels/test_mhc_kernels.py b/tests/kernels/test_mhc_kernels.py index a2070dd595e3..a146c1a5c973 100644 --- a/tests/kernels/test_mhc_kernels.py +++ b/tests/kernels/test_mhc_kernels.py @@ -5,13 +5,23 @@ import pytest import torch import torch.nn as nn +import torch.nn.functional as F import vllm.model_executor.kernels.mhc # noqa: F401 +import vllm.model_executor.layers.mhc as mhc_layers from vllm.model_executor.kernels.mhc.tilelang import ( _tilelang_hc_prenorm_gemm, _torch_hc_prenorm_gemm, ) -from vllm.model_executor.layers.mhc import HAS_TILELANG_MHC +from vllm.model_executor.layers.mhc import ( + HAS_AITER_MHC, + HAS_AITER_MHC_FUSED, + HAS_AITER_MHC_FUSED_NORM, + HAS_AITER_MHC_PRE_NORM, + HAS_TILELANG_MHC, + MHCFusedPostPreOp, + MHCPreOp, +) from vllm.models.deepseek_v4.nvidia.model import ( DeepseekV4DecoderLayer, DeepseekV4Model, @@ -291,6 +301,220 @@ def run_ref(): torch.testing.assert_close(x, layer_input_ref, atol=1e-2, rtol=1e-2) +def _rocm_mhc_inputs(num_tokens=2, hidden_size=256, hc_mult=4): + residual = torch.randn( + (num_tokens, hc_mult, hidden_size), dtype=torch.bfloat16, device=DEVICE + ) + hc_mult3 = 2 * hc_mult + hc_mult * hc_mult + fn = ( + torch.randn( + (hc_mult3, hc_mult * hidden_size), dtype=torch.float32, device=DEVICE + ) + * 1e-4 + ) + hc_scale = torch.randn((3,), dtype=torch.float32, device=DEVICE) * 0.1 + hc_base = torch.randn((hc_mult3,), dtype=torch.float32, device=DEVICE) * 0.1 + norm_weight = torch.randn(hidden_size, dtype=torch.bfloat16, device=DEVICE) + return residual, fn, hc_scale, hc_base, norm_weight + + +@pytest.mark.skipif(not current_platform.is_rocm(), reason="ROCm required") +def test_mhc_pre_rocm_fallback_applies_norm(monkeypatch): + set_random_seed(0) + residual, fn, hc_scale, hc_base, norm_weight = _rocm_mhc_inputs() + rms_eps = hc_pre_eps = hc_sinkhorn_eps = norm_eps = 1e-6 + sinkhorn_repeat = 20 + hc_post_alpha = 1.0 + ref = mhc_pre_ref( + residual, + fn, + hc_scale, + hc_base, + rms_eps, + hc_pre_eps, + hc_sinkhorn_eps, + hc_post_alpha, + sinkhorn_repeat, + ) + expected_layer_input = F.rms_norm( + ref[2], (ref[2].shape[-1],), norm_weight, norm_eps + ) + monkeypatch.setattr(mhc_layers, "HAS_AITER_MHC", True) + monkeypatch.setattr(mhc_layers, "HAS_AITER_MHC_PRE_NORM", False) + monkeypatch.setattr(mhc_layers, "HAS_TILELANG_MHC", False) + + out = object.__new__(MHCPreOp).forward_hip( + residual, + fn, + hc_scale, + hc_base, + rms_eps, + hc_pre_eps, + hc_sinkhorn_eps, + hc_post_alpha, + sinkhorn_repeat, + norm_weight=norm_weight, + norm_eps=norm_eps, + ) + + torch.testing.assert_close(out[0], ref[0]) + torch.testing.assert_close(out[1], ref[1]) + torch.testing.assert_close(out[2], expected_layer_input) + + +@pytest.mark.skipif(not current_platform.is_rocm(), reason="ROCm required") +def test_mhc_fused_rocm_fallback_applies_norm(monkeypatch): + set_random_seed(0) + residual, fn, hc_scale, hc_base, norm_weight = _rocm_mhc_inputs() + x = torch.randn((2, 256), dtype=torch.bfloat16, device=DEVICE) + post_layer_mix = torch.randn((2, 4, 1), dtype=torch.float32, device=DEVICE) + comb_res_mix = torch.randn((2, 4, 4), dtype=torch.float32, device=DEVICE) + rms_eps = hc_pre_eps = hc_sinkhorn_eps = norm_eps = 1e-6 + sinkhorn_repeat = 20 + hc_post_alpha = 1.0 + residual_ref = mhc_post_ref(x, residual, post_layer_mix, comb_res_mix) + pre_ref = mhc_pre_ref( + residual_ref, + fn, + hc_scale, + hc_base, + rms_eps, + hc_pre_eps, + hc_sinkhorn_eps, + hc_post_alpha, + sinkhorn_repeat, + ) + expected_layer_input = F.rms_norm( + pre_ref[2], (pre_ref[2].shape[-1],), norm_weight, norm_eps + ) + monkeypatch.setattr(mhc_layers, "HAS_AITER_MHC_FUSED", False) + monkeypatch.setattr(mhc_layers, "HAS_TILELANG_MHC", False) + + out = object.__new__(MHCFusedPostPreOp).forward_hip( + x, + residual, + post_layer_mix, + comb_res_mix, + fn, + hc_scale, + hc_base, + rms_eps, + hc_pre_eps, + hc_sinkhorn_eps, + hc_post_alpha, + sinkhorn_repeat, + norm_weight=norm_weight, + norm_eps=norm_eps, + ) + + torch.testing.assert_close(out[0], residual_ref) + torch.testing.assert_close(out[1], pre_ref[0]) + torch.testing.assert_close(out[2], pre_ref[1]) + torch.testing.assert_close(out[3], expected_layer_input) + + +@pytest.mark.skipif( + not (current_platform.is_rocm() and HAS_AITER_MHC and HAS_AITER_MHC_PRE_NORM), + reason="AITER mHC with fused RMSNorm required", +) +def test_mhc_pre_rocm_aiter_fuses_norm(): + set_random_seed(0) + residual, fn, hc_scale, hc_base, norm_weight = _rocm_mhc_inputs( + num_tokens=2, hidden_size=7168 + ) + rms_eps = hc_pre_eps = hc_sinkhorn_eps = norm_eps = 1e-6 + sinkhorn_repeat = 20 + hc_post_alpha = 1.0 + ref = mhc_pre_ref( + residual, + fn, + hc_scale, + hc_base, + rms_eps, + hc_pre_eps, + hc_sinkhorn_eps, + hc_post_alpha, + sinkhorn_repeat, + ) + expected_layer_input = F.rms_norm( + ref[2], (ref[2].shape[-1],), norm_weight, norm_eps + ) + + out = object.__new__(MHCPreOp).forward_hip( + residual, + fn, + hc_scale, + hc_base, + rms_eps, + hc_pre_eps, + hc_sinkhorn_eps, + hc_post_alpha, + sinkhorn_repeat, + norm_weight=norm_weight, + norm_eps=norm_eps, + ) + + torch.testing.assert_close(out[0], ref[0], atol=5e-2, rtol=1e-2) + torch.testing.assert_close(out[1], ref[1], atol=5e-2, rtol=1e-2) + torch.testing.assert_close(out[2], expected_layer_input, atol=5e-2, rtol=1e-2) + + +@pytest.mark.skipif( + not ( + current_platform.is_rocm() and HAS_AITER_MHC_FUSED and HAS_AITER_MHC_FUSED_NORM + ), + reason="AITER fused mHC with RMSNorm required", +) +def test_mhc_fused_rocm_aiter_fuses_norm(): + set_random_seed(0) + residual, fn, hc_scale, hc_base, norm_weight = _rocm_mhc_inputs( + num_tokens=2, hidden_size=7168 + ) + x = torch.randn((2, 7168), dtype=torch.bfloat16, device=DEVICE) + post_layer_mix = torch.randn((2, 4, 1), dtype=torch.float32, device=DEVICE) + comb_res_mix = torch.randn((2, 4, 4), dtype=torch.float32, device=DEVICE) + rms_eps = hc_pre_eps = hc_sinkhorn_eps = norm_eps = 1e-6 + sinkhorn_repeat = 20 + hc_post_alpha = 1.0 + residual_ref = mhc_post_ref(x, residual, post_layer_mix, comb_res_mix) + pre_ref = mhc_pre_ref( + residual_ref, + fn, + hc_scale, + hc_base, + rms_eps, + hc_pre_eps, + hc_sinkhorn_eps, + hc_post_alpha, + sinkhorn_repeat, + ) + expected_layer_input = F.rms_norm( + pre_ref[2], (pre_ref[2].shape[-1],), norm_weight, norm_eps + ) + + out = object.__new__(MHCFusedPostPreOp).forward_hip( + x, + residual, + post_layer_mix, + comb_res_mix, + fn, + hc_scale, + hc_base, + rms_eps, + hc_pre_eps, + hc_sinkhorn_eps, + hc_post_alpha, + sinkhorn_repeat, + norm_weight=norm_weight, + norm_eps=norm_eps, + ) + + torch.testing.assert_close(out[0], residual_ref, atol=5e-2, rtol=1e-2) + torch.testing.assert_close(out[1], pre_ref[0], atol=5e-2, rtol=1e-2) + torch.testing.assert_close(out[2], pre_ref[1], atol=5e-2, rtol=1e-2) + torch.testing.assert_close(out[3], expected_layer_input, atol=5e-2, rtol=1e-2) + + @pytest.mark.skipif( not current_platform.is_rocm(), reason="ROCm required", diff --git a/tests/models/registry.py b/tests/models/registry.py index 513c9a676039..ed228349dbde 100644 --- a/tests/models/registry.py +++ b/tests/models/registry.py @@ -294,6 +294,7 @@ def check_available_online( "GlmMoeDsaForCausalLM": _HfExamplesInfo( "zai-org/GLM-5", min_transformers_version="5.0.1", is_available_online=False ), + "Glm5NextForCausalLM": _HfExamplesInfo("zai-org/GLM-5.3-Flash"), "GPT2LMHeadModel": _HfExamplesInfo("openai-community/gpt2"), "GPTBigCodeForCausalLM": _HfExamplesInfo( "bigcode/starcoder", @@ -928,6 +929,7 @@ def check_available_online( "zai-org/GLM-OCR", min_transformers_version="5.1.0", ), + "Glm5NextForConditionalGeneration": _HfExamplesInfo("zai-org/GLM-5.3-Flash"), "H2OVLChatModel": _HfExamplesInfo( "h2oai/h2ovl-mississippi-800m", trust_remote_code=True, @@ -1698,6 +1700,10 @@ def check_available_online( speculative_model="zai-org/GLM-OCR", min_transformers_version="5.1.0", ), + "Glm5NextMTPModel": _HfExamplesInfo( + "zai-org/GLM-5.3-Flash", + speculative_model="zai-org/GLM-5.3-Flash", + ), "HYV3MTPModel": _HfExamplesInfo( "tencent/Hy3-preview", speculative_model="tencent/Hy3-preview", diff --git a/tests/models/test_initialization.py b/tests/models/test_initialization.py index 3566eb534133..4100d0c8cf93 100644 --- a/tests/models/test_initialization.py +++ b/tests/models/test_initialization.py @@ -142,7 +142,7 @@ def _initialize_kv_caches_v1(self, vllm_config): patch.object(V1EngineCore, "_initialize_kv_caches", _initialize_kv_caches_v1), monkeypatch.context() as m, ): - if requires_spawn_multiprocessing(): + if requires_spawn_multiprocessing() or model_arch == "Glm5NextForCausalLM": # The EngineCore subprocess re-imports the class and does not # inherit the KV-cache patch above, so it OOMs. Run in-process # so the patch applies. diff --git a/tests/multimodal/test_video.py b/tests/multimodal/test_video.py index 547e6dde230f..cd42c52b3f27 100644 --- a/tests/multimodal/test_video.py +++ b/tests/multimodal/test_video.py @@ -20,6 +20,7 @@ PYNVVIDEOCODEC_VIDEO_BACKEND, VIDEO_LOADER_REGISTRY, DynamicVideoBackend, + Glm5NextVideoBackend, GLM46VVideoBackend, Molmo2VideoBackend, Qwen2VLVideoBackend, @@ -1470,3 +1471,233 @@ def test_glm46v_duration_estimation_from_fps(): assert len(indices) > 0 assert len(indices) % 2 == 0 assert all(0 <= idx < 90 for idx in indices) + + +def test_glm5next_backend_selected_for_processor(): + """Glm5NextVideoProcessor maps to the glm5next loader so only the + sampled frames are decoded instead of the whole container. Both the + borrowed-config spelling and the dedicated Glm5next class name (landing + with the new checkpoint) must resolve.""" + for name in ("Glm5NextVideoProcessor", "Glm5nextVideoProcessor"): + assert VIDEO_LOADER_REGISTRY.get_backend_for_video_processor(name) == "glm5next" + + +@pytest.mark.parametrize( + ("total_frames", "original_fps", "duration", "fps", "max_frames"), + [ + (900, 30.0, 30.0, -1, None), # 30s at flat 2.0 raw fps -> 60 frames + (3000, 30.0, 100.0, -1, None), + (72000, 30.0, 2400.0, -1, None), # 2048 cap + (48, 2.0, 24.0, -1, None), + (7, 30.0, 10.0, -1, None), # short video -> uniform spread + dedup + (300, 25.0, 0, -1, None), # duration derived from frame count + (900, 30.0, 30.0, 4, None), # request fps override (raw fps) + (900, 30.0, 30.0, -1, 16), # request max_frames override + ], +) +def test_glm5next_backend_indices_match_sampler( + total_frames, original_fps, duration, fps, max_frames +): + """The loader must select exactly the frames the processor's sampler + would, with target.fps mapping onto the raw-fps override.""" + from vllm.transformers_utils.processors.glm5next import ( + glm_sample_frame_indices, + ) + + source = VideoSourceMetadata( + total_frames_num=total_frames, original_fps=original_fps, duration=duration + ) + target = VideoTargetMetadata(num_frames=-1, fps=fps, max_duration=-1) + + indices = Glm5NextVideoBackend.compute_frames_index_to_sample( + source, target, max_frames=max_frames + ) + + assert indices == glm_sample_frame_indices( + total_frames, + original_fps, + duration, + target_fps=fps if fps > 0 else None, + max_frame_count=max_frames, + ) + assert len(indices) % 2 == 0 + assert indices == sorted(indices) # pair padding may repeat the last frame + assert all(0 <= idx < total_frames for idx in indices) + + +def test_glm5next_backend_metadata_contract(): + """create_hf_metadata reports the subset so the processor skips + re-sampling (do_sample_frames=False) and keeps the original totals.""" + source = VideoSourceMetadata(total_frames_num=900, original_fps=30.0, duration=30.0) + target = VideoTargetMetadata(num_frames=-1, fps=-1, max_duration=-1) + indices = Glm5NextVideoBackend.compute_frames_index_to_sample(source, target) + + metadata = Glm5NextVideoBackend.create_hf_metadata( + source, indices, video_backend="glm5next" + ) + assert metadata["do_sample_frames"] is False + assert metadata["frames_indices"] == indices + assert metadata["total_num_frames"] == 900 + assert metadata["fps"] == 30.0 + assert metadata["duration"] == 30.0 + + # A fully-selected source keeps do_sample_frames=True so the processor's + # sampler takes over on the complete frame set. + full = list(range(48)) + assert ( + Glm5NextVideoBackend.create_hf_metadata( + VideoSourceMetadata(total_frames_num=48, original_fps=2.0, duration=24.0), + full, + video_backend="glm5next", + )["do_sample_frames"] + is True + ) + + +def _write_gray_video(tmp_path, total_frames, fps, size=(32, 32)): + """Synthetic clip whose frame i is flat gray level i (near-lossless under + mp4v), so a decoded frame's level maps back to its source index.""" + cv2 = pytest.importorskip("cv2") + + path = tmp_path / f"gray_{total_frames}_{fps}.mp4" + writer = cv2.VideoWriter( + str(path), cv2.VideoWriter_fourcc(*"mp4v"), fps, size, isColor=False + ) + assert writer.isOpened() + for i in range(total_frames): + writer.write(np.full((*size, 1), i, dtype=np.uint8)) + writer.release() + return path + + +class _CountingCap: + """Proxy over a real capture that counts grab/seek decoding work.""" + + def __init__(self, cap): + self._cap = cap + self.grabs = 0 + self.seeks = 0 + + def get(self, prop): + return self._cap.get(prop) + + def grab(self): + self.grabs += 1 + return self._cap.grab() + + def set(self, prop, value): + self.seeks += 1 + return self._cap.set(prop, value) + + def read(self): + return self._cap.read() + + +@pytest.mark.parametrize("backend", ["opencv", "torchcodec"]) +def test_glm5next_backend_codec_parity(tmp_path, backend): + """Every codec samples the same GLM indices and decodes the same + frames; the OpenCV seek reader and torchcodec batched index-exact decode + must agree.""" + if backend == "torchcodec": + pytest.importorskip("torchcodec") + + from vllm.transformers_utils.processors.glm5next import ( + glm_sample_frame_indices, + ) + + total_frames, fps = 120, 10 + path = _write_gray_video(tmp_path, total_frames, fps) + # Dense default sampling (gap 5) and a sparse max_frames cap (gap 20). + for max_frames in (None, 6): + kwargs = {} if max_frames is None else {"max_frames": max_frames} + expected = glm_sample_frame_indices( + total_frames, float(fps), 12.0, max_frame_count=max_frames + ) + + frames, metadata = Glm5NextVideoBackend.load_bytes( + path.read_bytes(), backend=backend, **kwargs + ) + + assert metadata["frames_indices"] == expected + assert metadata["video_backend"].startswith(backend) + assert len(frames) == len(expected) + for i, idx in enumerate(expected): + assert abs(round(float(np.asarray(frames[i]).mean())) - idx) <= 1 + + +def test_glm5next_backend_decodes_only_sampled_frames(tmp_path): + """End to end over a synthetic clip: load_bytes returns exactly the + sampler's frame count, with the right frame content at each index.""" + pytest.importorskip("cv2") + + total_frames, fps = 60, 10 + path = _write_gray_video(tmp_path, total_frames, fps) + + from vllm.transformers_utils.processors.glm5next import ( + glm_sample_frame_indices, + ) + + expected = glm_sample_frame_indices(total_frames, float(fps), 6.0) + + frames, metadata = Glm5NextVideoBackend.load_bytes(path.read_bytes()) + + assert len(frames) == len(expected) + assert metadata["frames_indices"] == expected + assert metadata["do_sample_frames"] is (len(expected) == total_frames) + # Flat gray frames survive mp4v near-losslessly: each decoded frame's + # level maps back to its source frame index. + for frame, idx in zip(frames, expected): + decoded_idx = round(float(np.asarray(frame).mean())) + assert abs(decoded_idx - idx) <= 1 + + +def test_glm5next_read_frames_seeks_past_large_gaps(tmp_path): + """Sparse targets must not walk the container: the stock reader grabs + every frame up to the last index; the GLM reader seeks instead.""" + cv2 = pytest.importorskip("cv2") + + total_frames, fps = 200, 10 + path = _write_gray_video(tmp_path, total_frames, fps) + targets = [0, 80, 160, 190] # gaps of 80/80/30 -> only 30 <= threshold + + stock = cv2.VideoCapture(str(path)) + _, stock_indices = VideoBackend.read_frames(stock, targets, total_frames) + stock.release() + assert stock_indices == targets + + cap = _CountingCap(cv2.VideoCapture(str(path))) + frames, indices = Glm5NextVideoBackend.read_frames(cap, targets, total_frames) + cap._cap.release() + + assert indices == targets == stock_indices + for frame, idx in zip(frames, targets): + assert abs(round(float(np.asarray(frame).mean())) - idx) <= 1 + # Walks only the sub-threshold 30-frame hop; the two 80-frame gaps are + # seeks. The stock reader grabs all 190 preceding frames. + assert cap.grabs <= 29 + assert cap.seeks == 3 + + +def test_glm5next_read_frames_dense_walk_matches_stock(tmp_path): + """Dense targets keep the sequential walk (seeking would be slower) and + return the same frames as the stock reader.""" + cv2 = pytest.importorskip("cv2") + + total_frames, fps = 120, 10 + path = _write_gray_video(tmp_path, total_frames, fps) + targets = list(range(0, total_frames, 15)) # gaps of 15 -> all walking + + stock = cv2.VideoCapture(str(path)) + stock_frames, stock_indices = VideoBackend.read_frames(stock, targets, total_frames) + stock.release() + + cap = _CountingCap(cv2.VideoCapture(str(path))) + frames, indices = Glm5NextVideoBackend.read_frames(cap, targets, total_frames) + cap._cap.release() + + assert indices == stock_indices + for frame, stock_frame, idx in zip(frames, stock_frames, targets): + assert abs(float(np.asarray(frame).mean()) - float(stock_frame.mean())) <= 2.0 + assert abs(round(float(np.asarray(frame).mean())) - idx) <= 1 + # One initial seek, then pure walking -- no re-seek churn. + assert cap.seeks == 1 diff --git a/tests/transformers_utils/processors/test_glm5next.py b/tests/transformers_utils/processors/test_glm5next.py new file mode 100644 index 000000000000..917a3f627ce6 --- /dev/null +++ b/tests/transformers_utils/processors/test_glm5next.py @@ -0,0 +1,389 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +"""Unit tests for the vLLM-native GLM-5.3-Flash multimodal processor.""" + +import math + +import pytest +import torch +from PIL import Image +from transformers.image_utils import PILImageResampling +from transformers.video_utils import VideoMetadata + +from vllm.transformers_utils.processors.glm5next import ( + Glm5NextImageProcessor, + Glm5NextVideoProcessor, + _get_pad_content_size, + _resize_or_pad, + glm_sample_frame_indices, + smart_resize, +) + +PATCH_SIZE = 14 +MERGE_SIZE = 2 +PATCH_EXPAND_FACTOR = 1 # checkpoint processor_config.json +FACTOR = PATCH_SIZE * MERGE_SIZE * PATCH_EXPAND_FACTOR # 28 +PIXELS_PER_TOKEN = 2 * (PATCH_SIZE * MERGE_SIZE) ** 2 # 1568 +MIN_TOKENS = 16 +MAX_TOKENS = 800 # serving-cap equivalent budget, keeps test canvases small +MIN_PIXELS = MIN_TOKENS * PIXELS_PER_TOKEN +MAX_PIXELS = MAX_TOKENS * PIXELS_PER_TOKEN + + +def resize(height: int, width: int, t: int = 2) -> tuple[int, int]: + return smart_resize( + t, + height, + width, + t_factor=2, + h_factor=FACTOR, + w_factor=FACTOR, + min_pixels=MIN_PIXELS, + max_pixels=MAX_PIXELS, + ) + + +@pytest.mark.parametrize( + ("height", "width", "t", "expected"), + [ + (300, 400, 2, (308, 420)), # ceil alignment, no budget pressure + (50, 80, 2, (112, 168)), # below min budget: proportional upscale + (57, 3, 2, (504, 28)), # slim side ceils to one factor + (2160, 3840, 2, (588, 1064)), # shrink: max canvas under the cap + (112, 11200, 2, (84, 7420)), # extreme aspect stays proportional + (300, 400, 16, (252, 308)), # frame count eats into the budget + (112, 11200, 16, (28, 2800)), # slim side stays positive, budget holds + ], +) +def test_smart_resize_reference_values(height, width, t, expected): + got = smart_resize( + t, + height, + width, + t_factor=2, + h_factor=FACTOR, + w_factor=FACTOR, + min_pixels=MIN_PIXELS, + max_pixels=MAX_PIXELS, + ) + assert got == expected + # The aligned canvas never exceeds the budget it was fitted against. + t_bar = max(2, round(t / 2) * 2) + assert t_bar * got[0] * got[1] <= MAX_PIXELS + + +def test_smart_resize_stays_snapped_and_positive(): + heights = (1, 27, 57, 111, 113, 300, 720, 2160) + widths = (1, 27, 57, 111, 113, 400, 1280, 11200) + for h in heights: + for w in widths: + for t in (2, 4, 16, 64): + resized_h, resized_w = resize(h, w, t) + assert resized_h > 0 and resized_w > 0 + assert resized_h % FACTOR == 0 + assert resized_w % FACTOR == 0 + t_bar = max(2, round(t / 2) * 2) + assert t_bar * resized_h * resized_w <= MAX_PIXELS + + +@pytest.mark.parametrize( + ("height", "width", "match"), + [ + (0, 100, "must be positive"), + (100, 0, "must be positive"), + ], +) +def test_smart_resize_rejects_degenerate_inputs(height, width, match): + with pytest.raises(ValueError, match=match): + resize(height, width, 2) + + +def test_smart_resize_rejects_inverted_budget(): + with pytest.raises(ValueError, match="min_pixels must be less than or equal"): + smart_resize( + 2, + 100, + 100, + t_factor=2, + h_factor=FACTOR, + w_factor=FACTOR, + min_pixels=MAX_PIXELS, + max_pixels=MIN_PIXELS, + ) + + +def test_get_pad_content_size(): + # Oversized content shrinks proportionally, never upscales by default. + assert _get_pad_content_size(300, 400, 308, 420) == (300, 400) + assert _get_pad_content_size(600, 800, 308, 420) == (308, 410) + # allow_upscale enlarges small content toward the canvas. + assert _get_pad_content_size(28, 28, 112, 112, allow_upscale=True) == (112, 112) + + +def test_resize_or_pad_pads_right_and_bottom(): + stacked = torch.rand(1, 3, 300, 400) + + def identity_resize(x, size, resample=None): + assert (size.height, size.width) == (300, 400) + return x + + padded = _resize_or_pad(stacked, 308, 420, "pad", None, identity_resize) + assert padded.shape == (1, 3, 308, 420) + torch.testing.assert_close(padded[..., :300, :400], stacked) + assert padded[..., 300:, :].eq(0).all() + assert padded[..., :, 400:].eq(0).all() + + def force_resize(x, size, resample=None): + return torch.zeros(x.shape[0], x.shape[1], size.height, size.width) + + assert _resize_or_pad(stacked, 308, 420, "resize", None, force_resize).shape == ( + 1, + 3, + 308, + 420, + ) + with pytest.raises(ValueError, match="resize_mode"): + _resize_or_pad(stacked, 308, 420, "crop", None, identity_resize) + + +@pytest.fixture(scope="module") +def image_processor(): + return Glm5NextImageProcessor( + min_image_tokens=MIN_TOKENS, max_image_tokens=MAX_TOKENS + ) + + +@pytest.mark.parametrize( + ("height", "width"), + [(50, 80), (300, 400), (2160, 3840), (112, 11200)], +) +def test_image_grid_matches_smart_resize(image_processor, height, width): + out = image_processor(Image.new("RGB", (width, height)), return_tensors="pt") + resized_h, resized_w = resize(height, width) + grid = out["image_grid_thw"][0].tolist() + assert grid == [1, resized_h // PATCH_SIZE, resized_w // PATCH_SIZE] + assert out["pixel_values"].shape == ( + grid[1] * grid[2], + 3 * 2 * PATCH_SIZE * PATCH_SIZE, + ) + + +@pytest.mark.parametrize( + ("height", "width"), + [(50, 80), (300, 400), (2160, 3840), (112, 11200)], +) +def test_image_patch_count_matches_preprocess(image_processor, height, width): + out = image_processor(Image.new("RGB", (width, height)), return_tensors="pt") + grid = out["image_grid_thw"][0].tolist() + expected = grid[1] * grid[2] + assert image_processor.get_number_of_image_patches(height, width) == expected + + +@pytest.fixture(scope="module") +def video_processor(): + return Glm5NextVideoProcessor( + min_image_tokens=MIN_TOKENS, max_image_tokens=MAX_TOKENS + ) + + +def run_video_preprocess(video_processor, frames, **overrides): + kwargs = dict( + resample=PILImageResampling.BICUBIC, + image_mean=(0.48145466, 0.4578275, 0.40821073), + image_std=(0.26862954, 0.26130258, 0.27577711), + patch_size=PATCH_SIZE, + temporal_patch_size=2, + merge_size=MERGE_SIZE, + patch_expand_factor=PATCH_EXPAND_FACTOR, + return_tensors="pt", + ) + kwargs.update(overrides) + return video_processor._preprocess([frames], **kwargs) + + +def test_video_preprocess_pads_odd_frame_count(video_processor): + out = run_video_preprocess(video_processor, torch.rand(7, 3, 300, 400)) + grid = out["video_grid_thw"][0].tolist() + # 7 frames padded to 8 -> grid_t 4; canvas ceil-aligned (308, 420). + assert grid == [4, 22, 30] + assert out["pixel_values_videos"].shape == ( + 4 * 22 * 30, + 3 * 2 * PATCH_SIZE * PATCH_SIZE, + ) + + +def test_video_preprocess_keeps_min_side_under_frame_budget(video_processor): + out = run_video_preprocess(video_processor, torch.rand(16, 3, 112, 11200)) + # Budget-fitted canvas (28, 2800): 16 * 28 * 2800 == the pixel cap. + assert out["video_grid_thw"][0].tolist() == [8, 2, 200] # height 28, not 0 + + +def test_video_preprocess_pads_content_not_distort(video_processor): + frames = torch.zeros(4, 3, 300, 400) + frames[..., 50, 60] = 1.0 # marker inside the content area + # Bypass rescale/normalize so zero padding stays exactly zero. + out = run_video_preprocess( + video_processor, frames, do_rescale=False, do_normalize=False + ) + # Pad mode keeps the 300x400 content aspect on the (308, 420) canvas + # with zero padding on the right/bottom. + grid = out["video_grid_thw"][0].tolist() + assert grid == [2, 22, 30] + patches = out["pixel_values_videos"].view( + grid[0], + grid[1] // MERGE_SIZE, + grid[2] // MERGE_SIZE, + MERGE_SIZE, + MERGE_SIZE, + 3 * 2 * PATCH_SIZE * PATCH_SIZE, + ) + frame = patches[0].permute(0, 3, 1, 4, 2).reshape(grid[1], grid[2], -1) + # Patch columns fully right of the 400px content are pure padding. + assert frame[:, math.ceil(400 / PATCH_SIZE) :].abs().max() == 0 + + +@pytest.mark.parametrize( + ("total_frames", "fps", "duration", "expected_len", "first", "last"), + [ + # Dense source: the tp-scaled greedy overshoots extract_t, so the + # fixup spreads picks uniformly across the whole video (linspace). + (900, 30.0, 30.0, 60, 0, 899), + (3000, 30.0, 100.0, 200, 0, 2999), + # extract_t capped at 2048 frames. + (72000, 30.0, 2400.0, 2048, 0, 71999), + # Low container fps: the greedy picks every frame. + (48, 2.0, 24.0, 48, 0, 47), + # Duration derived from the frame count when metadata lacks it. + (300, 25.0, 0, 26, 0, 299), + ], +) +def test_glm_sample_frame_indices_behaviour( + total_frames, fps, duration, expected_len, first, last +): + indices = glm_sample_frame_indices(total_frames, fps, duration) + assert len(indices) == expected_len + assert indices[0] == first + assert indices[-1] == last + assert len(indices) % 2 == 0 + assert indices == sorted(indices) + assert all(0 <= idx < total_frames for idx in indices) + + +def test_glm_sample_frame_indices_short_clip_floor_spread(): + # extract_t (20) > total (7): evenly spaced timestamps, deduplicated, + # then pair-padded to an even count. + assert glm_sample_frame_indices(7, 30.0, 10.0) == [0, 1, 2, 3, 4, 5, 6, 6] + + +@pytest.mark.parametrize( + ("kwargs", "expected_len"), + [ + ({"target_fps": 0.5}, 16), + ({"max_frame_count": 16}, 16), + ({"target_fps": 8}, 240), # 30s * 8 = 240, under the 2048 cap + ], +) +def test_glm_sample_frame_indices_request_overrides(kwargs, expected_len): + indices = glm_sample_frame_indices(900, 30.0, 30.0, **kwargs) + assert len(indices) == expected_len + assert len(indices) % 2 == 0 + + +def test_sample_frames_fps_interval(video_processor): + indices = video_processor.sample_frames( + VideoMetadata(total_num_frames=900, fps=30.0, duration=30.0) + ) + assert len(indices) == 60 + assert indices[0] == 0 + assert indices[-1] == 899 # uniform spread reaches the final frame + + # Request overrides reach the sampler through both kwarg spellings. The + # 15-frame spread is odd, so pair-padding duplicates the last frame. + assert ( + len( + video_processor.sample_frames( + VideoMetadata(total_num_frames=900, fps=30.0, duration=30.0), + fps=0.5, + ) + ) + == 16 + ) + assert ( + len( + video_processor.sample_frames( + VideoMetadata(total_num_frames=900, fps=30.0, duration=30.0), + target_fps=0.5, + ) + ) + == 16 + ) + assert ( + len( + video_processor.sample_frames( + VideoMetadata(total_num_frames=900, fps=30.0, duration=30.0), + max_frames=16, + ) + ) + == 16 + ) + + +def test_defaults_mirror_checkpoint_config(): + """Bare instantiation matches the checkpoint's ``processor_config.json`` + token-budget style defaults; ``size`` carries no budget.""" + image_processor = Glm5NextImageProcessor() + assert image_processor.patch_expand_factor == 1 + assert image_processor.min_image_tokens == 16 + assert image_processor.max_image_tokens == 8000 + assert image_processor.resize_mode == "pad" + + video_processor = Glm5NextVideoProcessor() + assert video_processor.patch_expand_factor == 1 + assert video_processor.min_image_tokens == 16 + assert video_processor.max_image_tokens == 240000 + assert video_processor.fps_interval == 2.0 + assert video_processor.max_frame_count_dynamic == 2048 + + +def test_token_budgets_drive_geometry(): + """The token bounds fully determine the pixel budget and the grid.""" + token_proc = Glm5NextImageProcessor(min_image_tokens=64, max_image_tokens=512) + img = Image.new("RGB", (400, 300)) + grid = token_proc(img, return_tensors="pt")["image_grid_thw"][0].tolist() + # Budgets 100,352..802,816 px: the 300x400 canvas ceils to (308, 420) + # without hitting either bound -> 22x30 patches. + assert grid == [1, 22, 30] + assert token_proc.get_number_of_image_patches(300, 400) == 22 * 30 + + +def test_missing_token_budgets_rejected(): + proc = Glm5NextImageProcessor(min_image_tokens=None) + with pytest.raises(ValueError, match="min_image_tokens"): + proc(Image.new("RGB", (400, 300)), return_tensors="pt") + + +def test_video_config_fields_land(): + """fps_interval / max_frame_count_dynamic from the dedicated config + shape sampling without any request overrides.""" + proc = Glm5NextVideoProcessor( + fps_interval=4, + max_frame_count_dynamic=32, + ) + indices = proc.sample_frames( + VideoMetadata(total_num_frames=900, fps=30.0, duration=30.0) + ) + # 30s * 4 = 120 candidates, capped at 32 -> uniform spread of 32. + assert len(indices) == 32 + assert indices[0] == 0 + assert indices[-1] == 899 + + # Request overrides still win over the config values. + assert ( + len( + proc.sample_frames( + VideoMetadata(total_num_frames=900, fps=30.0, duration=30.0), + fps=0.5, + ) + ) + == 16 + ) diff --git a/tests/transformers_utils/test_config.py b/tests/transformers_utils/test_config.py index 35b5a697dbd8..17bb031e863d 100644 --- a/tests/transformers_utils/test_config.py +++ b/tests/transformers_utils/test_config.py @@ -10,6 +10,7 @@ from typing import cast from unittest.mock import MagicMock, patch +import pytest from transformers import PretrainedConfig from vllm.config.model import ModelConfig @@ -19,6 +20,59 @@ get_safetensors_params_metadata, try_get_generation_config, ) +from vllm.transformers_utils.configs.glm5_next import ( + Glm5NextConfig, + Glm5NextTextConfig, + Glm5NextVisionConfig, +) + + +def test_glm5_next_accepts_deepseek_sparse_attention_layers(): + layer_types = ["linear_attention", "deepseek_sparse_attention"] + + config = Glm5NextTextConfig( + num_hidden_layers=len(layer_types), layer_types=layer_types + ) + + assert config.layer_types == layer_types + assert config.layers_block_type == ["linear_attention", "attention"] + + +def test_glm5_next_accepts_prebuilt_subconfigs(): + text_config = Glm5NextTextConfig(hidden_size=1024) + vision_config = Glm5NextVisionConfig(hidden_size=768) + + config = Glm5NextConfig( + text_config=text_config, + vision_config=vision_config, + ) + + assert config.text_config is text_config + assert config.vision_config is vision_config + + +@pytest.mark.parametrize( + ("kwargs", "option"), + [ + ( + {"index_topk": 2048, "index_dsa_use_layernorm": False}, + "index_dsa_use_layernorm", + ), + ( + {"index_topk": 2048, "index_kpool_compress": False}, + "index_kpool_compress", + ), + ( + {"index_topk": 2048, "index_kpool_always_select_tail": False}, + "index_kpool_always_select_tail", + ), + ({"hres_vwnstyle": False}, "hres_vwnstyle"), + ({"mhc_no_norm_weight": True}, "mhc_no_norm_weight"), + ], +) +def test_glm5_next_rejects_unimplemented_config_options(kwargs, option): + with pytest.raises(NotImplementedError, match=option): + Glm5NextTextConfig(**kwargs) def test_get_llama3_eos_token(): diff --git a/tests/utils.py b/tests/utils.py index 71f9d4b669c8..0dd86a75af50 100644 --- a/tests/utils.py +++ b/tests/utils.py @@ -22,7 +22,7 @@ from contextlib import ExitStack, contextmanager from multiprocessing import Process, get_context from pathlib import Path -from typing import Any, Literal, cast +from typing import TYPE_CHECKING, Any, Literal, cast from unittest.mock import patch import anthropic @@ -46,10 +46,6 @@ from vllm.engine.arg_utils import AsyncEngineArgs from vllm.entrypoints.cli.serve import ServeSubcommand from vllm.logger import init_logger -from vllm.model_executor.kernels.linear import ( - _KernelT, - init_fp8_linear_kernel, -) from vllm.model_executor.layers.quantization.utils.quant_utils import ( QuantKey, ) @@ -65,6 +61,9 @@ ) from vllm.v1.engine.utils import get_engine_process_shutdown_timeout +if TYPE_CHECKING: + from vllm.model_executor.kernels.linear import _KernelT + logger = init_logger(__name__) FP8_DTYPE = current_platform.fp8_dtype() @@ -2349,9 +2348,11 @@ def __init__( out_dtype: torch.dtype | None = None, transpose_weights: bool = False, device: torch.device | None = None, - force_kernel: type[_KernelT] | None = None, + force_kernel: "type[_KernelT] | None" = None, ): super().__init__() + from vllm.model_executor.kernels.linear import init_fp8_linear_kernel + self.input_size_per_partition = weight_shape[1] self.output_size_per_partition = weight_shape[0] self.logical_widths = [self.output_size_per_partition] diff --git a/tests/v1/attention/test_cuda_backend_probe_errors.py b/tests/v1/attention/test_cuda_backend_probe_errors.py index dd05a0d4ba5f..f60430f855fa 100644 --- a/tests/v1/attention/test_cuda_backend_probe_errors.py +++ b/tests/v1/attention/test_cuda_backend_probe_errors.py @@ -1,6 +1,6 @@ # SPDX-License-Identifier: Apache-2.0 # SPDX-FileCopyrightText: Copyright contributors to the vLLM project -"""Tests for error handling in the CUDA attention backend probe. +"""Tests for CUDA attention backend selection and probe error handling. Environment-shaped probe failures (missing packages, unreadable caches, broken driver installs) must mark the backend unavailable rather than @@ -89,6 +89,48 @@ def test_get_valid_backends_keeps_probing_after_failure(): assert valid +def test_sm90_nope_mla_prefers_flashinfer_without_changing_rope_order(): + backend_cls = MagicMock() + backend_cls.validate_configuration.return_value = [] + sparse_backends = { + AttentionBackendEnum.FLASH_ATTN_MLA_SPARSE, + AttentionBackendEnum.FLASHMLA_SPARSE, + AttentionBackendEnum.FLASHINFER_MLA_SPARSE_SM90, + } + + def sparse_order(head_size: int) -> list[AttentionBackendEnum]: + config = SELECTOR_CONFIG._replace( + head_size=head_size, + use_mla=True, + use_sparse=True, + ) + with patch( + "vllm.platforms.cuda._get_attn_backend_class", + return_value=backend_cls, + ): + valid, _ = CudaPlatform.get_valid_backends( + device_capability=SM90, + attn_selector_config=config, + num_heads=32, + ) + return [ + candidate.backend + for candidate in valid + if candidate.backend in sparse_backends + ] + + assert sparse_order(512) == [ + AttentionBackendEnum.FLASHINFER_MLA_SPARSE_SM90, + AttentionBackendEnum.FLASH_ATTN_MLA_SPARSE, + AttentionBackendEnum.FLASHMLA_SPARSE, + ] + assert sparse_order(576) == [ + AttentionBackendEnum.FLASH_ATTN_MLA_SPARSE, + AttentionBackendEnum.FLASHMLA_SPARSE, + AttentionBackendEnum.FLASHINFER_MLA_SPARSE_SM90, + ] + + def test_selected_backend_probe_failure_raises_value_error_with_cause(): exc = OSError("libcuda.so.1: cannot open shared object file") with ( diff --git a/tests/v1/attention/test_flashinfer_mla_sparse_sm90.py b/tests/v1/attention/test_flashinfer_mla_sparse_sm90.py new file mode 100644 index 000000000000..b20c16330659 --- /dev/null +++ b/tests/v1/attention/test_flashinfer_mla_sparse_sm90.py @@ -0,0 +1,280 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +"""CPU tests for the FlashInfer SM90 sparse MLA backend wiring (no GPU). + +The FlashInfer wrapper and top-k conversion are replaced by CPU recorders; +the tests pin the contract between the impl and the kernel API: page_size=1 +varlen rows, reserved-buffer refresh, plan parameters (dims, NoPE/rope scale, +causality), ckv/kpe cache splitting, and the backend's model-shape gates. +""" + +from types import SimpleNamespace + +import pytest +import torch + +# isort: off +import vllm.v1.attention.backends.mla.flashinfer_mla_sparse_sm90 as sm90_mod +from vllm.v1.attention.backends.mla.flashinfer_mla_sparse_sm90 import ( + FlashInferMLASparseSM90Backend, + FlashInferMLASparseSM90Builder, + FlashInferMLASparseSM90Impl, +) +# isort: on + +BLOCK_SIZE = 64 +HEAD = 512 +TOPK = 128 # triton convert requires width % 128 == 0 + + +def ref_convert(req_id, block_table, token_indices, BLOCK_SIZE=64, **_): + out = torch.full_like(token_indices, -1) + counts = torch.zeros(token_indices.shape[0], dtype=torch.int32) + for t in range(token_indices.shape[0]): + vals = [] + for j in range(token_indices.shape[1]): + pos = int(token_indices[t, j]) + if pos == -1: + continue + blk = int(block_table[int(req_id[t]), pos // BLOCK_SIZE]) + if blk < 0: + continue + vals.append(blk * BLOCK_SIZE + pos % BLOCK_SIZE) + out[t, : len(vals)] = torch.tensor(vals, dtype=out.dtype) + counts[t] = len(vals) + return out, counts + + +class FakeWrapper: + def __init__(self): + self.plan_args = None + self.run_args = None + + def plan(self, *args, **kwargs): + self.plan_args = (args, kwargs) + + def run(self, *args, **kwargs): + q_nope, q_pe, ckv, kpe = args + self.run_args = (q_nope, q_pe, ckv, kpe, kwargs) + return torch.zeros( + q_nope.shape[0], q_nope.shape[1], ckv.shape[-1], dtype=torch.bfloat16 + ) + + +class FakeState: + def __init__(self, width, max_tokens=64): + self.kv_indices = torch.zeros(max_tokens * width, dtype=torch.int32) + self.kv_len_arr = torch.zeros(max_tokens, dtype=torch.int32) + self.wrapper = FakeWrapper() + self.plan_calls = [] + + def plan(self, num_tokens, kv_lens): + self.plan_calls.append((num_tokens, kv_lens)) + + +def make_impl(qk_rope, kv_dtype="fp8_e4m3", num_heads=2, topk_width=TOPK): + impl = object.__new__(FlashInferMLASparseSM90Impl) + impl.num_heads = num_heads + impl.head_size = HEAD + qk_rope + impl.scale = (HEAD + qk_rope) ** -0.5 + impl.kv_lora_rank = HEAD + impl.qk_rope_head_dim = qk_rope + impl.kv_cache_dtype = kv_dtype + impl.use_fp8_kv_cache = kv_dtype in ("fp8", "fp8_e4m3") + rows = 4 + impl.topk_indices_buffer = torch.full((rows, topk_width), -1, dtype=torch.int32) + return impl, rows + + +def make_batch(rows, topk_rows, own_blocks): + req_id = torch.tensor([0] * rows, dtype=torch.int32) + block_table = torch.zeros(1, 16, dtype=torch.int32) + block_table[:, 0] = own_blocks[0] + topk = torch.full((rows, TOPK), -1, dtype=torch.int32) + for t, row in enumerate(topk_rows): + topk[t, : len(row)] = torch.tensor(row, dtype=torch.int32) + return SimpleNamespace( + req_id_per_token=req_id, block_table=block_table, block_size=BLOCK_SIZE + ) + + +@pytest.mark.parametrize("qk_rope,kv_dtype", [(0, "fp8_e4m3"), (64, "auto")]) +def test_forward_wiring(monkeypatch, qk_rope, kv_dtype): + impl, rows = make_impl(qk_rope, kv_dtype) + state = FakeState(TOPK) + monkeypatch.setattr( + sm90_mod, "triton_convert_req_index_to_global_index", ref_convert + ) + + # req with context 10 < topk: 8 valid + -1 padding. + topk_rows = [ + [7, 3, 1, 9, 0, 2, 5, 8] + [-1] * (TOPK - 8), + [4, 0, 2, 3, 1] + [-1] * (TOPK - 5), + [6, 5, 4] + [-1] * (TOPK - 3), + [2] + [-1] * (TOPK - 1), + ] + meta = make_batch(rows, topk_rows, [3]) + meta.state = state + q_nope = torch.randn(rows, impl.num_heads, HEAD) + q_rope = torch.randn(rows, impl.num_heads, qk_rope) + cache = torch.zeros( + 8 * BLOCK_SIZE, + impl.head_size, + dtype=torch.uint8 if impl.use_fp8_kv_cache else torch.bfloat16, + ) + + out, lse = impl.forward_mqa( + (q_nope, q_rope), cache, meta, SimpleNamespace(_k_scale_float=0.5) + ) + assert lse is None and out.shape == (rows, impl.num_heads, HEAD) + + # Reserved buffers carry this step's slots; lengths are NOT refreshed + # here (the builder plans them host-side before capture/replay). + ref_slots, ref_counts = ref_convert( + meta.req_id_per_token, meta.block_table, impl.topk_indices_buffer + ) + width = TOPK + got_slots = state.kv_indices[: rows * width].view(rows, width) + for t in range(rows): + k = int(ref_counts[t]) + assert got_slots[t, :k].tolist() == ref_slots[t, :k].tolist() + assert state.plan_calls == [] + + assert state.wrapper.run_args is not None + q_pe, ckv, kpe, kwargs = state.wrapper.run_args[1:] + assert q_pe.shape == (rows, impl.num_heads, qk_rope) + assert ckv.shape == (8 * BLOCK_SIZE, 1, HEAD) + assert kpe.shape[-1] == qk_rope + if impl.use_fp8_kv_cache: + assert kwargs["ckv_scale"] == 0.5 and kwargs["kpe_scale"] == 1.0 + else: + assert kwargs == {} + + +def test_builder_attaches_its_state(monkeypatch): + builder = object.__new__(FlashInferMLASparseSM90Builder) + builder._index_topk = 2048 + builder._index_kpool = 4 + builder._async_scheduling = False + builder.state = FakeState(TOPK) + metadata = object.__new__(sm90_mod.FlashInferMLASparseSM90Metadata) + metadata.state = None + monkeypatch.setattr( + sm90_mod.FlashInferMLASparseMetadataBuilder, + "build", + lambda *_args, **_kwargs: metadata, + ) + cam = SimpleNamespace( + num_reqs=1, + query_start_loc_cpu=torch.tensor([0, 1], dtype=torch.int32), + seq_lens=torch.tensor([1], dtype=torch.int32), + seq_lens_cpu_upper_bound=torch.tensor([1], dtype=torch.int32), + positions=None, + ) + + result = builder.build(0, cam) + + assert result.state is builder.state + assert builder.state.plan_calls[0][0] == 1 + assert builder.state.plan_calls[0][1].tolist() == [1] + + +def test_plan_uses_state_params(monkeypatch): + """The NoPE/rope dims and scale live on the builder state, not the layer. + + plan() takes exact per-row KV lengths; the schedule is rebuilt on every + call (contexts grow between steps) and the indptrs are always full-size + with zero-query padding rows past num_tokens. + """ + impl, rows = make_impl(64, "auto") + wrapper = FakeWrapper() + state = sm90_mod._SM90State.__new__(sm90_mod._SM90State) + state.device = torch.device("cpu") + state.wrapper = wrapper + state.num_heads = 4 + state.kv_dtype = torch.bfloat16 + state.kv_lora_rank = HEAD + state.qk_rope_head_dim = 64 + state.sm_scale = 576**-0.5 + state.max_tokens = 4 + state.topk_width = TOPK + state.kv_indices = torch.zeros(4 * TOPK) + state._arange_cpu = torch.arange(5, dtype=torch.int32) + state._qo_cpu = torch.empty(5, dtype=torch.int32) + state._kv_cpu = torch.empty(5, dtype=torch.int32) + state._lens_cpu = torch.full((4,), TOPK, dtype=torch.int32) + + state.plan(3, torch.tensor([2, 5, 7], dtype=torch.int32)) + assert wrapper.plan_args is not None + args, kwargs = wrapper.plan_args + (qo, kv, indices, kv_len, heads, ckv, kpe, page, causal, scale) = args + assert qo.tolist() == [0, 1, 2, 3, 3] # clamp: rows past 3 have no queries + assert kv.tolist() == [i * TOPK for i in (0, 1, 2, 3, 3)] + assert kv_len.tolist() == [2, 5, 7, TOPK] # padded row keeps full width + assert (heads, ckv, kpe, page, causal) == (4, HEAD, 64, 1, False) + assert scale == 576**-0.5 + assert kwargs["q_data_type"] == torch.bfloat16 + assert kwargs["kv_data_type"] == torch.bfloat16 + + +def test_kv_lens_host_formula(): + """Per-row host lengths: context == position + 1; capped at + index_topk + trailing-pool remainder past the sparse threshold.""" + builder = object.__new__(FlashInferMLASparseSM90Builder) + builder._index_topk = 2048 + builder._index_kpool = 4 + builder._async_scheduling = False + cam = SimpleNamespace( + num_reqs=3, + query_start_loc_cpu=torch.tensor([0, 5, 7, 10], dtype=torch.int32), + seq_lens=torch.tensor([100, 9, 3000], dtype=torch.int32), + seq_lens_cpu_upper_bound=torch.tensor([100, 9, 3000], dtype=torch.int32), + positions=None, + ) + num_rows, lens = builder._kv_lens_host(cam) + assert num_rows == 10 + # req0: positions 95..99 -> ctx 96..100 (all <= 2048: full context) + # req1: positions 7,8 -> ctx 8,9 + # req2: positions 2997..2999 -> ctx 2998..3000 (> 2048: topk + ctx%4) + assert lens.tolist() == [96, 97, 98, 99, 100, 8, 9, 2050, 2051, 2048] + + +def test_kv_lens_host_empty(): + builder = object.__new__(FlashInferMLASparseSM90Builder) + builder._index_topk = 2048 + builder._index_kpool = 4 + cam = SimpleNamespace( + num_reqs=0, + query_start_loc_cpu=torch.tensor([0], dtype=torch.int32), + seq_lens=torch.zeros(0, dtype=torch.int32), + ) + num_rows, lens = builder._kv_lens_host(cam) + assert num_rows == 0 and lens.numel() == 0 + + +def test_supports_combination_gates(monkeypatch, default_vllm_config): + monkeypatch.setattr(sm90_mod, "has_flashinfer_sm90_nope_mla", lambda: True) + call = lambda **kw: FlashInferMLASparseSM90Backend.supports_combination( + head_size=576, + dtype=torch.bfloat16, + kv_cache_dtype="fp8_e4m3", + block_size=64, + use_mla=True, + has_sink=False, + use_sparse=True, + use_mm_prefix=False, + device_capability=SimpleNamespace(major=9), + **kw, + ) + assert call() is None # no model config: only the feature gate applies + + import vllm.config as cfg + + monkeypatch.setattr( + cfg, + "get_current_vllm_config", + lambda: SimpleNamespace(model_config=None), + ) + assert call() is None + monkeypatch.setattr(sm90_mod, "has_flashinfer_sm90_nope_mla", lambda: False) + assert "requires FlashInfer" in (call() or "") diff --git a/tests/v1/attention/test_kpool_tail_slot_mapping.py b/tests/v1/attention/test_kpool_tail_slot_mapping.py new file mode 100644 index 000000000000..0966f9b20a6b --- /dev/null +++ b/tests/v1/attention/test_kpool_tail_slot_mapping.py @@ -0,0 +1,369 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +"""CPU tests for the kpool tail slot mapping (no GPU required). + +The kpool tail cache is a 1-block-per-request circular ring addressed by +``pos % kpool`` (``KpoolTailSpec`` / ``KpoolTailManager``: exactly one block +allocated per request, never grown, so only column 0 of its block table is +ever written; the rest stays zero-initialized). + +The generic per-group slot kernel cannot express that layout: it maps +``pos -> bt[req][pos // bs] * bs + pos % bs`` (``_compute_slot_mappings_kernel`` +in vllm/v1/worker/gpu/block_table.py), so every token at ``pos >= kpool`` +reads a zero column and collapses onto physical tail block 0. All concurrent +requests then share one ``kpool``-slot ring and corrupt each other's pool +compression. + +These tests pin that defect's arithmetic, verify the circular replacement +(``compute_kpool_tail_slot_mapping``), and mirror the tail kernels' index +math to show cross-request pollution before the fix and isolation after. +""" + +from types import SimpleNamespace + +import pytest +import torch + +from vllm.v1.attention.backend import CommonAttentionMetadata +from vllm.v1.attention.backends.mla.indexer import ( + KpoolTailBackend, + KpoolTailMetadataBuilder, + compute_kpool_tail_slot_mapping, +) +from vllm.v1.kv_cache_interface import KpoolTailSpec, compute_layout_strides +from vllm.v1.kv_cache_layout import KVCacheLayout + +KPOOL = 4 + + +def test_tail_backend_layout_matches_kernel_pointer_arithmetic(): + (layout,) = KpoolTailBackend.supported_kv_cache_layouts() + spec = KpoolTailSpec( + block_size=KPOOL, + num_kv_heads=2, + head_size=128, + head_size_v=0, + dtype=torch.bfloat16, + sliding_window=KPOOL, + ) + strides = compute_layout_strides(spec, num_blocks=8, num_layers=3, layout=layout) + _, _, head_stride, state_stride, content_stride = strides + + assert layout is KVCacheLayout.LBHNC + assert head_stride == KPOOL * 128 * torch.bfloat16.itemsize + assert state_stride == 128 * torch.bfloat16.itemsize + assert content_stride == 1 + + +def make_tail_block_table(own_blocks, width=64): + """Tail-group block table as BlockTables produces it: column 0 holds the + request's single KpoolTailManager block, the remaining columns are never + written and stay zero.""" + bt = torch.zeros(len(own_blocks), width, dtype=torch.int32) + bt[:, 0] = torch.tensor(own_blocks, dtype=torch.int32) + return bt + + +def legacy_generic_tail_slots(block_table, query_start_loc, positions): + """Reference of the generic ``_compute_slot_mappings_kernel`` arithmetic + (block_table.py:305-313) applied to the tail group's table.""" + slots = [] + for req in range(block_table.shape[0]): + for i in range(query_start_loc[req], query_start_loc[req + 1]): + pos = int(positions[i]) + block_number = int(block_table[req, pos // KPOOL]) + slots.append(block_number * KPOOL + pos % KPOOL) + return torch.tensor(slots, dtype=torch.int64) + + +def circular_tail_slots( + slot_mapping, block_table, query_start_loc, positions, num_actual, num_reqs +): + return compute_kpool_tail_slot_mapping( + slot_mapping, + block_table, + query_start_loc, + positions, + num_actual, + num_reqs, + KPOOL, + ) + + +def make_batch(per_req_positions, padded_len=None): + positions = torch.cat( + [torch.tensor(p, dtype=torch.int64) for p in per_req_positions] + ) + num_actual = positions.numel() + num_reqs = len(per_req_positions) + lens = [len(p) for p in per_req_positions] + qsl = torch.zeros(num_reqs + 1, dtype=torch.int64) + torch.cumsum(torch.tensor(lens, dtype=torch.int64), 0, out=qsl[1:]) + if padded_len is None: + padded_len = num_actual + slot_mapping = torch.full((padded_len,), -1, dtype=torch.int64) + return positions, qsl, slot_mapping, num_actual, num_reqs + + +def test_legacy_generic_mapping_collapses_onto_block_zero(): + """The bug: with the manager's 1-column block table, the generic kernel + maps every pos >= kpool onto tail block 0, and distinct requests collide.""" + own_blocks = [5, 9] + per_req = [list(range(10)), list(range(12))] # prompts of len 10 and 12 + positions, qsl, _, num_actual, num_reqs = make_batch(per_req) + bt = make_tail_block_table(own_blocks) + + legacy = legacy_generic_tail_slots(bt, qsl, positions) + + # Every token at pos >= kpool resolves to block 0, not the request's own. + off = 0 + for req, prompt in enumerate(per_req): + req_slots = legacy[off : off + len(prompt)] + for pos in range(len(prompt)): + slot = int(req_slots[pos]) + if pos >= KPOOL: + assert slot // KPOOL == 0, ( + f"expected collapse onto block 0 at pos {pos}" + ) + assert slot // KPOOL != own_blocks[req] + off += len(prompt) + + # The two requests share ring slots -> cross-request pollution. + a_slots = set(legacy[: len(per_req[0])].tolist()) + b_slots = set(legacy[len(per_req[0]) :].tolist()) + assert a_slots & b_slots, "legacy mapping must collide across requests" + + +def test_circular_mapping_isolates_requests(): + """The fix: every token lands in its own request's block at pos % kpool, + and no slot is ever shared by two different requests (slots do recur + within a request every kpool positions -- that is the circular design).""" + own_blocks = [5, 9] + per_req = [list(range(10)), list(range(12))] + positions, qsl, slot_mapping, num_actual, num_reqs = make_batch(per_req) + bt = make_tail_block_table(own_blocks) + + out = circular_tail_slots(slot_mapping, bt, qsl, positions, num_actual, num_reqs) + + off = 0 + per_req_slots = [] + for req, prompt in enumerate(per_req): + req_slots = set() + for pos in range(len(prompt)): + slot = int(out[off + pos]) + assert slot // KPOOL == own_blocks[req], ( + f"req {req} pos {pos} left its tail block" + ) + assert slot % KPOOL == pos % KPOOL + req_slots.add(slot) + per_req_slots.append(req_slots) + off += len(prompt) + assert not per_req_slots[0] & per_req_slots[1] + + +@pytest.mark.parametrize("prompt_len", [1, 2, 3, 4]) +def test_circular_mapping_matches_generic_for_short_requests(prompt_len): + """For pos < kpool the generic kernel already picks the own block, so the + two mappings agree while every position fits the request's first block + (single-request behavior is unchanged).""" + own_blocks = [7] + per_req = [list(range(prompt_len))] + positions, qsl, slot_mapping, num_actual, num_reqs = make_batch(per_req) + bt = make_tail_block_table(own_blocks) + + legacy = legacy_generic_tail_slots(bt, qsl, positions) + out = circular_tail_slots(slot_mapping, bt, qsl, positions, num_actual, num_reqs) + assert torch.equal(out, legacy) + + +def test_circular_mapping_preserves_padding_and_empty_batch(): + own_blocks = [5, 9] + per_req = [list(range(10)), list(range(12))] + padded_len = sum(len(p) for p in per_req) + 8 + positions, qsl, slot_mapping, num_actual, num_reqs = make_batch( + per_req, padded_len=padded_len + ) + bt = make_tail_block_table(own_blocks) + + out = circular_tail_slots(slot_mapping, bt, qsl, positions, num_actual, num_reqs) + assert out.shape == slot_mapping.shape + assert torch.equal(out[num_actual:], torch.full_like(out[num_actual:], -1)) + + empty = circular_tail_slots(slot_mapping, bt, qsl, positions[:0], 0, num_reqs) + assert torch.equal(empty, slot_mapping) + + +def make_common_metadata(per_req_positions, own_blocks, with_positions=True): + positions, qsl, slot_mapping, num_actual, num_reqs = make_batch( + per_req_positions, padded_len=sum(len(p) for p in per_req_positions) + 4 + ) + bt = make_tail_block_table(own_blocks) + seq_lens = torch.tensor( + [max(p) + 1 if p else 1 for p in per_req_positions], dtype=torch.int64 + ) + return CommonAttentionMetadata( + query_start_loc=qsl, + query_start_loc_cpu=qsl.clone(), + seq_lens=seq_lens, + num_reqs=num_reqs, + num_actual_tokens=num_actual, + max_query_len=max((len(p) for p in per_req_positions), default=1), + max_seq_len=int(seq_lens.max()) if num_reqs else 1, + block_table_tensor=bt, + slot_mapping=slot_mapping, + positions=positions if with_positions else None, + ) + + +def make_tail_builder(block_size=KPOOL, max_num_batched_tokens=128): + builder = object.__new__(KpoolTailMetadataBuilder) + builder.kv_cache_spec = SimpleNamespace(block_size=block_size) + builder.slot_mapping_buffer = torch.empty(max_num_batched_tokens, dtype=torch.int64) + return builder + + +def test_builder_build_uses_circular_mapping(): + per_req = [list(range(10)), list(range(12))] + own_blocks = [5, 9] + cam = make_common_metadata(per_req, own_blocks) + meta = KpoolTailMetadataBuilder.build(make_tail_builder(), 0, cam) + + out = meta.slot_mapping + off = 0 + for req, prompt in enumerate(per_req): + for pos in range(len(prompt)): + slot = int(out[off + pos]) + assert slot // KPOOL == own_blocks[req] + assert slot % KPOOL == pos % KPOOL + off += len(prompt) + # Padding tail of the buffer keeps the -1 sentinel. + assert torch.equal( + out[cam.num_actual_tokens :], torch.full_like(out[cam.num_actual_tokens :], -1) + ) + + +def test_builder_build_falls_back_without_positions(): + """Capture / dummy builds without positions keep the generic mapping.""" + per_req = [list(range(10))] + cam = make_common_metadata(per_req, [5], with_positions=False) + meta = KpoolTailMetadataBuilder.build(make_tail_builder(), 0, cam) + assert meta.slot_mapping is cam.slot_mapping + + +def test_builder_reuses_slot_mapping_storage(): + builder = make_tail_builder() + first = make_common_metadata([list(range(10))], [5]) + first_meta = KpoolTailMetadataBuilder.build(builder, 0, first) + data_ptr = first_meta.slot_mapping.data_ptr() + + second = make_common_metadata([list(range(12))], [9]) + second_meta = KpoolTailMetadataBuilder.build(builder, 0, second) + + assert second_meta.slot_mapping.data_ptr() == data_ptr + assert second_meta.slot_mapping[:12].tolist() == [ + 9 * KPOOL + pos % KPOOL for pos in range(12) + ] + + +# --------------------------------------------------------------------------- +# Index-level mirror of the tail kernels: seed / stash / pool completion +# (addressing replicated from kpool_compress.py's Triton kernels). +# --------------------------------------------------------------------------- + + +class TailRingMirror: + """Mirror of _kpool_tail_seed_kernel / _kpool_decode_update_batched_kernel + addressing: block = tail_slot // kpool, ring offset = pos % kpool; a pool + completing at pos reads ring slots (pool_start + s) % kpool and uses the + current token's own K/score for the last member.""" + + def __init__(self, num_blocks, kpool=KPOOL): + self.kpool = kpool + self.k = torch.full((num_blocks, kpool, 3), float("nan")) + self.s = torch.full((num_blocks, kpool, 3), float("nan")) + + def stash(self, tail_slot, pos, k, s): + blk, off = tail_slot // self.kpool, pos % self.kpool + self.k[blk, off] = k + self.s[blk, off] = s + + seed = stash # the seed kernel writes with the same addressing + + def complete(self, tail_slot, pos, k, s): + blk = tail_slot // self.kpool + start = pos - (self.kpool - 1) + kk = torch.stack( + [self.k[blk, (start + i) % self.kpool] for i in range(self.kpool)] + ) + ss = torch.stack( + [self.s[blk, (start + i) % self.kpool] for i in range(self.kpool)] + ) + kk[-1], ss[-1] = k, s # is_current for the completing token + w = torch.softmax(ss, dim=0) + return (kk * w).sum(0) + + +def token_kv(req, pos): + k = torch.tensor([pos + 100.0 * req, pos + 0.5, 2.0 * pos + 0.25]) + s = torch.tensor([0.1 * (pos + 1) + req, 0.2 * pos, 0.05 * pos]) + return k, s + + +def tail_slot_for(mapping, req, pos, own_block): + if mapping == "legacy": + bt_val = own_block if pos < KPOOL else 0 + return bt_val * KPOOL + pos % KPOOL + return own_block * KPOOL + pos % KPOOL + + +def run_scenario(mapping, interleave): + """Requests A (block 5, prompt len 9) and B (block 9, prompt len 11) + decode concurrently; returns A's boundary pool [8, 9, 10, 11].""" + ring = TailRingMirror(num_blocks=16) + blocks = {"A": 5, "B": 9} + prompts = {"A": 9, "B": 11} + + def slot(req, pos): + return tail_slot_for(mapping, 0 if req == "A" else 1, pos, blocks[req]) + + # Prefill: seed each request's trailing incomplete pool. + for req, L in prompts.items(): + for pos in range(L - (L % KPOOL or KPOOL), L): + if pos < 0: + continue + ring.stash(slot(req, pos), pos, *token_kv(0 if req == "A" else 1, pos)) + + # Decode: A emits pos 9, 10, 11; B emits 11, 12, 13. `interleave` + # processes B before A within a step, which is exactly what concurrent + # Triton programs do when they share tail block 0. + order = ["B", "A"] if interleave else ["A", "B"] + decode = {"A": [9, 10, 11], "B": [11, 12, 13]} + result = None + for step in range(3): + for req in order: + pos = decode[req][step] + k, s = token_kv(0 if req == "A" else 1, pos) + if pos % KPOOL == KPOOL - 1: + pool = ring.complete(slot(req, pos), pos, k, s) + if req == "A": + result = pool + ring.stash(slot(req, pos), pos, k, s) + return result + + +def test_interleaved_decode_pollution_legacy_vs_circular(): + """Ground truth: request A decoding alone (its ring used exclusively).""" + torch.manual_seed(0) + ground_truth = run_scenario("circular", interleave=False) + + legacy = run_scenario("legacy", interleave=True) + circular = run_scenario("circular", interleave=True) + + # The old mapping lets request B's tokens into A's ring: A's boundary + # pool is compressed from 2 of B's tokens -> wrong. + assert not torch.allclose(legacy, ground_truth), ( + f"legacy mapping unexpectedly clean: {legacy} vs {ground_truth}" + ) + + # The circular mapping keeps the rings isolated under interleaving. + torch.testing.assert_close(circular, ground_truth) diff --git a/tests/v1/attention/test_rocm_glm5next_sparse.py b/tests/v1/attention/test_rocm_glm5next_sparse.py new file mode 100644 index 000000000000..35b51c855790 --- /dev/null +++ b/tests/v1/attention/test_rocm_glm5next_sparse.py @@ -0,0 +1,127 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project + +import pytest +import torch + +from vllm.platforms import current_platform +from vllm.triton_utils import tl, triton +from vllm.v1.attention.backends.mla.rocm_aiter_mla_sparse import ( + _use_rocm_sparse_triton, + fit_kpool_indices_to_aiter, +) +from vllm.v1.attention.ops.rocm_aiter_mla_sparse import ( + _sparse_kv_row_offset, + _validate_dsv4_sparse_dims, + _validate_sparse_dims, +) + + +@triton.jit +def _store_sparse_kv_row_offset_kernel(slot_ptr, output_ptr, stride: tl.constexpr): + slot = tl.load(slot_ptr) + tl.store(output_ptr, _sparse_kv_row_offset(slot, stride)) + + +def test_fit_kpool_indices_preserves_tail_and_best_history(): + token_indices = torch.tensor( + [ + [10, 9, 8, 7, 6, 5, 100, 101], + [10, 9, 8, -1, -1, -1, 100, -1], + [-1, -1, -1, -1, -1, -1, -1, -1], + ], + dtype=torch.int32, + ) + + fitted = fit_kpool_indices_to_aiter(token_indices, topk_tokens=6) + + assert fitted.tolist() == [ + [10, 9, 8, 7, 100, 101], + [10, 9, 8, 100, -1, -1], + [-1, -1, -1, -1, -1, -1], + ] + + +def test_fit_kpool_indices_exact_width_is_noop(): + token_indices = torch.tensor([[3, 2, 1, -1]], dtype=torch.int32) + + fitted = fit_kpool_indices_to_aiter(token_indices, topk_tokens=4) + + assert fitted.data_ptr() == token_indices.data_ptr() + + +def test_fit_kpool_indices_rejects_narrow_input(): + with pytest.raises(ValueError, match="at least topk_tokens"): + fit_kpool_indices_to_aiter( + torch.zeros((1, 3), dtype=torch.int32), topk_tokens=4 + ) + + +@pytest.mark.parametrize( + ( + "kv_cache_dtype", + "head_size", + "num_prefills", + "num_decodes", + "num_decode_tokens", + "max_query_len", + "expected", + ), + [ + ("auto", 512, 1, 0, 0, 32, True), + ("auto", 512, 1, 2, 2, 32, True), + ("auto", 512, 0, 2, 2, 1, True), + ("fp8", 512, 1, 0, 0, 32, False), + ("auto", 576, 1, 0, 0, 32, False), + ("auto", 512, 0, 2, 4, 2, False), + ], +) +def test_rocm_sparse_triton_route( + kv_cache_dtype, + head_size, + num_prefills, + num_decodes, + num_decode_tokens, + max_query_len, + expected, +): + assert ( + _use_rocm_sparse_triton( + kv_cache_dtype=kv_cache_dtype, + head_size=head_size, + kv_lora_rank=512, + num_prefills=num_prefills, + num_decodes=num_decodes, + num_decode_tokens=num_decode_tokens, + max_query_len=max_query_len, + ) + is expected + ) + + +def test_rocm_sparse_attention_accepts_glm_nope_dimensions(): + _validate_sparse_dims(512, 512, 0, "test") + + +def test_rocm_sparse_attention_rejects_inconsistent_dimensions(): + with pytest.raises(AssertionError, match="expected head_dim"): + _validate_sparse_dims(511, 512, 0, "test") + + +def test_dsv4_sparse_attention_keeps_layout_constraint(): + _validate_dsv4_sparse_dims(512, 448, 64, "test") + with pytest.raises(AssertionError, match="expects 448 NoPE dims"): + _validate_dsv4_sparse_dims(512, 512, 0, "test") + + +@pytest.mark.skipif(not current_platform.is_rocm(), reason="ROCm required") +def test_sparse_prefill_kv_row_offset_does_not_overflow_int32(): + # GLM's 640-token pages cross the signed-int32 address boundary at block + # 6554 for a 512-element KV row. The production kernel must promote the + # slot before multiplying by the row stride. + slot = torch.tensor([6554 * 640], dtype=torch.int32, device="cuda") + output = torch.empty(1, dtype=torch.int64, device="cuda") + + _store_sparse_kv_row_offset_kernel[(1,)](slot, output, stride=512) + + assert output.item() == 6554 * 640 * 512 diff --git a/tests/v1/attention/test_sparse_indexer_decode_seq_lens.py b/tests/v1/attention/test_sparse_indexer_decode_seq_lens.py new file mode 100644 index 000000000000..94fdadd6e95d --- /dev/null +++ b/tests/v1/attention/test_sparse_indexer_decode_seq_lens.py @@ -0,0 +1,184 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +"""CPU tests for the decode top-k row->token mapping (no GPU required). + +On non-uniform decode batches (``requires_padding`` -- mixed plain-decode and +spec-verify requests, or variable MTP verify lens; taken on Hopper where the +varlen/flatten logits path is unavailable), the pool-topk rows follow the +PADDED ``[batch_size, next_n]`` grid: row ``(b, t)`` is flat decode token +``offset_b + t``. The former inline ``dec_seq = positions[:n] + 1`` +(``n = batch_size * next_n``) indexes the flat per-token layout with padded +coordinates, so rows after the first non-uniform request read another +request's positions and rows past the decode region read prefill tokens -- +``expand_pools_and_append_tail`` then anchors the tail at a foreign length. + +These tests pin the defect with the exact production arithmetic and verify +the layout-aware replacement (``_decode_topk_seq_lens``), including the tail +expansion consequences via the pure-torch expand/append pair that the fused +kernel is documented to replicate. +""" + +import torch + +# Bootstrap the glm5next package before entering the indexer module: its +# kpool_compress import runs glm5next/__init__, which pulls model -> +# attention -> back into sparse_attn_indexer_kpool (attention.py imports the +# class at module scope). Production always enters via attention.py first. +# isort: off +import vllm.models.glm5next # noqa: F401 +from vllm.models.glm5next.nvidia.ops.kpool_compress import ( # noqa: E402 + append_tail_to_topk, + expand_pools_to_tokens, +) +# isort: on + +import vllm.model_executor.layers.sparse_attn_indexer_kpool as indexer_mod +from vllm.model_executor.layers.sparse_attn_indexer_kpool import ( + _decode_topk_seq_lens, + _fill_short_decode_causal_indices, +) +from vllm.platforms import current_platform + +KPOOL = 4 +TOPK_TOKENS = 16 +SELECT_K = TOPK_TOKENS // KPOOL + + +def test_kpool_ops_dispatch_matches_platform(): + expected_backend = ".amd." if current_platform.is_rocm() else ".nvidia." + assert expected_backend in indexer_mod.kpool_ops.__name__ + + +def test_short_decode_fills_exact_causal_rows(): + topk = torch.full((3, 8), 99, dtype=torch.int32) + positions = torch.tensor([0, 3, 7], dtype=torch.int64) + + assert _fill_short_decode_causal_indices(topk, positions, 3, 8, 8) + assert topk.tolist() == [ + [0, -1, -1, -1, -1, -1, -1, -1], + [0, 1, 2, 3, -1, -1, -1, -1], + [0, 1, 2, 3, 4, 5, 6, 7], + ] + + +def test_short_decode_leaves_buffer_unchanged_for_sparse_context(): + topk = torch.full((2, 8), 99, dtype=torch.int32) + before = topk.clone() + + assert not _fill_short_decode_causal_indices(topk, torch.tensor([7, 8]), 2, 9, 8) + assert torch.equal(topk, before) + + +def make_non_uniform_batch(): + """3 requests: plain decode (1 token), MTP verify (4), adaptive verify (3). + + Flat decode positions (production layout: decode tokens first, then any + prefill tokens of the same batch): + req0: [30] (context len 30) + req1: [100..103] (verify at context len 100) + req2: [7, 8, 9] (verify at context len 7) + Followed by 4 prefill tokens at positions 555..558 so the flat tensor is + at least ``n = 3 * 4`` long, as it is in a real mixed batch. + """ + decode_lens = torch.tensor([1, 4, 3], dtype=torch.int64) + per_req_positions = [[30], [100, 101, 102, 103], [7, 8, 9]] + flat_decode = [p for req in per_req_positions for p in req] + positions = torch.tensor(flat_decode + [555, 556, 557, 558], dtype=torch.int64) + return decode_lens, per_req_positions, positions + + +def expected_row_seq_lens(per_req_positions, batch_size, next_n): + """Ground truth: row (b, t) -> pos + 1 for real rows, 0 for pad rows.""" + out = torch.zeros(batch_size * next_n, dtype=torch.int32) + for b, req in enumerate(per_req_positions): + for t, pos in enumerate(req): + out[b * next_n + t] = pos + 1 + return out + + +def test_legacy_flat_layout_misaligns_and_bleeds(): + """The bug: ``positions[:n] + 1`` with padded-row coordinates reads other + requests' positions and prefill positions.""" + decode_lens, per_req, positions = make_non_uniform_batch() + batch_size = decode_lens.shape[0] + next_n = int(decode_lens.max()) + n = batch_size * next_n + assert n == 12 + + legacy = positions[:n].to(torch.int32) + 1 + expected = expected_row_seq_lens(per_req, batch_size, next_n) + + # Row (1, 0): req1's first verify token (pos 100, seq 101) reads flat + # index 4 -> req1's LAST verify token (pos 103). Its true tail token is + # dropped and the tail anchors 3 tokens late. + assert int(legacy[4]) == 104 and int(expected[4]) == 101 + + # Rows (2, *): req2 starts at flat offset 5, but padded coordinates point + # at flat indices 8..11 -- the batch's PREFILL positions. + for t in range(3): + assert int(legacy[2 * next_n + t]) >= 556, ( + "legacy row (2, t) should read prefill positions" + ) + assert not torch.equal(legacy, expected) + + +def test_helper_uniform_layout_matches_flat_slice(): + """Uniform batches keep the flat shortcut (zero behavior/perf change).""" + per_req = [[200, 201, 202, 203], [50, 51, 52, 53]] + positions = torch.tensor([p for r in per_req for p in r], dtype=torch.int64) + decode_lens = torch.tensor([4, 4], dtype=torch.int64) + out = _decode_topk_seq_lens(positions, decode_lens, 8, 2, 4, requires_padding=False) + assert torch.equal(out, positions[:8].to(torch.int32) + 1) + + +def test_helper_padded_layout_per_row(): + """Non-uniform batches map every padded row to its own token's position; + pad rows get 0 (empty tail).""" + decode_lens, per_req, positions = make_non_uniform_batch() + out = _decode_topk_seq_lens(positions, decode_lens, 8, 3, 4, requires_padding=True) + expected = expected_row_seq_lens(per_req, 3, 4) + assert torch.equal(out, expected) + # Pad rows (0, 1..3) and (2, 3) collapse to 0 -> no tail appended. + assert out[1] == 0 and out[2] == 0 and out[3] == 0 and out[11] == 0 + + +def expand_tail_region(dec_seq): + """Run the production pure-torch expand + append pair (the fused kernel + replicates it exactly on the identity path) and return the tail columns + [TOPK_TOKENS, TOPK_TOKENS + KPOOL - 1).""" + rows = dec_seq.shape[0] + pool_ids = torch.arange(SELECT_K, dtype=torch.int64).expand(rows, SELECT_K) + valid = torch.ones_like(pool_ids, dtype=torch.bool) + expanded = expand_pools_to_tokens(pool_ids, valid, TOPK_TOKENS, KPOOL) + seq_lens = dec_seq.to(torch.int32) + pool_lens = (seq_lens // KPOOL).to(torch.int32) + out = append_tail_to_topk(expanded, seq_lens, pool_lens, KPOOL) + return out[:, TOPK_TOKENS:] + + +def test_tail_expansion_legacy_vs_fixed(): + """End-to-end tail consequence: the fixed mapping appends exactly the + request's trailing incomplete pool; the legacy flat slice drops real tail + tokens and emits indices far past the request's own sequence.""" + decode_lens, per_req, positions = make_non_uniform_batch() + n = 12 + fixed = _decode_topk_seq_lens( + positions, decode_lens, 8, 3, 4, requires_padding=True + ) + legacy = positions[:n].to(torch.int32) + 1 + + fixed_tail = expand_tail_region(fixed) + legacy_tail = expand_tail_region(legacy) + + # Fixed: row (1, 0) (seq 101) keeps its single tail token 100; row (2, 2) + # (seq 10) keeps its full trailing pool [8, 9]. + assert fixed_tail[4, 0].item() == 100 + assert fixed_tail[10, :2].tolist() == [8, 9] + + # Legacy: row (1, 0) loses its tail token (anchored at 104, count 0)... + assert legacy_tail[4, 0].item() == -1 + # ...and row (2, 2) reads prefill position 557 -> tail indices 556/557, + # way past req2's 10-token sequence -> out-of-bounds block-table reads. + assert legacy_tail[10, :2].tolist() == [556, 557] + + assert not torch.equal(legacy_tail, fixed_tail) diff --git a/tests/v1/attention/test_sparse_mla_backends.py b/tests/v1/attention/test_sparse_mla_backends.py index 31192f3e76cb..b2490e391b01 100644 --- a/tests/v1/attention/test_sparse_mla_backends.py +++ b/tests/v1/attention/test_sparse_mla_backends.py @@ -40,11 +40,13 @@ allow_module_level=True, ) +import vllm.v1.attention.backends.mla.flashinfer_mla_sparse as flashinfer_sparse_mod from vllm.model_executor.layers.attention.mla_attention import ( _canonicalize_sparse_mla_kv_cache_dtype, ) from vllm.utils.math_utils import cdiv from vllm.v1.attention.backends.mla.flashinfer_mla_sparse import ( + FlashInferMLASparseImpl, FlashInferMLASparseTRTLLMBackend, ) from vllm.v1.attention.backends.mla.flashmla_sparse import ( @@ -82,6 +84,60 @@ DEVICE_TYPE = current_platform.device_type +def test_nope_flashinfer_sparse_mla_uses_model_scale(monkeypatch): + """Weight absorption must not change the model's attention temperature.""" + model_scale = 256**-0.5 + kv_lora_rank = 512 + topk = torch.zeros((1, 1), dtype=torch.int32) + metadata = SimpleNamespace( + req_id_per_token=torch.zeros(1, dtype=torch.int32), + block_table=torch.zeros((1, 1), dtype=torch.int32), + block_size=1, + ) + recorded_scale = None + + impl = object.__new__(FlashInferMLASparseImpl) + impl.scale = model_scale + impl.qk_nope_head_dim = 256 + impl.kv_lora_rank = kv_lora_rank + impl.qk_rope_head_dim = 0 + impl.kv_cache_dtype = "auto" + impl.topk_indices_buffer = topk + impl.dcp_world_size = 1 + impl._workspace_buffer = torch.empty(1) + impl.bmm1_scale = None + impl.bmm2_scale = None + impl.is_nope_mla = True + impl.need_to_return_lse_for_decode = False + monkeypatch.setattr( + flashinfer_sparse_mod, + "triton_convert_req_index_to_global_index", + lambda *args, **kwargs: (topk, torch.ones(1, dtype=torch.int32)), + ) + + import flashinfer.decode + + def fake_flashinfer(**kwargs): + nonlocal recorded_scale + recorded_scale = kwargs["bmm1_scale"] + return torch.zeros((1, 1, 1, kv_lora_rank)) + + monkeypatch.setattr( + flashinfer.decode, + "trtllm_batch_decode_with_kv_cache_mla", + fake_flashinfer, + ) + impl.forward_mqa( + torch.zeros(1, 1, kv_lora_rank), + torch.zeros(1, 1, kv_lora_rank), + metadata, + SimpleNamespace(), + ) + + assert recorded_scale == model_scale + assert recorded_scale != kv_lora_rank**-0.5 + + def _float_to_e8m0_truncate(f: float) -> float: """Simulate SM100's float -> e8m0 -> bf16 scale conversion. e8m0 format only stores the exponent (power of 2). @@ -1451,14 +1507,14 @@ def test_triton_convert_returns_valid_counts(num_topk_tokens: int): return_valid_counts=False, ) assert isinstance(result_only, torch.Tensor) - torch.testing.assert_close(result_only, result, rtol=0, atol=0) + for row, num_valid in enumerate(expected_valid): + compact_valid = result[row, :num_valid].sort().values + original_valid = result_only[row][result_only[row] >= 0].sort().values + torch.testing.assert_close(compact_valid, original_valid, rtol=0, atol=0) + assert torch.all(result[row, num_valid:] == -1) def test_flashmla_cache_dtype_aliases_use_ds_layout(): - from vllm.model_executor.layers.attention.mla_attention import ( - _canonicalize_sparse_mla_kv_cache_dtype, - ) - # kv-cache dtype aliases are canonicalized to fp8_ds_mla before the layer # stores kv_cache_dtype, so they cannot bypass the gate. for alias in ("fp8", "fp8_e4m3"): diff --git a/tests/v1/core/prefix_cache/test_partial_prefix_cache_hits.py b/tests/v1/core/prefix_cache/test_partial_prefix_cache_hits.py index 9558409ae9e9..864d88831ad4 100644 --- a/tests/v1/core/prefix_cache/test_partial_prefix_cache_hits.py +++ b/tests/v1/core/prefix_cache/test_partial_prefix_cache_hits.py @@ -23,6 +23,7 @@ from vllm.v1.core.sched.scheduler import Scheduler from vllm.v1.kv_cache_interface import ( FullAttentionSpec, + KpoolTailSpec, KVCacheConfig, KVCacheGroupSpec, MambaSpec, @@ -1801,6 +1802,127 @@ def test_hybrid_sliding_window_group_disables_partial_hash_hits(): assert len(computed_blocks.blocks[0]) * hash_block_size == num_computed +def test_opted_out_scratch_group_keeps_partial_hash_hits(): + hash_block_size = 2 + mamba_block_size = 2 * hash_block_size + kv_cache_config = KVCacheConfig( + num_blocks=24, + kv_cache_tensors=[], + kv_cache_groups=[ + KVCacheGroupSpec( + ["full"], + FullAttentionSpec( + block_size=hash_block_size, + num_kv_heads=1, + head_size=1, + dtype=torch.float32, + ), + ), + KVCacheGroupSpec( + ["mamba"], + MambaSpec( + block_size=mamba_block_size, + shapes=(1, 1), + dtypes=(torch.float32,), + mamba_cache_mode="align", + ), + ), + KVCacheGroupSpec( + ["tail"], + KpoolTailSpec( + block_size=mamba_block_size, + num_kv_heads=1, + head_size=1, + dtype=torch.float32, + sliding_window=mamba_block_size, + ), + ), + ], + ) + manager = make_kv_cache_manager( + kv_cache_config=kv_cache_config, + max_model_len=8192, + enable_caching=True, + hash_block_size=hash_block_size, + ) + + assert manager.coordinator.enable_partial_hash_hits + + req0 = make_request("0", [0, 0, 1, 1, 2, 2], hash_block_size, sha256) + computed_blocks, num_computed, _ = manager.get_computed_blocks(req0) + assert manager.allocate_slots(req0, 6, num_computed, computed_blocks) is not None + manager.free(req0) + manager.new_step_starts() + + req1 = make_request("1", [0, 0, 1, 1, 2, 2, 3, 3], hash_block_size, sha256) + _, num_computed, _ = manager.get_computed_blocks(req1) + assert num_computed == 6 + + +def test_kpool_tail_supports_128_token_partial_hash_hits(): + hash_block_size = 128 + cache_block_size = 9 * hash_block_size + kpool_block_size = 64 + kv_cache_config = KVCacheConfig( + num_blocks=24, + kv_cache_tensors=[], + kv_cache_groups=[ + KVCacheGroupSpec( + ["full"], + FullAttentionSpec( + block_size=cache_block_size, + num_kv_heads=1, + head_size=1, + dtype=torch.float32, + ), + ), + KVCacheGroupSpec( + ["mamba"], + MambaSpec( + block_size=cache_block_size, + shapes=(1, 1), + dtypes=(torch.float32,), + mamba_cache_mode="align", + ), + ), + KVCacheGroupSpec( + ["tail"], + KpoolTailSpec( + block_size=kpool_block_size, + num_kv_heads=1, + head_size=1, + dtype=torch.float32, + sliding_window=kpool_block_size, + ), + ), + ], + ) + manager = make_kv_cache_manager( + kv_cache_config=kv_cache_config, + max_model_len=8192, + enable_caching=True, + hash_block_size=hash_block_size, + ) + + assert manager.coordinator.scheduler_block_size == cache_block_size + assert manager.coordinator.hash_block_size == hash_block_size + assert manager.coordinator.enable_partial_hash_hits + + shared_prefix = [10] * (12 * hash_block_size) + req0 = make_request("0", shared_prefix, hash_block_size, sha256) + computed_blocks, num_computed, _ = manager.get_computed_blocks(req0) + assert ( + manager.allocate_slots(req0, len(shared_prefix), num_computed, computed_blocks) + is not None + ) + manager.free(req0) + manager.new_step_starts() + + req1 = make_request("1", shared_prefix + [12] * 64, hash_block_size, sha256) + _, num_computed, _ = manager.get_computed_blocks(req1) + assert num_computed == len(shared_prefix) + + @pytest.mark.parametrize("dcp_world_size", [1, 2, 4]) def test_hybrid_partial_hash_hit_uses_cow_under_dcp(dcp_world_size: int): hash_block_size = 2 diff --git a/tests/v1/core/test_kv_cache_utils.py b/tests/v1/core/test_kv_cache_utils.py index c5159e0d1627..438b8753f8b5 100644 --- a/tests/v1/core/test_kv_cache_utils.py +++ b/tests/v1/core/test_kv_cache_utils.py @@ -6,13 +6,19 @@ from collections.abc import Callable from dataclasses import replace from types import SimpleNamespace -from typing import Any +from typing import Any, cast import pytest import torch import vllm.v1.core.kv_cache_utils as kv_cache_utils -from vllm.config import CacheConfig, ModelConfig, SchedulerConfig, VllmConfig +from vllm.config import ( + CacheConfig, + KVTransferConfig, + ModelConfig, + SchedulerConfig, + VllmConfig, +) from vllm.config.kv_events import KVEventsConfig from vllm.lora.request import LoRARequest from vllm.multimodal.inputs import ( @@ -47,6 +53,7 @@ ChunkedLocalAttentionSpec, FullAttentionSpec, HiddenStateCacheSpec, + KpoolTailSpec, KVCacheConfig, KVCacheGroupSpec, KVCacheSpec, @@ -2191,34 +2198,567 @@ def test_generate_scheduler_kv_cache_config(): ) -def test_mixed_precision_kv_cache_with_uniform_type_specs(): - fp8_spec = new_kv_cache_spec(dtype=torch.float8_e4m3fn) - bf16_spec = new_kv_cache_spec(dtype=torch.bfloat16) - worker_config = KVCacheConfig( - num_blocks=10, +def _glm5_like_kv_cache_spec( + mamba_spec_factory=new_mamba_spec, +) -> tuple[dict[str, KVCacheSpec], list[str]]: + """(mamba, mamba, mamba, MLA + indexer) * 11 plus a trailing mamba.""" + kv_cache_spec: dict[str, KVCacheSpec] = {} + mamba_layers = [] + for i in range(45): + if i % 4 == 3: + kv_cache_spec[f"layers.{i}.attn"] = MLAAttentionSpec( + block_size=1024, + num_kv_heads=1, + head_size=576, + dtype=torch.bfloat16, + ) + kv_cache_spec[f"layers.{i}.indexer"] = MLAAttentionSpec( + block_size=1024, + num_kv_heads=1, + head_size=132, + dtype=torch.uint8, + tokens_per_state=16, + ) + else: + name = f"layers.{i}.linear_attn" + kv_cache_spec[name] = mamba_spec_factory() + mamba_layers.append(name) + return kv_cache_spec, mamba_layers + + +def _glm5_like_kv_cache_spec_with_tail( + mamba_spec_factory=new_mamba_spec, +) -> dict[str, KVCacheSpec]: + """Production kpool=4 proportions: (mamba*3, MLA + indexer + tail) * 11. + + The tail's logical page (2 * kpool * 2*indexer_head_dim * bf16) must fit + inside the indexer page it parasitizes (block_size//kpool * 132 B), which + is why the fixture drops the tokens_per_state=16 shape used above. + """ + kv_cache_spec, _ = _glm5_like_kv_cache_spec(mamba_spec_factory) + for i in range(3, 45, 4): + kpool = 4 + kv_cache_spec[f"layers.{i}.indexer"] = replace( + cast(MLAAttentionSpec, kv_cache_spec[f"layers.{i}.indexer"]), + tokens_per_state=kpool, + ) + kv_cache_spec[f"layers.{i}.tail"] = KpoolTailSpec( + block_size=kpool, + num_kv_heads=2, + head_size=128, + head_size_v=0, + dtype=torch.bfloat16, + sliding_window=kpool, + ) + return kv_cache_spec + + +def _tensor_by_layer(kv_cache_config: KVCacheConfig) -> dict[str, KVCacheTensor]: + return { + layer_name: tensor + for tensor in kv_cache_config.kv_cache_tensors + for layer_name in tensor.layers + } + + +def _layer_offset(tensor: KVCacheTensor, layer_name: str) -> int: + return tensor.offset + tensor.layers.index(layer_name) * tensor.layer_stride + + +def test_get_kv_cache_config_balanced_mamba_hybrid(): + """Hybrid slot sharing: mamba layers co-own the MLA slot tensors.""" + model_config = ModelConfig(max_model_len=8192) + vllm_config = VllmConfig(model_config=model_config) + + kv_cache_spec, mamba_layers = _glm5_like_kv_cache_spec() + mla_page = kv_cache_spec["layers.3.attn"].page_size_bytes + idx_page = kv_cache_spec["layers.3.indexer"].page_size_bytes + + groups = kv_cache_utils.get_kv_cache_groups(vllm_config, kv_cache_spec) + uniform_groups = [ + group + for group in groups + if isinstance(group.kv_cache_spec, UniformTypeKVCacheSpecs) + ] + mamba_groups = [ + group for group in groups if isinstance(group.kv_cache_spec, MambaSpec) + ] + + # Round-robin into G = ceil(34 / 11) = 4 groups so every mamba layer + # gets an MLA slot. + assert len(groups) == 5 + assert len(uniform_groups) == 1 + assert len(uniform_groups[0].layer_names) == 22 + assert uniform_groups[0].kv_cache_spec.get_max_layers_per_page_size() == 11 + assert [len(group.layer_names) for group in mamba_groups] == [9, 9, 8, 8] + for k, name in enumerate(mamba_layers): + assert name in mamba_groups[k % 4].layer_names + # Mamba pages are padded up to the MLA page. + for group in mamba_groups: + assert group.kv_cache_spec.page_size_padded == mla_page + assert group.kv_cache_spec.page_size_bytes == mla_page + + bytes_per_block = kv_cache_utils._pool_bytes_per_block(groups) + assert bytes_per_block == 11 * mla_page + 11 * idx_page + + # Every block id is charged the full per-block sum. + attn_blocks = uniform_groups[0].kv_cache_spec.max_memory_usage_pages(vllm_config) + mamba_blocks_per_group = 1 + new_mamba_spec().num_speculative_blocks + blocks_per_request = attn_blocks + 4 * mamba_blocks_per_group + assert ( + kv_cache_utils._max_memory_usage_bytes_from_groups(vllm_config, groups) + == blocks_per_request * bytes_per_block + ) + + available_memory = bytes_per_block * 100 + 1 + kv_cache_config = kv_cache_utils.get_kv_cache_config_from_groups( + vllm_config, groups, available_memory + ) + assert kv_cache_config.num_blocks == 100 + + # Every logical layer has a view into one backing allocation. Mamba views + # alias their corresponding MLA slot by using the same byte offset. + assert len(kv_cache_config.kv_cache_tensors) == 56 + assert {t.size for t in kv_cache_config.kv_cache_tensors} == {bytes_per_block * 100} + tensors = _tensor_by_layer(kv_cache_config) + mla_layer_names = [f"layers.{i}.attn" for i in range(3, 45, 4)] + for i, mla_name in enumerate(mla_layer_names): + mla_tensor = tensors[mla_name] + assert mla_tensor.block_stride == mla_page + for group in mamba_groups: + if i < len(group.layer_names): + assert tensors[group.layer_names[i]].offset == mla_tensor.offset + for i in range(11): + assert tensors[f"layers.{4 * i + 3}.indexer"].block_stride == idx_page + + total_allocated = next(iter({t.size for t in kv_cache_config.kv_cache_tensors})) + assert total_allocated == bytes_per_block * 100 + assert 0 <= available_memory - total_allocated < bytes_per_block + + assert get_max_concurrency_for_kv_cache_config( + vllm_config, kv_cache_config + ) == pytest.approx(100 / blocks_per_request) + + +def test_get_kv_cache_config_kpool_tail_coowns_indexer_tensor(): + """The kpool tail parasitizes the indexer tensors instead of getting its + own: sibling idx/tail tensors paired by layer order, zero standalone tail + bytes, one shared block per request, and no prefix-caching leakage from + the kpool-sized scratch group.""" + model_config = ModelConfig(max_model_len=8192) + vllm_config = VllmConfig(model_config=model_config) + + kv_cache_spec = _glm5_like_kv_cache_spec_with_tail() + mla_page = kv_cache_spec["layers.3.attn"].page_size_bytes + idx_page = kv_cache_spec["layers.3.indexer"].page_size_bytes + tail_logical_page = kv_cache_spec["layers.3.tail"].page_size_bytes + assert tail_logical_page < idx_page + + groups = kv_cache_utils.get_kv_cache_groups(vllm_config, kv_cache_spec) + tail_group = next( + group + for group in groups + if isinstance(group.kv_cache_spec, UniformTypeKVCacheSpecs) + and all( + isinstance(spec, KpoolTailSpec) + for spec in group.kv_cache_spec.kv_cache_specs.values() + ) + ) + # The tail never prefix-caches; its page is padded up to the indexer page + # so the runner's strided view rides the indexer storage. + assert not tail_group.kv_cache_spec.prefix_cacheable + tail_inner = cast( + KpoolTailSpec, tail_group.kv_cache_spec.kv_cache_specs["layers.3.tail"] + ) + assert tail_inner.page_size_padded == idx_page + + bytes_per_block = kv_cache_utils._pool_bytes_per_block(groups) + assert bytes_per_block == 11 * mla_page + 11 * idx_page + + kv_cache_config = kv_cache_utils.get_kv_cache_config_from_groups( + vllm_config, groups, bytes_per_block * 100 + 1 + ) + assert kv_cache_config.num_blocks == 100 + # Tail layers get logical views but no additional storage: each view aliases + # its sibling indexer at the same offset in the shared allocation. + assert len(kv_cache_config.kv_cache_tensors) == 67 + tensors = _tensor_by_layer(kv_cache_config) + for i in range(11): + idx_name = f"layers.{4 * i + 3}.indexer" + tail_name = f"layers.{4 * i + 3}.tail" + assert tensors[idx_name].offset == tensors[tail_name].offset + assert {t.size for t in kv_cache_config.kv_cache_tensors} == {bytes_per_block * 100} + + # The layout detector surfaces the sibling names in layer order for the + # accounting and connector paths. + layout = kv_cache_utils._glm5_next_tensor_layout(kv_cache_config.kv_cache_groups) + assert layout is not None + assert layout[3] == [f"layers.{4 * i + 3}.indexer" for i in range(11)] + assert layout[6] == [f"layers.{4 * i + 3}.tail" for i in range(11)] + + # Accounting charges exactly one extra shared block per request for the + # tail, not a per-sequence allocation. + attn_group = next( + group + for group in groups + if isinstance(group.kv_cache_spec, UniformTypeKVCacheSpecs) + and group is not tail_group + ) + attn_blocks = attn_group.kv_cache_spec.max_memory_usage_pages(vllm_config) + mamba_blocks_per_group = 1 + new_mamba_spec().num_speculative_blocks + blocks_per_request = attn_blocks + 4 * mamba_blocks_per_group + 1 + assert ( + kv_cache_utils._max_memory_usage_bytes_from_groups(vllm_config, groups) + == blocks_per_request * bytes_per_block + ) + + +def test_glm5_kpool_tail_does_not_drag_hash_block_size(): + """The tail's kpool-sized scratch block (4 tokens) must not constrain the + prefix-cache hash granularity: participating groups alone decide it.""" + model_config = ModelConfig(max_model_len=8192) + vllm_config = VllmConfig(model_config=model_config) + + def align_mamba(): + return new_mamba_spec(mamba_cache_mode="align") + + groups = kv_cache_utils.get_kv_cache_groups( + vllm_config, _glm5_like_kv_cache_spec_with_tail(align_mamba) + ) + kv_cache_config = KVCacheConfig( + num_blocks=1, kv_cache_tensors=[], - kv_cache_groups=[ - KVCacheGroupSpec( - ["fp8_layer"], - UniformTypeKVCacheSpecs( - block_size=16, kv_cache_specs={"fp8_layer": fp8_spec} - ), - ), - KVCacheGroupSpec( - ["bf16_layer"], - UniformTypeKVCacheSpecs( - block_size=16, kv_cache_specs={"bf16_layer": bf16_spec} - ), - ), - ], + kv_cache_groups=groups, + ) + hash_vllm_config = SimpleNamespace( + cache_config=SimpleNamespace( + block_size=16, + enable_prefix_caching=True, + prefix_match_unit=None, + ), + parallel_config=SimpleNamespace(decode_context_parallel_size=1), + kv_transfer_config=object(), + ) + # gcd(attn 1024, mamba 16) with the tail's 4 excluded; scheduler size is + # the lcm including the tail (1024 % 4 == 0, so it coincides). + assert kv_cache_utils.resolve_kv_cache_block_sizes( + kv_cache_config, hash_vllm_config + ) == (1024, 16) + + +def test_get_kv_cache_config_mamba_hybrid_sharing_infeasible(): + """Reject GLM-5.3-Flash layouts whose Mamba page exceeds the MLA page.""" + model_config = ModelConfig(max_model_len=8192) + vllm_config = VllmConfig(model_config=model_config) + + # (512, 1024) fp32 state = 2 MiB per page > 1.125 MiB MLA page. + def big_mamba_spec(): + return new_mamba_spec(shapes=((512, 1024),), dtypes=(torch.float32,)) + + kv_cache_spec, _ = _glm5_like_kv_cache_spec(big_mamba_spec) + mla_page = kv_cache_spec["layers.3.attn"].page_size_bytes + assert big_mamba_spec().page_size_bytes > mla_page + + with pytest.raises(ValueError, match="does not fit the MLA page"): + kv_cache_utils.get_kv_cache_groups(vllm_config, kv_cache_spec) + + +def test_get_kv_cache_config_mamba_hybrid_sharing_infeasible_no_indexer(): + """Use the generic layout error when no kpool indexer is present.""" + model_config = ModelConfig(max_model_len=8192) + vllm_config = VllmConfig(model_config=model_config) + vllm_config.cache_config.kv_cache_layout = "LBHNC" + + kv_cache_spec: dict[str, KVCacheSpec] = {} + for i in range(27): + if i % 4 == 3 or i == 26: + kv_cache_spec[f"layers.{i}.attn"] = MLAAttentionSpec( + block_size=1024, + num_kv_heads=1, + head_size=576, + dtype=torch.bfloat16, + ) + else: + kv_cache_spec[f"layers.{i}.linear_attn"] = new_mamba_spec( + shapes=((512, 1024),), dtypes=(torch.float32,) + ) + mla_page = kv_cache_spec["layers.3.attn"].page_size_bytes + assert kv_cache_spec["layers.0.linear_attn"].page_size_bytes > mla_page + + with pytest.raises(NotImplementedError, match="page size"): + kv_cache_utils.get_kv_cache_groups(vllm_config, kv_cache_spec) + + +def test_get_kv_cache_config_mamba_hybrid_sharing_prepadded_mamba(): + """Platform-prepadded mamba pages (mamba_page_size_padded hint) must not + disable slot sharing; the layout re-pads them to the MLA page.""" + model_config = ModelConfig(max_model_len=8192) + vllm_config = VllmConfig(model_config=model_config) + + def prepadded_mamba_spec(): + return new_mamba_spec(page_size_padded=294_912) + + kv_cache_spec, _ = _glm5_like_kv_cache_spec(prepadded_mamba_spec) + mla_page = kv_cache_spec["layers.3.attn"].page_size_bytes + assert prepadded_mamba_spec().page_size_bytes < mla_page + + groups = kv_cache_utils.get_kv_cache_groups(vllm_config, kv_cache_spec) + mamba_groups = [ + group for group in groups if isinstance(group.kv_cache_spec, MambaSpec) + ] + assert len(groups) == 5 + assert [len(group.layer_names) for group in mamba_groups] == [9, 9, 8, 8] + for group in mamba_groups: + assert group.kv_cache_spec.page_size_bytes == mla_page + + bytes_per_block = kv_cache_utils._pool_bytes_per_block(groups) + kv_cache_config = kv_cache_utils.get_kv_cache_config_from_groups( + vllm_config, groups, bytes_per_block * 100 + 1 + ) + assert kv_cache_config.num_blocks == 100 + assert len(kv_cache_config.kv_cache_tensors) == 56 + + +def test_get_kv_cache_config_mamba_hybrid_sharing_pp_balanced_projection(): + """Round-robin mamba grouping keeps practical PP splits balanced: every + stage's largest projected mamba group slice fits its projected MLA + layers, so per-stage slot tensors all have an MLA owner.""" + model_config = ModelConfig(max_model_len=8192) + vllm_config = VllmConfig(model_config=model_config) + + kv_cache_spec, _ = _glm5_like_kv_cache_spec() + global_groups = kv_cache_utils.get_kv_cache_groups(vllm_config, kv_cache_spec) + mla_page = kv_cache_spec["layers.3.attn"].page_size_bytes + idx_page = kv_cache_spec["layers.3.indexer"].page_size_bytes + + # PP=2 split at transformer layer 22/23. + # Stage 0 of a PP=2 split at transformer layer 22/23: 5 MLA(+indexer) + # layers, projected mamba groups [5, 5, 4, 4]. + worker_spec = {n: s for n, s in kv_cache_spec.items() if int(n.split(".")[1]) <= 22} + groups = kv_cache_utils._project_kv_cache_groups_to_worker( + global_groups, worker_spec + ) + mamba_groups = [ + group for group in groups if isinstance(group.kv_cache_spec, MambaSpec) + ] + assert [len(group.layer_names) for group in mamba_groups] == [5, 5, 4, 4] + + bytes_per_block = kv_cache_utils._pool_bytes_per_block(groups) + assert bytes_per_block == 5 * mla_page + 5 * idx_page + + kv_cache_config = kv_cache_utils.get_kv_cache_config_from_groups( + vllm_config, groups, bytes_per_block * 100 + 1 + ) + assert kv_cache_config.num_blocks == 100 + tensors = _tensor_by_layer(kv_cache_config) + # Every projected Mamba slot aliases an MLA owner. + assert tensors["layers.3.attn"].offset == tensors["layers.0.linear_attn"].offset + assert tensors["layers.3.attn"].offset == tensors["layers.1.linear_attn"].offset + assert tensors["layers.3.attn"].offset == tensors["layers.2.linear_attn"].offset + assert tensors["layers.3.attn"].offset == tensors["layers.4.linear_attn"].offset + # The last slot only hosts the two groups with a 5th projected layer. + assert tensors["layers.19.attn"].offset == tensors["layers.21.linear_attn"].offset + assert tensors["layers.19.attn"].offset == tensors["layers.22.linear_attn"].offset + assert {t.size for t in kv_cache_config.kv_cache_tensors} == {bytes_per_block * 100} + + +def test_get_kv_cache_config_mamba_hybrid_sharing_pp_group_count_bump(monkeypatch): + """PP=4's default partition [11,11,12,11] leaves stage 0 with 2 MLA + layers but a round-robin slice of 3 under the minimum 4 mamba groups; + grouping bumps to 5 groups so every stage's projection keeps sharing + on instead of silently falling back on stage 0.""" + from vllm.config import ParallelConfig + from vllm.distributed.utils import get_pp_indices + + monkeypatch.setattr(ModelConfig, "get_total_num_hidden_layers", lambda self: 45) + vllm_config = VllmConfig( + model_config=ModelConfig(max_model_len=8192), + parallel_config=ParallelConfig(pipeline_parallel_size=4), + ) + + kv_cache_spec, mamba_layers = _glm5_like_kv_cache_spec() + groups = kv_cache_utils.get_kv_cache_groups(vllm_config, kv_cache_spec) + mamba_groups = [g for g in groups if isinstance(g.kv_cache_spec, MambaSpec)] + assert [len(g.layer_names) for g in mamba_groups] == [7, 7, 7, 7, 6] + for k, name in enumerate(mamba_layers): + assert name in mamba_groups[k % 5].layer_names + + for rank in range(4): + start, end = get_pp_indices(45, rank, 4) + worker_spec = { + n: s + for n, s in kv_cache_spec.items() + if start <= int(n.split(".")[1]) < end + } + projected = kv_cache_utils._project_kv_cache_groups_to_worker( + groups, worker_spec + ) + assert kv_cache_utils._glm5_next_tensor_layout(projected) is not None + + +def test_get_kv_cache_config_mamba_hybrid_sharing_pp_starved_stage(monkeypatch): + """Reject PP stages whose Mamba layers have no MLA slot to share.""" + from vllm.config import ParallelConfig + + monkeypatch.setattr(ModelConfig, "get_total_num_hidden_layers", lambda self: 45) + monkeypatch.setenv("VLLM_PP_LAYER_PARTITION", "3,42") + vllm_config = VllmConfig( + model_config=ModelConfig(max_model_len=8192), + parallel_config=ParallelConfig(pipeline_parallel_size=2), ) - scheduler_config = generate_scheduler_kv_cache_config([worker_config]) - assert worker_config.needs_kv_cache_zeroing - assert scheduler_config.needs_kv_cache_zeroing + kv_cache_spec, _ = _glm5_like_kv_cache_spec() + with pytest.raises(ValueError, match="VLLM_PP_LAYER_PARTITION"): + kv_cache_utils.get_kv_cache_groups(vllm_config, kv_cache_spec) + + +def test_get_kv_cache_config_mamba_hybrid_sharing_beats_cross_layers_flag(): + """Hybrid slot sharing must take precedence over the experimental + enable_cross_layers_blocks packed layout: generic packing would give the + MLA slots a strided view, breaking the contiguous virtual split.""" + model_config = ModelConfig(max_model_len=8192) + kv_transfer_config = KVTransferConfig( + kv_connector="NixlConnector", + kv_role="kv_both", + kv_connector_extra_config={"enable_cross_layers_blocks": "true"}, + ) + vllm_config = VllmConfig( + model_config=model_config, kv_transfer_config=kv_transfer_config + ) + + kv_cache_spec, _ = _glm5_like_kv_cache_spec() + groups = kv_cache_utils.get_kv_cache_groups(vllm_config, kv_cache_spec) + + mla_page = kv_cache_spec["layers.3.attn"].page_size_bytes + idx_page = kv_cache_spec["layers.3.indexer"].page_size_bytes + bytes_per_block = kv_cache_utils._pool_bytes_per_block(groups) + assert bytes_per_block == 11 * mla_page + 11 * idx_page + + kv_cache_config = kv_cache_utils.get_kv_cache_config_from_groups( + vllm_config, groups, bytes_per_block * 100 + 1 + ) + assert kv_cache_config.num_blocks == 100 + assert len(kv_cache_config.kv_cache_tensors) == 56 + assert all( + tensor.block_stride in (mla_page, idx_page) + for tensor in kv_cache_config.kv_cache_tensors + ) + assert {t.size for t in kv_cache_config.kv_cache_tensors} == {bytes_per_block * 100} + + +def test_get_kv_cache_config_mamba_hybrid_sharing_no_indexer(): + """Kimi-Linear-like: MLA without indexer layers, idx_stride == 0.""" + model_config = ModelConfig(max_model_len=8192) + vllm_config = VllmConfig(model_config=model_config) + vllm_config.cache_config.kv_cache_layout = "LBNHC" + + # 20 mamba + 7 MLA layers -> G = ceil(20 / 7) = 3 groups of [7, 7, 6]. + kv_cache_spec: dict[str, KVCacheSpec] = {} + for i in range(27): + if i % 4 == 3 or i == 26: + kv_cache_spec[f"layers.{i}.attn"] = MLAAttentionSpec( + block_size=1024, + num_kv_heads=1, + head_size=576, + dtype=torch.bfloat16, + ) + else: + kv_cache_spec[f"layers.{i}.linear_attn"] = new_mamba_spec() + mla_page = kv_cache_spec["layers.3.attn"].page_size_bytes + + groups = kv_cache_utils.get_kv_cache_groups(vllm_config, kv_cache_spec) + mamba_groups = [ + group for group in groups if isinstance(group.kv_cache_spec, MambaSpec) + ] + assert len(groups) == 4 + assert [len(group.layer_names) for group in mamba_groups] == [7, 7, 6] + for group in mamba_groups: + assert group.kv_cache_spec.page_size_bytes == mla_page + + bytes_per_block = kv_cache_utils._pool_bytes_per_block(groups) + assert bytes_per_block == 7 * mla_page + + kv_cache_config = kv_cache_utils.get_kv_cache_config_from_groups( + vllm_config, groups, bytes_per_block * 100 + 1 + ) + assert kv_cache_config.num_blocks == 100 + # Each group has one tensor descriptor and all groups alias the same slots. + assert len(kv_cache_config.kv_cache_tensors) == 4 + tensors = _tensor_by_layer(kv_cache_config) + mla_names = [f"layers.{i}.attn" for i in (*range(3, 27, 4), 26)] + for index, mla_name in enumerate(mla_names): + mla_offset = _layer_offset(tensors[mla_name], mla_name) + for group in mamba_groups: + if index < len(group.layer_names): + mamba_name = group.layer_names[index] + assert _layer_offset(tensors[mamba_name], mamba_name) == mla_offset + assert {t.size for t in kv_cache_config.kv_cache_tensors} == {bytes_per_block * 100} + + +def test_get_kv_cache_capacity_after_scheduler_unwrap(): + """max_concurrency must survive the scheduler-config unwrap. + + Regression for the balanced-mamba hybrid layout: + ``generate_scheduler_kv_cache_config`` flattens the MLA + ``UniformTypeKVCacheSpecs`` group into a single ``MLAAttentionSpec``, so the + scheduler config holds an MLA spec (page size A) next to ``MambaSpec`` + groups (page size B). ``EngineCore._initialize_kv_caches`` calls + ``get_kv_cache_capacity`` on that unwrapped config; the uniform-page-size + path used to assert-fail (``assert len(page_sizes) == 1``) on this topology. + """ + model_config = ModelConfig(max_model_len=8192) + vllm_config = VllmConfig(model_config=model_config) + + kv_cache_spec: dict[str, KVCacheSpec] = {} + for i in range(45): + if i % 4 == 3: + kv_cache_spec[f"layers.{i}.attn"] = MLAAttentionSpec( + block_size=1024, + num_kv_heads=1, + head_size=576, + dtype=torch.bfloat16, + ) + kv_cache_spec[f"layers.{i}.indexer"] = MLAAttentionSpec( + block_size=1024, + num_kv_heads=1, + head_size=132, + dtype=torch.uint8, + tokens_per_state=16, + ) + else: + kv_cache_spec[f"layers.{i}.linear_attn"] = new_mamba_spec() + + groups = kv_cache_utils.get_kv_cache_groups(vllm_config, kv_cache_spec) + bytes_per_block = kv_cache_utils._pool_bytes_per_block(groups) + kv_cache_config = kv_cache_utils.get_kv_cache_config_from_groups( + vllm_config, groups, bytes_per_block * 100 + 1 + ) + + scheduler_config = generate_scheduler_kv_cache_config([kv_cache_config]) + # Confirm we are exercising the regression path: the MLA group is no longer + # UniformTypeKVCacheSpecs, so sibling specs have differing page sizes. + assert not any( + isinstance(g.kv_cache_spec, UniformTypeKVCacheSpecs) + for g in scheduler_config.kv_cache_groups + ) + + unwrapped_groups = scheduler_config.kv_cache_groups + expected_max_mem = kv_cache_utils._max_memory_usage_bytes_from_groups( + vllm_config, unwrapped_groups + ) + expected_pool = kv_cache_utils._pool_bytes_per_block(unwrapped_groups) + expected_blocks_per_request = ( + expected_max_mem + expected_pool - 1 + ) // expected_pool + + _, max_concurrency = get_kv_cache_capacity(vllm_config, scheduler_config) + assert max_concurrency > 0 + assert max_concurrency == pytest.approx( + scheduler_config.num_blocks / expected_blocks_per_request + ) -def new_mla_spec(cache_dtype_str=None, block_size=16): +def new_mla_spec(cache_dtype_str=None, block_size: int = 16): # head_size = kv_lora_rank(512) + qk_rope_head_dim(64) = 576 return MLAAttentionSpec( block_size=block_size, @@ -3305,9 +3845,10 @@ def _hybrid_specs_with_draft(draft: bool, draft_shares_target_spec: bool = False """A K3-shaped hybrid: MLA full attention + Mamba, optionally plus a DSpark-style draft MLA layer marked non_causal_multi_token_decode. - The target's fp8 KV dtype is what keeps the draft in its own bucket, as it - does on Kimi-K3 (target `--kv-cache-dtype fp8_e4m3`, draft `auto`). Pass - draft_shares_target_spec to collapse them into one group instead. + The target's fp8 KV dtype is one difference from the draft on Kimi-K3 + (target `--kv-cache-dtype fp8_e4m3`, draft `auto`). Pass + draft_shares_target_spec to match that dtype while retaining the draft + marker as a group boundary. """ target_dtype = None if draft_shares_target_spec else "fp8_e4m3" specs = { diff --git a/tests/v1/kv_connector/unit/test_mooncake_store_hma_e2e.py b/tests/v1/kv_connector/unit/test_mooncake_store_hma_e2e.py index 184115e5d5e0..d6b0becc1d2c 100644 --- a/tests/v1/kv_connector/unit/test_mooncake_store_hma_e2e.py +++ b/tests/v1/kv_connector/unit/test_mooncake_store_hma_e2e.py @@ -30,6 +30,7 @@ from vllm.v1.core.kv_cache_utils import BlockHash from vllm.v1.kv_cache_interface import ( FullAttentionSpec, + KpoolTailSpec, KVCacheConfig, KVCacheGroupSpec, KVCacheTensor, @@ -605,3 +606,195 @@ def batch_put_from_multi_buffers(self, keys, addrs, sizes, *a, **k): token_dbs[1].key_for(hs[2]): [10_000 + mamba_cow_block * 512], } assert store.puts == expected + + +def test_worker_lookup_hits_sub_block_partial_tail(): + """worker.lookup must query sub-block keys when partial hash hits are on. + + Regression test: ``lookup`` hard-coded ``fine_grained = False``, so the + sub-block keys persisted by ``_sub_block_tail_puts`` were never probed and + partial prefix hits silently returned 0 while the store side kept writing + them. Here the mamba block (16) exceeds the hash unit (4), so + ``enable_partial_hash_hits`` is on and the lookup must find the stored + boundary at 12. + """ + full = FullAttentionSpec(block_size=16, num_kv_heads=8, head_size=64, dtype=None) + mamba = MambaSpec( + block_size=16, + shapes=((1, 1),), + dtypes=(torch.float32,), + mamba_cache_mode="align", + ) + cfg = KVCacheConfig( + num_blocks=4, + kv_cache_tensors=[ + KVCacheTensor( + size=4 * full.page_size_bytes, + layers=["L0"], + layer_stride=4 * full.page_size_bytes, + block_stride=full.page_size_bytes, + ), + KVCacheTensor( + size=4 * mamba.page_size_bytes, + layers=["L1"], + layer_stride=4 * mamba.page_size_bytes, + block_stride=mamba.page_size_bytes, + ), + ], + kv_cache_groups=[ + KVCacheGroupSpec(["L0"], full), + KVCacheGroupSpec(["L1"], mamba), + ], + ) + vllm_config = _minimal_vllm_config(cache_block_size=16) + # Hash unit 4 < mamba block 16 -> partial hash hits are enabled. + vllm_config.cache_config.prefix_match_unit = 4 + store = _DictStore() + + worker = _build_worker_with_dict_store(vllm_config, cfg, store) + worker.tp_size = 1 + worker.pp_size = 1 + worker.num_kv_head = 8 + assert worker.coord.enable_partial_hash_hits + + for g_idx, db in enumerate(worker.token_dbs): + db.set_kv_caches_base_addr([g_idx * 10_000]) + db.set_block_len([512]) + + send_thread = KVCacheStoreSendingThread( + store=store, + token_databases=worker.token_dbs, + block_size=worker.block_size, + coord=worker.coord, + tp_rank=0, + group_put_steps=[1, 1], + kv_role="kv_both", + ready_event=threading.Event(), + replicate_config=MagicMock(), + ) + + # Persist the sub-block partial tail at boundary 12 (keyed by hs[12//4-1]). + hs = [BlockHash(bytes([i + 1]) * 4) for i in range(5)] + req = ReqMeta( + req_id="r0", + token_len_chunk=0, + block_ids=([1], [2]), + block_hashes=hs, + can_save=True, + num_prompt_tokens=20, + boundary_state_offloads=[(1, 7, 12)], + ) + send_thread._maybe_offload_boundary_states(req) + + worker.store = store + + # A 13-token prompt sharing the prefix must hit the stored boundary at 12. + assert worker.lookup(num_tokens=13, block_hashes=hs).hit_length == 12 + + +def test_worker_setup_tolerates_finer_scratch_group(): + """Setup and lookup must tolerate a non-prefix-cacheable scratch group. + + Regression: GLM-5.3-Flash carries a kpool-tail scratch group whose block + size (``index_kpool`` tokens) is finer than the hash unit, so the + coordinator's divisibility assert and the scratch group's + ``ChunkedTokenDatabase`` both rejected worker setup for any hash unit the + scratch block does not divide. Scratch groups never participate in + store/load/lookup, so setup must skip them and lookup must still hit + stored sub-block boundaries on the participating groups. + """ + full = FullAttentionSpec(block_size=16, num_kv_heads=8, head_size=64, dtype=None) + mamba = MambaSpec( + block_size=16, + shapes=((1, 1),), + dtypes=(torch.float32,), + mamba_cache_mode="align", + ) + # Scratch block 4 is not divisible by the hash unit 8 below. + scratch = KpoolTailSpec( + block_size=4, + num_kv_heads=2, + head_size=64, + head_size_v=0, + dtype=torch.bfloat16, + sliding_window=4, + ) + cfg = KVCacheConfig( + num_blocks=4, + kv_cache_tensors=[ + KVCacheTensor( + size=4 * full.page_size_bytes, + layers=["L0"], + layer_stride=4 * full.page_size_bytes, + block_stride=full.page_size_bytes, + ), + KVCacheTensor( + size=4 * mamba.page_size_bytes, + layers=["L1"], + layer_stride=4 * mamba.page_size_bytes, + block_stride=mamba.page_size_bytes, + ), + KVCacheTensor( + size=4 * scratch.page_size_bytes, + layers=["L2"], + layer_stride=4 * scratch.page_size_bytes, + block_stride=scratch.page_size_bytes, + ), + ], + kv_cache_groups=[ + KVCacheGroupSpec(["L0"], full), + KVCacheGroupSpec(["L1"], mamba), + KVCacheGroupSpec(["L2"], scratch), + ], + ) + vllm_config = _minimal_vllm_config(cache_block_size=16) + # Hash unit 8 divides the participating groups (16) but not the scratch + # group (4); mamba block 16 > 8 keeps partial hash hits on. + vllm_config.cache_config.prefix_match_unit = 8 + store = _DictStore() + + worker = _build_worker_with_dict_store(vllm_config, cfg, store) + worker.tp_size = 1 + worker.pp_size = 1 + worker.num_kv_head = 8 + assert worker.coord.enable_partial_hash_hits + # The scratch DB is keyed at its own block size and never probed. + assert worker.token_dbs[2].hash_block_size == 4 + + for g_idx, db in enumerate(worker.token_dbs): + db.set_kv_caches_base_addr([g_idx * 10_000]) + db.set_block_len([512]) + + send_thread = KVCacheStoreSendingThread( + store=store, + token_databases=worker.token_dbs, + block_size=worker.block_size, + coord=worker.coord, + tp_rank=0, + group_put_steps=[1, 1, 1], + kv_role="kv_both", + ready_event=threading.Event(), + replicate_config=MagicMock(), + group_participates=[True, True, False], + ) + + # Persist the sub-block partial tail at boundary 12 (keyed by hs[12//8-1]). + hs = [BlockHash(bytes([i + 1]) * 8) for i in range(3)] + req = ReqMeta( + req_id="r0", + token_len_chunk=0, + block_ids=([1], [2], [3]), + block_hashes=hs, + can_save=True, + num_prompt_tokens=20, + boundary_state_offloads=[(1, 7, 12)], + ) + send_thread._maybe_offload_boundary_states(req) + + worker.store = store + + # A 13-token prompt sharing the prefix must hit the first hash unit. + assert worker.lookup(num_tokens=13, block_hashes=hs).hit_length == 8 + # The scratch group's namespace never enters the store. + scratch_prefix = worker.token_dbs[2].key_for(hs[0]).rsplit("@", 1)[0] + assert not any(key.startswith(scratch_prefix) for key in store._data) diff --git a/tests/v1/kv_connector/unit/test_nixl_desc_geometry.py b/tests/v1/kv_connector/unit/test_nixl_desc_geometry.py index 561457157d44..712a9d30c08d 100644 --- a/tests/v1/kv_connector/unit/test_nixl_desc_geometry.py +++ b/tests/v1/kv_connector/unit/test_nixl_desc_geometry.py @@ -211,6 +211,167 @@ def _make_mla_hybrid_worker(local_block_size, kernel_block_size, num_logical_blo return worker +@pytest.mark.cpu_test +@pytest.mark.parametrize("logical_block_size", [1152, 640]) +@pytest.mark.parametrize("tail_first", [False, True]) +def test_register_compressed_indexer_uses_virtual_transfer_pages( + logical_block_size, tail_first +): + """Compressed indexer rows must split into contiguous NIXL transfer pages.""" + from unittest.mock import MagicMock + + from vllm.config import set_current_vllm_config + from vllm.distributed.kv_transfer.kv_connector.v1.nixl import ( + base_worker as bw, + ) + from vllm.distributed.kv_transfer.kv_connector.v1.nixl.worker import ( + NixlConnectorWorker, + ) + from vllm.v1.kv_cache_interface import ( + KpoolTailSpec, + KVCacheConfig, + KVCacheGroupSpec, + KVCacheLayout, + KVCacheTensor, + MLAAttentionSpec, + UniformTypeKVCacheSpecs, + create_kv_cache_views, + ) + + num_logical_blocks = 3 + transfer_block_size = 64 + kernel_block_size = 128 + tokens_per_state = 4 + state_content_bytes = 132 + + indexer_spec = MLAAttentionSpec( + block_size=logical_block_size, + num_kv_heads=1, + head_size=128, + head_size_v=0, + dtype=torch.uint8, + state_content_bytes=state_content_bytes, + tokens_per_state=tokens_per_state, + ) + indexer_page_size = indexer_spec.page_size_bytes + tail_spec = KpoolTailSpec( + block_size=tokens_per_state, + num_kv_heads=2, + head_size=128, + head_size_v=0, + dtype=torch.bfloat16, + page_size_padded=indexer_page_size, + sliding_window=tokens_per_state, + ) + + allocation_size = num_logical_blocks * indexer_page_size + indexer_tensor = KVCacheTensor( + size=allocation_size, + layers=["indexer"], + layer_stride=allocation_size, + block_stride=indexer_page_size, + ) + tail_tensor = KVCacheTensor( + size=allocation_size, + layers=["tail"], + layer_stride=allocation_size, + block_stride=indexer_page_size, + ) + kv_cache_config = KVCacheConfig( + num_blocks=num_logical_blocks, + kv_cache_tensors=[indexer_tensor, tail_tensor], + kv_cache_groups=[ + KVCacheGroupSpec( + ["indexer"], + UniformTypeKVCacheSpecs( + block_size=logical_block_size, + kv_cache_specs={"indexer": indexer_spec}, + ), + ), + KVCacheGroupSpec( + ["tail"], + UniformTypeKVCacheSpecs( + block_size=tokens_per_state, + kv_cache_specs={"tail": tail_spec}, + ), + ), + ], + ) + + raw = torch.zeros(allocation_size, dtype=torch.int8) + (indexer_cache,) = create_kv_cache_views( + raw, + indexer_spec, + num_logical_blocks, + KVCacheLayout.LBHNC, + indexer_tensor, + kernel_block_size=kernel_block_size, + ) + (tail_cache,) = create_kv_cache_views( + raw, + tail_spec, + num_logical_blocks, + KVCacheLayout.LBHNC, + tail_tensor, + ) + assert indexer_cache.data_ptr() == tail_cache.data_ptr() == raw.data_ptr() + + vllm_config = create_vllm_config(block_size=logical_block_size) + vllm_config.cache_config.kv_cache_layout = "LBHNC" + vllm_config.kv_transfer_config.kv_buffer_device = "cuda" + fake_backend = MagicMock() + fake_backend.get_supported_kernel_block_sizes.return_value = [transfer_block_size] + fake_backend.get_name.return_value = "DEEPSEEK_V32_INDEXER" + fake_backend.full_cls_name.return_value = "fake.DEEPSEEK_V32_INDEXER" + fake_platform = MagicMock() + fake_platform.device_type = "cuda" + fake_platform.get_nixl_memory_type.return_value = "VRAM" + + caches = [("indexer", indexer_cache), ("tail", tail_cache)] + if tail_first: + caches.reverse() + + with ( + patch.object(bw, "NixlWrapper", _RecordingNixl), + patch.object(bw, "get_tensor_model_parallel_rank", return_value=0), + patch.object(bw, "get_tensor_model_parallel_world_size", return_value=1), + patch.object(bw, "get_current_attn_backends", return_value=[fake_backend]), + patch.object(bw, "current_platform", fake_platform), + set_current_vllm_config(vllm_config), + ): + worker = NixlConnectorWorker(vllm_config, "local-engine", kv_cache_config) + worker.use_mla = True + worker.register_kv_caches(dict(caches)) + + transfer_page_size = transfer_block_size // tokens_per_state * state_content_bytes + num_transfer_blocks = num_logical_blocks * ( + logical_block_size // transfer_block_size + ) + expected_descs = np.asarray( + [ + [ + raw.data_ptr() + block_idx * transfer_page_size, + transfer_page_size, + 0, + ] + for block_idx in range(num_transfer_blocks) + ], + dtype=np.uint64, + ) + + assert worker.block_size == transfer_block_size + assert worker.num_regions == 1 + assert worker.block_len_per_layer == [transfer_page_size] + assert worker.block_stride_per_layer == [transfer_page_size] + assert worker._region_is_mla == [True] + assert worker.kv_caches_base_addr[worker.engine_id][0] == [raw.data_ptr()] + assert worker._registered_descs[0] == [(raw.data_ptr(), raw.nbytes, 0, "")] + np.testing.assert_array_equal(worker.src_blocks_data, expected_descs) + assert expected_descs[-1, 0] + expected_descs[-1, 1] == ( + raw.data_ptr() + raw.nbytes + ) + + def _make_remote_meta( worker, remote_block_size, diff --git a/tests/v1/worker/test_attn_utils.py b/tests/v1/worker/test_attn_utils.py index ba0f13d11972..9751865794b7 100644 --- a/tests/v1/worker/test_attn_utils.py +++ b/tests/v1/worker/test_attn_utils.py @@ -143,15 +143,22 @@ def test_reshape_padded_kv_cache_strides_by_padded_page(): @pytest.mark.parametrize( - ("kernel_block_sizes", "expected_num_blocks", "expected_num_states"), + ( + "kernel_block_sizes", + "storage_block_size", + "expected_num_blocks", + "expected_num_states", + ), [ - (None, 4, 64), - ([256], 4, 64), - ([64], 16, 16), + (None, None, 4, 64), + ([256], None, 4, 64), + ([64], None, 16, 16), + ([64], 256, 4, 64), ], ) def test_allocate_compressed_mla_cache( kernel_block_sizes: list[int] | None, + storage_block_size: int | None, expected_num_blocks: int, expected_num_states: int, ): @@ -161,6 +168,7 @@ def test_allocate_compressed_mla_cache( head_size=128, dtype=torch.bfloat16, tokens_per_state=4, + storage_block_size=storage_block_size, ) num_pages = 4 config = KVCacheConfig( diff --git a/tests/v1/worker/test_dsv4_packed_zeroer_geometry.py b/tests/v1/worker/test_dsv4_packed_zeroer_geometry.py index 400cfbfc2e11..f84c91872faa 100644 --- a/tests/v1/worker/test_dsv4_packed_zeroer_geometry.py +++ b/tests/v1/worker/test_dsv4_packed_zeroer_geometry.py @@ -76,6 +76,7 @@ def test_packed_dsv4_zeroer_zeroes_only_each_layers_page(): static_forward_context={ f"layer.{i}": SimpleNamespace(kv_cache=views[i]) for i in range(NUM_LAYERS) }, + num_blocks=NUM_BLOCKS, ) seg_addrs, seg_block_strides, seg_page_sizes, _, _, n_segs = zeroer._meta @@ -153,6 +154,7 @@ def make_spec(head_size): static_forward_context={ name: SimpleNamespace(kv_cache=views[name]) for name in views }, + num_blocks=config.num_blocks, ) seg_addrs, seg_block_strides, seg_page_sizes, _, _, n_segs = zeroer._meta diff --git a/tests/v1/worker/test_kv_block_zeroer.py b/tests/v1/worker/test_kv_block_zeroer.py index 945046ad6f83..64a81771ddd3 100644 --- a/tests/v1/worker/test_kv_block_zeroer.py +++ b/tests/v1/worker/test_kv_block_zeroer.py @@ -54,6 +54,7 @@ def test_attention_blocks_are_zeroed(spec): static_forward_context={ layer_name: SimpleNamespace(kv_cache=storage), }, + num_blocks=4, ) zeroer.zero_block_ids([1]) @@ -64,6 +65,56 @@ def test_attention_blocks_are_zeroed(spec): assert torch.equal(storage, expected) +@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA required") +def test_layers_in_one_group_may_use_different_kernel_pages_per_block(): + """Derive kernel pages per block from each layer's allocation.""" + device = torch.device("cuda") + num_blocks = 4 + spec = SlidingWindowSpec( + block_size=2, + num_kv_heads=1, + head_size=1, + dtype=torch.int32, + sliding_window=2, + ) + # Two kernel pages per logical block. + wide = torch.ones((num_blocks * 2, 3), dtype=torch.int32, device=device) + # One kernel page per logical block, carved out of a larger allocation so + # an over-strided write lands in the guard region instead of faulting or + # silently hitting another tensor. + narrow_backing = torch.ones((num_blocks * 3, 5), dtype=torch.int32, device=device) + narrow = narrow_backing[:num_blocks] + + zeroer = KVBlockZeroer( + device, + attn_groups_iter=[ + AttentionGroup( + None, + ["wide", "narrow"], + spec, + 0, + ) + ], + kernel_block_sizes=[1], + static_forward_context={ + "wide": SimpleNamespace(kv_cache=wide), + "narrow": SimpleNamespace(kv_cache=narrow), + }, + num_blocks=num_blocks, + ) + + zeroer.zero_block_ids([num_blocks - 1]) + torch.accelerator.synchronize() + + expected_wide = torch.ones_like(wide) + expected_wide[2 * (num_blocks - 1) :] = 0 + assert torch.equal(wide, expected_wide) + + expected_narrow = torch.ones_like(narrow_backing) + expected_narrow[num_blocks - 1] = 0 + assert torch.equal(narrow_backing, expected_narrow) + + @pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA required") def test_block_ids_are_not_overwritten_while_copy_is_in_flight(): device = torch.device("cuda") @@ -207,6 +258,7 @@ def test_large_dsv4_launch_geometry(monkeypatch): name: SimpleNamespace(kv_cache=storage) for name, storage in storages.items() }, + num_blocks=1, ) assert zeroer._meta is not None @@ -332,6 +384,7 @@ def test_zeroes_exactly_one_block_per_layer(layout: KVCacheLayout): attn_groups_iter=iter(groups), kernel_block_sizes=[spec.block_size], static_forward_context=ctx, + num_blocks=num_blocks, ) zeroer.zero_block_ids([2]) torch.accelerator.synchronize() diff --git a/tests/v1/worker/test_mamba_hybrid_model_state.py b/tests/v1/worker/test_mamba_hybrid_model_state.py index 749821274318..0cdd91e0d7b3 100644 --- a/tests/v1/worker/test_mamba_hybrid_model_state.py +++ b/tests/v1/worker/test_mamba_hybrid_model_state.py @@ -7,15 +7,57 @@ import pytest import torch +from vllm.config.compilation import CUDAGraphMode from vllm.platforms import current_platform from vllm.v1.attention.backends.recoverssm_metadata import ( RecoverSSMMetadata, RecoverSSMPostprocessMetadata, ) +from vllm.v1.worker.gpu.model_states import mamba_hybrid from vllm.v1.worker.gpu.model_states.mamba_hybrid import MambaHybridModelState from vllm.v1.worker.gpu.model_states.recoverssm import RecoverSSMState +def test_prepare_attn_forwards_positions(monkeypatch: pytest.MonkeyPatch) -> None: + state = object.__new__(MambaHybridModelState) + state.vllm_config = SimpleNamespace(num_speculative_tokens=0) + state.max_model_len = 8192 + state._align_mode = False + state.recoverssm = None + + positions = torch.tensor([1536], dtype=torch.int64) + input_batch = SimpleNamespace( + num_reqs=1, + num_tokens=1, + num_reqs_after_padding=1, + num_tokens_after_padding=1, + query_start_loc_np=torch.tensor([0, 1], dtype=torch.int32).numpy(), + query_start_loc=torch.tensor([0, 1], dtype=torch.int32), + num_scheduled_tokens=torch.tensor([1], dtype=torch.int32), + seq_lens_cpu_upper_bound=torch.tensor([1537], dtype=torch.int32), + seq_lens=torch.tensor([1537], dtype=torch.int32), + is_prefilling_np=torch.tensor([False]).numpy(), + dcp_local_seq_lens=None, + positions=positions, + prompt_lens=torch.tensor([1024], dtype=torch.int32), + ) + expected_metadata = {"layer": object()} + build_attn_metadata = Mock(return_value=expected_metadata) + monkeypatch.setattr(mamba_hybrid, "build_attn_metadata", build_attn_metadata) + + metadata = state.prepare_attn( + input_batch=input_batch, + cudagraph_mode=CUDAGraphMode.NONE, + block_tables=(), + slot_mappings=torch.empty(0, dtype=torch.int64), + attn_groups=[], + kv_cache_config=Mock(), + ) + + assert metadata is expected_metadata + assert build_attn_metadata.call_args.kwargs["positions"] is positions + + @pytest.mark.skipif(not current_platform.is_cuda(), reason="Requires CUDA") @pytest.mark.parametrize(("num_sampled", "expected_value"), [(0, 1), (3, 3)]) def test_postprocess_state_scalar_with_int32_mapping( diff --git a/vllm/_aiter_ops.py b/vllm/_aiter_ops.py index e865ea01b9cf..a68b18759857 100644 --- a/vllm/_aiter_ops.py +++ b/vllm/_aiter_ops.py @@ -3468,7 +3468,7 @@ def mhc_pre( # AITER's Python wrapper allocates intermediate/output tensors without # explicit device arguments, so run it under the residual tensor's device. with torch.device(residual_flat.device): - post_mix, comb_mix, layer_input = mhc_pre( + args = ( residual_flat, fn, hc_scale, @@ -3478,9 +3478,11 @@ def mhc_pre( hc_sinkhorn_eps, hc_post_mult_value, sinkhorn_repeat, - norm_weight, - norm_eps, ) + if norm_weight is None: + post_mix, comb_mix, layer_input = mhc_pre(*args) + else: + post_mix, comb_mix, layer_input = mhc_pre(*args, norm_weight, norm_eps) return ( post_mix.view(*outer_shape, hc_mult, 1), comb_mix.view(*outer_shape, hc_mult, hc_mult), @@ -3630,23 +3632,26 @@ def mhc_fused_post_pre( ), ) + args = ( + x_flat, + residual_flat, + post_flat, + comb_flat, + fn, + hc_scale, + hc_base, + rms_eps, + hc_pre_eps, + hc_sinkhorn_eps, + hc_post_mult_value, + sinkhorn_repeat, + ) with torch.device(residual_flat.device): - post_mix, comb_mix, layer_input, next_residual = mhc_fused_post_pre( - x_flat, - residual_flat, - post_flat, - comb_flat, - fn, - hc_scale, - hc_base, - rms_eps, - hc_pre_eps, - hc_sinkhorn_eps, - hc_post_mult_value, - sinkhorn_repeat, - norm_weight, - norm_eps, - ) + if norm_weight is None: + result = mhc_fused_post_pre(*args) + else: + result = mhc_fused_post_pre(*args, norm_weight, norm_eps) + post_mix, comb_mix, layer_input, next_residual = result return ( next_residual.view_as(residual), diff --git a/vllm/config/speculative.py b/vllm/config/speculative.py index 22196a2ddd78..49823bc03112 100644 --- a/vllm/config/speculative.py +++ b/vllm/config/speculative.py @@ -61,6 +61,7 @@ "hy_v4_mtp", "gemma4_mtp", "inkling_mtp", + "glm5_next_mtp", ] NgramGPUTypes = Literal["ngram_gpu"] DFlashModelTypes = Literal["dflash"] @@ -1016,6 +1017,12 @@ def hf_config_override(hf_config: PretrainedConfig) -> PretrainedConfig: hf_config.update( {"n_predict": n_predict, "architectures": ["MiniMaxM3MTP"]} ) + if hf_config.model_type == "glm5_next": + hf_config.model_type = "glm5_next_mtp" + n_predict = hf_config.num_nextn_predict_layers + hf_config.update( + {"n_predict": n_predict, "architectures": ["Glm5NextMTPModel"]} + ) return hf_config diff --git a/vllm/config/vllm.py b/vllm/config/vllm.py index e4c223daf69b..9c30a752f068 100644 --- a/vllm/config/vllm.py +++ b/vllm/config/vllm.py @@ -81,6 +81,9 @@ "DeepSeekV4MTPModel", "Dots3NoteForCausalLM", "Dots3NoteMTPModel", + "Glm5NextForCausalLM", + "Glm5NextForConditionalGeneration", + "Glm5NextMTPModel", "GlmMoeDsaForCausalLM", "HYV4ForCausalLM", "HYV4MTPModel", diff --git a/vllm/distributed/kv_transfer/kv_connector/v1/mooncake/mooncake_connector.py b/vllm/distributed/kv_transfer/kv_connector/v1/mooncake/mooncake_connector.py index 624bfb832152..e5f2e4f48c7c 100644 --- a/vllm/distributed/kv_transfer/kv_connector/v1/mooncake/mooncake_connector.py +++ b/vllm/distributed/kv_transfer/kv_connector/v1/mooncake/mooncake_connector.py @@ -56,6 +56,7 @@ from vllm.v1.kv_cache_interface import ( AttentionSpec, FullAttentionSpec, + KpoolTailSpec, KVCacheSpec, MambaSpec, SlidingWindowSpec, @@ -1675,7 +1676,9 @@ def register_kv_caches(self, kv_caches: dict[str, torch.Tensor]): block_len = region_cache.stride(0) * region_cache.element_size() region_base_addresses.append(base_addr) - if isinstance(layer_spec, AttentionSpec) and block_is_contiguous: + if isinstance(layer_spec, KpoolTailSpec): + kv_block_len = layer_spec.unpadded_page_size_bytes // 2 + elif isinstance(layer_spec, AttentionSpec) and block_is_contiguous: assert ( layer_spec.page_size_bytes % self._physical_blocks_per_logical_kv_block diff --git a/vllm/distributed/kv_transfer/kv_connector/v1/mooncake/store/coordinator.py b/vllm/distributed/kv_transfer/kv_connector/v1/mooncake/store/coordinator.py index c2f0e6640600..1d2fd7257c2a 100644 --- a/vllm/distributed/kv_transfer/kv_connector/v1/mooncake/store/coordinator.py +++ b/vllm/distributed/kv_transfer/kv_connector/v1/mooncake/store/coordinator.py @@ -75,8 +75,14 @@ def __init__( retention_interval: int | None = None, dcp_world_size: int = 1, ) -> None: + # Mirrors core's resolve_kv_cache_block_sizes: the hash unit only has + # to divide groups that participate in prefix caching. Non-shareable + # scratch groups (e.g. GLM-5.3-Flash's kpool tail) are skipped by + # _verify_and_split_kv_cache_groups and never probed for hits. assert all( - g.kv_cache_spec.block_size % hash_block_size == 0 for g in kv_cache_groups + g.kv_cache_spec.block_size % hash_block_size == 0 + for g in kv_cache_groups + if g.kv_cache_spec.prefix_cacheable ), "block_size must be divisible by hash_block_size" assert scheduler_block_size % hash_block_size == 0, ( f"scheduler_block_size ({scheduler_block_size}) must be a multiple of " @@ -116,6 +122,12 @@ def _verify_and_split_kv_cache_groups(self) -> None: """ attention_groups: list[SpecGroup] = [] for i, g in enumerate(self.kv_cache_groups): + # Skip groups that opt out of prefix caching (e.g. GLM-5.3-Flash + # kpool tail): per-request scratch, never shareable, so they must + # not participate in hit lookup. Mirrors core's + # KVCacheCoordinator.verify_and_split_kv_cache_groups. + if not g.kv_cache_spec.prefix_cacheable: + continue spec = _unwrap_spec(g.kv_cache_spec) manager_cls = KVCacheSpecRegistry.get_manager_class(spec) assert manager_cls is not None, ( diff --git a/vllm/distributed/kv_transfer/kv_connector/v1/mooncake/store/worker.py b/vllm/distributed/kv_transfer/kv_connector/v1/mooncake/store/worker.py index 89d5315aa913..870500e84068 100644 --- a/vllm/distributed/kv_transfer/kv_connector/v1/mooncake/store/worker.py +++ b/vllm/distributed/kv_transfer/kv_connector/v1/mooncake/store/worker.py @@ -523,6 +523,7 @@ def __init__( enable_group_semantics: bool = False, supports_group_ids: bool = False, record_operation: Callable[..., None] | None = None, + group_participates: Sequence[bool] | None = None, ): super().__init__( store, @@ -537,6 +538,11 @@ def __init__( self.group_put_steps = group_put_steps self.coord = coord self.kv_role = kv_role + self.group_participates = ( + list(group_participates) + if group_participates is not None + else [True] * len(token_databases) + ) # req_id -> ids of its store jobs that are still queued or running. # Keying by store_job_id, which never repeats for the engine's lifetime, # rather than counting jobs per request id makes the ledger immune to id @@ -744,6 +750,8 @@ def _sub_block_tail_puts( saved = self._saved_offset.get(req_meta.req_id, 0) puts: list[tuple[str, list[int], list[int], KeyMetadata]] = [] for g_idx, db in enumerate(self.token_databases): + if not self.group_participates[g_idx]: + continue group_blocks = req_meta.block_ids[g_idx] # Distribute across ranks by the same rule as normal chunks. put_step = self.group_put_steps[g_idx] @@ -983,6 +991,8 @@ def _handle_request(self, req_meta: ReqMeta): group_indices: list[int] = [] store_shard_ids: list[StoreShardId] = [] for g_idx, db in enumerate(self.token_databases): + if not self.group_participates[g_idx]: + continue # Rotate the stride phase per group to balance load across ranks. put_step = self.group_put_steps[g_idx] put_step_rank = (self.tp_rank + g_idx) % put_step @@ -1253,6 +1263,7 @@ def __init__( disk_offload_buffer_budget_bytes: int | None = None, record_operation: Callable[..., None] | None = None, request_queue: queue.Queue[Any] | None = None, + group_participates: Sequence[bool] | None = None, ): super().__init__( store, @@ -1264,6 +1275,11 @@ def __init__( record_operation=record_operation, request_queue=request_queue, ) + self.group_participates = ( + list(group_participates) + if group_participates is not None + else [True] * len(token_databases) + ) # _invalid_block_ids can be access by both the Worker and RecvingThread self._invalid_block_ids_lock = threading.Lock() self._invalid_block_ids: set[int] = set() @@ -1311,6 +1327,8 @@ def _handle_request(self, req_meta: ReqMeta): key_list: list[str] = [] block_id_list: list[int] = [] for g_idx, db in enumerate(self.token_databases): + if not self.group_participates[g_idx]: + continue mask = load_mask_per_group[g_idx] chunks: list[tuple[int, int]] = [] store_shard_ids: list[StoreShardId] = [] @@ -1791,6 +1809,11 @@ def _build_token_databases( """Construct token databases and their Store layouts.""" token_dbs: list[ChunkedTokenDatabase] = [] for group_idx, group in enumerate(self._kv_cache_groups): + hash_block_size = ( + self.hash_block_size + if group.kv_cache_spec.prefix_cacheable + else group.kv_cache_spec.block_size + ) group_tp_rank = self.tp_rank if layout_cls is None: group_tp_rank //= self._group_tp_replication_factors[group_idx] @@ -1805,7 +1828,7 @@ def _build_token_databases( store_layout = layout_cls( group_metadata, group.kv_cache_spec.block_size, - self.hash_block_size, + hash_block_size, local_tp_size=self.tp_size, store_tp_size=self.store_tp_size, tp_rank=self.tp_rank, @@ -1815,7 +1838,7 @@ def _build_token_databases( ChunkedTokenDatabase( group_metadata, group.kv_cache_spec.block_size, - hash_block_size=self.hash_block_size, + hash_block_size=hash_block_size, store_layout=store_layout, ) ) @@ -1937,6 +1960,10 @@ def register_kv_caches( enable_group_semantics=self.enable_group_semantics, supports_group_ids=self._supports_group_ids, record_operation=self._record_kv_connector_operation, + group_participates=[ + group.kv_cache_spec.prefix_cacheable + for group in self._kv_cache_groups + ], ) self.kv_send_thread.start() @@ -1954,6 +1981,10 @@ def register_kv_caches( disk_offload_buffer_budget_bytes=self.disk_offload_buffer_budget_bytes, record_operation=self._record_kv_connector_operation, request_queue=self.recv_request_queue, + group_participates=[ + group.kv_cache_spec.prefix_cacheable + for group in self._kv_cache_groups + ], ) recv_thread.name = f"KVCacheStoreRecvingThread-{i}" recv_thread.start() @@ -2126,6 +2157,8 @@ def lookup( fine_grained = self.coord.enable_partial_hash_hits lookup_masks = None if fine_grained else self.coord.lookup_mask(token_len) for g_idx, db in enumerate(self.token_dbs): + if not self._kv_cache_groups[g_idx].kv_cache_spec.prefix_cacheable: + continue spec_block_size = db.block_size key_prefixes = self._lookup_key_prefixes[g_idx] if fine_grained: @@ -2236,6 +2269,9 @@ def _tail_key_boundaries( boundaries = [] hit_boundary_hash_idx = hit_length // self.hash_block_size - 1 for group_id, db in enumerate(self.token_dbs): + if not self._kv_cache_groups[group_id].kv_cache_spec.prefix_cacheable: + # Scratch groups are never stored, so they have no tail key. + continue chunk_id = cdiv(hit_length, db.block_size) - 1 boundary_tokens = hit_length contains_hit_boundary = cached_block_pool.contains( diff --git a/vllm/distributed/kv_transfer/kv_connector/v1/nixl/base_worker.py b/vllm/distributed/kv_transfer/kv_connector/v1/nixl/base_worker.py index 46d0ec3092df..04c50df0f598 100644 --- a/vllm/distributed/kv_transfer/kv_connector/v1/nixl/base_worker.py +++ b/vllm/distributed/kv_transfer/kv_connector/v1/nixl/base_worker.py @@ -74,6 +74,7 @@ from vllm.v1.kv_cache_interface import ( CircularBufferSpec, FullAttentionSpec, + KpoolTailSpec, KVCacheLayout, KVCacheSpec, MambaSpec, @@ -99,6 +100,46 @@ def _share_storage_and_block_stride(caches: list[torch.Tensor]) -> bool: return len(block_strides) == len(storage_ptrs) == 1 +def _tensor_byte_span_end(cache: torch.Tensor) -> int: + """Return the exclusive end address touched by a nonnegative-stride view.""" + if cache.numel() == 0: + return cache.data_ptr() + if any(stride < 0 for stride in cache.stride()): + raise ValueError("NIXL cache views must have nonnegative strides") + max_element_offset = sum( + (size - 1) * stride for size, stride in zip(cache.shape, cache.stride()) + ) + return cache.data_ptr() + (max_element_offset + 1) * cache.element_size() + + +def _uses_dense_virtual_transfer_pages( + layer_spec: KVCacheSpec, + cache: torch.Tensor, + physical_page_size: int, + num_blocks: int, +) -> bool: + """Return whether a compressed kernel view can be split into NIXL pages.""" + if not ( + isinstance(layer_spec, MLAAttentionSpec) + and layer_spec.tokens_per_state > 1 + and cache.ndim == 4 + and cache.shape[1] == 1 + and cache.is_contiguous() + and physical_page_size > 0 + and layer_spec.state_content_size_bytes > 0 + ): + return False + + block_stride = cache.stride(0) * cache.element_size() + return ( + block_stride > physical_page_size + and block_stride % physical_page_size == 0 + and physical_page_size % layer_spec.state_content_size_bytes == 0 + and cache.shape[0] * (block_stride // physical_page_size) == num_blocks + and cache.nbytes == num_blocks * physical_page_size + ) + + class NixlBaseConnectorWorker: """Base implementation of Worker side methods shared by pull and push.""" @@ -1150,6 +1191,31 @@ def register_kv_caches(self, kv_caches: dict[str, torch.Tensor]): # aliases even though every view shares one block-major allocation. packed_storage = packed_storage and not self._is_csa_linear + layer_specs: dict[str, KVCacheSpec] = {} + compressed_region_owners: dict[int, torch.Tensor] = {} + for layer_name, cache in xfer_buffers.items(): + layer_spec = self._layer_specs.get(layer_name) + if isinstance(layer_spec, UniformTypeKVCacheSpecs): + layer_spec = layer_spec.kv_cache_specs[layer_name] + if layer_spec is None: + continue + layer_specs[layer_name] = layer_spec + physical_page_size = ( + layer_spec.page_size_bytes + if isinstance(layer_spec, MambaSpec) + else layer_spec.page_size_bytes + // self._physical_blocks_per_logical_kv_block + ) + num_blocks = ( + self._logical_num_blocks + if isinstance(layer_spec, MambaSpec) + else self.num_blocks + ) + if _uses_dense_virtual_transfer_pages( + layer_spec, cache, physical_page_size, num_blocks + ): + compressed_region_owners.setdefault(cache.data_ptr(), cache) + # K and V are packed into the content dim, so each attention layer is a # single NIXL region whose block transfers as one unit. Mamba layers instead # register separate conv/ssm sub-regions (see `_build_mamba_local`). @@ -1158,7 +1224,7 @@ def register_kv_caches(self, kv_caches: dict[str, torch.Tensor]): # kernel requires a specific block size. This leads to SSM and FA layers # having different num_blocks. # `_physical_blocks_per_logical_kv_block` ratio is used to adjust for this. - layer_spec = self._layer_specs.get(layer_name) + layer_spec = layer_specs.get(layer_name) if layer_spec is None: logger.debug( "Skipping layer %s as no KVCache spec is present. " @@ -1166,9 +1232,6 @@ def register_kv_caches(self, kv_caches: dict[str, torch.Tensor]): layer_name, ) continue - if isinstance(layer_spec, UniformTypeKVCacheSpecs): - # DSA Indexer case: UniformTypeKVCacheSpecs merges kv_cache_specs - layer_spec = layer_spec.kv_cache_specs[layer_name] # `layer_spec.page_size_bytes` only accounts for logical page_size, that is # the page_size assuming constant `self._logical_num_blocks`. physical_page_size = ( @@ -1187,6 +1250,33 @@ def register_kv_caches(self, kv_caches: dict[str, torch.Tensor]): ) storage = cache.untyped_storage() storage_addr = storage.data_ptr() + + if isinstance(layer_spec, KpoolTailSpec): + compressed_owner = compressed_region_owners.get(cache.data_ptr()) + if compressed_owner is not None: + owner_storage = compressed_owner.untyped_storage() + owner_end = compressed_owner.data_ptr() + compressed_owner.nbytes + tail_is_covered = ( + compressed_owner.is_contiguous() + and owner_storage.data_ptr() == storage_addr + and _tensor_byte_span_end(cache) <= owner_end + ) + if not tail_is_covered: + raise AssertionError( + "Kpool tail cache is not fully covered by its compressed " + f"indexer region: layer={layer_name}, " + f"tail_shape={tuple(cache.shape)}, " + f"tail_stride={tuple(cache.stride())}, " + f"owner_shape={tuple(compressed_owner.shape)}, " + f"owner_stride={tuple(compressed_owner.stride())}" + ) + logger.debug( + "Skipping layer %s because its compressed indexer region " + "covers the same storage", + layer_name, + ) + continue + # Memory registration follows allocations, while transfer regions follow # logical layers (or contiguous head segments). This keeps strided # cross-layer views inside their registered allocation. @@ -1228,7 +1318,15 @@ def register_kv_caches(self, kv_caches: dict[str, torch.Tensor]): and cache.stride(2) == cache.shape[3] and cache.stride(1) == cache.shape[2] * cache.shape[3] ) - if storage_is_block_major and ( + virtual_transfer_pages = _uses_dense_virtual_transfer_pages( + layer_spec, cache, physical_page_size, num_blocks + ) + if virtual_transfer_pages: + # A compressed kernel row can contain multiple NIXL transfer pages. + region_specs = [ + (cache.data_ptr(), physical_page_size, physical_page_size) + ] + elif storage_is_block_major and ( (packed_storage and is_mla_region) or (not hnc_contiguous and not self._is_csa_linear) ): @@ -1243,7 +1341,15 @@ def register_kv_caches(self, kv_caches: dict[str, torch.Tensor]): ] else: segment_bytes = num_blocks * block_stride - assert cache.nbytes % segment_bytes == 0 + if cache.nbytes % segment_bytes != 0: + raise AssertionError( + "KV cache view cannot be partitioned into NIXL regions: " + f"layer={layer_name}, cache_nbytes={cache.nbytes}, " + f"num_blocks={num_blocks}, block_stride={block_stride}, " + f"physical_page_size={physical_page_size}, " + f"cache_shape={tuple(cache.shape)}, " + f"cache_stride={tuple(cache.stride())}" + ) num_segments = cache.nbytes // segment_bytes region_block_len = ( block_stride if num_segments > 1 else physical_page_size diff --git a/vllm/distributed/kv_transfer/kv_connector/v1/nixl/metadata.py b/vllm/distributed/kv_transfer/kv_connector/v1/nixl/metadata.py index 99dae07f979a..840cb2aa73e0 100644 --- a/vllm/distributed/kv_transfer/kv_connector/v1/nixl/metadata.py +++ b/vllm/distributed/kv_transfer/kv_connector/v1/nixl/metadata.py @@ -45,8 +45,9 @@ # 7: Include NIXL transfer mode (push vs pull) in the compatibility hash # 8: Add dcp_size and pcp_size to NixlAgentMetadata # 9: Add block_strides +# 10: Add dense virtual transfer pages for compressed MLA caches # -NIXL_CONNECTOR_VERSION: int = 9 +NIXL_CONNECTOR_VERSION: int = 10 @dataclass diff --git a/vllm/model_executor/layers/attention/mla_attention.py b/vllm/model_executor/layers/attention/mla_attention.py index 55f7d82882d8..f0cdd803742e 100644 --- a/vllm/model_executor/layers/attention/mla_attention.py +++ b/vllm/model_executor/layers/attention/mla_attention.py @@ -1534,7 +1534,7 @@ def get_builder_cls() -> type["MLACommonMetadataBuilder"]: @classmethod def get_supported_head_sizes(cls) -> list[int]: - return [320, 576] + return [320, 512, 576] @classmethod def is_mla(cls) -> bool: diff --git a/vllm/model_executor/layers/mamba/ops/scatter_states.py b/vllm/model_executor/layers/mamba/ops/scatter_states.py new file mode 100644 index 000000000000..0102b6cf12b4 --- /dev/null +++ b/vllm/model_executor/layers/mamba/ops/scatter_states.py @@ -0,0 +1,72 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project + +import torch + +from vllm.platforms import current_platform +from vllm.triton_utils import tl, triton + + +@triton.jit +def _scatter_states_kernel( + state_ptr, + src_ptr, + indices_ptr, + stride_state_batch, + stride_src_batch, + stride_indices, + row_size: tl.constexpr, + BLOCK_SIZE: tl.constexpr, + launch_pdl: tl.constexpr, +): + block_idx = tl.program_id(0) + batch_idx = tl.program_id(1) + offsets = block_idx * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE) + mask = offsets < row_size + + if launch_pdl: + tl.extra.cuda.gdc_wait() + tl.extra.cuda.gdc_launch_dependents() + + state_idx = tl.load(indices_ptr + batch_idx * stride_indices).to(tl.int64) + values = tl.load(src_ptr + batch_idx * stride_src_batch + offsets, mask=mask) + tl.store(state_ptr + state_idx * stride_state_batch + offsets, values, mask=mask) + + +def scatter_states( + state: torch.Tensor, + src: torch.Tensor, + indices: torch.Tensor, +) -> None: + """Scatter ``src`` rows into ``state`` at ``indices`` (in place). + + Equivalent to ``state[indices] = src`` but non-atomic and bandwidth-bound, + since mamba cache slots are unique per sequence. ``gather_initial_states`` + is the read-side counterpart. + """ + assert state.ndim >= 2 + assert state.is_cuda + assert src.ndim == state.ndim + assert indices.ndim == 1 + assert indices.device == state.device + assert src.shape[1:] == state.shape[1:] + assert src.shape[0] == indices.shape[0] + assert indices.dtype in (torch.int32, torch.int64) + + row_size = state[0].numel() + assert state[0].is_contiguous() + assert src[0].is_contiguous() + block_size = min(triton.next_power_of_2(row_size), 1024) + grid = (triton.cdiv(row_size, block_size), indices.numel()) + _scatter_states_kernel[grid]( + state, + src, + indices, + state.stride(0), + src.stride(0), + indices.stride(0), + row_size=row_size, + BLOCK_SIZE=block_size, + num_warps=8, + launch_pdl=current_platform.is_arch_support_pdl(), + ) diff --git a/vllm/model_executor/layers/mhc.py b/vllm/model_executor/layers/mhc.py index 8f1c22f16ac3..d4bc224af474 100644 --- a/vllm/model_executor/layers/mhc.py +++ b/vllm/model_executor/layers/mhc.py @@ -1,5 +1,7 @@ # SPDX-License-Identifier: Apache-2.0 # SPDX-FileCopyrightText: Copyright contributors to the vLLM project +import inspect + import torch # this import will also register the custom ops @@ -29,6 +31,67 @@ def _has_tilelang_mhc() -> bool: HAS_AITER_MHC = is_aiter_found_and_supported() +def _has_aiter_mhc_fused() -> bool: + if not HAS_AITER_MHC: + return False + try: + from aiter.ops.mhc import mhc_fused_post_pre + except Exception: + return False + return callable(mhc_fused_post_pre) + + +def _aiter_mhc_op_accepts_norm(op_name: str) -> bool: + if not HAS_AITER_MHC: + return False + try: + from aiter.ops import mhc as aiter_mhc + + for candidate_name in (op_name, f"{op_name}_fake"): + op = getattr(aiter_mhc, candidate_name, None) + if op is None: + continue + parameters = inspect.signature(op).parameters + if "norm_weight" in parameters and "norm_eps" in parameters: + return True + except (AttributeError, ImportError, TypeError, ValueError): + pass + return False + + +HAS_AITER_MHC_FUSED = _has_aiter_mhc_fused() +HAS_AITER_MHC_PRE_NORM = _aiter_mhc_op_accepts_norm("mhc_pre") +HAS_AITER_MHC_FUSED_NORM = _aiter_mhc_op_accepts_norm("mhc_fused_post_pre") + + +def _aiter_mhc_supported( + residual: torch.Tensor, + norm_weight: torch.Tensor | None, + *, + supports_norm: bool, +) -> bool: + hidden_size = residual.shape[-1] + hc_mult = residual.shape[-2] + return ( + HAS_AITER_MHC + and hidden_size % 256 == 0 + and hc_mult == 4 + and (norm_weight is None or supports_norm) + ) + + +def _apply_mhc_norm( + layer_input: torch.Tensor, + norm_weight: torch.Tensor | None, + norm_eps: float, +) -> torch.Tensor: + if norm_weight is None: + return layer_input + from vllm import ir + + return ir.ops.rms_norm(layer_input, norm_weight, norm_eps) + + # --8<-- [start:mhc_pre] @CustomOp.register("mhc_pre") class MHCPreOp(CustomOp): @@ -89,9 +152,11 @@ def forward_hip( norm_weight: torch.Tensor | None = None, norm_eps: float = 0.0, ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: - hidden_size = residual.shape[-1] - hc_mult = residual.shape[-2] - if HAS_AITER_MHC and hidden_size % 256 == 0 and hc_mult == 4: + if _aiter_mhc_supported( + residual, + norm_weight, + supports_norm=HAS_AITER_MHC_PRE_NORM, + ): return torch.ops.vllm.mhc_pre_aiter( residual, fn, @@ -122,7 +187,7 @@ def forward_hip( norm_eps, ) else: - return self.forward_native( + post_mix, comb_mix, layer_input = self.forward_native( residual, fn, hc_scale, @@ -136,6 +201,11 @@ def forward_hip( norm_weight, norm_eps, ) + return ( + post_mix, + comb_mix, + _apply_mhc_norm(layer_input, norm_weight, norm_eps), + ) def forward_native( self, @@ -225,9 +295,7 @@ def forward_hip( post_layer_mix: torch.Tensor, comb_res_mix: torch.Tensor, ) -> torch.Tensor: - hidden_size = residual.shape[-1] - hc_mult = residual.shape[-2] - if HAS_AITER_MHC and hidden_size % 256 == 0 and hc_mult == 4: + if _aiter_mhc_supported(residual, None, supports_norm=True): return torch.ops.vllm.mhc_post_aiter( x, residual, @@ -449,9 +517,11 @@ def forward_hip( norm_weight: torch.Tensor | None = None, norm_eps: float = 0.0, ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]: - hidden_size = residual.shape[-1] - hc_mult = residual.shape[-2] - if HAS_AITER_MHC and hidden_size % 256 == 0 and hc_mult == 4: + if HAS_AITER_MHC_FUSED and _aiter_mhc_supported( + residual, + norm_weight, + supports_norm=HAS_AITER_MHC_FUSED_NORM, + ): return torch.ops.vllm.mhc_fused_post_pre_aiter( x, residual, @@ -489,7 +559,7 @@ def forward_hip( norm_weight, norm_eps, ) - return self.forward_native( + residual_cur, post_mix_cur, comb_mix_cur, layer_input_cur = self.forward_native( x, residual, post_layer_mix, @@ -507,6 +577,12 @@ def forward_hip( norm_weight, norm_eps, ) + return ( + residual_cur, + post_mix_cur, + comb_mix_cur, + _apply_mhc_norm(layer_input_cur, norm_weight, norm_eps), + ) def forward_native( self, @@ -577,3 +653,13 @@ def forward_xpu( hc_post_mult_value, sinkhorn_repeat, ) + + +def hc_expand(x: torch.Tensor, n: int) -> torch.Tensor: + """[s, hidden_size] -> [s, n * hidden_size] by replication.""" + return x.unsqueeze(1).expand(-1, n, -1).contiguous() + + +def hc_contract(x: torch.Tensor, n: int) -> torch.Tensor: + """[s, n * hidden_size] -> [s, hidden_size] by averaging.""" + return x.mean(dim=1) diff --git a/vllm/model_executor/layers/mla.py b/vllm/model_executor/layers/mla.py index 6c0e8e069fc7..421181306236 100644 --- a/vllm/model_executor/layers/mla.py +++ b/vllm/model_executor/layers/mla.py @@ -8,6 +8,7 @@ from vllm.model_executor.custom_op import PluggableLayer from vllm.model_executor.layers.attention import MLAAttention from vllm.model_executor.layers.quantization import QuantizationConfig +from vllm.models.common.ops import fused_q_kv_rmsnorm from vllm.platforms import current_platform @@ -17,7 +18,7 @@ class MLAModules: kv_a_layernorm: torch.nn.Module kv_b_proj: torch.nn.Module - rotary_emb: torch.nn.Module + rotary_emb: torch.nn.Module | None o_proj: torch.nn.Module fused_qkv_a_proj: torch.nn.Module | None kv_a_proj_with_mqa: torch.nn.Module | None @@ -69,6 +70,7 @@ def __init__( skip_topk: bool = False, non_causal_multi_token_decode: bool = False, allow_short_prefill_indexer_scoring_skip: bool = False, + fuse_qkv_rmsnorm: bool = False, ) -> None: super().__init__() self.hidden_size = hidden_size @@ -98,6 +100,9 @@ def __init__( # the topk_tokens buffer written by a previous layer in the same pass. # Refer: https://arxiv.org/abs/2603.12201 for more details. self.skip_topk = skip_topk + # When True, fuse the q_a and kv_a RMSNorms into a single kernel launch + # (MLA layers with q-LoRA). Opt-in; default False preserves other models. + self.fuse_qkv_rmsnorm = fuse_qkv_rmsnorm # qrep is active when the query projection is a DCP-group-sharded layer # that materializes the full group head set locally. q_proj_layer = self.q_b_proj if self.q_lora_rank is not None else self.q_proj @@ -172,8 +177,9 @@ def forward( [self.q_lora_rank, self.kv_lora_rank + self.qk_rope_head_dim], dim=-1, ) - q_c = self.q_a_layernorm(q_c) q_proj_layer = self.q_b_proj + if not self.fuse_qkv_rmsnorm: + q_c = self.q_a_layernorm(q_c) q_proj_input = q_c else: assert self.kv_a_proj_with_mqa is not None, ( @@ -187,7 +193,20 @@ def forward( q_proj_input = hidden_states kv_c, k_pe = kv_lora.split([self.kv_lora_rank, self.qk_rope_head_dim], dim=-1) - kv_c_normed = self.kv_a_layernorm(kv_c) + if self.fuse_qkv_rmsnorm and q_c is not None: + # Fuse q_a + kv_a RMSNorm into one kernel launch (q_c is still the + # pre-norm projection output here). + assert self.q_a_layernorm is not None + q_proj_input, kv_c_normed = fused_q_kv_rmsnorm( + q_c, + kv_c, + self.q_a_layernorm.weight.data, + self.kv_a_layernorm.weight.data, + self.q_a_layernorm.variance_epsilon, + ) + q_c = q_proj_input + else: + kv_c_normed = self.kv_a_layernorm(kv_c) # Add head dim of 1 to k_pe k_pe = k_pe.unsqueeze(1) diff --git a/vllm/model_executor/layers/quantization/utils/flashinfer_utils.py b/vllm/model_executor/layers/quantization/utils/flashinfer_utils.py index 7bccc373dde8..b31309679479 100644 --- a/vllm/model_executor/layers/quantization/utils/flashinfer_utils.py +++ b/vllm/model_executor/layers/quantization/utils/flashinfer_utils.py @@ -361,30 +361,27 @@ def _shuffle_deepseek_fp8_moe_weights( Returns 4D weight tensors in BlockMajorK layout (E, K/block_k, Mn, block_k) """ - from flashinfer import shuffle_matrix_a - from flashinfer.fused_moe import convert_to_block_layout + from flashinfer.utils import get_shuffle_matrix_a_row_indices epilogue_tile_m = 64 block_k = 128 - num_experts = w13.shape[0] - M13, K13 = w13.shape[1], w13.shape[2] - M2, K2 = w2.shape[1], w2.shape[2] - w13_out = torch.empty( - num_experts, K13 // block_k, M13, block_k, dtype=torch.uint8, device=w13.device - ) - w2_out = torch.empty( - num_experts, K2 // block_k, M2, block_k, dtype=torch.uint8, device=w2.device - ) - - for i in range(num_experts): - t13 = shuffle_matrix_a(w13[i].view(torch.uint8), epilogue_tile_m) - w13_out[i] = convert_to_block_layout(t13, block_k) - - t2 = shuffle_matrix_a(w2[i].view(torch.uint8), epilogue_tile_m) - w2_out[i] = convert_to_block_layout(t2, block_k) + def shuffle_to_block_major_k(w: torch.Tensor) -> torch.Tensor: + # shuffle_matrix_a's row permutation depends only on (M, + # epilogue_tile_m), so it is computed once and applied to every expert + # in a single gather instead of once per expert. Gathering through the + # BlockMajorK-permuted view also folds convert_to_block_layout into + # that same kernel. Per-expert loops here cost minutes for a MoE this + # wide (~24k tiny launches plus a host round-trip each). + num_experts, m, k = w.shape + rows = get_shuffle_matrix_a_row_indices( + w[0].view(torch.uint8), epilogue_tile_m + ).to(w.device) + blocked = w.view(torch.uint8).view(num_experts, m, k // block_k, block_k) + out = blocked.permute(0, 2, 1, 3)[:, :, rows, :].contiguous() + return out.view(torch.float8_e4m3fn) - return w13_out.view(torch.float8_e4m3fn), w2_out.view(torch.float8_e4m3fn) + return shuffle_to_block_major_k(w13), shuffle_to_block_major_k(w2) def _shuffle_mxfp8_moe_weights( @@ -396,70 +393,54 @@ def _shuffle_mxfp8_moe_weights( ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]: """Preprocess MXFP8 weights and scales for the FlashInfer TRT-LLM kernel. - Following flashinfer/tests/moe/test_trtllm_gen_fused_moe.py: - 1. reorder_rows_for_gated_act_gemm (interleave gate/up rows) - 2. shuffle_matrix_a (weight data layout shuffle) - 3. shuffle_matrix_sf_a (scale factor layout shuffle) + All three transforms (gate/up row reorder, ``shuffle_matrix_a`` weight + shuffle, ``shuffle_matrix_sf_a`` scale shuffle) are fixed row/index + permutations that depend only on the per-expert matrix shape, so the + permutation is computed once and applied to every expert in a single + gather instead of once per expert. ``block_scale_interleave`` accepts a + batched ``(E, M, K)`` scale tensor directly. Output is bit-identical to + the per-expert loop but ~20x faster (a 288-expert MoE otherwise costs + seconds per layer at load). """ - from flashinfer import ( - reorder_rows_for_gated_act_gemm, - shuffle_matrix_a, - shuffle_matrix_sf_a, + from flashinfer.fused_moe.core import ( + get_reorder_rows_for_gated_act_gemm_row_indices, ) + from flashinfer.quantization.fp4_quantization import block_scale_interleave + from flashinfer.utils import get_shuffle_matrix_a_row_indices epilogue_tile_m = 128 - num_experts = w13.shape[0] - intermediate_size = w13.shape[1] // 2 - hidden_size = w13.shape[2] - w13_interleaved: list[torch.Tensor] = [] - w13_scale_interleaved: list[torch.Tensor] = [] - for i in range(num_experts): - if is_gated: - w13_interleaved.append( - reorder_rows_for_gated_act_gemm( - w13[i].reshape(2 * intermediate_size, -1) - ) - ) - w13_scale_interleaved.append( - reorder_rows_for_gated_act_gemm( - w13_scale[i].reshape(2 * intermediate_size, -1) - ) - ) - else: - w13_interleaved.append(w13[i]) - w13_scale_interleaved.append(w13_scale[i]) - - w13_shuffled: list[torch.Tensor] = [] - w2_shuffled: list[torch.Tensor] = [] - w13_scale_shuffled: list[torch.Tensor] = [] - w2_scale_shuffled: list[torch.Tensor] = [] - for i in range(num_experts): - w13_shuffled.append( - shuffle_matrix_a(w13_interleaved[i].view(torch.uint8), epilogue_tile_m) - ) - w2_shuffled.append(shuffle_matrix_a(w2[i].view(torch.uint8), epilogue_tile_m)) - w13_scale_shuffled.append( - shuffle_matrix_sf_a( - w13_scale_interleaved[i] - .view(torch.uint8) - .reshape(2 * intermediate_size, -1), - epilogue_tile_m, - ) - ) - w2_scale_shuffled.append( - shuffle_matrix_sf_a( - w2_scale[i].view(torch.uint8).reshape(hidden_size, -1), - epilogue_tile_m, - ) - ) + w13_u = w13.view(torch.uint8) + w2_u = w2.view(torch.uint8) - w13_out = torch.stack(w13_shuffled).view(torch.float8_e4m3fn) - w2_out = torch.stack(w2_shuffled).view(torch.float8_e4m3fn) - w13_scale_out = torch.stack(w13_scale_shuffled).reshape(w13_scale.shape) - w2_scale_out = torch.stack(w2_scale_shuffled).reshape(w2_scale.shape) + # 1. Interleave gate/up rows (gated activation GEMM layout). + if is_gated: + gate_idx = get_reorder_rows_for_gated_act_gemm_row_indices( + w13_u[0].reshape(w13_u.shape[1], -1) + ) + w13_u = w13_u[:, gate_idx] + w13_scale = w13_scale[:, gate_idx] + + def shuffle_weights(t: torch.Tensor) -> torch.Tensor: + # Row permutation depends only on (M, epilogue_tile_m). + idx = get_shuffle_matrix_a_row_indices(t[0], epilogue_tile_m).to(t.device) + return t[:, idx].view(torch.float8_e4m3fn) + + def shuffle_scales(s: torch.Tensor) -> torch.Tensor: + # shuffle_matrix_sf_a == row-gather (same indices as the weight shuffle) + # followed by the 128x4 block-scale interleave, which is batch-capable. + idx = get_shuffle_matrix_a_row_indices( + s[0].view(torch.uint8).reshape(s.shape[1], -1), epilogue_tile_m + ).to(s.device) + interleaved = block_scale_interleave(s[:, idx]) + return interleaved.reshape(s.shape) - return w13_out, w2_out, w13_scale_out, w2_scale_out + return ( + shuffle_weights(w13_u), + shuffle_weights(w2_u), + shuffle_scales(w13_scale), + shuffle_scales(w2_scale), + ) def prepare_fp8_moe_layer_for_fi( diff --git a/vllm/model_executor/layers/sparse_attn_indexer_kpool.py b/vllm/model_executor/layers/sparse_attn_indexer_kpool.py new file mode 100644 index 000000000000..ba7a5adff160 --- /dev/null +++ b/vllm/model_executor/layers/sparse_attn_indexer_kpool.py @@ -0,0 +1,1053 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +"""Custom Sparse Attention Indexer layers.""" + +from typing import TYPE_CHECKING + +import torch + +import vllm.envs as envs +from vllm._aiter_ops import rocm_aiter_ops +from vllm.compilation.breakable_cudagraph import eager_break_during_capture +from vllm.config import get_current_vllm_config_or_none +from vllm.forward_context import get_forward_context +from vllm.logger import init_logger +from vllm.model_executor.custom_op import CustomOp +from vllm.platforms import current_platform + +if TYPE_CHECKING: + from vllm.models.glm5next.nvidia.ops import kpool_compress as kpool_ops +elif current_platform.is_rocm(): + from vllm.models.glm5next.amd.ops import kpool_compress as kpool_ops +else: + from vllm.models.glm5next.nvidia.ops import kpool_compress as kpool_ops + +from vllm.utils.deep_gemm import has_deep_gemm +from vllm.utils.torch_utils import ( + LayerNameType, + _encode_layer_name, + _resolve_layer_name, +) +from vllm.v1.attention.backends.mla.indexer import ( + DeepseekV32IndexerMetadata, +) +from vllm.v1.attention.ops.common import pack_seq_triton, unpack_seq_triton +from vllm.v1.worker.workspace import current_workspace_manager + +if current_platform.is_cuda_alike(): + from vllm import _custom_ops as ops +elif current_platform.is_xpu(): + from vllm._xpu_ops import xpu_ops + +logger = init_logger(__name__) + +RADIX_TOPK_WORKSPACE_SIZE = 1024 * 1024 + +# MXFP4 layout: 2 values packed per byte, ue8m0 (1-byte) scale per block of 32. +MXFP4_BLOCK_SIZE = 32 + +# kpool write helper: form pools from the current token batch and compress them +# into the index K cache via the fused Triton kernel. + + +def _kpool_compress_insert( + k: torch.Tensor, + gate_score: torch.Tensor, + ape: torch.Tensor, + kv_cache: torch.Tensor, + slot_mapping: torch.Tensor, + kpool: int, + head_dim: int, + round_scale: bool, +) -> None: + """Pool ``kpool`` consecutive tokens into one fp8 K and write at pool slots. + + ``slot_mapping`` is pool-granular (compress_ratio == kpool on the spec): + only the *last* token of each complete pool carries a valid (>=0) slot; + intra-pool tokens are -1. Every position is treated as a pool-completion + candidate and non-completions are masked off inside the kernel. Compacting + the valid rows first costs two device syncs on the eager prefill path and + buys nothing numerically. Assumes pool-aligned chunk starts. + """ + n = slot_mapping.shape[0] + # No pool can complete in a batch smaller than one pool; also keeps the + # clamped gather indices below in bounds. + if n < kpool: + return + pos = torch.arange(n, device=k.device) + valid = slot_mapping >= 0 + # Drop pools whose start falls before the batch (leading padding); their + # gate/k data is undefined anyway. + write_mask = valid & (pos >= kpool - 1) + offs = torch.arange(kpool, device=k.device) + idx = (pos - (kpool - 1)).clamp_min(0)[:, None] + offs[None, :] + kpool_ops.kpool_compress_and_write_cache( + kv_cache, + k[idx], # [n, kpool, head_dim] + gate_score[idx], + ape, + slot_mapping.to(torch.int64), + pool_size=kpool, + head_dim=head_dim, + write_mask=write_mask, + round_scale=round_scale, + write_cache=True, + return_compressed=False, + ) + + +def _build_decode_scatter_indices( + decode_lens: torch.Tensor, + num_requests: int, + n: int, +) -> tuple[torch.Tensor, torch.Tensor]: + """Per-token (request id, intra-request index) for a non-uniform decode + batch, with ``n == decode_lens.sum()`` as a host int (avoids a + device sync and keeps both repeat_interleaves sync-free). + + Shared by every ``_scatter_decode_tokens_by_request`` call in a step: + building it per call would repeat the same repeat_interleave/cumsum chain + up to 5x per layer on the eager decode break. + """ + device = decode_lens.device + dl = decode_lens.to(torch.int64) + req_id = torch.repeat_interleave( + torch.arange(num_requests, device=device, dtype=torch.int64), + dl, + output_size=n, + ) + req_starts = torch.cumsum( + torch.cat([torch.zeros(1, device=device, dtype=torch.int64), dl[:-1]]), + dim=0, + ) + # Broadcast the per-request start offsets to per-token (length n == + # dl.sum()) so each token's intra-request index subtracts its own + # request's start. + starts = torch.repeat_interleave(req_starts, dl, output_size=n) + intra = torch.arange(n, device=device, dtype=torch.int64) - starts + return req_id, intra + + +def _scatter_decode_tokens_by_request( + tokens: torch.Tensor, + pad_value, + num_requests: int, + lmax: int, + scatter_indices: tuple[torch.Tensor, torch.Tensor], +) -> torch.Tensor: + """Group ``[N, ...]`` decode tokens into a padded ``[num_requests, lmax, ...]`` + layout: request ``r``'s tokens at row ``r`` in order; short requests padded. + + Unlike ``pack_seq_triton`` this is dtype-agnostic (needed for the int32 + slot/pos tensors) — it scatters with the shared per-step indices from + ``_build_decode_scatter_indices``. Used only for the non-uniform + (``requires_padding``) decode batch; uniform batches use a zero-copy + reshape. + """ + req_id, intra = scatter_indices + out = torch.full( + (num_requests, lmax, *tokens.shape[1:]), + pad_value, + dtype=tokens.dtype, + device=tokens.device, + ) + out[req_id, intra] = tokens + return out + + +def _decode_topk_seq_lens( + positions: torch.Tensor, + decode_lens: torch.Tensor, + num_decode_tokens: int, + batch_size: int, + next_n: int, + requires_padding: bool, +) -> torch.Tensor: + """Token-granular seq_len (pos + 1) per pool-topk row, layout-aware. + + ``pool_topk`` (and the logits it comes from) follow the padded + ``[batch_size, next_n]`` grid whenever ``requires_padding`` is set, so row + ``(b, t)`` corresponds to flat decode token ``offset_b + t`` -- NOT + ``b * next_n + t``. Slicing flat ``positions[: batch_size * next_n]`` + (the uniform-layout shortcut) misaligns every row after the first + non-uniform request and, past the decode region, reads prefill tokens' + positions; ``expand_pools_and_append_tail`` then anchors the tail at + another request's length, dropping the row's real tail tokens or emitting + indices past its sequence (out-of-bounds block-table reads). Padded rows + get 0 (empty tail); they are dropped by ``unpack_seq_triton`` anyway. + """ + n = batch_size * next_n + if not requires_padding: + return positions[:n].to(torch.int32) + 1 + scatter_idx = _build_decode_scatter_indices( + decode_lens, batch_size, num_decode_tokens + ) + padded = _scatter_decode_tokens_by_request( + positions[:num_decode_tokens].to(torch.int32), + -1, + batch_size, + next_n, + scatter_idx, + ) + return padded.reshape(n) + 1 # pad rows: -1 + 1 = 0 -> empty tail + + +def _fill_causal_indices(rows: torch.Tensor, positions: torch.Tensor) -> None: + causal_range = torch.arange(rows.shape[1], device=rows.device, dtype=torch.int32) + positions = positions.to(torch.int32) + rows[:] = causal_range[None, :] + rows[causal_range[None, :] > positions[:, None]] = -1 + + +def _fill_short_decode_causal_indices( + topk_indices_buffer: torch.Tensor, + positions: torch.Tensor | None, + num_decode_tokens: int, + max_seq_len: int, + topk_tokens: int, +) -> bool: + """Fill exact causal rows when sparse decode would select every token.""" + if positions is None or positions.numel() == 0 or max_seq_len > topk_tokens: + return False + _fill_causal_indices( + topk_indices_buffer[:num_decode_tokens], positions[:num_decode_tokens] + ) + return True + + +def _gather_workspace_shapes( + total_seq_lens: int, + head_dim: int, + fp8_dtype: torch.dtype, + use_fp4_cache: bool, +) -> tuple[tuple[tuple[int, int], torch.dtype], tuple[tuple[int, int], torch.dtype]]: + """Return ((values_shape, values_dtype), (scales_shape, scales_dtype)) for + the K-gather workspace. FP8 path: (T, head_dim) fp8 + (T, 4) uint8 fp32 + scales. MXFP4 path: (T, head_dim // 2) uint8 packed mxfp4 + + (T, head_dim // MXFP4_BLOCK_SIZE) uint8 ue8m0 scales.""" + if use_fp4_cache: + return ( + ((total_seq_lens, head_dim // 2), torch.uint8), + ((total_seq_lens, head_dim // MXFP4_BLOCK_SIZE), torch.uint8), + ) + return ( + ((total_seq_lens, head_dim), fp8_dtype), + ((total_seq_lens, 4), torch.uint8), + ) + + +def kv_cache_as_quant_view( + kv_cache: torch.Tensor, + head_dim: int, + use_fp4_cache: bool, +) -> torch.Tensor: + """4D ``[num_blocks, block_size, 1, head_width]`` view expected by + DeepGEMM, from the 3D indexer kv-cache allocation.""" + if use_fp4_cache: + assert kv_cache.ndim == 3 and kv_cache.dtype == torch.uint8 + num_blocks, block_size, _ = kv_cache.shape + page_bytes = int(kv_cache.stride(0)) + fp4_bytes = head_dim // 2 + head_dim // MXFP4_BLOCK_SIZE + return torch.as_strided( + kv_cache, + size=(num_blocks, block_size, 1, fp4_bytes), + stride=(page_bytes, fp4_bytes, fp4_bytes, 1), + ) + return kv_cache.unsqueeze(-2) + + +@eager_break_during_capture +def sparse_attn_indexer_kpool( + hidden_states: torch.Tensor, + k_cache_prefix: LayerNameType, + kv_cache: torch.Tensor, + q_quant: torch.Tensor, + q_scale: torch.Tensor | None, + k: torch.Tensor, + weights: torch.Tensor, + quant_block_size: int, + scale_fmt: str | None, + topk_tokens: int, + head_dim: int, + max_model_len: int, + total_seq_lens: int, + topk_indices_buffer: torch.Tensor, + skip_k_cache_insert: bool, + use_fp4_cache: bool = False, + # kpool params (Plan-A: gate is consumed at write time and read back at + # topk time to softmax-weight the pool). + gate_score: torch.Tensor | None = None, + compress_ape: torch.Tensor | None = None, + index_kpool: int = 1, + positions: torch.Tensor | None = None, + # Paged tail cache (in-progress pool's raw K + gate score), replacing the + # transient _DECODE_TAIL ring. tail_prefix resolves attn_metadata[tail_prefix] + # for the tail group's token-granular slot_mapping. None on the dummy/profiling + # path and when the tail cache is disabled. + tail_kv_cache: torch.Tensor | None = None, + tail_prefix: str | None = None, +) -> torch.Tensor: + # careful! this will be None in dummy run + attn_metadata = get_forward_context().attn_metadata + fp8_dtype = current_platform.fp8_dtype() + k_cache_prefix = _resolve_layer_name(k_cache_prefix) + + # assert isinstance(attn_metadata, dict) + if not isinstance(attn_metadata, dict): + # Reserve workspace for indexer during profiling run + values_spec, scales_spec = _gather_workspace_shapes( + total_seq_lens, head_dim, fp8_dtype, use_fp4_cache + ) + current_workspace_manager().get_simultaneous( + values_spec, + scales_spec, + ((RADIX_TOPK_WORKSPACE_SIZE,), torch.uint8), + ) + + # Reserve profiler-visible memory for the worst-case decode logits, + # whose shape is [B * next_n, max_model_len]. This profiling branch + # returns before invoking the logits kernel itself. + cfg = get_current_vllm_config_or_none() + worst_decode_tokens = 0 + if cfg is not None: + sched = cfg.scheduler_config + num_spec = ( + cfg.speculative_config.num_speculative_tokens + if cfg.speculative_config is not None + else 0 + ) + worst_decode_tokens = min( + sched.max_num_seqs * (num_spec + 1), + sched.max_num_batched_tokens, + ) + # float32 logits -> 4 bytes/element; uint8 sentinel so elems == bytes. + decode_logits_elems = worst_decode_tokens * max_model_len * 4 + prefill_cap_elems = envs.VLLM_SPARSE_INDEXER_MAX_LOGITS_MB * 1024 * 1024 + max_logits_elems = max(decode_logits_elems, prefill_cap_elems) + _ = torch.empty( + max_logits_elems, dtype=torch.uint8, device=hidden_states.device + ) + + return topk_indices_buffer + attn_metadata_narrowed = attn_metadata[k_cache_prefix] + assert isinstance(attn_metadata_narrowed, DeepseekV32IndexerMetadata) + slot_mapping = attn_metadata_narrowed.slot_mapping + has_decode = attn_metadata_narrowed.num_decodes > 0 + has_prefill = attn_metadata_narrowed.num_prefills > 0 + num_decode_tokens = attn_metadata_narrowed.num_decode_tokens + + # q_scale is required iff the FP4 cache path is enabled; the FP8 path + # folds the Q scale into `weights` inside fused_indexer_q_rope_quant. + if use_fp4_cache: + assert q_scale is not None, "use_fp4_cache=True requires q_scale" + else: + assert q_scale is None, "q_scale must be None when use_fp4_cache=False" + + # During speculative decoding, k may be padded to the CUDA graph batch + # size while slot_mapping only covers actual tokens. Truncate k to avoid + # out-of-bounds reads in the kernel. + num_tokens = slot_mapping.shape[0] + if k is not None: + k = k[:num_tokens] + + if not skip_k_cache_insert: + assert not use_fp4_cache, "Unfused FP4 Insert is not supported yet" + if index_kpool > 1 and gate_score is not None and compress_ape is not None: + # kpool prefill write: pool kpool consecutive prefill tokens via + # softmax(gate+ape)-weighted sum -> Hadamard -> fp8 -> pool slots. + # Decode tokens (the first num_decode_tokens in the batch) cannot be + # pooled here — their pool's earlier tokens are not in this batch — + # so they are deferred to the tail-buffer kernel in has_decode. + # compress_ratio == index_kpool makes slot_mapping pool-granular. + n_prefill = num_tokens - num_decode_tokens + if n_prefill > 0: + # decode tokens are batched first; prefill tokens follow. + prefill_slice = slice(num_decode_tokens, num_tokens) + _kpool_compress_insert( + k[prefill_slice], + gate_score[prefill_slice], + compress_ape, + kv_cache, + slot_mapping[prefill_slice], + index_kpool, + head_dim, + round_scale=(scale_fmt is not None), + ) + # Persist each request's incomplete prefill pool so decode can + # finish it, including after PD transfer. Tail slots use + # ``pos % kpool`` within the request's tail block. Processing + # only the batch's trailing tokens would miss all but the last + # request in a multi-request prefill. + if tail_kv_cache is not None and tail_prefix is not None: + tail_meta = attn_metadata.get(_resolve_layer_name(tail_prefix)) + if tail_meta is not None: + assert isinstance(tail_meta, DeepseekV32IndexerMetadata) + kpool_ops.kpool_seed_tail_cache( + tail_kv_cache, + k[prefill_slice], + gate_score[prefill_slice], + tail_meta.slot_mapping[prefill_slice], + index_kpool, + head_dim, + ) + else: + # standard: per-token fp8 quant + scatter (all tokens). + assert scale_fmt is not None + if current_platform.is_rocm(): + from vllm.v1.attention.ops.rocm_aiter_mla_sparse import ( + indexer_k_quant_and_cache_triton, + ) + + indexer_k_quant_and_cache_triton( + k, + kv_cache, + slot_mapping, + quant_block_size, + scale_fmt, + ) + else: + ops.indexer_k_quant_and_cache( + k, + kv_cache, + slot_mapping, + quant_block_size, + scale_fmt, + ) + + topk_indices_buffer[: hidden_states.shape[0]] = -1 + if has_prefill: + prefill_metadata = attn_metadata_narrowed.prefill + assert prefill_metadata is not None + + # Short sequences select every pool, so skip sparse scoring and fill + # the top-k buffer with all causal token indices. The index-K cache was + # already written above. + n_prefill_sf = num_tokens - num_decode_tokens + # Host-side short-prefill predicate: max_prefill_seq_len is computed + # in the metadata builder (exact for prefill rows) and equals + # positions[prefill_slice].max() + 1, so this replaces a + # positions.max().item() device sync per layer. -1 (unknown metadata) + # falls back to the device-side check. + if prefill_metadata.max_prefill_seq_len >= 0: + short_prefill = ( + n_prefill_sf > 0 + and positions is not None + and prefill_metadata.max_prefill_seq_len <= topk_tokens + ) + else: + short_prefill = ( + n_prefill_sf > 0 + and positions is not None + and int(positions[num_decode_tokens:num_tokens].max().item()) + 1 + <= topk_tokens + ) + if short_prefill: + # short_prefill is only True when positions is not None (above), + # but narrow explicitly for the indexer below. + assert positions is not None + _pos = positions[num_decode_tokens:num_tokens].to(torch.int32) + _buf = topk_indices_buffer[num_decode_tokens:num_tokens] + _fill_causal_indices(_buf, _pos) + + # Get the full shared workspace buffers once (will allocate on first use). + # Layout switches between FP8 (head_dim bytes + 4-byte fp32 scale) and + # MXFP4 (head_dim/2 bytes packed + head_dim/MXFP4_BLOCK_SIZE ue8m0 + # scales) based on use_fp4_cache. + workspace_manager = current_workspace_manager() + values_spec, scales_spec = _gather_workspace_shapes( + total_seq_lens, head_dim, fp8_dtype, use_fp4_cache + ) + k_quant_full, k_scale_full = workspace_manager.get_simultaneous( + values_spec, + scales_spec, + ) + for chunk in prefill_metadata.chunks if not short_prefill else (): + k_quant = k_quant_full[: chunk.total_seq_lens] + k_scale = k_scale_full[: chunk.total_seq_lens] + + if not chunk.skip_kv_gather: + if current_platform.is_rocm(): + from vllm.v1.attention.ops.rocm_aiter_mla_sparse import ( + cp_gather_indexer_k_quant_cache_triton, + ) + + cp_gather_indexer_k_quant_cache_triton( + kv_cache, + k_quant, + k_scale, + chunk.block_table, + chunk.cu_seq_lens, + token_to_seq=chunk.token_to_seq, + ) + else: + ops.cp_gather_indexer_k_quant_cache( + kv_cache, + k_quant, + k_scale, + chunk.block_table, + chunk.cu_seq_lens, + ) + + q_slice = q_quant[chunk.token_start : chunk.token_end] + q_scale_slice = ( + q_scale[chunk.token_start : chunk.token_end] + if q_scale is not None + else None + ) + # DeepGEMM scalar-type tags (zero-copy): MXFP4 values → int8 + # (kPackedFP4), scales → int32 squeezed to 1-D kv_sf / 2-D q_sf. + if use_fp4_cache: + q_slice_cast = q_slice.view(torch.int8) + k_quant_cast = k_quant.view(torch.int8) + k_scale_cast = k_scale.view(torch.int32).squeeze(-1) + else: + q_slice_cast = q_slice + k_quant_cast = k_quant + k_scale_cast = k_scale.view(torch.float32).squeeze(-1) + if current_platform.is_rocm(): + from vllm.v1.attention.ops.rocm_aiter_mla_sparse import ( + rocm_fp8_mqa_logits, + ) + + assert q_scale_slice is None + logits = rocm_fp8_mqa_logits( + q_slice_cast, + (k_quant_cast, k_scale_cast), + weights[chunk.token_start : chunk.token_end], + chunk.cu_seqlen_ks, + chunk.cu_seqlen_ke, + ) + else: + from vllm.utils.deep_gemm import fp8_fp4_mqa_logits + + logits = fp8_fp4_mqa_logits( + (q_slice_cast, q_scale_slice), + (k_quant_cast, k_scale_cast), + weights[chunk.token_start : chunk.token_end], + chunk.cu_seqlen_ks, + chunk.cu_seqlen_ke, + clean_logits=False, + ) + num_rows = logits.shape[0] + + # kpool: logits are pool-granular (compress_ratio == index_kpool), + # so topk selects pools. We pick topk_tokens // kpool pools then + # expand each pool back to its kpool constituent tokens. + select_k = topk_tokens // index_kpool if index_kpool > 1 else topk_tokens + if index_kpool > 1: + pool_topk = torch.empty( + (num_rows, select_k), dtype=torch.int32, device=logits.device + ) + topk_dst = pool_topk + else: + topk_dst = topk_indices_buffer[ + chunk.token_start : chunk.token_end, :topk_tokens + ] + + if current_platform.is_xpu(): + xpu_ops.top_k_per_row_prefill( # type: ignore[attr-defined] + logits, + chunk.cu_seqlen_ks, + chunk.cu_seqlen_ke, + topk_dst, + num_rows, + logits.stride(0), + logits.stride(1), + select_k, + ) + else: + torch.ops._C.top_k_per_row_prefill( + logits, + chunk.cu_seqlen_ks, + chunk.cu_seqlen_ke, + topk_dst, + num_rows, + logits.stride(0), + logits.stride(1), + select_k, + ) + + if index_kpool > 1: + pool_ids = pool_topk.to(torch.int64) + if positions is not None: + # Fused expand-pools + append-tail into one Triton kernel + # (replaces ~25 elementwise ops). seq_len is token-granular + # (pos+1); the kernel derives pool_len internally. + q_seq = ( + positions[chunk.token_start : chunk.token_end].to(torch.int32) + + 1 + ) + expanded = kpool_ops.expand_pools_and_append_tail( + pool_ids, q_seq, index_kpool + ) + else: + valid = pool_ids >= 0 + expanded = kpool_ops.expand_pools_to_tokens( + pool_ids, valid, topk_tokens, index_kpool + ) + topk_indices_buffer[ + chunk.token_start : chunk.token_end, : expanded.shape[-1] + ] = expanded + + if has_decode: + decode_metadata = attn_metadata_narrowed.decode + assert decode_metadata is not None + kv_cache_raw = kv_cache # raw [num_blocks, block_size, head_dim+4] for writes + kv_cache = kv_cache_as_quant_view(kv_cache, head_dim, use_fp4_cache) + + # Update the tail before reading logits; completed pools are compressed + # into the slot supplied by slot_mapping. + # Spec verification groups tokens by request and preserves position + # order so each token is stashed before the next completes its pool. + # Positions must remain token-granular because the kernel derives the + # pool phase and tail index from ``pos % kpool``. + if ( + index_kpool > 1 + and gate_score is not None + and compress_ape is not None + and positions is not None + and not skip_k_cache_insert + ): + num_requests = attn_metadata_narrowed.num_decodes + # Kpool writes must recover the original request grouping after the + # indexer's flattened decode path. Host metadata avoids a CUDA graph + # sync when choosing the uniform or padded layout. + per_req_lens = decode_metadata.per_req_decode_lens + if per_req_lens is not None: + use_uniform = ( + decode_metadata.decode_is_uniform + and num_decode_tokens + == num_requests * decode_metadata.write_max_decode_len + ) + group_lens = per_req_lens + lmax = decode_metadata.write_max_decode_len + else: + # Legacy metadata without per-request lens: fall back to the + # host-side requires_padding flag. Unreached now (per-request + # lens is always populated for decode), kept defensive. + use_uniform = not decode_metadata.requires_padding + group_lens = decode_metadata.decode_lens + lmax = int(decode_metadata.decode_lens.max().item()) + if not use_uniform: + # Non-uniform decode_lens (mixed plain-decode + spec-verify, or + # a variable MTP-verify batch): scatter actual tokens into a + # padded [B, lmax] layout. int32 tensors can't go through + # pack_seq_triton (float/uint8 only). The scatter indices are + # shared by all five scatters below (and the tail slot one). + scatter_idx = _build_decode_scatter_indices( + group_lens, num_requests, num_decode_tokens + ) + dec_k = _scatter_decode_tokens_by_request( + k[:num_decode_tokens], 0, num_requests, lmax, scatter_idx + ) + dec_gate = _scatter_decode_tokens_by_request( + gate_score[:num_decode_tokens], + 0, + num_requests, + lmax, + scatter_idx, + ) + dec_slot = _scatter_decode_tokens_by_request( + slot_mapping[:num_decode_tokens], + -1, + num_requests, + lmax, + scatter_idx, + ) + dec_pos = _scatter_decode_tokens_by_request( + positions[:num_decode_tokens].to(torch.int32), + -1, + num_requests, + lmax, + scatter_idx, + ) + else: + next_n = num_decode_tokens // num_requests + shape2 = (num_requests, next_n) + dec_k = k[:num_decode_tokens].view(*shape2, head_dim) + dec_gate = gate_score[:num_decode_tokens].view(*shape2, head_dim) + dec_slot = slot_mapping[:num_decode_tokens].view(shape2) + dec_pos = positions[:num_decode_tokens].to(torch.int32).view(shape2) + tail_meta = ( + attn_metadata.get(_resolve_layer_name(tail_prefix)) + if tail_prefix is not None + else None + ) + # Paged tail cache replaces the transient _DECODE_TAIL ring. Group + # the tail group's token-granular slot_mapping per-request, mirroring + # dec_slot / dec_pos, so the kernel gets each request's current-token + # tail slot (block * kpool + pos % kpool). + if tail_meta is not None: + assert isinstance(tail_meta, DeepseekV32IndexerMetadata) + if tail_meta is None or tail_kv_cache is None: + dec_tail_slot = None + elif not use_uniform: + dec_tail_slot = _scatter_decode_tokens_by_request( + tail_meta.slot_mapping[:num_decode_tokens], + -1, + num_requests, + lmax, + scatter_idx, + ) + else: + dec_tail_slot = tail_meta.slot_mapping[:num_decode_tokens].view(shape2) + # The compress kernel writes the raw fp8 cache (not the quant view); + # pass the underlying kv_cache, not kv_cache_quant_view. + if dec_tail_slot is not None: + # Single batched launch over [num_requests, next_n] replaces the + # per-token sequential loop. The kernel iterates each request's + # tokens in position order internally, preserving the + # pool-completion read-after-stash dependency that the loop + # provided. Inputs are already grouped per request (uniform: + # view; non-uniform: _scatter_decode_tokens_by_request padded to + # [B, lmax]) — no per-token .contiguous() copies needed. + kpool_ops.kpool_decode_update_and_maybe_write_cache_batched( + kv_cache_raw, + tail_kv_cache, + dec_tail_slot, + dec_k, + dec_gate, + compress_ape, + dec_slot, + dec_pos, + index_kpool, + head_dim, + round_scale=(scale_fmt is not None), + ) + if current_platform.is_cuda_alike() and _fill_short_decode_causal_indices( + topk_indices_buffer, + positions, + num_decode_tokens, + attn_metadata_narrowed.max_seq_len, + topk_tokens, + ): + return topk_indices_buffer + decode_lens = decode_metadata.decode_lens + if decode_metadata.requires_padding: + # Padding also covers short chunked prefills classified as decode. + # MXFP4 uses zero-byte padding so padded slots dequantize to zero. + if q_scale is not None: + padded_q_quant_decode_tokens = pack_seq_triton( + q_quant[:num_decode_tokens], decode_lens, pad_value=0 + ) + padded_q_scale = pack_seq_triton( + q_scale[:num_decode_tokens], decode_lens, pad_value=0 + ) + else: + padded_q_quant_decode_tokens = pack_seq_triton( + q_quant[:num_decode_tokens], decode_lens + ) + padded_q_scale = None + padded_weights = pack_seq_triton( + weights[:num_decode_tokens], decode_lens, pad_value=0 + ).reshape(-1, *weights.shape[1:]) + else: + padded_q_quant_decode_tokens = q_quant[:num_decode_tokens].reshape( + decode_lens.shape[0], -1, *q_quant.shape[1:] + ) + if q_scale is not None: + padded_q_scale = q_scale[:num_decode_tokens].reshape( + decode_lens.shape[0], -1, *q_scale.shape[1:] + ) + else: + padded_q_scale = None + padded_weights = weights[:num_decode_tokens] + # TODO: move and optimize below logic with triton kernels + batch_size = padded_q_quant_decode_tokens.shape[0] + next_n = padded_q_quant_decode_tokens.shape[1] + num_padded_tokens = batch_size * next_n + seq_lens = decode_metadata.seq_lens[:batch_size] + # seq_lens is always 2D: (B, next_n) for native spec decode, (B, 1) + # otherwise. deep_gemm fp8_fp4_paged_mqa_logits requires 2D context_lens; + # the downstream topk kernels accept both 1D and 2D. + padded_q_quant_cast = ( + padded_q_quant_decode_tokens.view(torch.int8) + if use_fp4_cache + else padded_q_quant_decode_tokens + ) + if current_platform.is_rocm(): + from vllm.v1.attention.ops.rocm_aiter_mla_sparse import ( + rocm_fp8_paged_mqa_logits, + ) + + assert padded_q_scale is None + logits = rocm_fp8_paged_mqa_logits( + padded_q_quant_cast, + kv_cache, + padded_weights[:num_padded_tokens], + seq_lens, + decode_metadata.block_table, + decode_metadata.schedule_metadata, + max_model_len=max_model_len, + ) + else: + from vllm.utils.deep_gemm import fp8_fp4_paged_mqa_logits + + logits = fp8_fp4_paged_mqa_logits( + (padded_q_quant_cast, padded_q_scale), + kv_cache, + padded_weights[:num_padded_tokens], + seq_lens, + decode_metadata.block_table, + decode_metadata.schedule_metadata, + max_model_len=max_model_len, + clean_logits=False, + ) + num_rows = logits.shape[0] + # kpool: logits are pool-granular -> select topk_tokens//kpool pools, + # then expand each pool back to its kpool tokens. + select_k = topk_tokens // index_kpool if index_kpool > 1 else topk_tokens + if index_kpool > 1: + pool_topk = torch.empty( + (num_rows, select_k), dtype=torch.int32, device=logits.device + ) + topk_dst = pool_topk + else: + topk_dst = topk_indices_buffer[:num_padded_tokens, :topk_tokens] + + if current_platform.is_cuda() and select_k in (512, 1024, 2048): + workspace_manager = current_workspace_manager() + (topk_workspace,) = workspace_manager.get_simultaneous( + ((RADIX_TOPK_WORKSPACE_SIZE,), torch.uint8), + ) + torch.ops._C.persistent_topk( + logits, + seq_lens, + topk_dst, + topk_workspace, + select_k, + attn_metadata_narrowed.max_seq_len, + ) + else: + if current_platform.is_xpu(): + xpu_ops.top_k_per_row_decode( # type: ignore[attr-defined] + logits, + next_n, + seq_lens, + topk_dst, + num_rows, + logits.stride(0), + logits.stride(1), + select_k, + ) + else: + torch.ops._C.top_k_per_row_decode( + logits, + next_n, + seq_lens, + topk_dst, + num_rows, + logits.stride(0), + logits.stride(1), + select_k, + ) + + # Resolve to token-level indices in the output buffer. + if index_kpool > 1: + pool_ids = pool_topk.to(torch.int64) + n = pool_topk.shape[0] + # Decode seq_lens are pool-granular; recover token lengths from + # positions using the padded [B, next_n] row layout when needed. + if positions is not None: + dec_seq = _decode_topk_seq_lens( + positions, + decode_lens, + num_decode_tokens, + batch_size, + next_n, + decode_metadata.requires_padding, + ) + else: + dec_seq = decode_metadata.seq_lens[:n] + if dec_seq.ndim == 2: + dec_seq = dec_seq[:, -1] + dec_seq = dec_seq.to(torch.int32) + out = kpool_ops.expand_pools_and_append_tail(pool_ids, dec_seq, index_kpool) + else: + out = topk_dst + + if decode_metadata.requires_padding: + # Drop padded query rows introduced by the next_n padding above. + out = unpack_seq_triton( + out.reshape(batch_size, -1, out.shape[-1]), decode_lens + ) + topk_indices_buffer[: out.shape[0], : out.shape[-1]] = out + + return topk_indices_buffer + + +@CustomOp.register("sparse_attn_indexer_kpool") +class SparseAttnIndexerKpool(CustomOp): + """Sparse Attention Indexer Custom Op Layer. This layer is extracted as a + separate custom op since it involves heavy custom kernels like `mqa_logits`, + `paged_mqa_logits` and `top_k_per_row`, etc. Those kernels maybe requires + specific memory layout or implementation for different hardware backends to + achieve optimal performance. + + For now, the default native path will use CUDA backend path. Other platform + may requires add the corresponding Custom Op name `sparse_attn_indexer` to + `custom_ops` in `CompilationConfig` to enable the platform specific path. + """ + + def __init__( + self, + k_cache, + quant_block_size: int, + scale_fmt: str, + topk_tokens: int, + head_dim: int, + max_model_len: int, + max_total_seq_len: int, + topk_indices_buffer: torch.Tensor, + skip_k_cache_insert: bool = False, + use_fp4_cache: bool = False, + tail_cache=None, + ): + super().__init__() + self.k_cache = k_cache + self.tail_cache = tail_cache + self.quant_block_size = quant_block_size + self.scale_fmt = scale_fmt + self.topk_tokens = topk_tokens + self.head_dim = head_dim + self.max_model_len = max_model_len + self.max_total_seq_len = max_total_seq_len + self.topk_indices_buffer = topk_indices_buffer + self.skip_k_cache_insert = skip_k_cache_insert + self.use_fp4_cache = use_fp4_cache + if current_platform.is_cuda() and not has_deep_gemm(): + raise RuntimeError( + "Sparse Attention Indexer CUDA op requires DeepGEMM to be installed." + ) + + def forward_native( + self, + hidden_states: torch.Tensor, + q_quant: torch.Tensor | tuple[torch.Tensor, torch.Tensor], + k: torch.Tensor, + weights: torch.Tensor, + *, + gate_score: torch.Tensor | None = None, + compress_ape: torch.Tensor | None = None, + index_kpool: int = 1, + positions: torch.Tensor | None = None, + ): + if current_platform.is_cuda() or current_platform.is_xpu(): + return self.forward_cuda( + hidden_states, + q_quant, + k, + weights, + gate_score=gate_score, + compress_ape=compress_ape, + index_kpool=index_kpool, + positions=positions, + ) + elif current_platform.is_rocm(): + return self.forward_hip( + hidden_states, + q_quant, + k, + weights, + gate_score=gate_score, + compress_ape=compress_ape, + index_kpool=index_kpool, + positions=positions, + ) + else: + raise NotImplementedError( + "SparseAttnIndexer native forward is only implemented for " + "CUDA, ROCm and XPU platforms." + ) + + def forward_cuda( + self, + hidden_states: torch.Tensor, + q_quant: torch.Tensor | tuple[torch.Tensor, torch.Tensor], + k: torch.Tensor, + weights: torch.Tensor, + *, + gate_score: torch.Tensor | None = None, + compress_ape: torch.Tensor | None = None, + index_kpool: int = 1, + positions: torch.Tensor | None = None, + ): + # FP8 path: single tensor (per-token scale is folded into `weights`). + # FP4 path: (values, scales) tuple with scales required by the kernel. + if isinstance(q_quant, tuple): + q_values, q_scale = q_quant + else: + q_values, q_scale = q_quant, None + return sparse_attn_indexer_kpool( + hidden_states, + self.k_cache.prefix, + self.k_cache.kv_cache, + q_values, + q_scale, + k, + weights, + self.quant_block_size, + self.scale_fmt, + self.topk_tokens, + self.head_dim, + self.max_model_len, + self.max_total_seq_len, + self.topk_indices_buffer, + self.skip_k_cache_insert, + self.use_fp4_cache, + gate_score, + compress_ape, + index_kpool, + positions, + self.tail_cache.kv_cache if self.tail_cache is not None else None, + self.tail_cache.prefix if self.tail_cache is not None else None, + ) + + def forward_hip( + self, + hidden_states: torch.Tensor, + q_quant: torch.Tensor | tuple[torch.Tensor, torch.Tensor], + k: torch.Tensor, + weights: torch.Tensor, + *, + gate_score: torch.Tensor | None = None, + compress_ape: torch.Tensor | None = None, + index_kpool: int = 1, + positions: torch.Tensor | None = None, + ): + assert not self.use_fp4_cache, "AMD platform doesn't support fp4 cache yet" + assert isinstance(q_quant, torch.Tensor), ( + "AMD sparse_attn_indexer expects a single FP8 q_quant tensor" + ) + if rocm_aiter_ops.is_enabled(): + if index_kpool <= 1: + return torch.ops.vllm.rocm_aiter_sparse_attn_indexer( + hidden_states, + _encode_layer_name(self.k_cache.prefix), + self.k_cache.kv_cache, + q_quant, + k, + weights, + self.quant_block_size, + self.scale_fmt, + self.topk_tokens, + self.head_dim, + self.max_model_len, + self.max_total_seq_len, + self.topk_indices_buffer, + skip_k_cache_insert=self.skip_k_cache_insert, + ) + return self.forward_cuda( + hidden_states, + q_quant, + k, + weights, + gate_score=gate_score, + compress_ape=compress_ape, + index_kpool=index_kpool, + positions=positions, + ) + raise RuntimeError( + "Sparse attention indexer ROCm path is only supported on AITER. " + "Please enable aiter with VLLM_ROCM_USE_AITER=1" + ) diff --git a/vllm/model_executor/models/registry.py b/vllm/model_executor/models/registry.py index 9f77ea83e4b3..a0272fcc053c 100644 --- a/vllm/model_executor/models/registry.py +++ b/vllm/model_executor/models/registry.py @@ -120,6 +120,7 @@ "Glm4MoeForCausalLM": ("glm4_moe", "Glm4MoeForCausalLM"), "Glm4MoeLiteForCausalLM": ("glm4_moe_lite", "Glm4MoeLiteForCausalLM"), "GlmMoeDsaForCausalLM": ("vllm.models.deepseek_v32", "GlmMoeDsaForCausalLM"), + "Glm5NextForCausalLM": ("vllm.models.glm5next", "Glm5NextForCausalLM"), "GptOssForCausalLM": ("gpt_oss", "GptOssForCausalLM"), "GPT2LMHeadModel": ("gpt2", "GPT2LMHeadModel"), "GPTJForCausalLM": ("gpt_j", "GPTJForCausalLM"), @@ -412,6 +413,10 @@ "Glm4vForConditionalGeneration": ("glm4_1v", "Glm4vForConditionalGeneration"), "Glm4vMoeForConditionalGeneration": ("glm4_1v", "Glm4vMoeForConditionalGeneration"), "GlmOcrForConditionalGeneration": ("glm_ocr", "GlmOcrForConditionalGeneration"), + "Glm5NextForConditionalGeneration": ( + "vllm.models.glm5next", + "Glm5NextForConditionalGeneration", + ), "GraniteSpeechForConditionalGeneration": ( "granite_speech", "GraniteSpeechForConditionalGeneration", @@ -665,6 +670,7 @@ "Glm4MoeMTPModel": ("glm4_moe_mtp", "Glm4MoeMTP"), "Glm4MoeLiteMTPModel": ("glm4_moe_lite_mtp", "Glm4MoeLiteMTP"), "GlmOcrMTPModel": ("glm_ocr_mtp", "GlmOcrMTP"), + "Glm5NextMTPModel": ("vllm.models.glm5next", "Glm5NextMTP"), "MedusaModel": ("medusa", "Medusa"), "OpenPanguMTPModel": ("openpangu_mtp", "OpenPanguMTP"), "Qwen3NextMTP": ("qwen3_next_mtp", "Qwen3NextMTP"), diff --git a/vllm/model_executor/warmup/kernel_warmup.py b/vllm/model_executor/warmup/kernel_warmup.py index 3b2596c5e322..ddf375f9370f 100644 --- a/vllm/model_executor/warmup/kernel_warmup.py +++ b/vllm/model_executor/warmup/kernel_warmup.py @@ -41,6 +41,9 @@ from vllm.model_executor.warmup.replayssm_warmup import ( replayssm_autotune_warmup, ) +from vllm.model_executor.warmup.spec_decode_rejection_warmup import ( + spec_decode_rejection_warmup, +) from vllm.platforms import current_platform from vllm.utils.deep_gemm import is_deep_gemm_supported from vllm.utils.flashinfer import has_flashinfer @@ -163,6 +166,7 @@ def kernel_warmup(worker: "Worker", *, process_local_only: bool = False): if worker.vllm_config.kernel_config.enable_jit_warmup: kimi_k3_triton_warmup(worker) fa4_cutedsl_warmup(worker) + spec_decode_rejection_warmup(worker) qwen4_exp_qsa_triton_warmup(worker) if current_platform.has_device_capability(90): diff --git a/vllm/model_executor/warmup/spec_decode_rejection_warmup.py b/vllm/model_executor/warmup/spec_decode_rejection_warmup.py new file mode 100644 index 000000000000..78e12b6ff881 --- /dev/null +++ b/vllm/model_executor/warmup/spec_decode_rejection_warmup.py @@ -0,0 +1,111 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +"""Warm up spec-decode rejection-sampler Triton kernels. + +The rejection sampler kernels (``_compute_local_logits_stats_kernel``, +``_rejection_kernel``, ``_resample_kernel``) are JIT-compiled by Triton on +first use. Without warmup, the first spec-decode request pays a multi-second +compilation cost. This pre-compiles them with dummy data matching the +server's vocab size and speculative config. +""" + +from __future__ import annotations + +from typing import TYPE_CHECKING + +import torch + +from vllm.logger import init_logger + +if TYPE_CHECKING: + from vllm.v1.worker.gpu_worker import Worker + +logger = init_logger(__name__) + + +@torch.inference_mode() +def spec_decode_rejection_warmup(worker: Worker) -> None: + spec_config = worker.vllm_config.speculative_config + if spec_config is None: + return + + from vllm.v1.worker.gpu.spec_decode.rejection_sampler_utils import ( + rejection_sample, + ) + + model_config = worker.vllm_config.model_config + vocab_size = model_config.get_vocab_size() + num_spec = spec_config.num_speculative_tokens + if num_spec <= 0 or vocab_size <= 0: + return + + # Mirror the constexpr-relevant flags the runtime uses. + rejection_method = getattr(spec_config, "rejection_sample_method", None) + use_block_verification = rejection_method == "block" + use_synthetic = rejection_method == "synthetic" + + device = torch.device("cuda") + num_reqs = 1 + tokens_per_req = num_spec + 1 + num_logits = num_reqs * tokens_per_req + + # Triton JIT-specializes on tensor dtypes. The target logits may be fp32 + # (apply_sampling_params copies to fp32 when processing is needed) or the + # model dtype (pass-through otherwise), while draft logits are always the + # model dtype. Warm every (target, draft) combination the runtime can hit. + model_dtype = model_config.dtype + warmup_dtype_pairs = { + (model_dtype, model_dtype), + (torch.float32, torch.float32), + (torch.float32, model_dtype), + (model_dtype, torch.float32), + } + + logger.info( + "Warming up spec-decode rejection sampler kernels " + "(vocab=%d, num_spec=%d, dtype_pairs=%s, block_verify=%s).", + vocab_size, + num_spec, + [(str(t), str(d)) for t, d in warmup_dtype_pairs], + use_block_verification, + ) + for tgt_dtype, draft_dtype in warmup_dtype_pairs: + target_logits = torch.zeros( + (num_logits, vocab_size), dtype=tgt_dtype, device=device + ) + draft_logits = torch.zeros( + (num_reqs, num_spec, vocab_size), dtype=draft_dtype, device=device + ) + synthetic_rates = ( + torch.full((num_spec,), 0.5, dtype=torch.float32, device=device) + if use_synthetic + else None + ) + try: + rejection_sample( + target_logits=target_logits, + draft_logits=draft_logits, + draft_sampled=torch.zeros(num_logits, dtype=torch.int64, device=device), + cu_num_logits=torch.tensor( + [0, num_logits], dtype=torch.int32, device=device + ), + pos=torch.zeros(num_logits, dtype=torch.int64, device=device), + idx_mapping=torch.zeros(num_reqs, dtype=torch.int32, device=device), + expanded_idx_mapping=torch.zeros( + num_logits, dtype=torch.int32, device=device + ), + expanded_local_pos=torch.arange( + num_logits, dtype=torch.int32, device=device + ), + temperature=torch.zeros(num_reqs, dtype=torch.float32, device=device), + seed=torch.full((num_reqs,), 42, dtype=torch.int64, device=device), + num_speculative_steps=num_spec, + synthetic_conditional_rates=synthetic_rates, + use_fp64=False, + use_block_verification=use_block_verification, + ) + except Exception: + logger.warning( + "Skipping spec-decode rejection sampler warmup.", exc_info=True + ) + return diff --git a/vllm/models/glm5next/__init__.py b/vllm/models/glm5next/__init__.py new file mode 100644 index 000000000000..a1146e2602eb --- /dev/null +++ b/vllm/models/glm5next/__init__.py @@ -0,0 +1,16 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project + +from vllm.platforms import current_platform + +if current_platform.is_xpu(): + raise NotImplementedError("GLM-5.3-Flash does not currently support XPU.") + +from .nvidia.model import Glm5NextForCausalLM, Glm5NextForConditionalGeneration +from .nvidia.mtp import Glm5NextMTP + +__all__ = [ + "Glm5NextForCausalLM", + "Glm5NextForConditionalGeneration", + "Glm5NextMTP", +] diff --git a/vllm/models/glm5next/amd/__init__.py b/vllm/models/glm5next/amd/__init__.py new file mode 100644 index 000000000000..208f01a7cb5e --- /dev/null +++ b/vllm/models/glm5next/amd/__init__.py @@ -0,0 +1,2 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project diff --git a/vllm/models/glm5next/amd/ops/__init__.py b/vllm/models/glm5next/amd/ops/__init__.py new file mode 100644 index 000000000000..208f01a7cb5e --- /dev/null +++ b/vllm/models/glm5next/amd/ops/__init__.py @@ -0,0 +1,2 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project diff --git a/vllm/models/glm5next/amd/ops/kpool_compress.py b/vllm/models/glm5next/amd/ops/kpool_compress.py new file mode 100644 index 000000000000..7afc67e99e33 --- /dev/null +++ b/vllm/models/glm5next/amd/ops/kpool_compress.py @@ -0,0 +1,890 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +"""kpool (key-pooling) Triton kernels for the sparse-attention indexer. + +The cache stores POOLS (1 entry per ``pool_size`` consecutive tokens) rather +than individual tokens. ``compress_ratio == pool_size`` on the kv_cache_spec +makes the metadata builder emit pool-granular slot_mapping / seq_lens / +cu_seq_lens / page_table for free; this file supplies the compress-write +kernel (replacing ``indexer_k_quant_and_cache``) and the pool-level topk +helpers (select pools -> expand to tokens -> append tail). +""" + +from __future__ import annotations + +import torch + +from vllm.platforms import current_platform +from vllm.triton_utils import tl, triton + +# The GLM-5.3-Flash indexer head dimension is fixed at 128. +INDEX_HEAD_DIM = 128 +FP8_DTYPE = current_platform.fp8_dtype() +FP8_MAX = torch.finfo(FP8_DTYPE).max + + +@triton.jit +def _cache_k_offset( + token_offset, + dim_offset, + head_dim: tl.constexpr, + preshuffle: tl.constexpr, +): + if preshuffle: + return ( + (token_offset // 16) * 16 * head_dim + + (dim_offset // 16) * 16 * 16 + + (token_offset % 16) * 16 + + dim_offset % 16 + ) + return token_offset * head_dim + dim_offset + + +# Hadamard-128 rotation + + +@triton.jit +def _hadamard128_stage(x, GROUPS: tl.constexpr, STRIDE: tl.constexpr): + x3 = tl.reshape(x, (GROUPS, 2, STRIDE)) + x3 = tl.trans(x3, 0, 2, 1) + a, b = tl.split(x3) + x3 = tl.join(a + b, a - b) + x3 = tl.trans(x3, 0, 2, 1) + return tl.reshape(x3, (128,)) + + +@triton.jit +def _hadamard128(x): + x = _hadamard128_stage(x, 64, 1) + x = _hadamard128_stage(x, 32, 2) + x = _hadamard128_stage(x, 16, 4) + x = _hadamard128_stage(x, 8, 8) + x = _hadamard128_stage(x, 4, 16) + x = _hadamard128_stage(x, 2, 32) + x = _hadamard128_stage(x, 1, 64) + return x * 0.08838834764831845 # 1/sqrt(128) + + +# Map pool ids to physical cache slots. + + +def compute_pooled_write_locs( + page_table_64: torch.Tensor, + pool_ids: torch.Tensor, + pool_size: int, +) -> torch.Tensor: + """Map logical pooled-K ids to physical flat cache slots. + + ``pool_size`` consecutive tokens share one pool slot that lives at the + *first* token page of each page-group. ``page_table_64`` maps token pages + to physical block ids; we gather the block id of each pool's page-group + and add the in-block pool offset. + """ + assert page_table_64.ndim == 1 + pool_ids = pool_ids.to(torch.int64) + block_size = 64 + pool_page_group = torch.div(pool_ids, block_size, rounding_mode="floor") + token_page_row = pool_page_group * pool_size + packed_page = page_table_64.index_select(0, token_page_row.to(torch.int64)) + return packed_page.to(torch.int64) * block_size + torch.remainder( + pool_ids, block_size + ) + + +def build_pooled_page_table( + page_table: torch.Tensor, + pool_size: int, +) -> torch.Tensor: + """Build a pool-granular page table by taking every ``pool_size``-th + token-page column (one pool maps to ``pool_size`` token pages). + + Uses gather (not strided slicing) so the result is always a fresh + row-major tensor — some downstream kernels require stride(-1) == 1. + """ + block_size = page_table.shape[-1] + assert block_size % pool_size == 0, ( + f"pool_size ({pool_size}) must divide page columns ({block_size})" + ) + idx = torch.arange(0, block_size, pool_size, device=page_table.device) + return page_table[..., idx].contiguous() + + +# Fused pool compression and cache write. + + +@triton.jit +def _kpool_softmax_rotate_write_cache_kernel( + buf_fp8_ptr, + buf_fp32_ptr, + slot_k_ptr, + slot_score_ptr, + ape_ptr, + loc_ptr, + write_mask_ptr, + compressed_k_ptr, + compressed_scale_ptr, + slot_k_stride_0, + slot_k_stride_1, + slot_score_stride_0, + slot_score_stride_1, + ape_stride_0, + PAGE_SIZE: tl.constexpr, + BUF_NUMEL_PER_PAGE: tl.constexpr, + POOL_SIZE: tl.constexpr, + HEAD_DIM: tl.constexpr, + S_OFFSET_NBYTES_IN_PAGE: tl.constexpr, + FP8_MAX: tl.constexpr, + PRESHUFFLE: tl.constexpr, + ROUND_SCALE: tl.constexpr, + HAS_WRITE_MASK: tl.constexpr, + RETURN_COMPRESSED: tl.constexpr, + WRITE_CACHE: tl.constexpr, + BLOCK_D: tl.constexpr, +): + """One program per pool. softmax(slot_score+ape)-weighted sum of slot_k -> + Hadamard-128 -> per-vector fp8 absmax quant -> write to cache at ``loc``.""" + row = tl.program_id(0) + do_write = True + if HAS_WRITE_MASK: + do_write = tl.load(write_mask_ptr + row) + + offs = tl.arange(0, BLOCK_D) + mask = (offs < HEAD_DIM) & do_write + + # --- Pass 1: per-dim max over the pool (softmax numerical stability) --- + max_score = tl.full((BLOCK_D,), -float("inf"), tl.float32) + for slot in tl.static_range(0, POOL_SIZE): + score = tl.load( + slot_score_ptr + + row * slot_score_stride_0 + + slot * slot_score_stride_1 + + offs, + mask=mask, + other=0.0, + ).to(tl.float32) + score += tl.load(ape_ptr + slot * ape_stride_0 + offs, mask=mask, other=0.0).to( + tl.float32 + ) + max_score = tl.maximum(max_score, score) + + # --- Pass 2: softmax-weighted sum of K --- + acc = tl.full((BLOCK_D,), 0.0, tl.float32) + denom = tl.full((BLOCK_D,), 0.0, tl.float32) + for slot in tl.static_range(0, POOL_SIZE): + score = tl.load( + slot_score_ptr + + row * slot_score_stride_0 + + slot * slot_score_stride_1 + + offs, + mask=mask, + other=0.0, + ).to(tl.float32) + score += tl.load(ape_ptr + slot * ape_stride_0 + offs, mask=mask, other=0.0).to( + tl.float32 + ) + prob = tl.exp(score - max_score) + denom += prob + k = tl.load( + slot_k_ptr + row * slot_k_stride_0 + slot * slot_k_stride_1 + offs, + mask=mask, + other=0.0, + ).to(tl.float32) + acc += k * prob + + x = acc / denom + x = tl.where(do_write, x, 0.0).to(tl.bfloat16).to(tl.float32) + + # Match the unfused pooled-K path's bf16 precision before quantization. + x = _hadamard128(x).to(tl.bfloat16).to(tl.float32) + + # --- per-vector absmax fp8 quant --- + fp8_max_inv = 1.0 / FP8_MAX + absmax = tl.max(tl.abs(x), axis=0) + absmax = tl.maximum(absmax, 1e-4) + if ROUND_SCALE: + scale = tl.exp2(tl.ceil(tl.log2(absmax * fp8_max_inv))) + else: + scale = absmax * fp8_max_inv + quantized = x / scale + quantized = tl.minimum(tl.maximum(quantized, -FP8_MAX), FP8_MAX) + + if WRITE_CACHE: + loc = tl.load(loc_ptr + row, mask=do_write, other=0) + loc_page_index = loc // PAGE_SIZE + loc_token_offset_in_page = loc % PAGE_SIZE + out_k_offsets = loc_page_index * BUF_NUMEL_PER_PAGE + _cache_k_offset( + loc_token_offset_in_page, + offs, + HEAD_DIM, + PRESHUFFLE, + ) + out_s_offset = ( + loc_page_index * BUF_NUMEL_PER_PAGE // 4 + + S_OFFSET_NBYTES_IN_PAGE // 4 + + loc_token_offset_in_page + ) + tl.store(buf_fp8_ptr + out_k_offsets, quantized, mask=mask) + tl.store(buf_fp32_ptr + out_s_offset, scale, mask=do_write) + + if RETURN_COMPRESSED: + tl.store( + compressed_k_ptr + row * HEAD_DIM + offs, + quantized, + mask=offs < HEAD_DIM, + ) + tl.store(compressed_scale_ptr + row, scale) + + +def kpool_compress_and_write_cache( + kv_cache: torch.Tensor, + slot_k: torch.Tensor, + slot_score: torch.Tensor, + ape: torch.Tensor, + loc: torch.Tensor, + pool_size: int, + head_dim: int = INDEX_HEAD_DIM, + write_mask: torch.Tensor | None = None, + round_scale: bool = True, + return_compressed: bool = False, + write_cache: bool = True, +): + """Compress ``pool_size`` tokens into one fp8 K and write at ``loc``. + + Args: + kv_cache: indexer K cache ``[num_blocks, block_size, head_dim+4]`` uint8. + slot_k: ``[n_pools, pool_size, head_dim]`` bf16 — raw per-token K. + slot_score: ``[n_pools, pool_size, head_dim]`` — per-token gate score. + ape: ``[pool_size, head_dim]`` fp32 — per-slot position bias. + loc: ``[n_pools]`` int64 — flat physical slot per pool. + """ + assert slot_k.ndim == 3 + assert slot_score.shape == slot_k.shape + assert ape.shape == slot_k.shape[1:] + assert slot_k.shape[2] == head_dim + assert slot_k.dtype == torch.bfloat16 + assert ape.dtype == torch.float32 + assert kv_cache.dtype == torch.uint8 + assert loc.dtype == torch.int64 + assert write_cache or return_compressed + + page_size = kv_cache.shape[1] + buf = kv_cache + slot_k = slot_k.contiguous() + slot_score = slot_score.contiguous() + ape = ape.contiguous() + loc = loc.contiguous() + if write_mask is None: + write_mask = torch.empty((1,), dtype=torch.bool, device=slot_k.device) + has_write_mask = False + else: + assert write_mask.shape == (slot_k.shape[0],) + write_mask = write_mask.contiguous() + has_write_mask = True + assert not return_compressed + + if slot_k.shape[0] == 0: + if return_compressed: + return ( + torch.empty( + (0, head_dim), + dtype=FP8_DTYPE, + device=slot_k.device, + ), + torch.empty((0,), dtype=torch.float32, device=slot_k.device), + ) + return None + + buf_fp8 = buf.view(FP8_DTYPE) + buf_fp32 = buf.view(torch.float32) + # bytes per page (last dim of kv_cache) viewed as uint8 + buf_numel_per_page = buf.stride(0) + s_offset_nbytes_in_page = page_size * head_dim + + if return_compressed: + compressed_k = torch.empty( + (slot_k.shape[0], head_dim), + dtype=FP8_DTYPE, + device=slot_k.device, + ) + compressed_scale = torch.empty( + (slot_k.shape[0],), dtype=torch.float32, device=slot_k.device + ) + else: + compressed_k = buf_fp8 + compressed_scale = buf_fp32 + + if page_size > 1: + assert page_size % 16 == 0, "ROCm preshuffle requires 16-token tiles" + + _kpool_softmax_rotate_write_cache_kernel[(slot_k.shape[0],)]( + buf_fp8, + buf_fp32, + slot_k, + slot_score, + ape, + loc, + write_mask, + compressed_k, + compressed_scale, + slot_k.stride(0), + slot_k.stride(1), + slot_score.stride(0), + slot_score.stride(1), + ape.stride(0), + PAGE_SIZE=page_size, + BUF_NUMEL_PER_PAGE=buf_numel_per_page, + POOL_SIZE=slot_k.shape[1], + HEAD_DIM=head_dim, + S_OFFSET_NBYTES_IN_PAGE=s_offset_nbytes_in_page, + FP8_MAX=FP8_MAX, + PRESHUFFLE=page_size > 1, + ROUND_SCALE=round_scale, + HAS_WRITE_MASK=has_write_mask, + RETURN_COMPRESSED=return_compressed, + WRITE_CACHE=write_cache, + BLOCK_D=triton.next_power_of_2(head_dim), + ) + + if return_compressed: + return compressed_k, compressed_scale + return None + + +# Seed each request's incomplete pool into its paged tail during prefill. + + +@triton.jit +def _kpool_tail_seed_kernel( + key_ptr, + score_ptr, + tslot_ptr, + tail_ptr, + n_tokens, + TAIL_BLOCK_ELEMS: tl.constexpr, + KPOOL_HEAD: tl.constexpr, + HEAD_DIM: tl.constexpr, + KPOOL: tl.constexpr, + BLOCK_D: tl.constexpr, +): + """Copy token ``i``'s raw K + gate into its request's tail block. + + Token ``i`` is among its request's last KPOOL tokens iff the token KPOOL + ahead belongs to a different tail block (or is past the batch / padding, + slot < 0). ``tslot = block * KPOOL + pos % KPOOL``; the destination is + ``tail[block, {0:K, 1:score}, pos % KPOOL, :]``. + """ + i = tl.program_id(0) + t = tl.load(tslot_ptr + i).to(tl.int64) + if t < 0: + return + blk = t // KPOOL # t >= 0 here, so trunc == floor + ahead = tl.load(tslot_ptr + i + KPOOL, mask=i + KPOOL < n_tokens, other=-1).to( + tl.int64 + ) + # Match the torch semantics exactly: a negative ahead slot floors to a + # block id that differs from every real block -> token is in the tail. + # Only divide non-negative slots (Triton int div truncates, torch floors). + if ahead >= 0 and ahead // KPOOL == blk: + return + offs = tl.arange(0, BLOCK_D) + m = offs < HEAD_DIM + block_base = blk * TAIL_BLOCK_ELEMS + base = block_base + (t % KPOOL) * HEAD_DIM + k = tl.load(key_ptr + i * HEAD_DIM + offs, mask=m) + s = tl.load(score_ptr + i * HEAD_DIM + offs, mask=m) + tl.store(tail_ptr + base + offs, k, mask=m) + tl.store( + tail_ptr + block_base + KPOOL_HEAD + (t % KPOOL) * HEAD_DIM + offs, s, mask=m + ) + + +def kpool_seed_tail_cache( + tail_kv_cache: torch.Tensor, + key: torch.Tensor, + gate_score: torch.Tensor, + tslot: torch.Tensor, + kpool: int, + head_dim: int = INDEX_HEAD_DIM, +) -> None: + """Seed the paged tail cache from a prefill batch (see the kernel).""" + assert tail_kv_cache.dtype == torch.bfloat16 + assert key.dtype == torch.bfloat16 + n = tslot.shape[0] + if n == 0: + return + _kpool_tail_seed_kernel[(n,)]( + key, + gate_score, + tslot, + tail_kv_cache, + n, + TAIL_BLOCK_ELEMS=tail_kv_cache.stride(0), + KPOOL_HEAD=tail_kv_cache.stride(1), + HEAD_DIM=head_dim, + KPOOL=kpool, + BLOCK_D=triton.next_power_of_2(head_dim), + ) + + +# Update each request's tail during decode and write completed pools. + + +@triton.jit +def _kpool_decode_update_batched_kernel( + buf_fp8_ptr, + buf_fp32_ptr, + tail_kv_ptr, + tail_slot_mapping_ptr, # [B, NEXT_N] int32 + key_ptr, # [B, NEXT_N, HEAD_DIM] bf16 + key_stride_b, + key_stride_t, + slot_score_ptr, # [B, NEXT_N, HEAD_DIM] bf16 + ss_stride_b, + ss_stride_t, + ape_ptr, + ape_stride_0, + slot_mapping_ptr, # [B, NEXT_N] int32 + positions_ptr, # [B, NEXT_N] int32 + NEXT_N, # runtime token count per request (no .item() needed) + PAGE_SIZE: tl.constexpr, + BUF_NUMEL_PER_PAGE: tl.constexpr, + POOL_SIZE: tl.constexpr, + TAIL_BLOCK_ELEMS: tl.constexpr, + KPOOL_HEAD: tl.constexpr, + HEAD_DIM: tl.constexpr, + S_OFFSET_NBYTES_IN_PAGE: tl.constexpr, + FP8_MAX: tl.constexpr, + PRESHUFFLE: tl.constexpr, + ROUND_SCALE: tl.constexpr, + BLOCK_D: tl.constexpr, +): + """One program per request; iterates its NEXT_N verify tokens in order. + + Replaces the caller's per-token sequential launch loop. The intra-request + iteration MUST stay in position order: a pool-completion at token t* reads + the tail-ring slots that tokens t < t* (same request) just stashed in this + same invocation. ``tl.range`` iterates sequentially within the program, so + those stashes are visible to the later completion read. Cross-request + programs are independent (distinct tail blocks). With NEXT_N < POOL_SIZE + (the spec-verify case: NEXT_N ~= num_spec+1, POOL_SIZE=16) at most one + completion can occur per request per call, but the ordered loop is correct + for any NEXT_N. + """ + req = tl.program_id(0) + offs = tl.arange(0, BLOCK_D) + dim_mask = offs < HEAD_DIM + + for t in tl.range(0, NEXT_N): + idx = req * NEXT_N + t + cache_loc = tl.load(slot_mapping_ptr + idx) + pos = tl.load(positions_ptr + idx) + safe_pos = tl.maximum(pos, 0) + pos_valid = (cache_loc >= 0) & (pos >= 0) + + slot = safe_pos % POOL_SIZE + phys_slot = safe_pos % POOL_SIZE + + # Derive the tail block from THIS token's tail_slot (the request's block + # is constant across a pool, but a padded / invalid entry carries a + # negative sentinel -- reading it from token 0 would poison every + # token's base address). Clamp so an invalid entry can never form an + # out-of-bounds base; the accesses below are gated on pos_valid anyway. + tail_slot = tl.load(tail_slot_mapping_ptr + idx) + block = tl.maximum(tail_slot, 0).to(tl.int64) // POOL_SIZE + block_base = block * TAIL_BLOCK_ELEMS + + # The tail-ring stash must run for EVERY real token, so it is gated on + # the token-granular tail slot -- not on `pos_valid`, which keys off the + # POOL-granular `slot_mapping` and is therefore only true on the pool's + # last token. Gating the stash on pos_valid dropped every intra-pool + # token, so a decode-built pool compressed 3 stale ring entries (the + # prefill-seeded prompt tail, frozen forever) plus the current token. + stash_valid = (pos >= 0) & (tail_slot >= 0) + + key = tl.load( + key_ptr + req * key_stride_b + t * key_stride_t + offs, + mask=dim_mask, + other=0.0, + ).to(tl.float32) + score_current = tl.load( + slot_score_ptr + req * ss_stride_b + t * ss_stride_t + offs, + mask=dim_mask, + other=0.0, + ).to(tl.float32) + + if pos_valid & (slot == POOL_SIZE - 1): + pool_logical_start = safe_pos - slot + + max_score = tl.full((BLOCK_D,), -float("inf"), tl.float32) + for pool_slot in tl.static_range(0, POOL_SIZE): + is_current = pool_slot == slot + phys = (pool_logical_start + pool_slot) % POOL_SIZE + score_buf = tl.load( + tail_kv_ptr + block_base + KPOOL_HEAD + phys * HEAD_DIM + offs, + mask=dim_mask, + other=0.0, + ).to(tl.float32) + score = tl.where(is_current, score_current, score_buf) + score += tl.load( + ape_ptr + pool_slot * ape_stride_0 + offs, + mask=dim_mask, + other=0.0, + ).to(tl.float32) + max_score = tl.maximum(max_score, score) + + acc = tl.full((BLOCK_D,), 0.0, tl.float32) + denom = tl.full((BLOCK_D,), 0.0, tl.float32) + for pool_slot in tl.static_range(0, POOL_SIZE): + is_current = pool_slot == slot + phys = (pool_logical_start + pool_slot) % POOL_SIZE + score_buf = tl.load( + tail_kv_ptr + block_base + KPOOL_HEAD + phys * HEAD_DIM + offs, + mask=dim_mask, + other=0.0, + ).to(tl.float32) + score = tl.where(is_current, score_current, score_buf) + score += tl.load( + ape_ptr + pool_slot * ape_stride_0 + offs, + mask=dim_mask, + other=0.0, + ).to(tl.float32) + prob = tl.exp(score - max_score) + denom += prob + k_buf = tl.load( + tail_kv_ptr + block_base + phys * HEAD_DIM + offs, + mask=dim_mask, + other=0.0, + ).to(tl.float32) + k = tl.where(is_current, key, k_buf) + acc += k * prob + + x = (acc / denom).to(tl.bfloat16).to(tl.float32) + x = _hadamard128(x).to(tl.bfloat16).to(tl.float32) + + fp8_max_inv = 1.0 / FP8_MAX + absmax = tl.maximum(tl.max(tl.abs(x), axis=0), 1e-4) + if ROUND_SCALE: + scale = tl.exp2(tl.ceil(tl.log2(absmax * fp8_max_inv))) + else: + scale = absmax * fp8_max_inv + quantized = tl.minimum(tl.maximum(x / scale, -FP8_MAX), FP8_MAX) + + loc = cache_loc.to(tl.int64) + loc_page_index = loc // PAGE_SIZE + loc_token_offset_in_page = loc % PAGE_SIZE + out_k_offsets = loc_page_index * BUF_NUMEL_PER_PAGE + _cache_k_offset( + loc_token_offset_in_page, + offs, + HEAD_DIM, + PRESHUFFLE, + ) + out_s_offset = ( + loc_page_index * BUF_NUMEL_PER_PAGE // 4 + + S_OFFSET_NBYTES_IN_PAGE // 4 + + loc_token_offset_in_page + ) + tl.store(buf_fp8_ptr + out_k_offsets, quantized, mask=dim_mask) + tl.store(buf_fp32_ptr + out_s_offset, scale) + + # Stash the current token AFTER any completion read so the completion + # uses prior stashes (and the current token's own key/score via + # is_current), then leaves this token for future pools. Order matches + # the per-token kernel: completion read first, stash second. + update_mask = dim_mask & stash_valid + tl.store( + tail_kv_ptr + block_base + phys_slot * HEAD_DIM + offs, + key, + mask=update_mask, + ) + tl.store( + tail_kv_ptr + block_base + KPOOL_HEAD + phys_slot * HEAD_DIM + offs, + score_current, + mask=update_mask, + ) + + +def kpool_decode_update_and_maybe_write_cache_batched( + kv_cache: torch.Tensor, + tail_kv_cache: torch.Tensor, + tail_slot_mapping: torch.Tensor, + key: torch.Tensor, + slot_score: torch.Tensor, + ape: torch.Tensor, + slot_mapping: torch.Tensor, + positions: torch.Tensor, + pool_size: int, + head_dim: int = INDEX_HEAD_DIM, + round_scale: bool = True, +) -> None: + """Batched decode-step kpool update for spec verify (``next_n > 1``). + + One launch replaces the caller's per-token loop. Inputs are grouped per + request: ``[num_requests, next_n, ...]``. Each program handles one + request's ``next_n`` tokens in position order (see the kernel docstring for + why ordering is required for pool-completion correctness). + + Plain decode (``next_n == 1``) is handled here too — the kernel collapses + to a single-iteration loop. + + Args: + kv_cache: indexer K cache ``[num_blocks, block_size, head_dim+4]`` uint8. + tail_kv_cache: paged tail cache ``[num_blocks, 2, pool_size, head_dim]`` + bf16 (K at half 0, gate score at half 1). + tail_slot_mapping: ``[num_requests, next_n]`` int32. + key / slot_score: ``[num_requests, next_n, head_dim]`` bf16. + ape: ``[pool_size, head_dim]`` fp32. + slot_mapping / positions: ``[num_requests, next_n]`` int32. + """ + num_requests, next_n = key.shape[0], key.shape[1] + if num_requests == 0 or next_n == 0: + return + assert tail_kv_cache.ndim == 4 + assert tail_kv_cache.shape[1] == 2 + assert tail_kv_cache.shape[2] == pool_size + assert tail_kv_cache.shape[3] == head_dim + assert tail_kv_cache.dtype == torch.bfloat16 + assert key.ndim == 3 and key.shape[2] == head_dim + assert slot_score.shape == key.shape + assert ape.shape == (pool_size, head_dim) + assert tail_slot_mapping.shape == (num_requests, next_n) + assert slot_mapping.shape == (num_requests, next_n) + assert positions.shape == (num_requests, next_n) + assert key.dtype == torch.bfloat16 + assert slot_score.dtype == torch.bfloat16 + assert ape.dtype == torch.float32 + assert kv_cache.dtype == torch.uint8 + + page_size = kv_cache.shape[1] + buf = kv_cache + buf_fp8 = buf.view(FP8_DTYPE) + buf_fp32 = buf.view(torch.float32) + + # The kernel indexes the int tensors as ``req * next_n + t`` (row-major), + # so they must be contiguous. Callers pass either a view of a contiguous + # slice or a freshly scattered tensor, making these no-ops; the calls guard + # against a future caller handing over a strided view. + tail_slot_mapping = tail_slot_mapping.contiguous() + slot_mapping = slot_mapping.contiguous() + positions = positions.contiguous() + + if page_size > 1: + assert page_size % 16 == 0, "ROCm preshuffle requires 16-token tiles" + + _kpool_decode_update_batched_kernel[(num_requests,)]( + buf_fp8, + buf_fp32, + tail_kv_cache, + tail_slot_mapping, + key, + key.stride(0), + key.stride(1), + slot_score, + slot_score.stride(0), + slot_score.stride(1), + ape, + ape.stride(0), + slot_mapping, + positions, + next_n, + PAGE_SIZE=page_size, + BUF_NUMEL_PER_PAGE=buf.stride(0), + POOL_SIZE=pool_size, + TAIL_BLOCK_ELEMS=tail_kv_cache.stride(0), + KPOOL_HEAD=tail_kv_cache.stride(1), + HEAD_DIM=head_dim, + S_OFFSET_NBYTES_IN_PAGE=page_size * head_dim, + FP8_MAX=FP8_MAX, + PRESHUFFLE=page_size > 1, + ROUND_SCALE=round_scale, + BLOCK_D=triton.next_power_of_2(head_dim), + ) + + +# Pool-level top-k helpers. + + +def history_group_budget_for_topk(topk: int, pool_size: int) -> int: + """Number of pools to select so that expanding yields ``topk`` tokens.""" + assert topk % pool_size == 0 + return topk // pool_size + + +def expand_pools_to_tokens( + group_ids: torch.Tensor, + group_valid: torch.Tensor, + topk: int, + pool_size: int, + page_table: torch.Tensor | None = None, + topk_offsets: torch.Tensor | None = None, +) -> torch.Tensor: + """Expand selected full-pool ids to a strict-width token topk tensor.""" + assert group_ids.ndim == 2 + assert group_valid.shape == group_ids.shape + assert topk % pool_size == 0 + assert group_ids.shape[1] == history_group_budget_for_topk(topk, pool_size) + assert page_table is None or topk_offsets is None + + device = group_ids.device + offsets = torch.arange(pool_size, device=device, dtype=torch.int64) + token_ids = group_ids.to(torch.int64).unsqueeze(-1) * pool_size + offsets + token_ids = token_ids.reshape(group_ids.shape[0], topk) + valid = ( + group_valid.unsqueeze(-1) + .expand(-1, -1, pool_size) + .reshape(group_ids.shape[0], topk) + ) + + if page_table is not None: + assert page_table.ndim == 2 + safe_ids = token_ids.clamp(min=0, max=page_table.shape[1] - 1) + output = torch.gather(page_table, dim=1, index=safe_ids).to(torch.int32) + elif topk_offsets is not None: + if topk_offsets.ndim == 2: + assert topk_offsets.shape[1] == 1 + topk_offsets = topk_offsets.squeeze(1) + output = (token_ids + topk_offsets.to(torch.int64).unsqueeze(1)).to(torch.int32) + else: + output = token_ids.to(torch.int32) + + return torch.where(valid, output, torch.full_like(output, -1)) + + +def append_tail_to_topk( + topk_result: torch.Tensor, + seq_lens: torch.Tensor, + pool_lens: torch.Tensor, + pool_size: int, + page_table: torch.Tensor | None = None, + topk_offsets: torch.Tensor | None = None, +) -> torch.Tensor: + """Append non-pooled tail tokens after expanded history tokens. + + ``index_kpool_always_select_tail`` keeps the (incomplete) trailing pool so + the most recent tokens are always attended to. + """ + assert topk_result.dtype == torch.int32 + assert seq_lens.ndim == 1 + assert pool_lens.ndim == 1 + + tail_pool = pool_size - 1 + if tail_pool == 0: + return topk_result + + rows, n_cols = topk_result.shape + out_cols = n_cols + tail_pool + out = torch.empty( + (rows, out_cols), dtype=topk_result.dtype, device=topk_result.device + ) + + # tail tokens: [pool_len*pool_size, seq_len) for each row. + pool_len = pool_lens.to(torch.int32) + tail_start = pool_len * pool_size + seq_len = seq_lens.to(torch.int32) + tail_count = seq_len - tail_start # in [0, pool_size) + + cols = torch.arange(out_cols, device=topk_result.device)[None, :] + history_len = n_cols + is_history = cols < history_len + tail_off = cols - history_len + is_tail = (tail_off >= 0) & (tail_off < tail_count[:, None]) + + # safe_hist must be per-row [rows, out_cols] so the gather reads each row's + # OWN history. cols is [1, out_cols]; if used directly, gather (which does + # NOT broadcast the index) would read only row 0 of topk_result, making every + # query inherit row 0's history (empty for the first token) and lose all its + # selected tokens — only the per-row tail would survive. This only manifests + # for multi-row sparse PREFILL (decode has 1 row, so it reads its own row 0). + safe_hist = torch.minimum(cols, torch.full_like(cols, n_cols - 1)).expand( + rows, out_cols + ) + history_val = torch.gather(topk_result, 1, safe_hist) + + tail_raw = tail_start[:, None] + tail_off + tail_val = tail_raw.to(torch.int32) + if page_table is not None: + safe_tail = tail_raw.clamp(min=0, max=page_table.shape[1] - 1) + tail_val = torch.gather(page_table, 1, safe_tail).to(torch.int32) + elif topk_offsets is not None: + tail_val = (tail_raw + topk_offsets.to(torch.int64).unsqueeze(1)).to( + torch.int32 + ) + + out = torch.where(is_history, history_val, -1) + out = torch.where(is_tail, tail_val, out) + return out + + +@triton.jit +def _expand_pools_and_append_tail_kernel( + pool_ids_ptr, # [rows, n_groups], int (any int dtype) + seq_lens_ptr, # [rows], int32 (token-granular seq_len) + out_ptr, # [rows, out_cols], int32 + topk, # n_groups * pool_size + out_cols, # topk + pool_size - 1 + POOL_SIZE: tl.constexpr, + BLOCK_COLS: tl.constexpr, + pid_s0, + out_s0, +): + # Fuses expand_pools_to_tokens + append_tail_to_topk (identity path) into a + # single kernel. Each program writes one (row, column-tile) of the output. + row = tl.program_id(0) + tile = tl.program_id(1) + cols = tile * BLOCK_COLS + tl.arange(0, BLOCK_COLS) + mask = cols < out_cols + + seq_len = tl.load(seq_lens_ptr + row) + pool_len = seq_len // POOL_SIZE + tail_start = pool_len * POOL_SIZE + tail_count = seq_len - tail_start # in [0, POOL_SIZE) + + # History region [0, topk): expand selected pool g = cols // POOL_SIZE. + is_history = cols < topk + g = cols // POOL_SIZE + o = cols % POOL_SIZE + pid = tl.load(pool_ids_ptr + row * pid_s0 + g, mask=mask & is_history, other=-1) + hist_val = (pid * POOL_SIZE + o).to(tl.int32) + hist_out = tl.where(pid >= 0, hist_val, -1) + + # Tail region [topk, out_cols): the request's trailing incomplete pool. + tail_off = cols - topk + is_tail = (tail_off >= 0) & (tail_off < tail_count) + tail_val = (tail_start + tail_off).to(tl.int32) + tail_out = tl.where(is_tail, tail_val, -1) + + result = tl.where(is_history, hist_out, tail_out) + tl.store(out_ptr + row * out_s0 + cols, result, mask=mask) + + +def expand_pools_and_append_tail( + pool_ids: torch.Tensor, + seq_lens: torch.Tensor, + pool_size: int, +) -> torch.Tensor: + """Fuse ``expand_pools_to_tokens`` + ``append_tail_to_topk`` (identity path). + + Produces the same ``[rows, topk + pool_size - 1]`` int32 output as calling + the two functions in sequence when neither ``page_table`` nor + ``topk_offsets`` is passed — the only path used by the GLM-5.3-Flash indexer. + The kernel derives ``pool_len = seq_len // pool_size`` internally, so the + caller no longer needs to precompute it. Replaces ~25 elementwise kernels + with one Triton launch. + """ + rows, n_groups = pool_ids.shape + topk = n_groups * pool_size + out_cols = topk + pool_size - 1 + out = torch.empty((rows, out_cols), dtype=torch.int32, device=pool_ids.device) + BLOCK_COLS = 128 + n_tiles = triton.cdiv(out_cols, BLOCK_COLS) + _expand_pools_and_append_tail_kernel[(rows, n_tiles)]( + pool_ids, + seq_lens, + out, + topk, + out_cols, + POOL_SIZE=pool_size, + BLOCK_COLS=BLOCK_COLS, + pid_s0=pool_ids.stride(0), + out_s0=out.stride(0), + ) + return out diff --git a/vllm/models/glm5next/amd/ops/third_party/__init__.py b/vllm/models/glm5next/amd/ops/third_party/__init__.py new file mode 100644 index 000000000000..208f01a7cb5e --- /dev/null +++ b/vllm/models/glm5next/amd/ops/third_party/__init__.py @@ -0,0 +1,2 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project diff --git a/vllm/models/glm5next/amd/ops/third_party/kda/__init__.py b/vllm/models/glm5next/amd/ops/third_party/kda/__init__.py new file mode 100644 index 000000000000..a073bd8021d4 --- /dev/null +++ b/vllm/models/glm5next/amd/ops/third_party/kda/__init__.py @@ -0,0 +1,6 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project + +from .kernels import chunk_kda_with_fused_gate, fused_recurrent_kda + +__all__ = ["chunk_kda_with_fused_gate", "fused_recurrent_kda"] diff --git a/vllm/models/glm5next/amd/ops/third_party/kda/fused_recurrent.py b/vllm/models/glm5next/amd/ops/third_party/kda/fused_recurrent.py new file mode 100644 index 000000000000..d7de85984f85 --- /dev/null +++ b/vllm/models/glm5next/amd/ops/third_party/kda/fused_recurrent.py @@ -0,0 +1,656 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +# SPDX-FileCopyrightText: Songlin Yang, Yu Zhang +# mypy: ignore-errors +# +# This file contains code copied from the flash-linear-attention project. +# The original source code was licensed under the MIT license and included +# the following copyright notice: +# Copyright (c) 2023-2025, Songlin Yang, Yu Zhang +# ruff: noqa: E501 + +import torch + +from vllm.third_party.flash_linear_attention.ops.op import exp +from vllm.triton_utils import tl, triton + + +@triton.heuristics( + { + "USE_INITIAL_STATE": lambda args: args["h0"] is not None, + "IS_VARLEN": lambda args: args["cu_seqlens"] is not None, + "IS_CONTINUOUS_BATCHING": lambda args: args["ssm_state_indices"] is not None, + "IS_SPEC_DECODING": lambda args: args["num_accepted_tokens"] is not None, + } +) +@triton.jit(do_not_specialize=["N", "T"]) +def fused_recurrent_gated_delta_rule_fwd_kernel( + q, + k, + v, + g, + beta, + o, + h0, + ht, + cu_seqlens, + ssm_state_indices, + num_accepted_tokens, + a_log, + g_bias, + scale, + N: tl.int64, # num of sequences + T: tl.int64, # num of tokens + B: tl.constexpr, + H: tl.constexpr, + HV: tl.constexpr, + K: tl.constexpr, + V: tl.constexpr, + BK: tl.constexpr, + BV: tl.constexpr, + stride_init_state_token: tl.constexpr, + stride_final_state_token: tl.constexpr, + stride_indices_seq: tl.constexpr, + stride_indices_tok: tl.constexpr, + USE_INITIAL_STATE: tl.constexpr, # whether to use initial state + INPLACE_FINAL_STATE: tl.constexpr, # whether to store final state inplace + IS_BETA_HEADWISE: tl.constexpr, # whether beta is headwise vector or scalar, + USE_QK_L2NORM_IN_KERNEL: tl.constexpr, + IS_VARLEN: tl.constexpr, + IS_CONTINUOUS_BATCHING: tl.constexpr, + IS_SPEC_DECODING: tl.constexpr, + IS_KDA: tl.constexpr, + SIGMOID_BETA: tl.constexpr, # beta holds raw logits; sigmoid at fp32 load + COMPUTE_GATE: tl.constexpr, # g holds raw logits; KDA gate computed in-kernel + SAFE_GATE: tl.constexpr, # bounded gate variant (only branch implemented) + LOWER_BOUND: tl.constexpr, +): + i_k, i_v, i_nh = tl.program_id(0), tl.program_id(1), tl.program_id(2) + i_n, i_hv = i_nh // HV, i_nh % HV + i_h = i_hv // (HV // H) + if IS_VARLEN: + bos, eos = ( + tl.load(cu_seqlens + i_n).to(tl.int64), + tl.load(cu_seqlens + i_n + 1).to(tl.int64), + ) + all = T + T = eos - bos + else: + bos, eos = i_n * T, i_n * T + T + all = B * T + + if T == 0: + # no tokens to process for this sequence + return + + o_k = i_k * BK + tl.arange(0, BK) + o_v = i_v * BV + tl.arange(0, BV) + + p_q = q + (bos * H + i_h) * K + o_k + p_k = k + (bos * H + i_h) * K + o_k + p_v = v + (bos * HV + i_hv) * V + o_v + if IS_BETA_HEADWISE: + p_beta = beta + (bos * HV + i_hv) * V + o_v + else: + p_beta = beta + bos * HV + i_hv + + if not IS_KDA: + p_g = g + bos * HV + i_hv + else: + p_gk = g + (bos * HV + i_hv) * K + o_k + + # Per-head gate amplitude, hoisted out of the token loop (COMPUTE_GATE). + if COMPUTE_GATE: + b_a_log = tl.exp(tl.load(a_log + i_h).to(tl.float32)) + + p_o = o + ((i_k * all + bos) * HV + i_hv) * V + o_v + + mask_k = o_k < K + mask_v = o_v < V + mask_h = mask_v[:, None] & mask_k[None, :] + + b_h = tl.zeros([BV, BK], dtype=tl.float32) + if USE_INITIAL_STATE: + if IS_CONTINUOUS_BATCHING: + if IS_SPEC_DECODING: + i_t = tl.load(num_accepted_tokens + i_n).to(tl.int64) - 1 + else: + i_t = 0 + # Load state index and check for invalid entries + state_idx = tl.load(ssm_state_indices + i_n * stride_indices_seq + i_t).to( + tl.int64 + ) + # Skip if state index is invalid (NULL_BLOCK_ID=0) + if state_idx <= 0: + return + p_h0 = h0 + state_idx * stride_init_state_token + else: + p_h0 = h0 + bos * HV * V * K + p_h0 = p_h0 + i_hv * V * K + o_v[:, None] * K + o_k[None, :] + b_h += tl.load(p_h0, mask=mask_h, other=0).to(tl.float32) + + for i_t in range(0, T): + b_q = tl.load(p_q, mask=mask_k, other=0).to(tl.float32) + b_k = tl.load(p_k, mask=mask_k, other=0).to(tl.float32) + b_v = tl.load(p_v, mask=mask_v, other=0).to(tl.float32) + + if USE_QK_L2NORM_IN_KERNEL: + b_q = b_q / tl.sqrt(tl.sum(b_q * b_q) + 1e-6) + b_k = b_k / tl.sqrt(tl.sum(b_k * b_k) + 1e-6) + b_q = b_q * scale + # [BV, BK] + if not IS_KDA: + b_g = tl.load(p_g).to(tl.float32) + b_h *= exp(b_g) + else: + b_gk = tl.load(p_gk).to(tl.float32) + if COMPUTE_GATE: + # Replicates kda_gate_fwd_kernel's SAFE_GATE branch + # bit-for-bit (same tl.exp, same fp32 math; the intermediate + # gate value this replaces was stored/reloaded as fp32, + # which is lossless): y = lb / (1 + exp(-exp(A)*(g+bias))). + b_gk += tl.load(g_bias + i_h * K + o_k, mask=mask_k, other=0.0).to( + tl.float32 + ) + b_gk = LOWER_BOUND / (1.0 + tl.exp(-(b_a_log * b_gk))) + b_h *= exp(b_gk[None, :]) + # [BV] + b_v -= tl.sum(b_h * b_k[None, :], 1) + if IS_BETA_HEADWISE: + b_beta = tl.load(p_beta, mask=mask_v, other=0).to(tl.float32) + else: + b_beta = tl.load(p_beta).to(tl.float32) + # Matches torch's `x.float().sigmoid()` pre-computation bit-for-bit + # on the input side (bf16->fp32 is exact); only the sigmoid impl itself + # can differ by <=1 ULP. + if SIGMOID_BETA: + b_beta = tl.sigmoid(b_beta) + b_v *= b_beta + # [BV, BK] + b_h += b_v[:, None] * b_k[None, :] + # [BV] + b_o = tl.sum(b_h * b_q[None, :], 1) + tl.store(p_o, b_o.to(p_o.dtype.element_ty), mask=mask_v) + + # keep the states for multi-query tokens + if INPLACE_FINAL_STATE: + # Load state index and check for invalid entries + final_state_idx = tl.load( + ssm_state_indices + i_n * stride_indices_seq + i_t + ).to(tl.int64) + # Only store if state index is valid (not NULL_BLOCK_ID=0) + if final_state_idx > 0: + p_ht = ht + final_state_idx * stride_final_state_token + p_ht = p_ht + i_hv * V * K + o_v[:, None] * K + o_k[None, :] + tl.store(p_ht, b_h.to(p_ht.dtype.element_ty), mask=mask_h) + else: + p_ht = ht + (bos + i_t) * stride_final_state_token + p_ht = p_ht + i_hv * V * K + o_v[:, None] * K + o_k[None, :] + tl.store(p_ht, b_h.to(p_ht.dtype.element_ty), mask=mask_h) + + p_q += H * K + p_k += H * K + p_o += HV * V + p_v += HV * V + if not IS_KDA: + p_g += HV + else: + p_gk += HV * K + p_beta += HV * (V if IS_BETA_HEADWISE else 1) + + +def fused_recurrent_gated_delta_rule_fwd( + q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, + g: torch.Tensor, + beta: torch.Tensor, + scale: float, + initial_state: torch.Tensor, + inplace_final_state: bool = True, + cu_seqlens: torch.Tensor | None = None, + ssm_state_indices: torch.Tensor | None = None, + num_accepted_tokens: torch.Tensor | None = None, + use_qk_l2norm_in_kernel: bool = False, +) -> tuple[torch.Tensor, torch.Tensor]: + B, T, H, K, V = *k.shape, v.shape[-1] + HV = v.shape[2] + N = B if cu_seqlens is None else len(cu_seqlens) - 1 + BK, BV = triton.next_power_of_2(K), min(triton.next_power_of_2(V), 32) + NK, NV = triton.cdiv(K, BK), triton.cdiv(V, BV) + assert NK == 1, "NK > 1 is not supported yet" + num_stages = 3 + num_warps = 1 + + o = q.new_empty(NK, *v.shape) + if inplace_final_state: + final_state = initial_state + else: + final_state = q.new_empty(T, HV, V, K, dtype=initial_state.dtype) + + stride_init_state_token = initial_state.stride(0) + stride_final_state_token = final_state.stride(0) + + if ssm_state_indices is None: + stride_indices_seq, stride_indices_tok = 1, 1 + elif ssm_state_indices.ndim == 1: + stride_indices_seq, stride_indices_tok = ssm_state_indices.stride(0), 1 + else: + stride_indices_seq, stride_indices_tok = ssm_state_indices.stride() + + grid = (NK, NV, N * HV) + fused_recurrent_gated_delta_rule_fwd_kernel[grid]( + q=q, + k=k, + v=v, + g=g, + beta=beta, + o=o, + h0=initial_state, + ht=final_state, + cu_seqlens=cu_seqlens, + ssm_state_indices=ssm_state_indices, + num_accepted_tokens=num_accepted_tokens, + scale=scale, + N=N, + T=T, + B=B, + H=H, + HV=HV, + K=K, + V=V, + BK=BK, + BV=BV, + stride_init_state_token=stride_init_state_token, + stride_final_state_token=stride_final_state_token, + stride_indices_seq=stride_indices_seq, + stride_indices_tok=stride_indices_tok, + IS_BETA_HEADWISE=beta.ndim == v.ndim, + USE_QK_L2NORM_IN_KERNEL=use_qk_l2norm_in_kernel, + INPLACE_FINAL_STATE=inplace_final_state, + IS_KDA=False, + SIGMOID_BETA=False, + a_log=None, + g_bias=None, + COMPUTE_GATE=False, + SAFE_GATE=True, + LOWER_BOUND=-5.0, + num_warps=num_warps, + num_stages=num_stages, + ) + o = o.squeeze(0) + return o, final_state + + +@triton.jit +def fused_recurrent_gated_delta_rule_packed_decode_kernel( + mixed_qkv, + a, + b, + A_log, + dt_bias, + o, + h0, + ht, + ssm_state_indices, + scale, + stride_mixed_qkv_tok: tl.constexpr, + stride_a_tok: tl.constexpr, + stride_b_tok: tl.constexpr, + stride_init_state_token: tl.constexpr, + stride_final_state_token: tl.constexpr, + stride_indices_seq: tl.constexpr, + H: tl.constexpr, + HV: tl.constexpr, + K: tl.constexpr, + V: tl.constexpr, + BK: tl.constexpr, + BV: tl.constexpr, + SOFTPLUS_THRESHOLD: tl.constexpr, + USE_QK_L2NORM_IN_KERNEL: tl.constexpr, + SPLIT_BATCH_HEAD_GRID: tl.constexpr, +): + if SPLIT_BATCH_HEAD_GRID: + i_v, i_hv, i_n = tl.program_id(0), tl.program_id(1), tl.program_id(2) + else: + i_v, i_nh = tl.program_id(0), tl.program_id(1) + i_n, i_hv = i_nh // HV, i_nh % HV + i_h = i_hv // (HV // H) + + o_k = tl.arange(0, BK) + o_v = i_v * BV + tl.arange(0, BV) + mask_k = o_k < K + mask_v = o_v < V + mask_h = mask_v[:, None] & mask_k[None, :] + + state_idx = tl.load(ssm_state_indices + i_n * stride_indices_seq).to(tl.int64) + p_o = o + (i_n * HV + i_hv) * V + o_v + + # Skip if state index is invalid (NULL_BLOCK_ID=0) + if state_idx <= 0: + zero = tl.zeros([BV], dtype=tl.float32).to(p_o.dtype.element_ty) + tl.store(p_o, zero, mask=mask_v) + return + + p_h0 = h0 + state_idx * stride_init_state_token + p_h0 = p_h0 + i_hv * V * K + o_v[:, None] * K + o_k[None, :] + b_h = tl.load(p_h0, mask=mask_h, other=0).to(tl.float32) + + p_mixed = mixed_qkv + i_n * stride_mixed_qkv_tok + q_off = i_h * K + o_k + k_off = (H * K) + i_h * K + o_k + v_off = (2 * H * K) + i_hv * V + o_v + b_q = tl.load(p_mixed + q_off, mask=mask_k, other=0).to(tl.float32) + b_k = tl.load(p_mixed + k_off, mask=mask_k, other=0).to(tl.float32) + b_v = tl.load(p_mixed + v_off, mask=mask_v, other=0).to(tl.float32) + + if USE_QK_L2NORM_IN_KERNEL: + b_q = b_q / tl.sqrt(tl.sum(b_q * b_q) + 1e-6) + b_k = b_k / tl.sqrt(tl.sum(b_k * b_k) + 1e-6) + b_q = b_q * scale + + a_val = tl.load(a + i_n * stride_a_tok + i_hv).to(tl.float32) + b_val = tl.load(b + i_n * stride_b_tok + i_hv).to(tl.float32) + A_log_val = tl.load(A_log + i_hv).to(tl.float32) + dt_bias_val = tl.load(dt_bias + i_hv).to(tl.float32) + x = a_val + dt_bias_val + softplus_x = tl.where(x <= SOFTPLUS_THRESHOLD, tl.log(1.0 + tl.exp(x)), x) + g_val = -tl.exp(A_log_val) * softplus_x + beta_val = tl.sigmoid(b_val) + + b_h *= exp(g_val) + b_v -= tl.sum(b_h * b_k[None, :], 1) + b_v *= beta_val + b_h += b_v[:, None] * b_k[None, :] + b_o = tl.sum(b_h * b_q[None, :], 1) + tl.store(p_o, b_o.to(p_o.dtype.element_ty), mask=mask_v) + + p_ht = ht + state_idx * stride_final_state_token + p_ht = p_ht + i_hv * V * K + o_v[:, None] * K + o_k[None, :] + tl.store(p_ht, b_h.to(p_ht.dtype.element_ty), mask=mask_h) + + +def fused_recurrent_gated_delta_rule_packed_decode( + mixed_qkv: torch.Tensor, + a: torch.Tensor, + b: torch.Tensor, + A_log: torch.Tensor, + dt_bias: torch.Tensor, + scale: float, + initial_state: torch.Tensor, + out: torch.Tensor, + ssm_state_indices: torch.Tensor, + use_qk_l2norm_in_kernel: bool = False, +) -> tuple[torch.Tensor, torch.Tensor]: + if mixed_qkv.ndim != 2: + raise ValueError( + f"`mixed_qkv` must be a 2D tensor (got ndim={mixed_qkv.ndim})." + ) + if mixed_qkv.stride(-1) != 1: + raise ValueError("`mixed_qkv` must be contiguous in the last dim.") + if a.ndim != 2 or b.ndim != 2: + raise ValueError( + f"`a` and `b` must be 2D tensors (got a.ndim={a.ndim}, b.ndim={b.ndim})." + ) + if a.stride(-1) != 1 or b.stride(-1) != 1: + raise ValueError("`a`/`b` must be contiguous in the last dim.") + if A_log.ndim != 1 or dt_bias.ndim != 1: + raise ValueError("`A_log`/`dt_bias` must be 1D tensors.") + if A_log.stride(0) != 1 or dt_bias.stride(0) != 1: + raise ValueError("`A_log`/`dt_bias` must be contiguous.") + if ssm_state_indices.ndim != 1: + raise ValueError( + f"`ssm_state_indices` must be 1D for packed decode (got ndim={ssm_state_indices.ndim})." + ) + if not out.is_contiguous(): + raise ValueError("`out` must be contiguous.") + + dev = mixed_qkv.device + if ( + a.device != dev + or b.device != dev + or A_log.device != dev + or dt_bias.device != dev + or initial_state.device != dev + or out.device != dev + or ssm_state_indices.device != dev + ): + raise ValueError("All inputs must be on the same device.") + + B = mixed_qkv.shape[0] + if a.shape[0] != B or b.shape[0] != B: + raise ValueError( + "Mismatched batch sizes: " + f"mixed_qkv.shape[0]={B}, a.shape[0]={a.shape[0]}, b.shape[0]={b.shape[0]}." + ) + if ssm_state_indices.shape[0] != B: + raise ValueError( + f"`ssm_state_indices` must have shape [B] (got {tuple(ssm_state_indices.shape)}; expected ({B},))." + ) + + if initial_state.ndim != 4: + raise ValueError( + f"`initial_state` must be a 4D tensor (got ndim={initial_state.ndim})." + ) + if initial_state.stride(-1) != 1: + raise ValueError("`initial_state` must be contiguous in the last dim.") + HV, V, K = initial_state.shape[-3:] + if a.shape[1] != HV or b.shape[1] != HV: + raise ValueError( + f"`a`/`b` must have shape [B, HV] with HV={HV} (got a.shape={tuple(a.shape)}, b.shape={tuple(b.shape)})." + ) + if A_log.numel() != HV or dt_bias.numel() != HV: + raise ValueError( + f"`A_log` and `dt_bias` must have {HV} elements (got A_log.numel()={A_log.numel()}, dt_bias.numel()={dt_bias.numel()})." + ) + if out.shape != (B, 1, HV, V): + raise ValueError( + f"`out` must have shape {(B, 1, HV, V)} (got out.shape={tuple(out.shape)})." + ) + + qkv_dim = mixed_qkv.shape[1] + qk_dim = qkv_dim - HV * V + if qk_dim <= 0 or qk_dim % 2 != 0: + raise ValueError( + f"Invalid packed `mixed_qkv` last dim={qkv_dim} for HV={HV}, V={V}." + ) + q_dim = qk_dim // 2 + if q_dim % K != 0: + raise ValueError(f"Invalid packed Q size {q_dim}: must be divisible by K={K}.") + H = q_dim // K + if H <= 0 or HV % H != 0: + raise ValueError( + f"Invalid head config inferred from mixed_qkv: H={H}, HV={HV}." + ) + + BK = triton.next_power_of_2(K) + if triton.cdiv(K, BK) != 1: + raise ValueError( + f"Packed decode kernel only supports NK=1 (got K={K}, BK={BK})." + ) + BV = min(triton.next_power_of_2(V), 32) + num_stages = 3 + num_warps = 1 + + stride_mixed_qkv_tok = mixed_qkv.stride(0) + stride_a_tok = a.stride(0) + stride_b_tok = b.stride(0) + stride_init_state_token = initial_state.stride(0) + stride_final_state_token = initial_state.stride(0) + stride_indices_seq = ssm_state_indices.stride(0) + + NV = triton.cdiv(V, BV) + # CUDA limits grid Y/Z dimensions to 65535. + split_batch_head_grid = B * HV > 65535 + grid = (NV, HV, B) if split_batch_head_grid else (NV, B * HV) + fused_recurrent_gated_delta_rule_packed_decode_kernel[grid]( + mixed_qkv=mixed_qkv, + a=a, + b=b, + A_log=A_log, + dt_bias=dt_bias, + o=out, + h0=initial_state, + ht=initial_state, + ssm_state_indices=ssm_state_indices, + scale=scale, + stride_mixed_qkv_tok=stride_mixed_qkv_tok, + stride_a_tok=stride_a_tok, + stride_b_tok=stride_b_tok, + stride_init_state_token=stride_init_state_token, + stride_final_state_token=stride_final_state_token, + stride_indices_seq=stride_indices_seq, + H=H, + HV=HV, + K=K, + V=V, + BK=BK, + BV=BV, + SOFTPLUS_THRESHOLD=20.0, + USE_QK_L2NORM_IN_KERNEL=use_qk_l2norm_in_kernel, + SPLIT_BATCH_HEAD_GRID=split_batch_head_grid, + num_warps=num_warps, + num_stages=num_stages, + ) + return out, initial_state + + +class FusedRecurrentFunction(torch.autograd.Function): + @staticmethod + def forward( + ctx, + q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, + g: torch.Tensor, + beta: torch.Tensor, + scale: float, + initial_state: torch.Tensor, + inplace_final_state: bool = True, + cu_seqlens: torch.Tensor | None = None, + ssm_state_indices: torch.Tensor | None = None, + num_accepted_tokens: torch.Tensor | None = None, + use_qk_l2norm_in_kernel: bool = False, + ): + o, final_state = fused_recurrent_gated_delta_rule_fwd( + q=q.contiguous(), + k=k.contiguous(), + v=v.contiguous(), + g=g.contiguous(), + beta=beta.contiguous(), + scale=scale, + initial_state=initial_state, + inplace_final_state=inplace_final_state, + cu_seqlens=cu_seqlens, + ssm_state_indices=ssm_state_indices, + num_accepted_tokens=num_accepted_tokens, + use_qk_l2norm_in_kernel=use_qk_l2norm_in_kernel, + ) + + return o, final_state + + +def fused_recurrent_gated_delta_rule( + q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, + g: torch.Tensor, + beta: torch.Tensor = None, + scale: float = None, + initial_state: torch.Tensor = None, + inplace_final_state: bool = True, + cu_seqlens: torch.Tensor | None = None, + ssm_state_indices: torch.Tensor | None = None, + num_accepted_tokens: torch.Tensor | None = None, + use_qk_l2norm_in_kernel: bool = False, +) -> tuple[torch.Tensor, torch.Tensor]: + r""" + Args: + q (torch.Tensor): + queries of shape `[B, T, H, K]`. + k (torch.Tensor): + keys of shape `[B, T, H, K]`. + v (torch.Tensor): + values of shape `[B, T, HV, V]`. + GVA is applied if `HV > H`. + g (torch.Tensor): + g (decays) of shape `[B, T, HV]`. + beta (torch.Tensor): + betas of shape `[B, T, HV]`. + scale (Optional[int]): + Scale factor for the RetNet attention scores. + If not provided, it will default to `1 / sqrt(K)`. Default: `None`. + initial_state (Optional[torch.Tensor]): + Initial state of shape `[N, HV, V, K]` for `N` input sequences. + For equal-length input sequences, `N` equals the batch size `B`. + Default: `None`. + inplace_final_state: bool: + Whether to store the final state in-place to save memory. + Default: `True`. + cu_seqlens (torch.Tensor): + Cumulative sequence lengths of shape `[N+1]` used for variable-length training, + consistent with the FlashAttention API. + ssm_state_indices (Optional[torch.Tensor]): + Indices to map the input sequences to the initial/final states. + num_accepted_tokens (Optional[torch.Tensor]): + Number of accepted tokens for each sequence during decoding. + + Returns: + o (torch.Tensor): + Outputs of shape `[B, T, HV, V]`. + final_state (torch.Tensor): + Final state of shape `[N, HV, V, K]`. + + Examples:: + >>> import torch + >>> import torch.nn.functional as F + >>> from einops import rearrange + >>> from fla.ops.gated_delta_rule import fused_recurrent_gated_delta_rule + # inputs with equal lengths + >>> B, T, H, HV, K, V = 4, 2048, 4, 8, 512, 512 + >>> q = torch.randn(B, T, H, K, device='cuda') + >>> k = F.normalize(torch.randn(B, T, H, K, device='cuda'), p=2, dim=-1) + >>> v = torch.randn(B, T, HV, V, device='cuda') + >>> g = F.logsigmoid(torch.rand(B, T, HV, device='cuda')) + >>> beta = torch.rand(B, T, HV, device='cuda').sigmoid() + >>> h0 = torch.randn(B, HV, V, K, device='cuda') + >>> o, ht = fused_gated_recurrent_delta_rule( + q, k, v, g, beta, + initial_state=h0, + ) + # for variable-length inputs, the batch size `B` is expected to be 1 and `cu_seqlens` is required + >>> q, k, v, g, beta = map(lambda x: rearrange(x, 'b t ... -> 1 (b t) ...'), (q, k, v, g, beta)) + # for a batch with 4 sequences, `cu_seqlens` with 5 start/end positions are expected + >>> cu_seqlens = q.new_tensor([0, 2048, 4096, 6144, 8192], dtype=torch.int32) + >>> o_var, ht_var = fused_gated_recurrent_delta_rule( + q, k, v, g, beta, + initial_state=h0, + cu_seqlens=cu_seqlens + ) + """ + if cu_seqlens is not None and q.shape[0] != 1: + raise ValueError( + f"The batch size is expected to be 1 rather than {q.shape[0]} when using `cu_seqlens`." + f"Please flatten variable-length inputs before processing." + ) + if scale is None: + scale = k.shape[-1] ** -0.5 + else: + assert scale > 0, "scale must be positive" + if beta is None: + beta = torch.ones_like(q[..., 0]) + o, final_state = FusedRecurrentFunction.apply( + q, + k, + v, + g, + beta, + scale, + initial_state, + inplace_final_state, + cu_seqlens, + ssm_state_indices, + num_accepted_tokens, + use_qk_l2norm_in_kernel, + ) + return o, final_state diff --git a/vllm/models/glm5next/amd/ops/third_party/kda/kernels.py b/vllm/models/glm5next/amd/ops/third_party/kda/kernels.py new file mode 100644 index 000000000000..be417d6b0526 --- /dev/null +++ b/vllm/models/glm5next/amd/ops/third_party/kda/kernels.py @@ -0,0 +1,1362 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +# SPDX-FileCopyrightText: Songlin Yang, Yu Zhang +# mypy: ignore-errors +# +# This file contains code copied from the flash-linear-attention project. +# The original source code was licensed under the MIT license and included +# the following copyright notice: +# Copyright (c) 2023-2025, Songlin Yang, Yu Zhang +# ruff: noqa: E501 + + +import torch + +from vllm.third_party.flash_linear_attention.ops.chunk_delta_h import ( + chunk_gated_delta_rule_fwd_h, +) +from vllm.third_party.flash_linear_attention.ops.cumsum import chunk_local_cumsum +from vllm.third_party.flash_linear_attention.ops.index import prepare_chunk_indices +from vllm.third_party.flash_linear_attention.ops.l2norm import l2norm_fwd +from vllm.third_party.flash_linear_attention.ops.op import exp2, log +from vllm.third_party.flash_linear_attention.ops.solve_tril import solve_tril +from vllm.third_party.flash_linear_attention.ops.utils import FLA_CHUNK_SIZE, is_amd +from vllm.triton_utils import tl, triton +from vllm.utils.math_utils import RCP_LN2, cdiv, next_power_of_2 + +from .fused_recurrent import fused_recurrent_gated_delta_rule_fwd_kernel + +BT_LIST_AUTOTUNE = [32, 64, 128] +NUM_WARPS_AUTOTUNE = [2, 4, 8, 16] if is_amd else [4, 8, 16, 32] + + +def fused_recurrent_kda_fwd( + q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, + g: torch.Tensor, + beta: torch.Tensor, + scale: float, + initial_state: torch.Tensor, + inplace_final_state: bool = True, + cu_seqlens: torch.Tensor | None = None, + ssm_state_indices: torch.Tensor | None = None, + num_accepted_tokens: torch.Tensor | None = None, + use_qk_l2norm_in_kernel: bool = False, + out: torch.Tensor | None = None, + sigmoid_beta: bool = False, + a_log: torch.Tensor | None = None, + g_bias: torch.Tensor | None = None, + compute_gate: bool = False, + lower_bound: float | None = -5.0, +) -> tuple[torch.Tensor, torch.Tensor]: + B, T, H, K, V = *k.shape, v.shape[-1] + HV = v.shape[2] + N = B if cu_seqlens is None else len(cu_seqlens) - 1 + BK, BV = next_power_of_2(K), min(next_power_of_2(V), 8) + NK, NV = cdiv(K, BK), cdiv(V, BV) + assert NK == 1, "NK > 1 is not supported yet" + num_stages = 3 + num_warps = 1 + + if compute_gate: + assert a_log is not None and g_bias is not None, ( + "compute_gate requires a_log and g_bias" + ) + assert lower_bound is not None, ( + "compute_gate implements the bounded (safe_gate) branch only" + ) + a_log = a_log.reshape(-1).contiguous() + g_bias = g_bias.reshape(-1).contiguous() + + if out is None: + o = torch.empty_like(k) + else: + # Caller-provided output buffer; must be layout-compatible with the + # tensor the kernel indexes (contiguous, same shape/dtype as k). + assert out.shape == k.shape and out.dtype == k.dtype + assert out.is_contiguous() + o = out + if inplace_final_state: + final_state = initial_state + else: + final_state = q.new_empty(T, HV, V, K, dtype=initial_state.dtype) + + stride_init_state_token = initial_state.stride(0) + stride_final_state_token = final_state.stride(0) + + if ssm_state_indices is None: + stride_indices_seq, stride_indices_tok = 1, 1 + elif ssm_state_indices.ndim == 1: + stride_indices_seq, stride_indices_tok = ssm_state_indices.stride(0), 1 + else: + stride_indices_seq, stride_indices_tok = ssm_state_indices.stride() + + grid = (NK, NV, N * HV) + fused_recurrent_gated_delta_rule_fwd_kernel[grid]( + q=q, + k=k, + v=v, + g=g, + beta=beta, + o=o, + h0=initial_state, + ht=final_state, + cu_seqlens=cu_seqlens, + ssm_state_indices=ssm_state_indices, + num_accepted_tokens=num_accepted_tokens, + scale=scale, + N=N, + T=T, + B=B, + H=H, + HV=HV, + K=K, + V=V, + BK=BK, + BV=BV, + stride_init_state_token=stride_init_state_token, + stride_final_state_token=stride_final_state_token, + stride_indices_seq=stride_indices_seq, + stride_indices_tok=stride_indices_tok, + IS_BETA_HEADWISE=beta.ndim == v.ndim, + USE_QK_L2NORM_IN_KERNEL=use_qk_l2norm_in_kernel, + INPLACE_FINAL_STATE=inplace_final_state, + IS_KDA=True, + SIGMOID_BETA=sigmoid_beta, + a_log=a_log, + g_bias=g_bias, + COMPUTE_GATE=compute_gate, + SAFE_GATE=True, + LOWER_BOUND=lower_bound if lower_bound is not None else -5.0, + num_warps=num_warps, + num_stages=num_stages, + ) + + return o, final_state + + +def fused_recurrent_kda( + q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, + g: torch.Tensor, + beta: torch.Tensor = None, + scale: float = None, + initial_state: torch.Tensor = None, + inplace_final_state: bool = True, + use_qk_l2norm_in_kernel: bool = True, + cu_seqlens: torch.Tensor | None = None, + ssm_state_indices: torch.LongTensor | None = None, + num_accepted_tokens: torch.Tensor | None = None, + out: torch.Tensor | None = None, + sigmoid_beta: bool = False, + a_log: torch.Tensor | None = None, + g_bias: torch.Tensor | None = None, + compute_gate: bool = False, + lower_bound: float | None = -5.0, + **kwargs, +) -> tuple[torch.Tensor, torch.Tensor]: + if cu_seqlens is not None and q.shape[0] != 1: + raise ValueError( + f"The batch size is expected to be 1 rather than {q.shape[0]} when using `cu_seqlens`." + f"Please flatten variable-length inputs before processing." + ) + if scale is None: + scale = k.shape[-1] ** -0.5 + + o, final_state = fused_recurrent_kda_fwd( + q=q.contiguous(), + k=k.contiguous(), + v=v.contiguous(), + g=g.contiguous(), + beta=beta.contiguous(), + scale=scale, + initial_state=initial_state, + inplace_final_state=inplace_final_state, + cu_seqlens=cu_seqlens, + ssm_state_indices=ssm_state_indices, + num_accepted_tokens=num_accepted_tokens, + use_qk_l2norm_in_kernel=use_qk_l2norm_in_kernel, + out=out, + sigmoid_beta=sigmoid_beta, + a_log=a_log, + g_bias=g_bias, + compute_gate=compute_gate, + lower_bound=lower_bound, + ) + return o, final_state + + +@triton.heuristics({"IS_VARLEN": lambda args: args["cu_seqlens"] is not None}) +@triton.autotune( + configs=[ + triton.Config({"BK": BK}, num_warps=num_warps, num_stages=num_stages) + for BK in [32, 64] + for num_warps in [1, 2, 4, 8] + for num_stages in [2, 3, 4] + ], + key=["BC"], +) +@triton.jit(do_not_specialize=["T"]) +def chunk_kda_scaled_dot_kkt_fwd_kernel_intra_sub_inter( + q, + k, + g, + beta, + A, + Aqk, + scale, + cu_seqlens, + chunk_indices, + T, + H: tl.constexpr, + K: tl.constexpr, + BT: tl.constexpr, + BC: tl.constexpr, + BK: tl.constexpr, + NC: tl.constexpr, + IS_VARLEN: tl.constexpr, +): + i_t, i_c, i_bh = tl.program_id(0), tl.program_id(1), tl.program_id(2) + i_b, i_h = i_bh // H, i_bh % H + i_i, i_j = i_c // NC, i_c % NC + if IS_VARLEN: + i_n, i_t = ( + tl.load(chunk_indices + i_t * 2).to(tl.int32), + tl.load(chunk_indices + i_t * 2 + 1).to(tl.int32), + ) + bos, eos = ( + tl.load(cu_seqlens + i_n).to(tl.int32), + tl.load(cu_seqlens + i_n + 1).to(tl.int32), + ) + T = eos - bos + else: + bos, eos = i_b * T, i_b * T + T + + if i_t * BT + i_i * BC >= T: + return + if i_i <= i_j: + return + + q += (bos * H + i_h) * K + k += (bos * H + i_h) * K + g += (bos * H + i_h) * K + A += (bos * H + i_h) * BT + Aqk += (bos * H + i_h) * BT + + p_b = tl.make_block_ptr( + beta + bos * H + i_h, (T,), (H,), (i_t * BT + i_i * BC,), (BC,), (0,) + ) + b_b = tl.load(p_b, boundary_check=(0,)) + + b_A = tl.zeros([BC, BC], dtype=tl.float32) + b_Aqk = tl.zeros([BC, BC], dtype=tl.float32) + for i_k in range(tl.cdiv(K, BK)): + p_q = tl.make_block_ptr( + q, (T, K), (H * K, 1), (i_t * BT + i_i * BC, i_k * BK), (BC, BK), (1, 0) + ) + p_k = tl.make_block_ptr( + k, (T, K), (H * K, 1), (i_t * BT + i_i * BC, i_k * BK), (BC, BK), (1, 0) + ) + p_g = tl.make_block_ptr( + g, (T, K), (H * K, 1), (i_t * BT + i_i * BC, i_k * BK), (BC, BK), (1, 0) + ) + b_kt = tl.make_block_ptr( + k, (K, T), (1, H * K), (i_k * BK, i_t * BT + i_j * BC), (BK, BC), (0, 1) + ) + p_gk = tl.make_block_ptr( + g, (K, T), (1, H * K), (i_k * BK, i_t * BT + i_j * BC), (BK, BC), (0, 1) + ) + + o_k = i_k * BK + tl.arange(0, BK) + m_k = o_k < K + # [BK,] + b_gn = tl.load(g + (i_t * BT + i_i * BC) * H * K + o_k, mask=m_k, other=0) + # [BC, BK] + b_g = tl.load(p_g, boundary_check=(0, 1)) + b_k = tl.load(p_k, boundary_check=(0, 1)) * exp2(b_g - b_gn[None, :]) + # [BK, BC] + b_gk = tl.load(p_gk, boundary_check=(0, 1)) + b_kt = tl.load(b_kt, boundary_check=(0, 1)) + # [BC, BC] + b_ktg = b_kt * exp2(b_gn[:, None] - b_gk) + b_A += tl.dot(b_k, b_ktg) + + b_q = tl.load(p_q, boundary_check=(0, 1)) + b_qg = b_q * exp2(b_g - b_gn[None, :]) * scale + b_Aqk += tl.dot(b_qg, b_ktg) + + b_A *= b_b[:, None] + + p_A = tl.make_block_ptr( + A, (T, BT), (H * BT, 1), (i_t * BT + i_i * BC, i_j * BC), (BC, BC), (1, 0) + ) + tl.store(p_A, b_A.to(A.dtype.element_ty), boundary_check=(0, 1)) + p_Aqk = tl.make_block_ptr( + Aqk, (T, BT), (H * BT, 1), (i_t * BT + i_i * BC, i_j * BC), (BC, BC), (1, 0) + ) + tl.store(p_Aqk, b_Aqk.to(Aqk.dtype.element_ty), boundary_check=(0, 1)) + + +@triton.heuristics({"IS_VARLEN": lambda args: args["cu_seqlens"] is not None}) +@triton.autotune( + configs=[triton.Config({}, num_warps=num_warps) for num_warps in [1, 2, 4, 8]], + key=["BK", "BT"], +) +@triton.jit(do_not_specialize=["T"]) +def chunk_kda_scaled_dot_kkt_fwd_kernel_intra_sub_intra( + q, + k, + g, + beta, + A, + Aqk, + scale, + cu_seqlens, + chunk_indices, + T, + H: tl.constexpr, + K: tl.constexpr, + BT: tl.constexpr, + BC: tl.constexpr, + BK: tl.constexpr, + IS_VARLEN: tl.constexpr, +): + i_t, i_i, i_bh = tl.program_id(0), tl.program_id(1), tl.program_id(2) + i_b, i_h = i_bh // H, i_bh % H + if IS_VARLEN: + i_n, i_t = ( + tl.load(chunk_indices + i_t * 2).to(tl.int32), + tl.load(chunk_indices + i_t * 2 + 1).to(tl.int32), + ) + bos, eos = ( + tl.load(cu_seqlens + i_n).to(tl.int32), + tl.load(cu_seqlens + i_n + 1).to(tl.int32), + ) + T = eos - bos + else: + bos, eos = i_b * T, i_b * T + T + + if i_t * BT + i_i * BC >= T: + return + + o_i = tl.arange(0, BC) + o_k = tl.arange(0, BK) + m_k = o_k < K + m_A = (i_t * BT + i_i * BC + o_i) < T + o_A = (bos + i_t * BT + i_i * BC + o_i) * H * BT + i_h * BT + i_i * BC + + p_q = tl.make_block_ptr( + q + (bos * H + i_h) * K, + (T, K), + (H * K, 1), + (i_t * BT + i_i * BC, 0), + (BC, BK), + (1, 0), + ) + p_k = tl.make_block_ptr( + k + (bos * H + i_h) * K, + (T, K), + (H * K, 1), + (i_t * BT + i_i * BC, 0), + (BC, BK), + (1, 0), + ) + p_g = tl.make_block_ptr( + g + (bos * H + i_h) * K, + (T, K), + (H * K, 1), + (i_t * BT + i_i * BC, 0), + (BC, BK), + (1, 0), + ) + b_q = tl.load(p_q, boundary_check=(0, 1)) + b_k = tl.load(p_k, boundary_check=(0, 1)) + b_g = tl.load(p_g, boundary_check=(0, 1)) + + p_b = beta + (bos + i_t * BT + i_i * BC + o_i) * H + i_h + b_k = b_k * tl.load(p_b, mask=m_A, other=0)[:, None] + + p_kt = k + (bos + i_t * BT + i_i * BC) * H * K + i_h * K + o_k + p_gk = g + (bos + i_t * BT + i_i * BC) * H * K + i_h * K + o_k + + for j in range(0, min(BC, T - i_t * BT - i_i * BC)): + b_kt = tl.load(p_kt, mask=m_k, other=0).to(tl.float32) + b_gk = tl.load(p_gk, mask=m_k, other=0).to(tl.float32) + b_ktg = b_kt[None, :] * exp2(b_g - b_gk[None, :]) + b_A = tl.sum(b_k * b_ktg, 1) + b_A = tl.where(o_i > j, b_A, 0.0) + b_Aqk = tl.sum(b_q * b_ktg, 1) + b_Aqk = tl.where(o_i >= j, b_Aqk * scale, 0.0) + tl.store(A + o_A + j, b_A, mask=m_A) + tl.store(Aqk + o_A + j, b_Aqk, mask=m_A) + p_kt += H * K + p_gk += H * K + + +def chunk_kda_scaled_dot_kkt_fwd( + q: torch.Tensor, + k: torch.Tensor, + gk: torch.Tensor | None = None, + beta: torch.Tensor | None = None, + scale: float | None = None, + cu_seqlens: torch.Tensor | None = None, + chunk_indices: torch.Tensor | None = None, + chunk_size: int = FLA_CHUNK_SIZE, + output_dtype: torch.dtype = torch.float32, +) -> tuple[torch.Tensor, torch.Tensor]: + r""" + Compute beta * K * K^T. + + Args: + k (torch.Tensor): + The key tensor of shape `[B, T, H, K]`. + beta (torch.Tensor): + The beta tensor of shape `[B, T, H]`. + gk (torch.Tensor): + The cumulative sum of the gate tensor of shape `[B, T, H, K]` applied to the key tensor. Default: `None`. + cu_seqlens (torch.Tensor): + The cumulative sequence lengths of the input tensor. + Default: None + chunk_size (int): + The chunk size. Default: 64. + output_dtype (torch.dtype): + The dtype of the output tensor. Default: `torch.float32` + + Returns: + beta * K * K^T of shape `[B, T, H, BT]` where `BT` is the chunk size. + """ + B, T, H, K = k.shape + assert K <= 256 + BT = chunk_size + if chunk_indices is None and cu_seqlens is not None: + chunk_indices = prepare_chunk_indices(cu_seqlens, BT) + NT = cdiv(T, BT) if cu_seqlens is None else len(chunk_indices) + + BC = min(16, BT) + NC = cdiv(BT, BC) + BK = max(next_power_of_2(K), 16) + A = torch.zeros(B, T, H, BT, device=k.device, dtype=output_dtype) + Aqk = torch.zeros(B, T, H, BT, device=k.device, dtype=output_dtype) + grid = (NT, NC * NC, B * H) + chunk_kda_scaled_dot_kkt_fwd_kernel_intra_sub_inter[grid]( + q=q, + k=k, + g=gk, + beta=beta, + A=A, + Aqk=Aqk, + scale=scale, + cu_seqlens=cu_seqlens, + chunk_indices=chunk_indices, + T=T, + H=H, + K=K, + BT=BT, + BC=BC, + NC=NC, + ) + + grid = (NT, NC, B * H) + chunk_kda_scaled_dot_kkt_fwd_kernel_intra_sub_intra[grid]( + q=q, + k=k, + g=gk, + beta=beta, + A=A, + Aqk=Aqk, + scale=scale, + cu_seqlens=cu_seqlens, + chunk_indices=chunk_indices, + T=T, + H=H, + K=K, + BT=BT, + BC=BC, + BK=BK, + ) + return A, Aqk + + +@triton.heuristics( + { + "STORE_QG": lambda args: args["qg"] is not None, + "STORE_KG": lambda args: args["kg"] is not None, + "IS_VARLEN": lambda args: args["cu_seqlens"] is not None, + } +) +@triton.autotune( + configs=[ + triton.Config({}, num_warps=num_warps, num_stages=num_stages) + for num_warps in [2, 4, 8] + for num_stages in [2, 3, 4] + ], + key=["H", "K", "V", "BT", "BK", "BV", "IS_VARLEN"], +) +@triton.jit(do_not_specialize=["T"]) +def recompute_w_u_fwd_kernel( + q, + k, + qg, + kg, + v, + beta, + w, + u, + A, + gk, + cu_seqlens, + chunk_indices, + T, + H: tl.constexpr, + K: tl.constexpr, + V: tl.constexpr, + BT: tl.constexpr, + BK: tl.constexpr, + BV: tl.constexpr, + STORE_QG: tl.constexpr, + STORE_KG: tl.constexpr, + IS_VARLEN: tl.constexpr, + DOT_PRECISION: tl.constexpr, +): + i_t, i_bh = tl.program_id(0), tl.program_id(1) + i_b, i_h = i_bh // H, i_bh % H + if IS_VARLEN: + i_n, i_t = ( + tl.load(chunk_indices + i_t * 2).to(tl.int32), + tl.load(chunk_indices + i_t * 2 + 1).to(tl.int32), + ) + bos, eos = ( + tl.load(cu_seqlens + i_n).to(tl.int32), + tl.load(cu_seqlens + i_n + 1).to(tl.int32), + ) + T = eos - bos + else: + bos, eos = i_b * T, i_b * T + T + p_b = tl.make_block_ptr(beta + bos * H + i_h, (T,), (H,), (i_t * BT,), (BT,), (0,)) + b_b = tl.load(p_b, boundary_check=(0,)) + + p_A = tl.make_block_ptr( + A + (bos * H + i_h) * BT, (T, BT), (H * BT, 1), (i_t * BT, 0), (BT, BT), (1, 0) + ) + b_A = tl.load(p_A, boundary_check=(0, 1)) + + for i_v in range(tl.cdiv(V, BV)): + p_v = tl.make_block_ptr( + v + (bos * H + i_h) * V, + (T, V), + (H * V, 1), + (i_t * BT, i_v * BV), + (BT, BV), + (1, 0), + ) + p_u = tl.make_block_ptr( + u + (bos * H + i_h) * V, + (T, V), + (H * V, 1), + (i_t * BT, i_v * BV), + (BT, BV), + (1, 0), + ) + b_v = tl.load(p_v, boundary_check=(0, 1)) + b_vb = (b_v * b_b[:, None]).to(b_v.dtype) + b_u = tl.dot(b_A, b_vb, input_precision=DOT_PRECISION) + tl.store(p_u, b_u.to(p_u.dtype.element_ty), boundary_check=(0, 1)) + + for i_k in range(tl.cdiv(K, BK)): + p_w = tl.make_block_ptr( + w + (bos * H + i_h) * K, + (T, K), + (H * K, 1), + (i_t * BT, i_k * BK), + (BT, BK), + (1, 0), + ) + p_k = tl.make_block_ptr( + k + (bos * H + i_h) * K, + (T, K), + (H * K, 1), + (i_t * BT, i_k * BK), + (BT, BK), + (1, 0), + ) + b_k = tl.load(p_k, boundary_check=(0, 1)) + b_kb = b_k * b_b[:, None] + + p_gk = tl.make_block_ptr( + gk + (bos * H + i_h) * K, + (T, K), + (H * K, 1), + (i_t * BT, i_k * BK), + (BT, BK), + (1, 0), + ) + b_gk = tl.load(p_gk, boundary_check=(0, 1)) + b_kb *= exp2(b_gk) + if STORE_QG: + p_q = tl.make_block_ptr( + q + (bos * H + i_h) * K, + (T, K), + (H * K, 1), + (i_t * BT, i_k * BK), + (BT, BK), + (1, 0), + ) + p_qg = tl.make_block_ptr( + qg + (bos * H + i_h) * K, + (T, K), + (H * K, 1), + (i_t * BT, i_k * BK), + (BT, BK), + (1, 0), + ) + b_q = tl.load(p_q, boundary_check=(0, 1)) + b_qg = b_q * exp2(b_gk) + tl.store(p_qg, b_qg.to(p_qg.dtype.element_ty), boundary_check=(0, 1)) + if STORE_KG: + last_idx = min(i_t * BT + BT, T) - 1 + + o_k = i_k * BK + tl.arange(0, BK) + m_k = o_k < K + b_gn = tl.load( + gk + ((bos + last_idx) * H + i_h) * K + o_k, mask=m_k, other=0.0 + ) + b_kg = b_k * exp2(b_gn - b_gk) + + p_kg = tl.make_block_ptr( + kg + (bos * H + i_h) * K, + (T, K), + (H * K, 1), + (i_t * BT, i_k * BK), + (BT, BK), + (1, 0), + ) + tl.store(p_kg, b_kg.to(p_kg.dtype.element_ty), boundary_check=(0, 1)) + + b_w = tl.dot(b_A, b_kb.to(b_k.dtype)) + tl.store(p_w, b_w.to(p_w.dtype.element_ty), boundary_check=(0, 1)) + + +def recompute_w_u_fwd( + k: torch.Tensor, + v: torch.Tensor, + beta: torch.Tensor, + A: torch.Tensor, + q: torch.Tensor | None = None, + gk: torch.Tensor | None = None, + cu_seqlens: torch.Tensor | None = None, + chunk_indices: torch.Tensor | None = None, +) -> tuple[torch.Tensor, torch.Tensor]: + B, T, H, K, V = *k.shape, v.shape[-1] + BT = A.shape[-1] + BK = 64 + BV = 64 + + if chunk_indices is None and cu_seqlens is not None: + chunk_indices = prepare_chunk_indices(cu_seqlens, BT) + NT = cdiv(T, BT) if cu_seqlens is None else len(chunk_indices) + + w = torch.empty_like(k) + u = torch.empty_like(v) + kg = torch.empty_like(k) if gk is not None else None + recompute_w_u_fwd_kernel[(NT, B * H)]( + q=q, + k=k, + qg=None, + kg=kg, + v=v, + beta=beta, + w=w, + u=u, + A=A, + gk=gk, + cu_seqlens=cu_seqlens, + chunk_indices=chunk_indices, + T=T, + H=H, + K=K, + V=V, + BT=BT, + BK=BK, + BV=BV, + DOT_PRECISION="ieee", + ) + return w, u, None, kg + + +@triton.heuristics({"IS_VARLEN": lambda args: args["cu_seqlens"] is not None}) +@triton.autotune( + configs=[ + triton.Config({"BK": BK, "BV": BV}, num_warps=num_warps, num_stages=num_stages) + for BK in [32, 64] + for BV in [64, 128] + for num_warps in [2, 4, 8] + for num_stages in [2, 3, 4] + ], + key=["BT"], +) +@triton.jit(do_not_specialize=["T"]) +def chunk_gla_fwd_kernel_o( + q, + v, + g, + h, + o, + A, + cu_seqlens, + chunk_indices, + scale, + T, + H: tl.constexpr, + K: tl.constexpr, + V: tl.constexpr, + BT: tl.constexpr, + BK: tl.constexpr, + BV: tl.constexpr, + IS_VARLEN: tl.constexpr, +): + i_v, i_t, i_bh = tl.program_id(0), tl.program_id(1), tl.program_id(2) + i_b, i_h = i_bh // H, i_bh % H + if IS_VARLEN: + i_tg = i_t + i_n, i_t = ( + tl.load(chunk_indices + i_t * 2).to(tl.int32), + tl.load(chunk_indices + i_t * 2 + 1).to(tl.int32), + ) + bos, eos = ( + tl.load(cu_seqlens + i_n).to(tl.int32), + tl.load(cu_seqlens + i_n + 1).to(tl.int32), + ) + T = eos - bos + NT = tl.cdiv(T, BT) + else: + NT = tl.cdiv(T, BT) + i_tg = i_b * NT + i_t + bos, eos = i_b * T, i_b * T + T + + m_s = tl.arange(0, BT)[:, None] >= tl.arange(0, BT)[None, :] + + b_o = tl.zeros([BT, BV], dtype=tl.float32) + for i_k in range(tl.cdiv(K, BK)): + p_q = tl.make_block_ptr( + q + (bos * H + i_h) * K, + (T, K), + (H * K, 1), + (i_t * BT, i_k * BK), + (BT, BK), + (1, 0), + ) + p_g = tl.make_block_ptr( + g + (bos * H + i_h) * K, + (T, K), + (H * K, 1), + (i_t * BT, i_k * BK), + (BT, BK), + (1, 0), + ) + p_h = tl.make_block_ptr( + h + (i_tg * H + i_h) * K * V, + (V, K), + (K, 1), + (i_v * BV, i_k * BK), + (BV, BK), + (1, 0), + ) + + # [BT, BK] + b_q = tl.load(p_q, boundary_check=(0, 1)) + b_q = (b_q * scale).to(b_q.dtype) + # [BT, BK] + b_g = tl.load(p_g, boundary_check=(0, 1)) + # [BT, BK] + b_qg = (b_q * exp2(b_g)).to(b_q.dtype) + # [BV, BK] + b_h = tl.load(p_h, boundary_check=(0, 1)) + # [BT, BV] + if i_k >= 0: + b_o += tl.dot(b_qg, tl.trans(b_h).to(b_qg.dtype)) + p_v = tl.make_block_ptr( + v + (bos * H + i_h) * V, + (T, V), + (H * V, 1), + (i_t * BT, i_v * BV), + (BT, BV), + (1, 0), + ) + p_o = tl.make_block_ptr( + o + (bos * H + i_h) * V, + (T, V), + (H * V, 1), + (i_t * BT, i_v * BV), + (BT, BV), + (1, 0), + ) + p_A = tl.make_block_ptr( + A + (bos * H + i_h) * BT, (T, BT), (H * BT, 1), (i_t * BT, 0), (BT, BT), (1, 0) + ) + # [BT, BV] + b_v = tl.load(p_v, boundary_check=(0, 1)) + # [BT, BT] + b_A = tl.load(p_A, boundary_check=(0, 1)) + b_A = tl.where(m_s, b_A, 0.0).to(b_v.dtype) + b_o += tl.dot(b_A, b_v, allow_tf32=False) + tl.store(p_o, b_o.to(p_o.dtype.element_ty), boundary_check=(0, 1)) + + +def chunk_gla_fwd_o_gk( + q: torch.Tensor, + v: torch.Tensor, + g: torch.Tensor, + A: torch.Tensor, + h: torch.Tensor, + o: torch.Tensor, + scale: float, + cu_seqlens: torch.Tensor | None = None, + chunk_indices: torch.Tensor | None = None, + chunk_size: int = FLA_CHUNK_SIZE, +): + B, T, H, K, V = *q.shape, v.shape[-1] + BT = chunk_size + + if chunk_indices is None and cu_seqlens is not None: + chunk_indices = prepare_chunk_indices(cu_seqlens, chunk_size) + NT = cdiv(T, BT) if cu_seqlens is None else len(chunk_indices) + + def grid(meta): + return (cdiv(V, meta["BV"]), NT, B * H) + + chunk_gla_fwd_kernel_o[grid]( + q=q, + v=v, + g=g, + h=h, + o=o, + A=A, + cu_seqlens=cu_seqlens, + chunk_indices=chunk_indices, + scale=scale, + T=T, + H=H, + K=K, + V=V, + BT=BT, + ) + return o + + +@triton.heuristics( + { + "HAS_BIAS": lambda args: args["g_bias"] is not None, + "IS_VARLEN": lambda args: args["cu_seqlens"] is not None, + } +) +@triton.autotune( + configs=[ + triton.Config({"BD": BD}, num_warps=num_warps) + for BD in [32, 64] + for num_warps in [2, 4, 8] + ], + key=["H", "D", "BT", "IS_VARLEN"], +) +@triton.jit(do_not_specialize=["T"]) +def kda_gate_cumsum_fwd_kernel( + g, + A, + y, + g_bias, + cu_seqlens, + chunk_indices, + cumsum_scale, + beta, + threshold, + SAFE_GATE: tl.constexpr, + LOWER_BOUND: tl.constexpr, + T, + H: tl.constexpr, + D: tl.constexpr, + BT: tl.constexpr, + BD: tl.constexpr, + HAS_BIAS: tl.constexpr, + IS_VARLEN: tl.constexpr, +): + i_d, i_t, i_bh = tl.program_id(0), tl.program_id(1), tl.program_id(2) + i_b, i_h = i_bh // H, i_bh % H + if IS_VARLEN: + i_n, i_t = ( + tl.load(chunk_indices + i_t * 2).to(tl.int32), + tl.load(chunk_indices + i_t * 2 + 1).to(tl.int32), + ) + bos, eos = ( + tl.load(cu_seqlens + i_n).to(tl.int32), + tl.load(cu_seqlens + i_n + 1).to(tl.int32), + ) + T = eos - bos + else: + bos = i_b * T + + p_g = tl.make_block_ptr( + g + (bos * H + i_h) * D, + (T, D), + (H * D, 1), + (i_t * BT, i_d * BD), + (BT, BD), + (1, 0), + ) + p_y = tl.make_block_ptr( + y + (bos * H + i_h) * D, + (T, D), + (H * D, 1), + (i_t * BT, i_d * BD), + (BT, BD), + (1, 0), + ) + + b_g = tl.load(p_g, boundary_check=(0, 1)).to(tl.float32) + if HAS_BIAS: + o_d = i_d * BD + tl.arange(0, BD) + b_bias = tl.load(g_bias + i_h * D + o_d, mask=o_d < D, other=0.0).to(tl.float32) + b_g = b_g + b_bias[None, :] + + b_a = tl.load(A + i_h).to(tl.float32) + b_a = tl.exp(b_a) if SAFE_GATE else -tl.exp(b_a) + if SAFE_GATE: + # y = lower_bound * sigmoid(exp(A) * (g + g_bias)), bounded to + # (lower_bound, 0) for safe-gate checkpoints. + b_gate = LOWER_BOUND / (1.0 + tl.exp(-(b_a * b_g))) + else: + b_g_scaled = b_g * beta + b_softplus = tl.where( + b_g_scaled > threshold, + b_g, + (1.0 / beta) * log(1.0 + tl.exp(b_g_scaled)), + ) + b_gate = b_a * b_softplus + + # Out-of-bounds rows (load returns 0, but softplus/bias can still make + # b_gate non-zero) participate in the dot product. They only contribute to + # out-of-bounds output rows, which are masked away by `boundary_check` on + # the store, so visible output matches unfused gate + chunk-local cumsum. + o_t = tl.arange(0, BT) + m_cumsum = tl.where(o_t[:, None] >= o_t[None, :], 1.0, 0.0) + b_y = tl.dot(m_cumsum, b_gate, allow_tf32=False) * cumsum_scale + tl.store(p_y, b_y.to(p_y.dtype.element_ty), boundary_check=(0, 1)) + + +def fused_kda_gate_chunk_cumsum( + raw_g: torch.Tensor, + A_log: torch.Tensor, + g_bias: torch.Tensor | None = None, + beta: float = 1.0, + threshold: float = 20.0, + cu_seqlens: torch.Tensor | None = None, + chunk_indices: torch.Tensor | None = None, + chunk_size: int = FLA_CHUNK_SIZE, + output_dtype: torch.dtype | None = torch.float, + safe_gate: bool = False, + lower_bound: float = -5.0, +) -> torch.Tensor: + if cu_seqlens is not None: + assert raw_g.shape[0] == 1, ( + "Only batch size 1 is supported when cu_seqlens are provided" + ) + B, T, H, D = raw_g.shape + if chunk_indices is None and cu_seqlens is not None: + chunk_indices = prepare_chunk_indices(cu_seqlens, chunk_size) + NT = cdiv(T, chunk_size) if cu_seqlens is None else len(chunk_indices) + + A_log = A_log.reshape(-1) + if g_bias is not None: + g_bias = g_bias.reshape(-1) + y = torch.empty_like(raw_g, dtype=output_dtype or raw_g.dtype) + + def grid(meta): + return (cdiv(meta["D"], meta["BD"]), NT, B * H) + + kda_gate_cumsum_fwd_kernel[grid]( + g=raw_g, + A=A_log, + y=y, + g_bias=g_bias, + cu_seqlens=cu_seqlens, + chunk_indices=chunk_indices, + # RCP_LN2 folds in the natural-log -> log2 conversion so downstream + # exp2-based kernels reproduce exp(g). Keep this in sync with the + # `use_exp2=True` path in `_chunk_kda_fwd_with_cumulative_g`. + cumsum_scale=RCP_LN2, + beta=beta, + threshold=threshold, + SAFE_GATE=safe_gate, + LOWER_BOUND=lower_bound, + T=T, + H=H, + D=D, + BT=chunk_size, + ) + return y + + +def _chunk_kda_fwd_with_cumulative_g( + q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, + g: torch.Tensor, + beta: torch.Tensor, + scale: float, + initial_state: torch.Tensor, + output_final_state: bool, + cu_seqlens: torch.Tensor | None = None, + chunk_indices: torch.Tensor | None = None, + chunk_size: int = FLA_CHUNK_SIZE, +): + # `g` must already be chunk-local cumulatively-summed AND scaled by + # RCP_LN2 (so the downstream exp2-based kernels reproduce exp(g)). + # Use `chunk_kda_fwd` or `chunk_kda_with_fused_gate_fwd` instead of + # calling this helper directly unless that invariant is upheld. + # the intra Aqk is kept in fp32 + # the computation has very marginal effect on the entire throughput + A, Aqk = chunk_kda_scaled_dot_kkt_fwd( + q=q, + k=k, + gk=g, + beta=beta, + scale=scale, + cu_seqlens=cu_seqlens, + chunk_indices=chunk_indices, + chunk_size=chunk_size, + output_dtype=torch.float32, + ) + A = solve_tril(A=A, cu_seqlens=cu_seqlens, output_dtype=k.dtype) + w, u, _, kg = recompute_w_u_fwd( + k=k, + v=v, + beta=beta, + A=A, + gk=g, + cu_seqlens=cu_seqlens, + chunk_indices=chunk_indices, + ) + del A + h, v_new, final_state = chunk_gated_delta_rule_fwd_h( + k=kg, + w=w, + u=u, + gk=g, + initial_state=initial_state, + output_final_state=output_final_state, + cu_seqlens=cu_seqlens, + chunk_indices=chunk_indices, + use_exp2=True, + ) + del w, u, kg + o = chunk_gla_fwd_o_gk( + q=q, + v=v_new, + g=g, + A=Aqk, + h=h, + o=v, + scale=scale, + cu_seqlens=cu_seqlens, + chunk_indices=chunk_indices, + chunk_size=chunk_size, + ) + del Aqk, v_new, h + return o, final_state + + +def chunk_kda_fwd( + q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, + g: torch.Tensor, + beta: torch.Tensor, + scale: float, + initial_state: torch.Tensor, + output_final_state: bool, + cu_seqlens: torch.Tensor | None = None, +): + chunk_size = FLA_CHUNK_SIZE + chunk_indices = ( + prepare_chunk_indices(cu_seqlens, chunk_size) + if cu_seqlens is not None + else None + ) + g = chunk_local_cumsum( + g, + chunk_size=chunk_size, + cu_seqlens=cu_seqlens, + chunk_indices=chunk_indices, + ) + # KDA evaluates cumulative gate decays with exp2. Convert from natural-log + # space so exp(x) is preserved as exp2(x / ln(2)). + g = g * RCP_LN2 + return _chunk_kda_fwd_with_cumulative_g( + q=q, + k=k, + v=v, + g=g, + beta=beta, + scale=scale, + initial_state=initial_state, + output_final_state=output_final_state, + cu_seqlens=cu_seqlens, + chunk_indices=chunk_indices, + chunk_size=chunk_size, + ) + + +def chunk_kda_with_fused_gate_fwd( + q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, + raw_g: torch.Tensor, + beta: torch.Tensor, + A_log: torch.Tensor, + g_bias: torch.Tensor | None, + scale: float, + initial_state: torch.Tensor, + output_final_state: bool, + cu_seqlens: torch.Tensor | None = None, + safe_gate: bool = False, + lower_bound: float = -5.0, +): + chunk_size = FLA_CHUNK_SIZE + chunk_indices = ( + prepare_chunk_indices(cu_seqlens, chunk_size) + if cu_seqlens is not None + else None + ) + g = fused_kda_gate_chunk_cumsum( + raw_g, + A_log=A_log, + g_bias=g_bias, + cu_seqlens=cu_seqlens, + chunk_indices=chunk_indices, + chunk_size=chunk_size, + safe_gate=safe_gate, + lower_bound=lower_bound, + ) + return _chunk_kda_fwd_with_cumulative_g( + q=q, + k=k, + v=v, + g=g, + beta=beta, + scale=scale, + initial_state=initial_state, + output_final_state=output_final_state, + cu_seqlens=cu_seqlens, + chunk_indices=chunk_indices, + chunk_size=chunk_size, + ) + + +def chunk_kda( + q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, + g: torch.Tensor, + beta: torch.Tensor, + scale: float = None, + initial_state: torch.Tensor = None, + output_final_state: bool = False, + use_qk_l2norm_in_kernel: bool = False, + cu_seqlens: torch.Tensor | None = None, + **kwargs, +): + if scale is None: + scale = k.shape[-1] ** -0.5 + + if use_qk_l2norm_in_kernel: + q = l2norm_fwd(q.contiguous()) + k = l2norm_fwd(k.contiguous()) + + o, final_state = chunk_kda_fwd( + q=q, + k=k, + v=v.contiguous(), + g=g.contiguous(), + beta=beta.contiguous(), + scale=scale, + initial_state=initial_state.contiguous(), + output_final_state=output_final_state, + cu_seqlens=cu_seqlens, + ) + return o, final_state + + +def chunk_kda_with_fused_gate( + q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, + raw_g: torch.Tensor, + beta: torch.Tensor, + A_log: torch.Tensor, + g_bias: torch.Tensor | None, + scale: float | None = None, + initial_state: torch.Tensor | None = None, + output_final_state: bool = False, + use_qk_l2norm_in_kernel: bool = False, + cu_seqlens: torch.Tensor | None = None, + safe_gate: bool = False, + lower_bound: float = -5.0, + **kwargs, +): + """Run chunk KDA from raw gate projection using fused gate+cumsum.""" + if scale is None: + scale = k.shape[-1] ** -0.5 + + if use_qk_l2norm_in_kernel: + q = l2norm_fwd(q.contiguous()) + k = l2norm_fwd(k.contiguous()) + + o, final_state = chunk_kda_with_fused_gate_fwd( + q=q, + k=k, + v=v.contiguous(), + raw_g=raw_g.contiguous(), + beta=beta.contiguous(), + A_log=A_log, + g_bias=g_bias, + scale=scale, + initial_state=initial_state.contiguous() if initial_state is not None else None, + output_final_state=output_final_state, + cu_seqlens=cu_seqlens, + safe_gate=safe_gate, + lower_bound=lower_bound, + ) + return o, final_state + + +@triton.autotune( + configs=[ + triton.Config({"BT": bt}, num_warps=nw, num_stages=ns) + for bt in BT_LIST_AUTOTUNE + for nw in NUM_WARPS_AUTOTUNE + for ns in [2, 3] + ], + key=["H", "D"], +) +@triton.jit +def kda_gate_fwd_kernel( + g, + A, + y, + g_bias, + beta: tl.constexpr, + threshold: tl.constexpr, + SAFE_GATE: tl.constexpr, + LOWER_BOUND: tl.constexpr, + T, + H, + D: tl.constexpr, + BT: tl.constexpr, + BD: tl.constexpr, + HAS_BIAS: tl.constexpr, +): + i_t, i_h = tl.program_id(0), tl.program_id(1) + n_t = i_t * BT + + b_a = tl.load(A + i_h).to(tl.float32) + b_a = tl.exp(b_a) if SAFE_GATE else -tl.exp(b_a) + + stride_row = H * D + stride_col = 1 + + g_ptr = tl.make_block_ptr( + base=g + i_h * D, + shape=(T, D), + strides=(stride_row, stride_col), + offsets=(n_t, 0), + block_shape=(BT, BD), + order=(1, 0), + ) + + y_ptr = tl.make_block_ptr( + base=y + i_h * D, + shape=(T, D), + strides=(stride_row, stride_col), + offsets=(n_t, 0), + block_shape=(BT, BD), + order=(1, 0), + ) + + b_g = tl.load(g_ptr, boundary_check=(0, 1)).to(tl.float32) + + if HAS_BIAS: + n_d = tl.arange(0, BD) + bias_mask = n_d < D + b_bias = tl.load(g_bias + i_h * D + n_d, mask=bias_mask, other=0.0).to( + tl.float32 + ) + b_g = b_g + b_bias[None, :] + + if SAFE_GATE: + # y = lower_bound * sigmoid(exp(A) * (g + g_bias)), bounded to + # (lower_bound, 0) for safe-gate checkpoints. + b_y = LOWER_BOUND / (1.0 + tl.exp(-(b_a * b_g))) + else: + # softplus(x, beta) = (1/beta) * log(1 + exp(beta * x)) + # When beta * x > threshold, use linear approximation x + # Use threshold to switch to linear when beta*x > threshold + g_scaled = b_g * beta + use_linear = g_scaled > threshold + sp = tl.where(use_linear, b_g, (1.0 / beta) * log(1.0 + tl.exp(g_scaled))) + b_y = b_a * sp + + tl.store(y_ptr, b_y.to(y.dtype.element_ty), boundary_check=(0, 1)) + + +def fused_kda_gate( + g: torch.Tensor, + A: torch.Tensor, + head_k_dim: int, + g_bias: torch.Tensor | None = None, + beta: float = 1.0, + threshold: float = 20.0, + safe_gate: bool = False, + lower_bound: float | None = -5.0, +) -> torch.Tensor: + """ + Forward pass for KDA gate: + input g: [..., H*D] + param A: [H] or [1, 1, H, 1] + beta: softplus beta parameter (softplus branch only) + threshold: softplus threshold parameter (softplus branch only) + safe_gate: when False (default) compute y = -exp(A)*softplus(g+g_bias); + when True compute the bounded y = lower_bound*sigmoid(exp(A)*(g+g_bias)) + lower_bound: floor for the safe_gate branch (default -5.0) + return : [..., H, D] + """ + orig_shape = g.shape[:-1] + + g = g.view(-1, g.shape[-1]) + T = g.shape[0] + HD = g.shape[1] + H = A.numel() + assert H * head_k_dim == HD + + y = torch.empty_like(g, dtype=torch.float32) + + def grid(meta): + return (cdiv(T, meta["BT"]), H) + + kda_gate_fwd_kernel[grid]( + g, + A, + y, + g_bias, + beta, + threshold, + safe_gate, + lower_bound if lower_bound is not None else -5.0, + T, + H, + head_k_dim, + BD=next_power_of_2(head_k_dim), + HAS_BIAS=g_bias is not None, + ) + + y = y.view(*orig_shape, H, head_k_dim) + return y diff --git a/vllm/models/glm5next/nvidia/attention.py b/vllm/models/glm5next/nvidia/attention.py new file mode 100644 index 000000000000..0be749348494 --- /dev/null +++ b/vllm/models/glm5next/nvidia/attention.py @@ -0,0 +1,590 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project + +import torch +import torch.nn.functional as F +from torch import nn + +from vllm.config import ( + CacheConfig, + VllmConfig, +) +from vllm.distributed import ( + get_tensor_model_parallel_world_size, +) +from vllm.logger import init_logger +from vllm.model_executor.layers.layernorm import LayerNorm, RMSNorm +from vllm.model_executor.layers.linear import ( + ColumnParallelLinear, + MergedColumnParallelLinear, + ReplicatedLinear, + RowParallelLinear, +) +from vllm.model_executor.layers.mla import MLAModules, MultiHeadLatentAttentionWrapper +from vllm.model_executor.layers.quantization.base_config import QuantizationConfig +from vllm.model_executor.layers.rotary_embedding import RotaryEmbedding, get_rope +from vllm.model_executor.layers.sparse_attn_indexer_kpool import SparseAttnIndexerKpool +from vllm.model_executor.models.deepseek_v2 import ( + DeepSeekV2FusedQkvAProjLinear, + DeepseekV32IndexerCache, + yarn_get_mscale, +) +from vllm.model_executor.utils import maybe_disable_graph_partition +from vllm.models.glm5next.nvidia.ops.kpool_compress import fwht128_quant_fp8 +from vllm.platforms import current_platform +from vllm.transformers_utils.configs.glm5_next import Glm5NextConfig +from vllm.utils.deep_gemm import PAGED_MQA_PAGE_SIZES +from vllm.v1.kv_cache_interface import KpoolTailSpec, MLAAttentionSpec + +logger = init_logger(__name__) + +# Shared torch.compile config for the indexer's small-kernel leaves. The MLA +# indexer runs under breakable-CG (CompilationMode.NONE), which blocks FX-graph +# fusion of the surrounding eager ops; carving each cluster into its own +# @torch.compile leaf (backend==inductor) still fuses them. Matches the +# grouped_topk / _cast_sigmoid leaf pattern. +_INDEXER_COMPILE = dict( + dynamic=True, + backend=current_platform.simple_compile_backend, + options=maybe_disable_graph_partition(current_platform.simple_compile_backend), +) + + +@torch.compile(**_INDEXER_COMPILE) +def _fused_indexer_k_norm( + x: torch.Tensor, weight: torch.Tensor, bias: torch.Tensor, dim: int, eps: float +) -> torch.Tensor: + # Fuse fp32 cast + layer_norm + cast-back (was 3 kernels) into one. + return F.layer_norm(x.float(), (dim,), weight, bias, eps).type_as(x) + + +@torch.compile(**_INDEXER_COMPILE) +def _fused_indexer_weight_scale( + weights: torch.Tensor, q_scale: torch.Tensor, scale: float +) -> torch.Tensor: + # Fuse the weight-scaling muls (was 2 kernels) into one. `scale` folds + # softmax_scale (head_dim**-0.5) and n_head**-0.5 into a single constant. + return (weights.unsqueeze(-1) * q_scale * scale).squeeze(-1) + + +@torch.compile(**_INDEXER_COMPILE) +def _pad_indexer_heads(x: torch.Tensor, pad: int) -> torch.Tensor: + # DeepGEMM MQA-logits needs num_heads in {32,64}; zero-pad the head dim. + # Fuse new_zeros + cat (was 2 kernels) into one. Pad values are zero (exact + # in fp8 e4m3 and zero-weight in the logits sum), so numerically a no-op. + return torch.cat([x, x.new_zeros(x.shape[0], pad, *x.shape[2:])], dim=1) + + +class Glm5NextIndexerCache(DeepseekV32IndexerCache): + """Indexer K cache that stores kpool-compressed entries. + + Setting ``tokens_per_state = index_kpool`` on the KV cache spec makes vLLM's + indexer metadata builder emit pool-granular ``slot_mapping`` / + ``seq_lens`` / ``cu_seq_lens`` / ``page_table`` for free, and shrinks the + cache allocation store one state per ``index_kpool`` tokens. The pool + *content* (softmax-weighted sum vs keep-every-Nth) is computed by the + kpool compress kernel inside the indexer op — the cache only provides the + addressing, which is identical for both schemes. + + The indexer shares one block with the co-located MLA (a single + ``MLAAttentionSpec`` / block_table), so ``block_size`` is the model-wide + ``cache_config.block_size``. DeepGEMM's paged-MQA kernel + (``csrc/apis/attention.hpp``) requires ``block_kv`` to be exactly 32 or + 64, so the storage block is virtually split into pool pages of the + largest such size that tiles it (``storage_kernel_block_size``); this + needs ``block_size`` to be a multiple of ``index_kpool * 32`` (512 for + ``index_kpool = 16``). A smaller block (e.g. the default 64) silently + collapses ``storage_block_size`` (64 // 16 = 4) and only fails later at + the opaque C++ assert; ``get_kv_cache_spec`` guards this up front + instead. + """ + + def __init__( + self, + *, + head_dim: int, + dtype: torch.dtype, + prefix: str, + cache_config, + index_kpool: int, + ): + super().__init__( + head_dim=head_dim, dtype=dtype, prefix=prefix, cache_config=cache_config + ) + assert index_kpool > 1, "Glm5NextIndexerCache expects index_kpool > 1" + # Keep chunked-prefill boundaries aligned to complete pools. + assert cache_config.block_size % index_kpool == 0, ( + "Glm5NextIndexerCache: cache_config.block_size " + f"({cache_config.block_size}) must be a multiple of index_kpool " + f"({index_kpool}) so chunked-prefill boundaries stay pool-aligned." + ) + self._index_kpool = index_kpool + + def get_kv_cache_spec(self, vllm_config: VllmConfig): + from dataclasses import replace + + spec = super().get_kv_cache_spec(vllm_config) + # ``tokens_per_state`` is the KV-spec representation of kpool + # compression in the current cache-layout API. + assert isinstance(spec, MLAAttentionSpec) + spec = replace(spec, tokens_per_state=self._index_kpool) + + # DeepGEMM paged-MQA takes block_kv in {32, 64}; the storage block + # (= block_size // index_kpool) is virtually split into pool pages of + # the largest such size that tiles it, so it must be a multiple of 32. + storage_block_size = spec.block_size // self._index_kpool + assert ( + spec.block_size % self._index_kpool == 0 and storage_block_size % 32 == 0 + ), ( + "Glm5NextIndexerCache: kpool indexer requires cache block_size to " + f"be a multiple of index_kpool * 32 ({self._index_kpool * 32}) so " + "that DeepGEMM paged-MQA pool pages (32 or 64 entries) tile the " + f"storage block, got block_size={spec.block_size} -> " + f"storage_block_size={storage_block_size}." + ) + max_page_size = max(PAGED_MQA_PAGE_SIZES) + min_page_size = min(PAGED_MQA_PAGE_SIZES) + if storage_block_size <= max_page_size: + page_size = storage_block_size + elif storage_block_size % max_page_size == 0: + page_size = max_page_size + else: + page_size = min_page_size + return replace( + spec, + storage_block_size=page_size * self._index_kpool, + ) + + +class Glm5NextTailCache(DeepseekV32IndexerCache): + """Paged circular buffer for the kpool indexer's in-progress (tail) pool. + + Holds the trailing incomplete pool's raw K + gate score: one block of + ``index_kpool`` slots per request, overwritten in place by ``pos % kpool`` + as decode/spec-decode advances. Prefill seeds it (instead of discarding the + tail raw K+gate); the connector transfers it across PD; decode reads it to + compress the boundary pool correctly. ``KpoolTailSpec`` / + ``KpoolTailManager`` provide the no-prune, 1-block/req allocation that lets + the in-progress pool survive across steps and across transfer. + + Stores raw bf16 K (``head_dim``) as the "K" half of each block and the + bf16 gate score (``head_dim``) as the "V" half -- not the fp8-compressed + entry, which lives in ``Glm5NextIndexerCache``. + """ + + def __init__( + self, + *, + head_dim: int, + dtype: torch.dtype, + prefix: str, + cache_config, + index_kpool: int, + ): + super().__init__( + head_dim=head_dim, dtype=dtype, prefix=prefix, cache_config=cache_config + ) + assert index_kpool > 1, "Glm5NextTailCache expects index_kpool > 1" + self._index_kpool = index_kpool + + def get_kv_cache_spec(self, vllm_config: VllmConfig): + # The two head slots form [K, gate score] in the generic + # [block, head, state, content] cache view. + return KpoolTailSpec( + block_size=self._index_kpool, + num_kv_heads=2, + head_size=self.head_dim, + head_size_v=0, + dtype=torch.bfloat16, + sliding_window=self._index_kpool, + ) + + def get_attn_backend(self): + from vllm.v1.attention.backends.mla.indexer import KpoolTailBackend + + return KpoolTailBackend + + +class Indexer(nn.Module): + def __init__( + self, + vllm_config: VllmConfig, + config: Glm5NextConfig, + hidden_size: int, + q_lora_rank: int, + quant_config: QuantizationConfig | None, + cache_config: CacheConfig | None, + topk_indices_buffer: torch.Tensor | None, + prefix: str = "", + ): + super().__init__() + self.vllm_config = vllm_config + self.config = config + self.quant_config = quant_config + # self.indexer_cfg = config.attn_module_list_cfg[0]["attn_index"] + # Indexer is only constructed for v32 configs, where these sparse-indexer + # fields are guaranteed populated; narrow away the `int | None` declared + # on Glm5NextConfig for the optional-indexer case. + assert config.index_topk is not None + assert config.index_n_heads is not None + assert config.index_head_dim is not None + assert config.index_kpool is not None + self.topk_tokens = config.index_topk + self.n_head = config.index_n_heads # 64 + self.head_dim = config.index_head_dim # 128 + self.rope_dim = config.qk_rope_head_dim # 64 + self.index_kpool = config.index_kpool + self.q_lora_rank = q_lora_rank # 1536 + + # kpool + self.index_kpool_compress_ape = nn.Parameter( + torch.zeros(self.index_kpool, self.head_dim, dtype=torch.float32) + ) + # Keep the checkpoint name ``index_kpool_compress_gate`` without a + # ``.weight`` suffix. F.linear consumes its [head_dim, hidden_size] shape. + self.index_kpool_compress_gate = nn.Parameter( + torch.empty(self.head_dim, hidden_size, dtype=torch.bfloat16) + ) + + # no tensor parallel, just replicated + self.wq_b = ReplicatedLinear( + self.q_lora_rank, + self.head_dim * self.n_head, + bias=False, + quant_config=quant_config, + prefix=f"{prefix}.wq_b", + ) + # Fused wk + weights_proj: single GEMM producing [head_dim + n_head]. + # FP8 wk weights are upcasted to BF16 during loading to maintain fusion. + self.wk_weights_proj = MergedColumnParallelLinear( + hidden_size, + [self.head_dim, self.n_head], + bias=False, + quant_config=None, + disable_tp=True, + prefix=f"{prefix}.wk_weights_proj", + ) + self.k_norm = LayerNorm(self.head_dim, eps=1e-6) + self.softmax_scale = self.head_dim**-0.5 + + # Hadamard-128 rotation of the indexer query is fused with the FP8 + # quant (see forward: fwht128_quant_fp8) -- no precomputed matrix. + + self.scale_fmt = "ue8m0" + self.quant_block_size = 128 # TODO: get from config + self.topk_indices_buffer = topk_indices_buffer + self._wp_fp32: torch.Tensor | None = None + + # NOTE: (zyongye) we use fp8 naive cache, + # where we store value in fp8 and scale in fp32 + # per self.quant_block_size element + self.k_cache = Glm5NextIndexerCache( + head_dim=self.head_dim + self.head_dim // self.quant_block_size * 4, + dtype=torch.uint8, + prefix=f"{prefix}.k_cache", + cache_config=cache_config, + index_kpool=self.index_kpool, + ) + # Paged tail cache (in-progress pool's raw K + gate score). Written by + # prefill (seeds the boundary pool) and decode (per-step stash); read by + # the decode kernel to compress the boundary pool. Transferred across PD + # so the decode side sees the prefill tail. See KpoolTailSpec/Manager. + self.tail_cache = Glm5NextTailCache( + head_dim=self.head_dim, + dtype=torch.bfloat16, + prefix=f"{prefix}.tail_cache", + cache_config=cache_config, + index_kpool=self.index_kpool, + ) + self.max_model_len = vllm_config.model_config.max_model_len + self.prefix = prefix + from vllm.v1.attention.backends.mla.indexer import get_max_prefill_buffer_size + + self.max_total_seq_len = get_max_prefill_buffer_size(vllm_config) + self.indexer_op = SparseAttnIndexerKpool( + self.k_cache, + self.quant_block_size, + self.scale_fmt, + self.topk_tokens, + self.head_dim, + self.max_model_len, + self.max_total_seq_len, + self.topk_indices_buffer, + tail_cache=self.tail_cache, + ) + + def forward( + self, hidden_states: torch.Tensor, qr: torch.Tensor, positions, rotary_emb + ) -> torch.Tensor: + q, _ = self.wq_b(qr) + q = q.view(-1, self.n_head, self.head_dim) + + # Compute the head gate in fp32; bf16 error can change near-tie pool + # rankings on long-context tasks. Cache it after weights are loaded. + kw, _ = self.wk_weights_proj(hidden_states) + k = kw[:, : self.head_dim] + if self._wp_fp32 is None: + self._wp_fp32 = ( + self.wk_weights_proj.weight.data[self.head_dim :, :] + .t() + .contiguous() + .float() + ) + weights = torch.mm(hidden_states.float(), self._wp_fp32) + + k = _fused_indexer_k_norm( + k, self.k_norm.weight, self.k_norm.bias, self.head_dim, self.k_norm.eps + ) + + if self.rope_dim > 0: + q_pe, q_nope = torch.split( + q, [self.rope_dim, self.head_dim - self.rope_dim], dim=-1 + ) + k_pe, k_nope = torch.split( + k, [self.rope_dim, self.head_dim - self.rope_dim], dim=-1 + ) + + q_pe, k_pe = rotary_emb(positions, q_pe, k_pe.unsqueeze(1)) + # Note: RoPE (NeoX) can introduce extra leading dimensions during + # compilation so we need to reshape back to token-flattened shapes + q_pe = q_pe.reshape(-1, self.n_head, self.rope_dim) + k_pe = k_pe.reshape(-1, 1, self.rope_dim) + + # `rotary_emb` is shape-preserving; `q_pe` is already + # [num_tokens, n_head, rope_dim]. + q = torch.cat([q_pe, q_nope], dim=-1) + # `k_pe` is [num_tokens, 1, rope_dim] (MQA). + k = torch.cat([k_pe.squeeze(-2), k_nope], dim=-1) + # else: qk_rope_head_dim=0 — no rope component. q is already + # [num_tokens, n_head, head_dim] and k is [num_tokens, head_dim] (all + # nope), so skip the rope split / rotary / cat entirely; otherwise the + # split/reshape would build 0-element tensors (breaks dynamo tracing). + + # Rotate Q into the cached K basis before computing fp8 MQA logits. + # Fusing the fp32 FWHT and quantization avoids an intermediate HBM + # round-trip and bf16 matrix-rounding bias. + assert self.head_dim == 128 and self.quant_block_size == 128 + assert self.scale_fmt == "ue8m0" + q = q.view(-1, self.head_dim) + q_fp8, q_scale = fwht128_quant_fp8(q) + q_fp8 = q_fp8.view(-1, self.n_head, self.head_dim) + q_scale = q_scale.view(-1, self.n_head, 1) + + weights = _fused_indexer_weight_scale( + weights, q_scale, self.softmax_scale * self.n_head**-0.5 + ) + + # kpool: per-token gate score driving the softmax-weighted pool. Computed + # from the same hidden_states that produced `k`, so it stays token-aligned. + # F.linear(x, gate) = x @ gate.T with gate [head_dim, hidden_size]. + gate_score = F.linear(hidden_states, self.index_kpool_compress_gate) + + # DeepGEMM's MQA-logits kernels (fp8_mqa_logits / + # fp8_fp4_paged_mqa_logits) require num_heads in {32, 64}; this + # checkpoint uses index_n_heads=16. Zero-pad q and the per-head + # weights: logits are a weights-weighted sum over heads, so + # zero-weight padded heads contribute exactly nothing. + if self.n_head < 32: + pad = 32 - self.n_head + q_fp8 = _pad_indexer_heads(q_fp8, pad) + weights = _pad_indexer_heads(weights, pad) + + return self.indexer_op( + hidden_states, + q_fp8, + k, + weights, + gate_score=gate_score, + compress_ape=self.index_kpool_compress_ape, + index_kpool=self.index_kpool, + positions=positions, + ) + + +class Glm5NextMLAAttention(nn.Module): + def __init__( + self, + vllm_config: VllmConfig, + config: Glm5NextConfig, + hidden_size: int, + num_heads: int, + qk_nope_head_dim: int, + qk_rope_head_dim: int, + v_head_dim: int, + q_lora_rank: int | None, + kv_lora_rank: int, + max_position_embeddings: int = 8192, + cache_config: CacheConfig | None = None, + quant_config: QuantizationConfig | None = None, + prefix: str = "", + topk_indices_buffer: torch.Tensor | None = None, + input_size: int | None = None, + skip_rope: bool | None = False, + ) -> None: + super().__init__() + self.hidden_size = hidden_size + self.qk_nope_head_dim = qk_nope_head_dim + self.qk_rope_head_dim = qk_rope_head_dim + self.qk_head_dim = qk_nope_head_dim + qk_rope_head_dim + self.v_head_dim = v_head_dim + + self.q_lora_rank = q_lora_rank + self.kv_lora_rank = kv_lora_rank + + self.num_heads = num_heads + tp_size = get_tensor_model_parallel_world_size() + assert num_heads % tp_size == 0 + self.num_local_heads = num_heads // tp_size + + self.scaling = self.qk_head_dim**-0.5 + self.max_position_embeddings = max_position_embeddings + + # Use input_size for projection input dimensions if provided, + # otherwise default to hidden_size (used in Eagle3 Deepseek with MLA) + proj_input_size = input_size if input_size is not None else self.hidden_size + + if self.q_lora_rank is not None: + self.fused_qkv_a_proj = DeepSeekV2FusedQkvAProjLinear( + proj_input_size, + [self.q_lora_rank, self.kv_lora_rank + self.qk_rope_head_dim], + quant_config=quant_config, + prefix=f"{prefix}.fused_qkv_a_proj", + ) + else: + self.kv_a_proj_with_mqa = ReplicatedLinear( + proj_input_size, + self.kv_lora_rank + self.qk_rope_head_dim, + bias=False, + quant_config=quant_config, + prefix=f"{prefix}.kv_a_proj_with_mqa", + ) + + if self.q_lora_rank is not None: + self.q_a_layernorm = RMSNorm(self.q_lora_rank, eps=config.rms_norm_eps) + self.q_b_proj = ColumnParallelLinear( + self.q_lora_rank, + self.num_heads * self.qk_head_dim, + bias=False, + quant_config=quant_config, + prefix=f"{prefix}.q_b_proj", + ) + else: + self.q_proj = ColumnParallelLinear( + proj_input_size, + self.num_heads * self.qk_head_dim, + bias=False, + quant_config=quant_config, + prefix=f"{prefix}.q_proj", + ) + self.kv_a_layernorm = RMSNorm(self.kv_lora_rank, eps=config.rms_norm_eps) + self.kv_b_proj = ColumnParallelLinear( + self.kv_lora_rank, + self.num_heads * (self.qk_nope_head_dim + self.v_head_dim), + bias=False, + quant_config=quant_config, + prefix=f"{prefix}.kv_b_proj", + ) + self.o_proj = RowParallelLinear( + self.num_heads * self.v_head_dim, + self.hidden_size, + bias=False, + quant_config=quant_config, + prefix=f"{prefix}.o_proj", + ) + + if not skip_rope: + assert config.rope_parameters is not None + if config.rope_parameters["rope_type"] != "default": + config.rope_parameters["rope_type"] = ( + "deepseek_yarn" + if config.rope_parameters.get("apply_yarn_scaling", True) + else "deepseek_llama_scaling" + ) + + self.rotary_emb: RotaryEmbedding | None = get_rope( + qk_rope_head_dim, + max_position=max_position_embeddings, + rope_parameters=config.rope_parameters, + is_neox_style=False, + ) + + if ( + config.rope_parameters["rope_type"] != "default" + and config.rope_parameters["rope_type"] == "deepseek_yarn" + ): + mscale_all_dim = config.rope_parameters.get("mscale_all_dim", False) + scaling_factor = config.rope_parameters["factor"] + mscale = yarn_get_mscale(scaling_factor, float(mscale_all_dim)) + self.scaling = self.scaling * mscale * mscale + else: + self.rotary_emb = None + + self.is_v32 = config.index_topk is not None + + if self.is_v32: + self.indexer_rope_emb: RotaryEmbedding | None = get_rope( + qk_rope_head_dim, + max_position=max_position_embeddings, + rope_parameters=config.rope_parameters, + is_neox_style=not config.indexer_rope_interleave, + ) + # The sparse indexer projects from the MLA q-lora rank, which is + # always set for v32 MLA configs; narrow away the `int | None`. + assert q_lora_rank is not None + self.indexer: Indexer | None = Indexer( + vllm_config, + config, + hidden_size, + q_lora_rank, + quant_config, + cache_config, + topk_indices_buffer, + f"{prefix}.indexer", + ) + + else: + self.indexer_rope_emb = None + self.indexer = None + + mla_modules = MLAModules( + kv_a_layernorm=self.kv_a_layernorm, + kv_b_proj=self.kv_b_proj, + rotary_emb=self.rotary_emb, + o_proj=self.o_proj, + fused_qkv_a_proj=self.fused_qkv_a_proj + if self.q_lora_rank is not None + else None, + kv_a_proj_with_mqa=self.kv_a_proj_with_mqa + if self.q_lora_rank is None + else None, + q_a_layernorm=self.q_a_layernorm if self.q_lora_rank is not None else None, + q_b_proj=self.q_b_proj if self.q_lora_rank is not None else None, + q_proj=self.q_proj if self.q_lora_rank is None else None, + indexer=self.indexer, + indexer_rotary_emb=self.indexer_rope_emb, + is_sparse=self.is_v32, + topk_indices_buffer=topk_indices_buffer, + ) + + self.mla_attn = MultiHeadLatentAttentionWrapper( + self.hidden_size, + self.num_local_heads, + self.scaling, + self.qk_nope_head_dim, + self.qk_rope_head_dim, + self.v_head_dim, + self.q_lora_rank, + self.kv_lora_rank, + mla_modules, + cache_config, + quant_config, + prefix, + skip_topk=False, + fuse_qkv_rmsnorm=True, + ) + + def forward( + self, hidden_states: torch.Tensor, positions: torch.Tensor + ) -> torch.Tensor: + # The wrapper also runs the sparse indexer before MLA attention. + return self.mla_attn(positions, hidden_states) diff --git a/vllm/models/glm5next/nvidia/kda.py b/vllm/models/glm5next/nvidia/kda.py new file mode 100644 index 000000000000..e6ba6ac5c659 --- /dev/null +++ b/vllm/models/glm5next/nvidia/kda.py @@ -0,0 +1,611 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +"""GLM-5.3-Flash KDA layer with separate convolutions and a bounded safe gate.""" + +import torch +from torch import nn + +from vllm.compilation.breakable_cudagraph import eager_break_during_capture +from vllm.config import VllmConfig, get_current_vllm_config +from vllm.distributed import divide +from vllm.forward_context import get_forward_context +from vllm.model_executor.layers.linear import ( + ColumnParallelLinear, + MergedColumnParallelLinear, + RowParallelLinear, +) +from vllm.model_executor.layers.mamba.gdn.base import GatedDeltaNetAttention +from vllm.model_executor.layers.mamba.mamba_utils import ( + MambaStateDtypeCalculator, + MambaStateShapeCalculator, + is_conv_state_dim_first, +) +from vllm.model_executor.layers.mamba.ops.causal_conv1d import ( + causal_conv1d_fn, + causal_conv1d_update, +) +from vllm.model_executor.layers.mamba.ops.gather_initial_states import ( + gather_initial_states, +) +from vllm.model_executor.layers.mamba.ops.scatter_states import scatter_states +from vllm.model_executor.model_loader.weight_utils import sharded_weight_loader +from vllm.model_executor.utils import ( + maybe_disable_graph_partition, + set_weight_attrs, +) +from vllm.platforms import current_platform +from vllm.third_party.flash_linear_attention.ops.kda import FusedRMSNormGated +from vllm.transformers_utils.configs.glm5_next import Glm5NextConfig +from vllm.v1.attention.backends.gdn_attn import GDNAttentionMetadata + +if current_platform.is_rocm(): + from vllm.models.glm5next.amd.ops.third_party.kda import ( + chunk_kda_with_fused_gate, + fused_recurrent_kda, + ) +else: + from vllm.models.glm5next.nvidia.ops.third_party.kda import ( + chunk_kda_with_fused_gate, + fused_recurrent_kda, + ) + + +class _Glm5NextMergedColumnParallelLinear(MergedColumnParallelLinear): + """Merged projection with multiple replicated output shards. + + Extends K3's ``_KimiGDNMergedColumnParallelLinear`` to support two + replicated shards (f_a, g_a) instead of one. Pre-multiplies each + replicated entry's output_size by tp_size so the per-rank shard + divides back to the full size, and forces tp_rank=0 during weight + loading for replicated shards. + """ + + def __init__( + self, + input_size: int, + output_sizes: list[int], + replicated_shard_ids: tuple[int, ...], + tp_size: int, + **kwargs, + ) -> None: + self.replicated_shard_ids = set(replicated_shard_ids) + output_sizes = output_sizes.copy() + for sid in self.replicated_shard_ids: + output_sizes[sid] *= tp_size + super().__init__(input_size, output_sizes, **kwargs) + + def weight_loader( + self, + param: nn.Parameter, + loaded_weight: torch.Tensor, + loaded_shard_id: tuple[int, ...] | int | None = None, + ) -> None: + tp_rank = self.tp_rank + param_tp_rank = getattr(param, "tp_rank", None) + if loaded_shard_id in self.replicated_shard_ids: + self.tp_rank = 0 + if param_tp_rank is not None: + param.tp_rank = 0 + try: + super().weight_loader(param, loaded_weight, loaded_shard_id) + finally: + self.tp_rank = tp_rank + if param_tp_rank is not None: + param.tp_rank = param_tp_rank + + def weight_loader_v2( + self, + param: nn.Parameter, + loaded_weight: torch.Tensor, + loaded_shard_id: tuple[int, ...] | int | None = None, + ) -> None: + tp_rank = self.tp_rank + param_tp_rank = getattr(param, "tp_rank", None) + if loaded_shard_id in self.replicated_shard_ids: + self.tp_rank = 0 + if param_tp_rank is not None: + param.tp_rank = 0 + try: + super().weight_loader_v2(param, loaded_weight, loaded_shard_id) + finally: + self.tp_rank = tp_rank + if param_tp_rank is not None: + param.tp_rank = param_tp_rank + + +@torch.compile( + dynamic=True, + backend=current_platform.simple_compile_backend, + options=maybe_disable_graph_partition(current_platform.simple_compile_backend), +) +def _cast_sigmoid(x: torch.Tensor) -> torch.Tensor: + """Fuse the fp32 cast + sigmoid into one Inductor kernel.""" + return x.float().sigmoid() + + +class Glm5NextLinearAttention(GatedDeltaNetAttention): + head_dim: int + num_heads: int + conv_size: int + + def get_state_dtype( + self, + ) -> tuple[torch.dtype, torch.dtype]: + if self.model_config is None or self.cache_config is None: + raise ValueError("model_config and cache_config must be set") + return MambaStateDtypeCalculator.kda_state_dtype( + self.model_config.dtype, self.cache_config.mamba_cache_dtype + ) + + def get_state_shape( + self, + ) -> tuple[tuple[int, ...], tuple[int, ...]]: + # conv_state width must include num_spec so the spec-decode conv update + # (causal_conv1d_update with num_accepted_tokens + max_query_len) can + # slide the window across the draft-verify tokens without reading past + # the allocated width. Matches qwen_gdn_linear_attn.get_state_shape. + return MambaStateShapeCalculator.kda_state_shape( + self.tp_size, + self.num_heads, + self.head_dim, + conv_kernel_size=self.conv_size, + num_spec=self.num_spec, + ) + + def __init__( + self, + config: Glm5NextConfig, + vllm_config: VllmConfig, + prefix: str = "", + ) -> None: + # KDA projections remain BF16 because fp8 checkpoints omit their scales. + saved_quant_config = vllm_config.quant_config + try: + vllm_config.quant_config = None + super().__init__(config, vllm_config, prefix) + finally: + vllm_config.quant_config = saved_quant_config + + self.head_dim = config.linear_head_dim + self.num_heads = config.linear_num_heads + self.conv_size = config.linear_conv_kernel_dim + assert self.num_heads % self.tp_size == 0 + self.local_num_heads = divide(self.num_heads, self.tp_size) + + projection_size = self.head_dim * self.num_heads + self.local_projection_size = divide(projection_size, self.tp_size) + + # Merge q, k, v, b, f_a, g_a projections into one GEMM (6→1 launches). + # Order matches checkpoint's fused_qkvbfg_a_proj convention. + # Shards 4 (f_a) and 5 (g_a) are replicated across TP ranks. + self.in_proj_qkvbfg_a = _Glm5NextMergedColumnParallelLinear( + self.hidden_size, + [ + projection_size, # q (shard 0) + projection_size, # k (shard 1) + projection_size, # v (shard 2) + self.num_heads, # b (shard 3) + self.head_dim, # f_a (shard 4, replicated) + self.head_dim, # g_a (shard 5, replicated) + ], + replicated_shard_ids=(4, 5), + tp_size=self.tp_size, + bias=False, + quant_config=self.quant_config, + prefix=f"{prefix}.in_proj_qkvbfg_a", + ) + + self.f_b_proj = ColumnParallelLinear( + self.head_dim, + projection_size, + bias=False, + quant_config=self.quant_config, + prefix=f"{prefix}.f_b_proj", + ) + self.dt_bias = nn.Parameter( + torch.empty(divide(projection_size, self.tp_size), dtype=torch.float32) + ) + + set_weight_attrs(self.dt_bias, {"weight_loader": sharded_weight_loader(0)}) + + self.q_conv1d = ColumnParallelLinear( + input_size=self.conv_size, + output_size=projection_size, + bias=False, + params_dtype=torch.float32, + prefix=f"{prefix}.q_conv1d", + ) + self.k_conv1d = ColumnParallelLinear( + input_size=self.conv_size, + output_size=projection_size, + bias=False, + params_dtype=torch.float32, + prefix=f"{prefix}.k_conv1d", + ) + self.v_conv1d = ColumnParallelLinear( + input_size=self.conv_size, + output_size=projection_size, + bias=False, + params_dtype=torch.float32, + prefix=f"{prefix}.v_conv1d", + ) + # unsqueeze to fit conv1d weights shape into the linear weights shape. + # Can't do this in `weight_loader` since it already exists in + # `ColumnParallelLinear` and `set_weight_attrs` + # doesn't allow to override it + self.q_conv1d.weight.data = self.q_conv1d.weight.data.unsqueeze(1) + self.k_conv1d.weight.data = self.k_conv1d.weight.data.unsqueeze(1) + self.v_conv1d.weight.data = self.v_conv1d.weight.data.unsqueeze(1) + # Lazily-built merged q|k|v conv weight (built on first forward, after + # weights are loaded). See _forward. + self._merged_conv_weight: torch.Tensor | None = None + + self.A_log = nn.Parameter( + torch.empty(1, 1, self.local_num_heads, 1, dtype=torch.float32) + ) + set_weight_attrs(self.A_log, {"weight_loader": sharded_weight_loader(2)}) + + self.g_b_proj = ColumnParallelLinear( + self.head_dim, + projection_size, + bias=False, + quant_config=self.quant_config, + prefix=f"{prefix}.g_b_proj", + ) + self.o_norm = FusedRMSNormGated(self.head_dim, activation="sigmoid") + self.o_proj = RowParallelLinear( + projection_size, + self.hidden_size, + bias=False, + quant_config=self.quant_config, + prefix=f"{prefix}.o_proj", + ) + + 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 + + # Checkpoints store A_log as 1-D; the model parameter is 4-D. + def _a_log_weight_loader(param, loaded_weight): + if loaded_weight.dim() == 1: + loaded_weight = loaded_weight.view([1, 1, -1, 1]) + return sharded_weight_loader(2)(param, loaded_weight) + + self.A_log.weight_loader = _a_log_weight_loader + + # GLM-5.3-Flash uses a bounded sigmoid gate instead of the default + # unbounded softplus gate. + self.kda_safe_gate = True + self.kda_lower_bound = config.linear_lower_bound + # Process-global conv-state layout, resolved once here instead of on + # every _forward call (it reads an env-derived flag each time). + self._conv_state_dim_first = is_conv_state_dim_first() + + def forward( + self, + hidden_states: torch.Tensor, + positions: torch.Tensor, + ) -> torch.Tensor: + num_tokens = hidden_states.size(0) + # One merged GEMM for q, k, v, b, f_a, g_a (replaces 6 separate GEMMs). + projected = self.in_proj_qkvbfg_a(hidden_states)[0] + qkv, beta_raw, f_a, g_a = projected.split( + [ + 3 * self.local_projection_size, + self.local_num_heads, + self.head_dim, + self.head_dim, + ], + dim=-1, + ) + + # Beta stays raw (bf16) here: the recurrent kernel sigmoids it in fp32 + # at load (SIGMOID_BETA), and only the chunked prefill path needs the + # pre-computed fp32 sigmoid — computed lazily in _forward. Pure decode + # / spec-verify steps then skip the _cast_sigmoid kernel and its fp32 + # intermediate entirely. + beta = beta_raw.unsqueeze(0) + g1 = self.f_b_proj(f_a)[0] + g1 = g1.reshape(1, -1, self.local_num_heads, self.head_dim) + + g_proj_states = self.g_b_proj(g_a)[0] + # Must stay 3D: rms_norm_gated reads H from g.shape[-2]. + g2 = g_proj_states.reshape(-1, self.local_num_heads, self.head_dim) + + core_attn_out = torch.empty( + (1, num_tokens, self.local_num_heads, self.head_dim), + dtype=hidden_states.dtype, + device=hidden_states.device, + ) + # Call the decorated eager break directly so host-side prefill branches + # are not captured by PIECEWISE CUDA graphs. + self._forward( + qkv_proj_states=qkv, + g1=g1, + beta=beta, + core_attn_out=core_attn_out, + ) + core_attn_out = self.o_norm(core_attn_out, g2) + core_attn_out = core_attn_out.reshape(core_attn_out.size(1), -1) + return self.o_proj(core_attn_out)[0] + + @eager_break_during_capture + def _forward( + self, + qkv_proj_states: torch.Tensor, + g1: torch.Tensor, + beta: torch.Tensor, + core_attn_out: torch.Tensor, + ) -> None: + forward_context = get_forward_context() + attn_metadata_raw = forward_context.attn_metadata + + if attn_metadata_raw is None: + return + + assert isinstance(attn_metadata_raw, dict) + attn_metadata_narrowed = attn_metadata_raw.get(self.prefix) + if attn_metadata_narrowed is None: + # Profile/warmup dummy runs may omit mamba-family metadata. + return + assert isinstance(attn_metadata_narrowed, GDNAttentionMetadata) + has_initial_state = attn_metadata_narrowed.has_initial_state + non_spec_query_start_loc = attn_metadata_narrowed.non_spec_query_start_loc + non_spec_state_indices_tensor = ( + attn_metadata_narrowed.non_spec_state_indices_tensor + ) # noqa: E501 + num_actual_tokens = attn_metadata_narrowed.num_actual_tokens + # Spec-decode metadata (all None when speculative decoding is disabled). + spec_sequence_masks = attn_metadata_narrowed.spec_sequence_masks + spec_query_start_loc = attn_metadata_narrowed.spec_query_start_loc + spec_state_indices_tensor = attn_metadata_narrowed.spec_state_indices_tensor + spec_token_indx = attn_metadata_narrowed.spec_token_indx + non_spec_token_indx = attn_metadata_narrowed.non_spec_token_indx + num_accepted_tokens = attn_metadata_narrowed.num_accepted_tokens + num_spec_decodes = attn_metadata_narrowed.num_spec_decodes + use_spec = spec_sequence_masks is not None and num_spec_decodes > 0 + # Safe-gate checkpoints use the bounded sigmoid variant. + safe_gate = self.kda_safe_gate + lower_bound = self.kda_lower_bound + constant_caches = self.kv_cache + + qkv_proj_states = qkv_proj_states[:num_actual_tokens] + g1 = g1[:, :num_actual_tokens] + beta = beta[:, :num_actual_tokens] + + (conv_state, recurrent_state) = constant_caches + # conv_state must be (..., dim, width-1) for the conv kernels. + # DS layout stores it that way directly; SD layout needs a transpose. + # Layout is process-global and resolved once at init (see __init__). + if not self._conv_state_dim_first: + conv_state = conv_state.transpose(-1, -2) + + # One merged short-conv over q|k|v instead of three separate calls. The + # 1D conv is independent per channel, so concatenating q/k/v along the + # channel dim and running a single causal_conv1d is bit-identical to + # three calls. The merged weight is q|k|v conv weights concatenated; + # built once and cached (params are fixed after load). conv_state is + # already stored as the merged q|k|v state, so it is used directly. + if self._merged_conv_weight is None: + + def _w(m): + return m.weight.view(m.weight.size(0), m.weight.size(2)) + + self._merged_conv_weight = torch.cat( + [_w(self.q_conv1d), _w(self.k_conv1d), _w(self.v_conv1d)], + dim=0, + ).contiguous() + conv_weights = self._merged_conv_weight + conv_bias = self.q_conv1d.bias + + # Split projections / gating into spec (draft-verify) and non-spec token + # groups when speculative decoding is active. Spec tokens carry + # num_spec+1 recurrent-state columns each and are advanced with + # num_accepted_tokens for rejection-sampling rollback; non-spec tokens + # are one-per-request. Mirrors olmo_gdn_linear_attn.py. Projections are + # [n, *] (token dim 0); g1/beta are [1, n, h, d] (token dim 1). + if use_spec: + # In a pure spec-verify step (no non-spec tokens) the metadata + # builder sets spec_token_indx = arange(num_actual_tokens), making + # the index_select calls below identity copies. Skip them on this + # steady-state decode hot path. The outputs alias the inputs here; + # the downstream conv/recurrent kernels read them without mutating + # in place, so the aliasing is safe. + if non_spec_token_indx is None or non_spec_token_indx.numel() == 0: + qkv_spec = qkv_proj_states + g1_spec = g1 + beta_spec = beta + else: + qkv_spec = qkv_proj_states.index_select(0, spec_token_indx) + g1_spec = g1.index_select(1, spec_token_indx) + beta_spec = beta.index_select(1, spec_token_indx) + if non_spec_token_indx is not None and non_spec_token_indx.numel() > 0: + qkv_ns = qkv_proj_states.index_select(0, non_spec_token_indx) + g1_ns = g1.index_select(1, non_spec_token_indx) + beta_ns = beta.index_select(1, non_spec_token_indx) + else: + qkv_ns = g1_ns = beta_ns = None + else: + qkv_spec = g1_spec = beta_spec = None + qkv_ns, g1_ns, beta_ns = qkv_proj_states, g1, beta + + # --- causal conv1d: spec (draft-verify) path --- + if use_spec: + assert spec_state_indices_tensor is not None + assert num_accepted_tokens is not None + conv_idx = spec_state_indices_tensor[:, 0][:num_spec_decodes] + conv_mql = spec_state_indices_tensor.size(-1) + qkv_spec = causal_conv1d_update( + qkv_spec, + conv_state, + conv_weights, + conv_bias, + activation="silu", + conv_state_indices=conv_idx, + num_accepted_tokens=num_accepted_tokens, + query_start_loc=spec_query_start_loc, + max_query_len=conv_mql, + ) + q_spec, k_spec, v_spec = qkv_spec.split(self.local_projection_size, dim=-1) + + # --- causal conv1d: non-spec path (prefill or plain decode) --- + q_ns = k_ns = v_ns = None + if attn_metadata_narrowed.num_prefills > 0: + assert qkv_ns is not None + qkv_ns = causal_conv1d_fn( + qkv_ns.transpose(0, 1), + conv_weights, + conv_bias, + activation="silu", + conv_states=conv_state, + has_initial_state=has_initial_state, + cache_indices=non_spec_state_indices_tensor, + query_start_loc=non_spec_query_start_loc, + metadata=attn_metadata_narrowed, + ).transpose(0, 1) + q_ns, k_ns, v_ns = qkv_ns.split(self.local_projection_size, dim=-1) + elif attn_metadata_narrowed.num_decodes > 0: + assert non_spec_state_indices_tensor is not None + decode_conv_indices = non_spec_state_indices_tensor[ + : attn_metadata_narrowed.num_decodes + ] + qkv_ns = causal_conv1d_update( + qkv_ns, + conv_state, + conv_weights, + conv_bias, + activation="silu", + conv_state_indices=decode_conv_indices, + ) + q_ns, k_ns, v_ns = qkv_ns.split(self.local_projection_size, dim=-1) + + def _rearr(x): + return x.reshape(1, -1, self.local_num_heads, self.head_dim) + + # --- core attention: spec (draft-verify) path --- + core_attn_out_spec = None + # In a pure spec-verify step (no non-spec tokens) the recurrent kernel + # can write straight into the layer output buffer, skipping the + # fresh allocation + copy below. Mixed steps must scatter via + # spec_token_indx, so they keep the kernel-managed output. + spec_out = ( + core_attn_out[0, :num_actual_tokens].unsqueeze(0) + if non_spec_token_indx is None or non_spec_token_indx.numel() == 0 + else None + ) + if use_spec: + assert spec_state_indices_tensor is not None + assert num_accepted_tokens is not None + assert spec_query_start_loc is not None + # Gate computed inside the recurrent kernel (COMPUTE_GATE) from + # raw g1 — replicates fused_kda_gate's arithmetic bit-for-bit and + # skips its launch + fp32 [n, H, D] intermediate per layer. + core_attn_out_spec, _ = fused_recurrent_kda( + q=_rearr(q_spec), + k=_rearr(k_spec), + v=_rearr(v_spec), + g=g1_spec, + beta=beta_spec, + initial_state=recurrent_state, + use_qk_l2norm_in_kernel=True, + cu_seqlens=spec_query_start_loc[: num_spec_decodes + 1], + ssm_state_indices=spec_state_indices_tensor, + num_accepted_tokens=num_accepted_tokens, + out=spec_out, + sigmoid_beta=True, + a_log=self.A_log, + g_bias=self.dt_bias, + compute_gate=True, + lower_bound=lower_bound, + ) + + # --- core attention: non-spec path (prefill or plain decode) --- + core_attn_out_non_spec = None + # Only the plain-decode recurrent kernel can write straight into the + # layer output buffer; the chunked prefill kernel cannot, so this + # stays None there and the merge copy below runs as before. + ns_out = None + if attn_metadata_narrowed.num_prefills > 0: + assert q_ns is not None + assert non_spec_state_indices_tensor is not None + assert has_initial_state is not None + initial_state = gather_initial_states( + recurrent_state, non_spec_state_indices_tensor, has_initial_state + ) + ( + core_attn_out_non_spec, + last_recurrent_state, + ) = chunk_kda_with_fused_gate( + q=_rearr(q_ns), + k=_rearr(k_ns), + v=_rearr(v_ns), + raw_g=g1_ns, + # Chunk path wants the pre-sigmoided fp32 beta (its kernels + # don't sigmoid); beta_ns is raw bf16 from forward. + beta=_cast_sigmoid(beta_ns.squeeze(0)).unsqueeze(0), + A_log=self.A_log, + g_bias=self.dt_bias, + initial_state=initial_state, + output_final_state=True, + use_qk_l2norm_in_kernel=True, + cu_seqlens=non_spec_query_start_loc, + safe_gate=safe_gate, + lower_bound=lower_bound, + ) + # Init cache + scatter_states( + recurrent_state, + last_recurrent_state, + non_spec_state_indices_tensor, + ) + elif attn_metadata_narrowed.num_decodes > 0: + assert non_spec_query_start_loc is not None + assert non_spec_state_indices_tensor is not None + # Plain decode step (no spec tokens): token order is dense, so the + # kernel can write straight into the layer output buffer. A mixed + # step scatters non-spec output via non_spec_token_indx instead. + # Gate computed in-kernel (COMPUTE_GATE), beta sigmoided in-kernel. + if not use_spec: + ns_out = spec_out + core_attn_out_non_spec, _ = fused_recurrent_kda( + q=_rearr(q_ns), + k=_rearr(k_ns), + v=_rearr(v_ns), + g=g1_ns, + beta=beta_ns, + initial_state=recurrent_state, + use_qk_l2norm_in_kernel=True, + cu_seqlens=non_spec_query_start_loc[ + : attn_metadata_narrowed.num_decodes + 1 + ], + ssm_state_indices=non_spec_state_indices_tensor, + out=ns_out, + sigmoid_beta=True, + a_log=self.A_log, + g_bias=self.dt_bias, + compute_gate=True, + lower_bound=lower_bound, + ) + + # --- merge spec / non-spec outputs back into token order --- + if use_spec and core_attn_out_non_spec is not None: + assert core_attn_out_spec is not None + merged = torch.empty( + (1, num_actual_tokens, *core_attn_out_spec.shape[2:]), + dtype=core_attn_out_non_spec.dtype, + device=core_attn_out_non_spec.device, + ) + merged.index_copy_(1, spec_token_indx, core_attn_out_spec) + merged.index_copy_(1, non_spec_token_indx, core_attn_out_non_spec) + core_attn_out[0, :num_actual_tokens] = merged.squeeze(0) + elif use_spec: + assert core_attn_out_spec is not None + if spec_out is None: + core_attn_out[0, :num_actual_tokens] = core_attn_out_spec.squeeze(0) + else: + assert core_attn_out_non_spec is not None + if ns_out is None: + core_attn_out[0, :num_actual_tokens] = core_attn_out_non_spec[ + 0, :num_actual_tokens + ] diff --git a/vllm/models/glm5next/nvidia/model.py b/vllm/models/glm5next/nvidia/model.py new file mode 100644 index 000000000000..9988dd84ea74 --- /dev/null +++ b/vllm/models/glm5next/nvidia/model.py @@ -0,0 +1,1196 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project + +from collections.abc import Iterable +from typing import ClassVar, Literal + +import torch +from torch import nn + +from vllm.config import ParallelConfig, VllmConfig +from vllm.distributed import ( + get_ep_group, + get_pp_group, + get_tensor_model_parallel_rank, + get_tensor_model_parallel_world_size, + tensor_model_parallel_all_gather, +) +from vllm.logger import init_logger +from vllm.model_executor.layers.activation import SiluAndMul, SiluAndMulWithClamp +from vllm.model_executor.layers.fused_moe import ( + FusedMoEFactory, + GateLinear, + fused_moe_make_expert_params_mapping, +) +from vllm.model_executor.layers.layernorm import RMSNorm +from vllm.model_executor.layers.linear import ( + MergedColumnParallelLinear, + RowParallelLinear, +) +from vllm.model_executor.layers.logits_processor import LogitsProcessor +from vllm.model_executor.layers.mamba.mamba_utils import ( + MambaStateCopyFunc, + MambaStateCopyFuncCalculator, + MambaStateDtypeCalculator, + MambaStateShapeCalculator, +) +from vllm.model_executor.layers.mhc import ( + MHCFusedPostPreOp, + MHCPostOp, + MHCPreOp, + hc_contract, + hc_expand, +) +from vllm.model_executor.layers.quantization import QuantizationConfig +from vllm.model_executor.layers.quantization.utils.quant_utils import ( + GroupShape, + scaled_dequantize, +) +from vllm.model_executor.layers.vocab_parallel_embedding import ( + ParallelLMHead, + VocabParallelEmbedding, +) +from vllm.model_executor.model_loader.weight_utils import ( + default_weight_loader, + maybe_remap_kv_scale_name, +) +from vllm.model_executor.models.deepseek_v2 import _get_moe_router_dtype +from vllm.model_executor.models.glm4_1v import ( + Glm4vDummyInputsBuilder, + Glm4vForConditionalGeneration, +) +from vllm.model_executor.models.interfaces import ( + HasInnerState, + IsHybrid, + MixtureOfExperts, + SupportsPP, +) +from vllm.model_executor.models.utils import ( + AutoWeightsLoader, + PPMissingLayer, + init_vllm_registered_model, + is_pp_missing_parameter, + make_layers, + maybe_prefix, + sequence_parallel_chunk, +) +from vllm.models.common.ops.sequence_parallel import ( + sp_all_gather, + sp_reduce_scatter, + sp_shard, +) +from vllm.multimodal import MULTIMODAL_REGISTRY +from vllm.platforms import current_platform +from vllm.sequence import IntermediateTensors +from vllm.transformers_utils.configs.glm5_next import Glm5NextConfig + +from .attention import Glm5NextMLAAttention +from .kda import Glm5NextLinearAttention +from .multimodal import ( + Glm5NextMultiModalProcessor, + Glm5NextProcessingInfo, + Glm5NextVisionTransformer, +) + +logger = init_logger(__name__) + + +class Glm5NextMLP(nn.Module): + def __init__( + self, + hidden_size: int, + intermediate_size: int, + hidden_act: str, + quant_config: QuantizationConfig | None = None, + reduce_results: bool = True, + is_sequence_parallel=False, + prefix: str = "", + swiglu_limit: float | None = None, + ) -> None: + super().__init__() + + # If is_sequence_parallel, the input and output tensors are sharded + # across the ranks within the tp_group. In this case the weights are + # replicated and no collective ops are needed. + # Otherwise we use standard TP with an allreduce at the end. + self.gate_up_proj = MergedColumnParallelLinear( + hidden_size, + [intermediate_size] * 2, + bias=False, + quant_config=quant_config, + disable_tp=is_sequence_parallel, + prefix=f"{prefix}.gate_up_proj", + ) + self.down_proj = RowParallelLinear( + intermediate_size, + hidden_size, + bias=False, + quant_config=quant_config, + reduce_results=reduce_results, + disable_tp=is_sequence_parallel, + prefix=f"{prefix}.down_proj", + ) + if hidden_act != "silu": + raise ValueError( + f"Unsupported activation: {hidden_act}. Only silu is supported for now." + ) + + self.swiglu_limit = swiglu_limit + if self.swiglu_limit is not None: + self.act_fn = SiluAndMulWithClamp(swiglu_limit=self.swiglu_limit) + else: + self.act_fn = SiluAndMul() + + def forward(self, x): + gate_up, _ = self.gate_up_proj(x) + x = self.act_fn(gate_up) + x, _ = self.down_proj(x) + return x + + +class Glm5NextMoE(nn.Module): + def __init__( + self, + config: Glm5NextConfig, + parallel_config: ParallelConfig, + quant_config: QuantizationConfig | None = None, + prefix: str = "", + apply_routed_scale_to_output: bool = False, + ): + super().__init__() + self.tp_size = get_tensor_model_parallel_world_size() + self.tp_rank = get_tensor_model_parallel_rank() + + self.routed_scaling_factor = config.routed_scaling_factor + + self.ep_group = get_ep_group().device_group + self.ep_rank = get_ep_group().rank_in_group + self.ep_size = self.ep_group.size() + self.n_routed_experts: int = config.n_routed_experts + self.n_shared_experts: int = config.n_shared_experts + + self.is_sequence_parallel = parallel_config.use_sequence_parallel_moe + + if config.hidden_act != "silu": + raise ValueError( + f"Unsupported activation: {config.hidden_act}. " + "Only silu is supported for now." + ) + + self.router_dtype = _get_moe_router_dtype(config) + self.gate = GateLinear( + config.hidden_size, + config.n_routed_experts, + out_dtype=self.router_dtype, + prefix=f"{prefix}.gate", + ) + if config.topk_method == "noaux_tc": + self.gate.e_score_correction_bias = nn.Parameter( + torch.empty(config.n_routed_experts, dtype=torch.float32) + ) + else: + self.gate.e_score_correction_bias = None + + # Load balancing settings. + eplb_config = parallel_config.eplb_config + self.enable_eplb = parallel_config.enable_eplb + + self.n_redundant_experts = eplb_config.num_redundant_experts + self.n_logical_experts = self.n_routed_experts + self.n_physical_experts = self.n_logical_experts + self.n_redundant_experts + self.n_local_physical_experts = self.n_physical_experts // self.ep_size + + self.physical_expert_start = self.ep_rank * self.n_local_physical_experts + self.physical_expert_end = ( + self.physical_expert_start + self.n_local_physical_experts + ) + + swiglu_limit = config.swiglu_limit + if config.n_shared_experts is None: + self.shared_experts = None + else: + intermediate_size = config.moe_intermediate_size * config.n_shared_experts + + self.shared_experts = Glm5NextMLP( + hidden_size=config.hidden_size, + intermediate_size=intermediate_size, + hidden_act=config.hidden_act, + quant_config=quant_config, + is_sequence_parallel=self.is_sequence_parallel, + reduce_results=False, + prefix=f"{prefix}.shared_experts", + swiglu_limit=swiglu_limit, + ) + + self.experts = FusedMoEFactory( + shared_experts=self.shared_experts, + gate=self.gate, + num_experts=config.n_routed_experts, + top_k=config.num_experts_per_token, + hidden_size=config.hidden_size, + intermediate_size=config.moe_intermediate_size, + renormalize=config.moe_renormalize, + quant_config=quant_config, + use_grouped_topk=True, + num_expert_group=config.n_group, + topk_group=config.topk_group, + prefix=f"{prefix}.experts", + scoring_func=config.scoring_func, + routed_scaling_factor=self.routed_scaling_factor, + apply_routed_scale_to_output=apply_routed_scale_to_output, + e_score_correction_bias=self.gate.e_score_correction_bias, + enable_eplb=self.enable_eplb, + num_redundant_experts=self.n_redundant_experts, + is_sequence_parallel=self.is_sequence_parallel, + n_shared_experts=None, + router_logits_dtype=self.gate.out_dtype, + swiglu_limit=swiglu_limit, + ) + + def forward( + self, + hidden_states: torch.Tensor, + already_sequence_parallel: bool = False, + ) -> torch.Tensor: + num_tokens, hidden_dim = hidden_states.shape + + # Chunk the hidden states so they aren't replicated across TP ranks. + # This avoids duplicate computation in self.experts. + if self.is_sequence_parallel and not already_sequence_parallel: + hidden_states = sequence_parallel_chunk(hidden_states) + + # The router is always external (self.gate); main's MoERunner expects + # pre-computed router_logits, so compute them here unconditionally. + router_logits, _ = self.gate(hidden_states) + final_hidden_states = self.experts( + hidden_states=hidden_states, router_logits=router_logits + ) + + if self.is_sequence_parallel and not already_sequence_parallel: + final_hidden_states = tensor_model_parallel_all_gather( + final_hidden_states, 0 + ) + final_hidden_states = final_hidden_states[:num_tokens] + + return final_hidden_states.view(num_tokens, hidden_dim) + + +class Glm5NextDecoderLayer(nn.Module): + def __init__( + self, + vllm_config: VllmConfig, + config: Glm5NextConfig, + layer_idx: int, + prefix: str = "", + topk_indices_buffer: torch.Tensor | None = None, + is_mtp_layer: bool = False, + **kwargs, + ) -> None: + super().__init__() + + cache_config = vllm_config.cache_config + quant_config = vllm_config.quant_config + parallel_config = vllm_config.parallel_config + + self.hidden_size = config.hidden_size + self.layer_idx = layer_idx + self.is_moe = config.is_moe + self.num_hidden_layers = config.num_hidden_layers + self.rms_norm_eps = config.rms_norm_eps + self.num_experts = config.n_routed_experts + self.is_mtp_layer = is_mtp_layer + self.mhc = config.mhc + is_kda_layer = not is_mtp_layer and config.is_kda_layer(layer_idx) + self.layer_kind = "kda" if is_kda_layer else "mla" + self.is_sequence_parallel = parallel_config.use_sequence_parallel_moe + + if is_kda_layer: + self.self_attn = Glm5NextLinearAttention( + config=config, + vllm_config=vllm_config, + prefix=f"{prefix}.self_attn", + ) + else: + # MLA layers require the latent head dims, which are guaranteed set + # on MLA configs; narrow away the `int | None`. + assert config.v_head_dim is not None + assert config.kv_lora_rank is not None + self.self_attn = Glm5NextMLAAttention( + vllm_config=vllm_config, + config=config, + hidden_size=self.hidden_size, + num_heads=config.num_attention_heads, + qk_nope_head_dim=config.qk_nope_head_dim, + qk_rope_head_dim=config.qk_rope_head_dim, + v_head_dim=config.v_head_dim, + q_lora_rank=config.q_lora_rank, + kv_lora_rank=config.kv_lora_rank, + max_position_embeddings=config.max_position_embeddings, + cache_config=cache_config, + quant_config=None, # MLA projections are BF16 in checkpoint + prefix=f"{prefix}.self_attn", + topk_indices_buffer=topk_indices_buffer, + skip_rope=config.mla_nope, + ) + + # MTP layers sit past the base model's hidden layers (layer_idx >= + # num_hidden_layers), so they're outside mlp_layer_types; default them + # to the last base layer's MLP type (sparse/MoE for these checkpoints). + mlp_layer_types = config.mlp_layer_types + mlp_type = ( + mlp_layer_types[layer_idx] + if layer_idx < len(mlp_layer_types) + else (mlp_layer_types[-1] if mlp_layer_types else "sparse") + ) + if self.is_moe and self.num_experts is not None and mlp_type == "sparse": + self.mlp = Glm5NextMoE( + config=config, + parallel_config=parallel_config, + quant_config=quant_config, + prefix=f"{prefix}.mlp", + ) + else: + self.mlp = Glm5NextMLP( + hidden_size=self.hidden_size, + intermediate_size=config.intermediate_size, + hidden_act=config.hidden_act, + quant_config=quant_config, + prefix=f"{prefix}.mlp", + swiglu_limit=config.swiglu_limit, + ) + self.input_layernorm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps) + # Cached for the hot forward path (isinstance per layer per step). + self._mlp_is_moe = isinstance(self.mlp, Glm5NextMoE) + # In SP, the attention output projection leaves a partial sum; the + # decoder-layer reduce_scatter after attention completes it (DSv4 pattern). + # MTP layers use the non-mHC path which has no sp_reduce_scatter, so + # their o_proj must still reduce normally. + if self.is_sequence_parallel and not is_mtp_layer: + self.self_attn.o_proj.reduce_results = False + self.post_attention_layernorm = RMSNorm( + config.hidden_size, eps=config.rms_norm_eps + ) + + if self.mhc and not is_mtp_layer: + # mhc config + self.mhc_num_residual_streams = config.mhc_num_residual_streams + self.mhc_tau = config.mhc_tau + self.hc_eps = config.hc_eps + self.mhc_sinkhorn_iterations = config.mhc_sinkhorn_iterations + self.mhc_post_mult_value = config.mhc_post_mult_value + + n = config.mhc_num_residual_streams + d_model = n * self.hidden_size + mix_hc = (2 + n) * n + + self.n = n + + # attn hc + self.hc_attn_fn = nn.Parameter( + torch.empty(mix_hc, d_model, dtype=torch.float32) + ) + self.hc_attn_base = nn.Parameter(torch.empty(mix_hc, dtype=torch.float32)) + self.hc_attn_scale = nn.Parameter(torch.empty(3, dtype=torch.float32)) + + # ffn hc + self.hc_ffn_fn = nn.Parameter( + torch.empty(mix_hc, d_model, dtype=torch.float32) + ) + self.hc_ffn_base = nn.Parameter(torch.empty(mix_hc, dtype=torch.float32)) + self.hc_ffn_scale = nn.Parameter(torch.empty(3, dtype=torch.float32)) + + self.mhc_pre_op = MHCPreOp() + self.mhc_post_op = MHCPostOp() + self.mhc_fused_post_pre_op = MHCFusedPostPreOp() + + def forward( + self, + positions: torch.Tensor, + hidden_states: torch.Tensor, + residual: torch.Tensor | None = None, + post: torch.Tensor | None = None, + comb: torch.Tensor | None = None, + ) -> tuple[ + torch.Tensor, + torch.Tensor | None, + torch.Tensor | None, + torch.Tensor | None, + ]: + # 70B or MTP layers: KDA + MoE without HC. + if not self.mhc or self.is_mtp_layer: + residual = hidden_states + hidden_states = self.input_layernorm(hidden_states) + + attn_output = self.self_attn( + hidden_states=hidden_states, + positions=positions, + ) + hidden_states, residual = self.post_attention_layernorm( + attn_output, residual=residual + ) + hidden_states = self.mlp(hidden_states) + if self.is_mtp_layer: + # Return the unsummed pair: the MTP caller feeds it straight + # into shared_head's fused_add_rms_norm (one kernel instead of + # a separate residual-add + norm). The sum itself is unchanged + # (fp32-accumulated inside the fused kernel). + return hidden_states, residual, None, None + hidden_states = residual + hidden_states + return hidden_states, residual, None, None + + # mHC start. `post`/`comb` carry the previous layer's deferred + # hc_post inputs (its ffn-pre outputs); when present, fuse that + # hc_post with this layer's attn hc_pre into one kernel (inter-layer + # fusion). Layer 0 has no incoming state -> standalone hc_pre. + x = hidden_states + if post is None: + if self.layer_idx == 0: + x = hc_expand(x, self.n) + residual = x + post, comb, x = self.hc_pre( + x, + self.hc_attn_fn, + self.hc_attn_scale, + self.hc_attn_base, + norm_weight=self.input_layernorm.weight.data, + norm_eps=self.input_layernorm.variance_epsilon, + ) + else: + residual, post, comb, x = self.hc_fused_post_pre( + x, + residual, + post, + comb, + self.hc_attn_fn, + self.hc_attn_scale, + self.hc_attn_base, + norm_weight=self.input_layernorm.weight.data, + norm_eps=self.input_layernorm.variance_epsilon, + ) + + # Attention needs the full token sequence; mHC above ran on the SP + # shard. Gather for attention, scatter back afterward (DSv4 pattern). + if self.is_sequence_parallel: + x = sp_all_gather(x)[: positions.shape[0]] + + x = self.self_attn( + hidden_states=x, + positions=positions, + ) + + if self.is_sequence_parallel: + x = sp_reduce_scatter(x) + + # Fuse post-attn hc_post + pre-FFN hc_pre (+ RMSNorm) into one kernel. + residual, post, comb, x = self.hc_fused_post_pre( + x, + residual, + post, + comb, + self.hc_ffn_fn, + self.hc_ffn_scale, + self.hc_ffn_base, + norm_weight=self.post_attention_layernorm.weight.data, + norm_eps=self.post_attention_layernorm.variance_epsilon, + ) + + # Fully Connected + if self._mlp_is_moe: + x = self.mlp(x, already_sequence_parallel=self.is_sequence_parallel) + else: + x = self.mlp(x) + + # mHC end. The last mHC layer materializes its final hc_post (nothing + # to fuse with) then contracts; every other layer defers its hc_post to + # the next layer's fused pre, returning the state. + if self.layer_idx == self.num_hidden_layers - 1: + x = self.hc_post(x, residual, post, comb) + x = hc_contract(x, self.n) + return x, None, None, None + + return x, residual, post, comb + + def hc_pre( + self, + x: torch.Tensor, + hc_fn: torch.Tensor, + hc_scale: torch.Tensor, + hc_base: torch.Tensor, + norm_weight: torch.Tensor | None = None, + norm_eps: float = 0.0, + ): + post_mix, res_mix, layer_input = self.mhc_pre_op( + residual=x, + fn=hc_fn, + hc_scale=hc_scale, + hc_base=hc_base, + rms_eps=self.rms_norm_eps, + hc_pre_eps=self.hc_eps, + hc_sinkhorn_eps=self.hc_eps, + hc_post_mult_value=self.mhc_post_mult_value, + sinkhorn_repeat=self.mhc_sinkhorn_iterations, + norm_weight=norm_weight, + norm_eps=norm_eps, + ) + return post_mix, res_mix, layer_input + + def hc_post( + self, + x: torch.Tensor, + residual: torch.Tensor, + post: torch.Tensor, + comb: torch.Tensor, + ): + return self.mhc_post_op(x, residual, post, comb) + + def hc_fused_post_pre( + self, + x: torch.Tensor, + residual: torch.Tensor, + post: torch.Tensor, + comb: torch.Tensor, + hc_fn: torch.Tensor, + hc_scale: torch.Tensor, + hc_base: torch.Tensor, + norm_weight: torch.Tensor | None = None, + norm_eps: float = 0.0, + ): + return self.mhc_fused_post_pre_op( + x=x, + residual=residual, + post_layer_mix=post, + comb_res_mix=comb, + fn=hc_fn, + hc_scale=hc_scale, + hc_base=hc_base, + rms_eps=self.rms_norm_eps, + hc_pre_eps=self.hc_eps, + hc_sinkhorn_eps=self.hc_eps, + hc_post_mult_value=self.mhc_post_mult_value, + sinkhorn_repeat=self.mhc_sinkhorn_iterations, + n_splits=1, + tile_n=1, + norm_weight=norm_weight, + norm_eps=norm_eps, + ) + + +class Glm5NextModel(nn.Module): + def __init__(self, *, vllm_config: VllmConfig, prefix: str = ""): + super().__init__() + + config = vllm_config.model_config.hf_config + self.config = config + + self.vocab_size = config.vocab_size + self.device = current_platform.device_type + + self.is_v32 = config.index_topk is not None + if self.is_v32: + topk_tokens = config.index_topk + assert topk_tokens is not None + # Reserve room for the incomplete pool tail. + kpool = config.index_kpool + assert kpool is not None + buffer_width = topk_tokens + (kpool - 1 if kpool > 1 else 0) + # Sparse MLA tiles top-k in 128 columns; padded slots remain masked. + sparse_topk_block_n = 128 + buffer_width = ( + (buffer_width + sparse_topk_block_n - 1) // sparse_topk_block_n + ) * sparse_topk_block_n + topk_indices_buffer = torch.empty( + vllm_config.scheduler_config.max_num_batched_tokens, + buffer_width, + dtype=torch.int32, + device=self.device, + ) + else: + # Full-MLA config (no kpool sparse indexer): no topk buffer. + topk_indices_buffer = None + + if get_pp_group().is_first_rank: + self.embed_tokens = VocabParallelEmbedding( + config.vocab_size, + config.hidden_size, + prefix=f"{prefix}.embed_tokens", + ) + else: + self.embed_tokens = PPMissingLayer() + + def get_layer(prefix: str): + layer_idx = int(prefix.rsplit(".", 1)[1]) + return Glm5NextDecoderLayer( + vllm_config=vllm_config, + config=config, + layer_idx=layer_idx, + prefix=prefix, + topk_indices_buffer=topk_indices_buffer, + ) + + self.start_layer, self.end_layer, self.layers = make_layers( + config.num_hidden_layers, + get_layer, + prefix=f"{prefix}.layers", + ) + # The active slice is fixed after construction; cache it so forward + # doesn't rebuild the slice (a fresh list) every step. + self._active_layers = self.layers[self.start_layer : self.end_layer] + + if get_pp_group().is_last_rank: + self.norm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps) + else: + self.norm = PPMissingLayer() + + self.is_sequence_parallel = ( + vllm_config.parallel_config.use_sequence_parallel_moe + ) + + world_size = get_tensor_model_parallel_world_size() + assert config.num_attention_heads % world_size == 0, ( + "num_attention_heads must be divisible by world_size" + ) + + def embed_input_ids(self, input_ids: torch.Tensor) -> torch.Tensor: + return self.embed_tokens(input_ids) + + def forward( + self, + input_ids: torch.Tensor | None, + positions: torch.Tensor, + intermediate_tensors: IntermediateTensors | None, + inputs_embeds: torch.Tensor | None = None, + **kwargs, + ) -> torch.Tensor: + if get_pp_group().is_first_rank: + if inputs_embeds is not None: + hidden_states = inputs_embeds + else: + hidden_states = self.embed_input_ids(input_ids) + residual = None + post = None + comb = None + else: + assert intermediate_tensors is not None + hidden_states = intermediate_tensors["hidden_states"] + residual = intermediate_tensors["residual"] + # post/comb (deferred mHC hc_post state) are not propagated across + # PP ranks; the receiving rank's first mHC layer uses standalone pre. + post = None + comb = None + + full_num_tokens = positions.shape[0] + if self.is_sequence_parallel: + hidden_states = sp_shard(hidden_states) + + for layer in self._active_layers: + hidden_states, residual, post, comb = layer( + positions, hidden_states, residual, post, comb + ) + + if not get_pp_group().is_last_rank: + # PP is gated off for GLM-5.3-Flash (no make_empty_intermediate_tensors), + # so this branch is not exercised. post/comb are the deferred + # hc_post state of this rank's last mHC layer; a future PP path + # would need to propagate them, but for now they are dropped (the + # receiving rank's first layer would fall back to standalone pre). + return IntermediateTensors( + {"hidden_states": hidden_states, "residual": residual} + ) + + if self.is_sequence_parallel: + hidden_states = sp_all_gather(hidden_states)[:full_num_tokens] + + hidden_states = self.norm(hidden_states) + return hidden_states + + def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]: + stacked_params_mapping = [ + # (param_name, shard_name, shard_id) + (".gate_up_proj", ".gate_proj", 0), + (".gate_up_proj", ".up_proj", 1), + # MLA: fuse q_a_proj and kv_a_proj_with_mqa + (".fused_qkv_a_proj", ".q_a_proj", 0), + (".fused_qkv_a_proj", ".kv_a_proj_with_mqa", 1), + # Indexer: fuse wk and weights_proj + (".wk_weights_proj", ".wk", 0), + (".wk_weights_proj", ".weights_proj", 1), + # KDA: merge q, k, v, b, f_a, g_a projections into one GEMM + (".in_proj_qkvbfg_a", ".q_proj", 0), + (".in_proj_qkvbfg_a", ".k_proj", 1), + (".in_proj_qkvbfg_a", ".v_proj", 2), + (".in_proj_qkvbfg_a", ".b_proj", 3), + (".in_proj_qkvbfg_a", ".f_a_proj", 4), + (".in_proj_qkvbfg_a", ".g_a_proj", 5), + ] + if self.config.is_moe: + # Params for weights, fp8 weight scales, fp8 activation scales + # (param_name, weight_name, expert_id, shard_id) + expert_params_mapping = fused_moe_make_expert_params_mapping( + self, + ckpt_gate_proj_name="gate_proj", + ckpt_down_proj_name="down_proj", + ckpt_up_proj_name="up_proj", + num_experts=self.config.n_routed_experts, + ) + else: + expert_params_mapping = [] + params_dict = dict(self.named_parameters()) + loaded_params: set[str] = set() + + # GLM-5.3-Flash NoPE checkpoints omit the RoPE rows from + # ``kv_a_proj_with_mqa``; pad them with zeros for the model shape. + kv_a_pad_size = 0 + if self.config.mla_nope and self.config.qk_rope_head_dim > 0: + kv_a_pad_size = self.config.qk_rope_head_dim + + _pending_wk_fp8: dict = {} + + for args in weights: + name, loaded_weight = args[:2] + kwargs: dict = args[2] if len(args) > 2 else {} + if "rotary_emb.inv_freq" in name: + continue + + spec_layer = get_spec_layer_idx_from_weight_name(self.config, name) + if spec_layer is not None: + continue # skip spec decode layers for main model + if "rotary_emb.cos_cached" in name or "rotary_emb.sin_cached" in name: + # Models trained using ColossalAI may include these tensors in + # the checkpoint. Skip them. + continue + + # Handle FP8 indexer WK: dequantize to BF16 for fusion with + # weights_proj into wk_weights_proj. + if _try_load_fp8_indexer_wk( + name, + loaded_weight, + _pending_wk_fp8, + params_dict, + loaded_params, + ): + continue + + # FP8 checkpoint: dequantize BF16-kept MLA projections + # (q_a_proj / kv_a_proj_with_mqa / o_proj) to BF16. + if _try_load_fp8_attn_proj( + name, + loaded_weight, + _pending_wk_fp8, + params_dict, + loaded_params, + kv_a_pad_size, + ): + continue + + # Pad kv_a_proj_with_mqa for NoPE models + if kv_a_pad_size > 0 and ".kv_a_proj_with_mqa." in name: + pad = torch.zeros( + kv_a_pad_size, + *loaded_weight.shape[1:], + dtype=loaded_weight.dtype, + device=loaded_weight.device, + ) + loaded_weight = torch.cat([loaded_weight, pad], dim=0) + + for param_name, weight_name, shard_id in stacked_params_mapping: + if weight_name not in name: + continue + # We have mlp.experts[0].gate_proj in the checkpoint. + # Since we handle the experts below in expert_params_mapping, + # we need to skip here BEFORE we update the name, otherwise + # name will be updated to mlp.experts[0].gate_up_proj, which + # will then be updated below in expert_params_mapping + # for mlp.experts[0].gate_gate_up_proj, which breaks load. + if ("mlp.experts." in name) and name not in params_dict: + continue + name_mapped = name.replace(weight_name, param_name) + # QKV fusion: skip if fused module doesn't exist in model + if param_name == ".fused_qkv_a_proj" and name_mapped not in params_dict: + continue + name = name_mapped + # Skip loading extra bias for GPTQ models. + if name.endswith(".bias") and name not in params_dict: + continue + if is_pp_missing_parameter(name, self): + continue + param = params_dict[name] + weight_loader = param.weight_loader + weight_loader(param, loaded_weight, shard_id) + break + else: + for idx, ( + param_name, + weight_name, + expert_id, + expert_shard_id, + ) in enumerate(expert_params_mapping): + if weight_name not in name: + continue + name = name.replace(weight_name, param_name) + if is_pp_missing_parameter(name, self): + continue + param = params_dict[name] + weight_loader = param.weight_loader + weight_loader( + param, + loaded_weight, + name, + expert_id=expert_id, + shard_id=expert_shard_id, + ) + break + else: + # Skip loading extra bias for GPTQ models. + if ( + name.endswith(".bias") + and name not in params_dict + and not self.config.is_linear_attn + ): # noqa: E501 + continue + # Remapping the name of FP8 kv-scale. + remapped_name = maybe_remap_kv_scale_name(name, params_dict) + if remapped_name is None: + continue + name = remapped_name + if is_pp_missing_parameter(name, self): + continue + + param = params_dict[name] + weight_loader = getattr( + param, "weight_loader", default_weight_loader + ) + weight_loader(param, loaded_weight, **kwargs) + loaded_params.add(name) + return loaded_params + + +class Glm5NextForCausalLM( + nn.Module, HasInnerState, SupportsPP, MixtureOfExperts, IsHybrid +): + def __init__(self, *, vllm_config: VllmConfig, prefix: str = ""): + super().__init__() + self.model_config = vllm_config.model_config + self.vllm_config = vllm_config + self.config = self.model_config.hf_config + quant_config = vllm_config.quant_config + self.quant_config = quant_config + self.model = Glm5NextModel( + vllm_config=vllm_config, prefix=maybe_prefix(prefix, "model") + ) + if get_pp_group().is_last_rank: + self.lm_head = ParallelLMHead( + self.config.vocab_size, + self.config.hidden_size, + quant_config=quant_config, + prefix=maybe_prefix(prefix, "lm_head"), + ) + else: + self.lm_head = PPMissingLayer() + self.logits_processor = LogitsProcessor( + self.config.vocab_size, scale=self.config.logit_scale + ) + + def embed_input_ids(self, input_ids: torch.Tensor) -> torch.Tensor: + return self.model.embed_input_ids(input_ids) + + def forward( + self, + input_ids: torch.Tensor | None, + positions: torch.Tensor, + intermediate_tensors: IntermediateTensors | None = None, + inputs_embeds: torch.Tensor | None = None, + **kwargs, + ) -> torch.Tensor | IntermediateTensors: + hidden_states = self.model( + input_ids, positions, intermediate_tensors, inputs_embeds, **kwargs + ) + return hidden_states + + @classmethod + def get_mamba_state_dtype_from_config( + cls, + vllm_config: "VllmConfig", + ) -> tuple[torch.dtype, torch.dtype]: + return MambaStateDtypeCalculator.kda_state_dtype( + vllm_config.model_config.dtype, vllm_config.cache_config.mamba_cache_dtype + ) + + @classmethod + def get_mamba_state_shape_from_config( + cls, vllm_config: "VllmConfig" + ) -> tuple[tuple[int, int], tuple[int, int, int]]: + parallel_config = vllm_config.parallel_config + hf_config = vllm_config.model_config.hf_config + tp_size = parallel_config.tensor_parallel_size + num_spec = ( + vllm_config.speculative_config.num_speculative_tokens + if vllm_config.speculative_config + else 0 + ) + return MambaStateShapeCalculator.kda_state_shape( + tp_size, + hf_config.linear_num_heads, + hf_config.linear_head_dim, + conv_kernel_size=hf_config.linear_conv_kernel_dim, + num_spec=num_spec, + ) + + @classmethod + def get_mamba_state_copy_func( + cls, + ) -> tuple[ + MambaStateCopyFunc, MambaStateCopyFunc, MambaStateCopyFunc, MambaStateCopyFunc + ]: + return MambaStateCopyFuncCalculator.kda_state_copy_func() + + def compute_logits( + self, + hidden_states: torch.Tensor, + ) -> torch.Tensor | None: + logits = self.logits_processor(self.lm_head, hidden_states) + return logits + + def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]: + loader = AutoWeightsLoader(self) + return loader.load_weights(weights) + + +@MULTIMODAL_REGISTRY.register_processor( + Glm5NextMultiModalProcessor, + info=Glm5NextProcessingInfo, + dummy_inputs=Glm4vDummyInputsBuilder, +) +class Glm5NextForConditionalGeneration( + Glm4vForConditionalGeneration, HasInnerState, IsHybrid +): + # The text model (KDA + dense-MLA + MoE) is a hybrid mamba model. The + # multimodal wrapper must declare the same interfaces so vLLM treats it as + # hybrid (auto-aligns mamba/attention block sizes, sizes the mamba state + # cache); the mamba-state classmethods delegate to the text model. + has_inner_state: ClassVar[Literal[True]] = True + is_hybrid: ClassVar[Literal[True]] = True + + # NOTE: weight-prefix mapping is inherited from Glm4vForConditionalGeneration + # (``model.visual.`` -> ``visual.``, ``model.language_model.`` -> + # ``language_model.model.``, ``lm_head.`` -> ``language_model.lm_head.``), + # matching the GLM-OCR / GLM-4V serialization convention. If the real + # checkpoint's safetensors keys differ (e.g. ``language_model.model.`` with + # no outer ``model.``), override ``hf_to_vllm_mapper`` accordingly. + + @classmethod + def get_mamba_state_dtype_from_config(cls, vllm_config: VllmConfig): + from .model import Glm5NextForCausalLM + + return Glm5NextForCausalLM.get_mamba_state_dtype_from_config(vllm_config) + + @classmethod + def get_mamba_state_shape_from_config(cls, vllm_config: VllmConfig): + from .model import Glm5NextForCausalLM + + return Glm5NextForCausalLM.get_mamba_state_shape_from_config(vllm_config) + + @classmethod + def get_mamba_state_copy_func(cls): + from .model import Glm5NextForCausalLM + + return Glm5NextForCausalLM.get_mamba_state_copy_func() + + def __init__(self, *, vllm_config: VllmConfig, prefix: str = ""): + super(Glm4vForConditionalGeneration, self).__init__() + config = vllm_config.model_config.hf_config + multimodal_config = vllm_config.model_config.multimodal_config + assert multimodal_config is not None + + self.config = config + self.model_config = vllm_config.model_config + self.multimodal_config = multimodal_config + self.use_data_parallel = multimodal_config.mm_encoder_tp_mode == "data" + self.is_multimodal_pruning_enabled = ( + multimodal_config.is_multimodal_pruning_enabled() + ) + + with self._mark_tower_model(vllm_config, {"image", "video"}): + self.visual = Glm5NextVisionTransformer( + config.text_config, + config.vision_config, + # Read eps from the VISION sub-config, not the top-level + # `config.rms_norm_eps`: Glm5NextConfig.__getattribute__ mirrors + # the latter onto text_config (1e-5), silently ignoring the + # vision tower's own (1e-6) rms_norm_eps. + norm_eps=config.vision_config.rms_norm_eps, + # Vision tower ships BF16 weights in this fp8 checkpoint (no + # weight_scale_inv for visual.*), so it must NOT inherit the + # global fp8 quant_config -- doing so incorrectly quantizes + # the tower + # and yields NaN image features. Mirrors the MLA/KDA proj + # pattern (quant_config=None for BF16 submodules). + quant_config=None, + prefix=maybe_prefix(prefix, "visual"), + ) + + with self._mark_language_model(vllm_config): + self.language_model = init_vllm_registered_model( + vllm_config=vllm_config, + hf_config=config.text_config, + prefix=maybe_prefix(prefix, "language_model"), + architectures=["Glm5NextForCausalLM"], + ) + + # Glm5NextForCausalLM does not implement make_empty_intermediate_tensors, + # so pipeline parallelism is gated off (consistent with the text-only + # model) and we intentionally do not alias it here. + + def get_encoder_cudagraph_config(self): + # This vision tower does not produce the absolute position embedding + # buffer used by GLM4V. + config = super().get_encoder_cudagraph_config() + config.buffer_keys = [k for k in config.buffer_keys if k != "pos_embeds"] + return config + + +def get_spec_layer_idx_from_weight_name( + config: Glm5NextConfig, weight_name: str +) -> int | None: + if hasattr(config, "num_nextn_predict_layers") and ( + config.num_nextn_predict_layers > 0 + ): + layer_idx = config.num_hidden_layers + for i in range(config.num_nextn_predict_layers): + if weight_name.startswith( + f"model.layers.{layer_idx + i}." + ) or weight_name.startswith(f"layers.{layer_idx + i}."): + return layer_idx + i + return None + + +def _try_load_fp8_indexer_wk(name, tensor, buf, params_dict, loaded_params): + if "indexer.wk." not in name or "wk_weights" in name: + return False + is_weight = name.endswith(".weight") and tensor.dtype == torch.float8_e4m3fn + is_scale = "weight_scale_inv" in name + if not is_weight and not is_scale: + return False + layer_prefix = name.rsplit(".wk.", 1)[0] + entry = buf.setdefault(layer_prefix, {}) + entry["weight" if is_weight else "scale"] = tensor + if "weight" not in entry or "scale" not in entry: + return True + + weight_fp8, scale_inv = entry["weight"], entry["scale"] + del buf[layer_prefix] + block_size = weight_fp8.shape[1] // scale_inv.shape[1] + weight_bf16 = scaled_dequantize( + weight_fp8, + scale_inv, + group_shape=GroupShape(block_size, block_size), + out_dtype=torch.bfloat16, + ) + + fused_name = f"{layer_prefix}.wk_weights_proj.weight" + param = params_dict[fused_name] + param.weight_loader(param, weight_bf16, 0) + loaded_params.add(fused_name) + return True + + +def _dequant_fp8_block( + weight_fp8: torch.Tensor, + scale_inv: torch.Tensor, + block_size: int = 128, +) -> torch.Tensor: + """Dequantize a block-FP8 (e4m3) weight with per-block scale to BF16. + + Unlike ``scaled_dequantize`` this tolerates a non-divisible (partial last + block) shape by zero-padding to a multiple of ``block_size`` before the + scale broadcast and trimming back afterwards (e.g. kv_a_proj_with_mqa is + 576 rows = 4*128 + 64). + """ + out_dim, in_dim = weight_fp8.shape + pad_out = (-out_dim) % block_size + pad_in = (-in_dim) % block_size + w = weight_fp8 + if pad_out or pad_in: + w = torch.nn.functional.pad(w, (0, pad_in, 0, pad_out)) + # scale_inv is (ceil(out/block), ceil(in/block)); broadcast to (out, in). + s = scale_inv.to(torch.float32) + s_full = s.repeat_interleave(block_size, dim=0).repeat_interleave(block_size, dim=1) + out = (w.to(torch.float32) * s_full).to(torch.bfloat16) + return out[:out_dim, :in_dim].contiguous() + + +# FP8 checkpoint projections that the MODEL keeps in BF16, so the block-FP8 +# (weight + weight_scale_inv) must be dequantized to BF16 on load. +# Maps checkpoint proj-suffix -> (buffer key, model target base, fused shard id +# or None for a direct projection, whether NoPE rope-padding applies). +_FP8_ATTN_PROJS = { + ".q_a_proj.": ("q_a", "fused_qkv_a_proj", 0, False), + ".kv_a_proj_with_mqa.": ("kv_a", "fused_qkv_a_proj", 1, True), + ".q_b_proj.": ("q_b", "q_b_proj", None, False), + ".o_proj.": ("o_proj", "o_proj", None, False), +} + + +def _try_load_fp8_attn_proj( + name, + tensor, + buf, + params_dict, + loaded_params, + kv_a_pad_size: int, +) -> bool: + """Dequantize FP8 q_a_proj / kv_a_proj_with_mqa / o_proj to BF16 on load. + + The FP8 checkpoint stores these as block-FP8 (weight + weight_scale_inv), + but the model holds them in BF16 (``fused_qkv_a_proj`` is always BF16 via + DeepSeekV2FusedQkvAProjLinear; ``o_proj`` is excluded by + modules_to_not_convert). When the model target is BF16 (no + ``weight_scale_inv`` param) we dequantize; otherwise we return False so the + normal stacked/direct path loads the FP8 tensor as-is. + """ + matched = None + for suffix, info in _FP8_ATTN_PROJS.items(): + if suffix in name: + matched = (suffix, info) + break + if matched is None: + return False + suffix, (key, target_base, shard_id, is_kva) = matched + is_weight = name.endswith(".weight") and tensor.dtype == torch.float8_e4m3fn + is_scale = "weight_scale_inv" in name + if not is_weight and not is_scale: + return False + + layer_prefix = name.rsplit(suffix, 1)[0] + target_w = f"{layer_prefix}.{target_base}.weight" + target_s = f"{layer_prefix}.{target_base}.weight_scale_inv" + # If the model actually kept this projection in FP8, let the normal path + # handle it (it has a weight_scale_inv param). + if target_s in params_dict: + return False + + entry = buf.setdefault(layer_prefix, {}).setdefault(key, {}) + entry["weight" if is_weight else "scale"] = tensor + if "weight" not in entry or "scale" not in entry: + return True + + weight_fp8, scale_inv = entry["weight"], entry["scale"] + buf[layer_prefix].pop(key, None) + block_size = weight_fp8.shape[1] // scale_inv.shape[1] + weight_bf16 = _dequant_fp8_block(weight_fp8, scale_inv, block_size) + # NoPE: pad kv_a rope portion (kv_lora_rank -> kv_lora_rank + qk_rope_head_dim). + if is_kva and kv_a_pad_size > 0: + pad = torch.zeros( + kv_a_pad_size, + weight_bf16.shape[1], + dtype=weight_bf16.dtype, + device=weight_bf16.device, + ) + weight_bf16 = torch.cat([weight_bf16, pad], dim=0) + + param = params_dict[target_w] + if shard_id is None: + param.weight_loader(param, weight_bf16) + else: + param.weight_loader(param, weight_bf16, shard_id) + loaded_params.add(target_w) + return True diff --git a/vllm/models/glm5next/nvidia/mtp.py b/vllm/models/glm5next/nvidia/mtp.py new file mode 100644 index 000000000000..ae1f7f56b5ca --- /dev/null +++ b/vllm/models/glm5next/nvidia/mtp.py @@ -0,0 +1,429 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +import typing +from collections.abc import Callable, Iterable + +import torch +import torch.nn as nn + +from vllm.config import VllmConfig +from vllm.model_executor.layers.fused_moe import ( + fused_moe_make_expert_params_mapping, +) +from vllm.model_executor.layers.layernorm import RMSNorm +from vllm.model_executor.layers.logits_processor import LogitsProcessor +from vllm.model_executor.layers.vocab_parallel_embedding import ( + VocabParallelEmbedding, +) +from vllm.model_executor.model_loader.weight_utils import ( + default_weight_loader, + maybe_remap_kv_scale_name, +) +from vllm.model_executor.models.deepseek_mtp import SharedHead +from vllm.model_executor.models.deepseek_v2 import DeepseekV2MixtureOfExperts +from vllm.model_executor.models.utils import maybe_prefix +from vllm.platforms import current_platform +from vllm.sequence import IntermediateTensors + +from .model import ( + Glm5NextDecoderLayer, + Glm5NextMLAAttention, + Glm5NextMoE, + _try_load_fp8_attn_proj, + _try_load_fp8_indexer_wk, + get_spec_layer_idx_from_weight_name, +) +from .ops.fused_eh_norm import fused_eh_norm + + +class Glm5NextMultiTokenPredictorLayer(nn.Module): + def __init__(self, vllm_config: VllmConfig, prefix: str) -> None: + super().__init__() + assert vllm_config.speculative_config is not None + config = vllm_config.speculative_config.draft_model_config.hf_config + self.config = config + quant_config = vllm_config.quant_config + + self.enorm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps) + self.hnorm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps) + self.eh_proj = nn.Linear(config.hidden_size * 2, config.hidden_size, bias=False) + + # Reserve room for the incomplete pool tail and align the sparse MLA + # buffer width to BLOCK_N=128. + topk_tokens = config.index_topk + assert topk_tokens is not None + kpool = config.index_kpool + assert kpool is not None + buffer_width = topk_tokens + (kpool - 1 if kpool > 1 else 0) + sparse_topk_block_n = 128 + buffer_width = ( + (buffer_width + sparse_topk_block_n - 1) // sparse_topk_block_n + ) * sparse_topk_block_n + topk_indices_buffer = torch.empty( + vllm_config.scheduler_config.max_num_batched_tokens, + buffer_width, + dtype=torch.int32, + device=current_platform.device_type, + ) + self.shared_head = SharedHead( + config=config, prefix=prefix, quant_config=quant_config + ) + # MTP layers sit past the base model's hidden layers; parse the index + # from the prefix (e.g. "...layers.32") so the decoder builds an MLA + # (DSA) layer rather than KDA for the MTP path. + layer_idx = int(prefix.rsplit(".", 1)[-1]) + self.mtp_block = Glm5NextDecoderLayer( + vllm_config=vllm_config, + config=config, + layer_idx=layer_idx, + prefix=prefix, + topk_indices_buffer=topk_indices_buffer, + is_mtp_layer=True, + ) + + def forward( + self, + input_ids: torch.Tensor, + positions: torch.Tensor, + previous_hidden_states: torch.Tensor, + inputs_embeds: torch.Tensor | None = None, + spec_step_index: int = 0, + ) -> torch.Tensor: + assert inputs_embeds is not None + # Fused: zero pos-0 embeds + enorm(embeds) + hnorm(prev) + cat -> [N, 2H]. + eh_input = fused_eh_norm( + positions, + inputs_embeds, + previous_hidden_states, + self.enorm.weight, + self.hnorm.weight, + self.enorm.variance_epsilon, + ) + hidden_states = self.eh_proj(eh_input) + # Fuse the residual add and final RMSNorm. Glm5NextMoE already performs + # its all-reduce, so no collective is needed here. The post-norm result + # feeds both draft logits and the next recycled hidden state. + hidden_states, residual, _, _ = self.mtp_block( + positions=positions, hidden_states=hidden_states, residual=None + ) + hidden_states, _ = self.shared_head.norm(hidden_states, residual=residual) + return hidden_states, hidden_states + + +class Glm5NextMultiTokenPredictor(nn.Module): + def __init__(self, *, vllm_config: VllmConfig, prefix: str = ""): + super().__init__() + config = vllm_config.model_config.hf_config + self.mtp_start_layer_idx = config.num_hidden_layers + self.num_mtp_layers = config.num_nextn_predict_layers + self.layers = torch.nn.ModuleDict( + { + str(idx): Glm5NextMultiTokenPredictorLayer( + vllm_config, f"{prefix}.layers.{idx}" + ) + for idx in range( + self.mtp_start_layer_idx, + self.mtp_start_layer_idx + self.num_mtp_layers, + ) + } + ) + self.embed_tokens = VocabParallelEmbedding( + config.vocab_size, + config.hidden_size, + prefix=maybe_prefix(prefix, "embed_tokens"), + ) + # Plain list for the per-propose lookup: ModuleDict[str(...)] builds a + # string and hashes it on every draft step. + self._mtp_layers = list(self.layers.values()) + self._mtp_mla_attns = [] + for layer in self._mtp_layers: + self_attn = layer.mtp_block.self_attn + assert isinstance(self_attn, Glm5NextMLAAttention) + self._mtp_mla_attns.append(self_attn.mla_attn) + self.logits_processor = LogitsProcessor(config.vocab_size) + + def set_skip_topk(self, skip: bool): + # index_share_for_mtp_iteration: step 0 computes top-k, steps 1+ reuse. + for mla_attn in self._mtp_mla_attns: + mla_attn.skip_topk = skip + + def compact_topk_indices(self, slot_ids: torch.Tensor): + """Gather the top-k index rows at ``slot_ids`` to the front of the buffer.""" + num_slots = slot_ids.numel() + for mla_attn in self._mtp_mla_attns: + topk_indices_buffer = mla_attn.topk_indices_buffer + assert topk_indices_buffer is not None + topk_indices_buffer[:num_slots] = topk_indices_buffer[slot_ids] + + def embed_input_ids(self, input_ids: torch.Tensor) -> torch.Tensor: + return self.embed_tokens(input_ids) + + def forward( + self, + input_ids: torch.Tensor, + positions: torch.Tensor, + previous_hidden_states: torch.Tensor, + inputs_embeds: torch.Tensor | None = None, + spec_step_idx: int = 0, + ) -> torch.Tensor: + if inputs_embeds is None: + inputs_embeds = self.embed_tokens(input_ids) + current_step_idx = spec_step_idx % self.num_mtp_layers + return self._mtp_layers[current_step_idx]( + input_ids, + positions, + previous_hidden_states, + inputs_embeds, + current_step_idx, + ) + + def compute_logits( + self, + hidden_states: torch.Tensor, + spec_step_idx: int = 0, + ) -> torch.Tensor: + current_step_idx = spec_step_idx % self.num_mtp_layers + mtp_layer = self._mtp_layers[current_step_idx] + # hidden_states is already post-final-norm (produced in the layer + # forward and recycled as-is); apply the LM head only, without a + # second RMSNorm. + return self.logits_processor(mtp_layer.shared_head.head, hidden_states) + + def get_top_tokens( + self, + hidden_states: torch.Tensor, + spec_step_idx: int = 0, + ) -> torch.Tensor: + current_step_idx = spec_step_idx % self.num_mtp_layers + mtp_layer = self._mtp_layers[current_step_idx] + # Vocab-parallel argmax for the greedy draft: per-rank head projection + # + local argmax + a [batch, 2*tp] (value, index) reduce, instead of + # materializing and all-gathering full [N, vocab] logits per draft + # step. Tie-breaking matches the full argmax (shards are contiguous + # and rank-ordered, so the lowest-rank winner is the lowest global + # index), so greedy draft tokens are unchanged. + return self.logits_processor.get_top_tokens( + mtp_layer.shared_head.head, hidden_states + ) + + +class Glm5NextMTP(nn.Module, DeepseekV2MixtureOfExperts): + def __init__(self, *, vllm_config: VllmConfig, prefix: str = ""): + super().__init__() + self.config = vllm_config.model_config.hf_config + self.quant_config = vllm_config.quant_config + self.model = Glm5NextMultiTokenPredictor( + vllm_config=vllm_config, prefix=maybe_prefix(prefix, "model") + ) + self.set_moe_parameters() + + def set_moe_parameters(self): + self.num_moe_layers = self.config.num_nextn_predict_layers + self.num_expert_groups = self.config.n_group + self.moe_layers = [] + self.moe_mlp_layers = [] + example_moe = None + for layer in self.model.layers.values(): + mlp = layer.mtp_block.mlp + if isinstance(mlp, Glm5NextMoE): + example_moe = mlp + self.moe_mlp_layers.append(mlp) + self.moe_layers.append(mlp.experts) + self.extract_moe_parameters(example_moe) + + def embed_input_ids(self, input_ids: torch.Tensor) -> torch.Tensor: + return self.model.embed_input_ids(input_ids) + + def forward( + self, + input_ids: torch.Tensor | None, + positions: torch.Tensor, + hidden_states: torch.Tensor, + intermediate_tensors: IntermediateTensors | None = None, + inputs_embeds: torch.Tensor | None = None, + spec_step_idx: int = 0, + ) -> torch.Tensor: + return self.model( + input_ids, positions, hidden_states, inputs_embeds, spec_step_idx + ) + + def compute_logits( + self, + hidden_states: torch.Tensor, + spec_step_idx: int = 0, + ) -> torch.Tensor | None: + return self.model.compute_logits(hidden_states, spec_step_idx) + + def get_top_tokens( + self, + hidden_states: torch.Tensor, + spec_step_idx: int = 0, + ) -> torch.Tensor: + # Greedy-draft path used when use_local_argmax_reduction is enabled: + # vocab-parallel argmax, no full-vocab logits. + return self.model.get_top_tokens(hidden_states, spec_step_idx) + + def _rewrite_spec_layer_name(self, spec_layer: int, name: str) -> str: + spec_layer_weight_names = [ + "embed_tokens", + "enorm", + "hnorm", + "eh_proj", + "shared_head", + ] + shared_weight_names = ["embed_tokens"] + spec_layer_weight = False + shared_weight = False + for weight_name in spec_layer_weight_names: + if weight_name in name: + spec_layer_weight = True + if weight_name in shared_weight_names: + shared_weight = True + break + if not spec_layer_weight: + name = name.replace( + f"model.layers.{spec_layer}.", f"model.layers.{spec_layer}.mtp_block." + ) + elif shared_weight: + name = name.replace(f"model.layers.{spec_layer}.", "model.") + return name + + def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]: + stacked_params_mapping = [ + ("gate_up_proj", "gate_proj", 0), + ("gate_up_proj", "up_proj", 1), + ("fused_qkv_a_proj", "q_a_proj", 0), + ("fused_qkv_a_proj", "kv_a_proj_with_mqa", 1), + ("wk_weights_proj", "wk", 0), + ("wk_weights_proj", "weights_proj", 1), + ] + expert_params_mapping = fused_moe_make_expert_params_mapping( + self, + ckpt_gate_proj_name="gate_proj", + ckpt_down_proj_name="down_proj", + ckpt_up_proj_name="up_proj", + num_experts=self.config.n_routed_experts, + ) + + params_dict = dict(self.named_parameters()) + loaded_params: set[str] = set() + _pending_wk_fp8: dict = {} + # GLM-5.3-Flash NoPE checkpoints omit the RoPE rows from + # ``kv_a_proj_with_mqa``; the FP8-to-BF16 path pads them for the model. + kv_a_pad_size = 0 + if self.config.mla_nope and self.config.qk_rope_head_dim > 0: + kv_a_pad_size = self.config.qk_rope_head_dim + for name, loaded_weight in weights: + if "rotary_emb.inv_freq" in name: + continue + # Multimodal (Glm5NextForConditionalGeneration) checkpoints prefix + # the text-tower weights with "model.language_model."; the MTP head + # is built as a text-only model (model.layers.*), so strip the + # prefix to match. + if name.startswith("model.language_model."): + name = name.replace("model.language_model.", "model.", 1) + spec_layer = get_spec_layer_idx_from_weight_name(self.config, name) + if spec_layer is None: + continue + name = self._rewrite_spec_layer_name(spec_layer, name) + + if _try_load_fp8_indexer_wk( + name, + loaded_weight, + _pending_wk_fp8, + params_dict, + loaded_params, + ): + continue + + # FP8 checkpoint: dequantize the BF16-kept MLA projections + # (q_a_proj / kv_a_proj_with_mqa / o_proj) to BF16, mirroring the + # target model. The model holds fused_qkv_a_proj / o_proj in BF16, + # so the checkpoint's block-FP8 weight + weight_scale_inv for these + # has no param home and would KeyError without this dequant. + if _try_load_fp8_attn_proj( + name, + loaded_weight, + _pending_wk_fp8, + params_dict, + loaded_params, + kv_a_pad_size, + ): + continue + + for param_name, weight_name, shard_id in stacked_params_mapping: + if weight_name not in name: + continue + if ("mlp.experts." in name) and name not in params_dict: + continue + name_mapped = name.replace(weight_name, param_name) + if ( + param_name == "fused_qkv_a_proj" + ) and name_mapped not in params_dict: + continue + else: + name = name_mapped + if name.endswith(".bias") and name not in params_dict: + continue + param = params_dict[name] + weight_loader = param.weight_loader + weight_loader(param, loaded_weight, shard_id) + break + else: + is_expert_weight = False + for mapping in expert_params_mapping: + param_name, weight_name, expert_id, shard_id = mapping # type: ignore[assignment] + if weight_name not in name: + continue + is_expert_weight = True + name_mapped = name.replace(weight_name, param_name) + param = params_dict[name_mapped] + weight_loader = typing.cast( + Callable[..., bool], param.weight_loader + ) + success = weight_loader( + param, + loaded_weight, + name_mapped, + shard_id=shard_id, + expert_id=expert_id, + return_success=True, + ) + if success: + name = name_mapped + break + else: + if is_expert_weight: + continue + if name.endswith(".bias") and name not in params_dict: + continue + name = maybe_remap_kv_scale_name(name, params_dict) # type: ignore[assignment] + if name is None: + continue + if ( + spec_layer != self.model.mtp_start_layer_idx + and ".layers" not in name + ): + continue + param = params_dict[name] + weight_loader = getattr( + param, "weight_loader", default_weight_loader + ) + weight_loader(param, loaded_weight) + loaded_params.add(name) + + loaded_layers: set[int] = set() + for param_name in loaded_params: + spec_layer = get_spec_layer_idx_from_weight_name(self.config, param_name) + if spec_layer is not None: + loaded_layers.add(spec_layer) + for layer_idx in range( + self.model.mtp_start_layer_idx, + self.model.mtp_start_layer_idx + self.model.num_mtp_layers, + ): + if layer_idx not in loaded_layers: + raise ValueError( + f"MTP speculative decoding layer {layer_idx} weights " + f"missing from checkpoint." + ) + return loaded_params diff --git a/vllm/models/glm5next/nvidia/multimodal.py b/vllm/models/glm5next/nvidia/multimodal.py new file mode 100644 index 000000000000..d2b21345d052 --- /dev/null +++ b/vllm/models/glm5next/nvidia/multimodal.py @@ -0,0 +1,732 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +"""GLM-5.3-Flash vision tower and multimodal processor.""" + +from collections.abc import Mapping +from functools import cached_property, partial + +import numpy as np +import torch +import torch.nn as nn +from einops import rearrange + +from vllm.distributed import ( + get_tensor_model_parallel_world_size, + parallel_state, +) +from vllm.distributed import utils as dist_utils +from vllm.model_executor.layers.activation import SiluAndMulWithClamp +from vllm.model_executor.layers.attention import MMEncoderAttention +from vllm.model_executor.layers.conv import Conv2dLayer, Conv3dLayer +from vllm.model_executor.layers.layernorm import RMSNorm +from vllm.model_executor.layers.linear import ( + ColumnParallelLinear, + MergedColumnParallelLinear, + QKVParallelLinear, + RowParallelLinear, +) +from vllm.model_executor.layers.quantization import QuantizationConfig +from vllm.model_executor.layers.rotary_embedding import get_rope +from vllm.model_executor.layers.rotary_embedding.common import ApplyRotaryEmb +from vllm.model_executor.models.glm4_1v import ( + Glm4vMultiModalProcessor, + Glm4vProcessingInfo, +) +from vllm.model_executor.models.utils import ( + AutoWeightsLoader, + WeightsMapper, +) +from vllm.model_executor.models.vision import ( + get_vit_attn_backend, + is_vit_use_data_parallel, +) +from vllm.models.common.ops import fused_q_kv_rmsnorm +from vllm.multimodal.parse import ImageSize, MultiModalDataItems +from vllm.v1.attention.backends.registry import AttentionBackendEnum + + +class Glm5NextVisionPatchEmbed(nn.Module): + def __init__( + self, + patch_size: int = 14, + temporal_patch_size: int = 1, + in_channels: int = 3, + hidden_size: int = 1536, + ) -> None: + super().__init__() + self.patch_size = patch_size + self.temporal_patch_size = temporal_patch_size + self.hidden_size = hidden_size + + kernel_size = (temporal_patch_size, patch_size, patch_size) + self.proj = Conv3dLayer( + in_channels, + hidden_size, + kernel_size=kernel_size, + stride=kernel_size, + bias=True, + ) + + def forward(self, x: torch.Tensor) -> torch.Tensor: + L, C = x.shape + x = x.view(L, -1, self.temporal_patch_size, self.patch_size, self.patch_size) + x = self.proj(x).view(L, self.hidden_size) + return x + + +class Glm5NextVisionMLP(nn.Module): + def __init__( + self, + in_features: int, + hidden_features: int, + swiglu_limit: float, + bias: bool = True, + quant_config: QuantizationConfig | None = None, + prefix: str = "", + ): + super().__init__() + use_data_parallel = is_vit_use_data_parallel() + self.gate_up_proj = MergedColumnParallelLinear( + input_size=in_features, + output_sizes=[hidden_features] * 2, + bias=bias, + quant_config=quant_config, + prefix=f"{prefix}.gate_up_proj", + disable_tp=use_data_parallel, + ) + self.down_proj = RowParallelLinear( + hidden_features, + in_features, + bias=bias, + quant_config=quant_config, + prefix=f"{prefix}.down_proj", + disable_tp=use_data_parallel, + ) + # GLM-5.3-Flash clamps the vision SwiGLU gate/up unlike GLM-OCR/GLM-4V. + self.act_fn = SiluAndMulWithClamp(swiglu_limit=swiglu_limit) + + def forward(self, x: torch.Tensor): + x, _ = self.gate_up_proj(x) + x = self.act_fn(x) + x, _ = self.down_proj(x) + return x + + +class Glm5NextVisionAttention(nn.Module): + def __init__( + self, + embed_dim: int, + num_heads: int, + projection_size: int, + quant_config: QuantizationConfig | None = None, + prefix: str = "", + ) -> None: + super().__init__() + use_data_parallel = is_vit_use_data_parallel() + self.tp_size = ( + 1 if use_data_parallel else get_tensor_model_parallel_world_size() + ) + self.tp_rank = ( + 0 if use_data_parallel else parallel_state.get_tensor_model_parallel_rank() + ) + self.hidden_size_per_attention_head = dist_utils.divide( + projection_size, num_heads + ) + self.num_attention_heads_per_partition = dist_utils.divide( + num_heads, self.tp_size + ) + + self.head_dim = embed_dim // num_heads + + # q/k norm eps hard-coded 1e-5 — distinct from block/post norm eps. + self.q_norm = RMSNorm(self.head_dim, eps=1e-5) + self.k_norm = RMSNorm(self.head_dim, eps=1e-5) + + self.qkv = QKVParallelLinear( + hidden_size=embed_dim, + head_size=self.hidden_size_per_attention_head, + total_num_heads=num_heads, + total_num_kv_heads=num_heads, + bias=True, + quant_config=quant_config, + prefix=f"{prefix}.qkv_proj" if quant_config else f"{prefix}.qkv", + disable_tp=use_data_parallel, + ) + self.proj = RowParallelLinear( + input_size=projection_size, + output_size=embed_dim, + quant_config=quant_config, + prefix=f"{prefix}.proj", + bias=True, + disable_tp=use_data_parallel, + ) + + self.attn = MMEncoderAttention( + num_heads=self.num_attention_heads_per_partition, + head_size=self.hidden_size_per_attention_head, + scale=self.hidden_size_per_attention_head**-0.5, + prefix=f"{prefix}.attn", + ) + self.apply_rotary_emb = ApplyRotaryEmb(enforce_enable=True) + + def split_qkv(self, qkv: torch.Tensor) -> tuple[torch.Tensor, ...]: + seq_len, bs, _ = qkv.shape + q, k, v = qkv.chunk(3, dim=2) + new_shape = ( + seq_len, + bs, + self.num_attention_heads_per_partition, + self.hidden_size_per_attention_head, + ) + q, k, v = (x.view(*new_shape) for x in (q, k, v)) + return q, k, v + + def forward( + self, + x: torch.Tensor, + cu_seqlens: torch.Tensor, + rotary_pos_emb_cos: torch.Tensor, + rotary_pos_emb_sin: torch.Tensor, + max_seqlen: torch.Tensor | None = None, + ) -> torch.Tensor: + x, _ = self.qkv(x) + q, k, v = self.split_qkv(x) + + # P1: fused q/k RMSNorm (two distinct weights, one launch; fp32, bit-identical). + q_shape, k_shape = q.shape, k.shape + q_flat = q.reshape(-1, self.head_dim) + k_flat = k.reshape(-1, self.head_dim) + q, k = fused_q_kv_rmsnorm( + q_flat, + k_flat, + self.q_norm.weight, + self.k_norm.weight, + self.q_norm.variance_epsilon, + ) + q = q.view(q_shape) + k = k.view(k_shape) + + q, k, v = (rearrange(t, "s b ... -> b s ...").contiguous() for t in (q, k, v)) + if rotary_pos_emb_cos is not None and rotary_pos_emb_sin is not None: + qk_concat = torch.cat([q, k], dim=0) + qk_rotated = self.apply_rotary_emb( + qk_concat, + rotary_pos_emb_cos, + rotary_pos_emb_sin, + ) + q, k = torch.chunk(qk_rotated, 2, dim=0) + + context_layer = self.attn( + query=q, + key=k, + value=v, + cu_seqlens=cu_seqlens, + max_seqlen=max_seqlen, + ) + context_layer = rearrange(context_layer, "b s h d -> s b (h d)").contiguous() + + output, _ = self.proj(context_layer) + return output + + +class Glm5NextVisionBlock(nn.Module): + def __init__( + self, + dim: int, + num_heads: int, + mlp_hidden_dim: int, + swiglu_limit: float, + norm_layer: partial[nn.Module] | None = None, + quant_config: QuantizationConfig | None = None, + prefix: str = "", + ) -> None: + super().__init__() + if norm_layer is None: + norm_layer = partial(nn.LayerNorm, eps=1e-6) + self.norm1 = norm_layer(dim) + self.norm2 = norm_layer(dim) + self.attn = Glm5NextVisionAttention( + embed_dim=dim, + num_heads=num_heads, + projection_size=dim, + quant_config=quant_config, + prefix=f"{prefix}.attn", + ) + self.mlp = Glm5NextVisionMLP( + dim, + mlp_hidden_dim, + swiglu_limit=swiglu_limit, + bias=True, + quant_config=quant_config, + prefix=f"{prefix}.mlp", + ) + + def forward( + self, + x: torch.Tensor, + cu_seqlens: torch.Tensor, + rotary_pos_emb_cos: torch.Tensor, + rotary_pos_emb_sin: torch.Tensor, + max_seqlen: int | None = None, + ) -> torch.Tensor: + x_attn = self.attn( + self.norm1(x), + cu_seqlens=cu_seqlens, + rotary_pos_emb_cos=rotary_pos_emb_cos, + rotary_pos_emb_sin=rotary_pos_emb_sin, + max_seqlen=max_seqlen, + ) + x_fused_norm, residual = self.norm2(x, residual=x_attn) + x = residual + self.mlp(x_fused_norm) + return x + + +class Glm5NextPatchMerger(nn.Module): + def __init__( + self, + d_model: int, + context_dim: int, + swiglu_limit: float, + quant_config: QuantizationConfig | None = None, + bias: bool = False, + prefix: str = "", + ) -> None: + super().__init__() + use_data_parallel = is_vit_use_data_parallel() + self.hidden_size = d_model + self.proj = ColumnParallelLinear( + self.hidden_size, + self.hidden_size, + bias=bias, + gather_output=True, + quant_config=quant_config, + prefix=f"{prefix}.proj", + disable_tp=use_data_parallel, + ) + self.post_projection_norm = nn.LayerNorm(self.hidden_size) + self.gate_up_proj = MergedColumnParallelLinear( + input_size=self.hidden_size, + output_sizes=[context_dim] * 2, + bias=bias, + quant_config=quant_config, + prefix=f"{prefix}.gate_up_proj", + disable_tp=use_data_parallel, + ) + self.down_proj = RowParallelLinear( + context_dim, + self.hidden_size, + bias=bias, + quant_config=quant_config, + prefix=f"{prefix}.down_proj", + disable_tp=use_data_parallel, + ) + # GLM-5.3-Flash also clamps the merger SwiGLU. + self.act_fn = SiluAndMulWithClamp(swiglu_limit=swiglu_limit) + self.extra_activation_func = nn.GELU() + + def forward(self, x: torch.Tensor): + x, _ = self.proj(x) + x = self.extra_activation_func(self.post_projection_norm(x)) + gate_up, _ = self.gate_up_proj(x) + x = self.act_fn(gate_up) + x, _ = self.down_proj(x) + return x + + +class Glm5NextVisionTransformer(nn.Module): + # Stacked-weight remap for the GLM-OCR/GLM-4V vision checkpoint layout. + hf_to_vllm_mapper = WeightsMapper( + orig_to_new_stacked={ + ".attn.q.": (".attn.qkv.", "q"), + ".attn.k.": (".attn.qkv.", "k"), + ".attn.v.": (".attn.qkv.", "v"), + ".gate_proj": (".gate_up_proj", 0), + ".up_proj": (".gate_up_proj", 1), + } + ) + + def __init__( + self, + text_config, # noqa: ANN001 + vision_config, + norm_eps: float = 1e-6, + quant_config: QuantizationConfig | None = None, + prefix: str = "", + ) -> None: + super().__init__() + use_data_parallel = is_vit_use_data_parallel() + self.tp_size = ( + 1 if use_data_parallel else get_tensor_model_parallel_world_size() + ) + + patch_size = vision_config.patch_size + temporal_patch_size = vision_config.temporal_patch_size + in_channels = vision_config.in_channels + depth = vision_config.depth + self.hidden_size = vision_config.hidden_size + self.num_heads = vision_config.num_heads + + self.patch_size = vision_config.patch_size + self.spatial_merge_size = vision_config.spatial_merge_size + self.out_hidden_size = vision_config.out_hidden_size + + swiglu_limit = vision_config.swiglu_limit + if swiglu_limit is None: + swiglu_limit = text_config.swiglu_limit + assert swiglu_limit is not None, ( + "GLM-5.3-Flash vision requires swiglu_limit (vision_config or text_config)" + ) + + # Single construction pass — no abs-pos embeddings / post-conv norm (OCR delta). + self.patch_embed = Glm5NextVisionPatchEmbed( + patch_size=patch_size, + temporal_patch_size=temporal_patch_size, + in_channels=in_channels, + hidden_size=self.hidden_size, + ) + + norm_layer = partial(RMSNorm, eps=norm_eps) + head_dim = self.hidden_size // self.num_heads + self.rotary_pos_emb = get_rope( + head_size=head_dim, + max_position=8192, + is_neox_style=True, + rope_parameters={"partial_rotary_factor": 0.5}, + ) + self.blocks = nn.ModuleList( + [ + Glm5NextVisionBlock( + dim=self.hidden_size, + num_heads=self.num_heads, + mlp_hidden_dim=vision_config.intermediate_size, + swiglu_limit=swiglu_limit, + norm_layer=norm_layer, + quant_config=quant_config, + prefix=f"{prefix}.blocks.{layer_idx}", + ) + for layer_idx in range(depth) + ] + ) + # GLM-5.3-Flash merger bottleneck width. + self.merger = Glm5NextPatchMerger( + d_model=vision_config.out_hidden_size, + context_dim=vision_config.projection_intermediate_size, + swiglu_limit=swiglu_limit, + quant_config=quant_config, + bias=False, + prefix=f"{prefix}.merger", + ) + + self.downsample = Conv2dLayer( + in_channels=vision_config.hidden_size, + out_channels=vision_config.out_hidden_size, + kernel_size=vision_config.spatial_merge_size, + stride=vision_config.spatial_merge_size, + ) + self.post_layernorm = RMSNorm( + vision_config.hidden_size, eps=vision_config.rms_norm_eps + ) + + self.attn_backend = get_vit_attn_backend( + head_size=head_dim, + dtype=torch.get_default_dtype(), + ) + + @property + def dtype(self) -> torch.dtype: + return self.patch_embed.proj.weight.dtype + + @property + def device(self) -> torch.device: + return self.patch_embed.proj.weight.device + + def rot_pos_emb( + self, grid_thw: list[list[int]] + ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: + pos_ids = [] + for t, h, w in grid_thw: + hpos_ids = torch.arange(h).unsqueeze(1).expand(-1, w) + wpos_ids = torch.arange(w).unsqueeze(0).expand(h, -1) + hpos_ids = ( + hpos_ids.reshape( + h // self.spatial_merge_size, + self.spatial_merge_size, + w // self.spatial_merge_size, + self.spatial_merge_size, + ) + .permute(0, 2, 1, 3) + .flatten() + ) + wpos_ids = ( + wpos_ids.reshape( + h // self.spatial_merge_size, + self.spatial_merge_size, + w // self.spatial_merge_size, + self.spatial_merge_size, + ) + .permute(0, 2, 1, 3) + .flatten() + ) + pos_ids.append(torch.stack([hpos_ids, wpos_ids], dim=-1).repeat(t, 1)) + pos_ids = torch.cat(pos_ids, dim=0) + max_grid_size = max(max(h, w) for _, h, w in grid_thw) + + cos, sin = self.rotary_pos_emb.get_cos_sin(max_grid_size) + + pos_ids = pos_ids.to(cos.device, non_blocking=True) + cos_combined = cos[pos_ids].flatten(1) + sin_combined = sin[pos_ids].flatten(1) + return cos_combined, sin_combined, pos_ids + + def compute_attn_mask_seqlen( + self, + cu_seqlens: torch.Tensor, + ) -> torch.Tensor | None: + max_seqlen = None + if self.attn_backend in { + AttentionBackendEnum.FLASH_ATTN, + AttentionBackendEnum.ROCM_AITER_FA, + AttentionBackendEnum.TRITON_ATTN, + }: + max_seqlen = (cu_seqlens[1:] - cu_seqlens[:-1]).max() + return max_seqlen + + def prepare_encoder_metadata( + self, + grid_thw_list: list[list[int]], + *, + max_batch_size: int | None = None, + max_frames_per_batch: int | None = None, + max_seqlen_override: int | None = None, + device: torch.device | None = None, + ) -> dict[str, torch.Tensor | None]: + """Compute encoder metadata for eager and CUDA graph execution.""" + if device is None: + device = self.device + + metadata: dict[str, torch.Tensor | None] = {} + + rotary_cos, rotary_sin, _ = self.rot_pos_emb(grid_thw_list) + metadata["rotary_pos_emb_cos"] = rotary_cos + metadata["rotary_pos_emb_sin"] = rotary_sin + + grid_thw_np = np.array(grid_thw_list, dtype=np.int32) + patches_per_frame = grid_thw_np[:, 1] * grid_thw_np[:, 2] + cu_seqlens = np.repeat(patches_per_frame, grid_thw_np[:, 0]).cumsum( + dtype=np.int32 + ) + cu_seqlens = np.concatenate([np.zeros(1, dtype=np.int32), cu_seqlens]) + + pad_to = ( + max_frames_per_batch if max_frames_per_batch is not None else max_batch_size + ) + if pad_to is not None: + num_seqs = len(cu_seqlens) - 1 + if num_seqs < pad_to: + cu_seqlens = np.concatenate( + [ + cu_seqlens, + np.full( + pad_to - num_seqs, + cu_seqlens[-1], + dtype=np.int32, + ), + ] + ) + + metadata["sequence_lengths"] = MMEncoderAttention.maybe_compute_seq_lens( + self.attn_backend, cu_seqlens, device + ) + + if max_seqlen_override is not None: + max_seqlen_val = max_seqlen_override + else: + max_seqlen_val = MMEncoderAttention.compute_max_seqlen( + self.attn_backend, cu_seqlens + ) + metadata["max_seqlen"] = torch.tensor(max_seqlen_val, dtype=torch.int32) + + metadata["cu_seqlens"] = MMEncoderAttention.maybe_recompute_cu_seqlens( + self.attn_backend, + cu_seqlens, + self.hidden_size, + self.tp_size, + device, + ) + + return metadata + + def forward( + self, + x: torch.Tensor, + grid_thw: torch.Tensor | list[list[int]], + *, + encoder_metadata: dict[str, torch.Tensor] | None = None, + ) -> torch.Tensor: + # patchify + x = x.to(device=self.device, dtype=self.dtype) + x = self.patch_embed(x) + + if encoder_metadata is not None: + # Encoder CUDA-graph path (PR #49852): rotary/cu_seqlens/max_seqlen are + # precomputed by prepare_encoder_metadata (which uses rot_pos_emb exactly + # as the eager rebuild does), so reuse them and skip the per-call CPU + # rebuild (the low-GPU-util culprit on multimodal workloads). + rotary_pos_emb_cos = encoder_metadata["rotary_pos_emb_cos"] + rotary_pos_emb_sin = encoder_metadata["rotary_pos_emb_sin"] + cu_seqlens = encoder_metadata["cu_seqlens"] + max_seqlen = encoder_metadata["max_seqlen"] + else: + if isinstance(grid_thw, list): + grid_thw = torch.tensor(grid_thw, dtype=torch.int32) + rotary_pos_emb_cos, rotary_pos_emb_sin, _ = self.rot_pos_emb(grid_thw) + cu_seqlens = torch.repeat_interleave( + grid_thw[:, 1] * grid_thw[:, 2], grid_thw[:, 0] + ).cumsum(dim=0, dtype=torch.int32) + cu_seqlens = torch.cat([cu_seqlens.new_zeros(1), cu_seqlens]) + cu_seqlens = cu_seqlens.to(self.device, non_blocking=True) + max_seqlen = self.compute_attn_mask_seqlen(cu_seqlens) + + # transformers + x = x.unsqueeze(1) + for blk in self.blocks: + x = blk( + x, + cu_seqlens=cu_seqlens, + rotary_pos_emb_cos=rotary_pos_emb_cos, + rotary_pos_emb_sin=rotary_pos_emb_sin, + max_seqlen=max_seqlen, + ) + + # adapter + x = self.post_layernorm(x) + x = x.view(-1, self.spatial_merge_size, self.spatial_merge_size, x.shape[-1]) + x = x.permute(0, 3, 1, 2) + x = self.downsample(x).view(-1, self.out_hidden_size) + x = self.merger(x) + return x + + def load_weights(self, weights) -> set[str]: + loader = AutoWeightsLoader(self) + return loader.load_weights(weights, mapper=self.hf_to_vllm_mapper) + + +class Glm5NextProcessingInfo(Glm4vProcessingInfo): + """Wires up the vLLM-native processor for the multimodal checkpoint. + + The checkpoint's ``processor_config.json`` declares a custom ``processor_class`` + and stores its image/video processor configs inline (no standalone + ``preprocessor_config.json``), so ``AutoProcessor`` cannot resolve the + config. We bypass it and build our own ``Glm5NextProcessor`` + (``vllm/transformers_utils/processors/glm5next.py``), a port of the + training-side pipeline that no longer imports transformers' GLM processor + classes. The port applies ``patch_expand_factor`` (checkpoint ships 1) + inside ``smart_resize``'s spatial factor. + """ + + @cached_property + def _glm5_hf_processor(self): + from vllm.transformers_utils.processors.glm5next import Glm5NextProcessor + + return Glm5NextProcessor.from_pretrained(self.ctx.model_config.model) + + def get_hf_processor(self, **kwargs: object): + return self._glm5_hf_processor + + def _processor_pixel_budget(self, proc) -> tuple[int, int]: + from vllm.transformers_utils.processors.glm5next import _pixel_budget + + return _pixel_budget( + proc.min_image_tokens, + proc.max_image_tokens, + proc.patch_size, + proc.merge_size, + proc.temporal_patch_size, + ) + + def _get_image_max_pixels(self) -> int: + mm_kwargs = self.ctx.get_merged_mm_kwargs({}) + if (override := mm_kwargs.get("max_pixels")) is not None: + return int(override) + return self._processor_pixel_budget(self.get_hf_processor().image_processor)[1] + + def _get_video_max_pixels(self) -> int: + mm_kwargs = self.ctx.get_merged_mm_kwargs({}) + if (override := mm_kwargs.get("max_pixels")) is not None: + return int(override) + return self._processor_pixel_budget(self.get_hf_processor().video_processor)[1] + + def _get_vision_info( + self, + *, + image_width: int, + image_height: int, + num_frames: int = 16, + do_resize: bool = True, + max_image_pixels: int = 28 * 28 * 2 * 30000, + ) -> tuple[ImageSize, int]: + """GLM-5.3-Flash canvas geometry for token budgeting and dummy inputs. + + The inherited Glm4v path resolves the pixel budget from + ``size.longest_edge`` and resizes with GLM-4V's ``smart_resize``. This + checkpoint's ``processor_config.json`` ships the token-budget style + (``min_image_tokens`` / ``max_image_tokens``) with no ``size`` key, and + the alignment factor carries ``patch_expand_factor`` — resolve both + from the vLLM-native processor so profiling matches runtime geometry. + """ + from vllm.transformers_utils.processors.glm5next import smart_resize + + vision_config = self.get_hf_config().vision_config + patch_size = vision_config.patch_size + merge_size = vision_config.spatial_merge_size + temporal_patch_size = vision_config.temporal_patch_size + + image_processor = self.get_hf_processor().image_processor + factor = patch_size * merge_size * image_processor.patch_expand_factor + # Keep the profiling search viable when the caller's budget is below + # one aligned canvas of the requested duration. + max_image_pixels = max(max_image_pixels, temporal_patch_size * factor * factor) + + if do_resize: + t = num_frames if num_frames > temporal_patch_size else temporal_patch_size + resized_height, resized_width = smart_resize( + t=t, + h=image_height, + w=image_width, + t_factor=temporal_patch_size, + h_factor=factor, + w_factor=factor, + min_pixels=1, + max_pixels=max_image_pixels, + ) + preprocessed_size = ImageSize(width=resized_width, height=resized_height) + else: + preprocessed_size = ImageSize(width=image_width, height=image_height) + + padded_num_frames = num_frames + (-num_frames % temporal_patch_size) + grid_t = max(padded_num_frames // temporal_patch_size, 1) + grid_h = preprocessed_size.height // patch_size + grid_w = preprocessed_size.width // patch_size + + num_patches = grid_t * grid_h * grid_w + num_vision_tokens = num_patches // (merge_size**2) + + return preprocessed_size, num_vision_tokens + + +class Glm5NextMultiModalProcessor(Glm4vMultiModalProcessor): + """The vLLM-native ``Glm5NextProcessor`` extracts image/video features + only and passes the prompt text through unchanged, so prompt expansion + (image token repeat, video frame/timestamp structure) is owned by vLLM's + prompt-update machinery — the inherited ``_get_prompt_updates`` builds + the replacement content and the placeholder scan validates against + exactly that.""" + + def _hf_processor_applies_updates( + self, + prompt_text: str, + mm_items: MultiModalDataItems, + hf_processor_mm_kwargs: Mapping[str, object], + tokenization_kwargs: Mapping[str, object], + ) -> bool: + return False diff --git a/vllm/models/glm5next/nvidia/ops/fused_eh_norm.py b/vllm/models/glm5next/nvidia/ops/fused_eh_norm.py new file mode 100644 index 000000000000..095db3917b49 --- /dev/null +++ b/vllm/models/glm5next/nvidia/ops/fused_eh_norm.py @@ -0,0 +1,77 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +import torch + +from vllm.triton_utils import tl, triton + + +@triton.jit +def _rms_norm(x, w, eps, HIDDEN_SIZE: tl.constexpr): + x = x.to(tl.float32) + mean_sq = tl.sum(x * x, axis=0) / HIDDEN_SIZE + rrms = tl.rsqrt(mean_sq + eps) + w = w.to(tl.float32) + return (x * rrms) * w + + +@triton.jit +def _fused_eh_norm_kernel( + pos_ptr, + embeds_ptr, + embeds_stride, + prev_ptr, + prev_stride, + enorm_w_ptr, + hnorm_w_ptr, + eps, + out_ptr, + out_stride, + H: tl.constexpr, + BLOCK: tl.constexpr, +): + """MTP input fusion: zero embeds at position 0, RMSNorm(embeds) with enorm + and RMSNorm(prev_hidden) with hnorm, written side-by-side into ``out`` + ([N, 2H]) ready for the eh_proj GEMM. Replaces where + 2x RMSNorm + cat.""" + tok = tl.program_id(0) + off = tl.arange(0, BLOCK) + mask = off < H + + pos = tl.load(pos_ptr + tok) + e = tl.load(embeds_ptr + tok * embeds_stride + off, mask=mask, other=0.0) + e = tl.where(pos == 0, 0.0, e.to(tl.float32)) + ew = tl.load(enorm_w_ptr + off, mask=mask) + e_normed = _rms_norm(e, ew, eps, H) + tl.store(out_ptr + tok * out_stride + off, e_normed, mask=mask) + + p = tl.load(prev_ptr + tok * prev_stride + off, mask=mask, other=0.0) + hw = tl.load(hnorm_w_ptr + off, mask=mask) + p_normed = _rms_norm(p, hw, eps, H) + tl.store(out_ptr + tok * out_stride + H + off, p_normed, mask=mask) + + +def fused_eh_norm( + positions: torch.Tensor, + inputs_embeds: torch.Tensor, + previous_hidden: torch.Tensor, + enorm_w: torch.Tensor, + hnorm_w: torch.Tensor, + eps: float, +) -> torch.Tensor: + """Returns cat([enorm(masked embeds), hnorm(prev_hidden)]) -> [N, 2H].""" + n, h = inputs_embeds.shape + out = torch.empty(n, 2 * h, dtype=inputs_embeds.dtype, device=inputs_embeds.device) + _fused_eh_norm_kernel[(n,)]( + positions, + inputs_embeds, + inputs_embeds.stride(0), + previous_hidden, + previous_hidden.stride(0), + enorm_w, + hnorm_w, + eps, + out, + out.stride(0), + h, + triton.next_power_of_2(h), + ) + return out diff --git a/vllm/models/glm5next/nvidia/ops/kpool_compress.py b/vllm/models/glm5next/nvidia/ops/kpool_compress.py new file mode 100644 index 000000000000..35b51c6b74bb --- /dev/null +++ b/vllm/models/glm5next/nvidia/ops/kpool_compress.py @@ -0,0 +1,891 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +"""kpool (key-pooling) Triton kernels for the sparse-attention indexer. + +The cache stores POOLS (1 entry per ``pool_size`` consecutive tokens) rather +than individual tokens. ``compress_ratio == pool_size`` on the kv_cache_spec +makes the metadata builder emit pool-granular slot_mapping / seq_lens / +cu_seq_lens / page_table for free; this file supplies the compress-write +kernel (replacing ``indexer_k_quant_and_cache``) and the pool-level topk +helpers (select pools -> expand to tokens -> append tail). +""" + +from __future__ import annotations + +import torch + +from vllm.triton_utils import tl, triton + +# The GLM-5.3-Flash indexer head dimension is fixed at 128. +INDEX_HEAD_DIM = 128 + + +# Hadamard-128 rotation + + +@triton.jit +def _hadamard128_stage(x, GROUPS: tl.constexpr, STRIDE: tl.constexpr): + x3 = tl.reshape(x, (GROUPS, 2, STRIDE)) + x3 = tl.trans(x3, 0, 2, 1) + a, b = tl.split(x3) + x3 = tl.join(a + b, a - b) + x3 = tl.trans(x3, 0, 2, 1) + return tl.reshape(x3, (128,)) + + +@triton.jit +def _hadamard128(x): + x = _hadamard128_stage(x, 64, 1) + x = _hadamard128_stage(x, 32, 2) + x = _hadamard128_stage(x, 16, 4) + x = _hadamard128_stage(x, 8, 8) + x = _hadamard128_stage(x, 4, 16) + x = _hadamard128_stage(x, 2, 32) + x = _hadamard128_stage(x, 1, 64) + return x * 0.08838834764831845 # 1/sqrt(128) + + +@triton.jit +def _fwht_stage(x, N: tl.constexpr, GROUPS: tl.constexpr, STRIDE: tl.constexpr): + # One FWHT butterfly stage on a flat tensor of N = GROUPS*2*STRIDE elems; + # same construction as _hadamard128_stage but with a parametric N so it can + # process BLOCK_R rows at once. + x3 = tl.reshape(x, (GROUPS, 2, STRIDE)) + x3 = tl.trans(x3, 0, 2, 1) + a, b = tl.split(x3) + x3 = tl.join(a + b, a - b) + x3 = tl.trans(x3, 0, 2, 1) + return tl.reshape(x3, (N,)) + + +@triton.jit +def _fwht_quant_kernel( + q_ptr, + qout_ptr, + sout_ptr, + n_rows, + BLOCK_R: tl.constexpr, +): + """Fused Hadamard-128 rotation + per-row absmax FP8 (ue8m0) quant. + + Each row uses fp32 butterflies and scaling, rounds to bf16, then applies + absmax quantization with a power-of-two scale. + """ + pid = tl.program_id(0) + rows = pid * BLOCK_R + tl.arange(0, BLOCK_R) + rmask = rows < n_rows + offs = tl.arange(0, 128) + x = tl.load( + q_ptr + rows[:, None] * 128 + offs[None, :], mask=rmask[:, None], other=0.0 + ).to(tl.float32) + + # Flatten so each row's 128 lanes stay contiguous: every stage's + # (GROUPS, 2, STRIDE) tiling has 2*STRIDE dividing 128, so pairs never + # straddle a row boundary. GROUPS of each stage scales by BLOCK_R. + N: tl.constexpr = BLOCK_R * 128 + x = tl.reshape(x, (N,)) + x = _fwht_stage(x, N, BLOCK_R * 64, 1) + x = _fwht_stage(x, N, BLOCK_R * 32, 2) + x = _fwht_stage(x, N, BLOCK_R * 16, 4) + x = _fwht_stage(x, N, BLOCK_R * 8, 8) + x = _fwht_stage(x, N, BLOCK_R * 4, 16) + x = _fwht_stage(x, N, BLOCK_R * 2, 32) + x = _fwht_stage(x, N, BLOCK_R, 64) + x = x * 0.08838834764831845 # 1/sqrt(128), exact in fp32 + + # Match the unfused path's bf16 materialization before quantizing. + x = x.to(tl.bfloat16).to(tl.float32) + x = tl.reshape(x, (BLOCK_R, 128)) + + absmax = tl.maximum(tl.max(tl.abs(x), axis=1), 1e-4) + scale = tl.exp2(tl.ceil(tl.log2(absmax * (1.0 / 448.0)))) + y = tl.minimum(tl.maximum(x / scale[:, None], -448.0), 448.0) + + tl.store(qout_ptr + rows[:, None] * 128 + offs[None, :], y, mask=rmask[:, None]) + tl.store(sout_ptr + rows, scale, mask=rmask) + + +def fwht128_quant_fp8(q: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: + """Rotate each 128-wide row by the Hadamard-128 transform, then FP8-quant. + + The fused kernel avoids materializing the rotated tensor and uses an exact + fp32 ``1 / sqrt(128)`` scale before block-128 ue8m0 quantization. + + Args: + q: ``[rows, 128]`` bf16 — one head vector per row. + + Returns: + (q_fp8 ``[rows, 128]`` float8_e4m3fn, scale ``[rows, 1]`` float32). + """ + assert q.ndim == 2 and q.shape[1] == 128, q.shape + assert q.dtype == torch.bfloat16 + assert q.is_contiguous() + n_rows = q.shape[0] + q_fp8 = torch.empty((n_rows, 128), dtype=torch.float8_e4m3fn, device=q.device) + q_scale = torch.empty((n_rows, 1), dtype=torch.float32, device=q.device) + if n_rows == 0: + return q_fp8, q_scale + BLOCK_R = 32 + grid = (triton.cdiv(n_rows, BLOCK_R),) + _fwht_quant_kernel[grid](q, q_fp8, q_scale, n_rows, BLOCK_R=BLOCK_R, num_warps=2) + return q_fp8, q_scale + + +# Fused pool compression and cache write. + + +@triton.jit +def _kpool_softmax_rotate_write_cache_kernel( + buf_fp8_ptr, + buf_fp32_ptr, + slot_k_ptr, + slot_score_ptr, + ape_ptr, + loc_ptr, + write_mask_ptr, + compressed_k_ptr, + compressed_scale_ptr, + slot_k_stride_0, + slot_k_stride_1, + slot_score_stride_0, + slot_score_stride_1, + ape_stride_0, + PAGE_SIZE: tl.constexpr, + BUF_NUMEL_PER_PAGE: tl.constexpr, + POOL_SIZE: tl.constexpr, + HEAD_DIM: tl.constexpr, + S_OFFSET_NBYTES_IN_PAGE: tl.constexpr, + ROUND_SCALE: tl.constexpr, + HAS_WRITE_MASK: tl.constexpr, + RETURN_COMPRESSED: tl.constexpr, + WRITE_CACHE: tl.constexpr, + BLOCK_D: tl.constexpr, +): + """One program per pool. softmax(slot_score+ape)-weighted sum of slot_k -> + Hadamard-128 -> per-vector fp8 absmax quant -> write to cache at ``loc``.""" + row = tl.program_id(0) + do_write = True + if HAS_WRITE_MASK: + do_write = tl.load(write_mask_ptr + row) + + offs = tl.arange(0, BLOCK_D) + mask = (offs < HEAD_DIM) & do_write + + # --- Pass 1: per-dim max over the pool (softmax numerical stability) --- + max_score = tl.full((BLOCK_D,), -float("inf"), tl.float32) + for slot in tl.static_range(0, POOL_SIZE): + score = tl.load( + slot_score_ptr + + row * slot_score_stride_0 + + slot * slot_score_stride_1 + + offs, + mask=mask, + other=0.0, + ).to(tl.float32) + score += tl.load(ape_ptr + slot * ape_stride_0 + offs, mask=mask, other=0.0).to( + tl.float32 + ) + max_score = tl.maximum(max_score, score) + + # --- Pass 2: softmax-weighted sum of K --- + acc = tl.full((BLOCK_D,), 0.0, tl.float32) + denom = tl.full((BLOCK_D,), 0.0, tl.float32) + for slot in tl.static_range(0, POOL_SIZE): + score = tl.load( + slot_score_ptr + + row * slot_score_stride_0 + + slot * slot_score_stride_1 + + offs, + mask=mask, + other=0.0, + ).to(tl.float32) + score += tl.load(ape_ptr + slot * ape_stride_0 + offs, mask=mask, other=0.0).to( + tl.float32 + ) + prob = tl.exp(score - max_score) + denom += prob + k = tl.load( + slot_k_ptr + row * slot_k_stride_0 + slot * slot_k_stride_1 + offs, + mask=mask, + other=0.0, + ).to(tl.float32) + acc += k * prob + + x = acc / denom + x = tl.where(do_write, x, 0.0).to(tl.bfloat16).to(tl.float32) + + # Match the unfused pooled-K path's bf16 precision before quantization. + x = _hadamard128(x).to(tl.bfloat16).to(tl.float32) + + # --- per-vector absmax fp8 quant --- + fp8_max = 448.0 + fp8_max_inv = 1.0 / fp8_max + absmax = tl.max(tl.abs(x), axis=0) + absmax = tl.maximum(absmax, 1e-4) + if ROUND_SCALE: + scale = tl.exp2(tl.ceil(tl.log2(absmax * fp8_max_inv))) + else: + scale = absmax * fp8_max_inv + quantized = x / scale + quantized = tl.minimum(tl.maximum(quantized, -fp8_max), fp8_max) + + if WRITE_CACHE: + loc = tl.load(loc_ptr + row, mask=do_write, other=0) + loc_page_index = loc // PAGE_SIZE + loc_token_offset_in_page = loc % PAGE_SIZE + out_k_offsets = ( + loc_page_index * BUF_NUMEL_PER_PAGE + + loc_token_offset_in_page * HEAD_DIM + + offs + ) + out_s_offset = ( + loc_page_index * BUF_NUMEL_PER_PAGE // 4 + + S_OFFSET_NBYTES_IN_PAGE // 4 + + loc_token_offset_in_page + ) + tl.store(buf_fp8_ptr + out_k_offsets, quantized, mask=mask) + tl.store(buf_fp32_ptr + out_s_offset, scale, mask=do_write) + + if RETURN_COMPRESSED: + tl.store( + compressed_k_ptr + row * HEAD_DIM + offs, + quantized, + mask=offs < HEAD_DIM, + ) + tl.store(compressed_scale_ptr + row, scale) + + +def kpool_compress_and_write_cache( + kv_cache: torch.Tensor, + slot_k: torch.Tensor, + slot_score: torch.Tensor, + ape: torch.Tensor, + loc: torch.Tensor, + pool_size: int, + head_dim: int = INDEX_HEAD_DIM, + write_mask: torch.Tensor | None = None, + round_scale: bool = True, + return_compressed: bool = False, + write_cache: bool = True, +): + """Compress ``pool_size`` tokens into one fp8 K and write at ``loc``. + + Args: + kv_cache: indexer K cache ``[num_blocks, block_size, head_dim+4]`` uint8. + slot_k: ``[n_pools, pool_size, head_dim]`` bf16 — raw per-token K. + slot_score: ``[n_pools, pool_size, head_dim]`` — per-token gate score. + ape: ``[pool_size, head_dim]`` fp32 — per-slot position bias. + loc: ``[n_pools]`` int64 — flat physical slot per pool. + """ + assert slot_k.ndim == 3 + assert slot_score.shape == slot_k.shape + assert ape.shape == slot_k.shape[1:] + assert slot_k.shape[2] == head_dim + assert slot_k.dtype == torch.bfloat16 + assert ape.dtype == torch.float32 + assert kv_cache.dtype == torch.uint8 + assert loc.dtype == torch.int64 + assert write_cache or return_compressed + + page_size = kv_cache.shape[1] + buf = kv_cache + slot_k = slot_k.contiguous() + slot_score = slot_score.contiguous() + ape = ape.contiguous() + loc = loc.contiguous() + if write_mask is None: + write_mask = torch.empty((1,), dtype=torch.bool, device=slot_k.device) + has_write_mask = False + else: + assert write_mask.shape == (slot_k.shape[0],) + write_mask = write_mask.contiguous() + has_write_mask = True + assert not return_compressed + + if slot_k.shape[0] == 0: + if return_compressed: + return ( + torch.empty( + (0, head_dim), + dtype=torch.float8_e4m3fn, + device=slot_k.device, + ), + torch.empty((0,), dtype=torch.float32, device=slot_k.device), + ) + return None + + buf_fp8 = buf.view(torch.float8_e4m3fn) + buf_fp32 = buf.view(torch.float32) + # bytes per page (last dim of kv_cache) viewed as uint8 + buf_numel_per_page = buf.stride(0) + s_offset_nbytes_in_page = page_size * head_dim + + if return_compressed: + compressed_k = torch.empty( + (slot_k.shape[0], head_dim), + dtype=torch.float8_e4m3fn, + device=slot_k.device, + ) + compressed_scale = torch.empty( + (slot_k.shape[0],), dtype=torch.float32, device=slot_k.device + ) + else: + compressed_k = buf_fp8 + compressed_scale = buf_fp32 + + _kpool_softmax_rotate_write_cache_kernel[(slot_k.shape[0],)]( + buf_fp8, + buf_fp32, + slot_k, + slot_score, + ape, + loc, + write_mask, + compressed_k, + compressed_scale, + slot_k.stride(0), + slot_k.stride(1), + slot_score.stride(0), + slot_score.stride(1), + ape.stride(0), + PAGE_SIZE=page_size, + BUF_NUMEL_PER_PAGE=buf_numel_per_page, + POOL_SIZE=slot_k.shape[1], + HEAD_DIM=head_dim, + S_OFFSET_NBYTES_IN_PAGE=s_offset_nbytes_in_page, + ROUND_SCALE=round_scale, + HAS_WRITE_MASK=has_write_mask, + RETURN_COMPRESSED=return_compressed, + WRITE_CACHE=write_cache, + BLOCK_D=triton.next_power_of_2(head_dim), + ) + + if return_compressed: + return compressed_k, compressed_scale + return None + + +# Seed each request's incomplete pool into its paged tail during prefill. + + +@triton.jit +def _kpool_tail_seed_kernel( + key_ptr, + score_ptr, + tslot_ptr, + tail_ptr, + n_tokens, + HEAD_DIM: tl.constexpr, + KPOOL: tl.constexpr, + BLOCK_D: tl.constexpr, +): + """Copy token ``i``'s raw K + gate into its request's tail block. + + Token ``i`` is among its request's last KPOOL tokens iff the token KPOOL + ahead belongs to a different tail block (or is past the batch / padding, + slot < 0). ``tslot = block * KPOOL + pos % KPOOL``; the destination is + ``tail[block, {0:K, 1:score}, pos % KPOOL, :]``. + """ + i = tl.program_id(0) + t = tl.load(tslot_ptr + i).to(tl.int64) + if t < 0: + return + blk = t // KPOOL # t >= 0 here, so trunc == floor + ahead = tl.load(tslot_ptr + i + KPOOL, mask=i + KPOOL < n_tokens, other=-1).to( + tl.int64 + ) + # Match the torch semantics exactly: a negative ahead slot floors to a + # block id that differs from every real block -> token is in the tail. + # Only divide non-negative slots (Triton int div truncates, torch floors). + if ahead >= 0 and ahead // KPOOL == blk: + return + offs = tl.arange(0, BLOCK_D) + m = offs < HEAD_DIM + base = (blk * 2 * KPOOL + t % KPOOL) * HEAD_DIM + k = tl.load(key_ptr + i * HEAD_DIM + offs, mask=m) + s = tl.load(score_ptr + i * HEAD_DIM + offs, mask=m) + tl.store(tail_ptr + base + offs, k, mask=m) + tl.store(tail_ptr + base + KPOOL * HEAD_DIM + offs, s, mask=m) + + +def kpool_seed_tail_cache( + tail_kv_cache: torch.Tensor, + key: torch.Tensor, + gate_score: torch.Tensor, + tslot: torch.Tensor, + kpool: int, + head_dim: int = INDEX_HEAD_DIM, +) -> None: + """Seed the paged tail cache from a prefill batch (see the kernel).""" + assert tail_kv_cache.dtype == torch.bfloat16 + assert key.dtype == torch.bfloat16 + n = tslot.shape[0] + if n == 0: + return + _kpool_tail_seed_kernel[(n,)]( + key, + gate_score, + tslot, + tail_kv_cache, + n, + HEAD_DIM=head_dim, + KPOOL=kpool, + BLOCK_D=triton.next_power_of_2(head_dim), + ) + + +# Update each request's tail during decode and write completed pools. + + +@triton.jit +def _kpool_decode_update_batched_kernel( + buf_fp8_ptr, + buf_fp32_ptr, + tail_kv_ptr, + tail_slot_mapping_ptr, # [B, NEXT_N] int32 + key_ptr, # [B, NEXT_N, HEAD_DIM] bf16 + key_stride_b, + key_stride_t, + slot_score_ptr, # [B, NEXT_N, HEAD_DIM] bf16 + ss_stride_b, + ss_stride_t, + ape_ptr, + ape_stride_0, + slot_mapping_ptr, # [B, NEXT_N] int32 + positions_ptr, # [B, NEXT_N] int32 + NEXT_N, # runtime token count per request (no .item() needed) + PAGE_SIZE: tl.constexpr, + BUF_NUMEL_PER_PAGE: tl.constexpr, + POOL_SIZE: tl.constexpr, + TAIL_BLOCK_ELEMS: tl.constexpr, + KPOOL_HEAD: tl.constexpr, + HEAD_DIM: tl.constexpr, + S_OFFSET_NBYTES_IN_PAGE: tl.constexpr, + ROUND_SCALE: tl.constexpr, + BLOCK_D: tl.constexpr, +): + """One program per request; iterates its NEXT_N verify tokens in order. + + Replaces the caller's per-token sequential launch loop. The intra-request + iteration MUST stay in position order: a pool-completion at token t* reads + the tail-ring slots that tokens t < t* (same request) just stashed in this + same invocation. ``tl.range`` iterates sequentially within the program, so + those stashes are visible to the later completion read. Cross-request + programs are independent (distinct tail blocks). With NEXT_N < POOL_SIZE + (the spec-verify case: NEXT_N ~= num_spec+1, POOL_SIZE=16) at most one + completion can occur per request per call, but the ordered loop is correct + for any NEXT_N. + """ + req = tl.program_id(0) + offs = tl.arange(0, BLOCK_D) + dim_mask = offs < HEAD_DIM + + for t in tl.range(0, NEXT_N): + idx = req * NEXT_N + t + cache_loc = tl.load(slot_mapping_ptr + idx) + pos = tl.load(positions_ptr + idx) + safe_pos = tl.maximum(pos, 0) + pos_valid = (cache_loc >= 0) & (pos >= 0) + + slot = safe_pos % POOL_SIZE + phys_slot = safe_pos % POOL_SIZE + + # Derive the tail block from THIS token's tail_slot (the request's block + # is constant across a pool, but a padded / invalid entry carries a + # negative sentinel -- reading it from token 0 would poison every + # token's base address). Clamp so an invalid entry can never form an + # out-of-bounds base; the accesses below are gated on pos_valid anyway. + tail_slot = tl.load(tail_slot_mapping_ptr + idx) + block = tl.maximum(tail_slot, 0).to(tl.int64) // POOL_SIZE + block_base = block * TAIL_BLOCK_ELEMS + + # The tail-ring stash must run for EVERY real token, so it is gated on + # the token-granular tail slot -- not on `pos_valid`, which keys off the + # POOL-granular `slot_mapping` and is therefore only true on the pool's + # last token. Gating the stash on pos_valid dropped every intra-pool + # token, so a decode-built pool compressed 3 stale ring entries (the + # prefill-seeded prompt tail, frozen forever) plus the current token. + stash_valid = (pos >= 0) & (tail_slot >= 0) + + key = tl.load( + key_ptr + req * key_stride_b + t * key_stride_t + offs, + mask=dim_mask, + other=0.0, + ).to(tl.float32) + score_current = tl.load( + slot_score_ptr + req * ss_stride_b + t * ss_stride_t + offs, + mask=dim_mask, + other=0.0, + ).to(tl.float32) + + if pos_valid & (slot == POOL_SIZE - 1): + pool_logical_start = safe_pos - slot + + max_score = tl.full((BLOCK_D,), -float("inf"), tl.float32) + for pool_slot in tl.static_range(0, POOL_SIZE): + is_current = pool_slot == slot + phys = (pool_logical_start + pool_slot) % POOL_SIZE + score_buf = tl.load( + tail_kv_ptr + block_base + KPOOL_HEAD + phys * HEAD_DIM + offs, + mask=dim_mask, + other=0.0, + ).to(tl.float32) + score = tl.where(is_current, score_current, score_buf) + score += tl.load( + ape_ptr + pool_slot * ape_stride_0 + offs, + mask=dim_mask, + other=0.0, + ).to(tl.float32) + max_score = tl.maximum(max_score, score) + + acc = tl.full((BLOCK_D,), 0.0, tl.float32) + denom = tl.full((BLOCK_D,), 0.0, tl.float32) + for pool_slot in tl.static_range(0, POOL_SIZE): + is_current = pool_slot == slot + phys = (pool_logical_start + pool_slot) % POOL_SIZE + score_buf = tl.load( + tail_kv_ptr + block_base + KPOOL_HEAD + phys * HEAD_DIM + offs, + mask=dim_mask, + other=0.0, + ).to(tl.float32) + score = tl.where(is_current, score_current, score_buf) + score += tl.load( + ape_ptr + pool_slot * ape_stride_0 + offs, + mask=dim_mask, + other=0.0, + ).to(tl.float32) + prob = tl.exp(score - max_score) + denom += prob + k_buf = tl.load( + tail_kv_ptr + block_base + phys * HEAD_DIM + offs, + mask=dim_mask, + other=0.0, + ).to(tl.float32) + k = tl.where(is_current, key, k_buf) + acc += k * prob + + x = (acc / denom).to(tl.bfloat16).to(tl.float32) + x = _hadamard128(x).to(tl.bfloat16).to(tl.float32) + + fp8_max = 448.0 + fp8_max_inv = 1.0 / fp8_max + absmax = tl.maximum(tl.max(tl.abs(x), axis=0), 1e-4) + if ROUND_SCALE: + scale = tl.exp2(tl.ceil(tl.log2(absmax * fp8_max_inv))) + else: + scale = absmax * fp8_max_inv + quantized = tl.minimum(tl.maximum(x / scale, -fp8_max), fp8_max) + + loc = cache_loc.to(tl.int64) + loc_page_index = loc // PAGE_SIZE + loc_token_offset_in_page = loc % PAGE_SIZE + out_k_offsets = ( + loc_page_index * BUF_NUMEL_PER_PAGE + + loc_token_offset_in_page * HEAD_DIM + + offs + ) + out_s_offset = ( + loc_page_index * BUF_NUMEL_PER_PAGE // 4 + + S_OFFSET_NBYTES_IN_PAGE // 4 + + loc_token_offset_in_page + ) + tl.store(buf_fp8_ptr + out_k_offsets, quantized, mask=dim_mask) + tl.store(buf_fp32_ptr + out_s_offset, scale) + + # Stash the current token AFTER any completion read so the completion + # uses prior stashes (and the current token's own key/score via + # is_current), then leaves this token for future pools. Order matches + # the per-token kernel: completion read first, stash second. + update_mask = dim_mask & stash_valid + tl.store( + tail_kv_ptr + block_base + phys_slot * HEAD_DIM + offs, + key, + mask=update_mask, + ) + tl.store( + tail_kv_ptr + block_base + KPOOL_HEAD + phys_slot * HEAD_DIM + offs, + score_current, + mask=update_mask, + ) + + +def kpool_decode_update_and_maybe_write_cache_batched( + kv_cache: torch.Tensor, + tail_kv_cache: torch.Tensor, + tail_slot_mapping: torch.Tensor, + key: torch.Tensor, + slot_score: torch.Tensor, + ape: torch.Tensor, + slot_mapping: torch.Tensor, + positions: torch.Tensor, + pool_size: int, + head_dim: int = INDEX_HEAD_DIM, + round_scale: bool = True, +) -> None: + """Batched decode-step kpool update for spec verify (``next_n > 1``). + + One launch replaces the caller's per-token loop. Inputs are grouped per + request: ``[num_requests, next_n, ...]``. Each program handles one + request's ``next_n`` tokens in position order (see the kernel docstring for + why ordering is required for pool-completion correctness). + + Plain decode (``next_n == 1``) is handled here too — the kernel collapses + to a single-iteration loop. + + Args: + kv_cache: indexer K cache ``[num_blocks, block_size, head_dim+4]`` uint8. + tail_kv_cache: paged tail cache ``[num_blocks, 2, pool_size, head_dim]`` + bf16 (K at half 0, gate score at half 1). + tail_slot_mapping: ``[num_requests, next_n]`` int32. + key / slot_score: ``[num_requests, next_n, head_dim]`` bf16. + ape: ``[pool_size, head_dim]`` fp32. + slot_mapping / positions: ``[num_requests, next_n]`` int32. + """ + num_requests, next_n = key.shape[0], key.shape[1] + if num_requests == 0 or next_n == 0: + return + assert tail_kv_cache.ndim == 4 + assert tail_kv_cache.shape[1] == 2 + assert tail_kv_cache.shape[2] == pool_size + assert tail_kv_cache.shape[3] == head_dim + assert tail_kv_cache.dtype == torch.bfloat16 + assert key.ndim == 3 and key.shape[2] == head_dim + assert slot_score.shape == key.shape + assert ape.shape == (pool_size, head_dim) + assert tail_slot_mapping.shape == (num_requests, next_n) + assert slot_mapping.shape == (num_requests, next_n) + assert positions.shape == (num_requests, next_n) + assert key.dtype == torch.bfloat16 + assert slot_score.dtype == torch.bfloat16 + assert ape.dtype == torch.float32 + assert kv_cache.dtype == torch.uint8 + + page_size = kv_cache.shape[1] + buf = kv_cache + buf_fp8 = buf.view(torch.float8_e4m3fn) + buf_fp32 = buf.view(torch.float32) + + # The kernel indexes the int tensors as ``req * next_n + t`` (row-major), + # so they must be contiguous. Callers pass either a view of a contiguous + # slice or a freshly scattered tensor, making these no-ops; the calls guard + # against a future caller handing over a strided view. + tail_slot_mapping = tail_slot_mapping.contiguous() + slot_mapping = slot_mapping.contiguous() + positions = positions.contiguous() + + _kpool_decode_update_batched_kernel[(num_requests,)]( + buf_fp8, + buf_fp32, + tail_kv_cache, + tail_slot_mapping, + key, + key.stride(0), + key.stride(1), + slot_score, + slot_score.stride(0), + slot_score.stride(1), + ape, + ape.stride(0), + slot_mapping, + positions, + next_n, + PAGE_SIZE=page_size, + BUF_NUMEL_PER_PAGE=buf.stride(0), + POOL_SIZE=pool_size, + TAIL_BLOCK_ELEMS=tail_kv_cache.stride(0), + KPOOL_HEAD=tail_kv_cache.stride(1), + HEAD_DIM=head_dim, + S_OFFSET_NBYTES_IN_PAGE=page_size * head_dim, + ROUND_SCALE=round_scale, + BLOCK_D=triton.next_power_of_2(head_dim), + ) + + +# Pool-level top-k helpers. + + +def history_group_budget_for_topk(topk: int, pool_size: int) -> int: + """Number of pools to select so that expanding yields ``topk`` tokens.""" + assert topk % pool_size == 0 + return topk // pool_size + + +def expand_pools_to_tokens( + group_ids: torch.Tensor, + group_valid: torch.Tensor, + topk: int, + pool_size: int, + page_table: torch.Tensor | None = None, + topk_offsets: torch.Tensor | None = None, +) -> torch.Tensor: + """Expand selected full-pool ids to a strict-width token topk tensor.""" + assert group_ids.ndim == 2 + assert group_valid.shape == group_ids.shape + assert topk % pool_size == 0 + assert group_ids.shape[1] == history_group_budget_for_topk(topk, pool_size) + assert page_table is None or topk_offsets is None + + device = group_ids.device + offsets = torch.arange(pool_size, device=device, dtype=torch.int64) + token_ids = group_ids.to(torch.int64).unsqueeze(-1) * pool_size + offsets + token_ids = token_ids.reshape(group_ids.shape[0], topk) + valid = ( + group_valid.unsqueeze(-1) + .expand(-1, -1, pool_size) + .reshape(group_ids.shape[0], topk) + ) + + if page_table is not None: + assert page_table.ndim == 2 + safe_ids = token_ids.clamp(min=0, max=page_table.shape[1] - 1) + output = torch.gather(page_table, dim=1, index=safe_ids).to(torch.int32) + elif topk_offsets is not None: + if topk_offsets.ndim == 2: + assert topk_offsets.shape[1] == 1 + topk_offsets = topk_offsets.squeeze(1) + output = (token_ids + topk_offsets.to(torch.int64).unsqueeze(1)).to(torch.int32) + else: + output = token_ids.to(torch.int32) + + return torch.where(valid, output, torch.full_like(output, -1)) + + +def append_tail_to_topk( + topk_result: torch.Tensor, + seq_lens: torch.Tensor, + pool_lens: torch.Tensor, + pool_size: int, + page_table: torch.Tensor | None = None, + topk_offsets: torch.Tensor | None = None, +) -> torch.Tensor: + """Append non-pooled tail tokens after expanded history tokens. + + ``index_kpool_always_select_tail`` keeps the (incomplete) trailing pool so + the most recent tokens are always attended to. + """ + assert topk_result.dtype == torch.int32 + assert seq_lens.ndim == 1 + assert pool_lens.ndim == 1 + + tail_pool = pool_size - 1 + if tail_pool == 0: + return topk_result + + rows, n_cols = topk_result.shape + out_cols = n_cols + tail_pool + out = torch.empty( + (rows, out_cols), dtype=topk_result.dtype, device=topk_result.device + ) + + # tail tokens: [pool_len*pool_size, seq_len) for each row. + pool_len = pool_lens.to(torch.int32) + tail_start = pool_len * pool_size + seq_len = seq_lens.to(torch.int32) + tail_count = seq_len - tail_start # in [0, pool_size) + + cols = torch.arange(out_cols, device=topk_result.device)[None, :] + history_len = n_cols + is_history = cols < history_len + tail_off = cols - history_len + is_tail = (tail_off >= 0) & (tail_off < tail_count[:, None]) + + # safe_hist must be per-row [rows, out_cols] so the gather reads each row's + # OWN history. cols is [1, out_cols]; if used directly, gather (which does + # NOT broadcast the index) would read only row 0 of topk_result, making every + # query inherit row 0's history (empty for the first token) and lose all its + # selected tokens — only the per-row tail would survive. This only manifests + # for multi-row sparse PREFILL (decode has 1 row, so it reads its own row 0). + safe_hist = torch.minimum(cols, torch.full_like(cols, n_cols - 1)).expand( + rows, out_cols + ) + history_val = torch.gather(topk_result, 1, safe_hist) + + tail_raw = tail_start[:, None] + tail_off + tail_val = tail_raw.to(torch.int32) + if page_table is not None: + safe_tail = tail_raw.clamp(min=0, max=page_table.shape[1] - 1) + tail_val = torch.gather(page_table, 1, safe_tail).to(torch.int32) + elif topk_offsets is not None: + tail_val = (tail_raw + topk_offsets.to(torch.int64).unsqueeze(1)).to( + torch.int32 + ) + + out = torch.where(is_history, history_val, -1) + out = torch.where(is_tail, tail_val, out) + return out + + +@triton.jit +def _expand_pools_and_append_tail_kernel( + pool_ids_ptr, # [rows, n_groups], int (any int dtype) + seq_lens_ptr, # [rows], int32 (token-granular seq_len) + out_ptr, # [rows, out_cols], int32 + topk, # n_groups * pool_size + out_cols, # topk + pool_size - 1 + POOL_SIZE: tl.constexpr, + BLOCK_COLS: tl.constexpr, + pid_s0, + out_s0, +): + # Fuses expand_pools_to_tokens + append_tail_to_topk (identity path) into a + # single kernel. Each program writes one (row, column-tile) of the output. + row = tl.program_id(0) + tile = tl.program_id(1) + cols = tile * BLOCK_COLS + tl.arange(0, BLOCK_COLS) + mask = cols < out_cols + + seq_len = tl.load(seq_lens_ptr + row) + pool_len = seq_len // POOL_SIZE + tail_start = pool_len * POOL_SIZE + tail_count = seq_len - tail_start # in [0, POOL_SIZE) + + # History region [0, topk): expand selected pool g = cols // POOL_SIZE. + is_history = cols < topk + g = cols // POOL_SIZE + o = cols % POOL_SIZE + pid = tl.load(pool_ids_ptr + row * pid_s0 + g, mask=mask & is_history, other=-1) + hist_val = (pid * POOL_SIZE + o).to(tl.int32) + hist_out = tl.where(pid >= 0, hist_val, -1) + + # Tail region [topk, out_cols): the request's trailing incomplete pool. + tail_off = cols - topk + is_tail = (tail_off >= 0) & (tail_off < tail_count) + tail_val = (tail_start + tail_off).to(tl.int32) + tail_out = tl.where(is_tail, tail_val, -1) + + result = tl.where(is_history, hist_out, tail_out) + tl.store(out_ptr + row * out_s0 + cols, result, mask=mask) + + +def expand_pools_and_append_tail( + pool_ids: torch.Tensor, + seq_lens: torch.Tensor, + pool_size: int, +) -> torch.Tensor: + """Fuse ``expand_pools_to_tokens`` + ``append_tail_to_topk`` (identity path). + + Produces the same ``[rows, topk + pool_size - 1]`` int32 output as calling + the two functions in sequence when neither ``page_table`` nor + ``topk_offsets`` is passed — the only path used by the GLM-5.3-Flash indexer. + The kernel derives ``pool_len = seq_len // pool_size`` internally, so the + caller no longer needs to precompute it. Replaces ~25 elementwise kernels + with one Triton launch. + """ + rows, n_groups = pool_ids.shape + topk = n_groups * pool_size + out_cols = topk + pool_size - 1 + out = torch.empty((rows, out_cols), dtype=torch.int32, device=pool_ids.device) + BLOCK_COLS = 128 + n_tiles = triton.cdiv(out_cols, BLOCK_COLS) + _expand_pools_and_append_tail_kernel[(rows, n_tiles)]( + pool_ids, + seq_lens, + out, + topk, + out_cols, + POOL_SIZE=pool_size, + BLOCK_COLS=BLOCK_COLS, + pid_s0=pool_ids.stride(0), + out_s0=out.stride(0), + ) + return out diff --git a/vllm/models/glm5next/nvidia/ops/third_party/__init__.py b/vllm/models/glm5next/nvidia/ops/third_party/__init__.py new file mode 100644 index 000000000000..208f01a7cb5e --- /dev/null +++ b/vllm/models/glm5next/nvidia/ops/third_party/__init__.py @@ -0,0 +1,2 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project diff --git a/vllm/models/glm5next/nvidia/ops/third_party/kda/__init__.py b/vllm/models/glm5next/nvidia/ops/third_party/kda/__init__.py new file mode 100644 index 000000000000..a073bd8021d4 --- /dev/null +++ b/vllm/models/glm5next/nvidia/ops/third_party/kda/__init__.py @@ -0,0 +1,6 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project + +from .kernels import chunk_kda_with_fused_gate, fused_recurrent_kda + +__all__ = ["chunk_kda_with_fused_gate", "fused_recurrent_kda"] diff --git a/vllm/models/glm5next/nvidia/ops/third_party/kda/fused_recurrent.py b/vllm/models/glm5next/nvidia/ops/third_party/kda/fused_recurrent.py new file mode 100644 index 000000000000..d7de85984f85 --- /dev/null +++ b/vllm/models/glm5next/nvidia/ops/third_party/kda/fused_recurrent.py @@ -0,0 +1,656 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +# SPDX-FileCopyrightText: Songlin Yang, Yu Zhang +# mypy: ignore-errors +# +# This file contains code copied from the flash-linear-attention project. +# The original source code was licensed under the MIT license and included +# the following copyright notice: +# Copyright (c) 2023-2025, Songlin Yang, Yu Zhang +# ruff: noqa: E501 + +import torch + +from vllm.third_party.flash_linear_attention.ops.op import exp +from vllm.triton_utils import tl, triton + + +@triton.heuristics( + { + "USE_INITIAL_STATE": lambda args: args["h0"] is not None, + "IS_VARLEN": lambda args: args["cu_seqlens"] is not None, + "IS_CONTINUOUS_BATCHING": lambda args: args["ssm_state_indices"] is not None, + "IS_SPEC_DECODING": lambda args: args["num_accepted_tokens"] is not None, + } +) +@triton.jit(do_not_specialize=["N", "T"]) +def fused_recurrent_gated_delta_rule_fwd_kernel( + q, + k, + v, + g, + beta, + o, + h0, + ht, + cu_seqlens, + ssm_state_indices, + num_accepted_tokens, + a_log, + g_bias, + scale, + N: tl.int64, # num of sequences + T: tl.int64, # num of tokens + B: tl.constexpr, + H: tl.constexpr, + HV: tl.constexpr, + K: tl.constexpr, + V: tl.constexpr, + BK: tl.constexpr, + BV: tl.constexpr, + stride_init_state_token: tl.constexpr, + stride_final_state_token: tl.constexpr, + stride_indices_seq: tl.constexpr, + stride_indices_tok: tl.constexpr, + USE_INITIAL_STATE: tl.constexpr, # whether to use initial state + INPLACE_FINAL_STATE: tl.constexpr, # whether to store final state inplace + IS_BETA_HEADWISE: tl.constexpr, # whether beta is headwise vector or scalar, + USE_QK_L2NORM_IN_KERNEL: tl.constexpr, + IS_VARLEN: tl.constexpr, + IS_CONTINUOUS_BATCHING: tl.constexpr, + IS_SPEC_DECODING: tl.constexpr, + IS_KDA: tl.constexpr, + SIGMOID_BETA: tl.constexpr, # beta holds raw logits; sigmoid at fp32 load + COMPUTE_GATE: tl.constexpr, # g holds raw logits; KDA gate computed in-kernel + SAFE_GATE: tl.constexpr, # bounded gate variant (only branch implemented) + LOWER_BOUND: tl.constexpr, +): + i_k, i_v, i_nh = tl.program_id(0), tl.program_id(1), tl.program_id(2) + i_n, i_hv = i_nh // HV, i_nh % HV + i_h = i_hv // (HV // H) + if IS_VARLEN: + bos, eos = ( + tl.load(cu_seqlens + i_n).to(tl.int64), + tl.load(cu_seqlens + i_n + 1).to(tl.int64), + ) + all = T + T = eos - bos + else: + bos, eos = i_n * T, i_n * T + T + all = B * T + + if T == 0: + # no tokens to process for this sequence + return + + o_k = i_k * BK + tl.arange(0, BK) + o_v = i_v * BV + tl.arange(0, BV) + + p_q = q + (bos * H + i_h) * K + o_k + p_k = k + (bos * H + i_h) * K + o_k + p_v = v + (bos * HV + i_hv) * V + o_v + if IS_BETA_HEADWISE: + p_beta = beta + (bos * HV + i_hv) * V + o_v + else: + p_beta = beta + bos * HV + i_hv + + if not IS_KDA: + p_g = g + bos * HV + i_hv + else: + p_gk = g + (bos * HV + i_hv) * K + o_k + + # Per-head gate amplitude, hoisted out of the token loop (COMPUTE_GATE). + if COMPUTE_GATE: + b_a_log = tl.exp(tl.load(a_log + i_h).to(tl.float32)) + + p_o = o + ((i_k * all + bos) * HV + i_hv) * V + o_v + + mask_k = o_k < K + mask_v = o_v < V + mask_h = mask_v[:, None] & mask_k[None, :] + + b_h = tl.zeros([BV, BK], dtype=tl.float32) + if USE_INITIAL_STATE: + if IS_CONTINUOUS_BATCHING: + if IS_SPEC_DECODING: + i_t = tl.load(num_accepted_tokens + i_n).to(tl.int64) - 1 + else: + i_t = 0 + # Load state index and check for invalid entries + state_idx = tl.load(ssm_state_indices + i_n * stride_indices_seq + i_t).to( + tl.int64 + ) + # Skip if state index is invalid (NULL_BLOCK_ID=0) + if state_idx <= 0: + return + p_h0 = h0 + state_idx * stride_init_state_token + else: + p_h0 = h0 + bos * HV * V * K + p_h0 = p_h0 + i_hv * V * K + o_v[:, None] * K + o_k[None, :] + b_h += tl.load(p_h0, mask=mask_h, other=0).to(tl.float32) + + for i_t in range(0, T): + b_q = tl.load(p_q, mask=mask_k, other=0).to(tl.float32) + b_k = tl.load(p_k, mask=mask_k, other=0).to(tl.float32) + b_v = tl.load(p_v, mask=mask_v, other=0).to(tl.float32) + + if USE_QK_L2NORM_IN_KERNEL: + b_q = b_q / tl.sqrt(tl.sum(b_q * b_q) + 1e-6) + b_k = b_k / tl.sqrt(tl.sum(b_k * b_k) + 1e-6) + b_q = b_q * scale + # [BV, BK] + if not IS_KDA: + b_g = tl.load(p_g).to(tl.float32) + b_h *= exp(b_g) + else: + b_gk = tl.load(p_gk).to(tl.float32) + if COMPUTE_GATE: + # Replicates kda_gate_fwd_kernel's SAFE_GATE branch + # bit-for-bit (same tl.exp, same fp32 math; the intermediate + # gate value this replaces was stored/reloaded as fp32, + # which is lossless): y = lb / (1 + exp(-exp(A)*(g+bias))). + b_gk += tl.load(g_bias + i_h * K + o_k, mask=mask_k, other=0.0).to( + tl.float32 + ) + b_gk = LOWER_BOUND / (1.0 + tl.exp(-(b_a_log * b_gk))) + b_h *= exp(b_gk[None, :]) + # [BV] + b_v -= tl.sum(b_h * b_k[None, :], 1) + if IS_BETA_HEADWISE: + b_beta = tl.load(p_beta, mask=mask_v, other=0).to(tl.float32) + else: + b_beta = tl.load(p_beta).to(tl.float32) + # Matches torch's `x.float().sigmoid()` pre-computation bit-for-bit + # on the input side (bf16->fp32 is exact); only the sigmoid impl itself + # can differ by <=1 ULP. + if SIGMOID_BETA: + b_beta = tl.sigmoid(b_beta) + b_v *= b_beta + # [BV, BK] + b_h += b_v[:, None] * b_k[None, :] + # [BV] + b_o = tl.sum(b_h * b_q[None, :], 1) + tl.store(p_o, b_o.to(p_o.dtype.element_ty), mask=mask_v) + + # keep the states for multi-query tokens + if INPLACE_FINAL_STATE: + # Load state index and check for invalid entries + final_state_idx = tl.load( + ssm_state_indices + i_n * stride_indices_seq + i_t + ).to(tl.int64) + # Only store if state index is valid (not NULL_BLOCK_ID=0) + if final_state_idx > 0: + p_ht = ht + final_state_idx * stride_final_state_token + p_ht = p_ht + i_hv * V * K + o_v[:, None] * K + o_k[None, :] + tl.store(p_ht, b_h.to(p_ht.dtype.element_ty), mask=mask_h) + else: + p_ht = ht + (bos + i_t) * stride_final_state_token + p_ht = p_ht + i_hv * V * K + o_v[:, None] * K + o_k[None, :] + tl.store(p_ht, b_h.to(p_ht.dtype.element_ty), mask=mask_h) + + p_q += H * K + p_k += H * K + p_o += HV * V + p_v += HV * V + if not IS_KDA: + p_g += HV + else: + p_gk += HV * K + p_beta += HV * (V if IS_BETA_HEADWISE else 1) + + +def fused_recurrent_gated_delta_rule_fwd( + q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, + g: torch.Tensor, + beta: torch.Tensor, + scale: float, + initial_state: torch.Tensor, + inplace_final_state: bool = True, + cu_seqlens: torch.Tensor | None = None, + ssm_state_indices: torch.Tensor | None = None, + num_accepted_tokens: torch.Tensor | None = None, + use_qk_l2norm_in_kernel: bool = False, +) -> tuple[torch.Tensor, torch.Tensor]: + B, T, H, K, V = *k.shape, v.shape[-1] + HV = v.shape[2] + N = B if cu_seqlens is None else len(cu_seqlens) - 1 + BK, BV = triton.next_power_of_2(K), min(triton.next_power_of_2(V), 32) + NK, NV = triton.cdiv(K, BK), triton.cdiv(V, BV) + assert NK == 1, "NK > 1 is not supported yet" + num_stages = 3 + num_warps = 1 + + o = q.new_empty(NK, *v.shape) + if inplace_final_state: + final_state = initial_state + else: + final_state = q.new_empty(T, HV, V, K, dtype=initial_state.dtype) + + stride_init_state_token = initial_state.stride(0) + stride_final_state_token = final_state.stride(0) + + if ssm_state_indices is None: + stride_indices_seq, stride_indices_tok = 1, 1 + elif ssm_state_indices.ndim == 1: + stride_indices_seq, stride_indices_tok = ssm_state_indices.stride(0), 1 + else: + stride_indices_seq, stride_indices_tok = ssm_state_indices.stride() + + grid = (NK, NV, N * HV) + fused_recurrent_gated_delta_rule_fwd_kernel[grid]( + q=q, + k=k, + v=v, + g=g, + beta=beta, + o=o, + h0=initial_state, + ht=final_state, + cu_seqlens=cu_seqlens, + ssm_state_indices=ssm_state_indices, + num_accepted_tokens=num_accepted_tokens, + scale=scale, + N=N, + T=T, + B=B, + H=H, + HV=HV, + K=K, + V=V, + BK=BK, + BV=BV, + stride_init_state_token=stride_init_state_token, + stride_final_state_token=stride_final_state_token, + stride_indices_seq=stride_indices_seq, + stride_indices_tok=stride_indices_tok, + IS_BETA_HEADWISE=beta.ndim == v.ndim, + USE_QK_L2NORM_IN_KERNEL=use_qk_l2norm_in_kernel, + INPLACE_FINAL_STATE=inplace_final_state, + IS_KDA=False, + SIGMOID_BETA=False, + a_log=None, + g_bias=None, + COMPUTE_GATE=False, + SAFE_GATE=True, + LOWER_BOUND=-5.0, + num_warps=num_warps, + num_stages=num_stages, + ) + o = o.squeeze(0) + return o, final_state + + +@triton.jit +def fused_recurrent_gated_delta_rule_packed_decode_kernel( + mixed_qkv, + a, + b, + A_log, + dt_bias, + o, + h0, + ht, + ssm_state_indices, + scale, + stride_mixed_qkv_tok: tl.constexpr, + stride_a_tok: tl.constexpr, + stride_b_tok: tl.constexpr, + stride_init_state_token: tl.constexpr, + stride_final_state_token: tl.constexpr, + stride_indices_seq: tl.constexpr, + H: tl.constexpr, + HV: tl.constexpr, + K: tl.constexpr, + V: tl.constexpr, + BK: tl.constexpr, + BV: tl.constexpr, + SOFTPLUS_THRESHOLD: tl.constexpr, + USE_QK_L2NORM_IN_KERNEL: tl.constexpr, + SPLIT_BATCH_HEAD_GRID: tl.constexpr, +): + if SPLIT_BATCH_HEAD_GRID: + i_v, i_hv, i_n = tl.program_id(0), tl.program_id(1), tl.program_id(2) + else: + i_v, i_nh = tl.program_id(0), tl.program_id(1) + i_n, i_hv = i_nh // HV, i_nh % HV + i_h = i_hv // (HV // H) + + o_k = tl.arange(0, BK) + o_v = i_v * BV + tl.arange(0, BV) + mask_k = o_k < K + mask_v = o_v < V + mask_h = mask_v[:, None] & mask_k[None, :] + + state_idx = tl.load(ssm_state_indices + i_n * stride_indices_seq).to(tl.int64) + p_o = o + (i_n * HV + i_hv) * V + o_v + + # Skip if state index is invalid (NULL_BLOCK_ID=0) + if state_idx <= 0: + zero = tl.zeros([BV], dtype=tl.float32).to(p_o.dtype.element_ty) + tl.store(p_o, zero, mask=mask_v) + return + + p_h0 = h0 + state_idx * stride_init_state_token + p_h0 = p_h0 + i_hv * V * K + o_v[:, None] * K + o_k[None, :] + b_h = tl.load(p_h0, mask=mask_h, other=0).to(tl.float32) + + p_mixed = mixed_qkv + i_n * stride_mixed_qkv_tok + q_off = i_h * K + o_k + k_off = (H * K) + i_h * K + o_k + v_off = (2 * H * K) + i_hv * V + o_v + b_q = tl.load(p_mixed + q_off, mask=mask_k, other=0).to(tl.float32) + b_k = tl.load(p_mixed + k_off, mask=mask_k, other=0).to(tl.float32) + b_v = tl.load(p_mixed + v_off, mask=mask_v, other=0).to(tl.float32) + + if USE_QK_L2NORM_IN_KERNEL: + b_q = b_q / tl.sqrt(tl.sum(b_q * b_q) + 1e-6) + b_k = b_k / tl.sqrt(tl.sum(b_k * b_k) + 1e-6) + b_q = b_q * scale + + a_val = tl.load(a + i_n * stride_a_tok + i_hv).to(tl.float32) + b_val = tl.load(b + i_n * stride_b_tok + i_hv).to(tl.float32) + A_log_val = tl.load(A_log + i_hv).to(tl.float32) + dt_bias_val = tl.load(dt_bias + i_hv).to(tl.float32) + x = a_val + dt_bias_val + softplus_x = tl.where(x <= SOFTPLUS_THRESHOLD, tl.log(1.0 + tl.exp(x)), x) + g_val = -tl.exp(A_log_val) * softplus_x + beta_val = tl.sigmoid(b_val) + + b_h *= exp(g_val) + b_v -= tl.sum(b_h * b_k[None, :], 1) + b_v *= beta_val + b_h += b_v[:, None] * b_k[None, :] + b_o = tl.sum(b_h * b_q[None, :], 1) + tl.store(p_o, b_o.to(p_o.dtype.element_ty), mask=mask_v) + + p_ht = ht + state_idx * stride_final_state_token + p_ht = p_ht + i_hv * V * K + o_v[:, None] * K + o_k[None, :] + tl.store(p_ht, b_h.to(p_ht.dtype.element_ty), mask=mask_h) + + +def fused_recurrent_gated_delta_rule_packed_decode( + mixed_qkv: torch.Tensor, + a: torch.Tensor, + b: torch.Tensor, + A_log: torch.Tensor, + dt_bias: torch.Tensor, + scale: float, + initial_state: torch.Tensor, + out: torch.Tensor, + ssm_state_indices: torch.Tensor, + use_qk_l2norm_in_kernel: bool = False, +) -> tuple[torch.Tensor, torch.Tensor]: + if mixed_qkv.ndim != 2: + raise ValueError( + f"`mixed_qkv` must be a 2D tensor (got ndim={mixed_qkv.ndim})." + ) + if mixed_qkv.stride(-1) != 1: + raise ValueError("`mixed_qkv` must be contiguous in the last dim.") + if a.ndim != 2 or b.ndim != 2: + raise ValueError( + f"`a` and `b` must be 2D tensors (got a.ndim={a.ndim}, b.ndim={b.ndim})." + ) + if a.stride(-1) != 1 or b.stride(-1) != 1: + raise ValueError("`a`/`b` must be contiguous in the last dim.") + if A_log.ndim != 1 or dt_bias.ndim != 1: + raise ValueError("`A_log`/`dt_bias` must be 1D tensors.") + if A_log.stride(0) != 1 or dt_bias.stride(0) != 1: + raise ValueError("`A_log`/`dt_bias` must be contiguous.") + if ssm_state_indices.ndim != 1: + raise ValueError( + f"`ssm_state_indices` must be 1D for packed decode (got ndim={ssm_state_indices.ndim})." + ) + if not out.is_contiguous(): + raise ValueError("`out` must be contiguous.") + + dev = mixed_qkv.device + if ( + a.device != dev + or b.device != dev + or A_log.device != dev + or dt_bias.device != dev + or initial_state.device != dev + or out.device != dev + or ssm_state_indices.device != dev + ): + raise ValueError("All inputs must be on the same device.") + + B = mixed_qkv.shape[0] + if a.shape[0] != B or b.shape[0] != B: + raise ValueError( + "Mismatched batch sizes: " + f"mixed_qkv.shape[0]={B}, a.shape[0]={a.shape[0]}, b.shape[0]={b.shape[0]}." + ) + if ssm_state_indices.shape[0] != B: + raise ValueError( + f"`ssm_state_indices` must have shape [B] (got {tuple(ssm_state_indices.shape)}; expected ({B},))." + ) + + if initial_state.ndim != 4: + raise ValueError( + f"`initial_state` must be a 4D tensor (got ndim={initial_state.ndim})." + ) + if initial_state.stride(-1) != 1: + raise ValueError("`initial_state` must be contiguous in the last dim.") + HV, V, K = initial_state.shape[-3:] + if a.shape[1] != HV or b.shape[1] != HV: + raise ValueError( + f"`a`/`b` must have shape [B, HV] with HV={HV} (got a.shape={tuple(a.shape)}, b.shape={tuple(b.shape)})." + ) + if A_log.numel() != HV or dt_bias.numel() != HV: + raise ValueError( + f"`A_log` and `dt_bias` must have {HV} elements (got A_log.numel()={A_log.numel()}, dt_bias.numel()={dt_bias.numel()})." + ) + if out.shape != (B, 1, HV, V): + raise ValueError( + f"`out` must have shape {(B, 1, HV, V)} (got out.shape={tuple(out.shape)})." + ) + + qkv_dim = mixed_qkv.shape[1] + qk_dim = qkv_dim - HV * V + if qk_dim <= 0 or qk_dim % 2 != 0: + raise ValueError( + f"Invalid packed `mixed_qkv` last dim={qkv_dim} for HV={HV}, V={V}." + ) + q_dim = qk_dim // 2 + if q_dim % K != 0: + raise ValueError(f"Invalid packed Q size {q_dim}: must be divisible by K={K}.") + H = q_dim // K + if H <= 0 or HV % H != 0: + raise ValueError( + f"Invalid head config inferred from mixed_qkv: H={H}, HV={HV}." + ) + + BK = triton.next_power_of_2(K) + if triton.cdiv(K, BK) != 1: + raise ValueError( + f"Packed decode kernel only supports NK=1 (got K={K}, BK={BK})." + ) + BV = min(triton.next_power_of_2(V), 32) + num_stages = 3 + num_warps = 1 + + stride_mixed_qkv_tok = mixed_qkv.stride(0) + stride_a_tok = a.stride(0) + stride_b_tok = b.stride(0) + stride_init_state_token = initial_state.stride(0) + stride_final_state_token = initial_state.stride(0) + stride_indices_seq = ssm_state_indices.stride(0) + + NV = triton.cdiv(V, BV) + # CUDA limits grid Y/Z dimensions to 65535. + split_batch_head_grid = B * HV > 65535 + grid = (NV, HV, B) if split_batch_head_grid else (NV, B * HV) + fused_recurrent_gated_delta_rule_packed_decode_kernel[grid]( + mixed_qkv=mixed_qkv, + a=a, + b=b, + A_log=A_log, + dt_bias=dt_bias, + o=out, + h0=initial_state, + ht=initial_state, + ssm_state_indices=ssm_state_indices, + scale=scale, + stride_mixed_qkv_tok=stride_mixed_qkv_tok, + stride_a_tok=stride_a_tok, + stride_b_tok=stride_b_tok, + stride_init_state_token=stride_init_state_token, + stride_final_state_token=stride_final_state_token, + stride_indices_seq=stride_indices_seq, + H=H, + HV=HV, + K=K, + V=V, + BK=BK, + BV=BV, + SOFTPLUS_THRESHOLD=20.0, + USE_QK_L2NORM_IN_KERNEL=use_qk_l2norm_in_kernel, + SPLIT_BATCH_HEAD_GRID=split_batch_head_grid, + num_warps=num_warps, + num_stages=num_stages, + ) + return out, initial_state + + +class FusedRecurrentFunction(torch.autograd.Function): + @staticmethod + def forward( + ctx, + q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, + g: torch.Tensor, + beta: torch.Tensor, + scale: float, + initial_state: torch.Tensor, + inplace_final_state: bool = True, + cu_seqlens: torch.Tensor | None = None, + ssm_state_indices: torch.Tensor | None = None, + num_accepted_tokens: torch.Tensor | None = None, + use_qk_l2norm_in_kernel: bool = False, + ): + o, final_state = fused_recurrent_gated_delta_rule_fwd( + q=q.contiguous(), + k=k.contiguous(), + v=v.contiguous(), + g=g.contiguous(), + beta=beta.contiguous(), + scale=scale, + initial_state=initial_state, + inplace_final_state=inplace_final_state, + cu_seqlens=cu_seqlens, + ssm_state_indices=ssm_state_indices, + num_accepted_tokens=num_accepted_tokens, + use_qk_l2norm_in_kernel=use_qk_l2norm_in_kernel, + ) + + return o, final_state + + +def fused_recurrent_gated_delta_rule( + q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, + g: torch.Tensor, + beta: torch.Tensor = None, + scale: float = None, + initial_state: torch.Tensor = None, + inplace_final_state: bool = True, + cu_seqlens: torch.Tensor | None = None, + ssm_state_indices: torch.Tensor | None = None, + num_accepted_tokens: torch.Tensor | None = None, + use_qk_l2norm_in_kernel: bool = False, +) -> tuple[torch.Tensor, torch.Tensor]: + r""" + Args: + q (torch.Tensor): + queries of shape `[B, T, H, K]`. + k (torch.Tensor): + keys of shape `[B, T, H, K]`. + v (torch.Tensor): + values of shape `[B, T, HV, V]`. + GVA is applied if `HV > H`. + g (torch.Tensor): + g (decays) of shape `[B, T, HV]`. + beta (torch.Tensor): + betas of shape `[B, T, HV]`. + scale (Optional[int]): + Scale factor for the RetNet attention scores. + If not provided, it will default to `1 / sqrt(K)`. Default: `None`. + initial_state (Optional[torch.Tensor]): + Initial state of shape `[N, HV, V, K]` for `N` input sequences. + For equal-length input sequences, `N` equals the batch size `B`. + Default: `None`. + inplace_final_state: bool: + Whether to store the final state in-place to save memory. + Default: `True`. + cu_seqlens (torch.Tensor): + Cumulative sequence lengths of shape `[N+1]` used for variable-length training, + consistent with the FlashAttention API. + ssm_state_indices (Optional[torch.Tensor]): + Indices to map the input sequences to the initial/final states. + num_accepted_tokens (Optional[torch.Tensor]): + Number of accepted tokens for each sequence during decoding. + + Returns: + o (torch.Tensor): + Outputs of shape `[B, T, HV, V]`. + final_state (torch.Tensor): + Final state of shape `[N, HV, V, K]`. + + Examples:: + >>> import torch + >>> import torch.nn.functional as F + >>> from einops import rearrange + >>> from fla.ops.gated_delta_rule import fused_recurrent_gated_delta_rule + # inputs with equal lengths + >>> B, T, H, HV, K, V = 4, 2048, 4, 8, 512, 512 + >>> q = torch.randn(B, T, H, K, device='cuda') + >>> k = F.normalize(torch.randn(B, T, H, K, device='cuda'), p=2, dim=-1) + >>> v = torch.randn(B, T, HV, V, device='cuda') + >>> g = F.logsigmoid(torch.rand(B, T, HV, device='cuda')) + >>> beta = torch.rand(B, T, HV, device='cuda').sigmoid() + >>> h0 = torch.randn(B, HV, V, K, device='cuda') + >>> o, ht = fused_gated_recurrent_delta_rule( + q, k, v, g, beta, + initial_state=h0, + ) + # for variable-length inputs, the batch size `B` is expected to be 1 and `cu_seqlens` is required + >>> q, k, v, g, beta = map(lambda x: rearrange(x, 'b t ... -> 1 (b t) ...'), (q, k, v, g, beta)) + # for a batch with 4 sequences, `cu_seqlens` with 5 start/end positions are expected + >>> cu_seqlens = q.new_tensor([0, 2048, 4096, 6144, 8192], dtype=torch.int32) + >>> o_var, ht_var = fused_gated_recurrent_delta_rule( + q, k, v, g, beta, + initial_state=h0, + cu_seqlens=cu_seqlens + ) + """ + if cu_seqlens is not None and q.shape[0] != 1: + raise ValueError( + f"The batch size is expected to be 1 rather than {q.shape[0]} when using `cu_seqlens`." + f"Please flatten variable-length inputs before processing." + ) + if scale is None: + scale = k.shape[-1] ** -0.5 + else: + assert scale > 0, "scale must be positive" + if beta is None: + beta = torch.ones_like(q[..., 0]) + o, final_state = FusedRecurrentFunction.apply( + q, + k, + v, + g, + beta, + scale, + initial_state, + inplace_final_state, + cu_seqlens, + ssm_state_indices, + num_accepted_tokens, + use_qk_l2norm_in_kernel, + ) + return o, final_state diff --git a/vllm/models/glm5next/nvidia/ops/third_party/kda/kernels.py b/vllm/models/glm5next/nvidia/ops/third_party/kda/kernels.py new file mode 100644 index 000000000000..be417d6b0526 --- /dev/null +++ b/vllm/models/glm5next/nvidia/ops/third_party/kda/kernels.py @@ -0,0 +1,1362 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +# SPDX-FileCopyrightText: Songlin Yang, Yu Zhang +# mypy: ignore-errors +# +# This file contains code copied from the flash-linear-attention project. +# The original source code was licensed under the MIT license and included +# the following copyright notice: +# Copyright (c) 2023-2025, Songlin Yang, Yu Zhang +# ruff: noqa: E501 + + +import torch + +from vllm.third_party.flash_linear_attention.ops.chunk_delta_h import ( + chunk_gated_delta_rule_fwd_h, +) +from vllm.third_party.flash_linear_attention.ops.cumsum import chunk_local_cumsum +from vllm.third_party.flash_linear_attention.ops.index import prepare_chunk_indices +from vllm.third_party.flash_linear_attention.ops.l2norm import l2norm_fwd +from vllm.third_party.flash_linear_attention.ops.op import exp2, log +from vllm.third_party.flash_linear_attention.ops.solve_tril import solve_tril +from vllm.third_party.flash_linear_attention.ops.utils import FLA_CHUNK_SIZE, is_amd +from vllm.triton_utils import tl, triton +from vllm.utils.math_utils import RCP_LN2, cdiv, next_power_of_2 + +from .fused_recurrent import fused_recurrent_gated_delta_rule_fwd_kernel + +BT_LIST_AUTOTUNE = [32, 64, 128] +NUM_WARPS_AUTOTUNE = [2, 4, 8, 16] if is_amd else [4, 8, 16, 32] + + +def fused_recurrent_kda_fwd( + q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, + g: torch.Tensor, + beta: torch.Tensor, + scale: float, + initial_state: torch.Tensor, + inplace_final_state: bool = True, + cu_seqlens: torch.Tensor | None = None, + ssm_state_indices: torch.Tensor | None = None, + num_accepted_tokens: torch.Tensor | None = None, + use_qk_l2norm_in_kernel: bool = False, + out: torch.Tensor | None = None, + sigmoid_beta: bool = False, + a_log: torch.Tensor | None = None, + g_bias: torch.Tensor | None = None, + compute_gate: bool = False, + lower_bound: float | None = -5.0, +) -> tuple[torch.Tensor, torch.Tensor]: + B, T, H, K, V = *k.shape, v.shape[-1] + HV = v.shape[2] + N = B if cu_seqlens is None else len(cu_seqlens) - 1 + BK, BV = next_power_of_2(K), min(next_power_of_2(V), 8) + NK, NV = cdiv(K, BK), cdiv(V, BV) + assert NK == 1, "NK > 1 is not supported yet" + num_stages = 3 + num_warps = 1 + + if compute_gate: + assert a_log is not None and g_bias is not None, ( + "compute_gate requires a_log and g_bias" + ) + assert lower_bound is not None, ( + "compute_gate implements the bounded (safe_gate) branch only" + ) + a_log = a_log.reshape(-1).contiguous() + g_bias = g_bias.reshape(-1).contiguous() + + if out is None: + o = torch.empty_like(k) + else: + # Caller-provided output buffer; must be layout-compatible with the + # tensor the kernel indexes (contiguous, same shape/dtype as k). + assert out.shape == k.shape and out.dtype == k.dtype + assert out.is_contiguous() + o = out + if inplace_final_state: + final_state = initial_state + else: + final_state = q.new_empty(T, HV, V, K, dtype=initial_state.dtype) + + stride_init_state_token = initial_state.stride(0) + stride_final_state_token = final_state.stride(0) + + if ssm_state_indices is None: + stride_indices_seq, stride_indices_tok = 1, 1 + elif ssm_state_indices.ndim == 1: + stride_indices_seq, stride_indices_tok = ssm_state_indices.stride(0), 1 + else: + stride_indices_seq, stride_indices_tok = ssm_state_indices.stride() + + grid = (NK, NV, N * HV) + fused_recurrent_gated_delta_rule_fwd_kernel[grid]( + q=q, + k=k, + v=v, + g=g, + beta=beta, + o=o, + h0=initial_state, + ht=final_state, + cu_seqlens=cu_seqlens, + ssm_state_indices=ssm_state_indices, + num_accepted_tokens=num_accepted_tokens, + scale=scale, + N=N, + T=T, + B=B, + H=H, + HV=HV, + K=K, + V=V, + BK=BK, + BV=BV, + stride_init_state_token=stride_init_state_token, + stride_final_state_token=stride_final_state_token, + stride_indices_seq=stride_indices_seq, + stride_indices_tok=stride_indices_tok, + IS_BETA_HEADWISE=beta.ndim == v.ndim, + USE_QK_L2NORM_IN_KERNEL=use_qk_l2norm_in_kernel, + INPLACE_FINAL_STATE=inplace_final_state, + IS_KDA=True, + SIGMOID_BETA=sigmoid_beta, + a_log=a_log, + g_bias=g_bias, + COMPUTE_GATE=compute_gate, + SAFE_GATE=True, + LOWER_BOUND=lower_bound if lower_bound is not None else -5.0, + num_warps=num_warps, + num_stages=num_stages, + ) + + return o, final_state + + +def fused_recurrent_kda( + q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, + g: torch.Tensor, + beta: torch.Tensor = None, + scale: float = None, + initial_state: torch.Tensor = None, + inplace_final_state: bool = True, + use_qk_l2norm_in_kernel: bool = True, + cu_seqlens: torch.Tensor | None = None, + ssm_state_indices: torch.LongTensor | None = None, + num_accepted_tokens: torch.Tensor | None = None, + out: torch.Tensor | None = None, + sigmoid_beta: bool = False, + a_log: torch.Tensor | None = None, + g_bias: torch.Tensor | None = None, + compute_gate: bool = False, + lower_bound: float | None = -5.0, + **kwargs, +) -> tuple[torch.Tensor, torch.Tensor]: + if cu_seqlens is not None and q.shape[0] != 1: + raise ValueError( + f"The batch size is expected to be 1 rather than {q.shape[0]} when using `cu_seqlens`." + f"Please flatten variable-length inputs before processing." + ) + if scale is None: + scale = k.shape[-1] ** -0.5 + + o, final_state = fused_recurrent_kda_fwd( + q=q.contiguous(), + k=k.contiguous(), + v=v.contiguous(), + g=g.contiguous(), + beta=beta.contiguous(), + scale=scale, + initial_state=initial_state, + inplace_final_state=inplace_final_state, + cu_seqlens=cu_seqlens, + ssm_state_indices=ssm_state_indices, + num_accepted_tokens=num_accepted_tokens, + use_qk_l2norm_in_kernel=use_qk_l2norm_in_kernel, + out=out, + sigmoid_beta=sigmoid_beta, + a_log=a_log, + g_bias=g_bias, + compute_gate=compute_gate, + lower_bound=lower_bound, + ) + return o, final_state + + +@triton.heuristics({"IS_VARLEN": lambda args: args["cu_seqlens"] is not None}) +@triton.autotune( + configs=[ + triton.Config({"BK": BK}, num_warps=num_warps, num_stages=num_stages) + for BK in [32, 64] + for num_warps in [1, 2, 4, 8] + for num_stages in [2, 3, 4] + ], + key=["BC"], +) +@triton.jit(do_not_specialize=["T"]) +def chunk_kda_scaled_dot_kkt_fwd_kernel_intra_sub_inter( + q, + k, + g, + beta, + A, + Aqk, + scale, + cu_seqlens, + chunk_indices, + T, + H: tl.constexpr, + K: tl.constexpr, + BT: tl.constexpr, + BC: tl.constexpr, + BK: tl.constexpr, + NC: tl.constexpr, + IS_VARLEN: tl.constexpr, +): + i_t, i_c, i_bh = tl.program_id(0), tl.program_id(1), tl.program_id(2) + i_b, i_h = i_bh // H, i_bh % H + i_i, i_j = i_c // NC, i_c % NC + if IS_VARLEN: + i_n, i_t = ( + tl.load(chunk_indices + i_t * 2).to(tl.int32), + tl.load(chunk_indices + i_t * 2 + 1).to(tl.int32), + ) + bos, eos = ( + tl.load(cu_seqlens + i_n).to(tl.int32), + tl.load(cu_seqlens + i_n + 1).to(tl.int32), + ) + T = eos - bos + else: + bos, eos = i_b * T, i_b * T + T + + if i_t * BT + i_i * BC >= T: + return + if i_i <= i_j: + return + + q += (bos * H + i_h) * K + k += (bos * H + i_h) * K + g += (bos * H + i_h) * K + A += (bos * H + i_h) * BT + Aqk += (bos * H + i_h) * BT + + p_b = tl.make_block_ptr( + beta + bos * H + i_h, (T,), (H,), (i_t * BT + i_i * BC,), (BC,), (0,) + ) + b_b = tl.load(p_b, boundary_check=(0,)) + + b_A = tl.zeros([BC, BC], dtype=tl.float32) + b_Aqk = tl.zeros([BC, BC], dtype=tl.float32) + for i_k in range(tl.cdiv(K, BK)): + p_q = tl.make_block_ptr( + q, (T, K), (H * K, 1), (i_t * BT + i_i * BC, i_k * BK), (BC, BK), (1, 0) + ) + p_k = tl.make_block_ptr( + k, (T, K), (H * K, 1), (i_t * BT + i_i * BC, i_k * BK), (BC, BK), (1, 0) + ) + p_g = tl.make_block_ptr( + g, (T, K), (H * K, 1), (i_t * BT + i_i * BC, i_k * BK), (BC, BK), (1, 0) + ) + b_kt = tl.make_block_ptr( + k, (K, T), (1, H * K), (i_k * BK, i_t * BT + i_j * BC), (BK, BC), (0, 1) + ) + p_gk = tl.make_block_ptr( + g, (K, T), (1, H * K), (i_k * BK, i_t * BT + i_j * BC), (BK, BC), (0, 1) + ) + + o_k = i_k * BK + tl.arange(0, BK) + m_k = o_k < K + # [BK,] + b_gn = tl.load(g + (i_t * BT + i_i * BC) * H * K + o_k, mask=m_k, other=0) + # [BC, BK] + b_g = tl.load(p_g, boundary_check=(0, 1)) + b_k = tl.load(p_k, boundary_check=(0, 1)) * exp2(b_g - b_gn[None, :]) + # [BK, BC] + b_gk = tl.load(p_gk, boundary_check=(0, 1)) + b_kt = tl.load(b_kt, boundary_check=(0, 1)) + # [BC, BC] + b_ktg = b_kt * exp2(b_gn[:, None] - b_gk) + b_A += tl.dot(b_k, b_ktg) + + b_q = tl.load(p_q, boundary_check=(0, 1)) + b_qg = b_q * exp2(b_g - b_gn[None, :]) * scale + b_Aqk += tl.dot(b_qg, b_ktg) + + b_A *= b_b[:, None] + + p_A = tl.make_block_ptr( + A, (T, BT), (H * BT, 1), (i_t * BT + i_i * BC, i_j * BC), (BC, BC), (1, 0) + ) + tl.store(p_A, b_A.to(A.dtype.element_ty), boundary_check=(0, 1)) + p_Aqk = tl.make_block_ptr( + Aqk, (T, BT), (H * BT, 1), (i_t * BT + i_i * BC, i_j * BC), (BC, BC), (1, 0) + ) + tl.store(p_Aqk, b_Aqk.to(Aqk.dtype.element_ty), boundary_check=(0, 1)) + + +@triton.heuristics({"IS_VARLEN": lambda args: args["cu_seqlens"] is not None}) +@triton.autotune( + configs=[triton.Config({}, num_warps=num_warps) for num_warps in [1, 2, 4, 8]], + key=["BK", "BT"], +) +@triton.jit(do_not_specialize=["T"]) +def chunk_kda_scaled_dot_kkt_fwd_kernel_intra_sub_intra( + q, + k, + g, + beta, + A, + Aqk, + scale, + cu_seqlens, + chunk_indices, + T, + H: tl.constexpr, + K: tl.constexpr, + BT: tl.constexpr, + BC: tl.constexpr, + BK: tl.constexpr, + IS_VARLEN: tl.constexpr, +): + i_t, i_i, i_bh = tl.program_id(0), tl.program_id(1), tl.program_id(2) + i_b, i_h = i_bh // H, i_bh % H + if IS_VARLEN: + i_n, i_t = ( + tl.load(chunk_indices + i_t * 2).to(tl.int32), + tl.load(chunk_indices + i_t * 2 + 1).to(tl.int32), + ) + bos, eos = ( + tl.load(cu_seqlens + i_n).to(tl.int32), + tl.load(cu_seqlens + i_n + 1).to(tl.int32), + ) + T = eos - bos + else: + bos, eos = i_b * T, i_b * T + T + + if i_t * BT + i_i * BC >= T: + return + + o_i = tl.arange(0, BC) + o_k = tl.arange(0, BK) + m_k = o_k < K + m_A = (i_t * BT + i_i * BC + o_i) < T + o_A = (bos + i_t * BT + i_i * BC + o_i) * H * BT + i_h * BT + i_i * BC + + p_q = tl.make_block_ptr( + q + (bos * H + i_h) * K, + (T, K), + (H * K, 1), + (i_t * BT + i_i * BC, 0), + (BC, BK), + (1, 0), + ) + p_k = tl.make_block_ptr( + k + (bos * H + i_h) * K, + (T, K), + (H * K, 1), + (i_t * BT + i_i * BC, 0), + (BC, BK), + (1, 0), + ) + p_g = tl.make_block_ptr( + g + (bos * H + i_h) * K, + (T, K), + (H * K, 1), + (i_t * BT + i_i * BC, 0), + (BC, BK), + (1, 0), + ) + b_q = tl.load(p_q, boundary_check=(0, 1)) + b_k = tl.load(p_k, boundary_check=(0, 1)) + b_g = tl.load(p_g, boundary_check=(0, 1)) + + p_b = beta + (bos + i_t * BT + i_i * BC + o_i) * H + i_h + b_k = b_k * tl.load(p_b, mask=m_A, other=0)[:, None] + + p_kt = k + (bos + i_t * BT + i_i * BC) * H * K + i_h * K + o_k + p_gk = g + (bos + i_t * BT + i_i * BC) * H * K + i_h * K + o_k + + for j in range(0, min(BC, T - i_t * BT - i_i * BC)): + b_kt = tl.load(p_kt, mask=m_k, other=0).to(tl.float32) + b_gk = tl.load(p_gk, mask=m_k, other=0).to(tl.float32) + b_ktg = b_kt[None, :] * exp2(b_g - b_gk[None, :]) + b_A = tl.sum(b_k * b_ktg, 1) + b_A = tl.where(o_i > j, b_A, 0.0) + b_Aqk = tl.sum(b_q * b_ktg, 1) + b_Aqk = tl.where(o_i >= j, b_Aqk * scale, 0.0) + tl.store(A + o_A + j, b_A, mask=m_A) + tl.store(Aqk + o_A + j, b_Aqk, mask=m_A) + p_kt += H * K + p_gk += H * K + + +def chunk_kda_scaled_dot_kkt_fwd( + q: torch.Tensor, + k: torch.Tensor, + gk: torch.Tensor | None = None, + beta: torch.Tensor | None = None, + scale: float | None = None, + cu_seqlens: torch.Tensor | None = None, + chunk_indices: torch.Tensor | None = None, + chunk_size: int = FLA_CHUNK_SIZE, + output_dtype: torch.dtype = torch.float32, +) -> tuple[torch.Tensor, torch.Tensor]: + r""" + Compute beta * K * K^T. + + Args: + k (torch.Tensor): + The key tensor of shape `[B, T, H, K]`. + beta (torch.Tensor): + The beta tensor of shape `[B, T, H]`. + gk (torch.Tensor): + The cumulative sum of the gate tensor of shape `[B, T, H, K]` applied to the key tensor. Default: `None`. + cu_seqlens (torch.Tensor): + The cumulative sequence lengths of the input tensor. + Default: None + chunk_size (int): + The chunk size. Default: 64. + output_dtype (torch.dtype): + The dtype of the output tensor. Default: `torch.float32` + + Returns: + beta * K * K^T of shape `[B, T, H, BT]` where `BT` is the chunk size. + """ + B, T, H, K = k.shape + assert K <= 256 + BT = chunk_size + if chunk_indices is None and cu_seqlens is not None: + chunk_indices = prepare_chunk_indices(cu_seqlens, BT) + NT = cdiv(T, BT) if cu_seqlens is None else len(chunk_indices) + + BC = min(16, BT) + NC = cdiv(BT, BC) + BK = max(next_power_of_2(K), 16) + A = torch.zeros(B, T, H, BT, device=k.device, dtype=output_dtype) + Aqk = torch.zeros(B, T, H, BT, device=k.device, dtype=output_dtype) + grid = (NT, NC * NC, B * H) + chunk_kda_scaled_dot_kkt_fwd_kernel_intra_sub_inter[grid]( + q=q, + k=k, + g=gk, + beta=beta, + A=A, + Aqk=Aqk, + scale=scale, + cu_seqlens=cu_seqlens, + chunk_indices=chunk_indices, + T=T, + H=H, + K=K, + BT=BT, + BC=BC, + NC=NC, + ) + + grid = (NT, NC, B * H) + chunk_kda_scaled_dot_kkt_fwd_kernel_intra_sub_intra[grid]( + q=q, + k=k, + g=gk, + beta=beta, + A=A, + Aqk=Aqk, + scale=scale, + cu_seqlens=cu_seqlens, + chunk_indices=chunk_indices, + T=T, + H=H, + K=K, + BT=BT, + BC=BC, + BK=BK, + ) + return A, Aqk + + +@triton.heuristics( + { + "STORE_QG": lambda args: args["qg"] is not None, + "STORE_KG": lambda args: args["kg"] is not None, + "IS_VARLEN": lambda args: args["cu_seqlens"] is not None, + } +) +@triton.autotune( + configs=[ + triton.Config({}, num_warps=num_warps, num_stages=num_stages) + for num_warps in [2, 4, 8] + for num_stages in [2, 3, 4] + ], + key=["H", "K", "V", "BT", "BK", "BV", "IS_VARLEN"], +) +@triton.jit(do_not_specialize=["T"]) +def recompute_w_u_fwd_kernel( + q, + k, + qg, + kg, + v, + beta, + w, + u, + A, + gk, + cu_seqlens, + chunk_indices, + T, + H: tl.constexpr, + K: tl.constexpr, + V: tl.constexpr, + BT: tl.constexpr, + BK: tl.constexpr, + BV: tl.constexpr, + STORE_QG: tl.constexpr, + STORE_KG: tl.constexpr, + IS_VARLEN: tl.constexpr, + DOT_PRECISION: tl.constexpr, +): + i_t, i_bh = tl.program_id(0), tl.program_id(1) + i_b, i_h = i_bh // H, i_bh % H + if IS_VARLEN: + i_n, i_t = ( + tl.load(chunk_indices + i_t * 2).to(tl.int32), + tl.load(chunk_indices + i_t * 2 + 1).to(tl.int32), + ) + bos, eos = ( + tl.load(cu_seqlens + i_n).to(tl.int32), + tl.load(cu_seqlens + i_n + 1).to(tl.int32), + ) + T = eos - bos + else: + bos, eos = i_b * T, i_b * T + T + p_b = tl.make_block_ptr(beta + bos * H + i_h, (T,), (H,), (i_t * BT,), (BT,), (0,)) + b_b = tl.load(p_b, boundary_check=(0,)) + + p_A = tl.make_block_ptr( + A + (bos * H + i_h) * BT, (T, BT), (H * BT, 1), (i_t * BT, 0), (BT, BT), (1, 0) + ) + b_A = tl.load(p_A, boundary_check=(0, 1)) + + for i_v in range(tl.cdiv(V, BV)): + p_v = tl.make_block_ptr( + v + (bos * H + i_h) * V, + (T, V), + (H * V, 1), + (i_t * BT, i_v * BV), + (BT, BV), + (1, 0), + ) + p_u = tl.make_block_ptr( + u + (bos * H + i_h) * V, + (T, V), + (H * V, 1), + (i_t * BT, i_v * BV), + (BT, BV), + (1, 0), + ) + b_v = tl.load(p_v, boundary_check=(0, 1)) + b_vb = (b_v * b_b[:, None]).to(b_v.dtype) + b_u = tl.dot(b_A, b_vb, input_precision=DOT_PRECISION) + tl.store(p_u, b_u.to(p_u.dtype.element_ty), boundary_check=(0, 1)) + + for i_k in range(tl.cdiv(K, BK)): + p_w = tl.make_block_ptr( + w + (bos * H + i_h) * K, + (T, K), + (H * K, 1), + (i_t * BT, i_k * BK), + (BT, BK), + (1, 0), + ) + p_k = tl.make_block_ptr( + k + (bos * H + i_h) * K, + (T, K), + (H * K, 1), + (i_t * BT, i_k * BK), + (BT, BK), + (1, 0), + ) + b_k = tl.load(p_k, boundary_check=(0, 1)) + b_kb = b_k * b_b[:, None] + + p_gk = tl.make_block_ptr( + gk + (bos * H + i_h) * K, + (T, K), + (H * K, 1), + (i_t * BT, i_k * BK), + (BT, BK), + (1, 0), + ) + b_gk = tl.load(p_gk, boundary_check=(0, 1)) + b_kb *= exp2(b_gk) + if STORE_QG: + p_q = tl.make_block_ptr( + q + (bos * H + i_h) * K, + (T, K), + (H * K, 1), + (i_t * BT, i_k * BK), + (BT, BK), + (1, 0), + ) + p_qg = tl.make_block_ptr( + qg + (bos * H + i_h) * K, + (T, K), + (H * K, 1), + (i_t * BT, i_k * BK), + (BT, BK), + (1, 0), + ) + b_q = tl.load(p_q, boundary_check=(0, 1)) + b_qg = b_q * exp2(b_gk) + tl.store(p_qg, b_qg.to(p_qg.dtype.element_ty), boundary_check=(0, 1)) + if STORE_KG: + last_idx = min(i_t * BT + BT, T) - 1 + + o_k = i_k * BK + tl.arange(0, BK) + m_k = o_k < K + b_gn = tl.load( + gk + ((bos + last_idx) * H + i_h) * K + o_k, mask=m_k, other=0.0 + ) + b_kg = b_k * exp2(b_gn - b_gk) + + p_kg = tl.make_block_ptr( + kg + (bos * H + i_h) * K, + (T, K), + (H * K, 1), + (i_t * BT, i_k * BK), + (BT, BK), + (1, 0), + ) + tl.store(p_kg, b_kg.to(p_kg.dtype.element_ty), boundary_check=(0, 1)) + + b_w = tl.dot(b_A, b_kb.to(b_k.dtype)) + tl.store(p_w, b_w.to(p_w.dtype.element_ty), boundary_check=(0, 1)) + + +def recompute_w_u_fwd( + k: torch.Tensor, + v: torch.Tensor, + beta: torch.Tensor, + A: torch.Tensor, + q: torch.Tensor | None = None, + gk: torch.Tensor | None = None, + cu_seqlens: torch.Tensor | None = None, + chunk_indices: torch.Tensor | None = None, +) -> tuple[torch.Tensor, torch.Tensor]: + B, T, H, K, V = *k.shape, v.shape[-1] + BT = A.shape[-1] + BK = 64 + BV = 64 + + if chunk_indices is None and cu_seqlens is not None: + chunk_indices = prepare_chunk_indices(cu_seqlens, BT) + NT = cdiv(T, BT) if cu_seqlens is None else len(chunk_indices) + + w = torch.empty_like(k) + u = torch.empty_like(v) + kg = torch.empty_like(k) if gk is not None else None + recompute_w_u_fwd_kernel[(NT, B * H)]( + q=q, + k=k, + qg=None, + kg=kg, + v=v, + beta=beta, + w=w, + u=u, + A=A, + gk=gk, + cu_seqlens=cu_seqlens, + chunk_indices=chunk_indices, + T=T, + H=H, + K=K, + V=V, + BT=BT, + BK=BK, + BV=BV, + DOT_PRECISION="ieee", + ) + return w, u, None, kg + + +@triton.heuristics({"IS_VARLEN": lambda args: args["cu_seqlens"] is not None}) +@triton.autotune( + configs=[ + triton.Config({"BK": BK, "BV": BV}, num_warps=num_warps, num_stages=num_stages) + for BK in [32, 64] + for BV in [64, 128] + for num_warps in [2, 4, 8] + for num_stages in [2, 3, 4] + ], + key=["BT"], +) +@triton.jit(do_not_specialize=["T"]) +def chunk_gla_fwd_kernel_o( + q, + v, + g, + h, + o, + A, + cu_seqlens, + chunk_indices, + scale, + T, + H: tl.constexpr, + K: tl.constexpr, + V: tl.constexpr, + BT: tl.constexpr, + BK: tl.constexpr, + BV: tl.constexpr, + IS_VARLEN: tl.constexpr, +): + i_v, i_t, i_bh = tl.program_id(0), tl.program_id(1), tl.program_id(2) + i_b, i_h = i_bh // H, i_bh % H + if IS_VARLEN: + i_tg = i_t + i_n, i_t = ( + tl.load(chunk_indices + i_t * 2).to(tl.int32), + tl.load(chunk_indices + i_t * 2 + 1).to(tl.int32), + ) + bos, eos = ( + tl.load(cu_seqlens + i_n).to(tl.int32), + tl.load(cu_seqlens + i_n + 1).to(tl.int32), + ) + T = eos - bos + NT = tl.cdiv(T, BT) + else: + NT = tl.cdiv(T, BT) + i_tg = i_b * NT + i_t + bos, eos = i_b * T, i_b * T + T + + m_s = tl.arange(0, BT)[:, None] >= tl.arange(0, BT)[None, :] + + b_o = tl.zeros([BT, BV], dtype=tl.float32) + for i_k in range(tl.cdiv(K, BK)): + p_q = tl.make_block_ptr( + q + (bos * H + i_h) * K, + (T, K), + (H * K, 1), + (i_t * BT, i_k * BK), + (BT, BK), + (1, 0), + ) + p_g = tl.make_block_ptr( + g + (bos * H + i_h) * K, + (T, K), + (H * K, 1), + (i_t * BT, i_k * BK), + (BT, BK), + (1, 0), + ) + p_h = tl.make_block_ptr( + h + (i_tg * H + i_h) * K * V, + (V, K), + (K, 1), + (i_v * BV, i_k * BK), + (BV, BK), + (1, 0), + ) + + # [BT, BK] + b_q = tl.load(p_q, boundary_check=(0, 1)) + b_q = (b_q * scale).to(b_q.dtype) + # [BT, BK] + b_g = tl.load(p_g, boundary_check=(0, 1)) + # [BT, BK] + b_qg = (b_q * exp2(b_g)).to(b_q.dtype) + # [BV, BK] + b_h = tl.load(p_h, boundary_check=(0, 1)) + # [BT, BV] + if i_k >= 0: + b_o += tl.dot(b_qg, tl.trans(b_h).to(b_qg.dtype)) + p_v = tl.make_block_ptr( + v + (bos * H + i_h) * V, + (T, V), + (H * V, 1), + (i_t * BT, i_v * BV), + (BT, BV), + (1, 0), + ) + p_o = tl.make_block_ptr( + o + (bos * H + i_h) * V, + (T, V), + (H * V, 1), + (i_t * BT, i_v * BV), + (BT, BV), + (1, 0), + ) + p_A = tl.make_block_ptr( + A + (bos * H + i_h) * BT, (T, BT), (H * BT, 1), (i_t * BT, 0), (BT, BT), (1, 0) + ) + # [BT, BV] + b_v = tl.load(p_v, boundary_check=(0, 1)) + # [BT, BT] + b_A = tl.load(p_A, boundary_check=(0, 1)) + b_A = tl.where(m_s, b_A, 0.0).to(b_v.dtype) + b_o += tl.dot(b_A, b_v, allow_tf32=False) + tl.store(p_o, b_o.to(p_o.dtype.element_ty), boundary_check=(0, 1)) + + +def chunk_gla_fwd_o_gk( + q: torch.Tensor, + v: torch.Tensor, + g: torch.Tensor, + A: torch.Tensor, + h: torch.Tensor, + o: torch.Tensor, + scale: float, + cu_seqlens: torch.Tensor | None = None, + chunk_indices: torch.Tensor | None = None, + chunk_size: int = FLA_CHUNK_SIZE, +): + B, T, H, K, V = *q.shape, v.shape[-1] + BT = chunk_size + + if chunk_indices is None and cu_seqlens is not None: + chunk_indices = prepare_chunk_indices(cu_seqlens, chunk_size) + NT = cdiv(T, BT) if cu_seqlens is None else len(chunk_indices) + + def grid(meta): + return (cdiv(V, meta["BV"]), NT, B * H) + + chunk_gla_fwd_kernel_o[grid]( + q=q, + v=v, + g=g, + h=h, + o=o, + A=A, + cu_seqlens=cu_seqlens, + chunk_indices=chunk_indices, + scale=scale, + T=T, + H=H, + K=K, + V=V, + BT=BT, + ) + return o + + +@triton.heuristics( + { + "HAS_BIAS": lambda args: args["g_bias"] is not None, + "IS_VARLEN": lambda args: args["cu_seqlens"] is not None, + } +) +@triton.autotune( + configs=[ + triton.Config({"BD": BD}, num_warps=num_warps) + for BD in [32, 64] + for num_warps in [2, 4, 8] + ], + key=["H", "D", "BT", "IS_VARLEN"], +) +@triton.jit(do_not_specialize=["T"]) +def kda_gate_cumsum_fwd_kernel( + g, + A, + y, + g_bias, + cu_seqlens, + chunk_indices, + cumsum_scale, + beta, + threshold, + SAFE_GATE: tl.constexpr, + LOWER_BOUND: tl.constexpr, + T, + H: tl.constexpr, + D: tl.constexpr, + BT: tl.constexpr, + BD: tl.constexpr, + HAS_BIAS: tl.constexpr, + IS_VARLEN: tl.constexpr, +): + i_d, i_t, i_bh = tl.program_id(0), tl.program_id(1), tl.program_id(2) + i_b, i_h = i_bh // H, i_bh % H + if IS_VARLEN: + i_n, i_t = ( + tl.load(chunk_indices + i_t * 2).to(tl.int32), + tl.load(chunk_indices + i_t * 2 + 1).to(tl.int32), + ) + bos, eos = ( + tl.load(cu_seqlens + i_n).to(tl.int32), + tl.load(cu_seqlens + i_n + 1).to(tl.int32), + ) + T = eos - bos + else: + bos = i_b * T + + p_g = tl.make_block_ptr( + g + (bos * H + i_h) * D, + (T, D), + (H * D, 1), + (i_t * BT, i_d * BD), + (BT, BD), + (1, 0), + ) + p_y = tl.make_block_ptr( + y + (bos * H + i_h) * D, + (T, D), + (H * D, 1), + (i_t * BT, i_d * BD), + (BT, BD), + (1, 0), + ) + + b_g = tl.load(p_g, boundary_check=(0, 1)).to(tl.float32) + if HAS_BIAS: + o_d = i_d * BD + tl.arange(0, BD) + b_bias = tl.load(g_bias + i_h * D + o_d, mask=o_d < D, other=0.0).to(tl.float32) + b_g = b_g + b_bias[None, :] + + b_a = tl.load(A + i_h).to(tl.float32) + b_a = tl.exp(b_a) if SAFE_GATE else -tl.exp(b_a) + if SAFE_GATE: + # y = lower_bound * sigmoid(exp(A) * (g + g_bias)), bounded to + # (lower_bound, 0) for safe-gate checkpoints. + b_gate = LOWER_BOUND / (1.0 + tl.exp(-(b_a * b_g))) + else: + b_g_scaled = b_g * beta + b_softplus = tl.where( + b_g_scaled > threshold, + b_g, + (1.0 / beta) * log(1.0 + tl.exp(b_g_scaled)), + ) + b_gate = b_a * b_softplus + + # Out-of-bounds rows (load returns 0, but softplus/bias can still make + # b_gate non-zero) participate in the dot product. They only contribute to + # out-of-bounds output rows, which are masked away by `boundary_check` on + # the store, so visible output matches unfused gate + chunk-local cumsum. + o_t = tl.arange(0, BT) + m_cumsum = tl.where(o_t[:, None] >= o_t[None, :], 1.0, 0.0) + b_y = tl.dot(m_cumsum, b_gate, allow_tf32=False) * cumsum_scale + tl.store(p_y, b_y.to(p_y.dtype.element_ty), boundary_check=(0, 1)) + + +def fused_kda_gate_chunk_cumsum( + raw_g: torch.Tensor, + A_log: torch.Tensor, + g_bias: torch.Tensor | None = None, + beta: float = 1.0, + threshold: float = 20.0, + cu_seqlens: torch.Tensor | None = None, + chunk_indices: torch.Tensor | None = None, + chunk_size: int = FLA_CHUNK_SIZE, + output_dtype: torch.dtype | None = torch.float, + safe_gate: bool = False, + lower_bound: float = -5.0, +) -> torch.Tensor: + if cu_seqlens is not None: + assert raw_g.shape[0] == 1, ( + "Only batch size 1 is supported when cu_seqlens are provided" + ) + B, T, H, D = raw_g.shape + if chunk_indices is None and cu_seqlens is not None: + chunk_indices = prepare_chunk_indices(cu_seqlens, chunk_size) + NT = cdiv(T, chunk_size) if cu_seqlens is None else len(chunk_indices) + + A_log = A_log.reshape(-1) + if g_bias is not None: + g_bias = g_bias.reshape(-1) + y = torch.empty_like(raw_g, dtype=output_dtype or raw_g.dtype) + + def grid(meta): + return (cdiv(meta["D"], meta["BD"]), NT, B * H) + + kda_gate_cumsum_fwd_kernel[grid]( + g=raw_g, + A=A_log, + y=y, + g_bias=g_bias, + cu_seqlens=cu_seqlens, + chunk_indices=chunk_indices, + # RCP_LN2 folds in the natural-log -> log2 conversion so downstream + # exp2-based kernels reproduce exp(g). Keep this in sync with the + # `use_exp2=True` path in `_chunk_kda_fwd_with_cumulative_g`. + cumsum_scale=RCP_LN2, + beta=beta, + threshold=threshold, + SAFE_GATE=safe_gate, + LOWER_BOUND=lower_bound, + T=T, + H=H, + D=D, + BT=chunk_size, + ) + return y + + +def _chunk_kda_fwd_with_cumulative_g( + q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, + g: torch.Tensor, + beta: torch.Tensor, + scale: float, + initial_state: torch.Tensor, + output_final_state: bool, + cu_seqlens: torch.Tensor | None = None, + chunk_indices: torch.Tensor | None = None, + chunk_size: int = FLA_CHUNK_SIZE, +): + # `g` must already be chunk-local cumulatively-summed AND scaled by + # RCP_LN2 (so the downstream exp2-based kernels reproduce exp(g)). + # Use `chunk_kda_fwd` or `chunk_kda_with_fused_gate_fwd` instead of + # calling this helper directly unless that invariant is upheld. + # the intra Aqk is kept in fp32 + # the computation has very marginal effect on the entire throughput + A, Aqk = chunk_kda_scaled_dot_kkt_fwd( + q=q, + k=k, + gk=g, + beta=beta, + scale=scale, + cu_seqlens=cu_seqlens, + chunk_indices=chunk_indices, + chunk_size=chunk_size, + output_dtype=torch.float32, + ) + A = solve_tril(A=A, cu_seqlens=cu_seqlens, output_dtype=k.dtype) + w, u, _, kg = recompute_w_u_fwd( + k=k, + v=v, + beta=beta, + A=A, + gk=g, + cu_seqlens=cu_seqlens, + chunk_indices=chunk_indices, + ) + del A + h, v_new, final_state = chunk_gated_delta_rule_fwd_h( + k=kg, + w=w, + u=u, + gk=g, + initial_state=initial_state, + output_final_state=output_final_state, + cu_seqlens=cu_seqlens, + chunk_indices=chunk_indices, + use_exp2=True, + ) + del w, u, kg + o = chunk_gla_fwd_o_gk( + q=q, + v=v_new, + g=g, + A=Aqk, + h=h, + o=v, + scale=scale, + cu_seqlens=cu_seqlens, + chunk_indices=chunk_indices, + chunk_size=chunk_size, + ) + del Aqk, v_new, h + return o, final_state + + +def chunk_kda_fwd( + q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, + g: torch.Tensor, + beta: torch.Tensor, + scale: float, + initial_state: torch.Tensor, + output_final_state: bool, + cu_seqlens: torch.Tensor | None = None, +): + chunk_size = FLA_CHUNK_SIZE + chunk_indices = ( + prepare_chunk_indices(cu_seqlens, chunk_size) + if cu_seqlens is not None + else None + ) + g = chunk_local_cumsum( + g, + chunk_size=chunk_size, + cu_seqlens=cu_seqlens, + chunk_indices=chunk_indices, + ) + # KDA evaluates cumulative gate decays with exp2. Convert from natural-log + # space so exp(x) is preserved as exp2(x / ln(2)). + g = g * RCP_LN2 + return _chunk_kda_fwd_with_cumulative_g( + q=q, + k=k, + v=v, + g=g, + beta=beta, + scale=scale, + initial_state=initial_state, + output_final_state=output_final_state, + cu_seqlens=cu_seqlens, + chunk_indices=chunk_indices, + chunk_size=chunk_size, + ) + + +def chunk_kda_with_fused_gate_fwd( + q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, + raw_g: torch.Tensor, + beta: torch.Tensor, + A_log: torch.Tensor, + g_bias: torch.Tensor | None, + scale: float, + initial_state: torch.Tensor, + output_final_state: bool, + cu_seqlens: torch.Tensor | None = None, + safe_gate: bool = False, + lower_bound: float = -5.0, +): + chunk_size = FLA_CHUNK_SIZE + chunk_indices = ( + prepare_chunk_indices(cu_seqlens, chunk_size) + if cu_seqlens is not None + else None + ) + g = fused_kda_gate_chunk_cumsum( + raw_g, + A_log=A_log, + g_bias=g_bias, + cu_seqlens=cu_seqlens, + chunk_indices=chunk_indices, + chunk_size=chunk_size, + safe_gate=safe_gate, + lower_bound=lower_bound, + ) + return _chunk_kda_fwd_with_cumulative_g( + q=q, + k=k, + v=v, + g=g, + beta=beta, + scale=scale, + initial_state=initial_state, + output_final_state=output_final_state, + cu_seqlens=cu_seqlens, + chunk_indices=chunk_indices, + chunk_size=chunk_size, + ) + + +def chunk_kda( + q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, + g: torch.Tensor, + beta: torch.Tensor, + scale: float = None, + initial_state: torch.Tensor = None, + output_final_state: bool = False, + use_qk_l2norm_in_kernel: bool = False, + cu_seqlens: torch.Tensor | None = None, + **kwargs, +): + if scale is None: + scale = k.shape[-1] ** -0.5 + + if use_qk_l2norm_in_kernel: + q = l2norm_fwd(q.contiguous()) + k = l2norm_fwd(k.contiguous()) + + o, final_state = chunk_kda_fwd( + q=q, + k=k, + v=v.contiguous(), + g=g.contiguous(), + beta=beta.contiguous(), + scale=scale, + initial_state=initial_state.contiguous(), + output_final_state=output_final_state, + cu_seqlens=cu_seqlens, + ) + return o, final_state + + +def chunk_kda_with_fused_gate( + q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, + raw_g: torch.Tensor, + beta: torch.Tensor, + A_log: torch.Tensor, + g_bias: torch.Tensor | None, + scale: float | None = None, + initial_state: torch.Tensor | None = None, + output_final_state: bool = False, + use_qk_l2norm_in_kernel: bool = False, + cu_seqlens: torch.Tensor | None = None, + safe_gate: bool = False, + lower_bound: float = -5.0, + **kwargs, +): + """Run chunk KDA from raw gate projection using fused gate+cumsum.""" + if scale is None: + scale = k.shape[-1] ** -0.5 + + if use_qk_l2norm_in_kernel: + q = l2norm_fwd(q.contiguous()) + k = l2norm_fwd(k.contiguous()) + + o, final_state = chunk_kda_with_fused_gate_fwd( + q=q, + k=k, + v=v.contiguous(), + raw_g=raw_g.contiguous(), + beta=beta.contiguous(), + A_log=A_log, + g_bias=g_bias, + scale=scale, + initial_state=initial_state.contiguous() if initial_state is not None else None, + output_final_state=output_final_state, + cu_seqlens=cu_seqlens, + safe_gate=safe_gate, + lower_bound=lower_bound, + ) + return o, final_state + + +@triton.autotune( + configs=[ + triton.Config({"BT": bt}, num_warps=nw, num_stages=ns) + for bt in BT_LIST_AUTOTUNE + for nw in NUM_WARPS_AUTOTUNE + for ns in [2, 3] + ], + key=["H", "D"], +) +@triton.jit +def kda_gate_fwd_kernel( + g, + A, + y, + g_bias, + beta: tl.constexpr, + threshold: tl.constexpr, + SAFE_GATE: tl.constexpr, + LOWER_BOUND: tl.constexpr, + T, + H, + D: tl.constexpr, + BT: tl.constexpr, + BD: tl.constexpr, + HAS_BIAS: tl.constexpr, +): + i_t, i_h = tl.program_id(0), tl.program_id(1) + n_t = i_t * BT + + b_a = tl.load(A + i_h).to(tl.float32) + b_a = tl.exp(b_a) if SAFE_GATE else -tl.exp(b_a) + + stride_row = H * D + stride_col = 1 + + g_ptr = tl.make_block_ptr( + base=g + i_h * D, + shape=(T, D), + strides=(stride_row, stride_col), + offsets=(n_t, 0), + block_shape=(BT, BD), + order=(1, 0), + ) + + y_ptr = tl.make_block_ptr( + base=y + i_h * D, + shape=(T, D), + strides=(stride_row, stride_col), + offsets=(n_t, 0), + block_shape=(BT, BD), + order=(1, 0), + ) + + b_g = tl.load(g_ptr, boundary_check=(0, 1)).to(tl.float32) + + if HAS_BIAS: + n_d = tl.arange(0, BD) + bias_mask = n_d < D + b_bias = tl.load(g_bias + i_h * D + n_d, mask=bias_mask, other=0.0).to( + tl.float32 + ) + b_g = b_g + b_bias[None, :] + + if SAFE_GATE: + # y = lower_bound * sigmoid(exp(A) * (g + g_bias)), bounded to + # (lower_bound, 0) for safe-gate checkpoints. + b_y = LOWER_BOUND / (1.0 + tl.exp(-(b_a * b_g))) + else: + # softplus(x, beta) = (1/beta) * log(1 + exp(beta * x)) + # When beta * x > threshold, use linear approximation x + # Use threshold to switch to linear when beta*x > threshold + g_scaled = b_g * beta + use_linear = g_scaled > threshold + sp = tl.where(use_linear, b_g, (1.0 / beta) * log(1.0 + tl.exp(g_scaled))) + b_y = b_a * sp + + tl.store(y_ptr, b_y.to(y.dtype.element_ty), boundary_check=(0, 1)) + + +def fused_kda_gate( + g: torch.Tensor, + A: torch.Tensor, + head_k_dim: int, + g_bias: torch.Tensor | None = None, + beta: float = 1.0, + threshold: float = 20.0, + safe_gate: bool = False, + lower_bound: float | None = -5.0, +) -> torch.Tensor: + """ + Forward pass for KDA gate: + input g: [..., H*D] + param A: [H] or [1, 1, H, 1] + beta: softplus beta parameter (softplus branch only) + threshold: softplus threshold parameter (softplus branch only) + safe_gate: when False (default) compute y = -exp(A)*softplus(g+g_bias); + when True compute the bounded y = lower_bound*sigmoid(exp(A)*(g+g_bias)) + lower_bound: floor for the safe_gate branch (default -5.0) + return : [..., H, D] + """ + orig_shape = g.shape[:-1] + + g = g.view(-1, g.shape[-1]) + T = g.shape[0] + HD = g.shape[1] + H = A.numel() + assert H * head_k_dim == HD + + y = torch.empty_like(g, dtype=torch.float32) + + def grid(meta): + return (cdiv(T, meta["BT"]), H) + + kda_gate_fwd_kernel[grid]( + g, + A, + y, + g_bias, + beta, + threshold, + safe_gate, + lower_bound if lower_bound is not None else -5.0, + T, + H, + head_k_dim, + BD=next_power_of_2(head_k_dim), + HAS_BIAS=g_bias is not None, + ) + + y = y.view(*orig_shape, H, head_k_dim) + return y diff --git a/vllm/multimodal/video.py b/vllm/multimodal/video.py index 698afd397a24..a02daa096e10 100644 --- a/vllm/multimodal/video.py +++ b/vllm/multimodal/video.py @@ -140,6 +140,24 @@ def _prepare_source(cls, source: VideoSourceMetadata) -> VideoSourceMetadata: """Sampling-algorithm-specific metadata adjustment hook.""" return source + @classmethod + def read_frames( + cls, + cap: "cv2.VideoCapture", + frame_idx: list[int], + total_frames_num: int, + *, + frame_recovery: bool = False, + ) -> tuple[npt.NDArray, list[int]]: + from vllm.multimodal.video_decoders.opencv import OpenCVVideoBackendMixin + + return OpenCVVideoBackendMixin.read_frames( + cap, + frame_idx, + total_frames_num, + frame_recovery=frame_recovery, + ) + @classmethod @abstractmethod def load_bytes( @@ -664,6 +682,100 @@ def load_bytes( ) +@VIDEO_LOADER_REGISTRY.register( + "glm5next", + # Both spellings: the borrowed-config type string (``Glm5Next...``) and + # the dedicated transformers classes landing with the new checkpoint + # (``Glm5nextVideoProcessor``, matching ``Glm5nextImageProcessor``). + video_processor=("Glm5NextVideoProcessor", "Glm5nextVideoProcessor"), +) +class Glm5NextVideoBackend(VideoBackend): + """GLM-5.3-Flash fps-interval video backend. + + Selects frames with the same ``glm_sample_frame_indices`` sampler the + processor falls back to, so only the sampled frames are + materialized. ``fps_interval`` semantics (default 2.0) with a + temporal-patch-scaled greedy walk, frame count capped at 2048, temporal + pairs kept even. Request overrides: ``fps`` -> fps interval, + ``max_frames`` -> frame cap, ``temporal_patch_size`` (default 2). + """ + + _SEEK_GAP_THRESHOLD: ClassVar[int] = 64 + + @classmethod + def compute_frames_index_to_sample( + cls, + source: VideoSourceMetadata, + target: VideoTargetMetadata, + **kwargs, + ) -> list[int]: + # Lazy import: the processor module sits behind the + # transformers_utils package init, which multimodal must not pull in. + from vllm.transformers_utils.processors.glm5next import ( + glm_sample_frame_indices, + ) + + return glm_sample_frame_indices( + source.total_frames_num, + source.original_fps, + source.duration or 0, + target_fps=target.fps if target.fps > 0 else None, + max_frame_count=kwargs.get("max_frames"), + temporal_patch_size=kwargs.get("temporal_patch_size", 2), + ) + + @classmethod + def read_frames( + cls, + cap: "cv2.VideoCapture", + frame_idx: list[int], + total_frames_num: int, + *, + frame_recovery: bool = False, + ) -> tuple[npt.NDArray, list[int]]: + if frame_recovery: + return super().read_frames( + cap, frame_idx, total_frames_num, frame_recovery=frame_recovery + ) + + wanted = sorted(set(frame_idx)) + width = int(cap.get(cv2.CAP_PROP_FRAME_WIDTH)) + height = int(cap.get(cv2.CAP_PROP_FRAME_HEIGHT)) + frames = np.empty((len(wanted), height, width, 3), dtype=np.uint8) + valid_frame_indices: list[int] = [] + current: int | None = None + for target in wanted: + gap = target - current if current is not None else None + if gap is not None and 2 <= gap <= cls._SEEK_GAP_THRESHOLD: + for _ in range(gap - 1): + if not cap.grab(): + current = None + break + else: + current = target - 1 + if current != target - 1: + cap.set(cv2.CAP_PROP_POS_FRAMES, target) + current = target - 1 + ok, frame = cap.read() + if ok: + frames[len(valid_frame_indices)] = cv2.cvtColor( + frame, cv2.COLOR_BGR2RGB + ) + valid_frame_indices.append(target) + current = target + else: + current = None + + valid_num_frames = len(valid_frame_indices) + if valid_num_frames < len(wanted): + logger.warning( + "GLM video loading expected %d sampled frames but only loaded %d.", + len(wanted), + valid_num_frames, + ) + return frames[:valid_num_frames], valid_frame_indices + + @VIDEO_LOADER_REGISTRY.register( "glmga", video_processor="GlmgaVideoProcessor", diff --git a/vllm/multimodal/video_decoders/opencv.py b/vllm/multimodal/video_decoders/opencv.py index e46979dd1d4c..2b86d17a28d9 100644 --- a/vllm/multimodal/video_decoders/opencv.py +++ b/vllm/multimodal/video_decoders/opencv.py @@ -42,7 +42,7 @@ def decode_opencv( frame_idx = loader_cls.compute_frames_index_to_sample( source=source, target=target, **sampling_kwargs ) - frames, valid = OpenCVVideoBackendMixin.read_frames( + frames, valid = loader_cls.read_frames( cap, frame_idx, total_frames_num=source.total_frames_num, diff --git a/vllm/platforms/cuda.py b/vllm/platforms/cuda.py index 0aff4ff9bef7..d1df152c63de 100644 --- a/vllm/platforms/cuda.py +++ b/vllm/platforms/cuda.py @@ -86,6 +86,7 @@ def _get_backend_priorities( num_heads: int | None = None, kv_cache_dtype: CacheDType | None = None, use_non_causal: bool = False, + head_size: int | None = None, ) -> list[AttentionBackendEnum]: """Get backend priorities with lazy import to avoid circular dependency.""" from vllm.utils.torch_utils import is_quantized_kv_cache @@ -133,13 +134,21 @@ def _get_backend_priorities( AttentionBackendEnum.FLASHINFER_MLA_SPARSE_SM120, ] else: + sparse_tail = [ + AttentionBackendEnum.FLASH_ATTN_MLA_SPARSE, + AttentionBackendEnum.FLASHMLA_SPARSE, + ] + flashinfer_sparse = AttentionBackendEnum.FLASHINFER_MLA_SPARSE_SM90 + if head_size == 512: + sparse_tail.insert(0, flashinfer_sparse) + else: + sparse_tail.append(flashinfer_sparse) return [ AttentionBackendEnum.FLASH_ATTN_MLA, AttentionBackendEnum.FLASHMLA, AttentionBackendEnum.FLASHINFER_MLA, AttentionBackendEnum.TRITON_MLA, - AttentionBackendEnum.FLASH_ATTN_MLA_SPARSE, - AttentionBackendEnum.FLASHMLA_SPARSE, + *sparse_tail, ] else: # SM100f defaults to FlashInfer for TRTLLM causal attention, but its non-causal @@ -373,11 +382,12 @@ def get_valid_backends( invalid_reasons: dict[AttentionBackendEnum, tuple[int, list[str]]] = {} backend_priorities = _get_backend_priorities( - attn_selector_config.use_mla, - device_capability, - num_heads, - attn_selector_config.kv_cache_dtype, - attn_selector_config.use_non_causal, + use_mla=attn_selector_config.use_mla, + device_capability=device_capability, + num_heads=num_heads, + kv_cache_dtype=attn_selector_config.kv_cache_dtype, + use_non_causal=attn_selector_config.use_non_causal, + head_size=attn_selector_config.head_size, ) for priority, backend in enumerate(backend_priorities): try: @@ -400,6 +410,26 @@ def get_valid_backends( return valid_backends_priorities, invalid_reasons + @classmethod + def _get_indexer_block_alignment(cls, vllm_config: VllmConfig) -> int | None: + index_kpool = getattr( + vllm_config.model_config.hf_text_config, "index_kpool", None + ) + if not index_kpool or index_kpool <= 1: + return None + from vllm.utils.deep_gemm import PAGED_MQA_PAGE_SIZES + + # kpool paged-MQA indexer: the storage block (block_size / + # index_kpool) is virtually split into pool pages, so block_size + # must be a multiple of index_kpool times a legal pool page. + page = min(PAGED_MQA_PAGE_SIZES) + if cls.is_device_capability_family(120): + # On sm120 the DeepGEMM paged-MQA kernel only accepts block_kv + # 64 for the fp8 indexer cache, so align to the largest pool + # page here to make the page split land on 64 not the min 32. + page = max(PAGED_MQA_PAGE_SIZES) + return index_kpool * page + @classmethod def get_attn_backend_cls( cls, diff --git a/vllm/platforms/interface.py b/vllm/platforms/interface.py index d7a31323ba29..7ad0b42da774 100644 --- a/vllm/platforms/interface.py +++ b/vllm/platforms/interface.py @@ -763,6 +763,17 @@ def per_token_page_bytes(dtype: "torch.dtype", cache_dtype: str) -> int: if cache_config.mamba_page_size_padded is not None: cache_config.mamba_page_size_padded = shared_page + @classmethod + def _get_indexer_block_alignment(cls, vllm_config: "VllmConfig") -> int | None: + """Extra ``block_size`` multiple a sparse indexer needs, else ``None``. + + The CUDA kpool paged-MQA indexer virtually splits each storage block + into pool pages, so ``block_size`` must be a multiple of + ``index_kpool * min(PAGED_MQA_PAGE_SIZES)`` — implemented in the CUDA + platform override. Other platforms impose no extra constraint. + """ + return None + @classmethod def _align_hybrid_block_size( cls, @@ -805,6 +816,7 @@ def _align_hybrid_block_size( num_kv_heads=model_config.get_num_kv_heads(parallel_config), head_size=model_config.get_head_size(), dtype=kv_cache_dtype, + cache_dtype_str=cache_config.cache_dtype, kv_quant_mode=kv_quant_mode, ).page_size_bytes elif cache_config.cache_dtype.startswith("turboquant_"): @@ -912,6 +924,9 @@ def _align_hybrid_block_size( mamba_page_size, kernel_block_alignment_size * attn_page_size_1_token, ) + indexer_align = cls._get_indexer_block_alignment(vllm_config) + if indexer_align: + attn_block_size = indexer_align * cdiv(attn_block_size, indexer_align) if cache_config.block_size < attn_block_size: cache_config.block_size = attn_block_size diff --git a/vllm/transformers_utils/config.py b/vllm/transformers_utils/config.py index dcbc0ef9f79a..692d3050e111 100644 --- a/vllm/transformers_utils/config.py +++ b/vllm/transformers_utils/config.py @@ -104,6 +104,9 @@ def __getitem__(self, key): k3_dspark="K3DSparkConfig", funaudiochat="FunAudioChatConfig", granite4_vision="Granite4VisionConfig", + glm5_next="Glm5NextConfig", + glm5_next_text="Glm5NextTextConfig", + glm5_next_vision="Glm5NextVisionConfig", hyperclovax="HyperCLOVAXConfig", hy_v3="HYV3Config", hy_v4="HYV4Config", diff --git a/vllm/transformers_utils/configs/__init__.py b/vllm/transformers_utils/configs/__init__.py index 257b0e7792ff..9ddbb003fc68 100644 --- a/vllm/transformers_utils/configs/__init__.py +++ b/vllm/transformers_utils/configs/__init__.py @@ -40,6 +40,9 @@ "FunAudioChatConfig": "vllm.transformers_utils.configs.funaudiochat", "FunAudioChatAudioEncoderConfig": "vllm.transformers_utils.configs.funaudiochat", "Granite4VisionConfig": "vllm.transformers_utils.configs.granite4_vision", + "Glm5NextConfig": "vllm.transformers_utils.configs.glm5_next", + "Glm5NextTextConfig": "vllm.transformers_utils.configs.glm5_next", + "Glm5NextVisionConfig": "vllm.transformers_utils.configs.glm5_next", "HYV3Config": "vllm.transformers_utils.configs.hy_v3", "HYV4Config": "vllm.transformers_utils.configs.hy_v4", "HyperCLOVAXConfig": "vllm.transformers_utils.configs.hyperclovax", @@ -133,6 +136,9 @@ "FunAudioChatConfig", "FunAudioChatAudioEncoderConfig", "Granite4VisionConfig", + "Glm5NextConfig", + "Glm5NextTextConfig", + "Glm5NextVisionConfig", "HYV3Config", "HYV4Config", "HyperCLOVAXConfig", diff --git a/vllm/transformers_utils/configs/glm5_next.py b/vllm/transformers_utils/configs/glm5_next.py new file mode 100644 index 000000000000..95a0081bdd45 --- /dev/null +++ b/vllm/transformers_utils/configs/glm5_next.py @@ -0,0 +1,420 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project + +from transformers.configuration_utils import PretrainedConfig + +from vllm.logger import init_logger + +logger = init_logger(__name__) + + +class Glm5NextTextConfig(PretrainedConfig): + model_type = "glm5_next_text" + base_config_key = "text_config" + keys_to_ignore_at_inference = ["past_key_values"] + + def __init__( + self, + model_type="glm5_next_text", + vocab_size: int = 154880, + hidden_size: int = 4096, + head_dim: int | None = None, + intermediate_size: int = 12288, + num_hidden_layers: int = 45, + num_attention_heads: int = 64, + num_key_value_heads: int | None = None, + hidden_act: str = "silu", + rms_norm_eps: float = 1e-5, + pad_token_id: int | None = 151329, + bos_token_id: int | None = None, + eos_token_id: int | list[int] | None = None, + rope_parameters: dict | None = None, + max_position_embeddings: int = 1048576, + tie_word_embeddings: bool = False, + moe_intermediate_size: int = 2048, + moe_renormalize: bool = True, + scoring_func: str = "sigmoid", + n_routed_experts: int | None = 288, + num_experts_per_token: int = 7, + n_shared_experts: int = 1, + routed_scaling_factor: float = 2.5, + topk_method: str | None = None, + first_k_dense_replace: int = 0, + moe_layer_freq: int = 1, + use_grouped_topk: bool = True, + n_group: int = 1, + topk_group: int = 1, + mla: bool = True, + q_lora_rank: int | None = 1536, + kv_lora_rank: int | None = 512, + qk_nope_head_dim: int = 256, + qk_rope_head_dim: int = 0, + v_head_dim: int | None = 256, + mla_nope: bool | None = True, + num_nextn_predict_layers: int = 1, + # Per-layer layout: "linear_attention" | "deepseek_sparse_attention" + layer_types: list[str] | None = None, + # Per-layer MLP: "dense" | "sparse" + mlp_layer_types: list[str] | None = None, + # Linear-attention (KDA) head config (flattened from the old + # linear_attn_config dict). + linear_head_dim: int = 128, + linear_num_heads: int = 64, + linear_conv_kernel_dim: int = 4, + linear_lower_bound: float = -5.0, + index_head_dim: int | None = None, + index_topk: int | None = None, + index_n_heads: int | None = None, + index_dsa_use_layernorm: bool = True, + index_kpool_compress: bool = True, + # Every ``index_kpool`` indexer K entries pool into one stored entry + # (compress_ratio). The 300B checkpoint ships 4; topk runs at pool + # granularity (select_k = index_topk // index_kpool). + index_kpool: int | None = 4, + index_kpool_always_select_tail: bool = True, + indexer_rope_interleave: bool = False, + mhc: bool | None = True, + mhc_num_residual_streams: int = 4, + hc_eps: float | None = 1e-06, + mhc_tau: float = 0.05, + hres_vwnstyle: bool | None = True, + mhc_no_norm_weight: bool | None = False, + mhc_sinkhorn_iterations: int | None = 20, + mhc_post_mult_value: float | None = 2.0, + swiglu_limit: float | None = None, + logit_scale: float = 1.0, + **kwargs, + ): + # Preserve checkpoint field names and local aliases because their + # consumers use different spellings. + num_experts_per_token = kwargs.get("num_experts_per_tok", num_experts_per_token) + moe_renormalize = kwargs.get("norm_topk_prob", moe_renormalize) + mhc_num_residual_streams = kwargs.get("hc_mult", mhc_num_residual_streams) + mhc_sinkhorn_iterations = kwargs.get( + "hc_sinkhorn_iters", mhc_sinkhorn_iterations + ) + # Checkpoint ships ``mla_use_nope`` (not ``mla_nope``); without this + # alias self.mla_nope silently stays at the param default. + mla_nope = kwargs.get("mla_use_nope", mla_nope) + # Checkpoints ship the KDA head config as the ``linear_attn_config`` + # dict (head_dim / num_heads / short_conv_kernel_size / + # gate_lower_bound) rather than the flattened top-level fields; fold it + # in so the trained values are read instead of the param defaults + # (which only match this checkpoint by coincidence). + linear_cfg = kwargs.get("linear_attn_config") or {} + if linear_cfg: + linear_head_dim = linear_cfg.get("head_dim", linear_head_dim) + linear_num_heads = linear_cfg.get("num_heads", linear_num_heads) + linear_conv_kernel_dim = linear_cfg.get( + "short_conv_kernel_size", linear_conv_kernel_dim + ) + linear_lower_bound = linear_cfg.get("gate_lower_bound", linear_lower_bound) + + if index_topk is not None: + if index_dsa_use_layernorm is not True: + raise NotImplementedError( + "GLM-5.3 sparse indexer requires index_dsa_use_layernorm=True" + ) + if index_kpool_compress is not True: + raise NotImplementedError( + "GLM-5.3 sparse indexer requires index_kpool_compress=True" + ) + if index_kpool_always_select_tail is not True: + raise NotImplementedError( + "GLM-5.3 sparse indexer requires " + "index_kpool_always_select_tail=True" + ) + + if mhc: + if hres_vwnstyle is not True: + raise NotImplementedError("GLM-5.3 mHC requires hres_vwnstyle=True") + if mhc_no_norm_weight not in (False, None): + raise NotImplementedError( + "GLM-5.3 mHC requires mhc_no_norm_weight=False" + ) + + self.model_type = model_type + self.vocab_size = vocab_size + self.hidden_size = hidden_size + self.head_dim = ( + head_dim if head_dim is not None else hidden_size // num_attention_heads + ) + self.intermediate_size = intermediate_size + self.num_hidden_layers = num_hidden_layers + self.num_attention_heads = num_attention_heads + + # for backward compatibility + if num_key_value_heads is None: + num_key_value_heads = num_attention_heads + + self.num_key_value_heads = num_key_value_heads + self.hidden_act = hidden_act + self.rms_norm_eps = rms_norm_eps + self.max_position_embeddings = max_position_embeddings + self.rope_parameters = rope_parameters + + # mla config + self.mla = mla + self.q_lora_rank = q_lora_rank + self.kv_lora_rank = kv_lora_rank + self.qk_nope_head_dim = qk_nope_head_dim + self.qk_rope_head_dim = qk_rope_head_dim + self.v_head_dim = v_head_dim + self.mla_nope = mla_nope + # moe config + self.n_routed_experts = n_routed_experts + self.num_experts_per_token = num_experts_per_token + self.moe_renormalize = moe_renormalize + self.n_shared_experts = n_shared_experts + self.routed_scaling_factor = routed_scaling_factor + self.topk_method = topk_method + self.scoring_func = scoring_func + assert self.scoring_func in ("softmax", "sigmoid") + self.moe_intermediate_size = moe_intermediate_size + self.first_k_dense_replace = first_k_dense_replace + self.moe_layer_freq = moe_layer_freq + self.use_grouped_topk = use_grouped_topk + self.n_group = n_group + self.topk_group = topk_group + self.num_nextn_predict_layers = num_nextn_predict_layers + + # Per-layer attention / MLP layout. Normalize mlp_layer_types from + # first_k_dense_replace when the new-schema field is absent so layer + # construction sees a consistent layout (mirrors cohere2_moe). + self.layer_types = layer_types + if mlp_layer_types is None: + n = self.num_hidden_layers + if first_k_dense_replace is not None: + mlp_layer_types = ["dense"] * first_k_dense_replace + ["sparse"] * ( + n - first_k_dense_replace + ) + else: + mlp_layer_types = ["sparse"] * n + self.mlp_layer_types = mlp_layer_types + + # Linear-attention (KDA) head config. + self.linear_head_dim = linear_head_dim + self.linear_num_heads = linear_num_heads + self.linear_conv_kernel_dim = linear_conv_kernel_dim + self.linear_lower_bound = linear_lower_bound + + # dsa index config + self.index_head_dim = index_head_dim + self.index_topk = index_topk + self.index_n_heads = index_n_heads + self.index_dsa_use_layernorm = index_dsa_use_layernorm + self.index_kpool_compress = index_kpool_compress + self.index_kpool = index_kpool + self.index_kpool_always_select_tail = index_kpool_always_select_tail + self.indexer_rope_interleave = indexer_rope_interleave + + # mhc config + self.mhc = mhc + self.mhc_num_residual_streams = mhc_num_residual_streams + self.mhc_tau = mhc_tau + self.hres_vwnstyle = hres_vwnstyle + self.hc_eps = hc_eps + self.mhc_no_norm_weight = mhc_no_norm_weight + self.mhc_sinkhorn_iterations = mhc_sinkhorn_iterations + self.mhc_post_mult_value = mhc_post_mult_value + + self.swiglu_limit = swiglu_limit + self.logit_scale = logit_scale + + super().__init__( + pad_token_id=pad_token_id, + bos_token_id=bos_token_id, + eos_token_id=eos_token_id, + tie_word_embeddings=tie_word_embeddings, + **kwargs, + ) + + @property + def is_mla(self): + return ( + self.q_lora_rank is not None + or self.kv_lora_rank is not None + or self.qk_nope_head_dim is not None + or self.qk_rope_head_dim is not None + or self.v_head_dim is not None + or self.mla_nope is True + ) + + @property + def is_moe(self): + return self.n_routed_experts is not None + + @property + def is_linear_attn(self) -> bool: + return self.layer_types is not None and any( + t == "linear_attention" for t in self.layer_types + ) + + def is_kda_layer(self, layer_idx: int): + return ( + self.layer_types is not None + and layer_idx < len(self.layer_types) + and self.layer_types[layer_idx] == "linear_attention" + ) + + @property + def layers_block_type(self): + # Map the schema's per-layer types onto the block strings vLLM's hybrid + # accounting (get_num_layers_by_block_type) recognizes: linear-attention + # layers stay "linear_attention"; every other attention variant collapses + # to "attention". + if self.layer_types is None: + return ["attention"] * self.num_hidden_layers + return [ + "linear_attention" if t == "linear_attention" else "attention" + for t in self.layer_types + ] + + +class Glm5NextVisionConfig(PretrainedConfig): + model_type = "glm5_next_vision" + base_config_key = "vision_config" + + def __init__( + self, + depth: int = 24, + hidden_size: int = 1024, + hidden_act: str = "silu", + image_size: int = 448, + intermediate_size: int = 4096, + num_heads: int = 16, + out_hidden_size: int = 4096, + projection_intermediate_size: int = 10240, + in_channels: int = 3, + initializer_range: float = 0.02, + patch_size: int = 14, + rms_norm_eps: float = 1e-5, + spatial_merge_size: int = 2, + temporal_patch_size: int = 2, + attention_dropout: float = 0.0, + attention_bias: bool = True, + swiglu_limit: float | None = None, + **kwargs, + ): + super().__init__(**kwargs) + + self.depth = depth + self.hidden_size = hidden_size + self.hidden_act = hidden_act + self.image_size = image_size + self.intermediate_size = intermediate_size + self.num_heads = num_heads + self.out_hidden_size = out_hidden_size + # GLM-5.3-Flash merger bottleneck width (absent from the generic + # GLM-OCR vision config); the tower uses it as the PatchMerger + # context_dim instead of text_config.intermediate_size. + self.projection_intermediate_size = projection_intermediate_size + self.in_channels = in_channels + self.initializer_range = initializer_range + self.patch_size = patch_size + # GLM-5.3-Flash checkpoints ship vision_config.rms_norm_eps = 1e-5, + # but the vision tower was trained with 1e-6. Serving with 1e-5 drifts + # the RMSNorm and produces repetitive/degraded image descriptions, so + # force the trained value regardless of the checkpoint field. + self.rms_norm_eps = 1e-6 + self.spatial_merge_size = spatial_merge_size + self.temporal_patch_size = temporal_patch_size + self.attention_dropout = attention_dropout + self.attention_bias = attention_bias + self.swiglu_limit = swiglu_limit + + +class Glm5NextConfig(PretrainedConfig): + model_type = "glm5_next" + sub_configs = { + "vision_config": Glm5NextVisionConfig, + "text_config": Glm5NextTextConfig, + } + keys_to_ignore_at_inference = ["past_key_values"] + + def __init__( + self, + text_config=None, + vision_config=None, + image_token_id: int = 154854, + video_token_id: int = 154855, + image_start_token_id: int = 154830, + image_end_token_id: int = 154831, + video_start_token_id: int = 154832, + video_end_token_id: int = 154833, + **kwargs, + ): + # Init super() first so base-class defaults don't clobber text-config + # values set below (PretrainedConfig has many text-related defaults + # that differ from Glm5NextTextConfig). + super().__init__(**kwargs) + + if isinstance(vision_config, dict): + self.vision_config = self.sub_configs["vision_config"](**vision_config) + elif vision_config is None: + self.vision_config = self.sub_configs["vision_config"]() + else: + self.vision_config = vision_config + + if isinstance(text_config, dict): + self.text_config = self.sub_configs["text_config"](**text_config) + elif text_config is None: + # Backward compatibility: a flat top-level checkpoint (no nested + # text_config) folds its text fields into Glm5NextTextConfig. + self.text_config = self.sub_configs["text_config"](**kwargs) + else: + self.text_config = text_config + + self.image_token_id = image_token_id + self.video_token_id = video_token_id + self.image_start_token_id = image_start_token_id + self.image_end_token_id = image_end_token_id + self.video_start_token_id = video_start_token_id + self.video_end_token_id = video_end_token_id + + # Mirror attention implementation recursively onto sub-configs. + self._attn_implementation = kwargs.pop("attn_implementation", None) + + # Config-metadata fields that belong to the top-level (multimodal) config + # and must NOT be mirrored onto text_config: ``architectures`` / + # ``torch_dtype`` differ between the top-level config and the text + # sub-config, and mirroring them makes the top-level ``architectures`` + # silently read back as None (PretrainedConfig initializes both to None), + # which then fails model-class resolution ("No model architectures are + # specified"). + _UNMIRRORED_KEYS = [ + "_name_or_path", + "model_type", + "dtype", + "torch_dtype", + "architectures", + "_attn_implementation_internal", + ] + + def __setattr__(self, key, value): + unmirrored = type(self)._UNMIRRORED_KEYS + if ( + (text_config := super().__getattribute__("__dict__").get("text_config")) + is not None + and key not in unmirrored + and key in text_config.__dict__ + ): + setattr(text_config, key, value) + else: + super().__setattr__(key, value) + + def __getattribute__(self, key): + unmirrored = type(self)._UNMIRRORED_KEYS + if ( + "text_config" in super().__getattribute__("__dict__") + and key not in unmirrored + ): + text_config = super().__getattribute__("text_config") + # Forward both instance attributes AND class-defined properties/ + # methods of the text config, so a flat text-only checkpoint + # (model_type "glm5_next", no nested text_config) sees is_moe / + # is_kda_layer / layers_block_type like a Glm5NextTextConfig. + if key in text_config.__dict__ or key in type(text_config).__dict__: + return getattr(text_config, key) + + return super().__getattribute__(key) diff --git a/vllm/transformers_utils/model_arch_config_convertor.py b/vllm/transformers_utils/model_arch_config_convertor.py index 8baa9ee16288..fa1af93f9d59 100644 --- a/vllm/transformers_utils/model_arch_config_convertor.py +++ b/vllm/transformers_utils/model_arch_config_convertor.py @@ -320,6 +320,8 @@ def is_deepseek_mla(self) -> bool: "dots3_note", "deepseek_mtp", "k3_dspark", + "glm5_next", + "glm5_next_text", "glm_moe_dsa", "glm4_moe_lite", "glm4_moe_lite_mtp", diff --git a/vllm/transformers_utils/processors/__init__.py b/vllm/transformers_utils/processors/__init__.py index b5e46077615a..249a4e58e662 100644 --- a/vllm/transformers_utils/processors/__init__.py +++ b/vllm/transformers_utils/processors/__init__.py @@ -18,6 +18,7 @@ "FireRedASR2Processor", "FunASRProcessor", "GLM4VProcessor", + "Glm5NextProcessor", "Granite4VisionProcessor", "H2OVLProcessor", "Moondream3Processor", @@ -56,6 +57,7 @@ "FireRedASR2Processor": "vllm.transformers_utils.processors.fireredasr2", "FunASRProcessor": "vllm.transformers_utils.processors.funasr", "GLM4VProcessor": "vllm.transformers_utils.processors.glm4v", + "Glm5NextProcessor": "vllm.transformers_utils.processors.glm5next", "Granite4VisionProcessor": "vllm.transformers_utils.processors.granite4_vision", "H2OVLProcessor": "vllm.transformers_utils.processors.h2ovl", "InternVLProcessor": "vllm.transformers_utils.processors.internvl", diff --git a/vllm/transformers_utils/processors/glm5next.py b/vllm/transformers_utils/processors/glm5next.py new file mode 100644 index 000000000000..b6988d3b8fed --- /dev/null +++ b/vllm/transformers_utils/processors/glm5next.py @@ -0,0 +1,953 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +"""vLLM-native multimodal processor for GLM-5.3-Flash.""" + +import math + +import numpy as np +import torch +from torchvision.transforms.v2 import functional as tvF +from transformers.image_processing_utils import BatchFeature +from transformers.image_processing_utils_fast import ( + BaseImageProcessorFast, + group_images_by_shape, + reorder_images, +) +from transformers.image_utils import ( + OPENAI_CLIP_MEAN, + OPENAI_CLIP_STD, + ChannelDimension, + ImageInput, + PILImageResampling, + SizeDict, + get_image_size, +) +from transformers.models.auto.image_processing_auto import get_image_processor_config +from transformers.processing_utils import ( + ImagesKwargs, + MultiModalData, + ProcessingKwargs, + ProcessorMixin, + Unpack, + VideosKwargs, +) +from transformers.tokenization_utils_base import PreTokenizedInput, TextInput +from transformers.utils import TensorType, logging +from transformers.video_processing_utils import BaseVideoProcessor +from transformers.video_utils import ( + VideoInput, + VideoMetadata, + group_videos_by_shape, + reorder_videos, +) + +from vllm.transformers_utils.repo_utils import get_hf_file_to_dict + +logger = logging.get_logger(__name__) + +# Cap video inputs at 30,000 vision tokens to keep encoder profiling from +# starving the KV cache. Image inputs retain their checkpoint-defined budget. +_MAX_VIDEO_TOKENS = 30000 + +# Frame-sampler fallbacks (mirror the checkpoint's fps_interval=2 / +# max_frame_count_dynamic=2048); used when neither the request nor the +# processor config overrides them. +GLM_VIDEO_DEFAULT_FPS = 2.0 +GLM_VIDEO_DEFAULT_MAX_FRAMES = 2048 + + +def glm_sample_frame_indices( + total_frames: int, + fps: float, + duration: float, + *, + target_fps: float | None = None, + max_frame_count: int | None = None, + temporal_patch_size: int = 2, +) -> list[int]: + """GLM video frame sampling (training-reference parity). + + ``target_fps`` is the ``fps_interval`` request knob. The greedy walk + advances at ``1 / (temporal_patch_size * target_fps)`` seconds, so on + frame-dense sources it collects more candidates than ``extract_t`` and + the ``> extract_t`` fixup re-spreads the picks uniformly with + ``np.linspace`` -- that fallback is the intended reference behavior, not + an accident. Short clips (fewer frames than ``extract_t``) are spread at + evenly spaced timestamps (``floor`` sampling; the linspace variant + samples frames unevenly and cost 4 points on video grounding evals). + Request overrides: ``target_fps`` -> fps interval, ``max_frame_count`` + -> frame cap. + """ + max_frame_idx = total_frames - 1 + if not duration: + duration = (round(max_frame_idx / fps) + 1) if fps else 0 + if max_frame_count is None: + max_frame_count = GLM_VIDEO_DEFAULT_MAX_FRAMES + if target_fps is None: + target_fps = GLM_VIDEO_DEFAULT_FPS + + extract_t = int(duration * target_fps) + extract_t = min(extract_t, int(max_frame_count)) + + duration_per_frame = 1 / fps + timestamps = [i * duration_per_frame for i in range(total_frames)] + max_second = int(duration) + + if total_frames < extract_t: + frame_indices = [ + math.floor(_i * total_frames / extract_t) for _i in range(extract_t) + ] + else: + frame_indices = [] + current_second = 0.0 + inv_fps = 1 / (temporal_patch_size * target_fps) + for frame_index in range(total_frames): + if timestamps[frame_index] >= current_second: + current_second += inv_fps + frame_indices.append(frame_index) + if current_second >= max_second: + break + + if len(frame_indices) < extract_t: + if len(frame_indices) == 0: + start, end = 0, max(total_frames - 1, 0) + else: + start, end = frame_indices[0], frame_indices[-1] + frame_indices = np.linspace(start, end, extract_t, dtype=int).tolist() + elif len(frame_indices) > extract_t: + frame_indices = np.linspace(0, total_frames - 1, extract_t, dtype=int).tolist() + + seen, uniq = set(), [] + for idx in frame_indices: + if idx not in seen: + seen.add(idx) + uniq.append(int(idx)) + + if len(uniq) & 1: + uniq.append(uniq[-1]) + + return uniq + + +def _ceil_to_factor(value: int, factor: int) -> int: + """Round a positive integer upward to the nearest multiple of factor.""" + return math.ceil(value / factor) * factor + + +def _fit_aligned_size_within_budget( + t: int, + h: int, + w: int, + h_factor: int, + w_factor: int, + max_pixels: int, +) -> tuple[int, int]: + """Largest proportional size whose upward-aligned canvas fits the budget. + + Binary search on the unaligned content height; each candidate is rounded + upward to h_factor/w_factor, so the returned canvas always satisfies + ``t * aligned_h * aligned_w <= max_pixels``. + """ + minimum_pixels = t * h_factor * w_factor + if max_pixels < minimum_pixels: + raise ValueError( + f"max_pixels={max_pixels} is too small. At least " + f"{minimum_pixels} pixels are required for one aligned patch." + ) + + low, high = 1, h + best_h, best_w = h_factor, w_factor + while low <= high: + content_h = (low + high) // 2 + content_w = max(1, math.floor(w * content_h / h)) + aligned_h = _ceil_to_factor(content_h, h_factor) + aligned_w = _ceil_to_factor(content_w, w_factor) + if t * aligned_h * aligned_w <= max_pixels: + best_h, best_w = aligned_h, aligned_w + low = content_h + 1 + else: + high = content_h - 1 + return best_h, best_w + + +def smart_resize( + t: int, + h: int, + w: int, + t_factor: int = 1, + h_factor: int = 28, + w_factor: int = 28, + min_pixels: int = 56 * 56, + max_pixels: int = 14 * 14 * 4 * 1280, +) -> tuple[int, int]: + """GLM-5.3-Flash ``smart_resize``: upward-aligned canvas under a + ``t_bar * h_bar * w_bar`` pixel budget. + + Height/width always round UP to their factors (content is then padded, + never cropped or distorted); an over-budget canvas is refit by binary + search instead of one-shot square-root scaling. ``h_factor`` / + ``w_factor`` carry ``patch_expand_factor`` on top of + ``patch_size * merge_size``; ``t_factor`` is ``temporal_patch_size``. For + a still image ``t = t_factor = temporal_patch_size`` so ``t_bar = + temporal_patch_size``. + """ + if min(t, h, w, t_factor, h_factor, w_factor) <= 0: + raise ValueError("Image dimensions and alignment factors must be positive.") + if min_pixels <= 0 or max_pixels <= 0: + raise ValueError("min_pixels and max_pixels must be positive.") + if min_pixels > max_pixels: + raise ValueError("min_pixels must be less than or equal to max_pixels.") + + t_bar = max(t_factor, round(t / t_factor) * t_factor) + h_bar = _ceil_to_factor(h, h_factor) + w_bar = _ceil_to_factor(w, w_factor) + + if t_bar * h_bar * w_bar > max_pixels: + h_bar, w_bar = _fit_aligned_size_within_budget( + t=t_bar, + h=h, + w=w, + h_factor=h_factor, + w_factor=w_factor, + max_pixels=max_pixels, + ) + elif t_bar * h_bar * w_bar < min_pixels: + beta = math.sqrt(min_pixels / (t * h * w)) + h_bar = _ceil_to_factor(max(1, math.ceil(h * beta)), h_factor) + w_bar = _ceil_to_factor(max(1, math.ceil(w * beta)), w_factor) + + # Alignment can push a candidate slightly over a tight max_pixels + # budget. Refit it when that happens. + if t_bar * h_bar * w_bar > max_pixels: + h_bar, w_bar = _fit_aligned_size_within_budget( + t=t_bar, + h=h, + w=w, + h_factor=h_factor, + w_factor=w_factor, + max_pixels=max_pixels, + ) + + return h_bar, w_bar + + +def _get_pad_content_size( + image_height: int, + image_width: int, + canvas_height: int, + canvas_width: int, + allow_upscale: bool = False, +) -> tuple[int, int]: + """Aspect-ratio-preserving content size that fits the canvas. + + Oversized images are shrunk proportionally. Small images are enlarged + only when ``allow_upscale``. Padding is applied after the resize. + """ + scale = min(canvas_height / image_height, canvas_width / image_width) + if not allow_upscale: + scale = min(1.0, scale) + content_height = max(1, min(canvas_height, math.floor(image_height * scale))) + content_width = max(1, min(canvas_width, math.floor(image_width * scale))) + return content_height, content_width + + +def _resize_or_pad( + stacked_images: torch.Tensor, + target_height: int, + target_width: int, + resize_mode: str, + resample: "PILImageResampling | tvF.InterpolationMode | int | None", + resize, + allow_upscale: bool = False, +) -> torch.Tensor: + """Resize onto the aligned canvas, or keep the aspect ratio and + zero-pad the right/bottom sides (``resize_mode="pad"``).""" + height, width = stacked_images.shape[-2:] + + if resize_mode == "resize": + return resize( + stacked_images, + size=SizeDict(height=target_height, width=target_width), + resample=resample, + ) + + if resize_mode != "pad": + raise ValueError("resize_mode must be either 'resize' or 'pad'.") + + content_height, content_width = _get_pad_content_size( + image_height=height, + image_width=width, + canvas_height=target_height, + canvas_width=target_width, + allow_upscale=allow_upscale, + ) + + if (content_height, content_width) != (height, width): + stacked_images = resize( + stacked_images, + size=SizeDict(height=content_height, width=content_width), + resample=resample, + ) + + # torchvision padding order: [left, top, right, bottom] -> pad only the + # right and bottom sides. + return tvF.pad( + stacked_images, + padding=[0, 0, target_width - content_width, target_height - content_height], + fill=0, + ) + + +def _pixel_budget( + min_image_tokens: int | None, + max_image_tokens: int | None, + patch_size: int, + merge_size: int, + temporal_patch_size: int, +) -> tuple[int, int]: + """(min_pixels, max_pixels) from the token bounds of + ``processor_config.json``; one vision token covers + ``temporal_patch_size * (patch_size * merge_size) ** 2`` pixels.""" + if min_image_tokens is None or max_image_tokens is None: + raise ValueError( + "min_image_tokens and max_image_tokens must be provided by " + "processor_config.json (or per-call kwargs)." + ) + factor = temporal_patch_size * (patch_size * merge_size) ** 2 + return min_image_tokens * factor, max_image_tokens * factor + + +class Glm5NextImageProcessorKwargs(ImagesKwargs, total=False): # type: ignore[call-arg] + patch_size: int | None + temporal_patch_size: int | None + merge_size: int | None + patch_expand_factor: int | None + resize_mode: str | None + min_image_tokens: int | None + max_image_tokens: int | None + + +class Glm5NextImageProcessor(BaseImageProcessorFast): + """Fast torchvision image processor for GLM-5.3-Flash. + + ``patch_expand_factor`` multiplies into the ``smart_resize`` spatial + factor ``patch_size * merge_size``. ``resize_mode`` picks the geometry: + ``"pad"`` (default) preserves the aspect ratio and zero-pads the + right/bottom of the upward-aligned canvas, ``"resize"`` stretches onto + it. Defaults mirror the checkpoint's ``image_processor`` config. + """ + + do_resize = True + resample = PILImageResampling.BICUBIC + size = {"longest_edge": 1} # unused: budgets come from the token bounds + do_rescale = True + do_normalize = True + image_mean = OPENAI_CLIP_MEAN + image_std = OPENAI_CLIP_STD + do_convert_rgb = True + patch_size = 14 + temporal_patch_size = 2 + merge_size = 2 + patch_expand_factor = 1 + resize_mode = "pad" + min_image_tokens = 16 + max_image_tokens = 8000 + valid_kwargs = Glm5NextImageProcessorKwargs + model_input_names = ["pixel_values", "image_grid_thw"] + + def _preprocess( + self, + images: list[torch.Tensor], + do_resize: bool, + size: SizeDict, + resample: "PILImageResampling | tvF.InterpolationMode | int | None", + do_rescale: bool, + rescale_factor: float, + do_normalize: bool, + image_mean: float | list[float] | None, + image_std: float | list[float] | None, + patch_size: int, + temporal_patch_size: int, + merge_size: int, + patch_expand_factor: int, + resize_mode: str | None, + min_image_tokens: int | None, + max_image_tokens: int | None, + disable_grouping: bool | None, + return_tensors: str | TensorType | None, + **kwargs, + ) -> BatchFeature: + resize_mode = resize_mode if resize_mode is not None else self.resize_mode + min_pixels, max_pixels = _pixel_budget( + min_image_tokens if min_image_tokens is not None else self.min_image_tokens, + max_image_tokens if max_image_tokens is not None else self.max_image_tokens, + patch_size, + merge_size, + temporal_patch_size, + ) + grouped_images, grouped_images_index = group_images_by_shape( + images, disable_grouping=disable_grouping + ) + resized_images_grouped = {} + for shape, stacked_images in grouped_images.items(): + height, width = stacked_images.shape[-2:] + if do_resize: + resized_height, resized_width = smart_resize( + t=temporal_patch_size, + h=height, + w=width, + t_factor=temporal_patch_size, + h_factor=patch_size * merge_size * patch_expand_factor, + w_factor=patch_size * merge_size * patch_expand_factor, + min_pixels=min_pixels, + max_pixels=max_pixels, + ) + stacked_images = _resize_or_pad( + stacked_images, + target_height=resized_height, + target_width=resized_width, + resize_mode=resize_mode, + resample=resample, + resize=self.resize, + allow_upscale=(temporal_patch_size * height * width < min_pixels), + ) + resized_images_grouped[shape] = stacked_images + + resized_images = reorder_images(resized_images_grouped, grouped_images_index) + + grouped_images, grouped_images_index = group_images_by_shape( + resized_images, disable_grouping=disable_grouping + ) + processed_images_grouped = {} + processed_grids = {} + + for shape, stacked_images in grouped_images.items(): + resized_height, resized_width = stacked_images.shape[-2:] + + patches = self.rescale_and_normalize( + stacked_images, + do_rescale, + rescale_factor, + do_normalize, + image_mean, + image_std, + ) + if patches.ndim == 4: # (B, C, H, W) + patches = patches.unsqueeze(1) # (B, T=1, C, H, W) + + if patches.shape[1] % temporal_patch_size != 0: + repeats = patches[:, -1:].repeat( + 1, + temporal_patch_size - (patches.shape[1] % temporal_patch_size), + 1, + 1, + 1, + ) + patches = torch.cat([patches, repeats], dim=1) + + batch_size, t_len, channel = patches.shape[:3] + grid_t = t_len // temporal_patch_size + grid_h, grid_w = resized_height // patch_size, resized_width // patch_size + + patches = patches.view( + batch_size, + grid_t, + temporal_patch_size, + channel, + grid_h // merge_size, + merge_size, + patch_size, + grid_w // merge_size, + merge_size, + patch_size, + ) + # (B, grid_t, gh, gw, mh, mw, C, tp, ph, pw) + patches = patches.permute(0, 1, 4, 7, 5, 8, 3, 2, 6, 9) + + flatten_patches = patches.reshape( + batch_size, + grid_t * grid_h * grid_w, + channel * temporal_patch_size * patch_size * patch_size, + ) + + processed_images_grouped[shape] = flatten_patches + processed_grids[shape] = [[grid_t, grid_h, grid_w]] * batch_size + + processed_images = reorder_images( + processed_images_grouped, grouped_images_index + ) + processed_grids = reorder_images(processed_grids, grouped_images_index) + + pixel_values = torch.cat(processed_images, dim=0) + image_grid_thw = torch.tensor(processed_grids) + + return BatchFeature( + data={"pixel_values": pixel_values, "image_grid_thw": image_grid_thw}, + tensor_type=return_tensors, + ) + + def preprocess( + self, images: ImageInput, **kwargs: Unpack[Glm5NextImageProcessorKwargs] + ) -> BatchFeature: + return super().preprocess(images, **kwargs) + + def get_number_of_image_patches( + self, height: int, width: int, images_kwargs: dict | None = None + ) -> int: + """Number of image patches (pre-merge) for a given (height, width).""" + images_kwargs = images_kwargs or {} + patch_size = images_kwargs.get("patch_size", self.patch_size) + merge_size = images_kwargs.get("merge_size", self.merge_size) + patch_expand_factor = images_kwargs.get( + "patch_expand_factor", self.patch_expand_factor + ) + min_pixels, max_pixels = _pixel_budget( + images_kwargs.get("min_image_tokens", self.min_image_tokens), + images_kwargs.get("max_image_tokens", self.max_image_tokens), + patch_size, + merge_size, + self.temporal_patch_size, + ) + resized_height, resized_width = smart_resize( + t=self.temporal_patch_size, + h=height, + w=width, + t_factor=self.temporal_patch_size, + h_factor=patch_size * merge_size * patch_expand_factor, + w_factor=patch_size * merge_size * patch_expand_factor, + min_pixels=min_pixels, + max_pixels=max_pixels, + ) + grid_h, grid_w = resized_height // patch_size, resized_width // patch_size + return grid_h * grid_w + + +class Glm5NextVideoProcessorKwargs(VideosKwargs, total=False): # type: ignore[call-arg] + fps: list[float] | float + patch_size: int + temporal_patch_size: int + merge_size: int + patch_expand_factor: int + resize_mode: str | None + target_fps: float | None + max_frames: int | None + fps_interval: int | None + max_frame_count_dynamic: int | None + min_image_tokens: int | None + max_image_tokens: int | None + + +class Glm5NextVideoProcessor(BaseVideoProcessor): + """Fast video processor for GLM-5.3-Flash. + + Shares ``smart_resize`` / the pad-mode geometry / the patchify with the + image processor, and adds GLM-5.3-Flash frame sampling + (``glm_sample_frame_indices``: ``fps_interval`` semantics with a + temporal-patch-scaled greedy walk). Defaults mirror the checkpoint's + ``video_processor`` config. + """ + + resample = PILImageResampling.BICUBIC + size = {"longest_edge": 1} # unused: budgets come from the token bounds + image_mean = OPENAI_CLIP_MEAN + image_std = OPENAI_CLIP_STD + do_resize = True + do_rescale = True + do_normalize = True + do_convert_rgb = True + do_sample_frames = True + patch_size = 14 + temporal_patch_size = 2 + patch_expand_factor = 1 + merge_size = 2 + valid_kwargs = Glm5NextVideoProcessorKwargs + num_frames = 16 + fps = 2 + fps_interval = 2.0 + max_frame_count_dynamic = 2048 + resize_mode = "pad" + min_image_tokens = 16 + max_image_tokens = 240000 + model_input_names = ["pixel_values_videos", "video_grid_thw"] + + def sample_frames( + self, + metadata: VideoMetadata, + fps: int | float | None = None, + **kwargs, + ) -> np.ndarray: + """Sample frame indices with GLM's fps-interval policy. + + ``fps`` / ``target_fps``, ``max_frames`` and ``fps_interval`` / + ``max_frame_count_dynamic`` are the overrides described in + :func:`glm_sample_frame_indices`. + """ + if metadata is None or getattr(metadata, "fps", None) is None: + raise ValueError( + "Asked to sample frames per second but no video metadata was " + "provided which is required when sampling in GLM-5.3-Flash. Please " + "pass in `VideoMetadata` object or set `do_sample_frames=False`." + ) + + target_fps = fps if fps is not None else kwargs.get("target_fps") + if target_fps is None: + target_fps = self.fps_interval + indices = glm_sample_frame_indices( + metadata.total_num_frames, + metadata.fps, + metadata.duration or 0, + target_fps=target_fps, + max_frame_count=kwargs.get("max_frames") or self.max_frame_count_dynamic, + temporal_patch_size=self.temporal_patch_size, + ) + return np.array(indices) + + def _preprocess( + self, + videos: list[torch.Tensor], + do_convert_rgb: bool = True, + do_resize: bool = True, + size: SizeDict | None = None, + resample: "PILImageResampling | int | None" = PILImageResampling.BICUBIC, + do_rescale: bool = True, + rescale_factor: float = 1 / 255.0, + do_normalize: bool = True, + image_mean: float | list[float] | None = None, + image_std: float | list[float] | None = None, + patch_size: int | None = None, + temporal_patch_size: int | None = None, + patch_expand_factor: int | None = None, + merge_size: int | None = None, + resize_mode: str | None = None, + min_image_tokens: int | None = None, + max_image_tokens: int | None = None, + return_tensors: str | TensorType | None = None, + **kwargs, + ) -> BatchFeature: + patch_expand_factor = self.patch_expand_factor + patch_size = patch_size if patch_size is not None else self.patch_size + temporal_patch_size = ( + temporal_patch_size + if temporal_patch_size is not None + else self.temporal_patch_size + ) + merge_size = merge_size if merge_size is not None else self.merge_size + resize_mode = resize_mode if resize_mode is not None else self.resize_mode + min_pixels, max_pixels = _pixel_budget( + min_image_tokens if min_image_tokens is not None else self.min_image_tokens, + max_image_tokens if max_image_tokens is not None else self.max_image_tokens, + patch_size, + merge_size, + temporal_patch_size, + ) + grouped_videos, grouped_videos_index = group_videos_by_shape(videos) + resized_videos_grouped = {} + for shape, stacked_videos in grouped_videos.items(): + if do_convert_rgb: + stacked_videos = self.convert_to_rgb(stacked_videos) + b, t_len, c, h, w = stacked_videos.shape + num_frames, height, width = t_len, h, w + if do_resize: + resized_height, resized_width = smart_resize( + t=num_frames, + h=height, + w=width, + t_factor=temporal_patch_size, + h_factor=patch_size * merge_size * patch_expand_factor, + w_factor=patch_size * merge_size * patch_expand_factor, + min_pixels=min_pixels, + max_pixels=max_pixels, + ) + stacked_videos = stacked_videos.view(b * t_len, c, h, w) + stacked_videos = _resize_or_pad( + stacked_videos, + target_height=resized_height, + target_width=resized_width, + resize_mode=resize_mode, + resample=resample, + resize=self.resize, + allow_upscale=(num_frames * height * width < min_pixels), + ) + stacked_videos = stacked_videos.view( + b, t_len, c, resized_height, resized_width + ) + resized_videos_grouped[shape] = stacked_videos + resized_videos = reorder_videos(resized_videos_grouped, grouped_videos_index) + + grouped_videos, grouped_videos_index = group_videos_by_shape(resized_videos) + processed_videos_grouped = {} + processed_grids = {} + for shape, stacked_videos in grouped_videos.items(): + resized_height, resized_width = get_image_size( + stacked_videos[0], channel_dim=ChannelDimension.FIRST + ) + stacked_videos = self.rescale_and_normalize( + stacked_videos, + do_rescale, + rescale_factor, + do_normalize, + image_mean, + image_std, + ) + patches = stacked_videos + + if pad := -patches.shape[1] % temporal_patch_size: + repeats = patches[:, -1:].expand(-1, pad, -1, -1, -1) + patches = torch.cat((patches, repeats), dim=1) + batch_size, grid_t, channel = patches.shape[:3] + grid_t = grid_t // temporal_patch_size + grid_h, grid_w = resized_height // patch_size, resized_width // patch_size + + patches = patches.view( + batch_size, + grid_t, + temporal_patch_size, + channel, + grid_h // merge_size, + merge_size, + patch_size, + grid_w // merge_size, + merge_size, + patch_size, + ) + patches = patches.permute(0, 1, 4, 7, 5, 8, 3, 2, 6, 9) + flatten_patches = patches.reshape( + batch_size, + grid_t * grid_h * grid_w, + channel * temporal_patch_size * patch_size * patch_size, + ) + + processed_videos_grouped[shape] = flatten_patches + processed_grids[shape] = [[grid_t, grid_h, grid_w]] * batch_size + + processed_videos = reorder_videos( + processed_videos_grouped, grouped_videos_index + ) + processed_grids = reorder_videos(processed_grids, grouped_videos_index) + pixel_values_videos = torch.cat(processed_videos, dim=0) + video_grid_thw = torch.tensor(processed_grids) + return BatchFeature( + data={ + "pixel_values_videos": pixel_values_videos, + "video_grid_thw": video_grid_thw, + }, + tensor_type=return_tensors, + ) + + +class Glm5NextProcessorKwargs(ProcessingKwargs, total=False): # type: ignore[call-arg] + images_kwargs: Glm5NextImageProcessorKwargs + videos_kwargs: Glm5NextVideoProcessorKwargs + _defaults = { + "text_kwargs": { + "padding": False, + "return_token_type_ids": False, + "return_mm_token_type_ids": False, + }, + "videos_kwargs": {"return_metadata": True}, + } + + +class Glm5NextProcessor(ProcessorMixin): + """Wraps a GLM-5.3-Flash image processor, video processor and tokenizer. + + Token expansion per image = ``prod(image_grid_thw) // merge_size**2``; video + frames are expanded with ``<|begin_of_image|>...<|end_of_image|>{ts} seconds`` + structure (mrope timestamps). + """ + + attributes = ["image_processor", "tokenizer", "video_processor"] + image_processor_class = "AutoImageProcessor" + video_processor_class = "AutoVideoProcessor" + tokenizer_class = ("PreTrainedTokenizer", "PreTrainedTokenizerFast") + + def __init__( + self, + image_processor=None, + tokenizer=None, + video_processor=None, + chat_template=None, + **kwargs, + ) -> None: + super().__init__( + image_processor, tokenizer, video_processor, chat_template=chat_template + ) + self.image_token = ( + "<|image|>" + if not hasattr(tokenizer, "image_token") + else tokenizer.image_token + ) + self.video_token = ( + "<|video|>" + if not hasattr(tokenizer, "video_token") + else tokenizer.video_token + ) + self.image_token_id = ( + tokenizer.image_token_id + if getattr(tokenizer, "image_token_id", None) + else tokenizer.convert_tokens_to_ids(self.image_token) + ) + self.video_token_id = ( + tokenizer.video_token_id + if getattr(tokenizer, "video_token_id", None) + else tokenizer.convert_tokens_to_ids(self.video_token) + ) + + @classmethod + def from_pretrained(cls, pretrained_model_name_or_path, **kwargs): + """Build the processor directly from the checkpoint config. + + GLM-5.3-Flash stores nested image/video configs in + ``processor_config.json`` and declares a custom processor class. This + method reads those configs directly and caps only the video token budget. + """ + from transformers import AutoTokenizer + + model_path = pretrained_model_name_or_path + tokenizer = AutoTokenizer.from_pretrained(model_path, **kwargs) + + def _cap_cfg(cfg: dict, *, is_video: bool) -> dict: + # Video keeps a serving token cap (the checkpoint's 240k-token + # budget would starve the KV cache at startup profiling); images + # follow the checkpoint budget verbatim so preprocessing matches + # the HF reference exactly. + if is_video and cfg.get("max_image_tokens") is not None: + cfg["max_image_tokens"] = min( + cfg["max_image_tokens"], _MAX_VIDEO_TOKENS + ) + return cfg + + ip_cfg = _cap_cfg( + dict(get_image_processor_config(model_path, **kwargs)), is_video=False + ) + image_processor = Glm5NextImageProcessor( + **{k: v for k, v in ip_cfg.items() if k != "image_processor_type"} + ) + + processor_config = get_hf_file_to_dict( + "processor_config.json", + model_path, + revision=kwargs.get("revision", "main"), + ) + if processor_config is None: + raise ValueError(f"Missing processor_config.json for {model_path}") + vp_cfg = _cap_cfg(dict(processor_config["video_processor"]), is_video=True) + video_processor = Glm5NextVideoProcessor( + **{k: v for k, v in vp_cfg.items() if k != "video_processor_type"} + ) + + return cls( + image_processor=image_processor, + tokenizer=tokenizer, + video_processor=video_processor, + ) + + def __call__( + self, + images: ImageInput | None = None, + text: TextInput + | PreTokenizedInput + | list[TextInput] + | list[PreTokenizedInput] = None, + videos: VideoInput | None = None, + **kwargs: Unpack[Glm5NextProcessorKwargs], + ) -> BatchFeature: + output_kwargs = self._merge_kwargs( + Glm5NextProcessorKwargs, + tokenizer_init_kwargs=self.tokenizer.init_kwargs, + **kwargs, + ) + if images is not None: + image_inputs = self.image_processor( + images=images, **output_kwargs["images_kwargs"] + ) + else: + image_inputs = {} + + if videos is not None: + videos_inputs = self.video_processor( + videos=videos, **output_kwargs["videos_kwargs"] + ) + if "return_metadata" not in kwargs: + videos_inputs.pop("video_metadata") + else: + videos_inputs = {} + + if not isinstance(text, list): + text = [text] + + # Prompt updates expand the unchanged image/video markers after this call. + return_tensors = output_kwargs["text_kwargs"].pop("return_tensors", None) + return_mm_token_type_ids = output_kwargs["text_kwargs"].pop( + "return_mm_token_type_ids", False + ) + text_inputs = self.tokenizer(text, **output_kwargs["text_kwargs"]) + + if return_mm_token_type_ids: + array_ids = np.array(text_inputs["input_ids"]) + mm_token_type_ids = np.zeros_like(text_inputs["input_ids"]) + mm_token_type_ids[array_ids == self.image_token_id] = 1 + text_inputs["mm_token_type_ids"] = mm_token_type_ids.tolist() + return BatchFeature( + data={**text_inputs, **image_inputs, **videos_inputs}, + tensor_type=return_tensors, + ) + + def _get_num_multimodal_tokens(self, image_sizes=None, video_sizes=None, **kwargs): + vision_data = {} + if image_sizes is not None: + images_kwargs = Glm5NextProcessorKwargs._defaults.get("images_kwargs", {}) + images_kwargs.update(kwargs) + merge_size = ( + images_kwargs.get("merge_size", None) or self.image_processor.merge_size + ) + + num_image_patches = [ + self.image_processor.get_number_of_image_patches( + *image_size, images_kwargs + ) + for image_size in image_sizes + ] + num_image_tokens = [(n // merge_size**2) for n in num_image_patches] + vision_data.update( + { + "num_image_tokens": num_image_tokens, + "num_image_patches": num_image_patches, + } + ) + + if video_sizes is not None: + videos_kwargs = Glm5NextProcessorKwargs._defaults.get("videos_kwargs", {}) + videos_kwargs.update(kwargs) + num_video_patches = [ + self.video_processor.get_number_of_video_patches( + *video_size, videos_kwargs + ) + for video_size in video_sizes + ] + num_video_tokens = [(n // merge_size**2) for n in num_video_patches] + vision_data["num_video_tokens"] = num_video_tokens + + return MultiModalData(**vision_data) + + def post_process_image_text_to_text( + self, + generated_outputs, + skip_special_tokens=True, + clean_up_tokenization_spaces=False, + **kwargs, + ): + return self.tokenizer.batch_decode( + generated_outputs, + skip_special_tokens=skip_special_tokens, + clean_up_tokenization_spaces=clean_up_tokenization_spaces, + **kwargs, + ) + + +__all__ = [ + "Glm5NextImageProcessor", + "Glm5NextVideoProcessor", + "Glm5NextProcessor", + "smart_resize", +] diff --git a/vllm/utils/deep_gemm.py b/vllm/utils/deep_gemm.py index 4b78142c6f66..3fc70c0f5da3 100644 --- a/vllm/utils/deep_gemm.py +++ b/vllm/utils/deep_gemm.py @@ -29,6 +29,11 @@ "qwen3_5_moe_text", } +# KV page sizes (in cache entries) supported by the paged-MQA logits kernels +# (fp8_fp4_paged_mqa_logits / get_paged_mqa_logits_metadata). Larger storage +# blocks must be virtually split into one of these page sizes. +PAGED_MQA_PAGE_SIZES = (32, 64) + def should_auto_disable_deep_gemm(model_type: str | None) -> bool: """Check if DeepGemm should be auto-disabled for this model on Blackwell. @@ -648,6 +653,13 @@ def fp8_fp4_paged_mqa_logits( _lazy_init() if _fp8_fp4_paged_mqa_logits_impl is None: return _missing() + # DeepGEMM asserts block_tables.stride(-1)==1. A trailing size-1 dim + # (e.g. block_table shape [B,1] for short seqs under a large block_size) + # can be a transposed view where .contiguous() is a no-op (torch treats the + # size-1 dim's stride as irrelevant) yet stride(-1)!=1, failing the kernel. + # clone to contiguous format to force stride(-1)==1. + if block_tables.dim() >= 2 and block_tables.stride(-1) != 1: + block_tables = block_tables.clone(memory_format=torch.contiguous_format) kwargs = {} if indices is None else {"indices": indices} return _fp8_fp4_paged_mqa_logits_impl( q, diff --git a/vllm/utils/flashinfer.py b/vllm/utils/flashinfer.py index e20fac6a6ca4..248a9f5fceae 100644 --- a/vllm/utils/flashinfer.py +++ b/vllm/utils/flashinfer.py @@ -292,6 +292,32 @@ def has_flashinfer_moe() -> bool: ) +@functools.cache +def has_flashinfer_sm90_nope_mla() -> bool: + """FlashInfer SM90 NoPE MLA (FP8 KV with in-kernel dequant, kpe=0). + + Feature-detected via the ``ckv_scale_arr`` run() kwarg introduced with + the SM90 NoPE support (FlashInfer >= 0.6.18), so dev builds carry the + gate without a version parse. + """ + if not has_flashinfer(): + return False + try: + import inspect + + from flashinfer.mla import BatchMLAPagedAttentionWrapper + except ImportError: + return False + try: + params = inspect.signature(BatchMLAPagedAttentionWrapper.run).parameters + except (TypeError, ValueError): + return False + return ( + "ckv_scale_arr" in params + and params["ckv_scale_arr"].kind is inspect.Parameter.KEYWORD_ONLY + ) + + @functools.cache def has_flashinfer_sparse_mla_sm120() -> bool: """Return ``True`` if FlashInfer sparse MLA decode support is available.""" diff --git a/vllm/v1/attention/backend.py b/vllm/v1/attention/backend.py index 4fdfafaaea9a..2d545a8d8a7b 100644 --- a/vllm/v1/attention/backend.py +++ b/vllm/v1/attention/backend.py @@ -610,6 +610,10 @@ def __init__( self.layer_names = layer_names self.vllm_config = vllm_config self.device = device + self.kernel_block_size: int | None = None + + def set_kernel_block_size(self, kernel_block_size: int) -> None: + self.kernel_block_size = kernel_block_size @classmethod def get_cudagraph_support( diff --git a/vllm/v1/attention/backends/mla/flashinfer_mla_sparse.py b/vllm/v1/attention/backends/mla/flashinfer_mla_sparse.py index 3798000d7ded..194db8edc3a8 100644 --- a/vllm/v1/attention/backends/mla/flashinfer_mla_sparse.py +++ b/vllm/v1/attention/backends/mla/flashinfer_mla_sparse.py @@ -54,7 +54,10 @@ def get_builder_cls() -> type["FlashInferMLASparseMetadataBuilder"]: @classmethod def get_supported_head_sizes(cls) -> list[int]: - return [576] + # 576 = 512 NoPE + 64 RoPE (with-rope layout); 512 = 512 NoPE only + # (no-rope layout, qk_rope_head_dim == 0). Both share D_V = 512 and are + # served by the TRTLLM-GEN sparse MLA kernel. + return [512, 576] @classmethod def is_mla(cls) -> bool: @@ -117,8 +120,20 @@ def supports_combination( # FlashInfer MLA sparse SM10 kernel requires qk_nope_head_dim in [128, 192]. if vllm_config.model_config is not None: hf_text_config = vllm_config.model_config.hf_text_config - qk_nope_head_dim = getattr(hf_text_config, "qk_nope_head_dim", 1) - if qk_nope_head_dim not in [128, 192]: + qk_nope_head_dim = hf_text_config.qk_nope_head_dim + qk_rope_head_dim = hf_text_config.qk_rope_head_dim + kv_lora_rank = hf_text_config.kv_lora_rank + if qk_rope_head_dim == 0: + # Native no-rope MLA: FlashInfer only ships the one shape + # (nope_mla_dimensions in flashinfer.mla._core). + if qk_nope_head_dim != 256 or kv_lora_rank != 512: + return ( + "FlashInfer native no-rope MLA requires " + "qk_nope_head_dim=256 and kv_lora_rank=512, but got " + f"qk_nope_head_dim={qk_nope_head_dim}, " + f"kv_lora_rank={kv_lora_rank}" + ) + elif qk_nope_head_dim not in [128, 192]: return ( "FlashInfer MLA Sparse kernel requires qk_nope_head_dim " f"in [128, 192], but got {qk_nope_head_dim}" @@ -358,6 +373,10 @@ def __init__( self.bmm1_scale: float | None = None self.bmm2_scale: float | None = None + # Native no-rope MLA additionally requires a per-query-token active + # top-k length tensor. + self.is_nope_mla = self.qk_rope_head_dim == 0 + # fp8 query quantization is required when using fp8 kv_cache, # as the TRTLLM-GEN sparse MLA kernel requires matching dtypes # for query and kv_cache (mixed bf16+fp8 is not supported). @@ -429,6 +448,31 @@ def forward_mqa( block_tables = topk_indices_physical.unsqueeze(1) seq_lens_arg = seq_lens + # page_table width = topk buffer width, which kpool widens past + # index_topk (topk_tokens) and rounds up to a multiple of 128. The + # kernel treats sparse_mla_top_k as the page-table *capacity* and bounds + # the active per-query length by ``seq_lens`` (the compacted valid + # count), so the -1 padding slots past seq_lens are never attended to. + # Use the actual buffer width instead of the fixed topk_tokens, which + # mismatches the page_table when index_kpool > 1. + sparse_topk_capacity = topk_indices_physical.shape[1] + + extra_kwargs: dict[str, torch.Tensor] = {} + empty_rows: torch.Tensor | None = None + if self.is_nope_mla: + # The native no-rope kernel takes the active top-k length per query + # token (``seq_lens`` here is already the compacted per-token valid + # count, int32) and rejects zero-length rows. Point empty rows at a + # single valid dummy slot with length 1 and zero their output after + # the launch. ``triton_convert_req_index_to_global_index`` packs the + # valid indices into a contiguous prefix, which is what the kernel + # requires of the page table. + empty_rows = seq_lens == 0 + topk_indices_physical[:, 0] = topk_indices_physical[:, 0].masked_fill( + empty_rows, 0 + ) + extra_kwargs["sparse_mla_top_k_lens"] = seq_lens.clamp(min=1) + kernel_out = trtllm_batch_decode_with_kv_cache_mla( query=query, kv_cache=kv_c_and_k_pe_cache.unsqueeze(1), @@ -438,11 +482,12 @@ def forward_mqa( qk_rope_head_dim=self.qk_rope_head_dim, block_tables=block_tables, seq_lens=seq_lens_arg, - max_seq_len=attn_metadata.topk_tokens, + max_seq_len=sparse_topk_capacity, bmm1_scale=self.bmm1_scale, bmm2_scale=self.bmm2_scale, - sparse_mla_top_k=attn_metadata.topk_tokens, + sparse_mla_top_k=sparse_topk_capacity, return_lse=self.need_to_return_lse_for_decode, + **extra_kwargs, ) if self.need_to_return_lse_for_decode: assert isinstance(kernel_out, tuple) @@ -455,9 +500,12 @@ def forward_mqa( out = o.view(-1, o.shape[-2], o.shape[-1]) if lse is not None: lse = self._normalize_lse(lse, out.shape[0], out.shape[1]) + if empty_rows is None and lse is not None: empty_rows = (topk_indices_physical == -1).all(dim=-1) + if empty_rows is not None: out.masked_fill_(empty_rows.view(-1, 1, 1), 0.0) - lse.masked_fill_(empty_rows.view(-1, 1), float("-inf")) + if lse is not None: + lse.masked_fill_(empty_rows.view(-1, 1), float("-inf")) return out, lse @staticmethod diff --git a/vllm/v1/attention/backends/mla/flashinfer_mla_sparse_sm90.py b/vllm/v1/attention/backends/mla/flashinfer_mla_sparse_sm90.py new file mode 100644 index 000000000000..90dc4cd94738 --- /dev/null +++ b/vllm/v1/attention/backends/mla/flashinfer_mla_sparse_sm90.py @@ -0,0 +1,478 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +"""FlashInfer sparse MLA backend for SM90 (Hopper) NoPE models. + +Wraps FlashInfer's ``BatchMLAPagedAttentionWrapper`` (FA2/FA3 paths), which +as of FlashInfer 0.6.18 supports ``head_dim_kpe=0`` (GLM-5.3-Flash NoPE MLA) +and FP8 E4M3 KV caches on SM90 with in-kernel dequantization: the FP8 cache +is read directly (half the bf16 HBM traffic) and converted to BF16 in shared +memory, while queries stay BF16 (no query quantization). + +Sparsity rides the same trick the FA-based sparse backend uses: with +``page_size=1`` the per-token top-k slot indices ARE the page table, so each +query token becomes one varlen batch row whose ``kv_indices`` slice is its +top-k row and whose ``kv_len`` is its valid count. Causality is already +encoded by the indexer's selection, so ``causal=False``. + +CUDA-graph handling: ``plan()`` copies its inputs to host unconditionally, +so it must stay outside graph capture. Each metadata builder owns a wrapper, +reserved capture-stable device buffers, and the plan parameters. The wrapper +bakes the per-row ``kv_len`` into its int schedule +at plan() time — ``run()`` never reads the device-side buffer — so the +metadata builder replans every step (outside capture) with exact host-side +lengths derived from the batch's sequence lengths; a full-width schedule +would send the kernel past each row's valid count into the -1 tail of the +converted index buffer (illegal address). Per-step content (top-k slots) +is written into the reserved buffers by kernels inside the captured +forward, and captured runs read the refreshed plan buffers on replay. + +KV cache format: plain contiguous E4M3 ``[num_blocks, block_size, 512]`` +(uint8 storage) with a per-tensor ``k_scale``; BF16 caches also work. The +per-token x 128-channel-group ``ckv_scale_arr`` layout is supported by the +kernel but not wired yet (it needs a group-quantizing cache-write op). +""" + +from dataclasses import dataclass +from typing import Any, ClassVar + +import torch + +from vllm.config import VllmConfig +from vllm.config.cache import CacheDType +from vllm.model_executor.layers.attention.sparse_mla_attention import ( + SparseMLACommonImpl, +) +from vllm.platforms.interface import DeviceCapability +from vllm.utils.flashinfer import has_flashinfer_sm90_nope_mla +from vllm.v1.attention.backend import ( + AttentionBackend, + AttentionLayer, + CommonAttentionMetadata, + MLAAttentionImpl, + MultipleOf, +) +from vllm.v1.attention.backends.mla.flashinfer_mla_sparse import ( + FlashInferMLASparseMetadata, + FlashInferMLASparseMetadataBuilder, +) +from vllm.v1.attention.backends.mla.sparse_utils import ( + triton_convert_req_index_to_global_index, +) +from vllm.v1.kv_cache_interface import AttentionSpec, KVCacheLayout + +_FP8_KV_DTYPES = ("fp8", "fp8_e4m3") +_WORKSPACE_BYTES = 128 * 1024 * 1024 + + +class FlashInferMLASparseSM90Backend(AttentionBackend): + supported_dtypes: ClassVar[list[torch.dtype]] = [torch.bfloat16] + supported_kv_cache_dtypes: ClassVar[list[CacheDType]] = [ + "auto", + "bfloat16", + "fp8", + "fp8_e4m3", + ] + + @staticmethod + def get_supported_kernel_block_sizes() -> list[int | MultipleOf]: + return [MultipleOf(64)] + + @staticmethod + def get_name() -> str: + return "FLASHINFER_MLA_SPARSE_SM90" + + @staticmethod + def get_builder_cls() -> type["FlashInferMLASparseSM90Builder"]: + return FlashInferMLASparseSM90Builder + + @staticmethod + def get_impl_cls() -> type[MLAAttentionImpl]: + return FlashInferMLASparseSM90Impl + + @classmethod + def get_supported_head_sizes(cls) -> list[int]: + # 512 = ckv 512 + kpe 0 (NoPE); 576 = ckv 512 + kpe 64. + return [512, 576] + + @classmethod + def is_mla(cls) -> bool: + return True + + @classmethod + def is_sparse(cls) -> bool: + return True + + @classmethod + def supports_compute_capability(cls, capability: DeviceCapability) -> bool: + return capability.major == 9 + + @classmethod + def supports_combination( + cls, + head_size: int, + dtype: torch.dtype, + kv_cache_dtype: CacheDType | None, + block_size: int | None, + use_mla: bool, + has_sink: bool, + use_sparse: bool, + use_mm_prefix: bool, + device_capability: DeviceCapability, + ) -> str | None: + if not has_flashinfer_sm90_nope_mla(): + return ( + "FLASHINFER_MLA_SPARSE_SM90 requires FlashInfer with SM90 " + "MLA support (ckv_scale_arr in " + "BatchMLAPagedAttentionWrapper.run, FlashInfer >= 0.6.18)" + ) + if not use_sparse: + return "FLASHINFER_MLA_SPARSE_SM90 requires sparse MLA" + from vllm.config import get_current_vllm_config + + vllm_config = get_current_vllm_config() + if vllm_config.model_config is not None: + hf = vllm_config.model_config.hf_text_config + # The SM90 FA2/FA3 kernel covers ckv=512 with kpe in {0, 64} + # (NoPE models and DeepSeek-style rope MLA alike). + if hf.kv_lora_rank != 512: + return "FLASHINFER_MLA_SPARSE_SM90 requires kv_lora_rank=512" + if hf.qk_rope_head_dim not in (0, 64): + return "FLASHINFER_MLA_SPARSE_SM90 requires qk_rope_head_dim in (0, 64)" + return None + + @staticmethod + def get_kv_cache_shape( + num_blocks: int, + block_size: int, + num_kv_heads: int, + head_size: int, + cache_dtype_str: str = "auto", + ) -> tuple[int, ...]: + return (num_blocks, block_size, head_size) + + @classmethod + def supported_kv_cache_layouts(cls) -> tuple[KVCacheLayout, ...]: + return (KVCacheLayout.LBHNC,) + + +class _SM90State: + """Builder-owned wrapper, capture-stable buffers, and plan parameters. + + One instance serves every MLA layer in an attention group because the plan + depends only on the batch shape, not the layer. + """ + + def __init__( + self, + device: torch.device, + num_heads: int, + kv_dtype: torch.dtype, + max_tokens: int, + topk_width: int, + kv_lora_rank: int, + qk_rope_head_dim: int, + sm_scale: float, + ) -> None: + from flashinfer.mla import BatchMLAPagedAttentionWrapper + + self.workspace = torch.empty(_WORKSPACE_BYTES, dtype=torch.uint8, device=device) + self.device = device + self.num_heads = num_heads + self.kv_dtype = kv_dtype + self.max_tokens = max_tokens + self.topk_width = topk_width + self.kv_lora_rank = kv_lora_rank + self.qk_rope_head_dim = qk_rope_head_dim + self.sm_scale = sm_scale + # User-reserved buffers: with use_cuda_graph=True plan() refreshes + # these in place, so run()'s captured kernels always read them. + self.kv_indices = torch.zeros( + max_tokens * topk_width, dtype=torch.int32, device=device + ) + self.kv_len_arr = torch.full( + (max_tokens,), topk_width, dtype=torch.int32, device=device + ) + self.wrapper = BatchMLAPagedAttentionWrapper( + self.workspace, + qo_indptr=torch.zeros(max_tokens + 1, dtype=torch.int32, device=device), + kv_indptr=torch.zeros(max_tokens + 1, dtype=torch.int32, device=device), + kv_indices=self.kv_indices, + kv_len_arr=self.kv_len_arr, + use_cuda_graph=True, + backend="fa3", + ) + self._arange_cpu = torch.arange(self.max_tokens + 1, dtype=torch.int32) + self._qo_cpu = torch.empty(self.max_tokens + 1, dtype=torch.int32) + self._kv_cpu = torch.empty(self.max_tokens + 1, dtype=torch.int32) + self._lens_cpu = torch.full( + (self.max_tokens,), self.topk_width, dtype=torch.int32 + ) + + def plan(self, num_tokens: int, kv_lens: torch.Tensor) -> None: + """Replan with exact per-row KV lengths (CPU int32, ``[num_tokens]``). + + The wrapper bakes kv_len into its int schedule from host values; + run() never reads the device kv_len_arr buffer. Scheduling with the + full buffer width while kv_indices rows carry a -1 tail past each + row's valid count makes the kernel compute ``-1 * ckv_stride_page`` + (illegal address), so the lengths must be exact at plan time and + replanned every step as contexts grow. Must run outside CUDA graph + capture: the in-place refreshed plan_info/indptr buffers are what + captured runs read. + """ + if torch.cuda.is_current_stream_capturing(): + raise RuntimeError( + "FlashInferMLASparseSM90 plan() called inside CUDA graph " + "capture; lengths must be planned host-side before capture." + ) + # CPU staging buffers are filled in place: plan() runs per step + # (once per draft/verify metadata build), so per-call allocations and + # device round trips are on the hot path. Passing CPU tensors lets + # the wrapper's internal .to("cpu") no-op; its reserved-buffer + # copy_ then performs the single H2D transfer per tensor. + # use_cuda_graph=True makes the wrapper copy qo/kv indptr into its + # fixed (max_tokens+1)-sized buffers with exact-size copy_, so the + # indptr must always be full-size. Rows past num_tokens are padded + # empty (qo_indptr flat at num_tokens) — zero-query rows read no q + # and schedule no work. Padded rows keep the full width lens; the + # value is never dereferenced. + torch.clamp(self._arange_cpu, max=num_tokens, out=self._qo_cpu) + torch.mul(self._qo_cpu, self.topk_width, out=self._kv_cpu) + self._lens_cpu.fill_(self.topk_width) + self._lens_cpu[:num_tokens] = kv_lens.to(torch.int32) + self.wrapper.plan( + self._qo_cpu, + self._kv_cpu, + self.kv_indices, + self._lens_cpu, + self.num_heads, + self.kv_lora_rank, # head_dim_ckv + self.qk_rope_head_dim, # 0 (NoPE) or 64 (rope MLA) + 1, # page_size: top-k slots are the page table + False, # causal: encoded by the indexer's selection + self.sm_scale, + q_data_type=torch.bfloat16, + kv_data_type=self.kv_dtype, + ) + + +@dataclass +class FlashInferMLASparseSM90Metadata(FlashInferMLASparseMetadata): + state: _SM90State | None = None + + +class FlashInferMLASparseSM90Builder(FlashInferMLASparseMetadataBuilder): + """Reuse the common sparse metadata (req ids, topk buffer access).""" + + metadata_cls = FlashInferMLASparseSM90Metadata + + def __init__( + self, + kv_cache_spec: "AttentionSpec", + layer_names: list[str], + vllm_config: "VllmConfig", + device: torch.device, + ) -> None: + super().__init__(kv_cache_spec, layer_names, vllm_config, device) + attention_layer = vllm_config.compilation_config.static_forward_context[ + layer_names[0] + ] + impl = attention_layer.impl + if not isinstance(impl, FlashInferMLASparseSM90Impl): + raise TypeError( + "FlashInferMLASparseSM90Builder requires an SM90 FlashInfer " + f"implementation, got {type(impl).__name__}." + ) + topk_indices_buffer = impl.topk_indices_buffer + assert topk_indices_buffer is not None + self.state = _SM90State( + device, + impl.num_heads, + kv_cache_spec.dtype, + vllm_config.scheduler_config.max_num_batched_tokens, + topk_indices_buffer.shape[1], + kv_lora_rank=impl.kv_lora_rank, + qk_rope_head_dim=impl.qk_rope_head_dim, + sm_scale=impl.scale, + ) + # seq_lens_cpu_upper_bound is optimistic on decode rows under async + # spec decode, so the fast sync-free path is only safe without it; + # under async scheduling the exact (device) positions are used at + # the cost of one D2H sync per metadata build. + self._async_scheduling = bool(vllm_config.scheduler_config.async_scheduling) + hf_config = vllm_config.model_config.hf_text_config + assert hf_config.index_topk is not None + self._index_topk = int(hf_config.index_topk) + self._index_kpool = int(kv_cache_spec.tokens_per_state) + + def _kv_lens_host(self, cam: CommonAttentionMetadata) -> tuple[int, torch.Tensor]: + """Exact per-row KV lengths, host-side (the flashinfer wrapper bakes + them into its schedule at plan time; there is no device-side path). + + A row for the j-th query token of request i attends + ``seq_lens[i] - q_len[i] + j + 1`` tokens. The indexer's selection + then bounds the valid count: contexts up to ``index_topk`` select + everything (valid == context); longer contexts keep the top + ``index_topk`` pool-expanded tokens plus the trailing incomplete + pool (valid == ``index_topk + context % index_kpool``). Both match + the count of non -1 entries the convert kernel produces. + """ + num_reqs = cam.num_reqs + qsl = cam.query_start_loc_cpu[: num_reqs + 1] + num_rows = int(qsl[-1]) + if num_rows == 0: + return 0, torch.zeros(0, dtype=torch.int32) + # Row context == position + 1. Without async scheduling the + # maintained host upper bound is exact, giving a sync-free path; + # under async scheduling it is optimistic on decode rows, so fall + # back to the exact device positions (one D2H sync per build). The + # seq_lens derivation equals position + 1 because positions are + # contiguous per request. + sl_host = cam.seq_lens_cpu_upper_bound + positions = cam.positions + if not self._async_scheduling and sl_host is not None: + seq_lens = sl_host[:num_reqs].to(torch.int32) + q_lens = qsl[1:] - qsl[:-1] + first_pos = seq_lens - q_lens + req_of_row = torch.repeat_interleave( + torch.arange(num_reqs, dtype=torch.int64), q_lens.to(torch.int64) + ) + rows = torch.arange(num_rows, dtype=torch.int32) + ctx = ( + first_pos.to(torch.int64)[req_of_row] + + rows.to(torch.int64) + - qsl.to(torch.int64)[req_of_row] + + 1 + ) + elif positions is not None and num_rows <= positions.shape[0]: + ctx = positions[:num_rows].cpu().to(torch.int64) + 1 + else: + seq_lens = cam.seq_lens[:num_reqs].cpu().to(torch.int32) + q_lens = qsl[1:] - qsl[:-1] + first_pos = seq_lens - q_lens + req_of_row = torch.repeat_interleave( + torch.arange(num_reqs, dtype=torch.int64), q_lens.to(torch.int64) + ) + rows = torch.arange(num_rows, dtype=torch.int32) + ctx = ( + first_pos.to(torch.int64)[req_of_row] + + rows.to(torch.int64) + - qsl.to(torch.int64)[req_of_row] + + 1 + ) + topk = self._index_topk + kpool = max(self._index_kpool, 1) + lens = torch.where(ctx <= topk, ctx, topk + ctx % kpool) + return num_rows, lens.to(torch.int32) + + def build( + self, + common_prefix_len: int, + common_attn_metadata: CommonAttentionMetadata, + fast_build: bool = False, + ) -> FlashInferMLASparseSM90Metadata: + metadata = super().build(common_prefix_len, common_attn_metadata, fast_build) + assert isinstance(metadata, FlashInferMLASparseSM90Metadata) + # Replan every step outside any CUDA graph capture with this step's + # exact per-row lengths; captured runs read the refreshed buffers. + num_rows, kv_lens = self._kv_lens_host(common_attn_metadata) + self.state.plan(num_rows, kv_lens) + metadata.state = self.state + return metadata + + +class FlashInferMLASparseSM90Impl(SparseMLACommonImpl[FlashInferMLASparseSM90Metadata]): + def __init__( + self, + num_heads: int, + head_size: int, + scale: float, + num_kv_heads: int, + alibi_slopes: list[float] | None, + sliding_window: int | None, + kv_cache_dtype: str, + logits_soft_cap: float | None, + attn_type: str, + kv_sharing_target_layer_name: str | None, + topk_indices_buffer: torch.Tensor | None = None, + indexer: Any | None = None, + **mla_args: Any, + ) -> None: + if any([alibi_slopes, sliding_window, logits_soft_cap]): + raise NotImplementedError( + "FlashInferMLASparseSM90Impl does not support alibi, sliding " + "window, or logits soft cap." + ) + super().__init__( + num_heads, + head_size, + scale, + num_kv_heads, + alibi_slopes, + sliding_window, + kv_cache_dtype, + logits_soft_cap, + attn_type, + kv_sharing_target_layer_name, + indexer=indexer, + topk_indices_buffer=topk_indices_buffer, + **mla_args, + ) + assert self.topk_indices_buffer is not None + self.supports_quant_query_input = False + self.use_fp8_kv_cache = self.kv_cache_dtype in _FP8_KV_DTYPES + + def forward_mqa( + self, + q: torch.Tensor | tuple[torch.Tensor, torch.Tensor], + kv_c_and_k_pe_cache: torch.Tensor, + attn_metadata: FlashInferMLASparseSM90Metadata, + layer: AttentionLayer, + ) -> tuple[torch.Tensor, torch.Tensor | None]: + if not isinstance(q, tuple): + raise NotImplementedError( + "FlashInferMLASparseSM90Impl expects split (q_nope, q_rope)." + ) + q_nope, q_rope = q + num_tokens = q_rope.shape[0] + # NoPE models hand a zero-width rope tensor through; rope MLA hands + # the real 64-dim part. The kernel takes both as-is. + q_pe = q_rope.reshape(num_tokens, self.num_heads, self.qk_rope_head_dim) + + assert self.topk_indices_buffer is not None + topk_indices = self.topk_indices_buffer[:num_tokens] + # return_valid_counts=True keeps the compacted-prefix layout: valid + # entries at [0, valid_count), -1 past it — exactly the prefix the + # planned per-row lengths address. + topk_slots, _ = triton_convert_req_index_to_global_index( + attn_metadata.req_id_per_token[:num_tokens], + attn_metadata.block_table, + topk_indices, + BLOCK_SIZE=attn_metadata.block_size, + NUM_TOPK_TOKENS=topk_indices.shape[1], + return_valid_counts=True, + ) + state = attn_metadata.state + assert state is not None + # Refresh top-k rows in graph and clamp masked tails to a valid slot; + # per-row lengths are already baked into the host-side plan. + width = topk_slots.shape[1] + state.kv_indices[: num_tokens * width].copy_( + topk_slots.reshape(-1).clamp_(min=0).to(torch.int32) + ) + + flat = ( + kv_c_and_k_pe_cache.view(torch.float8_e4m3fn) + if self.use_fp8_kv_cache + else kv_c_and_k_pe_cache + ).reshape(-1, 1, self.head_size) + ckv = flat[..., : self.kv_lora_rank] + kpe = flat[..., self.kv_lora_rank :] + + scale_kwargs = ( + {"ckv_scale": float(layer._k_scale_float or 1.0), "kpe_scale": 1.0} + if self.use_fp8_kv_cache + else {} + ) + out = state.wrapper.run(q_nope, q_pe, ckv, kpe, **scale_kwargs) + return out, None diff --git a/vllm/v1/attention/backends/mla/indexer.py b/vllm/v1/attention/backends/mla/indexer.py index 206023f5da01..6198c14eb738 100644 --- a/vllm/v1/attention/backends/mla/indexer.py +++ b/vllm/v1/attention/backends/mla/indexer.py @@ -37,7 +37,12 @@ get_dcp_local_seq_lens, split_decodes_and_prefills, ) -from vllm.v1.kv_cache_interface import KVCacheLayout, KVCacheSpec, MLAAttentionSpec +from vllm.v1.kv_cache_interface import ( + AttentionSpec, + KVCacheLayout, + KVCacheSpec, + MLAAttentionSpec, +) logger = init_logger(__name__) @@ -252,6 +257,30 @@ def get_builder_cls() -> type["DeepseekV32IndexerMetadataBuilder"]: return DeepseekV32IndexerMetadataBuilder +class KpoolTailBackend(DeepseekV32IndexerBackend): + """Storage-only backend for the GLM-5.3-Flash kpool tail cache.""" + + @classmethod + def supported_kv_cache_layouts(cls) -> tuple[KVCacheLayout, ...]: + return (KVCacheLayout.LBHNC,) + + @staticmethod + def get_name() -> str: + return "KPOOL_TAIL" + + @classmethod + def get_supported_head_sizes(cls) -> list[int]: + return [] + + @staticmethod + def get_supported_kernel_block_sizes() -> list[int | MultipleOf]: + return [MultipleOf(1)] + + @staticmethod + def get_builder_cls() -> type["KpoolTailMetadataBuilder"]: # type: ignore[override] + return KpoolTailMetadataBuilder + + class DeepseekV4IndexerBackend(DeepseekV32IndexerBackend): @staticmethod def get_name() -> str: @@ -419,6 +448,9 @@ def get_warmup_keys(self, vllm_config: VllmConfig) -> list[CompileKey]: ) ) ) + index_kpool = getattr(hf_config, "index_kpool", None) + if index_kpool and index_kpool > 1 and index_kpool not in compress_ratios: + compress_ratios = compress_ratios + (index_kpool,) return self._trace_dispatch(self.dispatch)( # Cover Triton's divisible, exact-one, and generic i32 classes. query_slice_start=(0, 1, 2), @@ -483,6 +515,7 @@ def __call__( @dataclass class DeepseekV32IndexerPrefillMetadata: chunks: list[DeepseekV32IndexerPrefillChunkMetadata] + max_prefill_seq_len: int = -1 @dataclass @@ -497,6 +530,9 @@ class DeepSeekV32IndexerDecodeMetadata: requires_padding: bool schedule_metadata: torch.Tensor global_seq_lens: torch.Tensor | None = None + per_req_decode_lens: torch.Tensor | None = None + decode_is_uniform: bool = True + write_max_decode_len: int = 0 indices: torch.Tensor | None = None @@ -519,6 +555,90 @@ class DeepseekV32IndexerMetadata: prefill: DeepseekV32IndexerPrefillMetadata | None = None +def compute_kpool_tail_slot_mapping( + slot_mapping: torch.Tensor, + block_table: torch.Tensor, + query_start_loc: torch.Tensor, + positions: torch.Tensor, + num_actual_tokens: int, + num_reqs: int, + kpool: int, + out: torch.Tensor | None = None, +) -> torch.Tensor: + """Map every token to its request's one circular tail block.""" + if out is None: + out = slot_mapping.clone() + else: + assert out.shape == slot_mapping.shape + out.copy_(slot_mapping) + if num_actual_tokens == 0: + return out + tokens = torch.arange(num_actual_tokens, device=slot_mapping.device) + req = torch.searchsorted(query_start_loc, tokens, right=True) - 1 + req = req.clamp_(min=0, max=num_reqs - 1) + own_block = block_table[:num_reqs, 0].index_select(0, req).to(torch.int64) + pos = positions[:num_actual_tokens].to(torch.int64) + out[:num_actual_tokens] = own_block * kpool + torch.remainder(pos, kpool) + return out + + +class KpoolTailMetadataBuilder(AttentionMetadataBuilder): + """Build only the circular slot mapping needed by the storage-only tail.""" + + _cudagraph_support = AttentionCGSupport.ALWAYS + supports_update_block_table = False + reorder_batch_threshold = None + + def __init__( + self, + kv_cache_spec: AttentionSpec, + layer_names: list[str], + vllm_config: VllmConfig, + device: torch.device, + ): + super().__init__(kv_cache_spec, layer_names, vllm_config, device) + self.slot_mapping_buffer = torch.empty( + vllm_config.scheduler_config.max_num_batched_tokens, + dtype=torch.int64, + device=device, + ) + + def build( + self, + common_prefix_len: int, + common_attn_metadata: CommonAttentionMetadata, + fast_build: bool = False, + ) -> DeepseekV32IndexerMetadata: + num_decodes, num_prefills, num_decode_tokens, num_prefill_tokens = ( + split_decodes_and_prefills(common_attn_metadata) + ) + slot_mapping = common_attn_metadata.slot_mapping + positions = common_attn_metadata.positions + if positions is not None: + slot_mapping_buffer = self.slot_mapping_buffer[ + : slot_mapping.numel() + ].view_as(slot_mapping) + slot_mapping = compute_kpool_tail_slot_mapping( + slot_mapping, + common_attn_metadata.block_table_tensor, + common_attn_metadata.query_start_loc, + positions, + common_attn_metadata.num_actual_tokens, + common_attn_metadata.num_reqs, + self.kv_cache_spec.block_size, + out=slot_mapping_buffer, + ) + return DeepseekV32IndexerMetadata( + seq_lens=common_attn_metadata.seq_lens, + max_seq_len=common_attn_metadata.max_seq_len, + slot_mapping=slot_mapping, + num_decodes=num_decodes, + num_decode_tokens=num_decode_tokens, + num_prefills=num_prefills, + num_prefill_tokens=num_prefill_tokens, + ) + + def get_max_prefill_buffer_size(vllm_config: VllmConfig): max_model_len = vllm_config.model_config.max_model_len # NOTE(Chen): 40 is a magic number for controlling the prefill buffer size. @@ -640,6 +760,11 @@ def __init__(self, *args, block_table_width: int, **kwargs) -> None: dtype=torch.int32, device=self.device, ) + self.per_req_decode_lens_buffer = torch.zeros( + (scheduler_config.max_num_batched_tokens,), + dtype=torch.int32, + device=self.device, + ) # Shared workspace for decode seq_lens. Native MTP views this as # (B, max_decode_len) at runtime, keeping context_lens contiguous even # when max_decode_len is smaller than next_n. @@ -707,6 +832,8 @@ def __init__(self, *args, block_table_width: int, **kwargs) -> None: dtype=torch.int32, device=self.device, ) + self.indexer_decode_block_table_buffer: torch.Tensor | None = None + self._max_num_batched_tokens = scheduler_config.max_num_batched_tokens def _dcp_localize_decode_seq_lens( self, @@ -921,7 +1048,16 @@ def build( compressed_slot_mapping = slot_mapping compressed_seq_lens = seq_lens + indexer_block_table = block_table if self.compress_ratio > 1: + kernel_block_size = self.kernel_block_size + if ( + kernel_block_size is not None + and self.kv_cache_spec.block_size != kernel_block_size + and self.kv_cache_spec.block_size % kernel_block_size == 0 + ): + factor = self.kv_cache_spec.block_size // kernel_block_size + indexer_block_table = (block_table[:, ::factor] // factor).contiguous() padded_num_tokens = num_tokens if self.pcp_world_size > 1: padded_num_tokens = slot_mapping.shape[0] // self.pcp_world_size @@ -929,7 +1065,7 @@ def build( num_tokens, query_start_loc, seq_lens, - block_table, + indexer_block_table, self.kv_cache_spec.num_states, self.compress_ratio, out=self.compressed_slot_mapping_buffer, @@ -979,7 +1115,7 @@ def build( seq_lens, compressed_seq_lens, compressed_seq_lens_cpu, - common_attn_metadata.block_table_tensor, + indexer_block_table, self.compress_ratio, query_slice=query_slice, skip_kv_gather=query_slice.start > 0, @@ -990,7 +1126,14 @@ def build( # Skip when total_seq_lens is 0 (i.e., no compressed token). if metadata is not None: chunks.append(metadata) - prefill_metadata = DeepseekV32IndexerPrefillMetadata(chunks) + prefill_metadata = DeepseekV32IndexerPrefillMetadata( + chunks, + max_prefill_seq_len=( + int(seq_lens_cpu[num_decodes:].max().item()) + if num_prefills > 0 + else 0 + ), + ) decode_metadata = None if num_decodes > 0: @@ -999,6 +1142,7 @@ def build( out=self.decode_lens_buffer[:num_decodes], ) decode_lens = self.decode_lens_buffer[:num_decodes] + self.per_req_decode_lens_buffer[:num_decodes].copy_(decode_lens) decode_lens_cpu = torch.diff( common_attn_metadata.query_start_loc_cpu[: num_decodes + 1] ) @@ -1018,6 +1162,8 @@ def build( block_table = common_attn_metadata.block_table_tensor[:num_decodes, ...] max_decode_len = int(decode_lens_cpu.max().item()) + min_decode_len = int(decode_lens_cpu.min().item()) + write_is_uniform = min_decode_len == max_decode_len next_n = 1 + self.num_speculative_tokens # The kernel sees max_decode_len Q rows, not the configured next_n, # so legality is per-step: on SM90 a uniformly 3-deep batch has no @@ -1064,7 +1210,30 @@ def build( ) ) - seq_lens_is_buffer_view = not use_native or next_n > 1 + if self.compress_ratio > 1: + kernel_block_size = self.kernel_block_size + if ( + kernel_block_size is not None + and self.kv_cache_spec.block_size != kernel_block_size + and self.kv_cache_spec.block_size % kernel_block_size == 0 + ): + factor = self.kv_cache_spec.block_size // kernel_block_size + compressed = block_table[:, ::factor] // factor + rows, cols = compressed.shape + if self.indexer_decode_block_table_buffer is None: + self.indexer_decode_block_table_buffer = torch.zeros( + (self._max_num_batched_tokens, cols), + dtype=torch.int32, + device=self.device, + ) + self.indexer_decode_block_table_buffer[:rows, :cols].copy_( + compressed + ) + block_table = self.indexer_decode_block_table_buffer[:rows, :cols] + + seq_lens_is_buffer_view = (use_native and next_n > 1) or ( + not use_native and max_decode_len > 1 + ) # DCP: localize the now-expanded per-token global bounds to this # rank's owned KV. Done here (after expansion) so each token's global @@ -1113,6 +1282,9 @@ def build( schedule_metadata=schedule_metadata, indices=decode_indices, global_seq_lens=global_seq_lens_for_decode, + per_req_decode_lens=self.per_req_decode_lens_buffer[:num_decodes], + decode_is_uniform=write_is_uniform, + write_max_decode_len=max_decode_len, ) attn_metadata = DeepseekV32IndexerMetadata( diff --git a/vllm/v1/attention/backends/mla/rocm_aiter_mla.py b/vllm/v1/attention/backends/mla/rocm_aiter_mla.py index ddf8bbd54e7c..598527e31916 100644 --- a/vllm/v1/attention/backends/mla/rocm_aiter_mla.py +++ b/vllm/v1/attention/backends/mla/rocm_aiter_mla.py @@ -881,7 +881,11 @@ def _fill_dcp_verify_page_table( rows. A DCP rank holds ``1/dcp_world_size`` of the sequence, so the request's block table is always wider than the pages a row can reach. """ - pages_per_block = self.kernel_block_size // self._segmented_page_size + kernel_block_size = self.kernel_block_size + page_size = self._segmented_page_size + assert kernel_block_size is not None + assert page_size is not None + pages_per_block = kernel_block_size // page_size max_local_pages = row_block_table.shape[1] max_local_blocks = cdiv(max_local_pages, pages_per_block) assert max_local_blocks <= block_table.shape[1], ( diff --git a/vllm/v1/attention/backends/mla/rocm_aiter_mla_sparse.py b/vllm/v1/attention/backends/mla/rocm_aiter_mla_sparse.py index 54f716ea50c1..ef707d6ee1ae 100644 --- a/vllm/v1/attention/backends/mla/rocm_aiter_mla_sparse.py +++ b/vllm/v1/attention/backends/mla/rocm_aiter_mla_sparse.py @@ -32,6 +32,9 @@ AiterMLAHelper, ) from vllm.v1.attention.backends.utils import split_decodes_and_prefills +from vllm.v1.attention.ops.rocm_aiter_mla_sparse import ( + rocm_sparse_attn_prefill, +) from vllm.v1.kv_cache_interface import AttentionSpec from vllm.v1.worker.workspace import current_workspace_manager @@ -40,6 +43,51 @@ logger = init_logger(__name__) +def _use_rocm_sparse_triton( + *, + kv_cache_dtype: str, + head_size: int, + kv_lora_rank: int, + num_prefills: int, + num_decodes: int, + num_decode_tokens: int, + max_query_len: int, +) -> bool: + """Select the rope-free BF16 path not supported by AITER sparse MLA.""" + plain_decode = num_decode_tokens == num_decodes + return ( + not kv_cache_dtype.startswith("fp8") + and head_size == kv_lora_rank + and plain_decode + and (num_prefills > 0 or (num_decodes > 0 and max_query_len == 1)) + ) + + +def fit_kpool_indices_to_aiter( + token_indices: torch.Tensor, topk_tokens: int +) -> torch.Tensor: + """Keep the live kpool tail while fitting AITER's fixed top-k width.""" + if token_indices.shape[1] < topk_tokens: + raise ValueError("token_indices width must be at least topk_tokens") + if token_indices.shape[1] == topk_tokens: + return token_indices + + history = token_indices[:, :topk_tokens] + tail = token_indices[:, topk_tokens:] + valid_history = (history >= 0).sum(dim=1) + valid_tail = (tail >= 0).sum(dim=1) + keep_history = torch.minimum(valid_history, topk_tokens - valid_tail) + + columns = torch.arange(topk_tokens, device=token_indices.device).unsqueeze(0) + tail_offsets = columns - keep_history.unsqueeze(1) + tail_values = torch.gather( + tail, 1, tail_offsets.clamp(min=0, max=tail.shape[1] - 1) + ) + output = torch.where(columns < keep_history.unsqueeze(1), history, tail_values) + valid_output = columns < (keep_history + valid_tail).unsqueeze(1) + return torch.where(valid_output, output, -1) + + @triton.jit def _convert_req_index_to_global_index_kernel( req_id_ptr, # int32 [num_tokens] @@ -359,6 +407,7 @@ def __init__( self.kv_cache_spec = kv_cache_spec self.model_config = vllm_config.model_config self.model_dtype = vllm_config.model_config.dtype + self.kv_cache_dtype = vllm_config.cache_config.cache_dtype parallel_config = vllm_config.parallel_config self.device = device max_num_batched_tokens = vllm_config.scheduler_config.max_num_batched_tokens @@ -550,53 +599,74 @@ def build( # context lengths (both clamped to topk_tokens, past which per-token KV # length saturates) and num_heads; fingerprint those CPU-side and skip # the launch when nothing changed. - num_reqs = common_attn_metadata.num_reqs - clamped_seq_lens = np.minimum( - common_attn_metadata.seq_lens_cpu[:num_reqs].numpy(), - self.topk_tokens, - ) - clamped_context_lens = np.minimum( - common_attn_metadata.seq_lens_cpu[:num_reqs].numpy() - seg_lengths, - self.topk_tokens, - ) - metadata_key = ( - num_tokens, - int(common_attn_metadata.max_query_len), - self._num_attention_heads, - clamped_seq_lens.tobytes(), - clamped_context_lens.tobytes(), - seg_lengths.tobytes(), + head_size = self.mla_dims.kv_lora_rank + self.mla_dims.qk_rope_head_dim + use_triton_sparse = _use_rocm_sparse_triton( + kv_cache_dtype=self.kv_cache_dtype, + head_size=head_size, + kv_lora_rank=self.mla_dims.kv_lora_rank, + num_prefills=num_prefills, + num_decodes=num_decodes, + num_decode_tokens=num_decode_tokens, + max_query_len=common_attn_metadata.max_query_len, ) - if metadata_key != self._prev_metadata_key: - from aiter import get_mla_metadata_v1 - - max_split_per_batch = self._sparse_decode_max_split( - int(common_attn_metadata.max_seq_len) + work_meta_data = None + work_indptr = None + work_info_set = None + reduce_indptr = None + reduce_final_map = None + reduce_partial_map = None + if not use_triton_sparse: + num_reqs = common_attn_metadata.num_reqs + clamped_seq_lens = np.minimum( + common_attn_metadata.seq_lens_cpu[:num_reqs].numpy(), + self.topk_tokens, + ) + clamped_context_lens = np.minimum( + common_attn_metadata.seq_lens_cpu[:num_reqs].numpy() - seg_lengths, + self.topk_tokens, ) - get_mla_metadata_v1( - qo_indptr, - paged_kv_indptr, - paged_kv_last_page_len, + metadata_key = ( + num_tokens, + int(common_attn_metadata.max_query_len), self._num_attention_heads, - 1, - True, - self._mla_work_meta_data, - self._mla_work_info_set, - self._mla_work_indptr, - self._mla_reduce_indptr, - self._mla_reduce_final_map, - self._mla_reduce_partial_map, - page_size=1, - kv_granularity=16, - max_seqlen_qo=1, - uni_seqlen_qo=1, - fast_mode=True, - max_split_per_batch=max_split_per_batch, + clamped_seq_lens.tobytes(), + clamped_context_lens.tobytes(), + seg_lengths.tobytes(), ) - # The persistent metadata buffers are read by graph replay. Order - # the async metadata write before the graph-captured decode kernel. - torch.cuda.current_stream(self.device).synchronize() - self._prev_metadata_key = metadata_key + if metadata_key != self._prev_metadata_key: + from aiter import get_mla_metadata_v1 + + max_split_per_batch = self._sparse_decode_max_split( + int(common_attn_metadata.max_seq_len) + ) + get_mla_metadata_v1( + qo_indptr, + paged_kv_indptr, + paged_kv_last_page_len, + self._num_attention_heads, + 1, + True, + self._mla_work_meta_data, + self._mla_work_info_set, + self._mla_work_indptr, + self._mla_reduce_indptr, + self._mla_reduce_final_map, + self._mla_reduce_partial_map, + page_size=1, + kv_granularity=16, + max_seqlen_qo=1, + uni_seqlen_qo=1, + fast_mode=True, + max_split_per_batch=max_split_per_batch, + ) + torch.cuda.current_stream(self.device).synchronize() + self._prev_metadata_key = metadata_key + work_meta_data = self._mla_work_meta_data + work_indptr = self._mla_work_indptr + work_info_set = self._mla_work_info_set + reduce_indptr = self._mla_reduce_indptr + reduce_final_map = self._mla_reduce_final_map + reduce_partial_map = self._mla_reduce_partial_map metadata = ROCMAiterMLASparseMetadata( num_reqs=common_attn_metadata.num_reqs, @@ -617,12 +687,12 @@ def build( paged_kv_last_page_len=paged_kv_last_page_len, paged_kv_indices=paged_kv_indices, paged_kv_indptr=paged_kv_indptr, - work_meta_data=self._mla_work_meta_data, - work_indptr=self._mla_work_indptr, - work_info_set=self._mla_work_info_set, - reduce_indptr=self._mla_reduce_indptr, - reduce_final_map=self._mla_reduce_final_map, - reduce_partial_map=self._mla_reduce_partial_map, + work_meta_data=work_meta_data, + work_indptr=work_indptr, + work_info_set=work_info_set, + reduce_indptr=reduce_indptr, + reduce_final_map=reduce_final_map, + reduce_partial_map=reduce_partial_map, ) return metadata @@ -715,6 +785,31 @@ def _forward_mla( device=q.device, ) + if _use_rocm_sparse_triton( + kv_cache_dtype=self.kv_cache_dtype, + head_size=q.shape[-1], + kv_lora_rank=self.kv_lora_rank, + num_prefills=attn_metadata.num_prefills, + num_decodes=attn_metadata.num_decodes, + num_decode_tokens=attn_metadata.num_decode_tokens, + max_query_len=attn_metadata.max_query_len, + ): + rocm_sparse_attn_prefill( + q=q, + kv=kv_c_and_k_pe_cache.view(-1, 1, q.shape[-1]), + indices=None, + topk_length=None, + scale=self.scale, + head_dim=q.shape[-1], + nope_head_dim=self.kv_lora_rank, + rope_head_dim=q.shape[-1] - self.kv_lora_rank, + attn_sink=None, + output=output, + ragged_indices=attn_metadata.paged_kv_indices, + ragged_indptr=attn_metadata.paged_kv_indptr, + ) + return AiterMLAHelper.get_mla_unpadded_o(self.num_heads, output) + # Build kwargs and forward the persistent MLA metadata when it has # been computed. The aiter mla_decode_fwd switches to its # work-stealing persistent kernel path when work_meta_data is given. @@ -766,13 +861,18 @@ def forward_mqa( ) else: q = self.q_concat_buffer[: ql_nope.shape[0]] - ops.concat_mla_q(ql_nope, q_pe, q) + if q_pe.shape[-1] == 0: + q.copy_(ql_nope) + else: + ops.concat_mla_q(ql_nope, q_pe, q) num_actual_toks = attn_metadata.num_actual_tokens # Get topk indices assert self.topk_indices_buffer is not None - topk_indices = self.topk_indices_buffer[:num_actual_toks] + topk_indices = fit_kpool_indices_to_aiter( + self.topk_indices_buffer[:num_actual_toks], attn_metadata.topk_tokens + ) triton_convert_req_index_to_global_index( attn_metadata.req_id_per_token, @@ -795,5 +895,4 @@ def forward_mqa( attn_out = self._forward_mla( layer, mla_padded_q, kv_c_and_k_pe_cache, attn_metadata ) - return attn_out, None diff --git a/vllm/v1/attention/backends/mla/sparse_utils.py b/vllm/v1/attention/backends/mla/sparse_utils.py index 168a8c05c15c..83e0e780e4c0 100644 --- a/vllm/v1/attention/backends/mla/sparse_utils.py +++ b/vllm/v1/attention/backends/mla/sparse_utils.py @@ -251,7 +251,7 @@ def get_warmup_keys( dict( HAS_PREFILL_WORKSPACE=False, COUNT_VALID=True, - COMPACT_TO_FRONT=False, + COMPACT_TO_FRONT=True, DCP_SIZE=1, DCP_RANK=0, DCP_INTERLEAVE=1, @@ -259,7 +259,7 @@ def get_warmup_keys( dict( HAS_PREFILL_WORKSPACE=True, COUNT_VALID=True, - COMPACT_TO_FRONT=False, + COMPACT_TO_FRONT=True, DCP_SIZE=1, DCP_RANK=0, DCP_INTERLEAVE=1, @@ -458,7 +458,14 @@ def triton_convert_req_index_to_global_index( req_id_c = req_id.contiguous() block_table_c = block_table.contiguous() token_indices_c = token_indices.contiguous() - out = torch.empty_like(token_indices_c) + # When return_valid_counts, the kernel scatters valid entries to a + # contiguous prefix [0, valid_count) and leaves the tail unwritten, so + # pre-fill -1 there. flash_mla_sparse_fwd then bounds attention to + # [:topk_length] == exactly the valid set (no dropped tokens). + if return_valid_counts: + out = torch.full_like(token_indices_c, -1) + else: + out = torch.empty_like(token_indices_c) valid_counts: torch.Tensor | None = None if return_valid_counts: @@ -491,7 +498,7 @@ def triton_convert_req_index_to_global_index( NUM_TOPK_TOKENS=NUM_TOPK_TOKENS, HAS_PREFILL_WORKSPACE=HAS_PREFILL_WORKSPACE, COUNT_VALID=return_valid_counts, - COMPACT_TO_FRONT=False, + COMPACT_TO_FRONT=return_valid_counts, # DCP disabled (no-op de-interleave) DCP_SIZE=1, DCP_RANK=0, diff --git a/vllm/v1/attention/backends/registry.py b/vllm/v1/attention/backends/registry.py index fa8a445353ad..0c9b39c9886f 100644 --- a/vllm/v1/attention/backends/registry.py +++ b/vllm/v1/attention/backends/registry.py @@ -77,6 +77,10 @@ class AttentionBackendEnum(Enum, metaclass=_AttentionBackendEnumMeta): "vllm.v1.attention.backends.mla.flashinfer_mla_sparse." "FlashInferMLASparseSM120Backend" ) + FLASHINFER_MLA_SPARSE_SM90 = ( + "vllm.v1.attention.backends.mla.flashinfer_mla_sparse_sm90." + "FlashInferMLASparseSM90Backend" + ) TRITON_MLA = "vllm.v1.attention.backends.mla.triton_mla.TritonMLABackend" CUTLASS_MLA = "vllm.v1.attention.backends.mla.cutlass_mla.CutlassMLABackend" FLASHMLA = "vllm.v1.attention.backends.mla.flashmla.FlashMLABackend" diff --git a/vllm/v1/attention/ops/rocm_aiter_mla_sparse.py b/vllm/v1/attention/ops/rocm_aiter_mla_sparse.py index 7fef84fdc643..62f91f260e80 100644 --- a/vllm/v1/attention/ops/rocm_aiter_mla_sparse.py +++ b/vllm/v1/attention/ops/rocm_aiter_mla_sparse.py @@ -909,6 +909,7 @@ def rocm_aiter_sparse_attn_indexer( has_decode = layer_attn_metadata.num_decodes > 0 has_prefill = layer_attn_metadata.num_prefills > 0 num_decode_tokens = layer_attn_metadata.num_decode_tokens + topk_indices_buffer[: hidden_states.shape[0]] = -1 # during speculative decoding, k may be padded to the CUDA graph batch # size while slot_mapping only covers actual tokens. @@ -1251,15 +1252,31 @@ def rocm_inv_rope_einsum( _DSV4_SPARSE_ROPE_DIM = 64 -def _validate_dsv4_sparse_dims( +def _validate_sparse_dims( head_dim: int, nope_head_dim: int, rope_head_dim: int, op_name: str, ) -> None: + assert head_dim > 0, f"{op_name} expected a positive head_dim, got {head_dim}" + assert nope_head_dim > 0, ( + f"{op_name} expected a positive NoPE dimension, got {nope_head_dim}" + ) + assert rope_head_dim >= 0, ( + f"{op_name} expected a non-negative RoPE dimension, got {rope_head_dim}" + ) assert head_dim == nope_head_dim + rope_head_dim, ( f"{op_name} expected head_dim={nope_head_dim + rope_head_dim}, got {head_dim}" ) + + +def _validate_dsv4_sparse_dims( + head_dim: int, + nope_head_dim: int, + rope_head_dim: int, + op_name: str, +) -> None: + _validate_sparse_dims(head_dim, nope_head_dim, rope_head_dim, op_name) assert ( nope_head_dim == _DSV4_SPARSE_NOPE_DIM and rope_head_dim == _DSV4_SPARSE_ROPE_DIM @@ -1351,6 +1368,12 @@ def _as_int32_contiguous_1d(x: torch.Tensor) -> torch.Tensor: return x.to(torch.int32).contiguous() +@triton.jit +def _sparse_kv_row_offset(slot, stride): + # A global token slot fits in int32, but its byte/element offset may not. + return slot.to(tl.int64) * stride + + @triton.jit def _sparse_attn_prefill_ragged_kernel( q_ptr, @@ -1414,7 +1437,7 @@ def _sparse_attn_prefill_ragged_kernel( kv = tl.load( kv_ptr - + safe_slot[:, None] * kv_stride_n + + _sparse_kv_row_offset(safe_slot[:, None], kv_stride_n) + dim_offsets[None, :] * kv_stride_d, mask=valid[:, None] & dim_mask[None, :], other=0.0, @@ -2573,7 +2596,7 @@ def _rocm_sparse_attn_prefill_ragged_triton( assert indptr.numel() == num_queries + 1, ( f"expected indptr shape [{num_queries + 1}], got {indptr.shape}" ) - _validate_dsv4_sparse_dims( + _validate_sparse_dims( head_dim, nope_head_dim, rope_head_dim, @@ -3125,7 +3148,7 @@ def _rocm_sparse_attn_decode_triton( def rocm_sparse_attn_prefill( q: torch.Tensor, kv: torch.Tensor, - indices: torch.Tensor, + indices: torch.Tensor | None, topk_length: torch.Tensor | None, scale: float, head_dim: int, @@ -3139,7 +3162,7 @@ def rocm_sparse_attn_prefill( assert kv.ndim == 3 and kv.shape[1] == 1, ( f"ROCm Triton sparse prefill expects kv=[skv,1,d], got {kv.shape}" ) - _validate_dsv4_sparse_dims( + _validate_sparse_dims( head_dim, nope_head_dim, rope_head_dim, @@ -3157,6 +3180,7 @@ def rocm_sparse_attn_prefill( rope_head_dim=rope_head_dim, ) else: + assert indices is not None indices_2d = indices.reshape(indices.shape[0], -1) output_chunk = _rocm_sparse_attn_prefill_triton( q=q, diff --git a/vllm/v1/core/kv_cache_coordinator.py b/vllm/v1/core/kv_cache_coordinator.py index b3a9b4c03210..34711d13cc81 100644 --- a/vllm/v1/core/kv_cache_coordinator.py +++ b/vllm/v1/core/kv_cache_coordinator.py @@ -599,6 +599,9 @@ def __init__( # can be a multiple of hash_block_size. self.hash_block_size = hash_block_size self.dcp_world_size = dcp_world_size + # Only groups that participate in prefix caching must satisfy the + # divisibility constraint; groups that opt out (e.g. GLM-5.3-Flash kpool + # tail, block_size=kpool) are scratch buffers and excluded. group_block_sizes = [ manager.block_size for manager, group in zip( @@ -679,6 +682,11 @@ def verify_and_split_kv_cache_groups(self) -> None: """ self.attention_groups: list[SpecGroup] = [] for i, g in enumerate(self.kv_cache_config.kv_cache_groups): + # Skip groups that opt out of prefix caching (e.g. GLM-5.3-Flash + # kpool tail): their blocks are per-request scratch, never + # shareable, so they must not participate in hit lookup (their + # manager-level hooks already no-op). Their slot in the per-group + # hit tuple stays empty. if not g.kv_cache_spec.prefix_cacheable: continue manager_cls = self.single_type_managers[i].__class__ diff --git a/vllm/v1/core/kv_cache_utils.py b/vllm/v1/core/kv_cache_utils.py index 1bbe002c97ef..2fce0c0f4e2a 100644 --- a/vllm/v1/core/kv_cache_utils.py +++ b/vllm/v1/core/kv_cache_utils.py @@ -10,7 +10,7 @@ from collections.abc import Callable, Iterable, Iterator, Sequence from dataclasses import dataclass, replace from functools import partial -from typing import Any, NamedTuple, NewType, TypeAlias, overload +from typing import Any, NamedTuple, NewType, TypeAlias, cast, overload from vllm import envs from vllm.config import VllmConfig @@ -24,6 +24,7 @@ ChunkedLocalAttentionSpec, FullAttentionSpec, HiddenStateCacheSpec, + KpoolTailSpec, KVCacheConfig, KVCacheGroupSpec, KVCacheLayout, @@ -1132,6 +1133,205 @@ def _get_kv_cache_groups_uniform_type( return [KVCacheGroupSpec(list(spec.kv_cache_specs.keys()), spec)] +def _pp_balanced_mamba_group_count( + vllm_config: VllmConfig, + mamba_layer_names: list[str], + mla_layer_names: list[str], +) -> int | None: + """Return a Mamba group count whose PP projections fit the MLA slots.""" + num_groups = cdiv(len(mamba_layer_names), len(mla_layer_names)) + pp_size = vllm_config.parallel_config.pipeline_parallel_size + if pp_size == 1: + return num_groups + + from vllm.distributed.utils import get_pp_indices + from vllm.model_executor.models.utils import extract_layer_index + + total_layers = vllm_config.model_config.get_total_num_hidden_layers() + mamba_indices = [extract_layer_index(name) for name in mamba_layer_names] + mla_indices = [extract_layer_index(name) for name in mla_layer_names] + for rank in range(pp_size): + start, end = get_pp_indices(total_layers, rank, pp_size) + num_mamba = sum(start <= index < end for index in mamba_indices) + num_mla = sum(start <= index < end for index in mla_indices) + if not num_mamba: + continue + if not num_mla: + return None + num_groups = max(num_groups, cdiv(num_mamba, num_mla)) + return num_groups + + +def _get_kv_cache_groups_glm5_next( + vllm_config: VllmConfig, + kv_cache_spec: dict[str, KVCacheSpec], +) -> list[KVCacheGroupSpec] | None: + """Build GLM-5.3-Flash groups with Mamba/MLA and tail/indexer aliasing.""" + mamba_specs = { + name: spec + for name, spec in kv_cache_spec.items() + if isinstance(spec, MambaSpec) + } + tail_specs = { + name: spec + for name, spec in kv_cache_spec.items() + if isinstance(spec, KpoolTailSpec) + } + attn_specs = { + name: spec + for name, spec in kv_cache_spec.items() + if not isinstance(spec, (MambaSpec, KpoolTailSpec)) + } + if not mamba_specs or not all( + type(spec) is MLAAttentionSpec for spec in attn_specs.values() + ): + return None + + mla_specs = cast(dict[str, MLAAttentionSpec], attn_specs) + idx_pages = { + spec.page_size_bytes for spec in mla_specs.values() if spec.tokens_per_state > 1 + } + if not idx_pages: + return None + + assert all(spec.page_size_padded is None for spec in mla_specs.values()) + assert len(idx_pages) == 1 + mla_names = [name for name, spec in mla_specs.items() if spec.tokens_per_state == 1] + mla_pages = {mla_specs[name].page_size_bytes for name in mla_names} + assert len(mla_pages) == 1 + mla_page = mla_pages.pop() + uniform_spec = UniformTypeKVCacheSpecs.from_specs(attn_specs) + assert uniform_spec is not None + + tail_group: KVCacheGroupSpec | None = None + if tail_specs: + idx_page = next(iter(idx_pages)) + padded_tail_specs: dict[str, KVCacheSpec] = { + name: replace(spec, page_size_padded=idx_page) + for name, spec in tail_specs.items() + } + tail_uniform = UniformTypeKVCacheSpecs.from_specs(padded_tail_specs) + assert tail_uniform is not None + tail_group = KVCacheGroupSpec(list(padded_tail_specs), tail_uniform) + + any_mamba = next(iter(mamba_specs.values())) + assert all(spec == any_mamba for spec in mamba_specs.values()) + if any_mamba.real_page_size_bytes > mla_page: + raise ValueError( + f"the mamba state page ({any_mamba.real_page_size_bytes} bytes) " + f"does not fit the MLA page ({mla_page} bytes); increase tensor " + "parallelism or use a wider KV cache dtype" + ) + padded_specs: dict[str, KVCacheSpec] = { + name: replace(any_mamba, page_size_padded=mla_page) for name in mamba_specs + } + num_groups = _pp_balanced_mamba_group_count( + vllm_config, list(mamba_specs), mla_names + ) + if num_groups is None: + raise ValueError( + "a pipeline stage has mamba layers but no MLA layer to share " + "slots with; realign the stage boundaries (VLLM_PP_LAYER_PARTITION)" + ) + mamba_grouped_names: list[list[str]] = [[] for _ in range(num_groups)] + for index, name in enumerate(mamba_specs): + mamba_grouped_names[index % num_groups].append(name) + + return ( + [KVCacheGroupSpec(list(attn_specs), uniform_spec)] + + ([tail_group] if tail_group is not None else []) + + create_kv_cache_group_specs(padded_specs, mamba_grouped_names) + ) + + +def _glm5_next_tensor_layout( + kv_cache_groups: list[KVCacheGroupSpec], +) -> ( + tuple[ + KVCacheGroupSpec, + list[KVCacheGroupSpec], + list[str], + list[str], + int, + int, + list[str], + int, + ] + | None +): + """Recognize the GLM-5.3-Flash grouping after optional PP projection.""" + uniform_groups = [ + group + for group in kv_cache_groups + if isinstance(group.kv_cache_spec, UniformTypeKVCacheSpecs) + ] + mamba_groups = [ + group for group in kv_cache_groups if isinstance(group.kv_cache_spec, MambaSpec) + ] + attn_group: KVCacheGroupSpec | None = None + tail_group: KVCacheGroupSpec | None = None + for group in uniform_groups: + inner = cast(UniformTypeKVCacheSpecs, group.kv_cache_spec).kv_cache_specs + if all(type(spec) is MLAAttentionSpec for spec in inner.values()): + attn_group = group + elif all(isinstance(spec, KpoolTailSpec) for spec in inner.values()): + tail_group = group + if attn_group is None or not mamba_groups: + return None + if len(uniform_groups) + len(mamba_groups) != len(kv_cache_groups): + return None + + attn_uniform = cast(UniformTypeKVCacheSpecs, attn_group.kv_cache_spec) + mla_inner = cast(dict[str, MLAAttentionSpec], attn_uniform.kv_cache_specs) + if not all( + type(spec) is MLAAttentionSpec and spec.page_size_padded is None + for spec in mla_inner.values() + ): + return None + mla_names = [ + name for name in attn_group.layer_names if mla_inner[name].tokens_per_state == 1 + ] + idx_names = [ + name for name in attn_group.layer_names if mla_inner[name].tokens_per_state > 1 + ] + mla_pages = {mla_inner[name].page_size_bytes for name in mla_names} + idx_pages = {mla_inner[name].page_size_bytes for name in idx_names} + if len(mla_pages) != 1 or len(idx_pages) != 1: + return None + mla_page = mla_pages.pop() + idx_page = idx_pages.pop() + if any(group.kv_cache_spec.page_size_bytes != mla_page for group in mamba_groups): + return None + + tail_names: list[str] = [] + tail_page = 0 + if tail_group is not None: + tail_names = list(tail_group.layer_names) + tail_inner = cast( + UniformTypeKVCacheSpecs, tail_group.kv_cache_spec + ).kv_cache_specs + tail_pages = { + cast(KpoolTailSpec, spec).unpadded_page_size_bytes + for spec in tail_inner.values() + } + if len(tail_pages) != 1 or len(tail_names) != len(idx_names): + return None + tail_page = tail_pages.pop() + if tail_page > idx_page: + return None + + return ( + attn_group, + mamba_groups, + mla_names, + idx_names, + mla_page, + idx_page, + tail_names, + tail_page, + ) + + def unify_kv_cache_spec_page_size( kv_cache_spec: dict[str, KVCacheSpec], ) -> dict[str, KVCacheSpec]: @@ -1359,6 +1559,10 @@ def _get_kv_cache_bytes_per_block( kv_cache_groups: list[KVCacheGroupSpec], ) -> int: """Return the largest cache group's bytes per block.""" + if (glm5_layout := _glm5_next_tensor_layout(kv_cache_groups)) is not None: + _, _, mla_names, idx_names, mla_page, idx_page, _, _ = glm5_layout + return len(mla_names) * mla_page + len(idx_names) * idx_page + bytes_per_block = max( sum( _get_per_layer_spec(group, layer_name).page_size_bytes @@ -1432,6 +1636,69 @@ def get_kv_cache_config_from_groups( ), ) + if (glm5_layout := _glm5_next_tensor_layout(kv_cache_groups)) is not None: + ( + attn_group, + mamba_groups, + mla_names, + idx_names, + mla_page, + idx_page, + tail_names, + _, + ) = glm5_layout + bytes_per_block = len(mla_names) * mla_page + len(idx_names) * idx_page + num_blocks = may_override_num_blocks( + vllm_config, available_memory // bytes_per_block + ) + size = bytes_per_block * num_blocks + attn_specs = cast( + UniformTypeKVCacheSpecs, attn_group.kv_cache_spec + ).kv_cache_specs + + kv_cache_tensors: list[KVCacheTensor] = [] + + def add_tensor(layer_name: str, spec: KVCacheSpec, offset: int) -> None: + kv_cache_tensors.append( + KVCacheTensor( + size=size, + layers=[layer_name], + layer_stride=spec.page_size_bytes * num_blocks, + block_stride=spec.page_size_bytes, + offset=offset, + ) + ) + + for index, mla_name in enumerate(mla_names): + offset = index * mla_page * num_blocks + add_tensor(mla_name, attn_specs[mla_name], offset) + for group in mamba_groups: + if index < len(group.layer_names): + add_tensor(group.layer_names[index], group.kv_cache_spec, offset) + + idx_base = len(mla_names) * mla_page * num_blocks + for index, idx_name in enumerate(idx_names): + offset = idx_base + index * idx_page * num_blocks + add_tensor(idx_name, attn_specs[idx_name], offset) + if tail_names: + tail_name = tail_names[index] + tail_group = next( + group for group in kv_cache_groups if tail_name in group.layer_names + ) + tail_specs = cast( + UniformTypeKVCacheSpecs, tail_group.kv_cache_spec + ).kv_cache_specs + add_tensor(tail_name, tail_specs[tail_name], offset) + + return KVCacheConfig( + num_blocks=num_blocks, + kv_cache_tensors=kv_cache_tensors, + kv_cache_groups=kv_cache_groups, + prefix_cache_retention_interval=( + vllm_config.cache_config.prefix_cache_retention_interval + ), + ) + layout = vllm_config.cache_config.get_resolved_kv_cache_layout() validate_kv_cache_layout(layout, kv_cache_groups) bytes_per_block = _get_kv_cache_bytes_per_block(kv_cache_groups) @@ -1941,6 +2208,9 @@ def get_kv_cache_groups( # full attention, or all layers are sliding window attention with the # same window size). Put all layers into one group. return _get_kv_cache_groups_uniform_type(uniform_spec) + elif glm5_groups := _get_kv_cache_groups_glm5_next(vllm_config, kv_cache_spec): + return glm5_groups + # Hidden-state layers use their own block table and must not be absorbed # into a compatible attention bucket. hidden_specs = { @@ -2065,6 +2335,30 @@ def _max_memory_usage_bytes_from_groups( if not kv_cache_groups: return 0 + if (glm5_layout := _glm5_next_tensor_layout(kv_cache_groups)) is not None: + ( + attn_group, + mamba_groups, + mla_names, + idx_names, + mla_page, + idx_page, + tail_names, + _, + ) = glm5_layout + uniform_spec = cast(UniformTypeKVCacheSpecs, attn_group.kv_cache_spec) + total_blocks = uniform_spec.max_memory_usage_pages(vllm_config) + total_blocks += sum( + cdiv( + group.kv_cache_spec.max_memory_usage_bytes(vllm_config), + group.kv_cache_spec.page_size_bytes, + ) + for group in mamba_groups + ) + if tail_names: + total_blocks += 1 + return total_blocks * (len(mla_names) * mla_page + len(idx_names) * idx_page) + bytes_per_block = _pool_bytes_per_block(kv_cache_groups) total_blocks = 0 for group in kv_cache_groups: diff --git a/vllm/v1/core/single_type_kv_cache_manager.py b/vllm/v1/core/single_type_kv_cache_manager.py index 6da0b8c0705b..b3e0abe7e299 100644 --- a/vllm/v1/core/single_type_kv_cache_manager.py +++ b/vllm/v1/core/single_type_kv_cache_manager.py @@ -6,6 +6,7 @@ from collections.abc import Sequence from typing import ClassVar +from vllm.logger import init_logger from vllm.utils.math_utils import cdiv from vllm.v1.core.block_pool import BlockPool from vllm.v1.core.kv_cache_utils import ( @@ -22,6 +23,7 @@ CrossAttentionSpec, FullAttentionSpec, HiddenStateCacheSpec, + KpoolTailSpec, KVCacheSpec, MambaSpec, MLAAttentionSpec, @@ -33,6 +35,8 @@ from vllm.v1.kv_cache_spec_registry import KVCacheSpecRegistry from vllm.v1.request import Request +logger = init_logger(__name__) + class SingleTypeKVCacheManager(ABC): """ @@ -1210,6 +1214,10 @@ def get_num_skipped_tokens(self, num_computed_tokens: int) -> int: return 0 +class KpoolTailManager(CircularBufferManager): + """One-block circular scratch manager for ``KpoolTailSpec``.""" + + class ChunkedLocalAttentionManager(SingleTypeKVCacheManager): def __init__(self, kv_cache_spec: ChunkedLocalAttentionSpec, **kwargs) -> None: super().__init__(kv_cache_spec, **kwargs) @@ -2085,6 +2093,11 @@ def register_all_kvcache_specs(vllm_config): SlidingWindowManager, uniform_type_base_spec=SlidingWindowMLASpec, ) + KVCacheSpecRegistry.register( + KpoolTailSpec, + KpoolTailManager, + uniform_type_base_spec=KpoolTailSpec, + ) KVCacheSpecRegistry.register( MambaSpec, MambaManager, uniform_type_base_spec=MambaSpec diff --git a/vllm/v1/engine/core.py b/vllm/v1/engine/core.py index c734e67229c1..06ebfb730519 100644 --- a/vllm/v1/engine/core.py +++ b/vllm/v1/engine/core.py @@ -333,8 +333,19 @@ def _initialize_kv_caches(self, vllm_config: VllmConfig) -> KVCacheConfig: vllm_config.cache_config.num_gpu_blocks = scheduler_kv_cache_config.num_blocks kv_cache_groups = scheduler_kv_cache_config.kv_cache_groups if kv_cache_groups: + # Exclude groups that opt out of prefix caching (e.g. GLM-5.3-Flash + # kpool tail, a 1-block/req scratch buffer with block_size=kpool): + # their small block_size would otherwise drag the global block_size + # below the real allocator block size and desync it from mamba. + participating = [ + g.kv_cache_spec.block_size + for g in kv_cache_groups + if g.kv_cache_spec.prefix_cacheable + ] vllm_config.cache_config.block_size = min( - g.kv_cache_spec.block_size for g in kv_cache_groups + participating + if participating + else [g.kv_cache_spec.block_size for g in kv_cache_groups] ) update_kv_cache_capacity(vllm_config, scheduler_kv_cache_config) diff --git a/vllm/v1/kv_cache_interface.py b/vllm/v1/kv_cache_interface.py index fbb3a22611c0..d322a98d1e46 100644 --- a/vllm/v1/kv_cache_interface.py +++ b/vllm/v1/kv_cache_interface.py @@ -158,6 +158,7 @@ class KVCacheSpec: @property def prefix_cacheable(self) -> bool: + """Whether this spec's group participates in prefix caching.""" return True @property @@ -555,6 +556,8 @@ class MLAAttentionSpec(FullAttentionSpec): # DeepseekV4 only fields. Non-DeepseekV4 MLA models leave these at defaults. alignment: int | None = None # Default to None for no padding. model_version: str | None = None + storage_block_size: int | None = None + """Token width used to view storage when it differs from the kernel block.""" # Marks draft groups that flatten a non-causal query block into decode rows. non_causal_multi_token_decode: bool = False # MLA stores a single latent vector per state; there is no separate V. @@ -572,13 +575,16 @@ def merge(cls, specs: list[Self]) -> Self: cache_dtype_str_set = set(spec.cache_dtype_str for spec in specs) tokens_per_state_set = set(spec.tokens_per_state for spec in specs) model_version_set = set(spec.model_version for spec in specs) + storage_block_size_set = set(spec.storage_block_size for spec in specs) assert ( len(cache_dtype_str_set) == 1 and len(tokens_per_state_set) == 1 and len(model_version_set) == 1 + and len(storage_block_size_set) == 1 ), ( "All attention layers in the same KV cache group must use the same " - "quantization method, tokens per state, and model version." + "quantization method, tokens per state, model version, and storage " + "block size." ) non_causal_mtd_set = {spec.non_causal_multi_token_decode for spec in specs} assert len(non_causal_mtd_set) == 1, ( @@ -597,6 +603,7 @@ def merge(cls, specs: list[Self]) -> Self: cache_dtype_str=cache_dtype_str_set.pop(), tokens_per_state=tokens_per_state_set.pop(), model_version=model_version_set.pop(), + storage_block_size=storage_block_size_set.pop(), non_causal_multi_token_decode=non_causal_mtd_set.pop(), ) for spec in specs: @@ -853,6 +860,28 @@ def is_uniform_with_collection( ) +@dataclass(frozen=True, kw_only=True) +class KpoolTailSpec(SlidingWindowSpec): + """One-block circular scratch cache for a kpool indexer's raw tail.""" + + def max_admission_blocks_per_request( + self, max_in_flight_tokens: int, max_model_len: int + ) -> int: + return 1 + + def max_num_blocks_per_req(self, vllm_config: VllmConfig, max_len: int) -> int: + return 1 + + def is_uniform_with_collection( + self, kv_cache_specs: dict[str, KVCacheSpec] + ) -> bool: + return all(isinstance(spec, KpoolTailSpec) for spec in kv_cache_specs.values()) + + @property + def prefix_cacheable(self) -> bool: + return False + + @dataclass(frozen=True) class MambaSpec(KVCacheSpec): shapes: tuple[tuple[int, ...], ...] @@ -875,6 +904,10 @@ def state_content_size_bytes(self) -> int: for (shape, dtype) in zip(self.shapes, self.dtypes) ) + @property + def real_page_size_bytes(self) -> int: + return self.state_content_size_bytes + @property def page_size_bytes(self) -> int: page_size = sum( diff --git a/vllm/v1/spec_decode/llm_base_proposer.py b/vllm/v1/spec_decode/llm_base_proposer.py index 0eaba01294da..1556ddb4411b 100644 --- a/vllm/v1/spec_decode/llm_base_proposer.py +++ b/vllm/v1/spec_decode/llm_base_proposer.py @@ -1021,6 +1021,7 @@ def model_returns_tuple(self) -> bool: { "DeepSeekMTPModel", "DeepseekV32MTPModel", + "Glm5NextMTPModel", "KimiK3MTPModel", }.intersection(architectures) ) diff --git a/vllm/v1/worker/gpu/model_runner.py b/vllm/v1/worker/gpu/model_runner.py index 24e38d47a74b..d7dd9441d71a 100644 --- a/vllm/v1/worker/gpu/model_runner.py +++ b/vllm/v1/worker/gpu/model_runner.py @@ -696,6 +696,7 @@ def _init_kv_zero_meta(self) -> None: attn_groups_iter=(g for groups in self.attn_groups for g in groups), kernel_block_sizes=self.kernel_block_sizes, static_forward_context=self.compilation_config.static_forward_context, + num_blocks=self.kv_cache_config.num_blocks, ) @torch.inference_mode() diff --git a/vllm/v1/worker/gpu/model_states/mamba_hybrid.py b/vllm/v1/worker/gpu/model_states/mamba_hybrid.py index 6d07c1588ca0..9564120f2338 100644 --- a/vllm/v1/worker/gpu/model_states/mamba_hybrid.py +++ b/vllm/v1/worker/gpu/model_states/mamba_hybrid.py @@ -322,6 +322,7 @@ def prepare_attn( kv_cache_config=kv_cache_config, seq_lens_cpu_upper_bound=seq_lens_cpu_upper_bound, dcp_local_seq_lens=input_batch.dcp_local_seq_lens, + positions=input_batch.positions, model_specific_attn_metadata=mamba_attn_metadata, for_cudagraph_capture=for_capture, rswa_prefix_lens=input_batch.prompt_lens, diff --git a/vllm/v1/worker/gpu_model_runner.py b/vllm/v1/worker/gpu_model_runner.py index d2271cc7935b..72981125aca2 100644 --- a/vllm/v1/worker/gpu_model_runner.py +++ b/vllm/v1/worker/gpu_model_runner.py @@ -1160,6 +1160,7 @@ def _init_kv_zero_meta(self) -> None: kernel_block_sizes=self._kernel_block_sizes, runner_only_attn_layers=self.runner_only_attn_layers, static_forward_context=self.compilation_config.static_forward_context, + num_blocks=self.kv_cache_config.num_blocks, ) def _zero_block_ids(self, block_ids: list[int]) -> None: diff --git a/vllm/v1/worker/utils.py b/vllm/v1/worker/utils.py index d767a8fffd05..067a2b0d2267 100644 --- a/vllm/v1/worker/utils.py +++ b/vllm/v1/worker/utils.py @@ -35,6 +35,7 @@ KVCacheLayout, KVCacheSpec, MambaSpec, + MLAAttentionSpec, UniformTypeKVCacheSpecs, create_kv_cache_views, ) @@ -115,6 +116,7 @@ def __init__( attn_groups_iter: Iterable["AttentionGroup"], kernel_block_sizes: list[int], static_forward_context: dict[str, Any], + num_blocks: int, runner_only_attn_layers: set[str] | None = None, ) -> None: """Precompute the absolute-address table for the Triton zeroing kernel. @@ -158,8 +160,6 @@ def __init__( continue kernel_bs = kernel_block_sizes[group.kv_cache_group_id] assert spec.block_size % kernel_bs == 0 - ratio = spec.block_size // kernel_bs - for layer_name in group.layer_names: if layer_name in runner_only_attn_layers: continue @@ -168,6 +168,12 @@ def __init__( continue dp = kv.data_ptr() + assert kv.shape[0] % num_blocks == 0, ( + f"{layer_name}: {kv.shape[0]} kernel blocks is not a " + f"multiple of {num_blocks} logical blocks" + ) + ratio = kv.shape[0] // num_blocks + el = kv.element_size() block_stride_bytes = kv.stride(0) * el assert block_stride_bytes % 4 == 0 @@ -267,11 +273,19 @@ def create_metadata_builders( kernel_block_size: int | None = None, num_metadata_builders: int = 1, ): - kv_cache_spec_builder = ( - self.kv_cache_spec.copy_with_new_block_size(kernel_block_size) - if kernel_block_size is not None - else self.kv_cache_spec - ) + if kernel_block_size is None: + kv_cache_spec_builder = self.kv_cache_spec + elif ( + isinstance(self.kv_cache_spec, MLAAttentionSpec) + and self.kv_cache_spec.storage_block_size is not None + ): + kv_cache_spec_builder = self.kv_cache_spec.copy_with_new_block_size( + self.kv_cache_spec.storage_block_size + ) + else: + kv_cache_spec_builder = self.kv_cache_spec.copy_with_new_block_size( + kernel_block_size + ) builder_cls = self.backend.get_builder_cls() builder_kwargs = {} if builder_cls.requires_block_table_width: @@ -291,6 +305,9 @@ def create_metadata_builders( ) for _ in range(num_metadata_builders) ] + if kernel_block_size is not None: + for builder in self.metadata_builders: + builder.set_kernel_block_size(kernel_block_size) def get_metadata_builder(self, ubatch_id: int = 0) -> AttentionMetadataBuilder: assert len(self.metadata_builders) > ubatch_id @@ -427,6 +444,8 @@ def allocate_kv_cache( kernel_block_size = None if kernel_block_sizes is not None and group_id < len(kernel_block_sizes): kernel_block_size = kernel_block_sizes[group_id] + if isinstance(spec, MLAAttentionSpec) and spec.storage_block_size is not None: + kernel_block_size = spec.storage_block_size views = create_kv_cache_views( buf,