From f75ee63156477a35961b7da7b750821159ecdccd Mon Sep 17 00:00:00 2001 From: AlpinDale Date: Mon, 7 Sep 2026 19:06:35 +0000 Subject: [PATCH 01/36] [sync] [Kernel] SM 12.x blockwise FP8: swizzle the CTA raster when the weight exceeds the L2 (#55180) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Upstream-vLLM: 4df80187b5ca8453eb7b47d2914f0555acf9c12f Co-authored-by: Jürgen Schmied --- .sync/vllm-sha | 2 +- .../c3x/scaled_mm_blockwise_sm120_fp8.cu | 33 +++++++++++++++++-- ...scaled_mm_blockwise_sm120_fp8_dispatch.cuh | 20 +++++++---- .../quantization/test_cutlass_scaled_mm.py | 13 +++++++- 4 files changed, 57 insertions(+), 11 deletions(-) diff --git a/.sync/vllm-sha b/.sync/vllm-sha index c994ce70ed..d7a91604fc 100644 --- a/.sync/vllm-sha +++ b/.sync/vllm-sha @@ -1 +1 @@ -6865e67f0be02d53694517f6f71d7fb96492792d +4df80187b5ca8453eb7b47d2914f0555acf9c12f diff --git a/csrc/libtorch_stable/quantization/w8a8/cutlass/c3x/scaled_mm_blockwise_sm120_fp8.cu b/csrc/libtorch_stable/quantization/w8a8/cutlass/c3x/scaled_mm_blockwise_sm120_fp8.cu index f2d8277894..1c2f29cb0e 100644 --- a/csrc/libtorch_stable/quantization/w8a8/cutlass/c3x/scaled_mm_blockwise_sm120_fp8.cu +++ b/csrc/libtorch_stable/quantization/w8a8/cutlass/c3x/scaled_mm_blockwise_sm120_fp8.cu @@ -4,18 +4,45 @@ namespace aphrodite { +namespace { + +// CTA swizzle for SM 12.x parts whose L2 does not hold the weight operand. +// +// On a GB10 (24 MiB L2) the blockwise kernel loses most of its throughput once +// the weight is re-streamed from DRAM per row of M tiles: a 16384x2560 FP8 +// weight runs at 165 TFLOPS at M=4096 but 90 at M=8192 and 52 at M>=16384; +// 8192x8192 is at 54 from M=6144. With the tile scheduler's max_swizzle_size = +// 8 the same launches run at 150-174 TFLOPS at every M, bit-identical to the +// default order (ten N/K shapes, M 2048-16384, all cells identical). The one +// place the default order is better is a narrow band around M=4096 on the +// 2560-wide weights (167 vs 153 at 16384x2560, 160 vs 154 at 12288x2560); +// elsewhere the swizzled order is equal or up to 3.3x faster, so it is used +// whenever the weight exceeds the L2. Parts whose L2 holds the weight (RTX PRO +// 6000 Blackwell / GB202: 96-128 MiB) keep the default order, which is also +// the faster one there (2560x6144 at 15 MiB: 178 vs 163 at M=2048). +constexpr int kBlockwiseFp8SwizzleSize = 8; + +int blockwise_fp8_swizzle_size(int64_t weight_bytes) { + const int64_t l2_bytes = get_device_prop()->l2CacheSize; + return (l2_bytes > 0 && weight_bytes > l2_bytes) ? kBlockwiseFp8SwizzleSize + : 1; +} + +} // namespace + void cutlass_scaled_mm_blockwise_sm120_fp8( torch::stable::Tensor& out, torch::stable::Tensor const& a, torch::stable::Tensor const& b, torch::stable::Tensor const& a_scales, torch::stable::Tensor const& b_scales) { + // b is [K, N] FP8 (one byte per element). + const int swizzle = blockwise_fp8_swizzle_size(b.size(1) * b.size(0)); if (out.scalar_type() == torch::headeronly::ScalarType::BFloat16) { cutlass_gemm_blockwise_sm120_fp8_dispatch( - out, a, b, a_scales, b_scales); - + out, a, b, a_scales, b_scales, swizzle); } else { STD_TORCH_CHECK(out.scalar_type() == torch::headeronly::ScalarType::Half); cutlass_gemm_blockwise_sm120_fp8_dispatch( - out, a, b, a_scales, b_scales); + out, a, b, a_scales, b_scales, swizzle); } } diff --git a/csrc/libtorch_stable/quantization/w8a8/cutlass/c3x/scaled_mm_blockwise_sm120_fp8_dispatch.cuh b/csrc/libtorch_stable/quantization/w8a8/cutlass/c3x/scaled_mm_blockwise_sm120_fp8_dispatch.cuh index 1a7bf3db81..81b453dc4a 100644 --- a/csrc/libtorch_stable/quantization/w8a8/cutlass/c3x/scaled_mm_blockwise_sm120_fp8_dispatch.cuh +++ b/csrc/libtorch_stable/quantization/w8a8/cutlass/c3x/scaled_mm_blockwise_sm120_fp8_dispatch.cuh @@ -176,7 +176,8 @@ template void cutlass_gemm_caller_blockwise(torch::stable::Tensor& out, torch::stable::Tensor const& a, torch::stable::Tensor const& b, torch::stable::Tensor const& a_scales, - torch::stable::Tensor const& b_scales) { + torch::stable::Tensor const& b_scales, + int max_swizzle_size) { static constexpr bool swap_ab = Gemm::swap_ab; using GemmKernel = typename Gemm::GemmKernel; using StrideA = typename Gemm::GemmKernel::StrideA; @@ -238,8 +239,14 @@ void cutlass_gemm_caller_blockwise(torch::stable::Tensor& out, torch::stable::Te auto c_ptr = static_cast(out.data_ptr()); typename GemmKernel::EpilogueArguments epilogue_args{ {}, c_ptr, c_stride, c_ptr, c_stride}; + // CTA rasterization: max_swizzle_size > 1 groups nearby M/N tiles in the + // persistent scheduler's raster, improving the temporal locality of the + // shared weight (B) tiles in the L2. Bit-identical to the default order + // (each output tile's K-reduction is unchanged). + typename GemmKernel::TileSchedulerArguments scheduler{}; + scheduler.max_swizzle_size = max_swizzle_size; c3x::cutlass_gemm_caller(a.device(), prob_shape, mainloop_args, - epilogue_args); + epilogue_args, scheduler); } template @@ -247,7 +254,8 @@ void cutlass_gemm_blockwise_sm120_fp8_dispatch(torch::stable::Tensor& out, torch::stable::Tensor const& a, torch::stable::Tensor const& b, torch::stable::Tensor const& a_scales, - torch::stable::Tensor const& b_scales) { + torch::stable::Tensor const& b_scales, + int max_swizzle_size) { int M = a.size(0); // more heuristic tuning can be done here by checking N/K dimensions as well bool swap_ab = (M <= 64); @@ -256,18 +264,18 @@ void cutlass_gemm_blockwise_sm120_fp8_dispatch(torch::stable::Tensor& out, if (M <= 256) { using Gemm = typename sm120_blockwise_fp8_config_pingpong::Gemm; return cutlass_gemm_caller_blockwise( - out, a, b, a_scales, b_scales); + out, a, b, a_scales, b_scales, max_swizzle_size); } // M > 256: use default 128x128x128 config with Cooperative (Auto) schedule using Gemm = typename sm120_blockwise_fp8_config_default::Gemm; return cutlass_gemm_caller_blockwise( - out, a, b, a_scales, b_scales); + out, a, b, a_scales, b_scales, max_swizzle_size); } else { // Swap A/B for small M to improve performance // Use TILE_N=32 as the minimum compatible tile size. using Gemm = typename sm120_blockwise_fp8_config_swapab::Gemm; return cutlass_gemm_caller_blockwise( - out, a, b, a_scales, b_scales); + out, a, b, a_scales, b_scales, max_swizzle_size); } } diff --git a/tests/kernels/quantization/test_cutlass_scaled_mm.py b/tests/kernels/quantization/test_cutlass_scaled_mm.py index 26e7941172..c066aa915b 100644 --- a/tests/kernels/quantization/test_cutlass_scaled_mm.py +++ b/tests/kernels/quantization/test_cutlass_scaled_mm.py @@ -215,7 +215,18 @@ def test_cutlass_fp8_gemm_padded(m: int, n: int, k: int, a_scale_group_shape, b_ torch.testing.assert_close(out, baseline, rtol=5e-1, atol=1.5e-1) -@pytest.mark.parametrize("m,n,k", MNK_FACTORS) +# Two prefill-sized cases for the SM 12.x blockwise path, where the op switches +# the tile scheduler to a swizzled CTA order once the weight exceeds the L2 +# (16384x2560 = 40 MiB and 5120x5120 = 25 MiB on a 24 MiB L2 part); the odd M +# also covers the non-swap-AB dispatch. Elsewhere they run the default order +# like the rest of the list. +BLOCKWISE_PREFILL_FACTORS = [ + (8193, 16384, 2560), + (5120, 5120, 5120), +] + + +@pytest.mark.parametrize("m,n,k", MNK_FACTORS + BLOCKWISE_PREFILL_FACTORS) @pytest.mark.parametrize("a_scale_group_shape,b_scale_group_shape", [((1, 128), (128, 128))]) @pytest.mark.parametrize("use_bias", [False]) @pytest.mark.skipif( From 63b0403d6d30d72d59883be903868c600b072f98 Mon Sep 17 00:00:00 2001 From: AlpinDale Date: Mon, 7 Sep 2026 19:11:09 +0000 Subject: [PATCH 02/36] [sync] [Kernel] Remove unused fake implementation (#55535) Upstream-vLLM: 199cb9b964822e59ab9b58d88e7be31eb419a2ae Co-authored-by: Jee Jee Li --- .sync/vllm-sha | 2 +- aphrodite/_aiter_ops.py | 97 ------------- aphrodite/_custom_ops.py | 56 -------- aphrodite/_xpu_ops.py | 60 -------- .../passes/fusion/allreduce_rms_fusion.py | 19 --- .../ops/dynamic_per_token_scaled_fp8_quant.py | 10 -- .../kernels/helion/ops/fused_qk_norm_rope.py | 18 --- .../helion/ops/per_token_group_fp8_quant.py | 18 --- .../ops/rms_norm_dynamic_per_token_quant.py | 13 -- .../helion/ops/rms_norm_per_block_quant.py | 15 -- .../ops/silu_and_mul_per_block_quant.py | 12 -- aphrodite/kernels/helion/register.py | 9 +- .../ops/triton_ops/fused_moe_lora_fp8_op.py | 133 ------------------ .../lora/ops/triton_ops/fused_moe_lora_op.py | 113 --------------- .../lora/ops/triton_ops/lora_expand_fp8_op.py | 24 ---- .../lora/ops/triton_ops/lora_expand_op.py | 18 --- .../lora/ops/triton_ops/lora_shrink_fp8_op.py | 23 --- .../lora/ops/triton_ops/lora_shrink_op.py | 17 --- aphrodite/lora/ops/xpu_ops/lora_ops.py | 35 ----- .../kernels/linear/scaled_mm/deep_gemm.py | 12 -- .../layers/attention/attention.py | 14 -- .../layers/attention/mla_attention.py | 19 --- .../layers/attention/static_sink_attention.py | 8 -- .../model_executor/layers/hpc/rope_norm.py | 10 -- .../layers/mamba/gdn/olmo_gdn_linear_attn.py | 10 -- .../layers/mamba/gdn/qwen_gdn_linear_attn.py | 23 --- .../mamba/linear/minimax_linear_attn.py | 10 -- .../layers/mamba/mamba_mixer.py | 9 -- .../layers/mamba/mamba_mixer2.py | 9 -- .../layers/mamba/ops/cpu/gdn_attention.py | 12 -- .../model_executor/layers/mamba/short_conv.py | 9 -- .../layers/rotary_embedding/common.py | 12 -- .../model_executor/models/bailing_moe_v3.py | 13 -- .../model_executor/offloader/prefetch_ops.py | 18 --- aphrodite/models/deepseek_v4/xpu/model.py | 13 -- aphrodite/models/qwen4_exp/amd/ple_layer.py | 18 --- aphrodite/models/qwen4_exp/amd/qsa.py | 13 -- .../models/qwen4_exp/nvidia/ple_layer.py | 20 --- aphrodite/models/qwen4_exp/nvidia/qsa.py | 13 -- aphrodite/utils/flashinfer.py | 8 -- 40 files changed, 8 insertions(+), 957 deletions(-) diff --git a/.sync/vllm-sha b/.sync/vllm-sha index d7a91604fc..163c52426a 100644 --- a/.sync/vllm-sha +++ b/.sync/vllm-sha @@ -1 +1 @@ -4df80187b5ca8453eb7b47d2914f0555acf9c12f +199cb9b964822e59ab9b58d88e7be31eb419a2ae diff --git a/aphrodite/_aiter_ops.py b/aphrodite/_aiter_ops.py index 304681bbf6..e57f5ae586 100644 --- a/aphrodite/_aiter_ops.py +++ b/aphrodite/_aiter_ops.py @@ -375,18 +375,6 @@ def _rocm_aiter_topk_softmax_impl( ) -def _rocm_aiter_topk_softmax_fake( - topk_weights: torch.Tensor, - topk_indices: torch.Tensor, - token_expert_indices: torch.Tensor, - gating_output: torch.Tensor, - renormalize: bool, - num_shared_experts: int = 0, - shared_expert_scoring_func: str = "", -) -> None: - pass - - def _rocm_aiter_topk_sigmoid_impl( topk_weights: torch.Tensor, topk_indices: torch.Tensor, @@ -397,14 +385,6 @@ def _rocm_aiter_topk_sigmoid_impl( topk_sigmoid(topk_weights, topk_indices, gating_output) -def _rocm_aiter_topk_sigmoid_fake( - topk_weights: torch.Tensor, - topk_indices: torch.Tensor, - gating_output: torch.Tensor, -) -> None: - pass - - def _rocm_aiter_biased_grouped_topk_impl( gating_output: torch.Tensor, correction_bias: torch.Tensor, @@ -429,19 +409,6 @@ def _rocm_aiter_biased_grouped_topk_impl( ) -def _rocm_aiter_biased_grouped_topk_fake( - gating_output: torch.Tensor, - correction_bias: torch.Tensor, - topk_weights: torch.Tensor, - topk_ids: torch.Tensor, - num_expert_group: int, - topk_group: int, - need_renorm: bool, - routed_scaling_factor: float = 1.0, # mul to topk_weights -) -> None: - pass - - def _rocm_aiter_grouped_topk_impl( gating_output: torch.Tensor, topk_weights: torch.Tensor, @@ -467,19 +434,6 @@ def _rocm_aiter_grouped_topk_impl( ) -def _rocm_aiter_grouped_topk_fake( - gating_output: torch.Tensor, - topk_weights: torch.Tensor, - topk_ids: torch.Tensor, - num_expert_group: int, - topk_group: int, - need_renorm: bool, - scoring_func: str = "softmax", - routed_scaling_factor: float = 1.0, # mul to topk_weights -) -> None: - pass - - def _rocm_aiter_fused_topk_impl( x: torch.Tensor, router_logits: torch.Tensor, @@ -592,29 +546,6 @@ def _rocm_aiter_mla_decode_fwd_impl( ) -def _rocm_aiter_mla_decode_fwd_fake( - q: torch.Tensor, - kv_buffer: torch.Tensor, - o: torch.Tensor, - qo_indptr: torch.Tensor, - max_seqlen_qo: int, - kv_indptr: torch.Tensor | None = None, - kv_indices: torch.Tensor | None = None, - kv_last_page_lens: torch.Tensor | None = None, - sm_scale: float = 1.0, - logit_cap: float = 0.0, - q_scale: torch.Tensor | None = None, - kv_scale: torch.Tensor | None = None, - work_meta_data: torch.Tensor | None = None, - work_indptr: torch.Tensor | None = None, - work_info_set: torch.Tensor | None = None, - reduce_indptr: torch.Tensor | None = None, - reduce_final_map: torch.Tensor | None = None, - reduce_partial_map: torch.Tensor | None = None, -) -> None: - pass - - def _rocm_aiter_w8a8_gemm_impl( A: torch.Tensor, B: torch.Tensor, @@ -1076,15 +1007,6 @@ def _rocm_aiter_per_tensor_quant_impl( static_per_tensor_quant(out, x, scale) -def _rocm_aiter_per_tensor_quant_fake( - out: torch.Tensor, - x: torch.Tensor, - scale: torch.Tensor, - is_dynamic: bool, -) -> None: - pass - - def _rocm_aiter_per_token_quant_impl( x: torch.Tensor, quant_dtype: torch.dtype, scale: torch.Tensor | None = None ) -> tuple[torch.Tensor, torch.Tensor]: @@ -1566,18 +1488,6 @@ def _triton_rotary_embedding_impl( key = key.view(key_shape) -def _triton_rotary_embedding_fake( - positions: torch.Tensor, - query: torch.Tensor, - key: torch.Tensor, - head_size: int, - cos_sin_cache: torch.Tensor, - is_neox_style: bool, - offsets: torch.Tensor | None = None, -) -> None: - return - - def _rocm_aiter_fp8_attn_impl( q: torch.Tensor, k: torch.Tensor, @@ -2085,7 +1995,6 @@ def register_ops_once() -> None: op_name="rocm_aiter_topk_softmax", op_func=_rocm_aiter_topk_softmax_impl, mutates_args=["topk_weights", "topk_indices", "token_expert_indices"], - fake_impl=_rocm_aiter_topk_softmax_fake, dispatch_key=current_platform.dispatch_key, ) @@ -2093,7 +2002,6 @@ def register_ops_once() -> None: op_name="rocm_aiter_topk_sigmoid", op_func=_rocm_aiter_topk_sigmoid_impl, mutates_args=["topk_weights", "topk_indices"], - fake_impl=_rocm_aiter_topk_sigmoid_fake, dispatch_key=current_platform.dispatch_key, ) @@ -2101,7 +2009,6 @@ def register_ops_once() -> None: op_name="rocm_aiter_biased_grouped_topk", op_func=_rocm_aiter_biased_grouped_topk_impl, mutates_args=["topk_weights", "topk_ids"], - fake_impl=_rocm_aiter_biased_grouped_topk_fake, dispatch_key=current_platform.dispatch_key, ) @@ -2109,7 +2016,6 @@ def register_ops_once() -> None: op_name="rocm_aiter_grouped_topk", op_func=_rocm_aiter_grouped_topk_impl, mutates_args=["topk_weights", "topk_ids"], - fake_impl=_rocm_aiter_grouped_topk_fake, dispatch_key=current_platform.dispatch_key, ) @@ -2125,7 +2031,6 @@ def register_ops_once() -> None: op_name="rocm_aiter_mla_decode_fwd", op_func=_rocm_aiter_mla_decode_fwd_impl, mutates_args=["o"], - fake_impl=_rocm_aiter_mla_decode_fwd_fake, ) direct_register_custom_op( @@ -2219,7 +2124,6 @@ def register_ops_once() -> None: op_name="rocm_aiter_per_tensor_quant", op_func=_rocm_aiter_per_tensor_quant_impl, mutates_args=["out", "scale"], - fake_impl=_rocm_aiter_per_tensor_quant_fake, dispatch_key=current_platform.dispatch_key, ) @@ -2259,7 +2163,6 @@ def register_ops_once() -> None: op_name="rocm_aiter_triton_rotary_embedding", op_func=_triton_rotary_embedding_impl, mutates_args=["query", "key"], # These tensors are modified in-place - fake_impl=_triton_rotary_embedding_fake, ) direct_register_custom_op( diff --git a/aphrodite/_custom_ops.py b/aphrodite/_custom_ops.py index efd7f58215..dda48ddf59 100644 --- a/aphrodite/_custom_ops.py +++ b/aphrodite/_custom_ops.py @@ -99,17 +99,6 @@ def _scaled_fp4_quant_fake( m = input.numel() // n return create_fp4_output_tensors(m, n, input.device, is_sf_swizzled_layout) - @register_fake("_C::scaled_fp4_quant.out") - def _scaled_fp4_quant_out_fake( - input: torch.Tensor, - input_scale: torch.Tensor, - is_sf_swizzled_layout: bool, - *, - output: torch.Tensor, - output_scale: torch.Tensor, - ) -> None: - return None - def exl3_gemm( a: torch.Tensor, @@ -1020,27 +1009,6 @@ def moe_gptq_gemm_rdna3( ) -if hasattr(torch.ops, "_rocm_C") and hasattr(torch.ops._rocm_C, "moe_gptq_gemm_rdna3"): - - @register_fake("_rocm_C::moe_gptq_gemm_rdna3") - def _moe_gptq_gemm_rdna3_fake( - a: torch.Tensor, - c: torch.Tensor, - b_q_weight: torch.Tensor, - b_scales: torch.Tensor, - b_qzeros: torch.Tensor, - topk_weights: torch.Tensor, - sorted_token_ids: torch.Tensor, - expert_ids: torch.Tensor, - num_tokens_post_padded: torch.Tensor, - top_k: int, - block_size_m: int, - mul_topk_weight: bool, - output_topk: int = 0, - ) -> None: - return - - if hasattr(torch.ops._C, "allspark_w8a16_gemm"): @register_fake("_C::allspark_w8a16_gemm") @@ -2989,17 +2957,6 @@ def fp32_router_gemm( return output -if hasattr(torch.ops, "_C") and hasattr(torch.ops._C, "fp32_router_gemm"): - - @register_fake("_C::fp32_router_gemm") - def fp32_router_gemm_fake( - output: torch.Tensor, - mat_a: torch.Tensor, - mat_b: torch.Tensor, - ) -> None: - return - - def topk_softmax( topk_weights: torch.Tensor, topk_ids: torch.Tensor, @@ -4919,19 +4876,6 @@ def safeFusedQuantizeNv( torch.ops._qutlass_C.fusedQuantizeNvAbsMax(a, b, xh_e2m1, xh_e4m3, global_scale) -if hasattr(torch.ops._qutlass_C, "fusedQuantizeNv"): - - @register_fake("aphrodite::safeFusedQuantizeNv") - def _fake_fused_quantize_nv( - a: torch.Tensor, - b: torch.Tensor, - xh_e2m1: torch.Tensor, - xh_e4m3: torch.Tensor, - global_scale: torch.Tensor, - ) -> None: - return - - def hadacore_transform(x: torch.Tensor, inplace: bool = True) -> torch.Tensor: """ Perform Hadamard transforms using [Hadacore](https://arxiv.org/abs/2412.08832) diff --git a/aphrodite/_xpu_ops.py b/aphrodite/_xpu_ops.py index 3c5fa323d1..0dfc582594 100644 --- a/aphrodite/_xpu_ops.py +++ b/aphrodite/_xpu_ops.py @@ -125,15 +125,6 @@ def _gemma_rms_norm_impl( torch.ops._C.gemma_rms_norm(out, input, weight, epsilon) -def _gemma_rms_norm_fake( - out: torch.Tensor, - input: torch.Tensor, - weight: torch.Tensor, - epsilon: float, -) -> None: - return None - - def _fused_add_gemma_rms_norm_impl( input: torch.Tensor, residual: torch.Tensor, @@ -143,15 +134,6 @@ def _fused_add_gemma_rms_norm_impl( torch.ops._C.fused_add_gemma_rms_norm(input, residual, weight, epsilon) -def _fused_add_gemma_rms_norm_fake( - input: torch.Tensor, - residual: torch.Tensor, - weight: torch.Tensor, - epsilon: float, -) -> None: - return None - - def _gdn_attention_core_xpu_impl( core_attn_out: torch.Tensor, z: torch.Tensor, @@ -236,16 +218,6 @@ def _gdn_attention_core_xpu_impl( ) -def _gdn_attention_core_xpu_fake( - core_attn_out: torch.Tensor, - z: torch.Tensor, - projected_states_qkvz: torch.Tensor, - projected_states_ba: torch.Tensor, - layer_name: str, -) -> None: - return - - def _xpu_ops_deepseek_scaling_rope_impl( positions: torch.Tensor, query: torch.Tensor, @@ -465,19 +437,6 @@ def _xpu_deepseek_fused_indexer_q_rope_fp8_impl( ) -def _xpu_deepseek_fused_indexer_q_rope_fp8_fake( - index_q: torch.Tensor, - positions: torch.Tensor, - index_q_cos_sin_cache: torch.Tensor, - index_weights: torch.Tensor, - index_weights_softmax_scale: float, - index_weights_head_scale: float, - index_q_fp8: torch.Tensor, - index_weights_out: torch.Tensor, -) -> None: - return - - def _xpu_deepseek_fused_indexer_q_rope_mxfp4_impl( index_q: torch.Tensor, positions: torch.Tensor, @@ -522,20 +481,6 @@ def _xpu_deepseek_fused_indexer_q_rope_mxfp4_impl( ) -def _xpu_deepseek_fused_indexer_q_rope_mxfp4_fake( - index_q: torch.Tensor, - positions: torch.Tensor, - index_q_cos_sin_cache: torch.Tensor, - index_weights: torch.Tensor, - index_weights_softmax_scale: float, - index_weights_head_scale: float, - index_q_packed: torch.Tensor, - index_q_scale: torch.Tensor, - index_weights_out: torch.Tensor, -) -> None: - return - - def _xpu_mxfp8_quantize_impl(x: torch.Tensor, dtype: torch.dtype | None = None) -> tuple[torch.Tensor, torch.Tensor]: MXFP8_BLOCK_SIZE = 32 assert x.shape[-1] % MXFP8_BLOCK_SIZE == 0 @@ -1214,14 +1159,12 @@ def register_ops_once() -> None: op_name="xpu_gemma_rms_norm", op_func=_gemma_rms_norm_impl, mutates_args=["out"], - fake_impl=_gemma_rms_norm_fake, ) direct_register_custom_op( op_name="xpu_fused_add_gemma_rms_norm", op_func=_fused_add_gemma_rms_norm_impl, mutates_args=["input", "residual"], - fake_impl=_fused_add_gemma_rms_norm_fake, ) direct_register_custom_op( @@ -1266,7 +1209,6 @@ def register_ops_once() -> None: op_name="gdn_attention_core_xpu", op_func=eager_break_during_capture(_gdn_attention_core_xpu_impl), mutates_args=["core_attn_out", "z"], - fake_impl=_gdn_attention_core_xpu_fake, ) direct_register_custom_op( @@ -1279,7 +1221,6 @@ def register_ops_once() -> None: op_name="xpu_deepseek_fused_indexer_q_rope_fp8", op_func=_xpu_deepseek_fused_indexer_q_rope_fp8_impl, mutates_args=["index_q_fp8", "index_weights_out"], - fake_impl=_xpu_deepseek_fused_indexer_q_rope_fp8_fake, ) direct_register_custom_op( @@ -1290,7 +1231,6 @@ def register_ops_once() -> None: "index_q_scale", "index_weights_out", ], - fake_impl=_xpu_deepseek_fused_indexer_q_rope_mxfp4_fake, ) _OPS_REGISTERED = True diff --git a/aphrodite/compilation/passes/fusion/allreduce_rms_fusion.py b/aphrodite/compilation/passes/fusion/allreduce_rms_fusion.py index f1a78ba9c3..b3a7b209b5 100644 --- a/aphrodite/compilation/passes/fusion/allreduce_rms_fusion.py +++ b/aphrodite/compilation/passes/fusion/allreduce_rms_fusion.py @@ -287,24 +287,6 @@ def call_trtllm_fused_allreduce_norm( trigger_completion_at_end=(use_oneshot is True) or num_tokens > PDL_ADVANCE_LAUNCH_TOKENS, ) - def call_trtllm_fused_allreduce_norm_fake( - allreduce_in: torch.Tensor, - residual: torch.Tensor, - rms_gamma: torch.Tensor, - rms_eps: float, - world_size: int, - launch_with_pdl: bool, - fp32_acc: bool, - max_token_num: int, - pattern_code: int, - norm_out: torch.Tensor | None = None, - quant_out: torch.Tensor | None = None, - scale_out: torch.Tensor | None = None, - scale_factor: torch.Tensor | None = None, - weight_bias: float = 0.0, - ) -> None: - pass - direct_register_custom_op( op_name="flashinfer_trtllm_fused_allreduce_norm", op_func=call_trtllm_fused_allreduce_norm, @@ -315,7 +297,6 @@ def call_trtllm_fused_allreduce_norm_fake( "quant_out", "scale_out", ], - fake_impl=call_trtllm_fused_allreduce_norm_fake, ) flashinfer_trtllm_fused_allreduce_norm = torch.ops.aphrodite.flashinfer_trtllm_fused_allreduce_norm.default diff --git a/aphrodite/kernels/helion/ops/dynamic_per_token_scaled_fp8_quant.py b/aphrodite/kernels/helion/ops/dynamic_per_token_scaled_fp8_quant.py index 91c573f76e..b7aaa2f9f3 100644 --- a/aphrodite/kernels/helion/ops/dynamic_per_token_scaled_fp8_quant.py +++ b/aphrodite/kernels/helion/ops/dynamic_per_token_scaled_fp8_quant.py @@ -95,15 +95,6 @@ def pick_config(args: tuple[Any, ...], config_keys: list[CaseKey]) -> CaseKey | return result -def fake_impl( - result: torch.Tensor, # [num_tokens, hidden_size] - input: torch.Tensor, # [num_tokens, hidden_size] - scale: torch.Tensor, # [num_tokens, 1] - scale_ub: torch.Tensor | None = None, # scalar tensor -) -> None: - return - - def baseline( result: torch.Tensor, # [num_tokens, hidden_size] input: torch.Tensor, # [num_tokens, hidden_size] @@ -130,7 +121,6 @@ def baseline( mutates_args=["result", "scale"], config_picker=pick_config, input_generator=generate_inputs, - fake_impl=fake_impl, helion_settings=helion.Settings( autotune_baseline_fn=baseline, ignore_warnings=[helion.exc.TensorOperationInWrapper], diff --git a/aphrodite/kernels/helion/ops/fused_qk_norm_rope.py b/aphrodite/kernels/helion/ops/fused_qk_norm_rope.py index 727ebcf619..c06f785d24 100644 --- a/aphrodite/kernels/helion/ops/fused_qk_norm_rope.py +++ b/aphrodite/kernels/helion/ops/fused_qk_norm_rope.py @@ -139,23 +139,6 @@ def pick_config(args: tuple[Any, ...], config_keys: list[CaseKey]) -> CaseKey | return result -def fake_impl( - qkv: torch.Tensor, # [num_tokens, (num_heads_q+num_heads_k+num_heads_v)*head_dim] - num_heads_q: int, - num_heads_k: int, - num_heads_v: int, - head_dim: int, - eps: float, - q_weight: torch.Tensor, - k_weight: torch.Tensor, - cos_sin_cache: torch.Tensor, # [max_position, rotary_dim] - is_neox: bool, - position_ids: torch.Tensor, # [num_tokens], - forced_token_heads_per_warp: int = -1, # dummy -) -> None: - return - - def baseline( qkv: torch.Tensor, # [num_tokens, (num_heads_q+num_heads_k+num_heads_v)*head_dim] num_heads_q: int, @@ -194,7 +177,6 @@ def baseline( mutates_args=["qkv"], config_picker=pick_config, input_generator=generate_inputs, - fake_impl=fake_impl, helion_settings=helion.Settings( autotune_baseline_fn=baseline, autotune_baseline_atol=5e-2, diff --git a/aphrodite/kernels/helion/ops/per_token_group_fp8_quant.py b/aphrodite/kernels/helion/ops/per_token_group_fp8_quant.py index 7584b8e180..588ce1d7f3 100644 --- a/aphrodite/kernels/helion/ops/per_token_group_fp8_quant.py +++ b/aphrodite/kernels/helion/ops/per_token_group_fp8_quant.py @@ -126,23 +126,6 @@ def pick_config(args: tuple[Any, ...], config_keys: list[CaseKey]) -> CaseKey | return result -def fake_impl( - input: torch.Tensor, # [num_tokens, hidden_size] - output_q: torch.Tensor, # [num_tokens, hidden_size] - output_s: torch.Tensor, # [num_tokens, groups_per_row] - group_size: int, - eps: float, - fp8_min: float, - fp8_max: float, - scale_ue8m0: bool, - # Unused dummy args - # Kept for consistency with existing kernel interface - dummy_is_scale_transposed: bool = False, - dummy_is_tma_aligned: bool = False, -) -> None: - return - - def baseline( input: torch.Tensor, # [num_tokens, hidden_size] output_q: torch.Tensor, # [num_tokens, hidden_size] @@ -172,7 +155,6 @@ def baseline( mutates_args=["output_q", "output_s"], config_picker=pick_config, input_generator=generate_inputs, - fake_impl=fake_impl, helion_settings=helion.Settings( autotune_baseline_fn=baseline, ), diff --git a/aphrodite/kernels/helion/ops/rms_norm_dynamic_per_token_quant.py b/aphrodite/kernels/helion/ops/rms_norm_dynamic_per_token_quant.py index 8c68d5f3f4..375789fb51 100644 --- a/aphrodite/kernels/helion/ops/rms_norm_dynamic_per_token_quant.py +++ b/aphrodite/kernels/helion/ops/rms_norm_dynamic_per_token_quant.py @@ -114,18 +114,6 @@ def pick_config(args: tuple[Any, ...], config_keys: list[CaseKey]) -> CaseKey | return result -def fake_impl( - result: torch.Tensor, # [num_tokens, hidden_size] - input: torch.Tensor, # [num_tokens, hidden_size] - weight: torch.Tensor, # [hidden_size] - scale: torch.Tensor, # [num_tokens, 1] - epsilon: float, - scale_ub: torch.Tensor | None = None, # [] - residual: torch.Tensor | None = None, # [num_tokens, hidden_size] -) -> None: - return - - def baseline( result: torch.Tensor, # [num_tokens, hidden_size] input: torch.Tensor, # [num_tokens, hidden_size] @@ -173,7 +161,6 @@ def baseline( mutates_args=["result", "scale", "residual"], config_picker=pick_config, input_generator=generate_inputs, - fake_impl=fake_impl, helion_settings=helion.Settings( autotune_baseline_fn=baseline, ignore_warnings=[helion.exc.TensorOperationInWrapper], diff --git a/aphrodite/kernels/helion/ops/rms_norm_per_block_quant.py b/aphrodite/kernels/helion/ops/rms_norm_per_block_quant.py index a1e8b2826e..98f05594be 100644 --- a/aphrodite/kernels/helion/ops/rms_norm_per_block_quant.py +++ b/aphrodite/kernels/helion/ops/rms_norm_per_block_quant.py @@ -144,20 +144,6 @@ def pick_config(args: tuple[Any, ...], config_keys: list[CaseKey]) -> CaseKey | return result -def fake_impl( - result: torch.Tensor, # [num_tokens, hidden_size] - input: torch.Tensor, # [num_tokens, hidden_size] - weight: torch.Tensor, # [hidden_size] - scale: torch.Tensor, # [num_tokens, groups_per_row] - epsilon: float, - scale_ub: torch.Tensor | None, # [] - residual: torch.Tensor | None, # [num_tokens, hidden_size] - group_size: int, - is_scale_transposed: bool, # dummy -) -> None: - return - - def baseline( result: torch.Tensor, # [num_tokens, hidden_size] input: torch.Tensor, # [num_tokens, hidden_size] @@ -208,7 +194,6 @@ def baseline( mutates_args=["result", "scale", "residual"], config_picker=pick_config, input_generator=generate_inputs, - fake_impl=fake_impl, helion_settings=helion.Settings( autotune_baseline_fn=baseline, ignore_warnings=[helion.exc.TensorOperationInWrapper], diff --git a/aphrodite/kernels/helion/ops/silu_and_mul_per_block_quant.py b/aphrodite/kernels/helion/ops/silu_and_mul_per_block_quant.py index e1ce67354d..d48e35184f 100644 --- a/aphrodite/kernels/helion/ops/silu_and_mul_per_block_quant.py +++ b/aphrodite/kernels/helion/ops/silu_and_mul_per_block_quant.py @@ -124,17 +124,6 @@ def pick_config(args: tuple[Any, ...], config_keys: list[CaseKey]) -> CaseKey | return result -def fake_impl( - out: torch.Tensor, # [num_tokens, intermediate_size] - input: torch.Tensor, # [num_tokens, 2 * intermediate_size] - scales: torch.Tensor, # [num_tokens, groups_per_row] - group_size: int, - scale_ub: torch.Tensor | None = None, # scalar tensor - is_scale_transposed: bool = False, -) -> None: - return - - def baseline( out: torch.Tensor, # [num_tokens, intermediate_size] input: torch.Tensor, # [num_tokens, 2 * intermediate_size] @@ -176,7 +165,6 @@ def baseline( mutates_args=["out", "scales"], config_picker=pick_config, input_generator=generate_inputs, - fake_impl=fake_impl, helion_settings=helion.Settings( autotune_baseline_fn=baseline, ignore_warnings=[helion.exc.TensorOperationInWrapper], diff --git a/aphrodite/kernels/helion/register.py b/aphrodite/kernels/helion/register.py index f2df7c23c6..c17b1a7502 100644 --- a/aphrodite/kernels/helion/register.py +++ b/aphrodite/kernels/helion/register.py @@ -38,6 +38,7 @@ from __future__ import annotations +import inspect from collections.abc import Callable from typing import Any @@ -248,7 +249,7 @@ def __init__( self, raw_kernel_func: Callable, op_name: str, - fake_impl: Callable, + fake_impl: Callable | None, config_picker: ConfigPicker, mutates_args: list[str] | None = None, helion_settings: helion.Settings | None = None, @@ -428,8 +429,12 @@ def decorator(kernel_func: Callable) -> HelionKernelWrapper: f"Use a different op_name or check for duplicate registrations." ) + # A void-returning kernel that mutates its args needs no fake impl: + # PyTorch auto-generates a trivial one for mutable, no-return schemas. + returns_none = inspect.signature(kernel_func).return_annotation is None + final_fake_impl = fake_impl - if final_fake_impl is None: + if final_fake_impl is None and not (mutates_args and returns_none): final_fake_impl = infer_fake_impl(kernel_func, helion_settings) logger.debug( "Auto-generated fake_impl for Helion kernel '%s'", diff --git a/aphrodite/lora/ops/triton_ops/fused_moe_lora_fp8_op.py b/aphrodite/lora/ops/triton_ops/fused_moe_lora_fp8_op.py index 3ef40bd52c..3d403f68d5 100644 --- a/aphrodite/lora/ops/triton_ops/fused_moe_lora_fp8_op.py +++ b/aphrodite/lora/ops/triton_ops/fused_moe_lora_fp8_op.py @@ -810,156 +810,23 @@ def _fused_moe_lora_fp8( ) -def _fused_moe_lora_fp8_fake( - output: torch.Tensor, - qcurr_hidden_states: torch.Tensor, - lora_a_stacked: list[torch.Tensor], - lora_b_stacked: list[torch.Tensor], - topk_weights: torch.Tensor, - sorted_token_ids: torch.Tensor | None, - expert_ids: torch.Tensor, - num_tokens_post_padded: torch.Tensor | None, - token_lora_mapping: torch.Tensor, - max_lora_rank: int, - top_k_num: int, - lora_ids: torch.Tensor, - num_active_loras: int, - adapter_enabled: torch.Tensor, - shrink_block_size_m: int, - shrink_block_size_n: int, - shrink_block_size_k: int, - shrink_group_size_m: int, - shrink_num_warps: int, - shrink_num_stages: int, - shrink_split_k: int, - expand_block_size_m: int, - expand_block_size_n: int, - expand_block_size_k: int, - expand_group_size_m: int, - expand_num_warps: int, - expand_num_stages: int, - expand_split_k: int, - lora_a_scale_stacked: list[torch.Tensor], - lora_b_scale_stacked: list[torch.Tensor], - mul_routed_weight: bool = False, - fully_sharded: bool = False, - offset: int = 0, - shrink_act_scale: torch.Tensor | None = None, - expand_act_scale: torch.Tensor | None = None, - use_fp8_w8a8: bool = False, - use_int8_w8a8: bool = False, - use_int8_w8a16: bool = False, - per_channel_quant: bool = False, - block_shape: List[int] | None = None, # noqa: UP006, UP007 -) -> None: - return - - -def _fused_moe_lora_shrink_fp8_fake( - a_intermediate_cache1: torch.Tensor, - qcurr_hidden_states: torch.Tensor, - lora_a_stacked: list[torch.Tensor], - topk_weights: torch.Tensor, - sorted_token_ids: torch.Tensor | None, - expert_ids: torch.Tensor, - num_tokens_post_padded: torch.Tensor | None, - token_lora_mapping: torch.Tensor, - top_k_num: int, - lora_ids: torch.Tensor, - adapter_enabled: torch.Tensor, - device: torch.device, - N: int, - M: int, - EM: int, - K: int, - num_tokens: int, - num_experts: int, - num_slices: int, - block_size_m: int, - block_size_n: int, - block_size_k: int, - group_size_m: int, - num_warps: int, - num_stages: int, - split_k: int, - num_active_loras: int, - lora_a_scale_stacked: list[torch.Tensor], - mul_routed_weight: bool = False, - use_gdc: bool = False, - act_scale: torch.Tensor | None = None, - use_fp8_w8a8: bool = False, - use_int8_w8a8: bool = False, - use_int8_w8a16: bool = False, - per_channel_quant: bool = False, - block_shape: List[int] | None = None, # noqa: UP006, UP007 -) -> None: - return - - -def _fused_moe_lora_expand_fp8_fake( - output: torch.Tensor, - a_intermediate_cache1: torch.Tensor, - lora_b_stacked: list[torch.Tensor], - topk_weights: torch.Tensor, - sorted_token_ids: torch.Tensor | None, - expert_ids: torch.Tensor, - num_tokens_post_padded: torch.Tensor | None, - token_lora_mapping: torch.Tensor, - top_k_num: int, - lora_ids: torch.Tensor, - adapter_enabled: torch.Tensor, - device: torch.device, - N: int, - M: int, - EM: int, - K: int, - num_tokens: int, - num_experts: int, - num_slices: int, - max_lora_rank: int, - w1_output_dim_size: int, - block_size_m: int, - block_size_n: int, - block_size_k: int, - group_size_m: int, - num_warps: int, - num_stages: int, - split_k: int, - num_active_loras: int, - act_scale: torch.Tensor, - lora_b_scale_stacked: list[torch.Tensor], - mul_routed_weight: bool = False, - offset: int = 0, - use_fp8_w8a8: bool = False, - use_int8_w8a8: bool = False, - use_int8_w8a16: bool = False, - per_channel_quant: bool = False, - block_shape: List[int] | None = None, # noqa: UP006, UP007 - use_gdc: bool = False, -) -> None: - return - - try: direct_register_custom_op( op_name="fused_moe_lora_fp8", op_func=_fused_moe_lora_fp8, mutates_args=["output"], - fake_impl=_fused_moe_lora_fp8_fake, ) direct_register_custom_op( op_name="fused_moe_lora_shrink_fp8", op_func=_fused_moe_lora_shrink_fp8, mutates_args=["a_intermediate_cache1"], - fake_impl=_fused_moe_lora_shrink_fp8_fake, ) direct_register_custom_op( op_name="fused_moe_lora_expand_fp8", op_func=_fused_moe_lora_expand_fp8, mutates_args=["output"], - fake_impl=_fused_moe_lora_expand_fp8_fake, ) fused_moe_lora_fp8 = torch.ops.aphrodite.fused_moe_lora_fp8 diff --git a/aphrodite/lora/ops/triton_ops/fused_moe_lora_op.py b/aphrodite/lora/ops/triton_ops/fused_moe_lora_op.py index 1f3d166da1..2e3c9d52ba 100644 --- a/aphrodite/lora/ops/triton_ops/fused_moe_lora_op.py +++ b/aphrodite/lora/ops/triton_ops/fused_moe_lora_op.py @@ -1579,136 +1579,23 @@ def _fused_moe_lora( ) -def _fused_moe_lora_fake( - output: torch.Tensor, - qcurr_hidden_states: torch.Tensor, - lora_a_stacked: list[torch.Tensor], - lora_b_stacked: list[torch.Tensor], - topk_weights: torch.Tensor, - sorted_token_ids: torch.Tensor | None, - expert_ids: torch.Tensor, - num_tokens_post_padded: torch.Tensor | None, - token_lora_mapping: torch.Tensor, - max_lora_rank: int, - top_k_num: int, - lora_ids: torch.Tensor, - num_active_loras: torch.Tensor, # CPU tensor [1], number of active LoRAs - adapter_enabled: torch.Tensor, - shrink_block_size_m: int, - shrink_block_size_n: int, - shrink_block_size_k: int, - shrink_group_size_m: int, - shrink_num_warps: int, - shrink_num_stages: int, - shrink_split_k: int, - expand_block_size_m: int, - expand_block_size_n: int, - expand_block_size_k: int, - expand_group_size_m: int, - expand_num_warps: int, - expand_num_stages: int, - expand_split_k: int, - mul_routed_weight: bool = False, - fully_sharded: bool = False, - offset: int = 0, - add_inputs: bool = True, -) -> None: - return - - -def _fused_moe_lora_shrink_fake( - a_intermediate_cache1: torch.Tensor, - qcurr_hidden_states: torch.Tensor, - lora_a_stacked: list[torch.Tensor], - topk_weights: torch.Tensor, - sorted_token_ids: torch.Tensor | None, - expert_ids: torch.Tensor, - num_tokens_post_padded: torch.Tensor | None, - token_lora_mapping: torch.Tensor, - top_k_num: int, - lora_ids: torch.Tensor, - adapter_enabled: torch.Tensor, - device: torch.device, - N: int, - M: int, - EM: int, - K: int, - num_tokens: int, - num_experts: int, - num_slices: int, - block_size_m: int, - block_size_n: int, - block_size_k: int, - group_size_m: int, - num_warps: int, - num_stages: int, - split_k: int, - num_active_loras: torch.Tensor, # CPU tensor [1], number of active LoRAs - mul_routed_weight: bool = False, - use_gdc: bool = False, - use_tma: bool = False, -) -> None: - return - - -def _fused_moe_lora_expand_fake( - output: torch.Tensor, - a_intermediate_cache1: torch.Tensor, - lora_b_stacked: list[torch.Tensor], - topk_weights: torch.Tensor, - sorted_token_ids: torch.Tensor | None, - expert_ids: torch.Tensor, - num_tokens_post_padded: torch.Tensor | None, - token_lora_mapping: torch.Tensor, - top_k_num: int, - lora_ids: torch.Tensor, - adapter_enabled: torch.Tensor, - device: torch.device, - N: int, - M: int, - EM: int, - K: int, - num_tokens: int, - num_experts: int, - num_slices: int, - max_lora_rank: int, - w1_output_dim_size: int, - block_size_m: int, - block_size_n: int, - block_size_k: int, - group_size_m: int, - num_warps: int, - num_stages: int, - split_k: int, - num_active_loras: torch.Tensor, # CPU tensor [1], number of active LoRAs - mul_routed_weight: bool = False, - offset: int = 0, - use_gdc: bool = False, - use_tma: bool = False, -) -> None: - return - - try: direct_register_custom_op( op_name="fused_moe_lora", op_func=_fused_moe_lora, mutates_args=["output"], - fake_impl=_fused_moe_lora_fake, ) direct_register_custom_op( op_name="fused_moe_lora_shrink", op_func=_fused_moe_lora_shrink, mutates_args=["a_intermediate_cache1"], - fake_impl=_fused_moe_lora_shrink_fake, ) direct_register_custom_op( op_name="fused_moe_lora_expand", op_func=_fused_moe_lora_expand, mutates_args=["output"], - fake_impl=_fused_moe_lora_expand_fake, ) fused_moe_lora = torch.ops.aphrodite.fused_moe_lora diff --git a/aphrodite/lora/ops/triton_ops/lora_expand_fp8_op.py b/aphrodite/lora/ops/triton_ops/lora_expand_fp8_op.py index d9a750baf2..a1d3e51276 100644 --- a/aphrodite/lora/ops/triton_ops/lora_expand_fp8_op.py +++ b/aphrodite/lora/ops/triton_ops/lora_expand_fp8_op.py @@ -363,35 +363,11 @@ def _lora_expand_fp8( return -def _lora_expand_fp8_fake( - inputs: torch.Tensor, - lora_b_weights: list[torch.Tensor], - output_tensor: torch.Tensor, - token_lora_mapping: torch.Tensor, - token_indices_sorted_by_lora_ids: torch.Tensor, - num_tokens_per_lora: torch.Tensor, - lora_token_start_loc: torch.Tensor, - lora_ids: torch.Tensor, - no_lora_flag_cpu: torch.Tensor, - num_active_loras: int, - b_scale: list[torch.Tensor], - a_scale: torch.Tensor | None = None, - offset_start: int = 0, - add_inputs: bool = False, - group_k: int = 0, - group_n: int = 0, - use_fp8_w8a8: bool = False, - per_channel_quant: bool = False, -) -> None: - return - - try: direct_register_custom_op( op_name="lora_expand_fp8", op_func=_lora_expand_fp8, mutates_args=["output_tensor"], - fake_impl=_lora_expand_fp8_fake, ) lora_expand_fp8 = torch.ops.aphrodite.lora_expand_fp8 diff --git a/aphrodite/lora/ops/triton_ops/lora_expand_op.py b/aphrodite/lora/ops/triton_ops/lora_expand_op.py index d282dadd6c..aecde6da53 100644 --- a/aphrodite/lora/ops/triton_ops/lora_expand_op.py +++ b/aphrodite/lora/ops/triton_ops/lora_expand_op.py @@ -282,29 +282,11 @@ def _lora_expand( return -def _lora_expand_fake( - inputs: torch.Tensor, - lora_b_weights: list[torch.Tensor], - output_tensor: torch.Tensor, - token_lora_mapping: torch.Tensor, - token_indices_sorted_by_lora_ids: torch.Tensor, - num_tokens_per_lora: torch.Tensor, - lora_token_start_loc: torch.Tensor, - lora_ids: torch.Tensor, - no_lora_flag_cpu: torch.Tensor, - num_active_loras: torch.Tensor, # CPU tensor [1], number of active LoRAs - offset_start: int = 0, - add_inputs: bool = False, -) -> None: - return - - try: direct_register_custom_op( op_name="lora_expand", op_func=_lora_expand, mutates_args=["output_tensor"], - fake_impl=_lora_expand_fake, ) lora_expand = torch.ops.aphrodite.lora_expand diff --git a/aphrodite/lora/ops/triton_ops/lora_shrink_fp8_op.py b/aphrodite/lora/ops/triton_ops/lora_shrink_fp8_op.py index fd9f9ce5e1..f62c6a3a8c 100644 --- a/aphrodite/lora/ops/triton_ops/lora_shrink_fp8_op.py +++ b/aphrodite/lora/ops/triton_ops/lora_shrink_fp8_op.py @@ -378,34 +378,11 @@ def _lora_shrink_fp8( return -def _lora_shrink_fp8_fake( - inputs: torch.Tensor, - lora_a_weights: list[torch.Tensor], - output_tensor: torch.Tensor, - token_lora_mapping: torch.Tensor, - token_indices_sorted_by_lora_ids: torch.Tensor, - num_tokens_per_lora: torch.Tensor, - lora_token_start_loc: torch.Tensor, - lora_ids: torch.Tensor, - no_lora_flag_cpu: torch.Tensor, - num_active_loras: int, - scaling: float, - b_scale: list[torch.Tensor], # LoRA weight scale per slice - a_scale: torch.Tensor | None = None, # Activation scale - per-token or block-wise - group_k: int = 0, # Block size for K in block-wise quantization (0 = tensor-wise) - group_n: int = 0, # Block size for N in block-wise quantization - use_fp8_w8a8: bool = False, - per_channel_quant: bool = False, -) -> None: - return - - try: direct_register_custom_op( op_name="lora_shrink_fp8", op_func=_lora_shrink_fp8, mutates_args=["output_tensor"], - fake_impl=_lora_shrink_fp8_fake, ) lora_shrink_fp8 = torch.ops.aphrodite.lora_shrink_fp8 diff --git a/aphrodite/lora/ops/triton_ops/lora_shrink_op.py b/aphrodite/lora/ops/triton_ops/lora_shrink_op.py index f622b6c403..bc6a81dabc 100644 --- a/aphrodite/lora/ops/triton_ops/lora_shrink_op.py +++ b/aphrodite/lora/ops/triton_ops/lora_shrink_op.py @@ -263,28 +263,11 @@ def _lora_shrink( return -def _lora_shrink_fake( - inputs: torch.Tensor, - lora_a_weights: list[torch.Tensor], - output_tensor: torch.Tensor, - token_lora_mapping: torch.Tensor, - token_indices_sorted_by_lora_ids: torch.Tensor, - num_tokens_per_lora: torch.Tensor, - lora_token_start_loc: torch.Tensor, - lora_ids: torch.Tensor, - no_lora_flag_cpu: torch.Tensor, - num_active_loras: torch.Tensor, # CPU tensor [1], number of active LoRAs - scaling: float, -) -> None: - return - - try: direct_register_custom_op( op_name="lora_shrink", op_func=_lora_shrink, mutates_args=["output_tensor"], - fake_impl=_lora_shrink_fake, ) lora_shrink = torch.ops.aphrodite.lora_shrink diff --git a/aphrodite/lora/ops/xpu_ops/lora_ops.py b/aphrodite/lora/ops/xpu_ops/lora_ops.py index f0752a4985..7e409187cd 100644 --- a/aphrodite/lora/ops/xpu_ops/lora_ops.py +++ b/aphrodite/lora/ops/xpu_ops/lora_ops.py @@ -83,57 +83,22 @@ def _bgmv_expand_slice_impl( ) -def _bgmv_shrink_fake( - inputs: torch.Tensor, - lora_a_weights: torch.Tensor, - output_tensor: torch.Tensor, - lora_indices_tensor: torch.Tensor, - scaling: float, -) -> None: - return None - - -def _bgmv_expand_fake( - inputs: torch.Tensor, - lora_b_weights: torch.Tensor, - output_tensor: torch.Tensor, - lora_indices_tensor: torch.Tensor, - add_inputs: bool, -) -> None: - return None - - -def _bgmv_expand_slice_fake( - inputs: torch.Tensor, - lora_b_weights: torch.Tensor, - output_tensor: torch.Tensor, - lora_indices_tensor: torch.Tensor, - slice_offset: int, - slice_size: int, - add_inputs: bool, -) -> None: - return None - - direct_register_custom_op( op_name="xpu_bgmv_shrink", op_func=_bgmv_shrink_impl, mutates_args=["output_tensor"], - fake_impl=_bgmv_shrink_fake, ) direct_register_custom_op( op_name="xpu_bgmv_expand", op_func=_bgmv_expand_impl, mutates_args=["output_tensor"], - fake_impl=_bgmv_expand_fake, ) direct_register_custom_op( op_name="xpu_bgmv_expand_slice", op_func=_bgmv_expand_slice_impl, mutates_args=["output_tensor"], - fake_impl=_bgmv_expand_slice_fake, ) diff --git a/aphrodite/model_executor/kernels/linear/scaled_mm/deep_gemm.py b/aphrodite/model_executor/kernels/linear/scaled_mm/deep_gemm.py index abed76ff4e..96d90863bd 100644 --- a/aphrodite/model_executor/kernels/linear/scaled_mm/deep_gemm.py +++ b/aphrodite/model_executor/kernels/linear/scaled_mm/deep_gemm.py @@ -130,20 +130,8 @@ def _fp8_gemm_nt_op( ) -def _fp8_gemm_nt_op_fake( - q_input: torch.Tensor, - input_scale: torch.Tensor, - weight: torch.Tensor, - weight_scale: torch.Tensor, - output: torch.Tensor, - use_deep_gemm_e8m0: bool, -) -> None: - return None - - direct_register_custom_op( "fp8_gemm_nt_op", _fp8_gemm_nt_op, mutates_args=["output"], - fake_impl=_fp8_gemm_nt_op_fake, ) diff --git a/aphrodite/model_executor/layers/attention/attention.py b/aphrodite/model_executor/layers/attention/attention.py index dbbf3abea3..a5845a5468 100644 --- a/aphrodite/model_executor/layers/attention/attention.py +++ b/aphrodite/model_executor/layers/attention/attention.py @@ -723,22 +723,8 @@ def unified_attention_with_output( ) -def unified_attention_with_output_fake( - query: torch.Tensor, - key: torch.Tensor, - value: torch.Tensor, - output: torch.Tensor, - layer_name: LayerNameType, - output_scale: torch.Tensor | None = None, - output_block_scale: torch.Tensor | None = None, - kv_cache_dummy_dep: torch.Tensor | None = None, -) -> None: - return - - direct_register_custom_op( op_name="unified_attention_with_output", op_func=unified_attention_with_output, mutates_args=["output", "output_block_scale"], - fake_impl=unified_attention_with_output_fake, ) diff --git a/aphrodite/model_executor/layers/attention/mla_attention.py b/aphrodite/model_executor/layers/attention/mla_attention.py index 636007d66c..cb18b510aa 100644 --- a/aphrodite/model_executor/layers/attention/mla_attention.py +++ b/aphrodite/model_executor/layers/attention/mla_attention.py @@ -1394,29 +1394,10 @@ def unified_mla_attention_with_output( ) -def unified_mla_attention_with_output_fake( - q: torch.Tensor, - kv_c_normed: torch.Tensor, - k_pe: torch.Tensor, - output: torch.Tensor, - layer_name: LayerNameType, - output_scale: torch.Tensor | None = None, - output_block_scale: torch.Tensor | None = None, - kv_cache_dummy_dep: torch.Tensor | None = None, - quant_group_size: int | None = None, - quant_scale_ue8m0: bool | None = None, - quant_col_major: bool | None = None, - quant_tma_aligned: bool | None = None, - q_dcp_replicated: torch.Tensor | None = None, -) -> None: - return - - direct_register_custom_op( op_name="unified_mla_attention_with_output", op_func=unified_mla_attention_with_output, mutates_args=["output", "output_block_scale"], - fake_impl=unified_mla_attention_with_output_fake, dispatch_key=current_platform.dispatch_key, tags=(torch.Tag.flexible_layout,), ) diff --git a/aphrodite/model_executor/layers/attention/static_sink_attention.py b/aphrodite/model_executor/layers/attention/static_sink_attention.py index 4dfcab17bc..850d96ffff 100644 --- a/aphrodite/model_executor/layers/attention/static_sink_attention.py +++ b/aphrodite/model_executor/layers/attention/static_sink_attention.py @@ -229,16 +229,8 @@ def maybe_populate_sink( self.populate_sink_kv(self_kv_cache) -def maybe_populate_sink_fake( - self_kv_cache: torch.Tensor, - layer_name: LayerNameType, -) -> None: - return - - direct_register_custom_op( op_name="maybe_populate_sink", op_func=maybe_populate_sink, mutates_args=["self_kv_cache"], - fake_impl=maybe_populate_sink_fake, ) diff --git a/aphrodite/model_executor/layers/hpc/rope_norm.py b/aphrodite/model_executor/layers/hpc/rope_norm.py index da330457fd..4e98faa19e 100644 --- a/aphrodite/model_executor/layers/hpc/rope_norm.py +++ b/aphrodite/model_executor/layers/hpc/rope_norm.py @@ -76,20 +76,10 @@ def hpc_rope_norm_forward( rope_norm._forward_impl(qkv, kv_cache, attn_metadata, attn_layer, output) -def hpc_rope_norm_forward_fake( - qkv: torch.Tensor, - output: torch.Tensor, - layer_name: str, -) -> None: - """Fake impl for torch.compile trace; output is a mutated arg.""" - return - - direct_register_custom_op( op_name="hpc_rope_norm_forward", op_func=hpc_rope_norm_forward, mutates_args=["output"], - fake_impl=hpc_rope_norm_forward_fake, ) diff --git a/aphrodite/model_executor/layers/mamba/gdn/olmo_gdn_linear_attn.py b/aphrodite/model_executor/layers/mamba/gdn/olmo_gdn_linear_attn.py index 807c4def1c..6c8d8b9c54 100644 --- a/aphrodite/model_executor/layers/mamba/gdn/olmo_gdn_linear_attn.py +++ b/aphrodite/model_executor/layers/mamba/gdn/olmo_gdn_linear_attn.py @@ -522,20 +522,10 @@ def olmo_hybrid_gdn_full_forward( ) -def olmo_hybrid_gdn_full_forward_fake( - hidden_states: torch.Tensor, - output: torch.Tensor, - layer_name: str, -) -> None: - """Fake implementation for torch.compile.""" - return - - direct_register_custom_op( op_name="olmo_hybrid_gdn_full_forward", op_func=olmo_hybrid_gdn_full_forward, mutates_args=["output"], - fake_impl=olmo_hybrid_gdn_full_forward_fake, ) diff --git a/aphrodite/model_executor/layers/mamba/gdn/qwen_gdn_linear_attn.py b/aphrodite/model_executor/layers/mamba/gdn/qwen_gdn_linear_attn.py index 9bb61955c6..3e197d83c4 100644 --- a/aphrodite/model_executor/layers/mamba/gdn/qwen_gdn_linear_attn.py +++ b/aphrodite/model_executor/layers/mamba/gdn/qwen_gdn_linear_attn.py @@ -1855,23 +1855,10 @@ def qwen_gdn_attention_core( ) -def gdn_attention_core_fake( - qkv_or_qkvz: torch.Tensor, - b_or_ba: torch.Tensor, - a_or_z_out: torch.Tensor, - core_attn_out: torch.Tensor, - layer_name: LayerNameType, - use_aiter: bool = False, -) -> None: - """Fake implementation for torch.compile.""" - return - - direct_register_custom_op( op_name="qwen_gdn_attention_core", op_func=qwen_gdn_attention_core, mutates_args=["a_or_z_out", "core_attn_out"], - fake_impl=gdn_attention_core_fake, ) @@ -1891,20 +1878,10 @@ def qwen_gdn_attention_core_fused_norm_packed( ) -def gdn_attention_core_fused_norm_packed_fake( - mixed_qkvz: torch.Tensor, - ba: torch.Tensor, - core_attn_out: torch.Tensor, - layer_name: LayerNameType, -) -> None: - return - - direct_register_custom_op( op_name="qwen_gdn_attention_core_fused_norm_packed", op_func=qwen_gdn_attention_core_fused_norm_packed, mutates_args=["core_attn_out"], - fake_impl=gdn_attention_core_fused_norm_packed_fake, ) diff --git a/aphrodite/model_executor/layers/mamba/linear/minimax_linear_attn.py b/aphrodite/model_executor/layers/mamba/linear/minimax_linear_attn.py index 913ec60369..0d6866a3c9 100644 --- a/aphrodite/model_executor/layers/mamba/linear/minimax_linear_attn.py +++ b/aphrodite/model_executor/layers/mamba/linear/minimax_linear_attn.py @@ -306,18 +306,8 @@ def linear_attention( self._forward(hidden_states=hidden_states, output=output, positions=positions) -def linear_attention_fake( - hidden_states: torch.Tensor, - output: torch.Tensor, - positions: torch.Tensor, - layer_name: str, -) -> None: - return - - direct_register_custom_op( op_name="linear_attention", op_func=linear_attention, mutates_args=["output"], - fake_impl=linear_attention_fake, ) diff --git a/aphrodite/model_executor/layers/mamba/mamba_mixer.py b/aphrodite/model_executor/layers/mamba/mamba_mixer.py index 2738119944..3de8c4c4d6 100644 --- a/aphrodite/model_executor/layers/mamba/mamba_mixer.py +++ b/aphrodite/model_executor/layers/mamba/mamba_mixer.py @@ -506,17 +506,8 @@ def mamba_mixer( self.forward_impl(hidden_states=hidden_states, output=output) -def mamba_mixer_fake( - hidden_states: torch.Tensor, - output: torch.Tensor, - layer_name: LayerNameType, -) -> None: - return - - direct_register_custom_op( op_name="mamba_mixer", op_func=mamba_mixer, mutates_args=["output"], - fake_impl=mamba_mixer_fake, ) diff --git a/aphrodite/model_executor/layers/mamba/mamba_mixer2.py b/aphrodite/model_executor/layers/mamba/mamba_mixer2.py index 477cb3dfa8..27f9f83b0a 100644 --- a/aphrodite/model_executor/layers/mamba/mamba_mixer2.py +++ b/aphrodite/model_executor/layers/mamba/mamba_mixer2.py @@ -1194,17 +1194,8 @@ def mamba_mixer2( self.conv_ssm_forward(projected_states=projected_states, output=output) -def mamba_mixer2_fake( - projected_states: torch.Tensor, - output: torch.Tensor, - layer_name: LayerNameType, -) -> None: - return - - direct_register_custom_op( op_name="mamba_mixer2", op_func=mamba_mixer2, mutates_args=["output"], - fake_impl=mamba_mixer2_fake, ) diff --git a/aphrodite/model_executor/layers/mamba/ops/cpu/gdn_attention.py b/aphrodite/model_executor/layers/mamba/ops/cpu/gdn_attention.py index e947884b92..900b206602 100644 --- a/aphrodite/model_executor/layers/mamba/ops/cpu/gdn_attention.py +++ b/aphrodite/model_executor/layers/mamba/ops/cpu/gdn_attention.py @@ -664,17 +664,6 @@ def _spec_aware_nonspec_subset( return attn_out.squeeze(0) -def cpu_gdn_attention_core_fake( - mixed_qkv: torch.Tensor, - b: torch.Tensor, - a: torch.Tensor, - core_attn_out: torch.Tensor, - layer_name: LayerNameType, -) -> None: - """Fake implementation for torch.compile.""" - return - - def register_cpu_gdn_attention_ops() -> None: global _CPU_GDN_ATTENTION_OPS_REGISTERED if _CPU_GDN_ATTENTION_OPS_REGISTERED: @@ -684,6 +673,5 @@ def register_cpu_gdn_attention_ops() -> None: op_name="cpu_gdn_attention_core", op_func=cpu_gdn_attention_core, mutates_args=["core_attn_out"], - fake_impl=cpu_gdn_attention_core_fake, ) _CPU_GDN_ATTENTION_OPS_REGISTERED = True diff --git a/aphrodite/model_executor/layers/mamba/short_conv.py b/aphrodite/model_executor/layers/mamba/short_conv.py index 390a7e2488..0928350c14 100644 --- a/aphrodite/model_executor/layers/mamba/short_conv.py +++ b/aphrodite/model_executor/layers/mamba/short_conv.py @@ -356,17 +356,8 @@ def short_conv( self.forward_native(hidden_states=hidden_states, output=output) -def short_conv_fake( - hidden_states: torch.Tensor, - output: torch.Tensor, - layer_name: str, -) -> None: - return - - direct_register_custom_op( op_name="short_conv", op_func=short_conv, mutates_args=["output"], - fake_impl=short_conv_fake, ) diff --git a/aphrodite/model_executor/layers/rotary_embedding/common.py b/aphrodite/model_executor/layers/rotary_embedding/common.py index 75baf6e984..3ded8cc3cd 100644 --- a/aphrodite/model_executor/layers/rotary_embedding/common.py +++ b/aphrodite/model_executor/layers/rotary_embedding/common.py @@ -96,23 +96,11 @@ def _flashinfer_rotary_embedding( ) -def _flashinfer_rotary_embedding_fake( - positions: torch.Tensor, - query: torch.Tensor, - key: torch.Tensor, - head_size: int, - cos_sin_cache: torch.Tensor, - is_neox: bool, -) -> None: - return - - # Register flashinfer rotary embedding custom op direct_register_custom_op( op_name="flashinfer_rotary_embedding", op_func=_flashinfer_rotary_embedding, mutates_args=["query", "key"], # These tensors are modified in-place - fake_impl=_flashinfer_rotary_embedding_fake, ) diff --git a/aphrodite/model_executor/models/bailing_moe_v3.py b/aphrodite/model_executor/models/bailing_moe_v3.py index d0a980975e..c362b2ec85 100644 --- a/aphrodite/model_executor/models/bailing_moe_v3.py +++ b/aphrodite/model_executor/models/bailing_moe_v3.py @@ -115,23 +115,10 @@ def bailing_v3_kda_attention( ) -def bailing_v3_kda_attention_fake( - q_proj_states: torch.Tensor, - k_proj_states: torch.Tensor, - v_proj_states: torch.Tensor, - g1: torch.Tensor, - beta: torch.Tensor, - core_attn_out: torch.Tensor, - layer_name: str, -) -> None: - return - - direct_register_custom_op( op_name="bailing_v3_kda_attention", op_func=bailing_v3_kda_attention, mutates_args=["core_attn_out"], - fake_impl=bailing_v3_kda_attention_fake, ) diff --git a/aphrodite/model_executor/offloader/prefetch_ops.py b/aphrodite/model_executor/offloader/prefetch_ops.py index c77c2dbc34..360d41e334 100644 --- a/aphrodite/model_executor/offloader/prefetch_ops.py +++ b/aphrodite/model_executor/offloader/prefetch_ops.py @@ -33,14 +33,6 @@ def _wait_prefetch_impl( get_offloader()._wait_for_layer(layer_idx) -def _wait_prefetch_fake( - input_tensor: torch.Tensor, - layer_idx: int, -) -> None: - """Fake implementation for torch.compile tracing.""" - return - - # --- start_prefetch op --- @@ -61,14 +53,6 @@ def _start_prefetch_impl( get_offloader()._start_prefetch(layer_idx) -def _start_prefetch_fake( - output_tensor: torch.Tensor, - layer_idx: int, -) -> None: - """Fake implementation for torch.compile tracing.""" - return - - def register_prefetch_offloader_ops() -> None: """Register custom ops for prefetch offloader. @@ -79,14 +63,12 @@ def register_prefetch_offloader_ops() -> None: op_name="wait_prefetch", op_func=_wait_prefetch_impl, mutates_args=["input_tensor"], - fake_impl=_wait_prefetch_fake, ) direct_register_custom_op( op_name="start_prefetch", op_func=_start_prefetch_impl, mutates_args=["output_tensor"], - fake_impl=_start_prefetch_fake, ) diff --git a/aphrodite/models/deepseek_v4/xpu/model.py b/aphrodite/models/deepseek_v4/xpu/model.py index af8f93d9a6..041cfc8f5e 100644 --- a/aphrodite/models/deepseek_v4/xpu/model.py +++ b/aphrodite/models/deepseek_v4/xpu/model.py @@ -560,23 +560,10 @@ def _deepseek_v4_mega_moe_experts_op( ) -def _deepseek_v4_mega_moe_experts_op_fake( - hidden_states: torch.Tensor, - topk_weights: torch.Tensor, - topk_ids: torch.Tensor, - out: torch.Tensor, - layer_name: str, - activation_clamp: float | None, - fast_math: bool, -) -> None: - return None - - direct_register_custom_op( op_name="deepseek_v4_mega_moe_experts", op_func=_deepseek_v4_mega_moe_experts_op, mutates_args=["out"], - fake_impl=_deepseek_v4_mega_moe_experts_op_fake, ) diff --git a/aphrodite/models/qwen4_exp/amd/ple_layer.py b/aphrodite/models/qwen4_exp/amd/ple_layer.py index ef94bc2b73..911b1379a5 100644 --- a/aphrodite/models/qwen4_exp/amd/ple_layer.py +++ b/aphrodite/models/qwen4_exp/amd/ple_layer.py @@ -1012,14 +1012,6 @@ def qwen4_exp_amd_ple_ngram_embedding( output.copy_(result) -def qwen4_exp_amd_ple_ngram_embedding_fake( - ngram_ids: torch.Tensor, - output: torch.Tensor, - layer_name: str, -) -> None: - return - - def qwen4_exp_ple_short_conv( inputs: torch.Tensor, output: torch.Tensor, @@ -1030,19 +1022,10 @@ def qwen4_exp_ple_short_conv( output[: result.shape[0]].copy_(result) -def qwen4_exp_ple_short_conv_fake( - inputs: torch.Tensor, - output: torch.Tensor, - layer_name: str, -) -> None: - return - - direct_register_custom_op( op_name="qwen4_exp_amd_ple_ngram_embedding", op_func=qwen4_exp_amd_ple_ngram_embedding, mutates_args=["output"], - fake_impl=qwen4_exp_amd_ple_ngram_embedding_fake, ) @@ -1050,7 +1033,6 @@ def qwen4_exp_ple_short_conv_fake( op_name="qwen4_exp_ple_short_conv", op_func=qwen4_exp_ple_short_conv, mutates_args=["output"], - fake_impl=qwen4_exp_ple_short_conv_fake, ) diff --git a/aphrodite/models/qwen4_exp/amd/qsa.py b/aphrodite/models/qwen4_exp/amd/qsa.py index 5532b98f62..b5ce7474f7 100644 --- a/aphrodite/models/qwen4_exp/amd/qsa.py +++ b/aphrodite/models/qwen4_exp/amd/qsa.py @@ -434,23 +434,10 @@ def qwen4_exp_qsa_with_output( ) -def qwen4_exp_qsa_with_output_fake( - hidden_states: torch.Tensor, - positions: torch.Tensor, - query: torch.Tensor, - key: torch.Tensor, - value: torch.Tensor, - output: torch.Tensor, - layer_name: LayerNameType, -) -> None: - del hidden_states, positions, query, key, value, output, layer_name - - direct_register_custom_op( op_name="qwen4_exp_qsa_with_output", op_func=qwen4_exp_qsa_with_output, mutates_args=["output"], - fake_impl=qwen4_exp_qsa_with_output_fake, ) diff --git a/aphrodite/models/qwen4_exp/nvidia/ple_layer.py b/aphrodite/models/qwen4_exp/nvidia/ple_layer.py index 01681d709d..d2837d2d76 100644 --- a/aphrodite/models/qwen4_exp/nvidia/ple_layer.py +++ b/aphrodite/models/qwen4_exp/nvidia/ple_layer.py @@ -833,16 +833,6 @@ def qwen4_exp_compute_ple_ngram_ids( ) -def qwen4_exp_compute_ple_ngram_ids_fake( - input_ids: torch.Tensor, - query_start_loc: torch.Tensor, - ngram_context: torch.Tensor, - output: torch.Tensor, - layer_name: str, -) -> None: - return - - def qwen4_exp_ple_short_conv( inputs: torch.Tensor, residual_output: torch.Tensor, @@ -852,19 +842,10 @@ def qwen4_exp_ple_short_conv( layer._short_conv(inputs, residual_output) -def qwen4_exp_ple_short_conv_fake( - inputs: torch.Tensor, - residual_output: torch.Tensor, - layer_name: str, -) -> None: - return - - direct_register_custom_op( op_name="qwen4_exp_compute_ple_ngram_ids", op_func=qwen4_exp_compute_ple_ngram_ids, mutates_args=["output"], - fake_impl=qwen4_exp_compute_ple_ngram_ids_fake, ) @@ -872,7 +853,6 @@ def qwen4_exp_ple_short_conv_fake( op_name="qwen4_exp_ple_short_conv", op_func=qwen4_exp_ple_short_conv, mutates_args=["residual_output"], - fake_impl=qwen4_exp_ple_short_conv_fake, ) diff --git a/aphrodite/models/qwen4_exp/nvidia/qsa.py b/aphrodite/models/qwen4_exp/nvidia/qsa.py index 1d130ddb39..bbf9119e7b 100644 --- a/aphrodite/models/qwen4_exp/nvidia/qsa.py +++ b/aphrodite/models/qwen4_exp/nvidia/qsa.py @@ -454,23 +454,10 @@ def qwen4_exp_qsa_with_output( ) -def qwen4_exp_qsa_with_output_fake( - hidden_states: torch.Tensor, - positions: torch.Tensor, - query: torch.Tensor, - key: torch.Tensor, - value: torch.Tensor, - output: torch.Tensor, - layer_name: LayerNameType, -) -> None: - del hidden_states, positions, query, key, value, output, layer_name - - direct_register_custom_op( op_name="qwen4_exp_qsa_with_output", op_func=qwen4_exp_qsa_with_output, mutates_args=["output"], - fake_impl=qwen4_exp_qsa_with_output_fake, ) diff --git a/aphrodite/utils/flashinfer.py b/aphrodite/utils/flashinfer.py index d62f529af5..89d7a5766f 100644 --- a/aphrodite/utils/flashinfer.py +++ b/aphrodite/utils/flashinfer.py @@ -660,19 +660,11 @@ def _flashinfer_concat_mla_k( concat_mla_k(k, k_nope, k_pe) - def _flashinfer_concat_mla_k_fake( - k: torch.Tensor, - k_nope: torch.Tensor, - k_pe: torch.Tensor, - ) -> None: - return - # Register flashinfer concat_mla_k custom op direct_register_custom_op( op_name="flashinfer_concat_mla_k", op_func=_flashinfer_concat_mla_k, mutates_args=["k"], # k tensor is modified in-place - fake_impl=_flashinfer_concat_mla_k_fake, ) @torch.library.custom_op( From 9f27610544a3555286dbe7c46c787729dc5dfb16 Mon Sep 17 00:00:00 2001 From: AlpinDale Date: Mon, 7 Sep 2026 19:14:54 +0000 Subject: [PATCH 03/36] [sync] add 2/3/5/6/7 CUDA support in AutoRound format (#52890) Upstream-vLLM: 294fbb4f595f61743124fc2e56a99e128cbe10da Co-authored-by: Wenhua Cheng --- .sync/vllm-sha | 2 +- .../layers/quantization/inc/inc.py | 2 +- .../inc/schemes/inc_wna16_scheme.py | 58 ++++++++++++ tests/quantization/test_auto_round.py | 89 +++++++++++++++++++ 4 files changed, 149 insertions(+), 2 deletions(-) diff --git a/.sync/vllm-sha b/.sync/vllm-sha index 163c52426a..356bfc840e 100644 --- a/.sync/vllm-sha +++ b/.sync/vllm-sha @@ -1 +1 @@ -199cb9b964822e59ab9b58d88e7be31eb419a2ae +294fbb4f595f61743124fc2e56a99e128cbe10da diff --git a/aphrodite/model_executor/layers/quantization/inc/inc.py b/aphrodite/model_executor/layers/quantization/inc/inc.py index 0375b212b7..bf93cc1b49 100644 --- a/aphrodite/model_executor/layers/quantization/inc/inc.py +++ b/aphrodite/model_executor/layers/quantization/inc/inc.py @@ -36,7 +36,7 @@ class INCConfig(QuantizationConfig): DEFAULT_INT_PACKING_FORMAT = "auto_round:auto_gptq" - SUPPORTED_BITS = {2, 3, 4, 8} + SUPPORTED_BITS = {2, 3, 4, 5, 6, 7, 8} SUPPORTED_DTYPES = {"int", "mx_fp", "fp"} SUPPORTED_FORMATS = { "auto_round:auto_gptq", diff --git a/aphrodite/model_executor/layers/quantization/inc/schemes/inc_wna16_scheme.py b/aphrodite/model_executor/layers/quantization/inc/schemes/inc_wna16_scheme.py index 97e6d24910..4c837aada9 100644 --- a/aphrodite/model_executor/layers/quantization/inc/schemes/inc_wna16_scheme.py +++ b/aphrodite/model_executor/layers/quantization/inc/schemes/inc_wna16_scheme.py @@ -21,6 +21,11 @@ logger = init_logger(__name__) XPU_WNA16_SUPPORTED_BITS = {2, 4} +# On CUDA the Marlin/GPTQ/AWQ kernels only cover 4/8-bit. The remaining +# widths (2/3/5/6/7) have no dedicated CUDA kernel, so they are dispatched to +# the humming kernel instead. This mirrors compressed-tensors WNA16, whose +# 5/6/7-bit biased scalar types also fall back to humming at kernel selection. +CUDA_HUMMING_SUPPORTED_BITS = {2, 3, 5, 6, 7} # Backends selectable through APHRODITE_XPU_INC_WNA16_BACKEND that are served by the # oneDNN int4 GEMMs rather than by ARK. These only cover int4. @@ -129,6 +134,12 @@ def get_linear_method( return INCLinearMethod(INCWNA16LinearScheme(layer_config)) raise NotImplementedError(f"INC on CPU: unsupported config {layer_config}") + # CUDA low-bit (2/3/5/6/7): no Marlin/GPTQ/AWQ kernel, route to humming + # so a single model can mix 4/8-bit (Marlin) and 2/3/5/6/7-bit (humming) + # layers. + if current_platform.is_cuda() and layer_config.bits in CUDA_HUMMING_SUPPORTED_BITS: + return _build_humming_linear_method(layer_config) + from .inc_wna16_linear import INCWNA16LinearScheme return INCLinearMethod(INCWNA16LinearScheme(layer_config)) @@ -148,6 +159,9 @@ def get_moe_method( ) return UnquantizedFusedMoEMethod(layer.moe_config) + # CUDA low-bit (2/3/5/6/7): route to the humming MoE kernel (see above). + if current_platform.is_cuda() and layer_config.bits in CUDA_HUMMING_SUPPORTED_BITS: + return _build_humming_moe_method(layer, layer_config) if layer_config.is_gptq: return _resolve_gptq_moe(layer, layer_config) if layer_config.is_awq: @@ -241,3 +255,47 @@ def _resolve_awq_moe(layer: "torch.nn.Module", layer_config: "INCLayerConfig"): } ) return MoeWNA16Method(moe_config, layer.moe_config) + + +def _humming_weight_config(layer_config: "INCLayerConfig") -> dict: + """Build the humming weight-schema config for a WNA16 int checkpoint.""" + if layer_config.is_gptq: + return { + "quant_method": "gptq", + "bits": layer_config.bits, + "group_size": layer_config.group_size, + "desc_act": False, + "sym": layer_config.sym, + } + if layer_config.is_awq: + return { + "quant_method": "awq", + "bits": layer_config.bits, + "group_size": layer_config.group_size, + "zero_point": not layer_config.sym, + } + raise NotImplementedError( + f"INC humming dispatch only supports gptq/awq packed int checkpoints, but found {layer_config}." + ) + + +def _build_humming_quant_config(layer_config: "INCLayerConfig"): + from aphrodite.model_executor.layers.quantization.humming import ( + HummingLayerQuantizationConfig, + ) + from aphrodite.utils.humming import BaseWeightSchema + + weight_schema = BaseWeightSchema.from_config(_humming_weight_config(layer_config)) + return HummingLayerQuantizationConfig(weight_schema=weight_schema) + + +def _build_humming_linear_method(layer_config: "INCLayerConfig"): + from aphrodite.model_executor.layers.quantization.humming import HummingLinearMethod + + return HummingLinearMethod(_build_humming_quant_config(layer_config)) + + +def _build_humming_moe_method(layer: "torch.nn.Module", layer_config: "INCLayerConfig"): + from aphrodite.model_executor.layers.quantization.humming import HummingMoEMethod + + return HummingMoEMethod(_build_humming_quant_config(layer_config), layer.moe_config) diff --git a/tests/quantization/test_auto_round.py b/tests/quantization/test_auto_round.py index 34af75c41b..cd9397bc5f 100644 --- a/tests/quantization/test_auto_round.py +++ b/tests/quantization/test_auto_round.py @@ -46,6 +46,7 @@ INCXPULinearMethod, ) from aphrodite.model_executor.layers.quantization.inc.schemes.inc_wna16_scheme import ( + _humming_weight_config, _resolve_awq_moe, _resolve_gptq_moe, ) @@ -670,6 +671,94 @@ def test_inc_resolve_scheme_selects_wna16() -> None: assert isinstance(scheme, INCWna16Scheme) +@pytest.mark.parametrize("bits", [2, 3, 5, 6, 7]) +@pytest.mark.parametrize("sym", [True, False]) +@pytest.mark.parametrize("group_size", [-1, 128]) +@pytest.mark.parametrize("packing_format", ["auto_round:auto_gptq", "auto_round:auto_awq"]) +def test_wna16_humming_weight_config(bits, sym, group_size, packing_format) -> None: + config = make_config(weight_bits=bits, sym=sym, group_size=group_size, packing_format=packing_format) + layer_config = config.config_parser.resolve(object(), "model.layers.0.mlp.down_proj") + weight_config = _humming_weight_config(layer_config) + + expected = {"bits": bits, "group_size": group_size} + if packing_format == "auto_round:auto_gptq": + expected.update(quant_method="gptq", desc_act=False, sym=sym) + else: + expected.update(quant_method="awq", zero_point=not sym) + assert weight_config == expected + + +@pytest.mark.parametrize("bits", [2, 3, 5, 6, 7]) +@pytest.mark.parametrize("packing_format", ["auto_round:auto_gptq", "auto_round:auto_awq"]) +def test_wna16_cuda_low_bit_linear_routes_to_humming(monkeypatch, bits, packing_format) -> None: + expected_method = object() + captured = {} + + monkeypatch.setattr(current_platform, "is_cuda", lambda: True) + monkeypatch.setattr(current_platform, "is_xpu", lambda: False) + monkeypatch.setattr(current_platform, "is_cpu", lambda: False) + monkeypatch.setattr( + "aphrodite.model_executor.layers.quantization.inc.schemes.inc_wna16_scheme._build_humming_linear_method", + lambda layer_config: captured.update({"layer_config": layer_config}) or expected_method, + ) + + layer_config = make_layer_config(bits=bits, packing_format=packing_format) + method = INCWna16Scheme().get_linear_method(make_config(), object(), "model.layers.0.mlp.down_proj", layer_config) + + assert method is expected_method + assert captured["layer_config"] is layer_config + + +@pytest.mark.parametrize("bits", [2, 3, 5, 6, 7]) +def test_wna16_cuda_low_bit_moe_routes_to_humming(monkeypatch, bits) -> None: + expected_method = object() + captured = {} + + class DummyMoeConfig: + pass + + monkeypatch.setattr(current_platform, "is_cuda", lambda: True) + monkeypatch.setattr(current_platform, "is_cpu", lambda: False) + monkeypatch.setattr( + "aphrodite.model_executor.layers.quantization.inc.schemes.inc_wna16_scheme._build_humming_moe_method", + lambda layer, layer_config: captured.update({"layer": layer, "layer_config": layer_config}) or expected_method, + ) + + layer = object.__new__(RoutedExperts) + layer.moe_config = DummyMoeConfig() + layer_config = make_layer_config(bits=bits) + method = INCWna16Scheme().get_moe_method(make_config(), layer, "model.layers.0.mlp", layer_config) + + assert method is expected_method + assert captured["layer"] is layer + assert captured["layer_config"] is layer_config + + +@pytest.mark.parametrize("bits", [4, 8]) +def test_wna16_cuda_high_bit_skips_humming(monkeypatch, bits) -> None: + """4/8-bit int stays on the Marlin/GPTQ/AWQ path even on CUDA so a single + model can mix high-bit Marlin and low-bit humming layers.""" + called = {"humming": False} + + monkeypatch.setattr(current_platform, "is_cuda", lambda: True) + monkeypatch.setattr(current_platform, "is_xpu", lambda: False) + monkeypatch.setattr(current_platform, "is_cpu", lambda: False) + monkeypatch.setattr( + "aphrodite.model_executor.layers.quantization.inc.schemes.inc_wna16_scheme._build_humming_linear_method", + lambda layer_config: called.update({"humming": True}), + ) + + method = INCWna16Scheme().get_linear_method( + make_config(), + object(), + "model.layers.0.mlp.down_proj", + make_layer_config(bits=bits), + ) + + assert called["humming"] is False + assert isinstance(method, INCLinearMethod) + + def test_qwen3_1p7b_mxfp4_autoround_uses_mxfp4_linear_scheme( monkeypatch, ) -> None: From 0685fe5c344b41385d9b4ef9f280c466487b47c7 Mon Sep 17 00:00:00 2001 From: AlpinDale Date: Mon, 7 Sep 2026 19:21:04 +0000 Subject: [PATCH 04/36] [sync] [ROCm][Perf][DeepSeek V4] Fuse native FP8 shared expert with MXFP4 routed experts (#53161) Upstream-vLLM: de69e821b7c834bcf1ec96325ec8c410ac183680 Co-authored-by: Fangzhou Ai <31551580+Fangzhou-Ai@users.noreply.github.com> --- .sync/vllm-sha | 2 +- aphrodite/_aiter_ops.py | 101 +++++ .../fused_moe/experts/rocm_aiter_moe.py | 11 + aphrodite/models/deepseek_v4/amd/model.py | 349 +++++++++++++++++- .../layers/test_fused_shared_expert.py | 321 +++++++++++++++- 5 files changed, 777 insertions(+), 7 deletions(-) diff --git a/.sync/vllm-sha b/.sync/vllm-sha index 356bfc840e..7a08e801e9 100644 --- a/.sync/vllm-sha +++ b/.sync/vllm-sha @@ -1 +1 @@ -294fbb4f595f61743124fc2e56a99e128cbe10da +de69e821b7c834bcf1ec96325ec8c410ac183680 diff --git a/aphrodite/_aiter_ops.py b/aphrodite/_aiter_ops.py index e57f5ae586..3fd69c7c7e 100644 --- a/aphrodite/_aiter_ops.py +++ b/aphrodite/_aiter_ops.py @@ -202,6 +202,25 @@ def wrapper(*args, **kwargs): return wrapper +def _validate_rocm_aiter_fused_moe_shared_expert_args( + shared_w1: torch.Tensor | None, + shared_w2: torch.Tensor | None, + shared_w1_scale: torch.Tensor | None, + shared_w2_scale: torch.Tensor | None, + shared_expert_id: int, +) -> bool: + shared_tensors = (shared_w1, shared_w2, shared_w1_scale, shared_w2_scale) + has_shared_expert = any(tensor is not None for tensor in shared_tensors) + if has_shared_expert: + if not all(tensor is not None for tensor in shared_tensors): + raise ValueError("Heterogeneous fused MoE requires both shared weights and scales.") + if shared_expert_id < 0: + raise ValueError("Heterogeneous fused MoE requires a non-negative shared expert ID.") + elif shared_expert_id >= 0: + raise ValueError("A non-negative shared expert ID requires shared weights and scales.") + return has_shared_expert + + def _rocm_aiter_fused_moe_impl( hidden_states: torch.Tensor, w1: torch.Tensor, @@ -227,7 +246,20 @@ def _rocm_aiter_fused_moe_impl( swiglu_limit: float = 0.0, beta: float | None = None, linear_beta: float | None = None, + shared_w1: torch.Tensor | None = None, + shared_w2: torch.Tensor | None = None, + shared_w1_scale: torch.Tensor | None = None, + shared_w2_scale: torch.Tensor | None = None, + shared_expert_id: int = -1, ) -> torch.Tensor: + has_shared_expert = _validate_rocm_aiter_fused_moe_shared_expert_args( + shared_w1, + shared_w2, + shared_w1_scale, + shared_w2_scale, + shared_expert_id, + ) + from aiter import ActivationType, QuantType from aiter.fused_moe import fused_moe @@ -241,6 +273,15 @@ def _rocm_aiter_fused_moe_impl( extra_kwargs["beta"] = beta extra_kwargs["linear_beta"] = linear_beta + if has_shared_expert: + extra_kwargs.update( + shared_w1=shared_w1, + shared_w2=shared_w2, + shared_w1_scale=shared_w1_scale, + shared_w2_scale=shared_w2_scale, + shared_expert_id=shared_expert_id, + ) + return fused_moe( hidden_states, w1, @@ -292,6 +333,11 @@ def _rocm_aiter_fused_moe_fake( swiglu_limit: float = 0.0, beta: float | None = None, linear_beta: float | None = None, + shared_w1: torch.Tensor | None = None, + shared_w2: torch.Tensor | None = None, + shared_w1_scale: torch.Tensor | None = None, + shared_w2_scale: torch.Tensor | None = None, + shared_expert_id: int = -1, ) -> torch.Tensor: if output_dtype is not None: return torch.empty_like(hidden_states, dtype=output_dtype) @@ -1952,6 +1998,44 @@ def fused_moe_supports_gate_mode(cls) -> bool: return "gate_mode" in inspect.signature(fused_moe).parameters + @staticmethod + def _probe_dsv4_i384_fhmoe_capability(num_tokens: int) -> bool: + """Probe AITER's CSV-backed DSV4 native-I384 FHMoE contract.""" + if type(num_tokens) is not int or num_tokens <= 0: + return False + + try: + import inspect + + import aiter.fhmoe as fhmoe + from aiter.fused_moe import fused_moe + + params = inspect.signature(fused_moe).parameters + supports_dsv4_i384_fhmoe = getattr(fhmoe, "supports_dsv4_i384_fhmoe", None) + if not callable(supports_dsv4_i384_fhmoe): + return False + supports_num_tokens = supports_dsv4_i384_fhmoe(num_tokens) + except Exception: + return False + + return supports_num_tokens is True and all( + name in params + for name in ( + "shared_w1", + "shared_w2", + "shared_w1_scale", + "shared_w2_scale", + "shared_expert_id", + ) + ) + + @classmethod + @if_aiter_supported + @functools.cache + def fused_moe_supports_heterogeneous_shared_expert(cls, num_tokens: int) -> bool: + """Whether AITER has DSV4 native-I384 configs through the given M.""" + return cls._probe_dsv4_i384_fhmoe_capability(num_tokens) + @staticmethod def register_ops_once() -> None: global _OPS_REGISTERED @@ -2410,6 +2494,11 @@ def fused_moe( swiglu_limit: float = 0.0, beta: float | None = None, linear_beta: float | None = None, + shared_w1: torch.Tensor | None = None, + shared_w2: torch.Tensor | None = None, + shared_w1_scale: torch.Tensor | None = None, + shared_w2_scale: torch.Tensor | None = None, + shared_expert_id: int = -1, ) -> torch.Tensor: return torch.ops.aphrodite.rocm_aiter_fused_moe( hidden_states, @@ -2436,6 +2525,11 @@ def fused_moe( swiglu_limit, beta, linear_beta, + shared_w1, + shared_w2, + shared_w1_scale, + shared_w2_scale, + shared_expert_id, ) @staticmethod @@ -2997,6 +3091,13 @@ def shuffle_scale_a16w4( return shuffle_scale_a16w4(tensor, num_experts, gate_up) + @staticmethod + def shuffle_scale(tensor: torch.Tensor) -> torch.Tensor: + """Shuffle a non-gate/up E8M0 scale tensor for AITER.""" + from aiter.ops.shuffle import shuffle_scale + + return shuffle_scale(tensor) + @staticmethod def shuffle_weights(*tensors: torch.Tensor, layout: tuple[int, int] = (16, 16)) -> tuple[torch.Tensor, ...]: """ diff --git a/aphrodite/model_executor/layers/fused_moe/experts/rocm_aiter_moe.py b/aphrodite/model_executor/layers/fused_moe/experts/rocm_aiter_moe.py index 1f0116c2bc..4783341eae 100644 --- a/aphrodite/model_executor/layers/fused_moe/experts/rocm_aiter_moe.py +++ b/aphrodite/model_executor/layers/fused_moe/experts/rocm_aiter_moe.py @@ -229,6 +229,11 @@ def rocm_aiter_fused_experts( num_local_tokens: torch.Tensor | None = None, output_dtype: torch.dtype | None = None, moe_sorting_dispatch_policy: int = 0, + shared_w1: torch.Tensor | None = None, + shared_w2: torch.Tensor | None = None, + shared_w1_scale: torch.Tensor | None = None, + shared_w2_scale: torch.Tensor | None = None, + shared_expert_id: int = -1, ) -> torch.Tensor: """ROCm AITER fused MoE expert computation.""" if quant_config is None: @@ -375,8 +380,14 @@ def rocm_aiter_fused_experts( bias1=quant_config.w1_bias if quant_config.use_mxfp4_w4a16 else None, bias2=quant_config.w2_bias if quant_config.use_mxfp4_w4a16 else None, moe_sorting_dispatch_policy=moe_sorting_dispatch_policy, + swiglu_limit=(0.0 if moe_config.swiglu_limit is None else float(moe_config.swiglu_limit)), beta=moe_config.activation_situ_beta, linear_beta=moe_config.activation_situ_linear_beta, + shared_w1=shared_w1, + shared_w2=shared_w2, + shared_w1_scale=shared_w1_scale, + shared_w2_scale=shared_w2_scale, + shared_expert_id=shared_expert_id, ) diff --git a/aphrodite/models/deepseek_v4/amd/model.py b/aphrodite/models/deepseek_v4/amd/model.py index 3509d9adef..4179f0c395 100644 --- a/aphrodite/models/deepseek_v4/amd/model.py +++ b/aphrodite/models/deepseek_v4/amd/model.py @@ -1,7 +1,9 @@ # SPDX-License-Identifier: Apache-2.0 # SPDX-FileCopyrightText: Copyright contributors to the vLLM project import typing +import weakref from collections.abc import Callable, Iterable +from dataclasses import replace from itertools import islice import regex as re @@ -20,9 +22,12 @@ from aphrodite.model_executor.layers.activation import SiluAndMul, SiluAndMulWithClamp from aphrodite.model_executor.layers.fused_moe import ( FusedMoEFactory, + FusedMoEQuantConfig, GateLinear, + RoutedExperts, fused_moe_make_expert_params_mapping, ) +from aphrodite.model_executor.layers.fused_moe.experts.rocm_aiter_moe import rocm_aiter_fused_experts from aphrodite.model_executor.layers.fused_moe.utils import ( is_model_fused_shared_expert_compatible, ) @@ -41,6 +46,7 @@ MHCPreOp, ) from aphrodite.model_executor.layers.quantization import QuantizationConfig +from aphrodite.model_executor.layers.quantization.mxfp4 import Mxfp4MoEMethod from aphrodite.model_executor.layers.quantization.utils.config_utils import ( is_shared_expert_quant_fse_compatible, ) @@ -61,6 +67,7 @@ ) from aphrodite.models.deepseek_v4.amd.rocm import DeepseekV4ROCMAiterMLAAttention from aphrodite.platforms import current_platform +from aphrodite.platforms.rocm import on_gfx950 from aphrodite.sequence import IntermediateTensors logger = init_logger(__name__) @@ -156,6 +163,316 @@ def forward(self, x): return x +def _use_heterogeneous_fhmoe(num_tokens: int) -> bool: + return num_tokens > 0 and (rocm_aiter_ops.fused_moe_supports_heterogeneous_shared_expert(num_tokens) is True) + + +def _validate_heterogeneous_routes( + topk_weights: torch.Tensor, + topk_ids: torch.Tensor, + experts_per_token: int, +) -> None: + if topk_weights.ndim != 2 or topk_ids.ndim != 2 or topk_weights.shape != topk_ids.shape: + raise ValueError("Heterogeneous fused MoE routes must have equal two-dimensional shapes.") + expected_columns = experts_per_token + 1 + if topk_weights.shape[1] != expected_columns: + raise ValueError( + f"Heterogeneous fused MoE routes must have exactly {expected_columns} columns, got {topk_weights.shape[1]}." + ) + + +def _heterogeneous_shared_expert_enabled(aphrodite_config: AphroditeConfig) -> bool: + config = aphrodite_config.model_config.hf_config + quant_config = aphrodite_config.quant_config + parallel_config = aphrodite_config.parallel_config + reasons: list[str] = [] + + if not rocm_aiter_ops.is_fusion_moe_shared_experts_enabled(): + return False + if not on_gfx950(): + reasons.append("the device is not gfx950") + if parallel_config.enable_expert_parallel: + reasons.append("expert parallelism is enabled") + if parallel_config.enable_eplb: + reasons.append("EPLB is enabled") + if parallel_config.tensor_parallel_size != 8: + reasons.append("tensor parallelism is not 8") + if parallel_config.data_parallel_size != 1: + reasons.append("data parallelism is enabled") + if parallel_config.prefill_context_parallel_size != 1: + reasons.append("prefill context parallelism is enabled") + if aphrodite_config.kernel_config.moe_backend != "aiter": + reasons.append("the MoE backend is not AITER") + if aphrodite_config.model_config.dtype != torch.bfloat16: + reasons.append("the model dtype is not BF16") + if not rocm_aiter_ops.fused_moe_supports_heterogeneous_shared_expert(1): + reasons.append("AITER lacks CSV-backed DSV4 I384 heterogeneous shared-expert support") + + expected_model = { + "n_routed_experts": 384, + "num_experts_per_tok": 6, + "n_shared_experts": 1, + "hidden_size": 7168, + "moe_intermediate_size": 3072, + "hidden_act": "silu", + "expert_dtype": "fp4", + "topk_method": "noaux_tc", + } + for name, expected in expected_model.items(): + if getattr(config, name, None) != expected: + reasons.append(f"{name} is not {expected!r}") + + if quant_config is None or quant_config.get_name() != "deepseek_v4_fp8": + reasons.append("the DeepSeek V4 FP8 quantization config is not active") + else: + if getattr(quant_config, "moe_quant_algo", "").upper() == "NVFP4": + reasons.append("routed experts use NVFP4") + if getattr(quant_config, "weight_block_size", None) != [128, 128]: + reasons.append("shared experts are not block-128 FP8") + if not getattr(quant_config, "is_checkpoint_fp8_serialized", False): + reasons.append("shared experts are not serialized as FP8") + if not getattr(quant_config, "is_scale_e8m0", False): + reasons.append("shared-expert scales are not E8M0") + if getattr(quant_config, "ignored_layers", None): + reasons.append("the quantization config has ignored layers") + + offload_config = getattr(aphrodite_config, "offload_config", None) + if offload_config is not None and ( + getattr(getattr(offload_config, "uva", None), "cpu_offload_gb", 0) > 0 + or getattr( + getattr(offload_config, "prefetch", None), + "offload_group_size", + 0, + ) + > 0 + ): + reasons.append("weight offloading is enabled") + + if reasons: + logger.debug_once( + "DeepSeek V4 heterogeneous shared-expert fusion is unavailable: %s. " + "Continuing with the existing shared-expert path.", + "; ".join(reasons), + ) + return False + + logger.info_once( + "Using AITER DeepSeek V4 heterogeneous MoE for CSV-covered MoE input " + "rows (M) per invocation (native intermediate width 384)." + ) + return True + + +@torch.no_grad() +def _prepare_native_fp8_shared_expert( + w13: torch.Tensor, + w2: torch.Tensor, + w13_scale: torch.Tensor, + w2_scale: torch.Tensor, + intermediate_size: int, +) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]: + """Expand native block-128 E8M0 scales to FHMoE's 1x32 layout.""" + if w13.dtype != torch.float8_e4m3fn or w2.dtype != torch.float8_e4m3fn: + raise ValueError("Heterogeneous shared-expert weights must be FP8 E4M3.") + hidden_size = w13.shape[1] + if w13.shape != (2 * intermediate_size, hidden_size): + raise ValueError(f"Unexpected shared W13 shape {tuple(w13.shape)}.") + if w2.shape != (hidden_size, intermediate_size): + raise ValueError(f"Unexpected shared W2 shape {tuple(w2.shape)}.") + + def scale_bytes(scale: torch.Tensor) -> torch.Tensor: + if scale.dtype == torch.float8_e8m0fnu: + return scale.contiguous().view(torch.uint8) + if scale.dtype == torch.uint8: + return scale.contiguous() + if scale.dtype == torch.float32: + bits = scale.contiguous().view(torch.int32) + exponent = (bits >> 23).to(torch.uint8) + if not torch.equal(bits, exponent.to(torch.int32) << 23): + raise ValueError("FP32 shared-expert scales are not exact E8M0.") + return exponent + raise ValueError(f"Unsupported shared-expert scale dtype {scale.dtype}.") + + w13_scale = scale_bytes(w13_scale) + w2_scale = scale_bytes(w2_scale) + expected_w13_scale = (2 * intermediate_size // 128, hidden_size // 128) + expected_w2_scale = (hidden_size // 128, intermediate_size // 128) + if tuple(w13_scale.shape) != expected_w13_scale: + raise ValueError(f"Expected shared W13 scale shape {expected_w13_scale}, got {tuple(w13_scale.shape)}.") + if tuple(w2_scale.shape) != expected_w2_scale: + raise ValueError(f"Expected shared W2 scale shape {expected_w2_scale}, got {tuple(w2_scale.shape)}.") + + w13_scale = w13_scale.repeat_interleave(128, dim=0).repeat_interleave(4, dim=1) + w2_scale = w2_scale.repeat_interleave(128, dim=0).repeat_interleave(4, dim=1) + + scale_cols = ((w2_scale.shape[1] + 7) // 8) * 8 + if scale_cols != w2_scale.shape[1]: + padded_w2_scale = torch.full( + (w2_scale.shape[0], scale_cols), + 0x7F, + dtype=torch.uint8, + device=w2_scale.device, + ) + padded_w2_scale[:, : w2_scale.shape[1]].copy_(w2_scale) + w2_scale = padded_w2_scale + + return ( + w13.unsqueeze(0), + w2.unsqueeze(0), + w13_scale.view(torch.float8_e8m0fnu), + w2_scale.view(torch.float8_e8m0fnu), + ) + + +class DeepseekV4HeterogeneousMxfp4MoEMethod(Mxfp4MoEMethod): + def process_weights_after_loading(self, layer: RoutedExperts) -> None: + super().process_weights_after_loading(layer) + assert isinstance(layer, DeepseekV4HeterogeneousSharedRoutedExperts) + layer.prepare_heterogeneous_shared_expert() + + +class DeepseekV4HeterogeneousSharedRoutedExperts(RoutedExperts): + def __init__( + self, + *args, + shared_expert: DeepseekV4MLP, + shared_expert_id: int, + **kwargs, + ): + self._shared_expert_ref = weakref.ref(shared_expert) + self.shared_expert_id = shared_expert_id + self._routed_quant_config: FusedMoEQuantConfig | None = None + super().__init__(*args, **kwargs) + if self.expert_map_manager.num_fused_shared_experts != 1: + raise ValueError("Heterogeneous fusion requires one appended expert.") + if self.shared_expert_id != self.global_num_experts: + raise ValueError("The shared expert ID must name the appended slot.") + self.register_buffer("shared_w1", None, persistent=False) + self.register_buffer("shared_w2", None, persistent=False) + self.register_buffer("shared_w1_scale", None, persistent=False) + self.register_buffer("shared_w2_scale", None, persistent=False) + + def _get_quant_method(self, prefix, quant_config, moe_config): + quant_method = super()._get_quant_method(prefix, quant_config, moe_config) + if not isinstance(quant_method, Mxfp4MoEMethod): + raise ValueError("Heterogeneous fusion requires MXFP4 routed experts.") + return DeepseekV4HeterogeneousMxfp4MoEMethod(moe_config) + + @torch.no_grad() + def prepare_heterogeneous_shared_expert(self) -> None: + shared_expert = self._shared_expert_ref() + if shared_expert is None: + raise RuntimeError("The native shared-expert module was released.") + + shared = _prepare_native_fp8_shared_expert( + shared_expert.gate_up_proj.weight, + shared_expert.down_proj.weight, + shared_expert.gate_up_proj.weight_scale_inv, + shared_expert.down_proj.weight_scale_inv, + self.moe_config.intermediate_size_per_partition, + ) + self.shared_w1 = rocm_aiter_ops.shuffle_weight_a16w4(shared[0], 16, True) + self.shared_w2 = rocm_aiter_ops.shuffle_weight_a16w4(shared[1], 16, False) + self.shared_w1_scale = rocm_aiter_ops.shuffle_scale_a16w4(shared[2], 1, True) + self.shared_w2_scale = rocm_aiter_ops.shuffle_scale(shared[3]) + + self._ensure_moe_quant_config_init() + quant_config = self.quant_method.moe_quant_config + if quant_config is None: + raise RuntimeError("The routed MXFP4 quantization config is unavailable.") + if quant_config.w1_scale is None or quant_config.w2_scale is None: + raise RuntimeError("The routed MXFP4 scales are unavailable.") + + def routed_scales(scale: torch.Tensor) -> torch.Tensor: + scale = scale.reshape(self.local_num_experts, -1, scale.shape[-1])[: self.shared_expert_id] + return scale.reshape(-1, scale.shape[-1]).contiguous() + + self._routed_quant_config = replace( + quant_config, + _w1=replace( + quant_config._w1, + scale=routed_scales(quant_config.w1_scale), + ), + _w2=replace( + quant_config._w2, + scale=routed_scales(quant_config.w2_scale), + ), + ) + + def forward_modular( + self, + x: torch.Tensor, + topk_weights: torch.Tensor, + topk_ids: torch.Tensor, + shared_experts=None, + shared_experts_input: torch.Tensor | None = None, + ) -> torch.Tensor: + if shared_experts is not None or shared_experts_input is not None: + raise ValueError("The heterogeneous expert owns the shared path.") + _validate_heterogeneous_routes( + topk_weights, + topk_ids, + self.moe_config.experts_per_token, + ) + if any( + tensor is None + for tensor in ( + self.shared_w1, + self.shared_w2, + self.shared_w1_scale, + self.shared_w2_scale, + ) + ): + raise RuntimeError("Heterogeneous shared weights were not prepared.") + + if _use_heterogeneous_fhmoe(x.shape[0]): + self._ensure_moe_quant_config_init() + quant_config = self.quant_method.moe_quant_config + if quant_config is None: + raise RuntimeError("The routed MXFP4 config is unavailable.") + return rocm_aiter_fused_experts( + hidden_states=x, + w1=self.w13_weight, + w2=self.w2_weight, + topk_weights=topk_weights, + topk_ids=topk_ids, + moe_config=self.moe_config, + activation=self.activation, + expert_mask=self.expert_map, + quant_config=quant_config, + output_dtype=x.dtype, + moe_sorting_dispatch_policy=(rocm_aiter_ops.get_moe_dispatch_policy()), + shared_w1=self.shared_w1, + shared_w2=self.shared_w2, + shared_w1_scale=self.shared_w1_scale, + shared_w2_scale=self.shared_w2_scale, + shared_expert_id=self.shared_expert_id, + ) + + shared_expert = self._shared_expert_ref() + if shared_expert is None or self._routed_quant_config is None: + raise RuntimeError("The separate shared-expert fallback is unavailable.") + shared_out = shared_expert(x) + routed_w1 = self.w13_weight[: self.shared_expert_id] + routed_w2 = self.w2_weight[: self.shared_expert_id] + routed_w1.is_shuffled = True + routed_w2.is_shuffled = True + routed = rocm_aiter_fused_experts( + hidden_states=x, + w1=routed_w1, + w2=routed_w2, + topk_weights=topk_weights[:, :-1].contiguous(), + topk_ids=topk_ids[:, :-1].contiguous(), + moe_config=self.moe_config, + activation=self.activation, + expert_mask=self.expert_map, + quant_config=self._routed_quant_config, + output_dtype=x.dtype, + moe_sorting_dispatch_policy=rocm_aiter_ops.get_moe_dispatch_policy(), + ) + return routed + shared_out + + def _shared_experts_are_fp4(config, layer_idx: int | None = None) -> bool: """Whether the shared experts are MXFP4 and thus fusable. @@ -196,6 +513,7 @@ def __init__( self, aphrodite_config: AphroditeConfig, prefix: str = "", + fuse_heterogeneous_shared_expert: bool = False, ): super().__init__() @@ -250,7 +568,9 @@ def __init__( # TODO: Historically, only `APHRODITE_ROCM_USE_AITER_FUSION_SHARED_EXPERTS=1` # is checked to enable FSE for DeepSeek-v4, despite AITER not being used. # This should be cleaned up and use `resolve_layer_fused_shared_expert`. - fse_requested = _fuse_shared_experts_enabled(config) + self.fuse_heterogeneous_shared_expert = fuse_heterogeneous_shared_expert + fse_requested = _fuse_shared_experts_enabled(config) and not self.fuse_heterogeneous_shared_expert + fse_compatible = False if fse_requested: fse_compatible, fse_reason = is_shared_expert_quant_fse_compatible( quant_config, @@ -286,10 +606,11 @@ def __init__( self.experts_start_idx = self.tp_rank * self.n_local_experts self.experts_end_idx = self.experts_start_idx + self.n_local_experts + fuse_shared_into_routed = self.is_fused_shared_expert_enabled or self.fuse_heterogeneous_shared_expert self.experts = FusedMoEFactory( - shared_experts=self.shared_experts, - n_shared_experts=(config.n_shared_experts if self.is_fused_shared_expert_enabled else None), - fuse_shared_experts=self.is_fused_shared_expert_enabled, + shared_experts=(None if fuse_shared_into_routed else self.shared_experts), + n_shared_experts=(config.n_shared_experts if fuse_shared_into_routed else None), + fuse_shared_experts=fuse_shared_into_routed, gate=self.gate, num_experts=config.n_routed_experts, top_k=config.num_experts_per_tok, @@ -304,6 +625,17 @@ def __init__( hash_indices_table=self.gate.tid2eid, swiglu_limit=self.swiglu_limit, router_logits_dtype=torch.float32, + routed_experts_cls=( + DeepseekV4HeterogeneousSharedRoutedExperts if self.fuse_heterogeneous_shared_expert else None + ), + routed_experts_args=( + { + "shared_expert": self.shared_experts, + "shared_expert_id": config.n_routed_experts, + } + if self.fuse_heterogeneous_shared_expert + else None + ), ) def forward(self, hidden_states: torch.Tensor, input_ids: torch.Tensor | None = None) -> torch.Tensor: @@ -331,6 +663,7 @@ def __init__( prefix, topk_indices_buffer: torch.Tensor | None = None, aux_stream_list: list[torch.cuda.Stream] | None = None, + fuse_heterogeneous_shared_expert: bool = False, ): super().__init__() @@ -348,7 +681,11 @@ def __init__( topk_indices_buffer=topk_indices_buffer, aux_stream_list=aux_stream_list, ) - self.ffn = DeepseekV4MoE(aphrodite_config, prefix=f"{prefix}.ffn") + self.ffn = DeepseekV4MoE( + aphrodite_config, + prefix=f"{prefix}.ffn", + fuse_heterogeneous_shared_expert=fuse_heterogeneous_shared_expert, + ) self.attn_norm = RMSNorm(self.hidden_size, self.rms_norm_eps) self.ffn_norm = RMSNorm(self.hidden_size, self.rms_norm_eps) @@ -593,6 +930,7 @@ def __init__(self, *, aphrodite_config: AphroditeConfig, prefix: str = ""): else: self.embed_tokens = PPMissingLayer() + self.fuse_heterogeneous_shared_expert = _heterogeneous_shared_expert_enabled(aphrodite_config) self.start_layer, self.end_layer, self.layers = make_layers( config.num_hidden_layers, lambda prefix: DeepseekV4DecoderLayer( @@ -600,6 +938,7 @@ def __init__(self, *, aphrodite_config: AphroditeConfig, prefix: str = ""): prefix=prefix, topk_indices_buffer=self.topk_indices_buffer, aux_stream_list=aux_stream_list, + fuse_heterogeneous_shared_expert=(self.fuse_heterogeneous_shared_expert), ), prefix=f"{prefix}.layers", ) diff --git a/tests/model_executor/layers/test_fused_shared_expert.py b/tests/model_executor/layers/test_fused_shared_expert.py index 5ec80b420d..dff27ade8b 100644 --- a/tests/model_executor/layers/test_fused_shared_expert.py +++ b/tests/model_executor/layers/test_fused_shared_expert.py @@ -2,8 +2,10 @@ # SPDX-FileCopyrightText: Copyright contributors to the vLLM project import importlib +import inspect +import sys from copy import deepcopy -from types import SimpleNamespace +from types import ModuleType, SimpleNamespace from typing import Any, cast from unittest.mock import patch @@ -393,6 +395,323 @@ def _is_quark_mxfp4_ocp(self, hf_config: object) -> bool: assert reason is None +def test_deepseek_v4_heterogeneous_fhmoe_keeps_native_intermediate_width() -> None: + from aphrodite.models.deepseek_v4.amd.model import _prepare_native_fp8_shared_expert + + hidden_size = 7168 + intermediate_size = 384 + w13 = torch.empty((2 * intermediate_size, hidden_size), dtype=torch.float8_e4m3fn) + w2 = torch.empty((hidden_size, intermediate_size), dtype=torch.float8_e4m3fn) + w13_scale_bytes = (torch.arange(6 * 56).reshape(6, 56).remainder(20) + 0x60).to(torch.uint8) + w2_scale_bytes = (torch.arange(56 * 3).reshape(56, 3).remainder(30) + 0x50).to(torch.uint8) + w13_scale = w13_scale_bytes.view(torch.float8_e8m0fnu) + w2_scale = w2_scale_bytes.view(torch.float8_e8m0fnu) + + prepared = _prepare_native_fp8_shared_expert(w13, w2, w13_scale, w2_scale, intermediate_size) + + assert prepared[0].shape == (1, 768, 7168) + assert prepared[1].shape == (1, 7168, 384) + assert prepared[2].shape == (768, 224) + assert prepared[3].shape == (7168, 16) + expected_w13_scale = w13_scale_bytes.repeat_interleave(128, dim=0).repeat_interleave(4, dim=1) + expected_w2_scale = w2_scale_bytes.repeat_interleave(128, dim=0).repeat_interleave(4, dim=1) + assert torch.equal(prepared[2].view(torch.uint8), expected_w13_scale) + assert torch.equal(prepared[3].view(torch.uint8)[:, :12], expected_w2_scale) + assert torch.all(prepared[3].view(torch.uint8)[:, 12:] == 0x7F) + + +@pytest.mark.parametrize( + ("num_tokens", "supported_through", "expected"), + [ + (0, 4096, False), + (1, 4096, True), + (1536, 4096, True), + (2048, 4096, True), + (2049, 4096, True), + (4096, 4096, True), + (4097, 4096, False), + (1536, 0, False), + ], +) +def test_deepseek_v4_heterogeneous_fhmoe_token_policy( + monkeypatch: pytest.MonkeyPatch, + num_tokens: int, + supported_through: int, + expected: bool, +) -> None: + from aphrodite.models.deepseek_v4.amd import model as deepseek_v4_model + + checked_tokens: list[int] = [] + + def supports(num_tokens: int) -> bool: + checked_tokens.append(num_tokens) + return num_tokens <= supported_through + + monkeypatch.setattr( + deepseek_v4_model.rocm_aiter_ops, + "fused_moe_supports_heterogeneous_shared_expert", + supports, + ) + + assert deepseek_v4_model._use_heterogeneous_fhmoe(num_tokens) is expected + assert checked_tokens == ([] if num_tokens == 0 else [num_tokens]) + + +def _supported_fhmoe_signature( + shared_w1=None, + shared_w2=None, + shared_w1_scale=None, + shared_w2_scale=None, + shared_expert_id=-1, +) -> None: + pass + + +def _incomplete_fhmoe_signature( + shared_w1=None, + shared_w2=None, + shared_w1_scale=None, + shared_w2_scale=None, +) -> None: + pass + + +def _install_fake_aiter_fhmoe( + monkeypatch: pytest.MonkeyPatch, + supports_dsv4_i384_fhmoe: object, + fused_moe: object, +) -> None: + fake_aiter = ModuleType("aiter") + fake_aiter.__path__ = [] + fake_fhmoe = ModuleType("aiter.fhmoe") + if supports_dsv4_i384_fhmoe is not None: + fake_fhmoe.__dict__["supports_dsv4_i384_fhmoe"] = supports_dsv4_i384_fhmoe + fake_fused_moe = ModuleType("aiter.fused_moe") + fake_fused_moe.__dict__["fused_moe"] = fused_moe + fake_aiter.__dict__["fhmoe"] = fake_fhmoe + fake_aiter.__dict__["fused_moe"] = fake_fused_moe + monkeypatch.setitem(sys.modules, "aiter", fake_aiter) + monkeypatch.setitem(sys.modules, "aiter.fhmoe", fake_fhmoe) + monkeypatch.setitem(sys.modules, "aiter.fused_moe", fake_fused_moe) + + +def _supports_fhmoe_through_2047(max_tokens: int) -> bool: + return max_tokens <= 2047 + + +def _supports_fhmoe_through_2048(max_tokens: int) -> bool: + return max_tokens <= 2048 + + +def _supports_fhmoe_through_4096(max_tokens: int) -> bool: + return max_tokens <= 4096 + + +def _raises_fhmoe_config_error(max_tokens: int) -> bool: + raise OSError + + +def _returns_truthy_non_bool(max_tokens: int) -> int: + return 1 + + +@pytest.mark.parametrize( + ("num_tokens", "capability", "fused_moe", "expected"), + [ + (0, _supports_fhmoe_through_2048, _supported_fhmoe_signature, False), + (True, _supports_fhmoe_through_2048, _supported_fhmoe_signature, False), + (2048, None, _supported_fhmoe_signature, False), + (2048, True, _supported_fhmoe_signature, False), + (2048, _raises_fhmoe_config_error, _supported_fhmoe_signature, False), + (2048, _returns_truthy_non_bool, _supported_fhmoe_signature, False), + (2048, _supports_fhmoe_through_2047, _supported_fhmoe_signature, False), + (2048, _supports_fhmoe_through_2048, _incomplete_fhmoe_signature, False), + (2048, _supports_fhmoe_through_2048, _supported_fhmoe_signature, True), + (2049, _supports_fhmoe_through_2048, _supported_fhmoe_signature, False), + (4096, _supports_fhmoe_through_4096, _supported_fhmoe_signature, True), + (4097, _supports_fhmoe_through_4096, _supported_fhmoe_signature, False), + ], +) +def test_deepseek_v4_heterogeneous_fhmoe_aiter_capability( + monkeypatch: pytest.MonkeyPatch, + num_tokens: object, + capability: object, + fused_moe: object, + expected: bool, +) -> None: + from aphrodite._aiter_ops import rocm_aiter_ops + + _install_fake_aiter_fhmoe(monkeypatch, capability, fused_moe) + + assert rocm_aiter_ops._probe_dsv4_i384_fhmoe_capability(cast(int, num_tokens)) is expected + + +def test_deepseek_v4_heterogeneous_fhmoe_aiter_capability_catches_import_error( + monkeypatch: pytest.MonkeyPatch, +) -> None: + from aphrodite._aiter_ops import rocm_aiter_ops + + fake_aiter = ModuleType("aiter") + fake_aiter.__path__ = [] + monkeypatch.setitem(sys.modules, "aiter", fake_aiter) + monkeypatch.setitem(sys.modules, "aiter.fhmoe", None) + + assert not rocm_aiter_ops._probe_dsv4_i384_fhmoe_capability(2048) + + +@pytest.mark.parametrize("error", [TypeError, ValueError]) +def test_deepseek_v4_heterogeneous_fhmoe_aiter_capability_catches_signature_error( + monkeypatch: pytest.MonkeyPatch, + error: type[Exception], +) -> None: + from aphrodite._aiter_ops import rocm_aiter_ops + + _install_fake_aiter_fhmoe(monkeypatch, _supports_fhmoe_through_2048, _supported_fhmoe_signature) + + def raise_signature_error(_: object) -> inspect.Signature: + raise error + + monkeypatch.setattr(inspect, "signature", raise_signature_error) + + assert not rocm_aiter_ops._probe_dsv4_i384_fhmoe_capability(2048) + + +@pytest.mark.parametrize( + ("setting", "value", "expected"), + [ + ("data_parallel_size", 1, True), + ("data_parallel_size", 2, False), + ("prefill_context_parallel_size", 2, False), + ("topk_method", "greedy", False), + ("fhmoe_supported", False, False), + ], +) +def test_deepseek_v4_heterogeneous_fhmoe_compatibility_gates( + monkeypatch: pytest.MonkeyPatch, + setting: str, + value: object, + expected: bool, +) -> None: + from aphrodite.models.deepseek_v4.amd import model as deepseek_v4_model + + hf_config = SimpleNamespace( + n_routed_experts=384, + num_experts_per_tok=6, + n_shared_experts=1, + hidden_size=7168, + moe_intermediate_size=3072, + hidden_act="silu", + expert_dtype="fp4", + topk_method="noaux_tc", + ) + parallel_config = SimpleNamespace( + enable_expert_parallel=False, + enable_eplb=False, + tensor_parallel_size=8, + data_parallel_size=1, + prefill_context_parallel_size=1, + ) + if setting == "topk_method": + hf_config.topk_method = value + elif setting != "fhmoe_supported": + setattr(parallel_config, setting, value) + + quant_config = SimpleNamespace( + get_name=lambda: "deepseek_v4_fp8", + moe_quant_algo="", + weight_block_size=[128, 128], + is_checkpoint_fp8_serialized=True, + is_scale_e8m0=True, + ignored_layers=None, + ) + aphrodite_config = SimpleNamespace( + model_config=SimpleNamespace(hf_config=hf_config, dtype=torch.bfloat16), + quant_config=quant_config, + parallel_config=parallel_config, + kernel_config=SimpleNamespace(moe_backend="aiter"), + offload_config=None, + ) + monkeypatch.setattr(deepseek_v4_model.current_platform, "is_rocm", lambda: True) + monkeypatch.setattr( + deepseek_v4_model.envs, + "APHRODITE_ROCM_USE_AITER_FUSION_SHARED_EXPERTS", + True, + ) + monkeypatch.setattr(deepseek_v4_model, "on_gfx950", lambda: True) + monkeypatch.setattr( + deepseek_v4_model.rocm_aiter_ops, + "is_fusion_moe_shared_experts_enabled", + lambda: True, + ) + checked_tokens: list[int] = [] + fhmoe_supported = setting != "fhmoe_supported" or value is True + + def supports(num_tokens: int) -> bool: + checked_tokens.append(num_tokens) + return fhmoe_supported + + monkeypatch.setattr( + deepseek_v4_model.rocm_aiter_ops, + "fused_moe_supports_heterogeneous_shared_expert", + supports, + ) + + assert deepseek_v4_model._heterogeneous_shared_expert_enabled(cast(AphroditeConfig, aphrodite_config)) is expected + assert checked_tokens == [1] + + +@pytest.mark.parametrize( + ("weights_shape", "ids_shape", "message"), + [ + ((2, 7, 1), (2, 7, 1), "equal two-dimensional shapes"), + ((2, 7), (3, 7), "equal two-dimensional shapes"), + ((2, 6), (2, 6), "exactly 7 columns"), + ((2, 8), (2, 8), "exactly 7 columns"), + ], +) +def test_deepseek_v4_heterogeneous_fhmoe_rejects_invalid_routes( + weights_shape: tuple[int, ...], + ids_shape: tuple[int, ...], + message: str, +) -> None: + from aphrodite.models.deepseek_v4.amd.model import _validate_heterogeneous_routes + + with pytest.raises(ValueError, match=message): + _validate_heterogeneous_routes( + torch.empty(weights_shape), + torch.empty(ids_shape, dtype=torch.int64), + experts_per_token=6, + ) + + +def test_deepseek_v4_heterogeneous_fhmoe_accepts_appended_shared_route() -> None: + from aphrodite.models.deepseek_v4.amd.model import _validate_heterogeneous_routes + + _validate_heterogeneous_routes( + torch.empty((2, 7)), + torch.empty((2, 7), dtype=torch.int64), + experts_per_token=6, + ) + + +def test_aiter_fused_moe_validates_shared_expert_arguments() -> None: + from aphrodite._aiter_ops import ( + _validate_rocm_aiter_fused_moe_shared_expert_args, + ) + + validate = _validate_rocm_aiter_fused_moe_shared_expert_args + shared_tensor = torch.empty(0) + + assert not validate(None, None, None, None, -1) + with pytest.raises(ValueError, match="requires shared weights and scales"): + validate(None, None, None, None, 0) + with pytest.raises(ValueError, match="both shared weights and scales"): + validate(shared_tensor, None, None, None, 0) + with pytest.raises(ValueError, match="non-negative shared expert ID"): + validate(shared_tensor, shared_tensor, shared_tensor, shared_tensor, -1) + assert validate(shared_tensor, shared_tensor, shared_tensor, shared_tensor, 0) + + def test_is_model_fused_shared_expert_compatible() -> None: class MoE(nn.Module): def __init__(self, enabled: bool) -> None: From 112a1ee0deec2441f1d963ced3e71a883e9dc3fa Mon Sep 17 00:00:00 2001 From: AlpinDale Date: Mon, 7 Sep 2026 19:23:56 +0000 Subject: [PATCH 05/36] [sync] [Bugfix][Offloader] Preserve prefetch static-buffer slot ownership (#54975) Upstream-vLLM: 1f778486fc313f3599a6b59e67f65b1c186477fd Co-authored-by: yulun <100981785+Big2Wheel@users.noreply.github.com> --- .sync/vllm-sha | 2 +- .../model_executor/offloader/prefetch.py | 18 ++++++++++- .../model_executor/offloader/test_prefetch.py | 31 +++++++++++++++++++ 3 files changed, 49 insertions(+), 2 deletions(-) create mode 100644 tests/model_executor/offloader/test_prefetch.py diff --git a/.sync/vllm-sha b/.sync/vllm-sha index 7a08e801e9..e750a83a9d 100644 --- a/.sync/vllm-sha +++ b/.sync/vllm-sha @@ -1 +1 @@ -de69e821b7c834bcf1ec96325ec8c410ac183680 +1f778486fc313f3599a6b59e67f65b1c186477fd diff --git a/aphrodite/model_executor/offloader/prefetch.py b/aphrodite/model_executor/offloader/prefetch.py index ec0de7c737..89be4b1ec2 100644 --- a/aphrodite/model_executor/offloader/prefetch.py +++ b/aphrodite/model_executor/offloader/prefetch.py @@ -123,6 +123,18 @@ def get_buffer( return self._buffers[key][slot_idx % self.slot_capacity] +def _get_next_prefetch_index( + index: int, + prefetch_step: int, + module_count: int, +) -> int: + """Return a refill target that preserves static-buffer slot ownership.""" + next_index = (index + prefetch_step) % module_count + if prefetch_step < module_count and next_index % prefetch_step != index % prefetch_step: + next_index = index % prefetch_step + return next_index + + class PrefetchOffloader(BaseOffloader): """Prefetching-based offloader with group-based layer selection. @@ -226,7 +238,11 @@ def forward(*args, **kwargs): # Start prefetch for next layer (circular) # mutates_args on output_tensor creates ordering dependency - next_index = (index + self.prefetch_step) % len(self.module_offloaders) + next_index = _get_next_prefetch_index( + index, + self.prefetch_step, + len(self.module_offloaders), + ) # Handle tuple output (e.g., (hidden_states, residual)) if isinstance(output, tuple): torch.ops.aphrodite.start_prefetch(output[0], next_index) diff --git a/tests/model_executor/offloader/test_prefetch.py b/tests/model_executor/offloader/test_prefetch.py new file mode 100644 index 0000000000..b680c9dec1 --- /dev/null +++ b/tests/model_executor/offloader/test_prefetch.py @@ -0,0 +1,31 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project + +import pytest + +from aphrodite.model_executor.offloader.prefetch import _get_next_prefetch_index + +pytestmark = pytest.mark.skip_global_cleanup + + +@pytest.mark.parametrize( + ("module_count", "prefetch_step", "expected_targets"), + [ + (3, 1, [1, 2, 0]), + (3, 2, [2, 1, 0]), + (4, 2, [2, 3, 0, 1]), + (5, 2, [2, 3, 4, 1, 0]), + (3, 3, [0, 1, 2]), + (3, 4, [1, 2, 0]), + ], +) +def test_next_prefetch_index_preserves_slot_ownership( + module_count: int, + prefetch_step: int, + expected_targets: list[int], +) -> None: + targets = [_get_next_prefetch_index(index, prefetch_step, module_count) for index in range(module_count)] + + assert targets == expected_targets + if prefetch_step < module_count: + assert all(target % prefetch_step == index % prefetch_step for index, target in enumerate(targets)) From dc198f1ed1c36956b211ef0032a75b9278c47cf6 Mon Sep 17 00:00:00 2001 From: AlpinDale Date: Mon, 7 Sep 2026 19:26:49 +0000 Subject: [PATCH 06/36] [sync] [Bugfix] Fix cuda profiler missing bug (#55237) Upstream-vLLM: c3ec0d29f54cdd450576d819b1e5229ce6ef080e Co-authored-by: Wei Zhao <51183510+wzhao18@users.noreply.github.com> --- .sync/vllm-sha | 2 +- aphrodite/v1/engine/async_llm.py | 5 ++-- tests/v1/engine/test_async_llm.py | 44 ++++++++++++++++++++++++++++++- 3 files changed, 47 insertions(+), 4 deletions(-) diff --git a/.sync/vllm-sha b/.sync/vllm-sha index e750a83a9d..91387514bc 100644 --- a/.sync/vllm-sha +++ b/.sync/vllm-sha @@ -1 +1 @@ -1f778486fc313f3599a6b59e67f65b1c186477fd +c3ec0d29f54cdd450576d819b1e5229ce6ef080e diff --git a/aphrodite/v1/engine/async_llm.py b/aphrodite/v1/engine/async_llm.py index 8c1d129df8..b8bd64ddae 100644 --- a/aphrodite/v1/engine/async_llm.py +++ b/aphrodite/v1/engine/async_llm.py @@ -7,7 +7,7 @@ import warnings from collections.abc import AsyncGenerator, Iterable, Mapping from copy import copy -from typing import Any, Optional +from typing import Any import aphrodite.envs as envs from aphrodite import TokensPrompt @@ -91,7 +91,7 @@ def __init__( client_addresses: dict[str, Any] | None = None, client_count: int = 1, client_index: int = 0, - profiler: Optional[TorchProfilerWrapper] = None, # type: ignore # noqa + profiler: TorchProfilerWrapper | None = None, ) -> None: """ Create an AsyncLLM. @@ -190,6 +190,7 @@ def __init__( except RuntimeError: pass + self.profiler = profiler if ( aphrodite_config.profiler_config.profiler == "torch" and not aphrodite_config.profiler_config.ignore_frontend diff --git a/tests/v1/engine/test_async_llm.py b/tests/v1/engine/test_async_llm.py index 3e6014b7b0..1d2cbd122a 100644 --- a/tests/v1/engine/test_async_llm.py +++ b/tests/v1/engine/test_async_llm.py @@ -4,10 +4,11 @@ import asyncio import time from contextlib import ExitStack -from unittest.mock import MagicMock +from unittest.mock import AsyncMock, MagicMock, call import pytest +import aphrodite.v1.engine.async_llm as async_llm_module from aphrodite import SamplingParams from aphrodite.assets.image import ImageAsset from aphrodite.config import AphroditeConfig @@ -56,6 +57,47 @@ } +def test_cuda_profiler_requests_reach_engine_core(monkeypatch: pytest.MonkeyPatch): + aphrodite_config = MagicMock() + aphrodite_config.observability_config.otlp_traces_endpoint = None + aphrodite_config.scheduler_config.stream_interval = 1 + aphrodite_config.profiler_config.profiler = "cuda" + aphrodite_config.profiler_config.ignore_frontend = False + + renderer = MagicMock() + engine_core = MagicMock() + engine_core.profile_async = AsyncMock() + monkeypatch.setattr(async_llm_module, "maybe_register_config_serialize_by_value", MagicMock()) + monkeypatch.setattr( + async_llm_module, + "load_stat_logger_plugin_factories", + MagicMock(return_value=[]), + ) + monkeypatch.setattr( + async_llm_module, + "renderer_from_config", + MagicMock(return_value=renderer), + ) + monkeypatch.setattr(async_llm_module, "InputProcessor", MagicMock()) + monkeypatch.setattr(async_llm_module, "OutputProcessor", MagicMock()) + monkeypatch.setattr( + async_llm_module.EngineCoreClient, + "make_async_mp_client", + MagicMock(return_value=engine_core), + ) + + engine = AsyncLLM(aphrodite_config, MagicMock(), log_stats=False) + + async def profile(): + await engine.start_profile() + await engine.stop_profile() + + asyncio.run(profile()) + + assert engine.profiler is None + engine_core.profile_async.assert_has_awaits([call(True, None), call(False)]) + + async def generate( engine: AsyncLLM, request_id: str, From f53397351d70799f13c52b4ab53e10e0591b6c5c Mon Sep 17 00:00:00 2001 From: AlpinDale Date: Mon, 7 Sep 2026 19:29:47 +0000 Subject: [PATCH 07/36] [sync] [Bugfix] Fix Kimi K3 NVFP4 MoE weight conversion OOM (#55407) Upstream-vLLM: 6748217fb949952b4db418d09e64952cd7d8e352 Co-authored-by: Wei Zhao <51183510+wzhao18@users.noreply.github.com> --- .sync/vllm-sha | 2 +- .../quantization/utils/flashinfer_fp4_moe.py | 142 +++++++++--------- 2 files changed, 68 insertions(+), 76 deletions(-) diff --git a/.sync/vllm-sha b/.sync/vllm-sha index 91387514bc..38b8b90c7a 100644 --- a/.sync/vllm-sha +++ b/.sync/vllm-sha @@ -1 +1 @@ -c3ec0d29f54cdd450576d819b1e5229ce6ef080e +6748217fb949952b4db418d09e64952cd7d8e352 diff --git a/aphrodite/model_executor/layers/quantization/utils/flashinfer_fp4_moe.py b/aphrodite/model_executor/layers/quantization/utils/flashinfer_fp4_moe.py index 9befec508a..2a8bf50e6b 100644 --- a/aphrodite/model_executor/layers/quantization/utils/flashinfer_fp4_moe.py +++ b/aphrodite/model_executor/layers/quantization/utils/flashinfer_fp4_moe.py @@ -182,25 +182,23 @@ def nvfp4_swizzled_scale_to_cutedsl_mma_view(scale: torch.Tensor) -> torch.Tenso def prepare_static_weights_for_trtllm_fp4_moe( - # args_dequant, - # args, - gemm1_weights, - gemm2_weights, - gemm1_scales_linear_fp4_bytes, - gemm2_scales_linear_fp4_bytes, - hidden_size, - intermediate_size, - num_experts, + gemm1_weights: torch.Tensor, + gemm2_weights: torch.Tensor, + gemm1_scales_linear_fp4_bytes: torch.Tensor, + gemm2_scales_linear_fp4_bytes: torch.Tensor, + hidden_size: int, + intermediate_size: int, + num_experts: int, is_gated_activation: bool, -): +) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]: + """Shuffle NVFP4 weights into the FlashInfer TRT-LLM layout.""" from flashinfer import nvfp4_block_scale_interleave from flashinfer.fused_moe.core import ( _maybe_get_cached_w3_w1_permute_indices, get_w2_permute_indices_with_cache, ) - _cache_permute_indices: dict[torch.Size, torch.Tensor] = {} - """Prepare quantized weights for kernel (done offline with weights).""" + permute_indices_cache: dict[tuple[object, ...], torch.Tensor] = {} epilogue_tile_m = 128 # FIXME: this depends on the kernel internals gemm1_intermediate_size = 2 * intermediate_size if is_gated_activation else intermediate_size @@ -219,77 +217,71 @@ def prepare_static_weights_for_trtllm_fp4_moe( num_experts, hidden_size, intermediate_size // 16 ) # fp8 scaling factors - gemm1_weights_fp4_shuffled = [] - gemm1_scales_fp4_shuffled = [] - gemm2_weights_fp4_shuffled = [] - gemm2_scales_fp4_shuffled = [] - for i in range(num_experts): - # Calculate the permute indices for the following: - # 1. Reorder rows of W1 and scales for fused gated activation - # 2. Shuffle weights and scaling factors for transposed mma output - # for both w3_w1 and w2 weights and scale factors - permute_indices = _maybe_get_cached_w3_w1_permute_indices( - _cache_permute_indices, - gemm1_weights_fp4[i].view(torch.uint8), - epilogue_tile_m, - is_gated_act_gemm=is_gated_activation, - ) - gemm1_weights_fp4_shuffled.append( - gemm1_weights_fp4[i].view(torch.uint8)[permute_indices.to(gemm1_weights_fp4.device)].contiguous() - ) + gemm1_weight_indices = _maybe_get_cached_w3_w1_permute_indices( + permute_indices_cache, + gemm1_weights_fp4[0].view(torch.uint8), + epilogue_tile_m, + is_gated_act_gemm=is_gated_activation, + ) + gemm1_scale_indices = _maybe_get_cached_w3_w1_permute_indices( + permute_indices_cache, + gemm1_scales_linear_fp4[0].view(torch.uint8), + epilogue_tile_m, + num_elts_per_sf=16, + is_gated_act_gemm=is_gated_activation, + ) + gemm2_weight_indices = get_w2_permute_indices_with_cache( + permute_indices_cache, + gemm2_weights_fp4[0].view(torch.uint8), + epilogue_tile_m, + ) + gemm2_scale_indices = get_w2_permute_indices_with_cache( + permute_indices_cache, + gemm2_scales_linear_fp4[0].view(torch.uint8), + epilogue_tile_m, + num_elts_per_sf=16, + ) - permute_sf_indices = _maybe_get_cached_w3_w1_permute_indices( - _cache_permute_indices, - gemm1_scales_linear_fp4[i].view(torch.uint8), - epilogue_tile_m, - num_elts_per_sf=16, - is_gated_act_gemm=is_gated_activation, + gemm1_weights_fp4_shuffled = torch.empty_like(gemm1_weights_fp4, dtype=torch.uint8) + gemm1_scales_fp4_shuffled = torch.empty_like(gemm1_scales_linear_fp4) + gemm2_weights_fp4_shuffled = torch.empty_like(gemm2_weights_fp4, dtype=torch.uint8) + gemm2_scales_fp4_shuffled = torch.empty_like(gemm2_scales_linear_fp4) + gemm1_scale_scratch = torch.empty_like(gemm1_scales_linear_fp4[0], dtype=torch.uint8) + gemm2_scale_scratch = torch.empty_like(gemm2_scales_linear_fp4[0], dtype=torch.uint8) + + for expert_id in range(num_experts): + torch.index_select( + gemm1_weights_fp4[expert_id].view(torch.uint8), + 0, + gemm1_weight_indices, + out=gemm1_weights_fp4_shuffled[expert_id], ) - gemm1_scales_fp4_shuffled.append( - nvfp4_block_scale_interleave( - gemm1_scales_linear_fp4[i] - .view(torch.uint8)[permute_sf_indices.to(gemm1_scales_linear_fp4.device)] - .contiguous() - ) + torch.index_select( + gemm1_scales_linear_fp4[expert_id].view(torch.uint8), + 0, + gemm1_scale_indices, + out=gemm1_scale_scratch, ) - - permute_indices = get_w2_permute_indices_with_cache( - _cache_permute_indices, - gemm2_weights_fp4[i].view(torch.uint8), - epilogue_tile_m, - ) - gemm2_weights_fp4_shuffled.append( - gemm2_weights_fp4[i].view(torch.uint8)[permute_indices.to(gemm2_weights_fp4.device)].contiguous() + gemm1_scales_fp4_shuffled[expert_id].view(torch.uint8).reshape(-1).copy_( + nvfp4_block_scale_interleave(gemm1_scale_scratch) ) - permute_sf_indices = get_w2_permute_indices_with_cache( - _cache_permute_indices, - gemm2_scales_linear_fp4[i].view(torch.uint8), - epilogue_tile_m, - num_elts_per_sf=16, + torch.index_select( + gemm2_weights_fp4[expert_id].view(torch.uint8), + 0, + gemm2_weight_indices, + out=gemm2_weights_fp4_shuffled[expert_id], ) - gemm2_scales_fp4_shuffled.append( - nvfp4_block_scale_interleave( - gemm2_scales_linear_fp4[i] - .view(torch.uint8)[permute_sf_indices.to(gemm2_scales_linear_fp4.device)] - .contiguous() - ) + torch.index_select( + gemm2_scales_linear_fp4[expert_id].view(torch.uint8), + 0, + gemm2_scale_indices, + out=gemm2_scale_scratch, + ) + gemm2_scales_fp4_shuffled[expert_id].view(torch.uint8).reshape(-1).copy_( + nvfp4_block_scale_interleave(gemm2_scale_scratch) ) - # Stack weights for all experts - gemm1_weights_fp4_shuffled = torch.stack(gemm1_weights_fp4_shuffled) - gemm1_scales_fp4_shuffled = ( - torch.stack(gemm1_scales_fp4_shuffled) - .view(torch.float8_e4m3fn) - .reshape(num_experts, gemm1_intermediate_size, hidden_size // 16) - ) - - gemm2_weights_fp4_shuffled = torch.stack(gemm2_weights_fp4_shuffled) - gemm2_scales_fp4_shuffled = ( - torch.stack(gemm2_scales_fp4_shuffled) - .view(torch.float8_e4m3fn) - .reshape(num_experts, hidden_size, intermediate_size // 16) - ) return ( gemm1_weights_fp4_shuffled, gemm1_scales_fp4_shuffled, From aa31debc3c64e3eaa7e2ab1312f66acdd7f9c1dc Mon Sep 17 00:00:00 2001 From: AlpinDale Date: Mon, 7 Sep 2026 19:33:42 +0000 Subject: [PATCH 08/36] [sync] [Perf] Extend Qwen Triton warmup to avoid first-request latency spikes (#54797) Upstream-vLLM: 3dc7a68ce45ce1b98c2879139465833611f117cb Co-authored-by: vhagor <142512532+vhagor@users.noreply.github.com> --- .sync/vllm-sha | 2 +- .../model_executor/warmup/kernel_warmup.py | 4 + .../warmup/mamba_triton_warmup.py | 53 ++++++++ .../warmup/qwen_triton_warmup.py | 30 ++--- .../warmup/qwen_vl_triton_warmup.py | 110 +++++++++++++++++ .../test_mamba_triton_warmup.py | 14 +++ .../model_executor/test_qwen_triton_warmup.py | 53 ++++++++ .../test_qwen_vl_triton_warmup.py | 113 ++++++++++++++++++ 8 files changed, 359 insertions(+), 20 deletions(-) create mode 100644 aphrodite/model_executor/warmup/mamba_triton_warmup.py create mode 100644 aphrodite/model_executor/warmup/qwen_vl_triton_warmup.py create mode 100644 tests/model_executor/test_mamba_triton_warmup.py create mode 100644 tests/model_executor/test_qwen_triton_warmup.py create mode 100644 tests/model_executor/test_qwen_vl_triton_warmup.py diff --git a/.sync/vllm-sha b/.sync/vllm-sha index 38b8b90c7a..46b9b0351d 100644 --- a/.sync/vllm-sha +++ b/.sync/vllm-sha @@ -1 +1 @@ -6748217fb949952b4db418d09e64952cd7d8e352 +3dc7a68ce45ce1b98c2879139465833611f117cb diff --git a/aphrodite/model_executor/warmup/kernel_warmup.py b/aphrodite/model_executor/warmup/kernel_warmup.py index b16adcff42..f5db4a6176 100644 --- a/aphrodite/model_executor/warmup/kernel_warmup.py +++ b/aphrodite/model_executor/warmup/kernel_warmup.py @@ -34,10 +34,12 @@ from aphrodite.model_executor.warmup.kimi_k3_triton_warmup import ( kimi_k3_triton_warmup, ) +from aphrodite.model_executor.warmup.mamba_triton_warmup import mamba_triton_warmup from aphrodite.model_executor.warmup.qwen4_exp_qsa_warmup import ( qwen4_exp_qsa_triton_warmup, ) from aphrodite.model_executor.warmup.qwen_triton_warmup import qwen_triton_warmup +from aphrodite.model_executor.warmup.qwen_vl_triton_warmup import qwen_vl_triton_warmup from aphrodite.model_executor.warmup.replayssm_warmup import ( replayssm_autotune_warmup, ) @@ -153,6 +155,8 @@ def kernel_warmup(worker: "Worker", *, process_local_only: bool = False): ) qwen_triton_warmup(worker.model_runner, worker.aphrodite_config.model_config) + qwen_vl_triton_warmup(worker.model_runner) + mamba_triton_warmup(worker.model_runner) compilation_config = worker.aphrodite_config.compilation_config cudagraph_capture_sizes = list(compilation_config.cudagraph_capture_sizes or []) diff --git a/aphrodite/model_executor/warmup/mamba_triton_warmup.py b/aphrodite/model_executor/warmup/mamba_triton_warmup.py new file mode 100644 index 0000000000..c4907c5392 --- /dev/null +++ b/aphrodite/model_executor/warmup/mamba_triton_warmup.py @@ -0,0 +1,53 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +"""Warm Mamba-style Triton kernels shared across GDN / Mamba / KDA models.""" + +from typing import TYPE_CHECKING + +import torch + +from aphrodite.logger import init_logger + +if TYPE_CHECKING: + from aphrodite.v1.worker.gpu_model_runner import GPUModelRunner + +logger = init_logger(__name__) + + +def _has_mamba_style_cache(runner: "GPUModelRunner") -> bool: + from aphrodite.v1.kv_cache_interface import MambaSpec, UniformTypeKVCacheSpecs + + for group in runner.kv_cache_config.kv_cache_groups: + spec = group.kv_cache_spec + if isinstance(spec, UniformTypeKVCacheSpecs): + spec = spec.first_spec + if isinstance(spec, MambaSpec): + return True + return False + + +def _warm_batch_memcpy_kernel(device: torch.device) -> None: + """Warm the Mamba prefix-cache state copy specialization.""" + from aphrodite.v1.worker.mamba_utils import batch_memcpy + + src = torch.empty(1024, dtype=torch.uint8, device=device) + dst = torch.empty_like(src) + batch_memcpy( + torch.tensor([src.data_ptr()], dtype=torch.uint64, device=device), + torch.tensor([dst.data_ptr()], dtype=torch.uint64, device=device), + # Keep this int32: Triton's compile key includes pointer element dtypes, + # and the production prefix-cache path passes an int32 sizes tensor. + torch.tensor([src.numel()], dtype=torch.int32, device=device), + ) + logger.info("Warmed Mamba batch_memcpy_kernel.") + + +@torch.inference_mode() +def mamba_triton_warmup(runner: "GPUModelRunner") -> None: + """Warm prefix-cache memcpy for every Mamba-style KV cache group.""" + device = runner.device + if device.type != "cuda": + return + if not _has_mamba_style_cache(runner): + return + _warm_batch_memcpy_kernel(device) diff --git a/aphrodite/model_executor/warmup/qwen_triton_warmup.py b/aphrodite/model_executor/warmup/qwen_triton_warmup.py index 81e374006b..546fa61714 100644 --- a/aphrodite/model_executor/warmup/qwen_triton_warmup.py +++ b/aphrodite/model_executor/warmup/qwen_triton_warmup.py @@ -273,33 +273,25 @@ def qwen_triton_warmup( runner: "GPUModelRunner", model_config: "ModelConfig", ) -> None: - """Warm Qwen Triton kernels reported by the JIT monitor.""" - if runner.is_pooling_model: - return - - hf_text_config = getattr(model_config, "hf_text_config", None) - hf_config = getattr(model_config, "hf_config", None) - model_type = None - for config in (hf_text_config, hf_config): - model_type = getattr(config, "model_type", None) - if model_type is not None: - model_type = str(model_type) - break + """Warm Qwen GDN Triton kernels reported by the JIT monitor.""" + model_type = getattr(model_config.hf_text_config, "model_type", "") or getattr( + model_config.hf_config, "model_type", "" + ) if model_type not in _QWEN_MODEL_TYPES: return - device = getattr(runner, "device", torch.device("cuda")) - logger.info("Warming up Qwen Triton kernels for model_type=%s.", model_type) + device = runner.device + logger.info("Warming up Qwen GDN Triton kernels for model_type=%s.", model_type) - compilation_config = getattr(runner, "compilation_config", None) - static_forward_context = getattr(compilation_config, "static_forward_context", None) - gdn_config = _qwen_gdn_warmup_config(static_forward_context) + gdn_config = _qwen_gdn_warmup_config(runner.compilation_config.static_forward_context) if gdn_config is None: return - max_num_tokens = max(1, int(getattr(runner, "max_num_tokens", 1))) + max_num_tokens = max(1, int(runner.max_num_tokens)) _warm_gated_rms_norm_kernel(device, gdn_config, max_num_tokens, model_config.dtype) _warm_causal_conv1d_fwd_kernel(device, gdn_config) _warm_fused_post_conv_kernel(device, gdn_config) - _warm_fused_sigmoid_gating_delta_rule_update_kernel(device, gdn_config) + # Pooling only runs full prefills; the decode update kernel is unused. + if not runner.is_pooling_model: + _warm_fused_sigmoid_gating_delta_rule_update_kernel(device, gdn_config) _synchronize_device(device) diff --git a/aphrodite/model_executor/warmup/qwen_vl_triton_warmup.py b/aphrodite/model_executor/warmup/qwen_vl_triton_warmup.py new file mode 100644 index 0000000000..d69ad84e00 --- /dev/null +++ b/aphrodite/model_executor/warmup/qwen_vl_triton_warmup.py @@ -0,0 +1,110 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +"""Warm Qwen3-VL ViT kernels and M-RoPE for SupportsMRoPE models.""" + +from typing import TYPE_CHECKING + +import torch + +from aphrodite.logger import init_logger +from aphrodite.utils.math_utils import round_up + +if TYPE_CHECKING: + from aphrodite.v1.worker.gpu_model_runner import GPUModelRunner + +logger = init_logger(__name__) + +# Covers Triton's ==1 / %16==0 / other integer buckets. +_MROPE_TOKEN_COUNTS = (1, 2, 16) + + +def _warm_vision(model: torch.nn.Module, device: torch.device) -> None: + from aphrodite.model_executor.models.qwen3_vl import Qwen3_VisionTransformer + + for visual in model.modules(): + if not isinstance(visual, Qwen3_VisionTransformer): + continue + # CPU offloading keeps tower weights off the accelerator until the + # module forward hook copies them back, so warming here would launch + # Triton on host pointers. + if visual.pos_embed.weight.device.type != device.type: + continue + attention = visual.blocks[0].attn + num_heads = int(attention.num_attention_heads_per_partition) + head_size = int(attention.hidden_size_per_attention_head) + merge_size = int(visual.spatial_merge_size) + divisible = round_up(16, merge_size) + while divisible % 16: + divisible += merge_size + spatial_sizes = [divisible] + if merge_size % 16: + spatial_sizes.append(merge_size) + if merge_size == 1: + spatial_sizes.append(2) + grids = [(1, h, w) for h in spatial_sizes for w in spatial_sizes] + + for grid in grids: + grid_thw = [list(grid)] + visual.fast_pos_embed_interpolate(grid_thw) + cos, sin = visual.rot_pos_emb(grid_thw) + qk = torch.empty( + (2, grid[0] * grid[1] * grid[2], num_heads, head_size), + dtype=cos.dtype, + device=cos.device, + ) + attention.apply_rotary_emb(qk, cos, sin) + + logger.info( + "Warmed position embedding and vision rotary kernels on grids=%s.", + grids, + ) + + +def _warm_mrope(runner: "GPUModelRunner", model: torch.nn.Module) -> None: + from aphrodite.model_executor.layers.rotary_embedding.mrope import MRotaryEmbedding + from aphrodite.model_executor.models.interfaces import supports_mrope + + is_mrope_model = supports_mrope(model) + if not is_mrope_model: + return + + rope = next( + (module for module in model.modules() if isinstance(module, MRotaryEmbedding)), + None, + ) + if rope is None: + return + + model_config = runner.model_config + parallel_config = runner.parallel_config + num_query_heads = int(model_config.get_num_attention_heads(parallel_config)) + num_kv_heads = int(model_config.get_num_kv_heads(parallel_config)) + for num_tokens in _MROPE_TOKEN_COUNTS: + positions = torch.arange(num_tokens, dtype=torch.long, device=runner.device).expand(3, -1) + query = torch.empty( + (num_tokens, num_query_heads * rope.head_size), + dtype=runner.dtype, + device=runner.device, + ) + key_tensor = torch.empty( + (num_tokens, num_kv_heads * rope.head_size), + dtype=runner.dtype, + device=runner.device, + ) + rope(positions, query, key_tensor) + + logger.info("Warmed M-RoPE Triton kernels.") + + +def _synchronize_device(device: torch.device) -> None: + if device.type == "cuda": + torch.accelerator.synchronize(device) + + +@torch.inference_mode() +def qwen_vl_triton_warmup(runner: "GPUModelRunner") -> None: + """Warm Qwen3-VL ViT kernels and M-RoPE for ``SupportsMRoPE`` models.""" + model = runner.get_model() + _warm_vision(model, runner.device) + _warm_mrope(runner, model) + _synchronize_device(runner.device) diff --git a/tests/model_executor/test_mamba_triton_warmup.py b/tests/model_executor/test_mamba_triton_warmup.py new file mode 100644 index 0000000000..868dd5f09a --- /dev/null +++ b/tests/model_executor/test_mamba_triton_warmup.py @@ -0,0 +1,14 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project + +import pytest +import torch + +from aphrodite.model_executor.warmup.mamba_triton_warmup import _warm_batch_memcpy_kernel +from aphrodite.platforms import current_platform + + +@pytest.mark.skipif(not current_platform.is_cuda_alike(), reason="CUDA is required") +def test_mamba_batch_memcpy_kernel_compiles_on_gpu() -> None: + _warm_batch_memcpy_kernel(torch.device("cuda")) + torch.accelerator.synchronize("cuda") diff --git a/tests/model_executor/test_qwen_triton_warmup.py b/tests/model_executor/test_qwen_triton_warmup.py new file mode 100644 index 0000000000..c8bad730a2 --- /dev/null +++ b/tests/model_executor/test_qwen_triton_warmup.py @@ -0,0 +1,53 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project + +import pytest +import torch + +from aphrodite.model_executor.warmup.qwen_triton_warmup import ( + _FLA_POST_CONV_WARMUP_LENGTHS, + _QwenGDNWarmupConfig, + _warm_causal_conv1d_fwd_kernel, + _warm_fused_post_conv_kernel, + _warm_gated_rms_norm_kernel, +) +from aphrodite.platforms import current_platform + + +def _cuda_gdn_config() -> _QwenGDNWarmupConfig: + h, hv, k, v = 2, 2, 16, 16 + conv_kernel_size = 4 + conv_dim = 2 * h * k + hv * v + device = torch.device("cuda") + conv_state = torch.empty( + (8, conv_dim, conv_kernel_size - 1), + dtype=torch.bfloat16, + device=device, + ) + return _QwenGDNWarmupConfig( + h=h, + hv=hv, + k=k, + v=v, + conv_kernel_size=conv_kernel_size, + conv_state=conv_state, + conv_dtype=conv_state.dtype, + norm_weight_dtype=torch.bfloat16, + norm_before_gate=True, + norm_activation="silu", + a_log=torch.zeros(hv, dtype=torch.float32, device=device), + dt_bias=torch.zeros(hv, dtype=torch.float32, device=device), + state_stride_token=hv * v * k, + state_dtype=torch.float32, + ) + + +@pytest.mark.skipif(not current_platform.is_cuda_alike(), reason="CUDA is required") +def test_qwen_gdn_prefill_warmup_kernels_compile_on_gpu() -> None: + config = _cuda_gdn_config() + device = torch.device("cuda") + _warm_gated_rms_norm_kernel(device, config, max_num_tokens=16, x_dtype=config.conv_dtype) + _warm_causal_conv1d_fwd_kernel(device, config) + _warm_fused_post_conv_kernel(device, config) + assert _FLA_POST_CONV_WARMUP_LENGTHS == (1, 2, 16) + torch.accelerator.synchronize(device) diff --git a/tests/model_executor/test_qwen_vl_triton_warmup.py b/tests/model_executor/test_qwen_vl_triton_warmup.py new file mode 100644 index 0000000000..9a55b1b23c --- /dev/null +++ b/tests/model_executor/test_qwen_vl_triton_warmup.py @@ -0,0 +1,113 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project + +from types import SimpleNamespace + +import torch + +from aphrodite.model_executor.warmup.qwen_vl_triton_warmup import ( + _warm_mrope, + _warm_vision, +) + + +def test_vision_warmup_calls_only_position_and_rotary_paths() -> None: + from aphrodite.model_executor.models.qwen3_vl import Qwen3_VisionTransformer + + calls: list[tuple[str, object]] = [] + + class FakeAttention: + num_attention_heads_per_partition = 4 + hidden_size_per_attention_head = 8 + + def apply_rotary_emb(self, qk, cos, sin): + calls.append(("rotary", qk.shape)) + + class FakeVisual(Qwen3_VisionTransformer): + spatial_merge_size = 2 + + def __init__(self) -> None: + torch.nn.Module.__init__(self) + self.blocks = [SimpleNamespace(attn=FakeAttention())] + self.pos_embed = torch.nn.Embedding(4, 4, device="cpu") + + def fast_pos_embed_interpolate(self, grid_thw): + calls.append(("position", grid_thw[0])) + + def rot_pos_emb(self, grid_thw): + _, h, w = grid_thw[0] + shape = (h * w, 4) + return torch.empty(shape), torch.empty(shape) + + def forward(self, *args, **kwargs): + raise AssertionError("vision warmup must not run the full tower") + + model = torch.nn.Module() + model.visual = FakeVisual() + _warm_vision(model, torch.device("cpu")) + + grids = [(1, 16, 16), (1, 16, 2), (1, 2, 16), (1, 2, 2)] + assert [value for name, value in calls if name == "position"] == [list(grid) for grid in grids] + assert [value for name, value in calls if name == "rotary"] == [torch.Size((2, h * w, 4, 8)) for _, h, w in grids] + + # CPU offloading leaves tower weights off the accelerator, so the Triton + # paths must not be launched on host pointers. + calls.clear() + _warm_vision(model, torch.device("cuda")) + assert calls == [] + + +def test_mrope_warmup_launches_triton_shapes() -> None: + from aphrodite.model_executor.layers.rotary_embedding.mrope import MRotaryEmbedding + + launched: list[tuple[torch.Size, torch.Size]] = [] + + class FakeRope(MRotaryEmbedding): + def __init__(self) -> None: + torch.nn.Module.__init__(self) + self.head_size = 8 + self.rotary_dim = 8 + self.mrope_section = [2, 3, 3] + self.mrope_interleaved = False + self.is_neox_style = True + + def forward(self, positions, query, key): + launched.append((positions.shape, query.shape)) + return query, key + + class FakeMropeModel(torch.nn.Module): + supports_mrope = True + + def __init__(self) -> None: + super().__init__() + self.rotary_emb = FakeRope() + + def get_mrope_input_positions(self, input_tokens, mm_features): + raise AssertionError("warmup must not build live M-RoPE positions") + + runner = SimpleNamespace( + model_config=SimpleNamespace( + get_num_attention_heads=lambda parallel_config: 4, + get_num_kv_heads=lambda parallel_config: 2, + ), + parallel_config=object(), + dtype=torch.bfloat16, + device=torch.device("cpu"), + ) + _warm_mrope(runner, FakeMropeModel()) + assert [shape for shape, _ in launched] == [ + torch.Size((3, 1)), + torch.Size((3, 2)), + torch.Size((3, 16)), + ] + + +def test_mrope_warmup_skips_models_without_supports_mrope() -> None: + def fail(*_args, **_kwargs): + raise AssertionError("M-RoPE warmup must not run without SupportsMRoPE") + + runner = SimpleNamespace( + model_config=SimpleNamespace(get_num_attention_heads=fail), + get_model=fail, + ) + _warm_mrope(runner, torch.nn.Module()) From bb927cdddff6affe31f4b2cb6694235a2539f71d Mon Sep 17 00:00:00 2001 From: AlpinDale Date: Mon, 7 Sep 2026 19:37:50 +0000 Subject: [PATCH 09/36] [sync] [Bugfix] Gracefully handle unsupported reasoning_effort in chat templates (#54022) Upstream-vLLM: 9cc7793e32f8340e111a6c822f231aa526a8ea21 Co-authored-by: frankie --- .sync/vllm-sha | 2 +- aphrodite/renderers/hf.py | 22 +++++++- tests/renderers/test_hf.py | 100 +++++++++++++++++++++++++++++++++++++ 3 files changed, 121 insertions(+), 3 deletions(-) diff --git a/.sync/vllm-sha b/.sync/vllm-sha index 46b9b0351d..2438245378 100644 --- a/.sync/vllm-sha +++ b/.sync/vllm-sha @@ -1 +1 @@ -3dc7a68ce45ce1b98c2879139465833611f117cb +9cc7793e32f8340e111a6c822f231aa526a8ea21 diff --git a/aphrodite/renderers/hf.py b/aphrodite/renderers/hf.py index 6edf989b96..3b0bd630f8 100644 --- a/aphrodite/renderers/hf.py +++ b/aphrodite/renderers/hf.py @@ -27,6 +27,7 @@ parse_chat_messages, parse_chat_messages_async, ) +from aphrodite.exceptions import AphroditeValidationError from aphrodite.inputs import EmbedsPrompt from aphrodite.inputs.engine import MultiModalInput from aphrodite.logger import init_logger @@ -651,6 +652,18 @@ def resolve_chat_template_kwargs( return {k: v for k, v in chat_template_kwargs.items() if k in accept_vars} +def _template_error_reason(exc: BaseException) -> str: + # Extract the most specific reason from a chat template error chain. + seen: set[int] = set() + current: BaseException | None = exc + while current is not None and id(current) not in seen: + seen.add(id(current)) + if isinstance(current, jinja2.TemplateError): + return str(current) + current = current.__cause__ or current.__context__ + return str(exc) + + @overload def safe_apply_chat_template( model_config: ModelConfig, @@ -769,8 +782,13 @@ def safe_apply_chat_template( **resolved_kwargs, ) except Exception as e: - logger.exception("An error occurred in `transformers` while applying chat template") - raise ValueError(str(e)) from e + # Chat templates reject invalid user input (e.g. an unsupported + # `reasoning_effort` value) by raising from within the template. + # Surface those as a 400 Bad Request carrying the template's own + # reason (which typically lists the supported values) instead of a + # 500 or any generic upstream wrapper message. + logger.warning("Chat template rejected the request: %s", e) + raise AphroditeValidationError(_template_error_reason(e)) from e if return_assistant_tokens_mask: assert isinstance(plain, list), f"Expected list[int], got {type(plain)}" diff --git a/tests/renderers/test_hf.py b/tests/renderers/test_hf.py index d6ce2d9197..e1a50f29f6 100644 --- a/tests/renderers/test_hf.py +++ b/tests/renderers/test_hf.py @@ -1,6 +1,7 @@ # SPDX-License-Identifier: Apache-2.0 # SPDX-FileCopyrightText: Copyright contributors to the vLLM project +import jinja2 import pytest from huggingface_hub.errors import GatedRepoError @@ -13,6 +14,7 @@ _convert_developer_to_system, _detect_developer_role_support, _get_hf_base_chat_template_params, + _template_error_reason, _try_extract_ast, resolve_chat_template, resolve_chat_template_content_format, @@ -842,6 +844,104 @@ def test_developer_only_no_prior_system(self, model_config, tokenizer): ) +EFFORT_VALIDATING_TEMPLATE = ( + "{% if reasoning_effort is defined and " + "reasoning_effort not in ['xhigh', 'medium', 'low'] %}" + "{{ raise_exception('Unexpected reasoning effort ' + reasoning_effort + " + "'. Supported types are xhigh (default), medium, and low.') }}" + "{% endif %}" + "{% for message in messages %}" + "{{ message['role'] }}: {{ message['content'] }}\n" + "{% endfor %}" +) + +EFFORT_BREAKING_TEMPLATE = ( + "{% if reasoning_effort is defined %}" + "{{ raise_exception('Template broke for an unrelated reason') }}" + "{% endif %}" + "{% for message in messages %}" + "{{ message['role'] }}: {{ message['content'] }}\n" + "{% endfor %}" +) + + +class TestApplyChatTemplateEffortTolerant: + """Chat templates that reject unsupported reasoning_effort values should + surface a 400-style client error instead of crashing the request.""" + + @pytest.fixture + def model_config(self): + return ModelConfig( + "facebook/opt-125m", + tokenizer="facebook/opt-125m", + tokenizer_mode="auto", + trust_remote_code=False, + dtype="float16", + ) + + @pytest.fixture + def tokenizer(self): + return get_tokenizer("facebook/opt-125m") + + def test_unsupported_effort_raises_bad_request(self, model_config, tokenizer): + conversation = [{"role": "user", "content": "Hello"}] + with pytest.raises( + APHRODITEValidationError, + match="Unexpected reasoning effort high", + ) as excinfo: + safe_apply_chat_template( + model_config, + tokenizer, + conversation, + chat_template=EFFORT_VALIDATING_TEMPLATE, + tokenize=False, + reasoning_effort="high", + ) + # The client sees exactly the template's own reason, with no wrapper + # noise. It tells them which values are supported. + assert str(excinfo.value) == ( + "Unexpected reasoning effort high. Supported types are xhigh (default), medium, and low." + ) + + def test_supported_effort_accepted(self, model_config, tokenizer): + conversation = [{"role": "user", "content": "Hello"}] + result = safe_apply_chat_template( + model_config, + tokenizer, + conversation, + chat_template=EFFORT_VALIDATING_TEMPLATE, + tokenize=False, + reasoning_effort="medium", + ) + assert result == "user: Hello\n" + + def test_non_effort_template_error_is_bad_request(self, model_config, tokenizer): + conversation = [{"role": "user", "content": "Hello"}] + with pytest.raises(APHRODITEValidationError, match="unrelated reason"): + safe_apply_chat_template( + model_config, + tokenizer, + conversation, + chat_template=EFFORT_BREAKING_TEMPLATE, + tokenize=False, + reasoning_effort="high", + ) + + +def test_template_error_reason_prefers_template_error(): + # Upstream wrappers must not hide the template's own reason from clients. + reason = jinja2.TemplateError("Unexpected reasoning effort high") + wrapper = ValueError(f"An error occurred while rendering the template: {reason}") + wrapper.__cause__ = reason + assert _template_error_reason(wrapper) == str(reason) + assert _template_error_reason(reason) == str(reason) + + +def test_template_error_reason_falls_back_to_message(): + err = ValueError("plain failure") + assert _template_error_reason(err) == "plain failure" + + class TestConsolidateSystemMessages: def test_no_system_messages_unchanged(self): conversation = [ From daec778cabd8ba8a6e99029ba3264f5756d94c1b Mon Sep 17 00:00:00 2001 From: AlpinDale Date: Mon, 7 Sep 2026 19:42:33 +0000 Subject: [PATCH 10/36] [sync] [Qwen3.8-Flash-Next] Remove torch.compile for NVIDIA implementation (#55272) Upstream-vLLM: d9105ea8001e0a6d77a96327d17515bb5791fb36 Co-authored-by: Thien Tran --- .sync/vllm-sha | 2 +- aphrodite/config/aphrodite.py | 3 + aphrodite/config/compilation.py | 2 +- .../layers/mamba/gdn/qwen_gdn_linear_attn.py | 3 + .../models/qwen4_exp/nvidia/indexer_qsa.py | 14 ++-- aphrodite/models/qwen4_exp/nvidia/model.py | 12 --- .../models/qwen4_exp/nvidia/model_state.py | 18 +++-- aphrodite/models/qwen4_exp/nvidia/mtp.py | 19 ----- .../models/qwen4_exp/nvidia/ple_layer.py | 64 ++-------------- aphrodite/models/qwen4_exp/nvidia/qsa.py | 74 ++++--------------- tests/models/qwen4_exp/test_config.py | 53 +++++++++---- tests/models/qwen4_exp/test_ple.py | 12 +++ tests/models/qwen4_exp/test_qsa_reference.py | 64 +++++++++++++++- tests/test_config.py | 8 +- 14 files changed, 162 insertions(+), 186 deletions(-) diff --git a/.sync/vllm-sha b/.sync/vllm-sha index 2438245378..d789e14bd5 100644 --- a/.sync/vllm-sha +++ b/.sync/vllm-sha @@ -1 +1 @@ -9cc7793e32f8340e111a6c822f231aa526a8ea21 +d9105ea8001e0a6d77a96327d17515bb5791fb36 diff --git a/aphrodite/config/aphrodite.py b/aphrodite/config/aphrodite.py index ddf685a4cb..edb1bafdff 100644 --- a/aphrodite/config/aphrodite.py +++ b/aphrodite/config/aphrodite.py @@ -100,6 +100,9 @@ "KimiLinearForCausalLM", "MiniMaxM3SparseForCausalLM", "MiniMaxM3SparseForConditionalGeneration", + "Qwen4ExpForCausalLM", + "Qwen4ExpForConditionalGeneration", + "Qwen4ExpMTP", } ) diff --git a/aphrodite/config/compilation.py b/aphrodite/config/compilation.py index dfc0224062..dd42edc40c 100644 --- a/aphrodite/config/compilation.py +++ b/aphrodite/config/compilation.py @@ -750,7 +750,7 @@ class CompilationConfig: "aphrodite::mamba_mixer2", "aphrodite::mamba_mixer", "aphrodite::short_conv", - "aphrodite::qwen4_exp_compute_ple_ngram_ids", + # Qwen4Exp's AMD backend still uses these splitting ops. "aphrodite::qwen4_exp_ple_short_conv", "aphrodite::qwen4_exp_qsa_with_output", "aphrodite::linear_attention", diff --git a/aphrodite/model_executor/layers/mamba/gdn/qwen_gdn_linear_attn.py b/aphrodite/model_executor/layers/mamba/gdn/qwen_gdn_linear_attn.py index 3e197d83c4..f37cb967a9 100644 --- a/aphrodite/model_executor/layers/mamba/gdn/qwen_gdn_linear_attn.py +++ b/aphrodite/model_executor/layers/mamba/gdn/qwen_gdn_linear_attn.py @@ -12,6 +12,7 @@ from aphrodite import _custom_ops as ops from aphrodite import envs from aphrodite._aiter_ops import rocm_aiter_ops +from aphrodite.compilation.breakable_cudagraph import eager_break_during_capture from aphrodite.config import ( AphroditeConfig, get_current_aphrodite_config, @@ -1815,6 +1816,7 @@ def _forward_core_fused_norm( ) +@eager_break_during_capture def qwen_gdn_attention_core( qkv_or_qkvz: torch.Tensor, b_or_ba: torch.Tensor, @@ -1862,6 +1864,7 @@ def qwen_gdn_attention_core( ) +@eager_break_during_capture def qwen_gdn_attention_core_fused_norm_packed( mixed_qkvz: torch.Tensor, ba: torch.Tensor, diff --git a/aphrodite/models/qwen4_exp/nvidia/indexer_qsa.py b/aphrodite/models/qwen4_exp/nvidia/indexer_qsa.py index 816d0dcd73..c3fa9823af 100644 --- a/aphrodite/models/qwen4_exp/nvidia/indexer_qsa.py +++ b/aphrodite/models/qwen4_exp/nvidia/indexer_qsa.py @@ -88,7 +88,7 @@ def _supports_fused_pre_indexer( class QSAIndexer(nn.Module): - """Replicated Q/K projection plus paged, weight-free QSA selection. + """QSA projection weights, side caches, and paged, weight-free selection. ``prefix`` must be the checkpoint's indexer prefix, normally ``model.layers.N.self_attn.indexer``. Consequently the trainable names are @@ -214,11 +214,11 @@ def _metadata( def forward( self, - hidden_states: torch.Tensor, + projected_qk: torch.Tensor, positions: torch.Tensor, out: torch.Tensor | None = None, ) -> torch.Tensor: - """Select each query row's token indices. + """Update side caches and select token indices from pre-projected Q/K. Returns the packed buffer of shape [num_tokens, output_width + 1]: the leading ``output_width`` columns are ``-1``-padded @@ -233,10 +233,10 @@ def forward( if self.skip_topk and out is not None: return out result = torch.full( - (hidden_states.shape[0], self.packed_output_width), + (projected_qk.shape[0], self.packed_output_width), -1, dtype=torch.int32, - device=hidden_states.device, + device=projected_qk.device, ) # Inert rows carry a zero valid count (empty loop bound), not -1. result[:, -1] = 0 @@ -254,11 +254,9 @@ def forward( raw_metadata, compressed_metadata = metadata num_tokens = raw_metadata.num_actual_tokens - hidden_states = hidden_states[:num_tokens] + projected_qk = projected_qk[:num_tokens] positions = positions[..., :num_tokens] - # Q/K projection - projected_qk, _ = self.index_qk_proj(hidden_states) projected_q, raw_keys = projected_qk.split( ( self.index_n_heads * self.index_head_dim, diff --git a/aphrodite/models/qwen4_exp/nvidia/model.py b/aphrodite/models/qwen4_exp/nvidia/model.py index 0719a3b1f0..d7292335a3 100644 --- a/aphrodite/models/qwen4_exp/nvidia/model.py +++ b/aphrodite/models/qwen4_exp/nvidia/model.py @@ -8,7 +8,6 @@ import torch from torch import nn -from aphrodite.compilation.decorators import support_torch_compile from aphrodite.config import AphroditeConfig from aphrodite.distributed import get_pp_group from aphrodite.model_executor.layers.fused_moe.utils import ( @@ -367,17 +366,6 @@ def update_physical_experts_metadata( moe.experts.update_expert_map() -@support_torch_compile( - dynamic_arg_dims={ - "input_ids": 0, - "positions": -1, - "intermediate_tensors": 0, - "inputs_embeds": 0, - "query_start_loc": 0, - "ngram_context": 0, - "deepstack_input_embeds": 0, - } -) class Qwen4ExpModel(nn.Module): hf_to_aphrodite_mapper = Qwen3_5Model.hf_to_aphrodite_mapper | _EXTRA_WEIGHTS_MAPPER diff --git a/aphrodite/models/qwen4_exp/nvidia/model_state.py b/aphrodite/models/qwen4_exp/nvidia/model_state.py index ceac40cce4..9284c63cc5 100644 --- a/aphrodite/models/qwen4_exp/nvidia/model_state.py +++ b/aphrodite/models/qwen4_exp/nvidia/model_state.py @@ -44,6 +44,8 @@ def __init__( if self.ngram_context_len <= 0: raise ValueError("N-gram embedding requires context length >= 1.") self.ngram_eos_token_id = int(config.eos_token_id) + # PLE runs inside captured regions, so these buffers keep a fixed shape + # and address as the active request count changes between replays. self.ngram_context = torch.full( (self.max_num_reqs, self.ngram_context_len), self.ngram_eos_token_id, @@ -68,8 +70,7 @@ def _prepare_ngram_context( req_states: RequestState, ) -> torch.Tensor: num_reqs = input_batch.num_reqs - num_reqs_padded = input_batch.num_reqs_after_padding - context = self.ngram_context[:num_reqs_padded] + context = self.ngram_context context.fill_(self.ngram_eos_token_id) if num_reqs == 0: return context @@ -99,8 +100,10 @@ def prepare_inputs( return model_inputs num_reqs_padded = input_batch.num_reqs_after_padding - query_start_loc = self.ple_query_start_loc[: num_reqs_padded + 1] - query_start_loc.copy_(input_batch.query_start_loc[: num_reqs_padded + 1]) + query_start_loc = self.ple_query_start_loc + query_start_loc[: num_reqs_padded + 1].copy_(input_batch.query_start_loc) + # Represent unused capacity as trailing zero-length requests. + query_start_loc[num_reqs_padded + 1 :].copy_(input_batch.query_start_loc[-1]) model_inputs.update( query_start_loc=query_start_loc, ngram_context=self._prepare_ngram_context(input_batch, req_states), @@ -116,7 +119,7 @@ def prepare_dummy_inputs( if not self.uses_ngram_embedding: return model_inputs - query_start_loc = self.ple_query_start_loc[: num_reqs + 1] + query_start_loc = self.ple_query_start_loc query_start_loc[0] = 0 tokens_per_req, num_extra_tokens = divmod(num_tokens, num_reqs) query_lens = torch.full( @@ -127,9 +130,10 @@ def prepare_dummy_inputs( ) if num_extra_tokens > 0: query_lens[-num_extra_tokens:] += 1 - torch.cumsum(query_lens, dim=0, out=query_start_loc[1:]) + torch.cumsum(query_lens, dim=0, out=query_start_loc[1 : num_reqs + 1]) + query_start_loc[num_reqs + 1 :].fill_(num_tokens) - ngram_context = self.ngram_context[:num_reqs] + ngram_context = self.ngram_context ngram_context.fill_(self.ngram_eos_token_id) model_inputs.update( query_start_loc=query_start_loc, diff --git a/aphrodite/models/qwen4_exp/nvidia/mtp.py b/aphrodite/models/qwen4_exp/nvidia/mtp.py index fc9551c1b0..7e18af4aa9 100644 --- a/aphrodite/models/qwen4_exp/nvidia/mtp.py +++ b/aphrodite/models/qwen4_exp/nvidia/mtp.py @@ -19,7 +19,6 @@ import torch from torch import nn -from aphrodite.compilation.decorators import support_torch_compile from aphrodite.config import AphroditeConfig, replace, set_current_aphrodite_config from aphrodite.distributed import get_pp_group from aphrodite.model_executor.layers.fused_moe.utils import ( @@ -147,15 +146,6 @@ def _make_draft_aphrodite_config( return draft_aphrodite_config -@support_torch_compile( - dynamic_arg_dims={ - "input_ids": 0, - "positions": -1, - "intermediate_tensors": 0, - "inputs_embeds": 0, - "hidden_states": 0, - } -) class Qwen4ExpMultiTokenPredictor(nn.Module): hf_to_aphrodite_mapper = Qwen3_5Model.hf_to_aphrodite_mapper | _EXTRA_WEIGHTS_MAPPER @@ -342,15 +332,6 @@ def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]: return loader.load_weights(weights, mapper=mapper) -@support_torch_compile( - dynamic_arg_dims={ - "input_ids": 0, - "positions": -1, - "intermediate_tensors": 0, - "inputs_embeds": 0, - "hidden_states": 0, - } -) class Qwen4ExpMTP(nn.Module, SupportsPP, Qwen4ExpMixtureOfExperts): packed_modules_mapping = { "qkv_proj": ["q_proj", "k_proj", "v_proj"], diff --git a/aphrodite/models/qwen4_exp/nvidia/ple_layer.py b/aphrodite/models/qwen4_exp/nvidia/ple_layer.py index d2837d2d76..2948525e97 100644 --- a/aphrodite/models/qwen4_exp/nvidia/ple_layer.py +++ b/aphrodite/models/qwen4_exp/nvidia/ple_layer.py @@ -8,6 +8,7 @@ import torch.nn.functional as F from torch import nn +from aphrodite.compilation.breakable_cudagraph import eager_break_during_capture from aphrodite.config import AphroditeConfig, CacheConfig, ModelConfig, get_current_aphrodite_config from aphrodite.forward_context import get_forward_context from aphrodite.model_executor.layers.linear import MergedColumnParallelLinear @@ -38,7 +39,6 @@ from aphrodite.transformers_utils.configs.qwen4_exp import ( Qwen4ExpTextConfig, ) -from aphrodite.utils.torch_utils import direct_register_custom_op from aphrodite.v1.attention.backends.registry import MambaAttentionBackendEnum from aphrodite.v1.attention.backends.short_conv_attn import ( PleShortConvAttentionBackend, @@ -257,12 +257,10 @@ def __init__( embedding_dim: int, ple_dense_layer_id: int, prefix: str, - layer_name: str, quant_config: QuantizationConfig | None = None, params_dtype: torch.dtype | None = None, ) -> None: super().__init__() - self.layer_name = layer_name self.embedding_dim = embedding_dim self.ngram_size = int(config.ngram_size) self.heads_per_ngram = int(config.heads_per_ngram) @@ -425,18 +423,7 @@ def forward( query_start_loc: torch.Tensor, ngram_context: torch.Tensor, ) -> torch.Tensor: - # Keep num_reqs-dependent ID generation outside PIECEWISE CUDA graphs, - # which dispatch only on the padded token count. - # torch.compile requires the splitting op to write graph-owned storage. - # Once compilation is removed, the op can return the IDs directly. - ngram_ids = input_ids.new_empty((input_ids.numel(), self.ngram_heads), dtype=torch.long) - torch.ops.aphrodite.qwen4_exp_compute_ple_ngram_ids( - input_ids, - query_start_loc, - ngram_context, - ngram_ids, - self.layer_name, - ) + ngram_ids = self.compute_ngram_ids(input_ids, query_start_loc, ngram_context) return self.ngram_embedding(ngram_ids).flatten(-2) def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]: @@ -533,7 +520,6 @@ def __init__( int(config.ple_embed_dim), self.ple_dense_layer_id, prefix=f"{prefix}.ple_embedding", - layer_name=prefix, quant_config=quant_config, params_dtype=model_config.dtype, ) @@ -735,6 +721,8 @@ def _short_conv_dilated_dispatch( token_indices=non_spec_token_indices, ) + # State routing consumes the current request metadata on every replay. + @eager_break_during_capture def _short_conv(self, inputs: torch.Tensor, residual: torch.Tensor) -> None: forward_context = get_forward_context() attn_metadata = forward_context.attn_metadata @@ -810,52 +798,10 @@ def forward( self.norm_conv.weight, self.norm_key.eps, ) - # State routing depends on runtime request metadata and remains outside - # the piecewise graph; short convolution accumulates into gated_output. - torch.ops.aphrodite.qwen4_exp_ple_short_conv(conv_input, gated_output, self.prefix) + self._short_conv(conv_input, gated_output) return gated_output -def qwen4_exp_compute_ple_ngram_ids( - input_ids: torch.Tensor, - query_start_loc: torch.Tensor, - ngram_context: torch.Tensor, - output: torch.Tensor, - layer_name: str, -) -> None: - """Compute request-dependent PLE n-gram IDs outside piecewise graphs.""" - layer = get_forward_context().no_compile_layers[layer_name] - layer.ple_embedding.compute_ngram_ids( - input_ids, - query_start_loc, - ngram_context, - output, - ) - - -def qwen4_exp_ple_short_conv( - inputs: torch.Tensor, - residual_output: torch.Tensor, - layer_name: str, -) -> None: - layer = get_forward_context().no_compile_layers[layer_name] - layer._short_conv(inputs, residual_output) - - -direct_register_custom_op( - op_name="qwen4_exp_compute_ple_ngram_ids", - op_func=qwen4_exp_compute_ple_ngram_ids, - mutates_args=["output"], -) - - -direct_register_custom_op( - op_name="qwen4_exp_ple_short_conv", - op_func=qwen4_exp_ple_short_conv, - mutates_args=["residual_output"], -) - - __all__ = [ "Qwen4ExpNGramEmbedding", "Qwen4ExpPLEGroupedNorm", diff --git a/aphrodite/models/qwen4_exp/nvidia/qsa.py b/aphrodite/models/qwen4_exp/nvidia/qsa.py index bbf9119e7b..5d2a916f87 100644 --- a/aphrodite/models/qwen4_exp/nvidia/qsa.py +++ b/aphrodite/models/qwen4_exp/nvidia/qsa.py @@ -9,6 +9,7 @@ import torch from torch import nn +from aphrodite.compilation.breakable_cudagraph import eager_break_during_capture from aphrodite.config import AphroditeConfig from aphrodite.config.cache import CacheDType from aphrodite.distributed import get_tensor_model_parallel_world_size @@ -27,10 +28,6 @@ Qwen4ExpTextConfig, ) from aphrodite.utils.torch_utils import ( - LayerNameType, - _encode_layer_name, - _resolve_layer_name, - direct_register_custom_op, kv_cache_dtype_str_to_dtype, ) from aphrodite.v1.attention.backend import ( @@ -336,9 +333,10 @@ def get_kv_cache_spec(self, aphrodite_config: AphroditeConfig) -> KVCacheSpec: kv_quant_mode=get_kv_quant_mode(self.kv_cache_dtype), ) + @eager_break_during_capture def _run_qsa( self, - hidden_states: torch.Tensor, + projected_qk: torch.Tensor, positions: torch.Tensor, query: torch.Tensor, key: torch.Tensor, @@ -363,7 +361,7 @@ def _run_qsa( if side_metadata.num_actual_tokens != num_tokens: raise RuntimeError("QSA main and side metadata token counts disagree") selected = self.indexer( - hidden_states, + projected_qk, positions, self.topk_indices_buffer[:num_tokens], ) @@ -401,27 +399,16 @@ def forward( key = k.view(num_tokens, self.num_kv_heads, self.head_dim) value = v.view(num_tokens, self.num_kv_heads, self.head_dim) attn_output = torch.empty_like(query) - encoded_layer_name = _encode_layer_name(self.layer_name) - if current_platform.opaque_attention_op(): - torch.ops.aphrodite.qwen4_exp_qsa_with_output( - hidden_states, - positions, - query, - key, - value, - attn_output, - encoded_layer_name, - ) - else: - qwen4_exp_qsa_with_output( - hidden_states, - positions, - query, - key, - value, - attn_output, - encoded_layer_name, - ) + # Keep the index projection outside the eager break. + projected_qk, _ = self.indexer.index_qk_proj(hidden_states) + self._run_qsa( + projected_qk, + positions, + query, + key, + value, + attn_output, + ) flat_output = attn_output.view(num_tokens, -1) if gate is not None: flat_output = flat_output * torch.sigmoid(gate) @@ -429,42 +416,9 @@ def forward( return output -def qwen4_exp_qsa_with_output( - hidden_states: torch.Tensor, - positions: torch.Tensor, - query: torch.Tensor, - key: torch.Tensor, - value: torch.Tensor, - output: torch.Tensor, - layer_name: LayerNameType, -) -> None: - """Run the complete QSA state/update/attend transaction.""" - - layer_name = _resolve_layer_name(layer_name) - layer = get_forward_context().no_compile_layers[layer_name] - if not isinstance(layer, Qwen4ExpQSAAttention): - raise TypeError(f"{layer_name} is not a Qwen4Exp QSA owner") - layer._run_qsa( - hidden_states, - positions, - query, - key, - value, - output, - ) - - -direct_register_custom_op( - op_name="qwen4_exp_qsa_with_output", - op_func=qwen4_exp_qsa_with_output, - mutates_args=["output"], -) - - __all__ = [ "QSAIndexer", "Qwen4ExpQSAAttention", "Qwen4ExpQSAFlashAttentionBackend", "Qwen4ExpQSAFlashAttentionImpl", - "qwen4_exp_qsa_with_output", ] diff --git a/tests/models/qwen4_exp/test_config.py b/tests/models/qwen4_exp/test_config.py index de0b6b3f88..3ce8f63c3e 100644 --- a/tests/models/qwen4_exp/test_config.py +++ b/tests/models/qwen4_exp/test_config.py @@ -141,9 +141,9 @@ def test_qwen4_exp_model_state_prepares_ngram_context() -> None: model_state.uses_ngram_embedding = True model_state.ngram_context_len = 3 model_state.ngram_eos_token_id = 99 - model_state.ngram_context = torch.empty((4, 3), dtype=torch.int32) + model_state.ngram_context = torch.empty((8, 3), dtype=torch.int32) model_state.ngram_context_offsets = torch.arange(-3, 0, dtype=torch.int64) - model_state.ple_query_start_loc = torch.empty(5, dtype=torch.int32) + model_state.ple_query_start_loc = torch.empty(9, dtype=torch.int32) input_batch = SimpleNamespace( num_reqs=2, @@ -159,30 +159,51 @@ def test_qwen4_exp_model_state_prepares_ngram_context() -> None: with patch.object(MambaHybridModelState, "prepare_inputs", return_value={}): model_inputs = model_state.prepare_inputs(input_batch, req_states) - torch.testing.assert_close( - model_inputs["query_start_loc"], - torch.tensor([0, 2, 3, 3], dtype=torch.int32), - ) - torch.testing.assert_close( - model_inputs["ngram_context"], - torch.tensor([[99, 99, 20], [1, 2, 3], [99, 99, 99]], dtype=torch.int32), - ) + expected_query_start_loc = torch.full((9,), 3, dtype=torch.int32) + expected_query_start_loc[0] = 0 + expected_query_start_loc[1] = 2 + torch.testing.assert_close(model_inputs["query_start_loc"], expected_query_start_loc) + expected_context = torch.full((8, 3), 99, dtype=torch.int32) + expected_context[:2] = torch.tensor([[99, 99, 20], [1, 2, 3]]) + torch.testing.assert_close(model_inputs["ngram_context"], expected_context) + + # Retain the views to detect reallocations as the request layout changes. + query_start_loc = model_inputs["query_start_loc"] + ngram_context = model_inputs["ngram_context"] + input_batch.num_reqs = 1 + input_batch.num_reqs_after_padding = 1 + input_batch.idx_mapping = torch.tensor([0]) + input_batch.query_start_loc = torch.tensor([0, 3], dtype=torch.int32) + with patch.object(MambaHybridModelState, "prepare_inputs", return_value={}): + model_inputs = model_state.prepare_inputs(input_batch, req_states) + + expected_query_start_loc.fill_(3) + expected_query_start_loc[0] = 0 + torch.testing.assert_close(model_inputs["query_start_loc"], expected_query_start_loc) + expected_context.fill_(99) + expected_context[0] = torch.tensor([1, 2, 3]) + torch.testing.assert_close(model_inputs["ngram_context"], expected_context) + assert model_inputs["query_start_loc"].data_ptr() == query_start_loc.data_ptr() + assert model_inputs["ngram_context"].data_ptr() == ngram_context.data_ptr() def test_qwen4_exp_model_state_prepares_stable_dummy_ngram_inputs() -> None: model_state = object.__new__(Qwen4ExpModelState) model_state.uses_ngram_embedding = True model_state.ngram_eos_token_id = 99 - model_state.ngram_context = torch.empty((4, 3), dtype=torch.int32) - model_state.ple_query_start_loc = torch.empty(5, dtype=torch.int32) + model_state.ngram_context = torch.empty((8, 3), dtype=torch.int32) + model_state.ple_query_start_loc = torch.empty(9, dtype=torch.int32) with patch.object(MambaHybridModelState, "prepare_dummy_inputs", return_value={}): - first = model_state.prepare_dummy_inputs(num_reqs=3, num_tokens=8) + first = model_state.prepare_dummy_inputs(num_reqs=3, num_tokens=4) + # Dummy runs establish the addresses used during CUDA graph capture. query_start_loc_ptr = first["query_start_loc"].data_ptr() ngram_context_ptr = first["ngram_context"].data_ptr() - second = model_state.prepare_dummy_inputs(num_reqs=3, num_tokens=8) + second = model_state.prepare_dummy_inputs(num_reqs=3, num_tokens=4) - torch.testing.assert_close(second["query_start_loc"], torch.tensor([0, 2, 5, 8], dtype=torch.int32)) - torch.testing.assert_close(second["ngram_context"], torch.full((3, 3), 99, dtype=torch.int32)) + expected_query_start_loc = torch.full((9,), 4, dtype=torch.int32) + expected_query_start_loc[:4] = torch.tensor([0, 1, 2, 4], dtype=torch.int32) + torch.testing.assert_close(second["query_start_loc"], expected_query_start_loc) + torch.testing.assert_close(second["ngram_context"], torch.full((8, 3), 99, dtype=torch.int32)) assert second["query_start_loc"].data_ptr() == query_start_loc_ptr assert second["ngram_context"].data_ptr() == ngram_context_ptr diff --git a/tests/models/qwen4_exp/test_ple.py b/tests/models/qwen4_exp/test_ple.py index 60e4898ea8..ca6ddbe6e8 100644 --- a/tests/models/qwen4_exp/test_ple.py +++ b/tests/models/qwen4_exp/test_ple.py @@ -480,6 +480,17 @@ def _ngram_hash_params(device: torch.device, context_len: int) -> dict: ([3], [1], [[11]], 20), ([4, 4], [0, 7], [[11, 12], [13, 14]], 20), ([4, 0, 3], [], [[11, 12], [13, 14], [15, 16]], 20), + ( + [3, 2, 0, 0], + [], + [ + [11, 12], + [13, 14], + [_NGRAM_EOS_TOKEN_ID, _NGRAM_EOS_TOKEN_ID], + [_NGRAM_EOS_TOKEN_ID, _NGRAM_EOS_TOKEN_ID], + ], + 20, + ), ( [1, 33, 2], [5, 32], @@ -510,6 +521,7 @@ def _ngram_hash_params(device: torch.device, context_len: int) -> dict: "bigram-only", "power-of-two", "empty-request", + "trailing-padded-requests", "three-requests", "six-requests", "four-gram", diff --git a/tests/models/qwen4_exp/test_qsa_reference.py b/tests/models/qwen4_exp/test_qsa_reference.py index 13f3d2558c..09bb1072e6 100644 --- a/tests/models/qwen4_exp/test_qsa_reference.py +++ b/tests/models/qwen4_exp/test_qsa_reference.py @@ -45,7 +45,6 @@ def test_qsa_mtp_index_share_updates_cache_but_skips_selection( indexer = SimpleNamespace( skip_topk=True, _metadata=lambda: (raw_metadata, compressed_metadata), - index_qk_proj=lambda hidden: (torch.zeros(2, 2), None), index_n_heads=1, index_kv_heads=1, index_head_dim=1, @@ -80,7 +79,7 @@ def test_qsa_mtp_index_share_updates_cache_but_skips_selection( actual = indexer_qsa.QSAIndexer.forward( indexer, - torch.zeros(2, 4), + torch.zeros(2, 2), torch.tensor([7, 8]), rows, ) @@ -430,6 +429,67 @@ def test_qsa_compressed_metadata_keeps_dummy_slots_inert() -> None: assert metadata.k_work_metadata.tolist() == [[0, 0], [2, 0], [2, 1], [-1, -1]] +@requires_qsa_kernels +@pytest.mark.usefixtures("default_aphrodite_config") +def test_qsa_unfused_cache_update_ignores_padded_qk() -> None: + """Padded projected Q/K rows must not affect either side cache.""" + from aphrodite.model_executor.layers.rotary_embedding import get_rope + + device = torch.device("cuda") + # Five tokens complete one compressed group and retain four keys in the ring. + raw_metadata = SimpleNamespace( + num_actual_tokens=5, + slot_mapping=torch.tensor([-1, 1, 2, 3, 0], device=device), + block_table=torch.zeros((1, 1), dtype=torch.int32, device=device), + token_to_req=torch.zeros(5, dtype=torch.int32, device=device), + query_start_loc=torch.tensor([0, 5], dtype=torch.int32, device=device), + logical_positions=torch.arange(5, device=device), + ) + compressed_metadata = SimpleNamespace( + slot_mapping=torch.tensor([-1, -1, -1, 0, -1], device=device), + ) + with torch.device(device): + rope = get_rope( + head_size=128, + max_position=32, + rope_parameters={"rope_type": "default", "partial_rotary_factor": 0.5}, + dtype=torch.bfloat16, + ) + raw_cache = torch.zeros((1, 4, 1, 64), dtype=torch.bfloat16, device=device) + compressed_cache = torch.zeros((1, 2, 1, 64), dtype=torch.bfloat16, device=device) + norm = SimpleNamespace( + weight=torch.zeros(64, dtype=torch.bfloat16, device=device), + variance_epsilon=1e-6, + ) + indexer = SimpleNamespace( + _metadata=lambda: (raw_metadata, compressed_metadata), + skip_topk=True, + index_kv_heads=1, + use_fused_pre_indexer=False, + index_n_heads=1, + index_head_dim=64, + q_layernorm=norm, + k_layernorm=norm, + rotary_emb=rope, + compress_ratio=4, + raw_key_cache=SimpleNamespace(kv_cache=raw_cache, key_cache=raw_cache, rope_position_cache=None), + compressed_key_cache=SimpleNamespace(kv_cache=compressed_cache), + ) + keys = torch.arange(1, 6, dtype=torch.bfloat16, device=device)[:, None].expand(5, 64) + padded_keys = torch.full((8, 64), torch.nan, dtype=torch.bfloat16, device=device) + padded_keys[:5].copy_(keys) + indexer_qsa.QSAIndexer.forward( + indexer, + torch.cat((torch.ones_like(padded_keys), padded_keys), dim=-1), + torch.zeros(8, dtype=torch.long, device=device), + torch.full((5, 5), -1, dtype=torch.int32, device=device), + ) + torch.testing.assert_close(raw_cache[0, :, 0], keys[[4, 1, 2, 3]]) + expected_compressed = torch.zeros_like(compressed_cache) + expected_compressed[0, 0] = 1 + torch.testing.assert_close(compressed_cache, expected_compressed) + + @requires_qsa_kernels @pytest.mark.parametrize("compress_ratio", [1, 4]) @pytest.mark.parametrize("num_reqs", [2, 3, 4, 7, 8, 9]) diff --git a/tests/test_config.py b/tests/test_config.py index fdf50b8a78..3ada71c4ef 100644 --- a/tests/test_config.py +++ b/tests/test_config.py @@ -325,9 +325,15 @@ def test_dsa_models_default_to_mrv2_and_breakable_cudagraph(monkeypatch, model, ("DeepseekV32MTPModel", True, False), ("GlmMoeDsaForCausalLM", False, True), ("GlmMoeDsaForCausalLM", True, False), + ("Qwen4ExpForCausalLM", False, True), + ("Qwen4ExpForCausalLM", True, False), + ("Qwen4ExpForConditionalGeneration", False, True), + ("Qwen4ExpForConditionalGeneration", True, False), + ("Qwen4ExpMTP", False, True), + ("Qwen4ExpMTP", True, False), ], ) -def test_dsa_breakable_cudagraph_platform_default(monkeypatch, architecture, is_rocm, expected): +def test_breakable_cudagraph_platform_default(monkeypatch, architecture, is_rocm, expected): from aphrodite.config.aphrodite import default_breakable_cudagraph_architectures from aphrodite.platforms import current_platform From b5e85de0c1bbf44b2fb322094cbabe963ff4ee2e Mon Sep 17 00:00:00 2001 From: AlpinDale Date: Mon, 7 Sep 2026 19:45:37 +0000 Subject: [PATCH 11/36] [sync] [XPU] Use fused_input_norm kernel in FusedInputNorm (#52945) Upstream-vLLM: ed29dfae6e76235e01fcc0eca21cf0dcbfc40af1 Co-authored-by: zofia <110436990+zufangzhu@users.noreply.github.com> --- .sync/vllm-sha | 2 +- aphrodite/_xpu_ops.py | 37 +++++++++++++++++++++++ aphrodite/model_executor/models/vision.py | 8 +++++ 3 files changed, 46 insertions(+), 1 deletion(-) diff --git a/.sync/vllm-sha b/.sync/vllm-sha index d789e14bd5..f3a761277a 100644 --- a/.sync/vllm-sha +++ b/.sync/vllm-sha @@ -1 +1 @@ -d9105ea8001e0a6d77a96327d17515bb5791fb36 +ed29dfae6e76235e01fcc0eca21cf0dcbfc40af1 diff --git a/aphrodite/_xpu_ops.py b/aphrodite/_xpu_ops.py index 0dfc582594..279846bba8 100644 --- a/aphrodite/_xpu_ops.py +++ b/aphrodite/_xpu_ops.py @@ -566,6 +566,37 @@ def _xpu_mxfp4_quantize_fake( return x_q, x_s +def _xpu_fused_input_norm_impl( + x: torch.Tensor, + weight: torch.Tensor | None, + bias: torch.Tensor | None, + visual_dtype: torch.dtype, +) -> torch.Tensor: + patches, size = x.shape + out = torch.empty( + (patches, size), + dtype=visual_dtype, + device=x.device, + ) + torch.ops._xpu_C.fused_input_norm(out, x.contiguous(), weight, bias) + return out + + +def _xpu_fused_input_norm_fake( + x: torch.Tensor, + weight: torch.Tensor | None, + bias: torch.Tensor | None, + visual_dtype: torch.dtype, +) -> torch.Tensor: + patches, size = x.shape + out = torch.empty( + (patches, size), + dtype=visual_dtype, + device=x.device, + ) + return out + + @triton.jit def _softplus(x): return tl.where(x <= 20.0, tl.math.log(tl.math.exp(x) + 1.0), x) @@ -1233,6 +1264,12 @@ def register_ops_once() -> None: ], ) + direct_register_custom_op( + op_name="xpu_fused_input_norm", + op_func=_xpu_fused_input_norm_impl, + fake_impl=_xpu_fused_input_norm_fake, + ) + _OPS_REGISTERED = True diff --git a/aphrodite/model_executor/models/vision.py b/aphrodite/model_executor/models/vision.py index c6286e715e..a5176d9e8d 100644 --- a/aphrodite/model_executor/models/vision.py +++ b/aphrodite/model_executor/models/vision.py @@ -664,6 +664,14 @@ def forward( patches, size = grid_thw.shape patch_size = size // self.channel + # On XPU, fuse the whole rescale + normalize into a single custom + # kernel. The eager path below materializes an fp32 intermediate and + # then casts back, which adds device-side compute that cancels the + # bandwidth saving of transferring uint8 pixel_values. The fused + # kernel reads uint8 directly and writes ``visual_dtype`` in one pass. + if current_platform.is_xpu() and grid_thw.dtype == torch.uint8 and self.weight.dtype == torch.float32: + return torch.ops.aphrodite.xpu_fused_input_norm(grid_thw, self.weight, self.bias, visual_dtype) + # Apply the per-channel affine transform directly instead of via # F.batch_norm. batch_norm dispatches to cuDNN, whose batch-norm # kernels cap the batch dimension near the CUDA grid limit (~65535); From 7f5df6b6ea0a4e520f592fd4eef8f3e08f706119 Mon Sep 17 00:00:00 2001 From: AlpinDale Date: Mon, 7 Sep 2026 19:49:00 +0000 Subject: [PATCH 12/36] [sync] [Bugfix][Audio] Restore soundfile-first automatic decoding (#55642) Upstream-vLLM: f7f060d253fb67c9b83f45a22a1eeeba4d72a2be Co-authored-by: Andreas Karatzas --- .sync/vllm-sha | 2 +- aphrodite/multimodal/media/audio.py | 41 ++++++++++--------- docs/src/content/docs/features/multimodal.md | 32 +++++++++++++++ tests/multimodal/media/test_audio.py | 43 ++++++++++++++++++++ 4 files changed, 97 insertions(+), 21 deletions(-) diff --git a/.sync/vllm-sha b/.sync/vllm-sha index f3a761277a..757a5971a6 100644 --- a/.sync/vllm-sha +++ b/.sync/vllm-sha @@ -1 +1 @@ -ed29dfae6e76235e01fcc0eca21cf0dcbfc40af1 +f7f060d253fb67c9b83f45a22a1eeeba4d72a2be diff --git a/aphrodite/multimodal/media/audio.py b/aphrodite/multimodal/media/audio.py index a927f0700f..d4079f3efe 100644 --- a/aphrodite/multimodal/media/audio.py +++ b/aphrodite/multimodal/media/audio.py @@ -60,11 +60,11 @@ # Audio decoding backends selectable via `load_audio(backend=...)` or # `--media-io-kwargs '{"audio": {"audio_backend": ...}}'`. -# "auto" tries torchcodec, then soundfile, then PyAV. +# "auto" tries soundfile, then torchcodec, then PyAV. AUDIO_BACKENDS = ("auto", "soundfile", "pyav", "torchcodec") # Raised as ImportError when the torchcodec backend cannot be used, so -# `load_audio(backend="auto")` falls back to the soundfile → PyAV chain. +# `load_audio(backend="auto")` can continue to PyAV. _TORCHCODEC_UNAVAILABLE_MSG = ( "torchcodec audio backend is unavailable (requires the torchcodec package and a system ffmpeg installation)" ) @@ -395,8 +395,8 @@ def load_audio( Args: backend: One of ``AUDIO_BACKENDS``. ``None`` (default) selects - ``"auto"``, which tries torchcodec, then falls back to the - soundfile → PyAV chain; the other values select a single + ``"auto"``, which tries soundfile, then torchcodec, then PyAV; + the other values select a single backend with no fallback. """ backend = backend or "auto" @@ -412,25 +412,11 @@ def load_audio( max_decode_bytes=max_decode_bytes, ) - # "auto": torchcodec → soundfile → PyAV. Each backend is called by its + # Keep soundfile first to preserve decoding and padding of supported formats. + # "auto": soundfile → torchcodec → PyAV. Each backend is called by its # bare name so tests that monkeypatch it take effect. Only an ImportError # (backend missing) or, for soundfile, an unsupported-format error defers # to the next; any other error (decode failure, guard ValueError) raises. - if isinstance(path, BytesIO): - path.seek(0) - try: - return load_audio_torchcodec( - path, - sr=sr, - mono=mono, - max_duration_s=max_duration_s, - max_decode_bytes=max_decode_bytes, - ) - except ImportError as exc: - # Decode errors don't defer: torchcodec is FFmpeg-based like PyAV, so - # retrying would just fail the same way. - logger.warning("torchcodec unavailable (%r); falling back to soundfile/PyAV.", exc) - if isinstance(path, BytesIO): path.seek(0) try: @@ -454,6 +440,21 @@ def load_audio( if exc.code not in _BAD_SF_CODES: raise + if isinstance(path, BytesIO): + path.seek(0) + try: + return load_audio_torchcodec( + path, + sr=sr, + mono=mono, + max_duration_s=max_duration_s, + max_decode_bytes=max_decode_bytes, + ) + except ImportError as exc: + # Decode errors don't defer: torchcodec is FFmpeg-based like PyAV, so + # retrying would just fail the same way. + logger.warning("torchcodec unavailable (%r); falling back to PyAV.", exc) + # PyAV is terminal: nothing left to fall back to. Normalize an FFmpeg # failure to ValueError; let a missing soundfile/PyAV surface its # "install aphrodite[audio]" ImportError. diff --git a/docs/src/content/docs/features/multimodal.md b/docs/src/content/docs/features/multimodal.md index 0bd034188c..6d5c6ddc4c 100644 --- a/docs/src/content/docs/features/multimodal.md +++ b/docs/src/content/docs/features/multimodal.md @@ -124,6 +124,38 @@ Aphrodite divides the configured budget evenly among them. For streaming video sources, use the `deepstream` backend instead. +## Configure audio decoding + +Aphrodite decodes audio bytes into waveforms using a selectable decoding backend. + +| Backend | Description | +| --- | --- | +| `auto` (default) | soundfile, falling back to torchcodec, then PyAV | +| `soundfile` | libsndfile only, no fallback | +| `pyav` | PyAV (FFmpeg) only, no fallback | +| `torchcodec` | TorchCodec (PyTorch-native) only, no fallback | + +Select the decoder with `--media-io-kwargs`: + +```bash +aphrodite serve mistralai/Voxtral-Mini-3B-2507 \ + --media-io-kwargs '{"audio": {"audio_backend": "torchcodec"}}' +``` + +PyAV drives FFmpeg through a per-frame Python generator, so concurrent decoding +can contend on the GIL. TorchCodec decodes each stream in a single call that +releases the GIL for its duration. Select it explicitly for concurrent decoding +workloads. + +`auto` prefers soundfile for supported formats to preserve their existing +decoding behavior, including encoder padding. Audio extracted from formats +that soundfile cannot read, such as video containers, can use torchcodec +through the fallback chain. + +TorchCodec requires both the package and a system FFmpeg installation. Install +it manually if it is not included on your platform. If the package or FFmpeg +is unavailable, `auto` uses the soundfile → PyAV chain. + ## Benchmark media Text-only tests do not measure encoder or download cost. Use diff --git a/tests/multimodal/media/test_audio.py b/tests/multimodal/media/test_audio.py index c2bc001387..116ed7efbc 100644 --- a/tests/multimodal/media/test_audio.py +++ b/tests/multimodal/media/test_audio.py @@ -193,6 +193,49 @@ def test_load_audio_backend_matches_default(backend, dummy_audio_bytes): np.testing.assert_allclose(ref_audio[:n], audio[:n], atol=1e-4) +def test_load_audio_default_preserves_vorbis_length(dummy_audio_bytes): + """Default decoding must not append Vorbis padding to the speech waveform.""" + expected, expected_sr = load_audio_soundfile(BytesIO(dummy_audio_bytes), sr=None) + audio, sr = load_audio(BytesIO(dummy_audio_bytes), sr=None) + assert sr == expected_sr + assert audio.shape == expected.shape + + +@pytest.mark.parametrize("torchcodec_available", [True, False]) +@pytest.mark.parametrize("soundfile_error", [ImportError(), sf.LibsndfileError(1)]) +def test_load_audio_auto_fallback_order(monkeypatch, torchcodec_available, soundfile_error): + """Fallbacks rewind the input and preserve all decoding limits.""" + calls = [] + source = BytesIO(b"audio") + source.seek(2) + expected = (np.zeros(4, dtype=np.float32), 16000) + options = dict(sr=16000, mono=False, max_duration_s=3, max_decode_bytes=64) + + def decoder(name, error=None): + def decode(path, **kwargs): + calls.append(name) + assert path is source + assert path.tell() == 0 + assert kwargs == options + path.read() + if error is not None: + raise error + return expected + + return decode + + monkeypatch.setattr(audio_module, "load_audio_soundfile", decoder("soundfile", soundfile_error)) + monkeypatch.setattr( + audio_module, + "load_audio_torchcodec", + decoder("torchcodec", None if torchcodec_available else ImportError()), + ) + monkeypatch.setattr(audio_module, "load_audio_pyav", decoder("pyav")) + + assert load_audio(source, **options) is expected + assert calls == (["soundfile", "torchcodec"] if torchcodec_available else ["soundfile", "torchcodec", "pyav"]) + + def test_load_audio_unknown_backend_rejected(dummy_audio_bytes): """An unknown backend must fail loudly instead of silently degrading.""" with pytest.raises(ValueError, match="Unknown audio backend"): From 4b0668c01d5668626b10e7fee63d35fa7873246e Mon Sep 17 00:00:00 2001 From: AlpinDale Date: Mon, 7 Sep 2026 19:56:34 +0000 Subject: [PATCH 13/36] [sync] [EPD] Add ECMooncakeConnector for encoder cache over Mooncake TransferEngine (#41567) Upstream-vLLM: 6fbb00b18874e27ba7d7adc0a3b8e93fee763ab1 Co-authored-by: Teng Ma --- .sync/vllm-sha | 2 +- aphrodite/config/ec_transfer.py | 5 + .../ec_transfer/ec_connector/base.py | 19 + .../ec_connector/cpu/scheduler/__init__.py | 8 +- .../ec_transfer/ec_connector/factory.py | 6 + .../ec_connector/mooncake/__init__.py | 9 + .../ec_connector/mooncake/config.py | 110 ++ .../ec_connector/mooncake/control.py | 432 ++++++++ .../ec_connector/mooncake/memory.py | 484 +++++++++ .../ec_connector/mooncake/metadata.py | 75 ++ .../ec_connector/mooncake/producer.py | 393 +++++++ .../ec_connector/mooncake/reservation.py | 442 ++++++++ .../ec_connector/mooncake/scheduler.py | 510 +++++++++ .../ec_connector/mooncake/state.py | 428 ++++++++ .../ec_connector/mooncake/transfer.py | 198 ++++ .../ec_connector/mooncake/worker.py | 866 +++++++++++++++ .../ec_connector/mooncake_ec_connector.py | 120 +++ aphrodite/v1/core/sched/scheduler.py | 47 +- .../worker/ec_connector_model_runner_mixin.py | 8 + aphrodite/v1/worker/gpu/ec_connector.py | 3 + aphrodite/v1/worker/gpu_worker.py | 5 + .../disaggregated_encoder/disagg_epd_proxy.py | 990 ++++++++++++++++++ tests/v1/core/test_encoder_cache_manager.py | 32 + tests/v1/core/test_scheduler.py | 90 +- tests/v1/ec_connector/integration/README.md | 43 +- .../run_epd_mooncake_ec_full_pipeline.sh | 289 +++++ .../integration/test_epd_correctness.py | 103 +- .../ec_connector/unit/test_epd_proxy_retry.py | 164 +++ .../unit/test_epd_proxy_round_robin.py | 50 +- .../unit/test_worker_ec_connector.py | 24 + 30 files changed, 5920 insertions(+), 35 deletions(-) create mode 100644 aphrodite/distributed/ec_transfer/ec_connector/mooncake/__init__.py create mode 100644 aphrodite/distributed/ec_transfer/ec_connector/mooncake/config.py create mode 100644 aphrodite/distributed/ec_transfer/ec_connector/mooncake/control.py create mode 100644 aphrodite/distributed/ec_transfer/ec_connector/mooncake/memory.py create mode 100644 aphrodite/distributed/ec_transfer/ec_connector/mooncake/metadata.py create mode 100644 aphrodite/distributed/ec_transfer/ec_connector/mooncake/producer.py create mode 100644 aphrodite/distributed/ec_transfer/ec_connector/mooncake/reservation.py create mode 100644 aphrodite/distributed/ec_transfer/ec_connector/mooncake/scheduler.py create mode 100644 aphrodite/distributed/ec_transfer/ec_connector/mooncake/state.py create mode 100644 aphrodite/distributed/ec_transfer/ec_connector/mooncake/transfer.py create mode 100644 aphrodite/distributed/ec_transfer/ec_connector/mooncake/worker.py create mode 100644 aphrodite/distributed/ec_transfer/ec_connector/mooncake_ec_connector.py create mode 100644 examples/disaggregated/disaggregated_encoder/disagg_epd_proxy.py create mode 100755 tests/v1/ec_connector/integration/run_epd_mooncake_ec_full_pipeline.sh create mode 100644 tests/v1/ec_connector/unit/test_epd_proxy_retry.py diff --git a/.sync/vllm-sha b/.sync/vllm-sha index 757a5971a6..223e536eb2 100644 --- a/.sync/vllm-sha +++ b/.sync/vllm-sha @@ -1 +1 @@ -f7f060d253fb67c9b83f45a22a1eeeba4d72a2be +6fbb00b18874e27ba7d7adc0a3b8e93fee763ab1 diff --git a/aphrodite/config/ec_transfer.py b/aphrodite/config/ec_transfer.py index 2382b804c6..cdf84bae79 100644 --- a/aphrodite/config/ec_transfer.py +++ b/aphrodite/config/ec_transfer.py @@ -18,6 +18,11 @@ class ECTransferConfig: ec_connector: str | None = None """The EC connector for Aphrodite to transmit EC caches between Aphrodite instances. + + Built-in options include ``ECExampleConnector`` (shared filesystem via + safetensors) and ``ECMooncakeConnector`` (Mooncake TransferEngine RDMA; + requires ``mooncake-transfer-engine`` and matching producer/consumer + ``ec_connector_extra_config``; see ``mooncake_ec_connector`` module docstring). """ engine_id: str | None = None diff --git a/aphrodite/distributed/ec_transfer/ec_connector/base.py b/aphrodite/distributed/ec_transfer/ec_connector/base.py index e0ccd20208..da05496fc1 100644 --- a/aphrodite/distributed/ec_transfer/ec_connector/base.py +++ b/aphrodite/distributed/ec_transfer/ec_connector/base.py @@ -251,6 +251,25 @@ def ensure_cache_available(self, request: "Request", num_computed_tokens: int) - """ return True + def start_save_caches(self, **kwargs: Any) -> None: + """Prepare this step's outbound pushes before the model runs.""" + return + + def start_worker_services(self) -> None: + """Start Worker-side services once the model is resident.""" + return + + def take_unavailable_requests(self) -> set[str]: + """Request IDs whose encoder inputs the connector can no longer obtain. + + The Scheduler fails these; re-issuing the request re-runs the encode. + """ + return set() + + def update_state_after_free(self, request: "Request", index: int): + """Notify the connector that an encoder cache entry was released.""" + return + @abstractmethod def update_state_after_alloc(self, request: "Request", index: int): """ diff --git a/aphrodite/distributed/ec_transfer/ec_connector/cpu/scheduler/__init__.py b/aphrodite/distributed/ec_transfer/ec_connector/cpu/scheduler/__init__.py index 71d69260ab..9aca0b646f 100644 --- a/aphrodite/distributed/ec_transfer/ec_connector/cpu/scheduler/__init__.py +++ b/aphrodite/distributed/ec_transfer/ec_connector/cpu/scheduler/__init__.py @@ -7,6 +7,7 @@ for the ECCPUConnector. """ +from collections.abc import Collection from typing import TYPE_CHECKING, Any import torch @@ -174,7 +175,12 @@ def has_cache_item(self, identifier: str) -> bool: entry = self._cache.get(identifier) return entry is not None and entry.ready - def ensure_cache_available(self, request: "Request", num_computed_tokens: int) -> bool: + def ensure_cache_available( + self, + request: "Request", + num_computed_tokens: int, + local_cache_hashes: Collection[str] | None = None, + ) -> bool: if not self._nixl_enabled: return True # CPU offload never blocks. first = self._first_in_batch diff --git a/aphrodite/distributed/ec_transfer/ec_connector/factory.py b/aphrodite/distributed/ec_transfer/ec_connector/factory.py index c060c1701c..20ff0f3661 100644 --- a/aphrodite/distributed/ec_transfer/ec_connector/factory.py +++ b/aphrodite/distributed/ec_transfer/ec_connector/factory.py @@ -87,3 +87,9 @@ def get_connector_class(cls, ec_transfer_config: "ECTransferConfig") -> type[ECC "aphrodite.distributed.ec_transfer.ec_connector.cpu.connector", "ECCPUConnector", ) + +ECConnectorFactory.register_connector( + "ECMooncakeConnector", + "aphrodite.distributed.ec_transfer.ec_connector.mooncake_ec_connector", + "ECMooncakeConnector", +) diff --git a/aphrodite/distributed/ec_transfer/ec_connector/mooncake/__init__.py b/aphrodite/distributed/ec_transfer/ec_connector/mooncake/__init__.py new file mode 100644 index 0000000000..5cd38b5017 --- /dev/null +++ b/aphrodite/distributed/ec_transfer/ec_connector/mooncake/__init__.py @@ -0,0 +1,9 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +"""Internal building blocks for the Mooncake encoder-cache connector. + +The public connector remains in :mod:`mooncake_ec_connector`. This package +separates configuration, control-plane messaging, transfer state, registered +memory, and worker/scheduler orchestration so that each component has one +owner and can be tested independently. +""" diff --git a/aphrodite/distributed/ec_transfer/ec_connector/mooncake/config.py b/aphrodite/distributed/ec_transfer/ec_connector/mooncake/config.py new file mode 100644 index 0000000000..290fbc00ca --- /dev/null +++ b/aphrodite/distributed/ec_transfer/ec_connector/mooncake/config.py @@ -0,0 +1,110 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +from __future__ import annotations + +import math +from dataclasses import dataclass +from typing import TYPE_CHECKING + +from aphrodite.utils.network_utils import make_zmq_path + +if TYPE_CHECKING: + from aphrodite.config import AphroditeConfig + +# Scheduler and Worker must agree on when a remote reservation becomes stale. +_RESERVATION_TTL_SECONDS = 300 + + +def _positive_int(name: str, value: object) -> int: + message = f"ECMooncakeConnector requires {name} to be a positive integer." + if isinstance(value, bool): + raise ValueError(message) + if isinstance(value, int): + result = value + elif isinstance(value, float) and math.isfinite(value) and value.is_integer(): + result = int(value) + elif isinstance(value, str): + try: + result = int(value) + except ValueError as error: + raise ValueError(message) from error + else: + raise ValueError(message) + if result <= 0: + raise ValueError(f"ECMooncakeConnector requires {name} > 0.") + return result + + +def _positive_float(name: str, value: object) -> float: + message = f"ECMooncakeConnector requires {name} > 0." + if isinstance(value, bool): + raise ValueError(message) + try: + result = float(value) # type: ignore[arg-type] + except (TypeError, ValueError) as error: + raise ValueError(message) from error + if not math.isfinite(result) or result <= 0: + raise ValueError(message) + return result + + +@dataclass(frozen=True) +class MooncakeECConfig: + """Validated settings shared by the Scheduler and Worker roles. + + ``control_port`` is the first TP-shard port after the DP offset; + ``control_addr`` targets that shard, which advertises the full topology. + ``control_host`` is the address this instance both advertises and binds, + so reaching a Consumer means being routed to it rather than finding it on + every interface. + """ + + is_producer: bool + is_consumer: bool + protocol: str + buffer_device: str + control_host: str + control_port: int + control_addr: str + control_timeout_ms: int + push_wait_timeout_s: float + pool_size: int + + @classmethod + def from_aphrodite_config(cls, aphrodite_config: AphroditeConfig) -> MooncakeECConfig: + parallel_config = aphrodite_config.parallel_config + ec_config = aphrodite_config.ec_transfer_config + assert ec_config is not None + + if ec_config.is_ec_producer: + if parallel_config.tensor_parallel_size > 1: + raise ValueError("ECMooncakeConnector producers require tensor_parallel_size=1.") + if parallel_config.pipeline_parallel_size > 1: + raise ValueError("ECMooncakeConnector producers do not support pipeline parallelism.") + if parallel_config.data_parallel_size > 1: + raise ValueError("ECMooncakeConnector producers require data_parallel_size=1.") + + registered_buffer_size = _positive_int("ec_buffer_size", ec_config.ec_buffer_size) + get = ec_config.get_from_extra_config + control_port = int(ec_config.ec_port) + ( + parallel_config.data_parallel_index * parallel_config.tensor_parallel_size + ) + highest_port = control_port + parallel_config.tensor_parallel_size - 1 + if not 1 <= control_port <= highest_port <= 65535: + raise ValueError("ECMooncakeConnector ec_port must be in 1..65535.") + + return cls( + is_producer=ec_config.is_ec_producer, + is_consumer=ec_config.is_ec_consumer, + protocol=str(get("mooncake_protocol", "rdma")), + buffer_device=str(ec_config.ec_buffer_device or "cuda").lower(), + control_host=str(ec_config.ec_ip), + control_port=control_port, + control_addr=make_zmq_path("tcp", ec_config.ec_ip, control_port), + control_timeout_ms=max( + 1, + math.ceil(_positive_float("control_timeout_s", get("control_timeout_s", 30)) * 1000), + ), + push_wait_timeout_s=_positive_float("push_wait_timeout_s", get("push_wait_timeout_s", 60)), + pool_size=registered_buffer_size, + ) diff --git a/aphrodite/distributed/ec_transfer/ec_connector/mooncake/control.py b/aphrodite/distributed/ec_transfer/ec_connector/mooncake/control.py new file mode 100644 index 0000000000..eaf9a2f693 --- /dev/null +++ b/aphrodite/distributed/ec_transfer/ec_connector/mooncake/control.py @@ -0,0 +1,432 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +"""ZMQ control plane for Mooncake encoder-cache reservations and events. + +Tensor bytes never travel through this module. It coordinates destination +reservations, completion/cancellation, TP-shard discovery, and readiness +notifications while :mod:`transfer` owns the Mooncake data plane. +""" + +from __future__ import annotations + +import threading +import time +from collections import deque +from collections.abc import Callable +from concurrent.futures import Executor, Future +from typing import Any + +import torch +import zmq + +from aphrodite.logger import init_logger +from aphrodite.utils.network_utils import make_zmq_path + +logger = init_logger(__name__) + +_MAX_PENDING_EVENTS = 4096 +_RESERVATION_REAP_INTERVAL_SECONDS = 1 + + +ControlRequest = dict[str, Any] +ControlResponse = dict[str, Any] + + +class ControlClient: + """Send control requests through reusable, thread-local REQ sockets.""" + + def __init__(self, timeout_ms: int) -> None: + self._context = zmq.Context() + self._timeout_ms = timeout_ms + self._local = threading.local() + self._topologies: dict[str, list[str]] = {} + self._closed = False + + def _sockets(self) -> dict[str, zmq.Socket]: + sockets = getattr(self._local, "sockets", None) + if sockets is None: + sockets = {} + self._local.sockets = sockets + return sockets + + def _discard(self, addr: str) -> None: + socket = self._sockets().pop(addr, None) + if socket is not None: + socket.close(linger=0) + + def _exchange(self, addr: str, payload: ControlRequest) -> ControlResponse: + sockets = self._sockets() + socket = sockets.get(addr) + if socket is None: + socket = self._context.socket(zmq.REQ) + socket.setsockopt(zmq.IPV6, 1) + socket.setsockopt(zmq.RCVTIMEO, self._timeout_ms) + socket.setsockopt(zmq.SNDTIMEO, self._timeout_ms) + socket.setsockopt(zmq.LINGER, 0) + socket.connect(addr) + sockets[addr] = socket + try: + socket.send_json(payload) + response = socket.recv_json() + except Exception: + self._discard(addr) + raise + assert isinstance(response, dict) + return response + + def request(self, addr: str, payload: ControlRequest) -> Any: + response = self._exchange(addr, payload) + if not response.get("ok"): + raise RuntimeError(response.get("error", "EC control request failed")) + return response.get("result") + + def discover_shards(self, base_addr: str) -> list[str] | None: + if base_addr in self._topologies: + return self._topologies[base_addr] + try: + reply = self.request(base_addr, {"op": "peers"}) + ports = reply.get("ports") if isinstance(reply, dict) else None + if not isinstance(ports, list) or not ports: + raise ValueError("invalid or empty peer list") + prefix = base_addr.rsplit(":", 1)[0] + shards = [f"{prefix}:{int(port)}" for port in ports] + except Exception: + logger.warning( + "EC Mooncake consumer at %s did not report its shards", + base_addr, + exc_info=True, + ) + return None + self._topologies[base_addr] = shards + return shards + + def close(self) -> None: + if self._closed: + return + self._closed = True + self._context.destroy(linger=0) + + +def make_cancel_request( + transfer_id: str, + reservation_id: str = "", + *, + abandon: bool = False, + refresh: bool = False, +) -> ControlRequest: + request: ControlRequest = { + "op": "cancel", + "transfer_id": transfer_id, + "reservation_id": reservation_id, + } + if abandon: + request["abandon"] = True + if refresh: + request["refresh"] = True + return request + + +class EventInbox: + """Receive Consumer readiness events without blocking the Scheduler. + + Subscribing needs one blocking control request per shard, so it runs on + `executor` when one is given: an unreachable Consumer must cost the + Scheduler nothing, not one control timeout per drain. + """ + + def __init__(self, client: ControlClient, executor: Executor | None = None) -> None: + self._client = client + self._executor = executor + self._discovery: Future[list[str] | None] | None = None + self._context: zmq.Context | None = None + self._socket: zmq.Socket | None = None + self._closed = False + self.shard_count = 1 + + def _endpoints(self, base_addr: str) -> list[str] | None: + """Ask every shard where it publishes readiness events. + + Blocking: called on `self._executor` unless there is none. + """ + shards = self._client.discover_shards(base_addr) + if shards is None: + return None + endpoints = [] + for addr in shards: + try: + event_port = self._client.request(addr, {"op": "event_port"}) + address, _ = addr.rsplit(":", 1) + endpoints.append(f"{address}:{int(event_port)}") + except Exception: + logger.warning( + "EC Mooncake could not subscribe to the event channel of " + "consumer shard %s; retrying the complete topology later.", + addr, + ) + return None + return endpoints + + def _connect(self, base_addr: str) -> None: + if self._socket is not None or self._closed: + return + if self._executor is None: + self._install(self._endpoints(base_addr)) + return + if self._discovery is None: + self._discovery = self._executor.submit(self._endpoints, base_addr) + return + if not self._discovery.done(): + return + discovery, self._discovery = self._discovery, None + try: + self._install(discovery.result()) + except Exception: + logger.warning( + "EC Mooncake event-channel discovery for %s failed; retrying.", + base_addr, + exc_info=True, + ) + + def _install(self, endpoints: list[str] | None) -> None: + """Adopt a discovered topology. + + Runs on the caller's thread so the socket is only ever touched there. + """ + if not endpoints or self._closed: + return + context = zmq.Context() + socket = context.socket(zmq.PULL) + socket.setsockopt(zmq.IPV6, 1) + for endpoint in endpoints: + socket.connect(endpoint) + self._context = context + self._socket = socket + self.shard_count = len(endpoints) + + def drain(self, base_addr: str) -> list[dict[str, Any]]: + self._connect(base_addr) + if self._socket is None: + return [] + events = [] + while True: + try: + events.append(self._socket.recv_json(flags=zmq.DONTWAIT)) + except zmq.Again: + return events + except zmq.ZMQError: + logger.warning("EC Mooncake event channel failed", exc_info=True) + return events + except (ValueError, UnicodeError): + # The event channel is a plain PULL socket: an undecodable + # frame must cost one frame, not the engine. + logger.warning("Discarding an undecodable EC Mooncake event.", exc_info=True) + + def close(self) -> None: + if self._closed: + return + self._closed = True + if self._discovery is not None: + self._discovery.cancel() + self._discovery = None + if self._socket is not None: + self._socket.close(linger=0) + self._socket = None + if self._context is not None: + self._context.term() + + +class ConsumerControlServer: + """Expose Consumer reservations and readiness events over ZMQ. + + One server runs on every receiving TP rank. The REP channel handles + reservation operations, while a PUSH channel publishes newly ready items + to the Scheduler. + """ + + def __init__( + self, + host: str, + port: int, + reserve: Callable[[dict[str, Any]], dict[str, Any]], + status: Callable[[str], dict[str, Any] | None], + complete: Callable[[str, str], tuple[bool, bool]], + cancel: Callable[[str, str, bool, bool], bool], + reap: Callable[[], int], + peer_ports: list[int] | None = None, + device: torch.device | None = None, + drain_ready: Callable[[], list[str]] = lambda: [], + ) -> None: + self.host = host + self.port = port + self.peer_ports = peer_ports or [port] + self._device = device + self.event_port: int | None = None + self._reserve = reserve + self._status = status + self._complete = complete + self._cancel = cancel + self._reap = reap + self._drain_ready = drain_ready + self._stop = threading.Event() + self._started = threading.Event() + self._thread: threading.Thread | None = None + self._startup_error: Exception | None = None + + def start(self) -> None: + def loop() -> None: + if self._device is not None and self._device.type == "cuda": + torch.accelerator.set_device_index(self._device.index or 0) + context = zmq.Context() + socket = context.socket(zmq.REP) + event_socket = context.socket(zmq.PUSH) + socket.setsockopt(zmq.IPV6, 1) + event_socket.setsockopt(zmq.IPV6, 1) + pending_events: deque[dict[str, Any]] = deque() + + def queue_event(event: dict[str, Any]) -> None: + event["shard"] = self.port + if len(pending_events) >= _MAX_PENDING_EVENTS: + dropped = pending_events.popleft() + logger.warning( + "EC Mooncake event backlog full on port %d; dropping readiness for transfer_id=%s", + self.port, + dropped.get("transfer_id"), + ) + pending_events.append(event) + + def queue_ready(transfer_id: str) -> None: + status = self._status(transfer_id) + if status is not None: + queue_event({"transfer_id": transfer_id, **status}) + + last_reap_at = time.monotonic() + socket.setsockopt(zmq.RCVTIMEO, 100) + try: + socket.bind(make_zmq_path("tcp", self.host, self.port)) + event_socket.bind(make_zmq_path("tcp", self.host, 0)) + self.event_port = int(event_socket.getsockopt(zmq.LAST_ENDPOINT).decode().rsplit(":", 1)[1]) + except Exception as e: + self._startup_error = e + self._started.set() + socket.close(linger=0) + event_socket.close(linger=0) + context.term() + return + self._started.set() + try: + while not self._stop.is_set(): + while pending_events: + try: + event_socket.send_json(pending_events[0], flags=zmq.DONTWAIT) + except zmq.Again: + break + pending_events.popleft() + now = time.monotonic() + if now - last_reap_at >= _RESERVATION_REAP_INTERVAL_SECONDS: + self._reap() + last_reap_at = now + try: + request = socket.recv_json() + except zmq.Again: + continue + except Exception: + # The frame arrived but did not decode. REP still owes a + # reply, so answer before returning to the loop. + logger.exception( + "EC Mooncake control channel on port %d received an undecodable request", + self.port, + ) + socket.send_json({"ok": False, "error": "malformed control request"}) + continue + try: + op = request.get("op") + result: Any = None + if op in ("reserve", "reserve_batch"): + items = request["items"] if op == "reserve_batch" else [request] + results = [] + for item in items: + try: + reserved = self._reserve(item) + results.append({"ok": True, "result": reserved}) + if reserved.get("ready"): + queue_ready(str(item["transfer_id"])) + except Exception as exc: + results.append({"ok": False, "error": str(exc)}) + if op == "reserve_batch": + result = {"items": results} + elif not results[0]["ok"]: + raise RuntimeError(results[0]["error"]) + else: + result = results[0]["result"] + elif op == "status": + result = self._status(str(request["transfer_id"])) + elif op == "event_port": + result = self.event_port + elif op == "peers": + result = {"ports": self.peer_ports} + elif op in ("complete", "complete_batch"): + items = request["items"] if op == "complete_batch" else [request] + completions = [] + for item in items: + transfer_id = str(item["transfer_id"]) + accepted, became_ready = self._complete(transfer_id, str(item["reservation_id"])) + for ready_id in self._drain_ready(): + queue_ready(ready_id) + completions.append( + { + "completed": accepted, + "became_ready": became_ready, + } + ) + if not became_ready: + continue + queue_ready(transfer_id) + result = {"items": completions} if op == "complete_batch" else completions[0] + elif op == "cancel": + result = { + "cancelled": self._cancel( + str(request["transfer_id"]), + str(request.get("reservation_id", "")), + bool(request.get("abandon", False)), + bool(request.get("refresh", False)), + ) + } + else: + raise ValueError(f"unknown control op: {op!r}") + socket.send_json({"ok": True, "result": result}) + except Exception as e: + socket.send_json({"ok": False, "error": str(e)}) + except Exception: + # `finally` closes the sockets and the thread ends, so without + # this every later reserve against this shard would surface + # only as a control timeout with nothing to attribute it to. + logger.exception( + "EC Mooncake control channel on port %d stopped serving", + self.port, + ) + raise + finally: + socket.close(linger=0) + event_socket.close(linger=0) + context.term() + + self._thread = threading.Thread(target=loop, name="ec-mooncake-control", daemon=True) + self._thread.start() + if not self._started.wait(timeout=5): + raise RuntimeError("EC Mooncake control channel failed to start") + if self._startup_error is not None: + raise RuntimeError("EC Mooncake control channel failed to bind") from (self._startup_error) + logger.info( + "EC Mooncake control channel listening on tcp://%s:%d (events tcp://%s:%d)", + self.host, + self.port, + self.host, + self.event_port, + ) + + def close(self) -> None: + if self._thread is None: + return + self._stop.set() + self._thread.join() + self._thread = None diff --git a/aphrodite/distributed/ec_transfer/ec_connector/mooncake/memory.py b/aphrodite/distributed/ec_transfer/ec_connector/mooncake/memory.py new file mode 100644 index 0000000000..d5dbe4bc3c --- /dev/null +++ b/aphrodite/distributed/ec_transfer/ec_connector/mooncake/memory.py @@ -0,0 +1,484 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +"""Registered-memory allocation and residency for Mooncake transfers. + +The Producer pool stages source tensors in a registered slab. The Consumer +pool owns destination allocations from reservation through publication, +resident reuse, CUDA-safe retirement, and pressure-driven reclamation. +""" + +from __future__ import annotations + +import bisect +import threading +from collections import OrderedDict +from collections.abc import Callable +from dataclasses import dataclass +from typing import Generic, TypeVar + +import torch + +from aphrodite.distributed.ec_transfer.ec_connector.mooncake.transfer import ( + MooncakeTransfer, +) +from aphrodite.logger import init_logger + +logger = init_logger(__name__) + +_T = TypeVar("_T") + + +@dataclass +class MemoryAllocation: + """Describe a tensor view carved from the Consumer receive slab.""" + + offset: int + size: int + tensor: torch.Tensor + + +@dataclass +class _ResidentEntry(Generic[_T]): + """Store one resident value and its ownership accounting.""" + + value: _T + pinned: bool = True + leases: int = 0 + + +@dataclass +class ResidentLease(Generic[_T]): + """Represent a borrow of one resident entry.""" + + key: str + _entry: _ResidentEntry[_T] + _active: bool = True + + @property + def value(self) -> _T: + return self._entry.value + + +class ContiguousAllocator: + """Own a registered slab and its aligned regions; callers serialize access.""" + + def __init__(self, capacity: int, alignment: int = 256): + self.alignment = alignment + self._capacity = capacity + self._free = [(0, capacity)] + self.tensor: torch.Tensor | None = None + self._disabled = False + + def prepare(self, device: torch.device, transfer: MooncakeTransfer) -> None: + if self.tensor is not None or self._disabled: + return + try: + tensor = torch.empty(self._capacity, dtype=torch.uint8, device=device) + ret = transfer.register_memory(tensor) + if ret != 0: + raise RuntimeError(f"Mooncake returned {ret}") + except (RuntimeError, torch.OutOfMemoryError) as error: + self._disabled = True + logger.warning("Could not initialize the EC registered buffer: %s", error) + return + self.tensor = tensor + self._free = [(0, tensor.nbytes)] + logger.info("Registered %d-byte buffer for Mooncake EC", tensor.nbytes) + + def view(self, offset: int, nbytes: int, shape: tuple[int, ...], dtype: torch.dtype) -> torch.Tensor: + assert self.tensor is not None + return self.tensor.narrow(0, offset, nbytes).view(dtype).view(shape) + + def close(self, transfer: MooncakeTransfer) -> bool: + if self.tensor is not None and not transfer.unregister_memory(self.tensor): + return False + self.tensor = None + self._free.clear() + return True + + def allocate(self, nbytes: int) -> tuple[int, int] | None: + size = (nbytes + self.alignment - 1) // self.alignment * self.alignment + for index, (offset, available) in enumerate(self._free): + if size > available: + continue + if size == available: + self._free.pop(index) + else: + self._free[index] = (offset + size, available - size) + return offset, size + return None + + def free(self, offset: int, size: int) -> None: + index = bisect.bisect_left(self._free, (offset, size)) + self._free.insert(index, (offset, size)) + if index + 1 < len(self._free): + next_offset, next_size = self._free[index + 1] + if offset + size == next_offset: + self._free[index] = (offset, size + next_size) + self._free.pop(index + 1) + if index > 0: + previous_offset, previous_size = self._free[index - 1] + current_offset, current_size = self._free[index] + if previous_offset + previous_size == current_offset: + self._free[index - 1] = ( + previous_offset, + previous_size + current_size, + ) + self._free.pop(index) + + +class ResidentPool(Generic[_T]): + """Track resident values, active leases, and LRU eviction eligibility.""" + + def __init__(self): + self._entries: dict[str, _ResidentEntry[_T]] = {} + self._evictable: OrderedDict[str, None] = OrderedDict() + + def referenced(self) -> list[str]: + return [key for key, entry in self._entries.items() if entry.pinned] + + def get(self, key: str) -> _T | None: + entry = self._entries.get(key) + return entry.value if entry is not None else None + + def insert( + self, + key: str, + value: _T, + ) -> _T | None: + """Pin an entry and return a displaced value that has no owners.""" + previous = self._entries.get(key) + entry = _ResidentEntry(value) + self._entries[key] = entry + self._evictable.pop(key, None) + if previous is not None and previous.leases == 0: + return previous.value + return None + + def pin(self, key: str) -> _T | None: + entry = self._entries.get(key) + if entry is None: + return None + self._evictable.pop(key, None) + entry.pinned = True + return entry.value + + def acquire(self, key: str) -> ResidentLease[_T] | None: + entry = self._entries.get(key) + if entry is None: + return None + entry.leases += 1 + self._evictable.pop(key, None) + return ResidentLease(key, entry) + + def release(self, lease: ResidentLease[_T]) -> _T | None: + if not lease._active: + return None + lease._active = False + entry = lease._entry + entry.leases -= 1 + current = self._entries.get(lease.key) + if current is not entry: + if entry.leases == 0: + return entry.value + return None + if not entry.pinned and entry.leases == 0: + self._evictable[lease.key] = None + return None + + def consume(self, lease: ResidentLease[_T]) -> tuple[_T, _T | None]: + current = self._entries[lease.key] + current.pinned = True + self._evictable.pop(lease.key, None) + released = self.release(lease) + return current.value, released + + def retire(self, key: str) -> None: + if key not in self._entries: + return + entry = self._entries[key] + entry.pinned = False + if entry.leases == 0: + self._evictable[key] = None + + def evict_lru(self, evict: Callable[[str, _T], bool]) -> str | None: + for key in list(self._evictable): + entry = self._entries[key] + if not evict(key, entry.value): + continue + self._evictable.pop(key, None) + del self._entries[key] + return key + return None + + def clear(self) -> None: + self._entries.clear() + self._evictable.clear() + + +@dataclass +class StagedSources: + """Own Producer tensor views and their staging-slab regions.""" + + tensors: list[torch.Tensor] + regions: list[tuple[int, int]] + + +class ProducerMemoryPool: + """Own the Producer staging slab and regions carved from it.""" + + def __init__(self, capacity: int, transfer: MooncakeTransfer) -> None: + self._transfer = transfer + self._allocator = ContiguousAllocator(capacity) + self._lock = threading.Lock() + self._local = threading.local() + + @property + def tensor(self) -> torch.Tensor | None: + return self._allocator.tensor + + def _free_regions(self, regions: list[tuple[int, int]]) -> None: + for offset, size in regions: + self._allocator.free(offset, size) + + def stage(self, tensors: list[torch.Tensor]) -> StagedSources | None: + """Copy tensors into one registered slab, or return None for fallback.""" + if not tensors: + return StagedSources([], []) + allocator = self._allocator + staged: list[torch.Tensor] = [] + regions: list[tuple[int, int]] = [] + with self._lock: + allocator.prepare(tensors[0].device, self._transfer) + pool = allocator.tensor + if pool is None: + return None + for tensor in tensors: + region = allocator.allocate(tensor.nbytes) + if region is None: + self._free_regions(regions) + return None + regions.append(region) + staged.append(allocator.view(region[0], tensor.nbytes, tuple(tensor.shape), tensor.dtype)) + if pool.device.type == "cuda": + stream = getattr(self._local, "stream", None) + if stream is None: + stream = torch.cuda.Stream(device=pool.device) + self._local.stream = stream + with torch.cuda.stream(stream): + for destination, source in zip(staged, tensors): + destination.copy_(source, non_blocking=True) + stream.synchronize() + else: + for destination, source in zip(staged, tensors): + destination.copy_(source) + return StagedSources(staged, regions) + + def release(self, staged: StagedSources) -> None: + if not staged.regions: + return + with self._lock: + self._free_regions(staged.regions) + + def close(self) -> None: + with self._lock: + self._allocator.close(self._transfer) + + +class ConsumerMemoryPool: + """Own the registered receive slab and resident allocation lifecycle.""" + + def __init__( + self, + capacity: int, + transfer: MooncakeTransfer, + ) -> None: + self._transfer = transfer + self._allocator = ContiguousAllocator(capacity) + self._residents: ResidentPool[MemoryAllocation] = ResidentPool() + self._retire_events: dict[str, torch.Event] = {} + self._pending_frees: list[tuple[torch.Event, MemoryAllocation]] = [] + self._reclaimed: set[str] = set() + self.lock = threading.RLock() + + @property + def tensor(self) -> torch.Tensor | None: + return self._allocator.tensor + + def prepare( + self, + device: torch.device, + ) -> None: + with self.lock: + self._allocator.prepare(device, self._transfer) + + def _free(self, allocation: MemoryAllocation) -> None: + self._allocator.free(allocation.offset, allocation.size) + + def free(self, allocation: MemoryAllocation) -> None: + with self.lock: + self._free(allocation) + + def _defer_or_free(self, allocation: MemoryAllocation, event: torch.Event | None) -> None: + if event is None or event.query(): + self._free(allocation) + else: + self._pending_frees.append((event, allocation)) + + def _poll_frees_locked(self) -> None: + pending = [] + for event, allocation in self._pending_frees: + if event.query(): + self._free(allocation) + else: + pending.append((event, allocation)) + self._pending_frees = pending + + def _reclaim_locked(self, nbytes: int) -> tuple[int, int] | None: + def evict(mm_hash: str, allocation: MemoryAllocation) -> bool: + event = self._retire_events.pop(mm_hash, None) + self._defer_or_free(allocation, event) + self._reclaimed.add(mm_hash) + return True + + while self._residents.evict_lru(evict) is not None: + # Eviction defers any free whose CUDA event is still pending, so + # those bytes reach the allocator only once the event is polled. + self._poll_frees_locked() + region = self._allocator.allocate(nbytes) + if region is not None: + return region + self._poll_frees_locked() + return self._allocator.allocate(nbytes) + + def _make_allocation( + self, + region: tuple[int, int], + nbytes: int, + shape: tuple[int, ...], + dtype: torch.dtype, + ) -> MemoryAllocation: + offset, size = region + tensor = self._allocator.view(offset, nbytes, shape, dtype) + return MemoryAllocation(offset, size, tensor) + + def try_allocate(self, nbytes: int, shape: tuple[int, ...], dtype: torch.dtype) -> MemoryAllocation | None: + with self.lock: + allocator = self._allocator + assert allocator.tensor is not None + self._poll_frees_locked() + region = allocator.allocate(nbytes) + if region is None: + return None + return self._make_allocation(region, nbytes, shape, dtype) + + def reclaim_and_allocate(self, nbytes: int, shape: tuple[int, ...], dtype: torch.dtype) -> MemoryAllocation | None: + with self.lock: + region = self._reclaim_locked(nbytes) + if region is None: + return None + return self._make_allocation(region, nbytes, shape, dtype) + + def acquire_cached( + self, mm_hash: str, shape: tuple[int, ...], dtype: torch.dtype + ) -> ResidentLease[MemoryAllocation] | None: + with self.lock: + allocation = self._residents.get(mm_hash) + if allocation is None: + return None + if tuple(allocation.tensor.shape) != shape or allocation.tensor.dtype != dtype: + raise ValueError("conflicting cached tensor for mm_hash") + return self._residents.acquire(mm_hash) + + def take_resident(self, mm_hash: str, shape: tuple[int, ...], dtype_name: str) -> torch.Tensor | None: + with self.lock: + allocation = self._residents.get(mm_hash) + if allocation is None: + return None + tensor = allocation.tensor + if tuple(tensor.shape) != shape or str(tensor.dtype).split(".")[-1] != dtype_name: + return None + self._residents.pin(mm_hash) + self._retire_events.pop(mm_hash, None) + return tensor + + def _record_release_event(self) -> torch.Event | None: + pool = self.tensor + if pool is None or pool.device.type != "cuda": + return None + event = torch.Event() + event.record(torch.accelerator.current_stream(pool.device)) + return event + + def release_cached(self, lease: ResidentLease[MemoryAllocation]) -> None: + with self.lock: + released = self._residents.release(lease) + if released is not None: + self._defer_or_free(released, self._record_release_event()) + + def publish( + self, + mm_hash: str, + allocation: MemoryAllocation, + lease: ResidentLease[MemoryAllocation] | None = None, + *, + pin: bool = True, + ) -> MemoryAllocation: + with self.lock: + if lease is not None: + canonical, released = self._residents.consume(lease) + self._retire_events.pop(mm_hash, None) + if released is not None: + self._defer_or_free(released, self._record_release_event()) + return canonical + previous = self._residents.get(mm_hash) + displaced = self._residents.insert(mm_hash, allocation) + if not pin: + self._residents.retire(mm_hash) + event = None + if previous is not None and previous is not allocation: + event = self._retire_events.pop(mm_hash, None) + if displaced is None or displaced is allocation: + return allocation + if event is None: + event = self._record_release_event() + self._defer_or_free(displaced, event) + return allocation + + def retire_stale( + self, + encoder_cache: dict[str, torch.Tensor], + reserved_hashes: set[str] | None = None, + *, + freed: list[str] | None = None, + ) -> None: + if self.tensor is None: + return + with self.lock: + for mm_hash in self._residents.referenced() if freed is None else freed: + allocation = self._residents.get(mm_hash) + if allocation is None: + continue + if encoder_cache.get(mm_hash) is allocation.tensor: + continue + if reserved_hashes and mm_hash in reserved_hashes: + continue + event = self._record_release_event() + if event is not None: + self._retire_events[mm_hash] = event + self._residents.retire(mm_hash) + self._poll_frees_locked() + + def drain_reclaimed(self) -> set[str]: + with self.lock: + reclaimed = self._reclaimed + self._reclaimed = set() + return reclaimed + + def close(self) -> None: + with self.lock: + if not self._allocator.close(self._transfer): + return + self._residents.clear() + self._retire_events.clear() + self._pending_frees.clear() diff --git a/aphrodite/distributed/ec_transfer/ec_connector/mooncake/metadata.py b/aphrodite/distributed/ec_transfer/ec_connector/mooncake/metadata.py new file mode 100644 index 0000000000..5f5130d7f1 --- /dev/null +++ b/aphrodite/distributed/ec_transfer/ec_connector/mooncake/metadata.py @@ -0,0 +1,75 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +"""Metadata exchanged by the Mooncake encoder-cache connector.""" + +from __future__ import annotations + +from dataclasses import dataclass, field + +from aphrodite.distributed.ec_transfer.ec_connector.base import ( + ECConnectorMetadata, + ECConnectorWorkerMetadata, +) + + +@dataclass +class ECMooncakeLoadSpec: + """Describe a remote reservation or resident allocation to load.""" + + mm_hash: str + nbytes: int + shape: tuple[int, ...] + dtype: str + transfer_id: str + # The consumer pool still holds this item, so the load is a local handoff: + # no transfer, no producer. + local: bool = False + + +@dataclass +class ECMooncakePushSpec: + """Describe a destination reservation prepared before a tensor is ready.""" + + mm_hash: str + nbytes: int + shape: tuple[int, ...] + dtype: str + consumer_zmq: str + transfer_id: str + request_id: str = "" + + +@dataclass +class ECMooncakeConnectorMetadata(ECConnectorMetadata): + """Worker operations emitted for one Scheduler step.""" + + loads: list[ECMooncakeLoadSpec] = field(default_factory=list) + pushes: list[ECMooncakePushSpec] = field(default_factory=list) + freed: list[str] | None = None + + +@dataclass +class ECMooncakeWorkerMetadata(ECConnectorWorkerMetadata): + """Completion state reported from Workers to the Scheduler.""" + + loaded: set[str] = field(default_factory=set) + failed_loads: set[str] = field(default_factory=set) + # Items the receive pool dropped under pressure. The scheduler assumes an + # evicted item stays resident until told otherwise. + reclaimed: set[str] = field(default_factory=set) + pending_saves: bool = False + failed_saves: set[str] = field(default_factory=set) + + def aggregate(self, other: ECConnectorWorkerMetadata) -> ECMooncakeWorkerMetadata: + assert isinstance(other, ECMooncakeWorkerMetadata) + return ECMooncakeWorkerMetadata( + # Every tensor-parallel rank gathers the embedding from its own + # cache, so an item counts as loaded only where all of them have + # it; one rank falling short must fail the load rather than leave + # the scheduler believing it is ready. + loaded=self.loaded & other.loaded, + failed_loads=self.failed_loads | other.failed_loads, + reclaimed=self.reclaimed | other.reclaimed, + pending_saves=self.pending_saves or other.pending_saves, + failed_saves=self.failed_saves | other.failed_saves, + ) diff --git a/aphrodite/distributed/ec_transfer/ec_connector/mooncake/producer.py b/aphrodite/distributed/ec_transfer/ec_connector/mooncake/producer.py new file mode 100644 index 0000000000..465966df3c --- /dev/null +++ b/aphrodite/distributed/ec_transfer/ec_connector/mooncake/producer.py @@ -0,0 +1,393 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +"""Producer-side push lifecycle, futures, and source-tensor ownership. + +This module records when a push may advance or release its source. Worker +orchestration performs the actual control exchanges and Mooncake writes. +""" + +from __future__ import annotations + +import threading +from collections import OrderedDict +from collections.abc import Callable +from concurrent.futures import Future, ThreadPoolExecutor +from contextlib import suppress +from dataclasses import dataclass, field +from enum import Enum, auto +from typing import Any + +import torch + +from aphrodite.distributed.ec_transfer.ec_connector.mooncake.metadata import ( + ECMooncakePushSpec, +) + + +def _same_destination(existing: ECMooncakePushSpec, incoming: ECMooncakePushSpec) -> bool: + """Whether two specs describe one push, ignoring which request asked. + + The same encoding requested twice shares a transfer id legitimately; only + a different payload or destination is a collision. + """ + return ( + existing.mm_hash == incoming.mm_hash + and existing.nbytes == incoming.nbytes + and existing.shape == incoming.shape + and existing.dtype == incoming.dtype + and existing.consumer_zmq == incoming.consumer_zmq + ) + + +class ProducerPushState(Enum): + """Lifecycle of one Producer push from reservation to terminal state.""" + + # The reservation reply and source tensor may arrive in either order. + WAITING_INPUTS = auto() + WRITING = auto() + NOTIFYING = auto() + DONE = auto() + CANCEL_PENDING = auto() + CANCELLED = auto() + FAILED = auto() + + +@dataclass +class ProducerPushRecord: + """Collect all asynchronous state for one Producer push.""" + + spec: ECMooncakePushSpec + state: ProducerPushState + reservation_future: Future[list[dict[str, Any]]] + reservations: list[dict[str, Any]] = field(default_factory=list) + shard_futures: list[Future[Any]] = field(default_factory=list) + source_tensor: torch.Tensor | None = None + source_event: torch.Event | None = None + batch_future: Future[None] | None = None + error: str | None = None + + +_ALLOWED_TRANSITIONS = { + ProducerPushState.WAITING_INPUTS: { + ProducerPushState.WRITING, + ProducerPushState.CANCEL_PENDING, + ProducerPushState.FAILED, + }, + ProducerPushState.WRITING: { + ProducerPushState.NOTIFYING, + ProducerPushState.FAILED, + }, + ProducerPushState.NOTIFYING: { + ProducerPushState.DONE, + ProducerPushState.FAILED, + }, + ProducerPushState.CANCEL_PENDING: { + ProducerPushState.CANCELLED, + ProducerPushState.FAILED, + }, + ProducerPushState.DONE: set(), + ProducerPushState.CANCELLED: set(), + ProducerPushState.FAILED: set(), +} + +_TERMINAL_STATES = { + ProducerPushState.DONE, + ProducerPushState.CANCELLED, + ProducerPushState.FAILED, +} +_TERMINAL_LIMIT = 1 << 16 + + +class ProducerPushManager: + """Own Producer push records, transitions, and source tensor leases. + + Terminal records are retained for duplicate-metadata idempotency. Separate + ordered indexes keep hot polling paths from scanning those tombstones. + """ + + def __init__(self, wake: Callable[[], None] = lambda: None) -> None: + self._records: OrderedDict[str, ProducerPushRecord] = OrderedDict() + self._active_ids: OrderedDict[str, None] = OrderedDict() + self._reapable_terminal_ids: OrderedDict[str, None] = OrderedDict() + self._unreported_ids: OrderedDict[str, None] = OrderedDict() + self._batch_ids: OrderedDict[str, None] = OrderedDict() + self._source_waiters: dict[str, OrderedDict[str, None]] = {} + self._lock = threading.RLock() + self._wake = wake + + def reserve( + self, + spec: ECMooncakePushSpec, + submit: Callable[[], Future[list[dict[str, Any]]]], + ) -> tuple[ProducerPushRecord, bool]: + with self._lock: + existing = self._records.get(spec.transfer_id) + if existing is not None: + if not _same_destination(existing.spec, spec): + raise ValueError(f"Conflicting EC destination for transfer_id={spec.transfer_id}") + return existing, False + record = ProducerPushRecord( + spec=spec, + state=ProducerPushState.WAITING_INPUTS, + reservation_future=submit(), + ) + self._records[spec.transfer_id] = record + self._active_ids[spec.transfer_id] = None + self._source_waiters.setdefault(spec.mm_hash, OrderedDict())[spec.transfer_id] = None + record.reservation_future.add_done_callback(lambda future: self._reservation_done(record, future)) + return record, True + + def bind_source( + self, + mm_hash: str, + tensor: torch.Tensor, + ready_event: torch.Event | None, + ) -> None: + with self._lock: + waiters = self._source_waiters.pop(mm_hash, OrderedDict()) + for transfer_id in waiters: + record = self._records[transfer_id] + if record.source_tensor is not None or record.state is not ProducerPushState.WAITING_INPUTS: + continue + record.source_tensor = tensor + record.source_event = ready_event + self._wake() + + def submit_batches( + self, + executor: ThreadPoolExecutor, + run_batch: Callable[[list[ProducerPushRecord]], None], + *, + wait: bool = False, + ) -> bool: + pending_event = False + with self._lock: + grouped: dict[str, list[ProducerPushRecord]] = {} + for transfer_id in list(self._active_ids): + record = self._records[transfer_id] + if ( + record.source_tensor is not None + and record.batch_future is None + and record.state is ProducerPushState.WAITING_INPUTS + ): + if not record.reservation_future.done(): + continue + if record.source_event is not None: + if wait: + record.source_event.synchronize() + elif not record.source_event.query(): + pending_event = True + continue + grouped.setdefault(record.spec.consumer_zmq, []).append(record) + batches = list(grouped.values()) + for records in batches: + future = executor.submit(run_batch, records) + for record in records: + record.batch_future = future + self._batch_ids[record.spec.transfer_id] = None + return pending_event + + def resolve_reservations(self, record: ProducerPushRecord) -> list[dict[str, Any]]: + results = record.reservation_future.result() + with self._lock: + record.reservations = results + return list(results) + + def finish_reservations( + self, + outcomes: list[tuple[ProducerPushRecord, list[dict[str, Any]] | Exception]], + ) -> None: + with self._lock: + for record, result in outcomes: + if isinstance(result, Exception): + record.reservation_future.set_exception(result) + else: + record.reservation_future.set_result(result) + + def _reservation_done( + self, + record: ProducerPushRecord, + future: Future[list[dict[str, Any]]], + ) -> None: + try: + future.result() + except Exception as exc: + with self._lock: + self._set_error(record, exc) + if record.state is ProducerPushState.WAITING_INPUTS and record.source_tensor is None: + self._transition(record, ProducerPushState.FAILED) + finally: + self._wake() + + def settle_all(self, records: list[ProducerPushRecord]) -> None: + for record in records: + with suppress(Exception): + record.reservation_future.result() + + def replace_reservations( + self, + record: ProducerPushRecord, + reservations: list[dict[str, Any]], + ) -> None: + with self._lock: + record.reservations = reservations + + def track_shard_futures( + self, + records: list[ProducerPushRecord], + futures: list[Future[Any]], + ) -> None: + with self._lock: + for record in records: + record.shard_futures.extend(futures) + + def begin_writing(self, record: ProducerPushRecord) -> None: + with self._lock: + if record.source_tensor is None: + raise RuntimeError(f"Producer push {record.spec.transfer_id!r} has no source tensor") + self._transition(record, ProducerPushState.WRITING) + + def begin_notifying(self, records: list[ProducerPushRecord]) -> None: + with self._lock: + for record in records: + self._transition(record, ProducerPushState.NOTIFYING) + + def complete(self, records: list[ProducerPushRecord]) -> None: + with self._lock: + self._check_sources_releasable(records) + for record in records: + self._transition(record, ProducerPushState.DONE) + self._release_source(record) + + def fail(self, records: list[ProducerPushRecord], error: Exception) -> None: + with self._lock: + self._check_sources_releasable(records) + for record in records: + if record.state not in _TERMINAL_STATES: + self._set_error(record, error) + self._transition(record, ProducerPushState.FAILED) + self._release_source(record) + + def cancel_requests(self, request_ids: set[str] | None) -> list[ProducerPushRecord]: + """Cancel source-less waiters for requests, or all waiters at shutdown.""" + cancelled = [] + with self._lock: + for transfer_id in list(self._active_ids): + record = self._records[transfer_id] + if ( + (request_ids is not None and record.spec.request_id not in request_ids) + or record.source_tensor is not None + or record.batch_future is not None + ): + continue + if record.state is not ProducerPushState.WAITING_INPUTS: + continue + self._transition(record, ProducerPushState.CANCEL_PENDING) + cancelled.append(record) + return cancelled + + def finish_cancel(self, record: ProducerPushRecord) -> None: + with self._lock: + if record.state is ProducerPushState.CANCEL_PENDING: + self._transition(record, ProducerPushState.CANCELLED) + + def submit_cancel( + self, + record: ProducerPushRecord, + executor: ThreadPoolExecutor, + run_cancel: Callable[[ProducerPushRecord], None], + ) -> None: + with self._lock: + future = executor.submit(run_cancel, record) + record.batch_future = future + transfer_id = record.spec.transfer_id + self._batch_ids[transfer_id] = None + if record.state in _TERMINAL_STATES: + self._reapable_terminal_ids.pop(transfer_id, None) + + def poll(self) -> list[tuple[str, str]]: + failures = [] + with self._lock: + for transfer_id in list(self._batch_ids): + record = self._records[transfer_id] + future = record.batch_future + assert future is not None + if not future.done(): + continue + try: + future.result() + except Exception as exc: + self._set_error(record, exc) + self._batch_ids.pop(transfer_id) + self._mark_reapable(record) + for transfer_id in list(self._unreported_ids): + if transfer_id in self._batch_ids: + continue + record = self._records[transfer_id] + assert record.error is not None + failures.append((record.spec.mm_hash, record.error)) + self._unreported_ids.pop(transfer_id) + self._mark_reapable(record) + self._reap_terminals() + return failures + + @property + def pending(self) -> bool: + with self._lock: + return bool(self._active_ids or self._batch_ids) + + def _release_source(self, record: ProducerPushRecord) -> None: + record.source_tensor = None + record.source_event = None + + def _set_error(self, record: ProducerPushRecord, error: BaseException) -> None: + if record.error is None: + record.error = str(error) + if record.state in _TERMINAL_STATES: + transfer_id = record.spec.transfer_id + self._unreported_ids[transfer_id] = None + self._reapable_terminal_ids.pop(transfer_id, None) + + def _reap_terminals(self) -> None: + while len(self._reapable_terminal_ids) > _TERMINAL_LIMIT: + transfer_id, _ = self._reapable_terminal_ids.popitem(last=False) + self._records.pop(transfer_id) + + def _mark_reapable(self, record: ProducerPushRecord) -> None: + transfer_id = record.spec.transfer_id + if ( + record.state in _TERMINAL_STATES + and transfer_id not in self._batch_ids + and transfer_id not in self._unreported_ids + ): + self._reapable_terminal_ids[transfer_id] = None + + def _drop_source_waiter(self, record: ProducerPushRecord) -> None: + waiters = self._source_waiters.get(record.spec.mm_hash) + if waiters is None: + return + waiters.pop(record.spec.transfer_id, None) + if not waiters: + self._source_waiters.pop(record.spec.mm_hash) + + @staticmethod + def _check_sources_releasable(records: list[ProducerPushRecord]) -> None: + for record in records: + futures = [record.reservation_future, *record.shard_futures] + if record.source_tensor is not None and not all(future.done() for future in futures): + raise RuntimeError(f"Producer push {record.spec.transfer_id!r} released its source too early") + + def _transition(self, record: ProducerPushRecord, state: ProducerPushState) -> None: + if state not in _ALLOWED_TRANSITIONS[record.state]: + raise RuntimeError( + f"Cannot transition producer push {record.spec.transfer_id!r} from {record.state.name} to {state.name}" + ) + record.state = state + if state is not ProducerPushState.WAITING_INPUTS: + self._drop_source_waiter(record) + if state in _TERMINAL_STATES: + transfer_id = record.spec.transfer_id + self._active_ids.pop(transfer_id, None) + if record.error is not None: + self._unreported_ids[transfer_id] = None + self._mark_reapable(record) diff --git a/aphrodite/distributed/ec_transfer/ec_connector/mooncake/reservation.py b/aphrodite/distributed/ec_transfer/ec_connector/mooncake/reservation.py new file mode 100644 index 0000000000..da98a5dd72 --- /dev/null +++ b/aphrodite/distributed/ec_transfer/ec_connector/mooncake/reservation.py @@ -0,0 +1,442 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +"""Consumer-side reservation lifecycle and destination-memory ownership. + +Reservations make remote writes idempotent and ensure cancellation or expiry +cannot free a destination while Mooncake may still be writing into it. +""" + +from __future__ import annotations + +import threading +import time +import uuid +from collections import OrderedDict +from dataclasses import dataclass +from enum import Enum, auto + +import torch + +from aphrodite.distributed.ec_transfer.ec_connector.mooncake.memory import ( + ConsumerMemoryPool, + MemoryAllocation, + ResidentLease, +) + + +class ConsumerReservationState(Enum): + """Lifecycle of one Consumer destination reservation.""" + + WRITING = auto() + READY = auto() + CANCEL_PENDING = auto() + EXPIRE_PENDING = auto() + CANCELLED = auto() + EXPIRED = auto() + + +_ALLOWED_TRANSITIONS = { + ConsumerReservationState.WRITING: { + ConsumerReservationState.READY, + ConsumerReservationState.CANCEL_PENDING, + ConsumerReservationState.EXPIRE_PENDING, + ConsumerReservationState.CANCELLED, + }, + ConsumerReservationState.READY: { + ConsumerReservationState.CANCELLED, + ConsumerReservationState.EXPIRED, + }, + ConsumerReservationState.CANCEL_PENDING: {ConsumerReservationState.CANCELLED}, + ConsumerReservationState.EXPIRE_PENDING: { + ConsumerReservationState.CANCELLED, + ConsumerReservationState.EXPIRED, + }, + ConsumerReservationState.CANCELLED: set(), + ConsumerReservationState.EXPIRED: set(), +} + +_ACTIVE_STATES = { + ConsumerReservationState.WRITING, + ConsumerReservationState.READY, + ConsumerReservationState.CANCEL_PENDING, + ConsumerReservationState.EXPIRE_PENDING, +} +_DEFERRED_STATES = { + ConsumerReservationState.CANCEL_PENDING, + ConsumerReservationState.EXPIRE_PENDING, +} +_WRITER_OWNED_STATES = { + ConsumerReservationState.WRITING, + ConsumerReservationState.CANCEL_PENDING, + ConsumerReservationState.EXPIRE_PENDING, +} + + +@dataclass +class ConsumerReservation: + """Own the identity and destination allocation of one remote push.""" + + transfer_id: str + mm_hash: str + reservation_id: str + state: ConsumerReservationState + shape: tuple[int, ...] = () + dtype: str = "" + allocation: MemoryAllocation | None = None + lease: ResidentLease[MemoryAllocation] | None = None + expires_at: float = 0 + writer_id: str = "" + + +class ConsumerReservationManager: + """Own reservation transitions and destination allocation releases. + + A remote writer owns a ``WRITING`` allocation. Cancellation and expiry are + therefore deferred until completion, unless refresh explicitly abandons + that writer before replacing its reservation. + """ + + def __init__( + self, + memory: ConsumerMemoryPool, + lease_ttl: float, + tombstone_limit: int, + ) -> None: + self._memory = memory + self._lease_ttl = lease_ttl + self._tombstone_limit = tombstone_limit + self._records: dict[str, ConsumerReservation] = {} + self._active_ids: dict[str, None] = {} + self._tombstones: OrderedDict[str, None] = OrderedDict() + self._writers: dict[str, ConsumerReservation] = {} + self._followers: dict[str, dict[str, None]] = {} + self._ready: list[str] = [] + self._condition = threading.Condition(memory.lock) + self._shutting_down = False + + def reserve( + self, + transfer_id: str, + mm_hash: str, + nbytes: int, + shape: tuple[int, ...], + dtype_name: str, + dtype: torch.dtype, + ) -> tuple[ConsumerReservation | None, bool]: + with self._condition: + if self._shutting_down: + raise RuntimeError("Consumer reservation manager is shutting down") + existing = self._records.get(transfer_id) + if existing is not None and existing.state is ConsumerReservationState.CANCELLED: + return existing, False + if existing is not None and existing.state is ConsumerReservationState.EXPIRED: + self._remove(transfer_id) + existing = None + if existing is not None and existing.state is ConsumerReservationState.EXPIRE_PENDING: + raise RuntimeError(f"Transfer {transfer_id!r} still has an active writer") + if existing is not None: + if existing.mm_hash != mm_hash or existing.shape != shape or existing.dtype != dtype_name: + raise ValueError("conflicting reservation for transfer_id") + if existing.state not in _ACTIVE_STATES: + raise RuntimeError(f"Cannot reserve transfer {transfer_id!r} in {existing.state.name}") + if existing.state is ConsumerReservationState.WRITING: + existing.expires_at = time.monotonic() + self._lease_ttl + return existing, False + + lease = self._memory.acquire_cached(mm_hash, shape, dtype) + now = time.monotonic() + if lease is not None: + record = ConsumerReservation( + transfer_id, + mm_hash, + uuid.uuid4().hex, + ConsumerReservationState.READY, + shape, + dtype_name, + lease.value, + lease, + now + self._lease_ttl, + ) + self._insert(record) + return record, False + writer = self._writers.get(mm_hash) + if writer is not None and writer.state is ConsumerReservationState.WRITING: + if writer.shape != shape or writer.dtype != dtype_name: + raise ValueError("conflicting in-flight tensor for mm_hash") + record = ConsumerReservation( + transfer_id, + mm_hash, + uuid.uuid4().hex, + ConsumerReservationState.WRITING, + shape, + dtype_name, + writer.allocation, + expires_at=now + self._lease_ttl, + writer_id=writer.transfer_id, + ) + self._insert(record) + self._followers.setdefault(writer.transfer_id, {})[transfer_id] = None + return record, False + allocation = self._memory.try_allocate(nbytes, shape, dtype) + if allocation is None: + self._expire_locked(time.monotonic()) + allocation = self._memory.try_allocate(nbytes, shape, dtype) + if allocation is None: + allocation = self._memory.reclaim_and_allocate(nbytes, shape, dtype) + if allocation is None: + return None, False + record = ConsumerReservation( + transfer_id, + mm_hash, + uuid.uuid4().hex, + ConsumerReservationState.WRITING, + shape, + dtype_name, + allocation, + None, + now + self._lease_ttl, + ) + self._insert(record) + self._writers[mm_hash] = record + return record, True + + def status(self, transfer_id: str) -> ConsumerReservation | None: + with self._memory.lock: + record = self._records.get(transfer_id) + if record is None or record.state not in _ACTIVE_STATES: + return None + return record + + def complete(self, transfer_id: str, reservation_id: str) -> tuple[bool, bool]: + with self._condition: + try: + record = self._records.get(transfer_id) + if record is None or record.reservation_id != reservation_id: + return False, False + if record.state is ConsumerReservationState.READY: + return True, False + if record.writer_id: + return False, False + if record.state in _WRITER_OWNED_STATES: + self._publish_followers(record) + if record.state in _DEFERRED_STATES: + terminal = ( + ConsumerReservationState.CANCELLED + if record.state is ConsumerReservationState.CANCEL_PENDING + else ConsumerReservationState.EXPIRED + ) + self._terminate(record, terminal) + return True, False + if record.state is not ConsumerReservationState.WRITING: + return False, False + self._transition(record, ConsumerReservationState.READY) + record.expires_at = time.monotonic() + self._lease_ttl + return True, True + finally: + self._condition.notify_all() + + def _publish_followers(self, writer: ConsumerReservation) -> None: + if self._writers.get(writer.mm_hash) is writer: + self._writers.pop(writer.mm_hash) + followers = self._followers.pop(writer.transfer_id, {}) + if not followers: + return + assert writer.allocation is not None + allocation = self._memory.publish(writer.mm_hash, writer.allocation, pin=False) + for record in [writer, *(self._records[key] for key in followers)]: + record.allocation = allocation + record.lease = self._memory.acquire_cached(record.mm_hash, record.shape, allocation.tensor.dtype) + assert record.lease is not None + if record is not writer: + record.writer_id = "" + self._transition(record, ConsumerReservationState.READY) + record.expires_at = time.monotonic() + self._lease_ttl + self._ready.append(record.transfer_id) + + def drain_ready(self) -> list[str]: + with self._memory.lock: + ready, self._ready = self._ready, [] + return ready + + def begin_shutdown(self) -> None: + """Stop new reservations and cancel everything without a remote writer.""" + with self._condition: + if self._shutting_down: + return + self._shutting_down = True + for transfer_id in list(self._active_ids): + record = self._records[transfer_id] + if record.writer_id or record.state is ConsumerReservationState.READY: + self._terminate(record, ConsumerReservationState.CANCELLED) + elif record.state is ConsumerReservationState.WRITING: + self._defer(record, ConsumerReservationState.CANCEL_PENDING) + self._condition.notify_all() + + def wait_for_writers(self, timeout: float) -> bool: + """Wait until no reservation is still owned by a remote writer.""" + + def writers_finished() -> bool: + return not any(self._records[transfer_id].state in _WRITER_OWNED_STATES for transfer_id in self._active_ids) + + with self._condition: + return self._condition.wait_for(writers_finished, timeout=max(0, timeout)) + + def cancel( + self, + transfer_id: str, + reservation_id: str, + abandon: bool = False, + refresh: bool = False, + ) -> bool: + with self._memory.lock: + record = self._records.get(transfer_id) + if record is not None and reservation_id and record.reservation_id != reservation_id: + return False + if record is None: + record = ConsumerReservation( + transfer_id, + "", + "", + ConsumerReservationState.CANCELLED, + ) + self._insert(record) + self._set_tombstone_deadline(record) + self._reap_tombstones(time.monotonic()) + return True + if record.state is ConsumerReservationState.CANCELLED: + self._set_tombstone_deadline(record) + self._reap_tombstones(time.monotonic()) + return True + if refresh: + if not abandon or record.state not in { + ConsumerReservationState.WRITING, + ConsumerReservationState.EXPIRE_PENDING, + }: + return False + if record.state is ConsumerReservationState.WRITING: + self._transition(record, ConsumerReservationState.EXPIRE_PENDING) + self._terminate(record, ConsumerReservationState.EXPIRED) + self._condition.notify_all() + self._reap_tombstones(time.monotonic()) + return True + if record.writer_id: + self._terminate(record, ConsumerReservationState.CANCELLED) + return True + if record.state in _DEFERRED_STATES and not abandon: + return True + if record.state is ConsumerReservationState.WRITING and not abandon: + self._defer(record, ConsumerReservationState.CANCEL_PENDING) + return True + if record.state not in _ACTIVE_STATES: + return False + self._terminate(record, ConsumerReservationState.CANCELLED) + if self._shutting_down: + self._condition.notify_all() + self._reap_tombstones(time.monotonic()) + return True + + def take(self, transfer_id: str, mm_hash: str) -> MemoryAllocation: + with self._memory.lock: + record = self._records.get(transfer_id) + if ( + record is None + or record.state is not ConsumerReservationState.READY + or record.mm_hash != mm_hash + or record.allocation is None + ): + raise RuntimeError(f"Pushed EC tensor is not ready for mm_hash={mm_hash}") + allocation = self._memory.publish(mm_hash, record.allocation, record.lease) + record.allocation = None + record.lease = None + self._remove(transfer_id) + return allocation + + def expire(self) -> int: + with self._memory.lock: + return self._expire_locked(time.monotonic()) + + def retire_stale(self, encoder_cache: dict[str, torch.Tensor], freed: list[str] | None = None) -> None: + with self._memory.lock: + self._memory.retire_stale(encoder_cache, freed=freed) + + def _expire_locked(self, now: float) -> int: + expired = 0 + for transfer_id in list(self._active_ids): + record = self._records[transfer_id] + if record.expires_at > now: + continue + if record.state is ConsumerReservationState.READY: + self._terminate(record, ConsumerReservationState.EXPIRED) + expired += 1 + elif record.state is ConsumerReservationState.WRITING: + if record.writer_id: + self._terminate(record, ConsumerReservationState.CANCELLED) + else: + self._defer(record, ConsumerReservationState.EXPIRE_PENDING) + self._reap_tombstones(now) + return expired + + def _defer(self, record: ConsumerReservation, state: ConsumerReservationState) -> None: + """Keep the destination until the writer completes or abandons it.""" + self._transition(record, state) + record.expires_at = float("inf") + + def _terminate(self, record: ConsumerReservation, state: ConsumerReservationState) -> None: + self._transition(record, state) + self._release(record) + self._set_tombstone_deadline(record) + + def _release(self, record: ConsumerReservation) -> None: + if record.writer_id: + followers = self._followers.get(record.writer_id, {}) + followers.pop(record.transfer_id, None) + record.allocation = None + return + if self._writers.get(record.mm_hash) is record: + self._writers.pop(record.mm_hash) + for transfer_id in list(self._followers.pop(record.transfer_id, {})): + self._terminate(self._records[transfer_id], ConsumerReservationState.CANCELLED) + allocation = record.allocation + if allocation is None: + return + if record.lease is not None: + self._memory.release_cached(record.lease) + else: + self._memory.free(allocation) + record.allocation = None + record.lease = None + + def _set_tombstone_deadline(self, record: ConsumerReservation) -> None: + record.expires_at = time.monotonic() + self._lease_ttl + self._tombstones[record.transfer_id] = None + self._tombstones.move_to_end(record.transfer_id) + + def _reap_tombstones(self, now: float) -> int: + dropped = 0 + while self._tombstones: + transfer_id = next(iter(self._tombstones)) + record = self._records[transfer_id] + if record.expires_at > now and len(self._tombstones) <= self._tombstone_limit: + break + self._remove(transfer_id) + dropped += 1 + return dropped + + def _insert(self, record: ConsumerReservation) -> None: + self._records[record.transfer_id] = record + if record.state in _ACTIVE_STATES: + self._active_ids[record.transfer_id] = None + + def _remove(self, transfer_id: str) -> None: + self._records.pop(transfer_id, None) + self._active_ids.pop(transfer_id, None) + self._tombstones.pop(transfer_id, None) + + def _transition(self, record: ConsumerReservation, state: ConsumerReservationState) -> None: + if state not in _ALLOWED_TRANSITIONS[record.state]: + raise RuntimeError(f"Cannot transition {record.transfer_id!r} from {record.state.name} to {state.name}") + record.state = state + if state in _ACTIVE_STATES: + self._active_ids[record.transfer_id] = None + else: + self._active_ids.pop(record.transfer_id, None) diff --git a/aphrodite/distributed/ec_transfer/ec_connector/mooncake/scheduler.py b/aphrodite/distributed/ec_transfer/ec_connector/mooncake/scheduler.py new file mode 100644 index 0000000000..bcb6b92d53 --- /dev/null +++ b/aphrodite/distributed/ec_transfer/ec_connector/mooncake/scheduler.py @@ -0,0 +1,510 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +"""Scheduler-side planning and observation for Mooncake cache transfers. + +The Scheduler prepares Producer reservations, consumes Consumer readiness +events, emits per-step Worker metadata, and converts Worker results back into +request availability without owning tensor memory or running data transfers. +""" + +from __future__ import annotations + +import math +import time +from collections import OrderedDict +from concurrent.futures import Future, ThreadPoolExecutor +from typing import TYPE_CHECKING, Any + +import torch + +from aphrodite.distributed.ec_transfer.ec_connector.base import ( + ECConnectorMetadata, +) +from aphrodite.distributed.ec_transfer.ec_connector.cpu.common import ( + _get_encoder_cache_hidden_dim, +) +from aphrodite.distributed.ec_transfer.ec_connector.mooncake.config import ( + _RESERVATION_TTL_SECONDS, + MooncakeECConfig, +) +from aphrodite.distributed.ec_transfer.ec_connector.mooncake.control import ( + ControlClient, + EventInbox, + make_cancel_request, +) +from aphrodite.distributed.ec_transfer.ec_connector.mooncake.metadata import ( + ECMooncakeConnectorMetadata, + ECMooncakeLoadSpec, + ECMooncakePushSpec, + ECMooncakeWorkerMetadata, +) +from aphrodite.distributed.ec_transfer.ec_connector.mooncake.state import ( + SchedulerTransferState, + SchedulerTransferTable, +) +from aphrodite.distributed.ec_transfer.ec_connector.mooncake.transfer import ( + ensure_mooncake_available, +) +from aphrodite.distributed.ec_transfer.ec_connector.utils import ( + PlaceholderMetadataResolver, + collect_ec_item_metadata, +) +from aphrodite.logger import init_logger +from aphrodite.v1.core.sched.output import SchedulerOutput +from aphrodite.v1.outputs import ECConnectorOutput + +if TYPE_CHECKING: + from aphrodite.config import AphroditeConfig + +logger = init_logger(__name__) + +# Unmatched `ec_items` entries are a protocol desync, not a per-request +# event: warn on the first few, then only once in a while. +_MAX_UNRESOLVED_TRANSFER_ID_WARNINGS = 5 + +_DRAIN_MIN_INTERVAL = 0.005 +_MAX_PENDING_EVENTS = 4096 +_MAX_TERMINAL_TRANSFER_RECORDS = 1 << 16 +_CANCEL_ATTEMPTS = 2 +_CONTROL_WORKERS = 8 + + +class ECMooncakeScheduler: + """Coordinate Mooncake transfers from the Aphrodite Scheduler process.""" + + def __init__(self, aphrodite_config: AphroditeConfig) -> None: + ensure_mooncake_available() + config = MooncakeECConfig.from_aphrodite_config(aphrodite_config) + + self._is_producer = config.is_producer + self._is_consumer = config.is_consumer + self._control_addr = config.control_addr + self._push_wait_timeout = config.push_wait_timeout_s + self._encoder_cache_hidden_dim = _get_encoder_cache_hidden_dim(aphrodite_config) if config.is_producer else None + self._model_config = aphrodite_config.model_config + self._control_client = ControlClient(config.control_timeout_ms) + self._control_executor = ThreadPoolExecutor( + max_workers=_CONTROL_WORKERS, + thread_name_prefix="ec-mooncake-control", + ) + self._event_inbox = EventInbox(self._control_client, self._control_executor) + + self._metadata_resolver = PlaceholderMetadataResolver(aphrodite_config.model_config) + self._unresolved_transfer_ids = 0 + self._drain_pending = True + self._drained_at = 0.0 + self._pending_cancels: dict[str, Future[Any]] = {} + self._transfers = SchedulerTransferTable(config.pool_size, _RESERVATION_TTL_SECONDS) + self._scheduler_pending_work = False + self._pushes_to_prepare: dict[str, ECMooncakePushSpec] = {} + self._prepared_push_transfer_ids: set[str] = set() + self._event_ready_shards: OrderedDict[str, set[int]] = OrderedDict() + # Mirror of the engine's encoder cache, maintained from the alloc and + # free notifications the Scheduler already sends. Tracking it here + # keeps `ensure_cache_available` on the upstream two-argument shape. + self._local_cache: set[str] = set() + self._failed_saves: set[str] = set() + + def _cancel_remote(self, consumer_zmq: str, transfer_id: str) -> bool: + pending = None + for _ in range(_CANCEL_ATTEMPTS): + pending = self._control_client.discover_shards(consumer_zmq) + if pending is not None: + break + if pending is None: + raise RuntimeError(f"Could not discover every EC consumer shard at {consumer_zmq}") + + cancelled = False + error: BaseException | None = None + for _ in range(_CANCEL_ATTEMPTS): + failed = [] + for addr in pending: + try: + result = self._control_client.request( + addr, + make_cancel_request(transfer_id), + ) + except Exception as exc: + if error is None: + error = exc + failed.append(addr) + continue + cancelled |= isinstance(result, dict) and bool(result.get("cancelled")) + if not failed: + return cancelled + pending = failed + if error is not None: + raise error + return cancelled + + def _note_awaiting_push( + self, + mm_hash: str, + transfer_id: str, + request_id: str, + now: float, + ) -> None: + record = self._transfers.wait_for_event( + transfer_id, + request_id, + mm_hash, + now + self._push_wait_timeout, + ) + if record is None: + # The id names another encoding; the request is already failed. + return + if record.state is not SchedulerTransferState.WAITING_EVENT: + return + assert record.deadline is not None + if now < record.deadline: + return + elapsed = now - record.deadline + self._push_wait_timeout + self._transfers.mark_unavailable(transfer_id, now) + logger.warning( + "EC Mooncake waited %.1fs for a push of mm_hash=%s " + "(transfer_id=%s) that never arrived; " + "requests needing it fail with a retryable error.", + elapsed, + mm_hash, + transfer_id, + ) + + def take_unavailable_requests(self) -> set[str]: + failed, self._failed_saves = self._failed_saves, set() + return failed | self._transfers.take_unavailable_requests() + + def _poll_pending_cancels(self) -> None: + pending = {} + for transfer_id, future in self._pending_cancels.items(): + if not future.done(): + pending[transfer_id] = future + continue + try: + future.result() + except Exception: + logger.warning("EC Mooncake reservation cancellation failed", exc_info=True) + self._pending_cancels = pending + + def _note_shard_ready(self, data: dict[str, Any]) -> bool: + shard_count = self._event_inbox.shard_count + if shard_count <= 1: + return True + transfer_id = str(data["transfer_id"]) + record = self._transfers.get(transfer_id) + if record is not None and record.state is not SchedulerTransferState.WAITING_EVENT: + return False + shard = data.get("shard") + shards = self._event_ready_shards.setdefault(transfer_id, set()) + self._event_ready_shards.move_to_end(transfer_id) + shards.add(int(shard) if shard is not None else len(shards)) + if len(shards) < shard_count: + while len(self._event_ready_shards) > _MAX_PENDING_EVENTS: + self._event_ready_shards.popitem(last=False) + return False + self._event_ready_shards.pop(transfer_id, None) + return True + + def _forget_shard_readiness(self, transfer_id: str) -> None: + self._event_ready_shards.pop(transfer_id, None) + + def _store_pushed_spec(self, data: dict[str, Any]) -> bool: + transfer_id = str(data["transfer_id"]) + identifier = str(data["mm_hash"]) + _, accepted = self._transfers.observe_ready( + ECMooncakeLoadSpec( + mm_hash=identifier, + nbytes=int(data["nbytes"]), + shape=tuple(int(value) for value in data["shape"]), + dtype=str(data["dtype"]), + transfer_id=transfer_id, + ), + time.monotonic() + _RESERVATION_TTL_SECONDS, + ) + return accepted + + def _queue_cancel( + self, + transfer_id: str, + mm_hash: str = "", + request_id: str = "", + ) -> None: + if not self._transfers.cancel( + transfer_id, + time.monotonic(), + mm_hash=mm_hash, + request_id=request_id, + ): + return + self._forget_shard_readiness(transfer_id) + self._pending_cancels[transfer_id] = self._control_executor.submit( + self._cancel_remote, + self._control_addr, + transfer_id, + ) + + def _expire_transfers(self) -> None: + now = time.monotonic() + expired = self._transfers.expire(now, _MAX_TERMINAL_TRANSFER_RECORDS) + for record in expired: + self._queue_cancel(record.transfer_id) + + def _drain_push_notifications(self) -> None: + now = time.monotonic() + if not self._drain_pending and now - self._drained_at < _DRAIN_MIN_INTERVAL: + return + self._drain_pending = False + self._drained_at = now + self._poll_pending_cancels() + self._expire_transfers() + events = self._event_inbox.drain(self._control_addr) + for data in events: + try: + if not isinstance(data, dict): + raise TypeError("EC readiness event must be an object") + if not data.get("ready"): + continue + self._accept_ready_event(data) + except (KeyError, TypeError, ValueError): + # Readiness events cross a plain PULL socket, so a malformed + # one costs that event and not the engine. + logger.warning("Discarding a malformed EC readiness event.", exc_info=True) + + def _accept_ready_event(self, data: dict[str, Any]) -> None: + transfer_id = str(data["transfer_id"]) + record = self._transfers.get(transfer_id) + if record is not None and record.state in { + SchedulerTransferState.CANCELLED, + SchedulerTransferState.UNAVAILABLE, + SchedulerTransferState.EXPIRED, + SchedulerTransferState.FAILED, + }: + return + if not self._note_shard_ready(data): + return + self._store_pushed_spec(data) + + def has_cache_item(self, identifier: str) -> bool: + if not self._is_consumer: + return False + self._drain_push_notifications() + return self._transfers.has_state( + identifier, + ( + SchedulerTransferState.READY, + SchedulerTransferState.RESIDENT, + SchedulerTransferState.AVAILABLE, + ), + ) + + def _warn_unresolved_transfer_id(self, request: Any, index: int, where: str) -> None: + """Report an `ec_items` entry that cannot be matched to a feature. + + A caller that cannot resolve a transfer id has to fall back to a + locally invented one or skip its bookkeeping entirely, and either way + the two sides of a transfer stop agreeing on its name. That desyncs + silently -- reservations are never released and only surface minutes + later as a full buffer pool -- so say it out loud the first time. + + `mm_hash` in `ec_items` must be a value some engine reported, never one + the caller derived itself: `mm_features[i].identifier` folds in the + engine's media_io_kwargs and mm_processor_kwargs, so a media uuid alone + no longer equals it. Omit the field to match by position instead. + """ + self._unresolved_transfer_ids += 1 + count = self._unresolved_transfer_ids + if count > _MAX_UNRESOLVED_TRANSFER_ID_WARNINGS and count % 1000: + return + params = getattr(request, "ec_transfer_params", None) + items = (params or {}).get("ec_items") or [] + logger.warning( + "EC Mooncake could not resolve a transfer id at %s for req=%s " + "index=%d: feature identifier=%s does not match any of the %d " + "ec_items %s (ec_transfer_params present=%s). Occurrence %d. " + "ec_items[].mm_hash must echo an engine-reported identifier, or " + "be omitted so the entry is matched by position.", + where, + request.request_id, + index, + request.mm_features[index].identifier[:16], + len(items), + [str(item.get("mm_hash"))[:16] for item in items[:4]], + params is not None, + count, + ) + + @staticmethod + def _request_transfer_id(request: Any, index: int) -> str | None: + params = getattr(request, "ec_transfer_params", None) or {} + items = params.get("ec_items") or [] + mm_hash = request.mm_features[index].identifier + if index < len(items): + item = items[index] + if item.get("mm_hash") in (None, mm_hash) and item.get("transfer_id"): + return str(item["transfer_id"]) + for item in items: + if item.get("mm_hash") == mm_hash and item.get("transfer_id"): + return str(item["transfer_id"]) + return None + + def ensure_cache_available(self, request: Any, num_computed_tokens: int) -> bool: + if self._is_producer: + for index, feature in enumerate(request.mm_features): + if feature.mm_position.offset + feature.mm_position.length > num_computed_tokens: + self._prepare_push_spec(request, index) + if not self._is_consumer: + return True + + self._drain_push_notifications() + # One timestamp for the whole decision, so the deadlines it hands out + # cannot disagree between features. + now = time.monotonic() + all_ready = True + for index, feature in enumerate(request.mm_features): + if feature.mm_position.offset + feature.mm_position.length <= num_computed_tokens: + continue + mm_hash = feature.identifier + transfer_id = self._request_transfer_id(request, index) + if transfer_id is not None: + self._transfers.touch_available(transfer_id, now + _RESERVATION_TTL_SECONDS) + if mm_hash in self._local_cache: + continue + if self._transfers.has_state(mm_hash, (SchedulerTransferState.READY,)): + continue + if self._transfers.has_state(mm_hash, (SchedulerTransferState.LOADING,)): + all_ready = False + continue + record = self._transfers.begin_load( + mm_hash, + transfer_id, + request.request_id, + now + self._push_wait_timeout, + ) + if record is not None: + self._scheduler_pending_work = True + all_ready = False + else: + waiting_id = transfer_id or f"{request.request_id}:{index}" + self._note_awaiting_push(mm_hash, waiting_id, request.request_id, now) + all_ready = False + return all_ready + + def _prepare_push_spec(self, request: Any, index: int) -> None: + params = getattr(request, "ec_transfer_params", None) or {} + consumer_zmq = params.get("consumer_zmq") + mm_hash = request.mm_features[index].identifier + transfer_id = self._request_transfer_id(request, index) + if transfer_id is None: + # Still push: after the proxy has rewritten the item to embeds the + # consumer has no media left to fall back on, so a nameless push + # beats none. But the consumer knows this transfer by the id it + # sent, not by the one invented here, so its cancel will never + # reach the reservation this push is about to take. + if consumer_zmq: + self._warn_unresolved_transfer_id(request, index, "push prepare") + transfer_id = f"{request.request_id}:{index}" + if not consumer_zmq or transfer_id in self._prepared_push_transfer_ids: + return + num_tokens = request.get_num_encoder_embeds(index) + dtype = self._model_config.dtype + assert isinstance(dtype, torch.dtype) + assert self._encoder_cache_hidden_dim is not None + dtype_name = str(dtype).split(".")[-1] + shape = (num_tokens, self._encoder_cache_hidden_dim) + nbytes = math.prod(shape) * dtype.itemsize + self._pushes_to_prepare[transfer_id] = ECMooncakePushSpec( + mm_hash=mm_hash, + nbytes=nbytes, + shape=shape, + dtype=dtype_name, + consumer_zmq=str(consumer_zmq), + transfer_id=transfer_id, + request_id=request.request_id, + ) + self._prepared_push_transfer_ids.add(transfer_id) + + def update_state_after_alloc(self, request: Any, index: int) -> None: + self._local_cache.add(request.mm_features[index].identifier) + if self._is_producer: + self._prepare_push_spec(request, index) + + def update_state_after_free(self, request: Any, index: int) -> None: + if not self._is_consumer: + return + transfer_id = self._request_transfer_id(request, index) + if transfer_id is None: + self._warn_unresolved_transfer_id(request, index, "encoder-cache free") + return + self._queue_cancel( + transfer_id, + mm_hash=request.mm_features[index].identifier, + request_id=request.request_id, + ) + + def build_connector_meta(self, scheduler_output: SchedulerOutput) -> ECConnectorMetadata: + for mm_hash in scheduler_output.free_encoder_mm_hashes: + self._local_cache.discard(mm_hash) + self._transfers.release_ready(mm_hash, time.monotonic()) + for transfer_id in self._transfers.drain_orphaned(): + self._queue_cancel(transfer_id) + meta = ECMooncakeConnectorMetadata(freed=scheduler_output.free_encoder_mm_hashes) + for push_spec in self._pushes_to_prepare.values(): + meta.pushes.append(push_spec) + self._pushes_to_prepare.clear() + for record in self._transfers.take_loads_to_dispatch(): + assert record.spec is not None + meta.loads.append(record.spec) + self._poll_pending_cancels() + self._drain_pending = True + return meta + + def update_connector_output(self, connector_output: ECConnectorOutput) -> None: + meta = connector_output.ec_connector_worker_meta + if not isinstance(meta, ECMooncakeWorkerMetadata): + return + for mm_hash in meta.loaded: + self._transfers.complete_load(mm_hash) + for mm_hash in meta.failed_loads: + self._transfers.fail_load(mm_hash, time.monotonic()) + for mm_hash in meta.reclaimed: + self._transfers.reclaim(mm_hash, time.monotonic()) + self._scheduler_pending_work = meta.pending_saves + self._failed_saves.update(meta.failed_saves) + + def has_pending_push_work(self) -> bool: + return self._scheduler_pending_work + + def request_finished(self, request: Any) -> tuple[bool, dict[str, Any] | None]: + if self._is_consumer: + for index in range(len(request.mm_features)): + transfer_id = self._request_transfer_id(request, index) + if transfer_id is None: + self._warn_unresolved_transfer_id(request, index, "request finish") + continue + self._queue_cancel( + transfer_id, + mm_hash=request.mm_features[index].identifier, + request_id=request.request_id, + ) + if self._is_producer and self._prepared_push_transfer_ids: + for index in range(len(request.mm_features)): + transfer_id = self._request_transfer_id(request, index) + if transfer_id is None: + transfer_id = f"{request.request_id}:{index}" + self._prepared_push_transfer_ids.discard(transfer_id) + if not self._is_producer: + return False, None + + items = collect_ec_item_metadata(request.mm_features, self._metadata_resolver) + for index, feature in enumerate(request.mm_features): + transfer_id = self._request_transfer_id(request, index) + if transfer_id is not None: + items[feature.identifier]["transfer_id"] = transfer_id + + if not items: + return False, None + return False, items + + def close(self) -> None: + self._control_executor.shutdown(wait=True, cancel_futures=True) + self._control_client.close() + self._event_inbox.close() diff --git a/aphrodite/distributed/ec_transfer/ec_connector/mooncake/state.py b/aphrodite/distributed/ec_transfer/ec_connector/mooncake/state.py new file mode 100644 index 0000000000..5c3d955f70 --- /dev/null +++ b/aphrodite/distributed/ec_transfer/ec_connector/mooncake/state.py @@ -0,0 +1,428 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +"""Scheduler-owned lifecycle and indexes for Consumer-bound transfers.""" + +from __future__ import annotations + +from collections import OrderedDict, deque +from collections.abc import Iterable +from dataclasses import dataclass, field, replace +from enum import Enum, auto + +from aphrodite.distributed.ec_transfer.ec_connector.mooncake.metadata import ( + ECMooncakeLoadSpec, +) +from aphrodite.logger import init_logger + +logger = init_logger(__name__) + + +class SchedulerTransferState(Enum): + """Lifecycle states visible to the Scheduler. + + The states cover waiting for a Consumer event, dispatching a Worker load, + retaining a locally reusable tensor, and terminal failure conditions. + """ + + WAITING_EVENT = auto() + AVAILABLE = auto() + LOADING = auto() + READY = auto() + RESIDENT = auto() + UNAVAILABLE = auto() + EXPIRED = auto() + FAILED = auto() + CANCELLED = auto() + + +@dataclass +class SchedulerTransfer: + """Track one transfer as observed by the Scheduler.""" + + transfer_id: str + request_id: str + mm_hash: str + state: SchedulerTransferState + spec: ECMooncakeLoadSpec | None + deadline: float | None + notified_requests: set[str] = field(default_factory=set, repr=False) + + +_TERMINAL_STATES = { + SchedulerTransferState.UNAVAILABLE, + SchedulerTransferState.EXPIRED, + SchedulerTransferState.FAILED, + SchedulerTransferState.CANCELLED, +} + + +class SchedulerTransferTable: + """Own Scheduler transfer state, lookup indexes, and dispatch queues. + + The Scheduler is single-threaded, so direct transitions are sufficient; + the Producer and Consumer managers retain stricter transition matrices + because callbacks and control requests can race there. + """ + + def __init__(self, resident_capacity: int, tombstone_ttl: float) -> None: + self._resident_capacity = resident_capacity + self._tombstone_ttl = tombstone_ttl + self._records: OrderedDict[str, SchedulerTransfer] = OrderedDict() + self._active_ids: dict[str, None] = {} + self._terminal_ids: OrderedDict[str, None] = OrderedDict() + self._resident_ids: OrderedDict[str, None] = OrderedDict() + self._resident_bytes = 0 + self._hash_index: dict[str, deque[str]] = {} + self._loads_to_dispatch: OrderedDict[str, None] = OrderedDict() + self._unavailable_requests: set[str] = set() + # Records expired straight out of READY/RESIDENT: the consumer worker + # still holds their push reservation, and nothing else tells it to let + # go, so the destination buffer would sit pinned until the lease TTL. + self._orphaned: list[str] = [] + + def get(self, transfer_id: str) -> SchedulerTransfer | None: + return self._records.get(transfer_id) + + def drain_orphaned(self) -> list[str]: + """Transfer IDs expired out of READY/RESIDENT since the last drain. + + The caller must release the consumer worker's reservation for each: + `cancel()` still accepts an EXPIRED record, so the usual cancel path + applies. + """ + orphaned = self._orphaned + self._orphaned = [] + return orphaned + + def records_for_hash( + self, + mm_hash: str, + states: Iterable[SchedulerTransferState], + ) -> list[SchedulerTransfer]: + wanted = set(states) + return [ + record + for transfer_id in self._hash_index.get(mm_hash, ()) + if (record := self._records.get(transfer_id)) is not None and record.state in wanted + ] + + def first_for_hash( + self, + mm_hash: str, + states: Iterable[SchedulerTransferState], + ) -> SchedulerTransfer | None: + states = tuple(states) + if len(states) == 1: + wanted_state = states[0] + for transfer_id in self._hash_index.get(mm_hash, ()): + record = self._records.get(transfer_id) + if record is not None and record.state is wanted_state: + return record + return None + wanted_states = set(states) + for transfer_id in self._hash_index.get(mm_hash, ()): + record = self._records.get(transfer_id) + if record is not None and record.state in wanted_states: + return record + return None + + def has_state(self, mm_hash: str, states: Iterable[SchedulerTransferState]) -> bool: + return self.first_for_hash(mm_hash, states) is not None + + def wait_for_event( + self, + transfer_id: str, + request_id: str, + mm_hash: str, + deadline: float, + ) -> SchedulerTransfer | None: + """Track a request waiting for a push. + + Returns None when `transfer_id` already names a different encoding. + The id comes from the request, so a collision has to fail that one + request; re-issuing it re-runs the encode under a fresh id. + """ + record = self._records.get(transfer_id) + if record is None: + record = SchedulerTransfer( + transfer_id=transfer_id, + request_id=request_id, + mm_hash=mm_hash, + state=SchedulerTransferState.WAITING_EVENT, + spec=None, + deadline=deadline, + ) + self._insert(record) + elif not self._identity_matches(record, mm_hash): + self._refuse_colliding_id(record, mm_hash, request_id) + return None + else: + if not record.request_id: + record.request_id = request_id + if record.state in _TERMINAL_STATES and request_id: + self._notify_unavailable(record, request_id) + return record + + def observe_ready(self, spec: ECMooncakeLoadSpec, deadline: float) -> tuple[SchedulerTransfer, bool]: + transfer_id = spec.transfer_id + record = self._records.get(transfer_id) + if record is None: + record = SchedulerTransfer( + transfer_id=transfer_id, + request_id="", + mm_hash=spec.mm_hash, + state=SchedulerTransferState.WAITING_EVENT, + spec=None, + deadline=None, + ) + self._insert(record) + elif not self._identity_matches(record, spec.mm_hash): + self._refuse_colliding_id(record, spec.mm_hash, record.request_id) + return record, False + if record.state is not SchedulerTransferState.WAITING_EVENT: + return record, False + record.spec = spec + record.deadline = deadline + self._transition(record, SchedulerTransferState.AVAILABLE) + return record, True + + def touch_available(self, transfer_id: str, deadline: float) -> None: + record = self._records.get(transfer_id) + if record is not None and record.state is SchedulerTransferState.AVAILABLE: + record.deadline = deadline + + def begin_load( + self, + mm_hash: str, + transfer_id: str | None = None, + request_id: str = "", + deadline: float | None = None, + ) -> SchedulerTransfer | None: + record = self._records.get(transfer_id) if transfer_id else None + if record is not None and (record.mm_hash != mm_hash or record.state is not SchedulerTransferState.AVAILABLE): + record = None + if record is None: + record = self.first_for_hash( + mm_hash, + ( + SchedulerTransferState.AVAILABLE, + SchedulerTransferState.RESIDENT, + ), + ) + if record is None or record.spec is None: + return None + if not record.request_id: + record.request_id = request_id + # A dispatched load that never reports back would otherwise hold this + # hash in LOADING for good, deferring every later request for it. + record.deadline = deadline + self._transition(record, SchedulerTransferState.LOADING) + self._loads_to_dispatch[record.transfer_id] = None + return record + + def take_loads_to_dispatch(self) -> list[SchedulerTransfer]: + records = [ + record + for transfer_id in self._loads_to_dispatch + if (record := self._records.get(transfer_id)) is not None and record.state is SchedulerTransferState.LOADING + ] + self._loads_to_dispatch.clear() + return records + + def complete_load(self, mm_hash: str) -> bool: + record = self.first_for_hash(mm_hash, (SchedulerTransferState.LOADING,)) + if record is None: + return self.has_state(mm_hash, (SchedulerTransferState.READY,)) + self._transition(record, SchedulerTransferState.READY) + return True + + def fail_load(self, mm_hash: str, now: float) -> bool: + record = self.first_for_hash(mm_hash, (SchedulerTransferState.LOADING,)) + if record is None: + return self.has_state(mm_hash, (SchedulerTransferState.FAILED,)) + self._transition(record, SchedulerTransferState.FAILED, now=now) + return True + + def release_ready(self, mm_hash: str, now: float) -> None: + ready = self.records_for_hash(mm_hash, (SchedulerTransferState.READY,)) + if ready: + canonical = ready[-1] + for record in self.records_for_hash( + mm_hash, + (SchedulerTransferState.READY, SchedulerTransferState.RESIDENT), + ): + if record is not canonical: + self._transition(record, SchedulerTransferState.EXPIRED, now=now) + if canonical.spec is None: + self._transition(canonical, SchedulerTransferState.EXPIRED, now=now) + else: + canonical.spec = replace(canonical.spec, local=True) + self._transition(canonical, SchedulerTransferState.RESIDENT) + self._evict_residents(now) + + def reclaim(self, mm_hash: str, now: float) -> None: + for record in self.records_for_hash( + mm_hash, + (SchedulerTransferState.READY, SchedulerTransferState.RESIDENT), + ): + if record.state is SchedulerTransferState.READY: + record.spec = None + else: + self._transition(record, SchedulerTransferState.EXPIRED, now=now) + + def mark_unavailable(self, transfer_id: str, now: float) -> None: + record = self._records[transfer_id] + self._transition(record, SchedulerTransferState.UNAVAILABLE, now=now) + if record.request_id: + self._notify_unavailable(record, record.request_id) + + def cancel( + self, + transfer_id: str, + now: float, + mm_hash: str = "", + request_id: str = "", + ) -> bool: + record = self._records.get(transfer_id) + if record is not None and mm_hash and not self._identity_matches(record, mm_hash): + return False + if record is None: + record = SchedulerTransfer( + transfer_id=transfer_id, + request_id=request_id, + mm_hash=mm_hash, + state=SchedulerTransferState.WAITING_EVENT, + spec=None, + deadline=None, + ) + self._insert(record) + if record.state in { + SchedulerTransferState.CANCELLED, + SchedulerTransferState.READY, + SchedulerTransferState.RESIDENT, + SchedulerTransferState.FAILED, + }: + return False + self._transition(record, SchedulerTransferState.CANCELLED, now=now) + self._loads_to_dispatch.pop(transfer_id, None) + return True + + def expire(self, now: float, terminal_limit: int) -> list[SchedulerTransfer]: + if terminal_limit < 0: + raise ValueError("terminal_limit must be non-negative") + expired = [] + for transfer_id in list(self._active_ids): + record = self._records[transfer_id] + if record.deadline is None or record.deadline > now: + continue + if record.state is SchedulerTransferState.AVAILABLE: + self._transition( + record, + SchedulerTransferState.EXPIRED, + now=now, + ) + expired.append(record) + elif record.state is SchedulerTransferState.LOADING: + self._transition( + record, + SchedulerTransferState.UNAVAILABLE, + now=now, + ) + if record.request_id: + self._notify_unavailable(record, record.request_id) + expired.append(record) + while self._terminal_ids: + transfer_id = next(iter(self._terminal_ids)) + record = self._records[transfer_id] + if record.deadline is not None and record.deadline > now and len(self._terminal_ids) <= terminal_limit: + break + self._remove(transfer_id) + return expired + + def take_unavailable_requests(self) -> set[str]: + unavailable = self._unavailable_requests + self._unavailable_requests = set() + return unavailable + + def _notify_unavailable(self, record: SchedulerTransfer, request_id: str) -> None: + if request_id not in record.notified_requests: + record.notified_requests.add(request_id) + self._unavailable_requests.add(request_id) + + def _insert(self, record: SchedulerTransfer) -> None: + self._records[record.transfer_id] = record + self._active_ids[record.transfer_id] = None + if record.mm_hash: + self._hash_index.setdefault(record.mm_hash, deque()).append(record.transfer_id) + + @staticmethod + def _identity_matches(record: SchedulerTransfer, mm_hash: str) -> bool: + return not record.mm_hash or record.mm_hash == mm_hash + + def _refuse_colliding_id(self, record: SchedulerTransfer, mm_hash: str, request_id: str) -> None: + """Report a transfer id that already names another encoding. + + Transfer ids arrive on the request, so a collision is reachable from + outside and must not reach the engine as an exception. + """ + logger.warning( + "EC Mooncake transfer_id=%s already names mm_hash=%s; refusing mm_hash=%s for request %s", + record.transfer_id, + record.mm_hash[:16], + mm_hash[:16], + request_id or "", + ) + if request_id: + self._unavailable_requests.add(request_id) + + def _transition( + self, + record: SchedulerTransfer, + state: SchedulerTransferState, + now: float | None = None, + ) -> None: + if record.state is SchedulerTransferState.RESIDENT: + self._resident_ids.pop(record.transfer_id, None) + if record.spec is not None: + self._resident_bytes -= record.spec.nbytes + if state is SchedulerTransferState.RESIDENT: + self._resident_ids[record.transfer_id] = None + if record.spec is not None: + self._resident_bytes += record.spec.nbytes + if state in _TERMINAL_STATES: + if now is None: + raise ValueError("Terminal transition requires a timestamp") + # Keep a bounded tombstone so late events and repeat request IDs + # remain idempotent instead of reviving a finished transfer. + record.deadline = now + self._tombstone_ttl + self._active_ids.pop(record.transfer_id, None) + self._terminal_ids[record.transfer_id] = None + self._terminal_ids.move_to_end(record.transfer_id) + if state is SchedulerTransferState.EXPIRED and record.state in ( + SchedulerTransferState.READY, + SchedulerTransferState.RESIDENT, + ): + self._orphaned.append(record.transfer_id) + record.state = state + self._records.move_to_end(record.transfer_id) + + def _evict_residents(self, now: float) -> None: + while self._resident_bytes > self._resident_capacity and self._resident_ids: + record = self._records[next(iter(self._resident_ids))] + self._transition(record, SchedulerTransferState.EXPIRED, now=now) + + def _remove(self, transfer_id: str) -> None: + record = self._records.pop(transfer_id, None) + self._active_ids.pop(transfer_id, None) + self._terminal_ids.pop(transfer_id, None) + if transfer_id in self._resident_ids: + self._resident_ids.pop(transfer_id) + if record is not None and record.spec is not None: + self._resident_bytes -= record.spec.nbytes + self._loads_to_dispatch.pop(transfer_id, None) + if record is None or not record.mm_hash: + return + transfer_ids = self._hash_index[record.mm_hash] + transfer_ids.remove(transfer_id) + if not transfer_ids: + self._hash_index.pop(record.mm_hash) diff --git a/aphrodite/distributed/ec_transfer/ec_connector/mooncake/transfer.py b/aphrodite/distributed/ec_transfer/ec_connector/mooncake/transfer.py new file mode 100644 index 0000000000..d49b9d85b5 --- /dev/null +++ b/aphrodite/distributed/ec_transfer/ec_connector/mooncake/transfer.py @@ -0,0 +1,198 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +"""Mooncake data-plane engine and memory-registration ownership. + +The wrapper isolates optional-engine initialization, session addressing, +synchronous batched writes, and the lifetime of transient source registrations. +""" + +from __future__ import annotations + +import threading +from dataclasses import dataclass + +import torch + +from aphrodite.logger import init_logger + +logger = init_logger(__name__) + +_MOONCAKE_IMPORT_ERROR: ImportError | None = None +try: + from mooncake.engine import TransferEngine +except ImportError as error: + TransferEngine = None # type: ignore[misc, assignment] + _MOONCAKE_IMPORT_ERROR = error + + +def ensure_mooncake_available() -> None: + if _MOONCAKE_IMPORT_ERROR is not None: + raise ImportError("Install mooncake-transfer-engine to use ECMooncakeConnector.") from _MOONCAKE_IMPORT_ERROR + + +@dataclass +class _SourceRegistration: + """Retain one transient source range while batches reference it.""" + + tensor: torch.Tensor + nbytes: int + users: int = 1 + + +class MooncakeTransfer: + """Own a lazy Mooncake engine and transient memory registrations.""" + + def __init__(self, hostname: str, protocol: str) -> None: + self._hostname = hostname + self._protocol = protocol + self._engine: TransferEngine | None = None + self._engine_lock = threading.Lock() + self._source_registrations: dict[int, _SourceRegistration] = {} + self._pending_unregister: dict[int, torch.Tensor] = {} + self._registration_lock = threading.Lock() + self._closed = False + + def _ensure_engine(self) -> TransferEngine: + if self._engine is not None: + return self._engine + with self._engine_lock: + if self._engine is not None: + return self._engine + engine = TransferEngine() + ret = engine.initialize(self._hostname, "P2PHANDSHAKE", self._protocol, "") + if ret != 0: + raise RuntimeError("Mooncake TransferEngine initialization failed.") + self._engine = engine + logger.info( + "ECMooncakeConnector TransferEngine ready at %s:%d", + self._hostname, + engine.get_rpc_port(), + ) + return self._engine + + def ensure_ready(self) -> None: + self._ensure_engine() + + def local_session(self) -> str: + engine = self._ensure_engine() + return f"{self._hostname}:{engine.get_rpc_port()}" + + def register_memory(self, tensor: torch.Tensor) -> int: + return self._ensure_engine().batch_register_memory([tensor.data_ptr()], [tensor.nbytes]) + + def unregister_memory(self, tensor: torch.Tensor) -> bool: + engine = self._ensure_engine() + address = tensor.data_ptr() + ret = engine.unregister_memory(address) + with self._registration_lock: + if ret != 0: + logger.error( + "Mooncake EC memory unregistration failed for address %d: %d", + address, + ret, + ) + self._pending_unregister[address] = tensor + return False + self._pending_unregister.pop(address, None) + return True + + @staticmethod + def _source_range(tensor: torch.Tensor) -> tuple[int, int]: + # Encoder batches commonly split one storage into sibling tensor views. + # Registering each view's exact bytes avoids overlapping memory regions. + return tensor.data_ptr(), tensor.nbytes + + def acquire_sources(self, tensors: list[torch.Tensor]) -> list[int]: + ranges: dict[int, tuple[int, torch.Tensor]] = {} + for tensor in tensors: + address, nbytes = self._source_range(tensor) + ranges.setdefault(address, (nbytes, tensor)) + + engine = self._ensure_engine() + acquired: list[int] = [] + new_addresses: list[int] = [] + new_lengths: list[int] = [] + with self._registration_lock: + for address, (nbytes, tensor) in ranges.items(): + entry = self._source_registrations.get(address) + if entry is not None: + if entry.nbytes != nbytes: + raise RuntimeError("Mooncake EC source storage changed size while registered") + entry.users += 1 + acquired.append(address) + continue + new_addresses.append(address) + new_lengths.append(nbytes) + self._source_registrations[address] = _SourceRegistration( + tensor=tensor, + nbytes=nbytes, + ) + acquired.append(address) + + if new_addresses: + ret = engine.batch_register_memory(new_addresses, new_lengths) + if ret != 0: + for address in acquired: + entry = self._source_registrations[address] + entry.users -= 1 + if entry.users == 0: + del self._source_registrations[address] + raise RuntimeError("Mooncake EC source registration failed") + return acquired + + def release_sources(self, addresses: list[int]) -> bool: + if not addresses: + return True + with self._registration_lock: + unused = [] + for address in addresses: + entry = self._source_registrations.get(address) + if entry is None: + continue + entry.users -= 1 + if entry.users == 0: + unused.append(address) + if not unused: + return True + ret = self._ensure_engine().batch_unregister_memory(unused) + if ret != 0: + logger.warning( + "Keeping %d EC source tensors registered after Mooncake unregistration failure", + len(unused), + ) + return False + for address in unused: + del self._source_registrations[address] + self._pending_unregister.pop(address, None) + return True + + def write( + self, + session: str, + sources: list[int], + destinations: list[int], + lengths: list[int], + ) -> None: + """Write one synchronous batch, returning only at terminal status.""" + ret = self._ensure_engine().batch_transfer_sync_write(session, sources, destinations, lengths) + if ret != 0: + raise RuntimeError(f"Mooncake EC push to {session} failed with status {ret}") + + def close(self) -> None: + if self._closed: + return + self._closed = True + engine = self._engine + if engine is None: + return + with self._registration_lock: + addresses = list(self._source_registrations) + addresses.extend(self._pending_unregister) + if not addresses: + return + ret = engine.batch_unregister_memory(list(dict.fromkeys(addresses))) + if ret != 0: + logger.error("Mooncake EC batch memory unregistration failed: %d", ret) + return + self._source_registrations.clear() + self._pending_unregister.clear() diff --git a/aphrodite/distributed/ec_transfer/ec_connector/mooncake/worker.py b/aphrodite/distributed/ec_transfer/ec_connector/mooncake/worker.py new file mode 100644 index 0000000000..ad6ad8af99 --- /dev/null +++ b/aphrodite/distributed/ec_transfer/ec_connector/mooncake/worker.py @@ -0,0 +1,866 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +"""Worker-side orchestration of Mooncake control, memory, and data planes. + +Consumer Workers expose rank-local reservations and publish received tensors. +Producer Workers reserve every destination shard, bind computed sources, run +batched Mooncake writes, and report asynchronous completion to the Scheduler. +""" + +from __future__ import annotations + +import math +import threading +import time +from collections.abc import Callable +from concurrent.futures import Future, ThreadPoolExecutor +from functools import partial +from typing import TYPE_CHECKING, Any, TypeVar, cast + +import torch + +from aphrodite.distributed.ec_transfer.ec_connector.mooncake.config import ( + _RESERVATION_TTL_SECONDS, + MooncakeECConfig, +) +from aphrodite.distributed.ec_transfer.ec_connector.mooncake.control import ( + ConsumerControlServer, + ControlClient, + make_cancel_request, +) +from aphrodite.distributed.ec_transfer.ec_connector.mooncake.memory import ( + ConsumerMemoryPool, + ProducerMemoryPool, +) +from aphrodite.distributed.ec_transfer.ec_connector.mooncake.metadata import ( + ECMooncakeConnectorMetadata, + ECMooncakePushSpec, + ECMooncakeWorkerMetadata, +) +from aphrodite.distributed.ec_transfer.ec_connector.mooncake.producer import ( + ProducerPushManager, + ProducerPushRecord, +) +from aphrodite.distributed.ec_transfer.ec_connector.mooncake.reservation import ( + ConsumerReservationManager, + ConsumerReservationState, +) +from aphrodite.distributed.ec_transfer.ec_connector.mooncake.transfer import ( + MooncakeTransfer, + ensure_mooncake_available, +) +from aphrodite.logger import init_logger +from aphrodite.utils.network_utils import get_ip + +logger = init_logger(__name__) + +_T = TypeVar("_T") + +if TYPE_CHECKING: + from aphrodite.config import AphroditeConfig + +_RESERVATION_REFRESH_SECONDS = _RESERVATION_TTL_SECONDS / 2 +_MAX_CANCELLED_TRANSFER_IDS = 1 << 16 +_CANCEL_ATTEMPTS = 2 +_TRANSFER_WORKERS = 4 +_CONTROL_WORKERS = 8 + + +class _FanoutError(RuntimeError): + """Retain shard outcomes after every started task settles.""" + + def __init__(self, error: BaseException, results: list[Any | None]) -> None: + self.results = results + super().__init__(str(error)) + + +class ECMooncakeWorker: + """Orchestrate consumer reservations and producer push batches. + + Consumers may use TP, PP, and DP. Only the first PP stage receives encoder + outputs, and each TP rank exposes a consecutive control port and receives + the same source concurrently. Producers remain unsharded and unreplicated. + With DP, the caller must route both halves of a request to the same replica + and pass that replica's control address to the producer. + """ + + def __init__(self, aphrodite_config: AphroditeConfig) -> None: + ensure_mooncake_available() + config = MooncakeECConfig.from_aphrodite_config(aphrodite_config) + self.is_producer = config.is_producer + self.is_consumer = config.is_consumer + self._buffer_device = config.buffer_device + self._control_host = config.control_host + self._control_port = config.control_port + self._transfer = MooncakeTransfer(get_ip(), config.protocol) + self._consumer_memory = ConsumerMemoryPool( + config.pool_size, + self._transfer, + ) + self._reservations = ConsumerReservationManager( + self._consumer_memory, + _RESERVATION_TTL_SECONDS, + _MAX_CANCELLED_TRANSFER_IDS, + ) + self._consumer_rank_resolved = False + self._is_receiving_rank = True + self._tp_rank = 0 + self._tp_size = 1 + self._control_server: ConsumerControlServer | None = None + # Worker producer + self._producer_memory = ProducerMemoryPool( + config.pool_size, + self._transfer, + ) + self._control_client = ControlClient(config.control_timeout_ms) + self._shutdown_drain_timeout_s = config.control_timeout_ms / 1000 + self._io_executor = ThreadPoolExecutor( + max_workers=_TRANSFER_WORKERS, + thread_name_prefix="ec-mooncake-transfer", + ) + self._control_executor = ThreadPoolExecutor( + max_workers=_CONTROL_WORKERS, + thread_name_prefix="ec-mooncake-control", + ) + self._shard_pool: ThreadPoolExecutor | None = None + self._shard_pool_lock = threading.Lock() + self._push_ready = threading.Event() + self._producer_pushes = ProducerPushManager(self._push_ready.set) + self._dispatch_stop = threading.Event() + self._dispatcher: threading.Thread | None = None + self._failed_saves: set[str] = set() + self._collecting_sources = False + self._completed_loads: set[str] = set() + self._failed_loads: set[str] = set() + self._shutdown = False + if self.is_producer: + self._dispatcher = threading.Thread(target=self._dispatch_pushes, name="ec-mooncake-ready", daemon=True) + self._dispatcher.start() + + def _dispatch_pushes(self) -> None: + pending_event = False + while not self._dispatch_stop.is_set(): + self._push_ready.wait(timeout=0.001 if pending_event else None) + self._push_ready.clear() + if self._dispatch_stop.is_set(): + break + if self._collecting_sources: + pending_event = False + continue + pending_event = self._producer_pushes.submit_batches(self._io_executor, self._push_batch) + + def _resolve_consumer_rank(self) -> None: + """Place this worker in the consumer receive topology.""" + if self._consumer_rank_resolved: + return + self._consumer_rank_resolved = True + try: + from aphrodite.distributed.parallel_state import get_pp_group, get_tp_group + + tp_group = get_tp_group() + self._tp_rank = tp_group.rank_in_group + self._tp_size = tp_group.world_size + self._is_receiving_rank = get_pp_group().is_first_rank + except AssertionError: + # Groups are only absent outside a distributed run, where this + # worker is the whole consumer. + self._tp_rank = 0 + self._tp_size = 1 + self._is_receiving_rank = True + + def start_services(self) -> None: + if not self.is_consumer or self._control_server is not None: + return + self._resolve_consumer_rank() + if not self._is_receiving_rank: + # Later pipeline stages hold no encoder outputs, so they need + # neither a receive pool nor a control channel. + return + self._consumer_memory.prepare(torch.device(self._buffer_device)) + consumer_pool = self._consumer_memory.tensor + if consumer_pool is None: + raise RuntimeError("Mooncake push mode requires a registered consumer buffer pool.") + base_port = self._control_port + self._control_server = ConsumerControlServer( + # The reservation channel hands out registered-memory addresses + # and takes cancellations, so it listens where the Consumer + # advertises itself (`ec_ip`), not on every interface. + self._control_host, + base_port + self._tp_rank, + self._reserve_push_destination, + self._push_status, + self._reservations.complete, + self._reservations.cancel, + self._reservations.expire, + peer_ports=[base_port + rank for rank in range(self._tp_size)], + device=consumer_pool.device, + drain_ready=self._reservations.drain_ready, + ) + try: + self._control_server.start() + except Exception: + self._control_server.close() + self._control_server = None + raise + + def _reserve_push_destination(self, payload: dict[str, Any]) -> dict[str, Any]: + transfer_id = str(payload["transfer_id"]) + mm_hash = str(payload["mm_hash"]) + nbytes = int(payload["nbytes"]) + shape = tuple(int(value) for value in payload["shape"]) + dtype_name = str(payload["dtype"]) + dtype = getattr(torch, dtype_name, None) + if not isinstance(dtype, torch.dtype): + raise ValueError(f"Unsupported torch dtype string: {dtype_name!r}") + expected_nbytes = math.prod(shape) * dtype.itemsize + if expected_nbytes != nbytes: + raise ValueError("shape and dtype do not match nbytes") + + self._reservations.expire() + reservation, should_write = self._reservations.reserve(transfer_id, mm_hash, nbytes, shape, dtype_name, dtype) + if reservation is None: + raise RuntimeError("EC consumer buffer pool is full") + if reservation.state in { + ConsumerReservationState.CANCEL_PENDING, + ConsumerReservationState.CANCELLED, + }: + return { + "reservation_id": "", + "dst_session": "", + "dst_ptr": 0, + "nbytes": nbytes, + "write": False, + "ready": False, + "cancelled": True, + } + assert reservation.allocation is not None + + return { + "reservation_id": reservation.reservation_id, + "dst_session": self._transfer.local_session(), + "dst_ptr": reservation.allocation.tensor.data_ptr(), + "nbytes": reservation.allocation.tensor.nbytes, + "write": should_write, + "ready": reservation.state is ConsumerReservationState.READY, + "cached": reservation.lease is not None, + } + + def _push_status(self, transfer_id: str) -> dict[str, Any] | None: + reservation = self._reservations.status(transfer_id) + if reservation is None: + return None + assert reservation.allocation is not None + return { + "mm_hash": reservation.mm_hash, + "ready": reservation.state is ConsumerReservationState.READY, + "reservation_id": reservation.reservation_id, + "nbytes": reservation.allocation.tensor.nbytes, + "shape": list(reservation.shape), + "dtype": reservation.dtype, + } + + def _shard_executor(self) -> ThreadPoolExecutor: + """Use a separate pool so nested shard fan-out cannot deadlock.""" + with self._shard_pool_lock: + if self._shard_pool is None: + self._shard_pool = ThreadPoolExecutor(max_workers=32, thread_name_prefix="ec-mooncake-shard") + return self._shard_pool + + def _reserve_one(self, addr: str, spec: ECMooncakePushSpec) -> dict[str, Any]: + result = self._control_client.request( + addr, + { + "op": "reserve", + "transfer_id": spec.transfer_id, + "mm_hash": spec.mm_hash, + "nbytes": spec.nbytes, + "shape": list(spec.shape), + "dtype": spec.dtype, + }, + ) + if not isinstance(result, dict): + raise RuntimeError("Invalid EC reservation response") + result["_received_at"] = time.monotonic() + result["addr"] = addr + return result + + def _run_fanout( + self, + tasks: list[Callable[[], _T]], + on_submit: Callable[[int, Future[_T]], None] | None = None, + ) -> list[_T]: + """Run every started shard task and retain partial results on failure. + + Waiting for all submitted tasks is what makes source-memory release and + partial-reservation cleanup safe after one shard fails. + """ + if not tasks: + return [] + futures: list[tuple[int, Future[_T]]] = [] + results: list[_T | None] = [None] * len(tasks) + error: BaseException | None = None + for index, task in enumerate(tasks[1:], 1): + try: + future = self._shard_executor().submit(task) + except Exception as exc: + error = exc + break + futures.append((index, future)) + if on_submit is not None: + on_submit(index, future) + if error is None: + try: + results[0] = tasks[0]() + except Exception as exc: + error = exc + for index, future in futures: + try: + results[index] = future.result() + except Exception as exc: + if error is None: + error = exc + if error is not None: + raise _FanoutError(error, results) + return cast(list[_T], results) + + def _cancel_reservations( + self, + spec: ECMooncakePushSpec, + reservations: list[dict[str, Any]], + *, + refresh: bool = False, + record: ProducerPushRecord | None = None, + ) -> None: + reservations = [shard for shard in reservations if not shard.get("cancelled", False)] + if not reservations: + return + + def cancel(shard: dict[str, Any]) -> dict[str, Any]: + addr = shard.get("addr") + if not isinstance(addr, str) or not addr: + raise RuntimeError("EC reservation is missing a confirmed shard address") + result = self._control_client.request( + addr, + make_cancel_request( + spec.transfer_id, + str(shard.get("reservation_id", "")), + abandon=True, + refresh=refresh, + ), + ) + if not isinstance(result, dict) or not result.get("cancelled"): + raise RuntimeError(f"Could not cancel EC reservation for mm_hash={spec.mm_hash}") + return shard + + def track(_index: int, future: Future[dict[str, Any]]) -> None: + if record is not None: + self._producer_pushes.track_shard_futures([record], [future]) + + self._run_fanout([partial(cancel, shard) for shard in reservations], track) + + def _retry_cancel_reservations( + self, + spec: ECMooncakePushSpec, + reservations: list[dict[str, Any]], + *, + record: ProducerPushRecord | None = None, + ) -> None: + pending = [shard for shard in reservations if not shard.get("cancelled", False)] + error: _FanoutError | None = None + for _ in range(_CANCEL_ATTEMPTS): + try: + self._cancel_reservations(spec, pending, record=record) + except _FanoutError as exc: + error = exc + pending = [shard for index, shard in enumerate(pending) if exc.results[index] is None] + continue + return + assert error is not None + raise error + + def _reserve_remote(self, spec: ECMooncakePushSpec) -> list[dict[str, Any]]: + """Reserve a destination on every shard of the consumer.""" + result = self._reserve_remote_many([spec])[0] + if isinstance(result, Exception): + raise result + return result + + def _refresh_remote_reservations( + self, + spec: ECMooncakePushSpec, + reservations: list[dict[str, Any]], + record: ProducerPushRecord | None = None, + ) -> list[dict[str, Any]]: + stale = [ + shard + for shard in reservations + if not shard.get("ready", False) and not shard.get("cached", False) and not shard.get("cancelled", False) + ] + try: + self._cancel_reservations(spec, stale, refresh=True, record=record) + except _FanoutError as exc: + pending = [shard for index, shard in enumerate(stale) if exc.results[index] is None] + try: + self._retry_cancel_reservations(spec, pending, record=record) + except _FanoutError as cleanup_error: + raise exc from cleanup_error + raise + return self._reserve_remote(spec) + + @staticmethod + def _validate_push_source(push: ProducerPushRecord) -> None: + tensor = push.source_tensor + assert tensor is not None + spec = push.spec + if tuple(tensor.shape) != tuple(spec.shape): + raise ValueError(f"EC source shape mismatch for mm_hash={spec.mm_hash}") + if str(tensor.dtype).split(".")[-1] != spec.dtype: + raise ValueError(f"EC source dtype mismatch for mm_hash={spec.mm_hash}") + if not tensor.is_contiguous(): + raise ValueError(f"EC source must be contiguous for mm_hash={spec.mm_hash}") + if tensor.nbytes != spec.nbytes: + raise ValueError(f"EC source size mismatch for mm_hash={spec.mm_hash}") + + def start_save_caches( + self, + metadata: ECMooncakeConnectorMetadata, + encoder_cache: dict[str, torch.Tensor] | None = None, + **kwargs: Any, + ) -> None: + self._collecting_sources = True + new: dict[str, list[ProducerPushRecord]] = {} + for spec in metadata.pushes: + try: + record, created = self._producer_pushes.reserve(spec, Future) + if created: + new.setdefault(spec.consumer_zmq, []).append(record) + except ValueError: + logger.warning("Rejected conflicting EC push", exc_info=True) + self._failed_saves.add(spec.request_id) + for records in new.values(): + future = self._control_executor.submit(self._reserve_batch, records) + future.add_done_callback(partial(self._reservation_batch_done, records)) + if not isinstance(encoder_cache, dict): + return + for mm_hash in dict.fromkeys(spec.mm_hash for spec in metadata.pushes): + tensor = encoder_cache.get(mm_hash) + if tensor is not None: + self._bind_push_source(tensor, mm_hash) + + def _reserve_batch(self, records: list[ProducerPushRecord]) -> None: + """Resolve independent item futures with one reserve RPC per shard.""" + if len(records) == 1: + record = records[0] + try: + record.reservation_future.set_result(self._reserve_remote(record.spec)) + except Exception as exc: + record.reservation_future.set_exception(exc) + return + outcomes = self._reserve_remote_many([record.spec for record in records]) + self._producer_pushes.finish_reservations(list(zip(records, outcomes))) + + def _reserve_remote_many(self, specs: list[ECMooncakePushSpec]) -> list[list[dict[str, Any]] | Exception]: + shards = None + for _ in range(_CANCEL_ATTEMPTS): + shards = self._control_client.discover_shards(specs[0].consumer_zmq) + if shards is not None: + break + if shards is None: + raise RuntimeError(f"Could not discover every EC consumer shard at {specs[0].consumer_zmq}") + + def reserve(addr: str) -> list[dict[str, Any]]: + if len(specs) == 1: + return [{"ok": True, "result": self._reserve_one(addr, specs[0])}] + response = self._control_client.request( + addr, + { + "op": "reserve_batch", + "items": [ + { + "transfer_id": spec.transfer_id, + "mm_hash": spec.mm_hash, + "nbytes": spec.nbytes, + "shape": list(spec.shape), + "dtype": spec.dtype, + } + for spec in specs + ], + }, + ) + items = response["items"] + if len(items) != len(specs): + raise RuntimeError("Malformed EC reservation batch response") + for item in items: + if item.get("ok"): + item["result"]["_received_at"] = time.monotonic() + return items + + fanout_error = None + results: list[Any] + try: + results = self._run_fanout([partial(reserve, addr) for addr in shards]) + except _FanoutError as exc: + results = exc.results + fanout_error = exc + outcomes: list[list[dict[str, Any]] | Exception] = [] + for index, spec in enumerate(specs): + reservations: list[dict[str, Any]] = [] + error = None + for addr, items in zip(shards, results): + item = items[index] if items is not None else {} + if item.get("ok"): + reservations.append({**item["result"], "addr": addr}) + else: + error = fanout_error or RuntimeError(item["error"]) + reservations.append({"addr": addr, "reservation_id": ""}) + if error is None: + outcomes.append(reservations) + else: + failure = _FanoutError(error, list(reservations)) + try: + self._retry_cancel_reservations(spec, reservations) + except Exception as cleanup_error: + failure.__cause__ = cleanup_error + logger.exception("Failed to release partial EC reservation batch") + outcomes.append(failure) + return outcomes + + @staticmethod + def _reservation_batch_done(records: list[ProducerPushRecord], future: Future[None]) -> None: + error = future.exception() + if error is not None: + for record in records: + if not record.reservation_future.done(): + record.reservation_future.set_exception(error) + + def start_load_caches( + self, + metadata: ECMooncakeConnectorMetadata, + encoder_cache: dict[str, torch.Tensor], + **kwargs: Any, + ) -> None: + self._resolve_consumer_rank() + if not self._is_receiving_rank: + # Later pipeline stages never gather multimodal embeddings. + return + self._transfer.ensure_ready() + if self._buffer_device == "cuda" and not torch.accelerator.is_available(): + raise RuntimeError("ECMooncakeConnector requires CUDA for ec_buffer_device=cuda") + self._reservations.retire_stale(encoder_cache, metadata.freed) + + for spec in metadata.loads: + if spec.mm_hash in encoder_cache: + if not spec.local: + # The spec's id is one shard's; cancel by transfer. + self._reservations.cancel(spec.transfer_id, "") + self._completed_loads.add(spec.mm_hash) + continue + if spec.local: + tensor = self._consumer_memory.take_resident(spec.mm_hash, tuple(spec.shape), spec.dtype) + else: + try: + allocation = self._reservations.take(spec.transfer_id, spec.mm_hash) + tensor = allocation.tensor + except RuntimeError as e: + logger.warning("EC Mooncake pushed load failed: %s", e) + tensor = None + if tensor is None: + self._failed_loads.add(spec.mm_hash) + else: + encoder_cache[spec.mm_hash] = tensor + self._completed_loads.add(spec.mm_hash) + + def _push_batch(self, pushes: list[ProducerPushRecord]) -> None: + started_at = time.monotonic() + ready: list[tuple[ProducerPushRecord, dict[str, Any]]] = [] + written_pushes: dict[str, ProducerPushRecord] = {} + failure: Exception | None = None + try: + for push in pushes: + self._validate_push_source(push) + reservations = self._producer_pushes.resolve_reservations(push) + stale = [ + index + for index, shard in enumerate(reservations) + if not shard.get("ready", False) + and not shard.get("cancelled", False) + and time.monotonic() - float(shard.get("_received_at", started_at)) >= _RESERVATION_REFRESH_SECONDS + ] + if stale: + reservations = self._refresh_remote_reservations(push.spec, reservations, push) + self._producer_pushes.replace_reservations(push, reservations) + self._producer_pushes.begin_writing(push) + writable = [ + shard + for shard in reservations + if not shard.get("cached", False) and not shard.get("cancelled", False) and shard.get("write", True) + ] + source = push.source_tensor + assert source is not None + for shard in writable: + if int(shard["nbytes"]) != source.nbytes: + raise RuntimeError(f"Reserved EC size does not match tensor for mm_hash={push.spec.mm_hash}") + ready.append((push, shard)) + written_pushes.setdefault(push.spec.transfer_id, push) + if ready: + # Stage each source once, then write it to every destination. + source_index = {push.spec.transfer_id: index for index, push in enumerate(written_pushes.values())} + tensors = [ + cast(torch.Tensor, push.source_tensor) + for push in written_pushes.values() + if push.source_tensor is not None + ] + lengths = [tensor.nbytes for tensor in tensors] + staged = self._producer_memory.stage(tensors) + registered_sources: list[int] = [] + if staged is not None: + sources = staged.tensors + else: + sources = tensors + registered_sources = self._transfer.acquire_sources(tensors) + addresses = [tensor.data_ptr() for tensor in sources] + try: + by_session: dict[str, list[tuple[int, int]]] = {} + session_records: dict[str, dict[str, ProducerPushRecord]] = {} + for push, shard in ready: + session = str(shard["dst_session"]) + by_session.setdefault(session, []).append( + (source_index[push.spec.transfer_id], int(shard["dst_ptr"])) + ) + session_records.setdefault(session, {})[push.spec.transfer_id] = push + + def write(session: str, items: list[tuple[int, int]]) -> None: + self._transfer.write( + session, + [addresses[index] for index, _ in items], + [dst for _, dst in items], + [lengths[index] for index, _ in items], + ) + + sessions = list(by_session.items()) + + # Write shards concurrently to avoid serial TP latency. + def track_write(index: int, future: Future[None]) -> None: + session = sessions[index][0] + self._producer_pushes.track_shard_futures(list(session_records[session].values()), [future]) + + writes: list[Callable[[], None]] = [partial(write, *session) for session in sessions] + self._run_fanout(writes, track_write) + finally: + if staged is not None: + self._producer_memory.release(staged) + self._transfer.release_sources(registered_sources) + + self._producer_pushes.begin_notifying(pushes) + self._notify_completions(ready) + self._producer_pushes.complete(pushes) + except Exception as exc: + # Report asynchronously; raising here would fail EngineCore. + failure = exc + logger.exception( + "EC Mooncake push batch failed for mm_hashes=%s", + [push.spec.mm_hash for push in pushes], + ) + self._producer_pushes.settle_all(pushes) + self._abandon_pushes(pushes) + finally: + if failure is not None: + self._producer_pushes.fail(pushes, failure) + + def _notify_completions(self, notifications: list[tuple[ProducerPushRecord, dict[str, Any]]]) -> None: + """Tell the consumer, in one message per destination, what landed.""" + if not notifications: + return + by_destination: dict[str, list[tuple[ProducerPushRecord, dict[str, Any]]]] = {} + for push, reservation in notifications: + by_destination.setdefault(str(reservation.get("addr", push.spec.consumer_zmq)), []).append( + (push, reservation) + ) + destinations = list(by_destination.items()) + + def notify( + consumer_zmq: str, + items: list[tuple[ProducerPushRecord, dict[str, Any]]], + ) -> None: + result = self._control_client.request( + consumer_zmq, + { + "op": "complete_batch", + "items": [ + { + "transfer_id": push.spec.transfer_id, + "reservation_id": reservation["reservation_id"], + } + for push, reservation in items + ], + }, + ) + completions = result.get("items", []) if isinstance(result, dict) else [] + if len(completions) != len(items): + raise RuntimeError("Malformed EC completion response") + for (push, _), completion in zip(items, completions): + if not completion.get("completed"): + raise RuntimeError(f"Unknown EC reservation for mm_hash={push.spec.mm_hash}") + + def track(index: int, future: Future[None]) -> None: + records = {push.spec.transfer_id: push for push, _ in destinations[index][1]} + self._producer_pushes.track_shard_futures(list(records.values()), [future]) + + self._run_fanout( + [partial(notify, destination, items) for destination, items in destinations], + track, + ) + + @staticmethod + def _known_reservations(record: ProducerPushRecord) -> list[dict[str, Any]]: + if record.reservations: + return list(record.reservations) + try: + return list(record.reservation_future.result()) + except _FanoutError as exc: + return [result for result in exc.results if isinstance(result, dict)] + except Exception: + return [] + + def _abandon_pushes(self, pushes: list[ProducerPushRecord]) -> None: + """Release the consumer-side reservations of a batch that failed.""" + for push in pushes: + shards = self._known_reservations(push) + if not shards: + # No confirmed topology means there is no safe address to + # cancel. The original push failure remains observable via + # ProducerPushManager.fail below. + continue + try: + self._retry_cancel_reservations(push.spec, shards, record=push) + except _FanoutError: + logger.exception( + "Failed to abandon EC reservations for transfer_id=%s", + push.spec.transfer_id, + ) + + def _flush_pending_pushes(self) -> None: + self._producer_pushes.submit_batches( + self._io_executor, + self._push_batch, + ) + + def _bind_push_source(self, tensor: torch.Tensor, mm_hash: str) -> None: + ready_event = None + if tensor.device.type == "cuda": + ready_event = torch.Event() + ready_event.record(torch.accelerator.current_stream(tensor.device)) + self._producer_pushes.bind_source(mm_hash, tensor, ready_event) + + def _cancel_orphaned_reservation(self, record: ProducerPushRecord) -> None: + resolution_error: Exception | None = None + try: + reservations = self._producer_pushes.resolve_reservations(record) + except Exception as exc: + resolution_error = exc + reservations = self._known_reservations(record) + reservations = [shard for shard in reservations if not shard.get("cancelled", False)] + if not reservations: + if resolution_error is not None: + self._producer_pushes.fail([record], resolution_error) + else: + self._producer_pushes.finish_cancel(record) + return + try: + self._retry_cancel_reservations(record.spec, reservations, record=record) + except Exception as cleanup_error: + if resolution_error is not None: + combined_error = RuntimeError( + f"EC reservation resolution failed ({resolution_error}); cleanup also failed ({cleanup_error})" + ) + combined_error.__cause__ = resolution_error + cleanup_error = combined_error + self._producer_pushes.fail([record], cleanup_error) + return + if resolution_error is not None: + self._producer_pushes.fail([record], resolution_error) + else: + self._producer_pushes.finish_cancel(record) + + def get_finished(self, finished_req_ids: set[str]) -> tuple[set[str] | None, set[str] | None]: + if not self.is_producer: + return None, None + + for record in self._producer_pushes.cancel_requests(finished_req_ids): + self._producer_pushes.submit_cancel( + record, + self._io_executor, + self._cancel_orphaned_reservation, + ) + return None, None + + def save_caches(self, encoder_cache: dict[str, torch.Tensor], mm_hash: str, **kwargs: Any) -> None: + if not self.is_producer: + return + tensor = encoder_cache[mm_hash] + self._bind_push_source(tensor, mm_hash) + + def build_connector_worker_meta(self) -> ECMooncakeWorkerMetadata | None: + if self.is_consumer and not self._is_receiving_rank: + # `loaded` is intersected across reporting ranks, so a stage that + # never loads must not report at all rather than report nothing. + return None + + self._collecting_sources = False + self._push_ready.set() + self._flush_pending_pushes() + failures = self._producer_pushes.poll() + for mm_hash, error in failures: + logger.error( + "EC Mooncake async save failed for mm_hash=%s: %s", + mm_hash, + error, + ) + reclaimed = self._consumer_memory.drain_reclaimed() + meta = ECMooncakeWorkerMetadata( + loaded=self._completed_loads, + failed_loads=self._failed_loads, + reclaimed=reclaimed, + pending_saves=self._producer_pushes.pending, + failed_saves=self._failed_saves, + ) + self._failed_saves = set() + self._completed_loads = set() + self._failed_loads = set() + return meta + + def close(self) -> None: + if self._shutdown: + return + self._shutdown = True + dispatcher = getattr(self, "_dispatcher", None) + if dispatcher is not None: + self._dispatch_stop.set() + self._push_ready.set() + dispatcher.join() + if self._control_server is not None: + self._reservations.begin_shutdown() + self._flush_pending_pushes() + for record in self._producer_pushes.cancel_requests(None): + self._producer_pushes.submit_cancel( + record, + self._io_executor, + self._cancel_orphaned_reservation, + ) + self._control_executor.shutdown(wait=True) + self._producer_pushes.submit_batches(self._io_executor, self._push_batch, wait=True) + self._io_executor.shutdown(wait=True) + if self._shard_pool is not None: + self._shard_pool.shutdown(wait=True) + # Every producer-side thread that could hold a control socket is stopped. + self._control_client.close() + drained = True + if self._control_server is not None: + drained = self._reservations.wait_for_writers(self._shutdown_drain_timeout_s) + self._control_server.close() + if drained: + self._consumer_memory.close() + else: + logger.error("Timed out waiting for Mooncake EC writers; keeping the consumer receive pool registered") + self._producer_memory.close() + self._transfer.close() diff --git a/aphrodite/distributed/ec_transfer/ec_connector/mooncake_ec_connector.py b/aphrodite/distributed/ec_transfer/ec_connector/mooncake_ec_connector.py new file mode 100644 index 0000000000..c995c3e21c --- /dev/null +++ b/aphrodite/distributed/ec_transfer/ec_connector/mooncake_ec_connector.py @@ -0,0 +1,120 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +"""Encoder-cache connector backed by Mooncake TransferEngine.""" + +from __future__ import annotations + +from typing import TYPE_CHECKING, Any + +from aphrodite.distributed.ec_transfer.ec_connector.base import ( + ECConnectorBase, + ECConnectorMetadata, + ECConnectorRole, + ECConnectorWorkerMetadata, +) +from aphrodite.distributed.ec_transfer.ec_connector.mooncake.metadata import ( + ECMooncakeConnectorMetadata, +) +from aphrodite.distributed.ec_transfer.ec_connector.mooncake.scheduler import ( + ECMooncakeScheduler, +) +from aphrodite.distributed.ec_transfer.ec_connector.mooncake.worker import ECMooncakeWorker + +if TYPE_CHECKING: + import torch + + from aphrodite.config import AphroditeConfig + from aphrodite.v1.core.sched.output import SchedulerOutput + from aphrodite.v1.outputs import ECConnectorOutput + from aphrodite.v1.request import Request + + +class ECMooncakeConnector(ECConnectorBase): + """Preserve the public API while delegating to one process-role component.""" + + def __init__(self, aphrodite_config: AphroditeConfig, role: ECConnectorRole): + super().__init__(aphrodite_config=aphrodite_config, role=role) + + self._scheduler: ECMooncakeScheduler | None = None + self._worker: ECMooncakeWorker | None = None + self._closed = False + + if role == ECConnectorRole.SCHEDULER: + self._scheduler = ECMooncakeScheduler(aphrodite_config) + elif role == ECConnectorRole.WORKER: + self._worker = ECMooncakeWorker(aphrodite_config) + else: + raise ValueError(f"Unknown EC connector role: {role}") + + def start_worker_services(self) -> None: + assert self._worker is not None + self._worker.start_services() + + def start_save_caches(self, **kwargs: Any) -> None: + assert self._worker is not None + metadata = self._get_connector_metadata() + assert isinstance(metadata, ECMooncakeConnectorMetadata) + self._worker.start_save_caches(metadata, **kwargs) + + def start_load_caches(self, encoder_cache: dict[str, torch.Tensor], **kwargs: Any) -> None: + assert self._worker is not None + metadata = self._get_connector_metadata() + assert isinstance(metadata, ECMooncakeConnectorMetadata) + self._worker.start_load_caches(metadata, encoder_cache, **kwargs) + + def save_caches(self, encoder_cache: dict[str, torch.Tensor], mm_hash: str, **kwargs: Any) -> None: + assert self._worker is not None + self._worker.save_caches(encoder_cache, mm_hash, **kwargs) + + def get_finished(self, finished_req_ids: set[str]) -> tuple[set[str] | None, set[str] | None]: + assert self._worker is not None + return self._worker.get_finished(finished_req_ids) + + def build_connector_worker_meta(self) -> ECConnectorWorkerMetadata | None: + assert self._worker is not None + return self._worker.build_connector_worker_meta() + + def take_unavailable_requests(self) -> set[str]: + assert self._scheduler is not None + return self._scheduler.take_unavailable_requests() + + def has_cache_item(self, identifier: str) -> bool: + assert self._scheduler is not None + return self._scheduler.has_cache_item(identifier) + + def ensure_cache_available(self, request: Request, num_computed_tokens: int) -> bool: + assert self._scheduler is not None + return self._scheduler.ensure_cache_available(request, num_computed_tokens) + + def update_state_after_alloc(self, request: Request, index: int) -> None: + assert self._scheduler is not None + self._scheduler.update_state_after_alloc(request, index) + + def update_state_after_free(self, request: Request, index: int) -> None: + assert self._scheduler is not None + self._scheduler.update_state_after_free(request, index) + + def build_connector_meta(self, scheduler_output: SchedulerOutput) -> ECConnectorMetadata: + assert self._scheduler is not None + return self._scheduler.build_connector_meta(scheduler_output) + + def update_connector_output(self, connector_output: ECConnectorOutput) -> None: + assert self._scheduler is not None + self._scheduler.update_connector_output(connector_output) + + def has_pending_push_work(self) -> bool: + assert self._scheduler is not None + return self._scheduler.has_pending_push_work() + + def request_finished(self, request: Request) -> tuple[bool, dict[str, Any] | None]: + assert self._scheduler is not None + return self._scheduler.request_finished(request) + + def shutdown(self) -> None: + if self._closed: + return + self._closed = True + if self._scheduler is not None: + self._scheduler.close() + if self._worker is not None: + self._worker.close() diff --git a/aphrodite/v1/core/sched/scheduler.py b/aphrodite/v1/core/sched/scheduler.py index 9136c9cbcc..effe4d7cd3 100644 --- a/aphrodite/v1/core/sched/scheduler.py +++ b/aphrodite/v1/core/sched/scheduler.py @@ -560,6 +560,17 @@ def schedule(self, throttle_prefills: bool = False) -> SchedulerOutput: req_index += 1 continue + if ( + self.ec_connector is not None + and request.mm_features + and not self.ec_connector.ensure_cache_available( + request, + request.num_computed_tokens - request.num_output_placeholders, + ) + ): + req_index += 1 + continue + num_new_tokens = ( request.num_tokens_with_spec + request.num_output_placeholders - request.num_computed_tokens ) @@ -854,11 +865,7 @@ def schedule(self, throttle_prefills: bool = False) -> SchedulerOutput: assert num_computed_tokens <= request.num_tokens # Skip request with pending mm encoding prefetches - if ( - self.ec_connector is not None - and request.mm_features - and not self.ec_connector.ensure_cache_available(request, num_computed_tokens) - ): + if self._ec_transfer_pending(request, num_computed_tokens): request_queue.pop_request() step_skipped_waiting.prepend_request(request) continue @@ -873,11 +880,18 @@ def schedule(self, throttle_prefills: bool = False) -> SchedulerOutput: ) else: # KVTransfer: WAITING reqs have num_computed_tokens > 0 - # after async KV recvs are completed. + # after async KV recvs are completed. A streaming-input + # session resumes here too, carrying whatever media its + # latest chunk added, so this branch needs the same gate. new_computed_blocks = self.kv_cache_manager.empty_kv_cache_blocks num_new_local_computed_tokens = 0 num_computed_tokens = request.num_computed_tokens + if self._ec_transfer_pending(request, num_computed_tokens): + request_queue.pop_request() + step_skipped_waiting.prepend_request(request) + continue + encoder_inputs_to_schedule = None external_load_encoder_input = [] new_encoder_compute_budget = encoder_compute_budget @@ -1938,6 +1952,10 @@ def update_from_output( self.grammar_compile_error_reqs.clear() if failed_kv_load_req_ids and not self.recompute_kv_load_failures: error_req_ids.update(failed_kv_load_req_ids) + if self.ec_connector is not None: + # An encoder input the connector can no longer obtain. Failing is + # retryable: re-issuing the request re-runs the encode. + error_req_ids.update(self.ec_connector.take_unavailable_requests()) if error_req_ids: error_reqs = self.finish_requests(error_req_ids, RequestStatus.FINISHED_ERROR) @@ -2017,6 +2035,14 @@ def update_from_output( return engine_core_outputs + def _ec_transfer_pending(self, request: Request, num_computed_tokens: int) -> bool: + """Whether an encoder input this request needs is still in transit.""" + return ( + self.ec_connector is not None + and bool(request.mm_features) + and not self.ec_connector.ensure_cache_available(request, num_computed_tokens) + ) + @staticmethod def _is_blocked_waiting_status(status: RequestStatus) -> bool: return status in ( @@ -2101,14 +2127,19 @@ def _free_encoder_inputs(self, request: Request) -> None: # With Whisper, as soon as we've generated a single token, # we know we're done with the encoder input. Cross Attention # KVs have been calculated and cached already. - self.encoder_cache_manager.free_encoder_input(request, input_id) + self._free_encoder_input(request, input_id) elif ( start_pos + num_tokens + spec_lookahead <= request.num_computed_tokens - request.num_output_placeholders ): # Processed, stored in the decoder KV cache, and far enough past # the placeholder range (plus the drafter's look-ahead) that no # rejection or drafter gather can reference it. - self.encoder_cache_manager.free_encoder_input(request, input_id) + self._free_encoder_input(request, input_id) + + def _free_encoder_input(self, request: Request, input_id: int) -> None: + self.encoder_cache_manager.free_encoder_input(request, input_id) + if self.ec_connector is not None: + self.ec_connector.update_state_after_free(request, input_id) def update_draft_token_ids(self, draft_token_ids: DraftTokenIds) -> None: for req_id, spec_token_ids in zip( diff --git a/aphrodite/v1/worker/ec_connector_model_runner_mixin.py b/aphrodite/v1/worker/ec_connector_model_runner_mixin.py index 87015c3a49..966d4d6b44 100644 --- a/aphrodite/v1/worker/ec_connector_model_runner_mixin.py +++ b/aphrodite/v1/worker/ec_connector_model_runner_mixin.py @@ -62,6 +62,12 @@ def _get_ec_connector_output( assert scheduler_output.ec_connector_metadata is not None ec_connector.bind_connector_metadata(scheduler_output.ec_connector_metadata) + if ec_connector.is_producer: + # Pass the cache explicitly: a producer needs it to publish an item + # that is already cached from an earlier step, which this step will + # not re-encode -- `save_caches` never fires for those. + ec_connector.start_save_caches(encoder_cache=encoder_cache, **kwargs) + # Load caches for consumer or both roles if ec_connector.is_consumer: ec_connector.start_load_caches(encoder_cache, **kwargs) @@ -73,4 +79,6 @@ def _get_ec_connector_output( scheduler_output.finished_req_ids ) + output.ec_connector_worker_meta = ec_connector.build_connector_worker_meta() + ec_connector.clear_connector_metadata() diff --git a/aphrodite/v1/worker/gpu/ec_connector.py b/aphrodite/v1/worker/gpu/ec_connector.py index ce54484539..52b1bd5aff 100644 --- a/aphrodite/v1/worker/gpu/ec_connector.py +++ b/aphrodite/v1/worker/gpu/ec_connector.py @@ -58,6 +58,9 @@ def maybe_get_output(self, scheduler_output: "SchedulerOutput") -> Generator[ECC assert scheduler_output.ec_connector_metadata is not None ec_connector.bind_connector_metadata(scheduler_output.ec_connector_metadata) + if ec_connector.is_producer: + ec_connector.start_save_caches(encoder_cache=self.encoder_cache) + if ec_connector.is_consumer: ec_connector.start_load_caches(self.encoder_cache) diff --git a/aphrodite/v1/worker/gpu_worker.py b/aphrodite/v1/worker/gpu_worker.py index 98e49a3b6f..c740135760 100644 --- a/aphrodite/v1/worker/gpu_worker.py +++ b/aphrodite/v1/worker/gpu_worker.py @@ -28,6 +28,8 @@ from aphrodite.distributed.ec_transfer import ( ensure_ec_transfer_initialized, ensure_ec_transfer_shutdown, + get_ec_transfer, + has_ec_transfer, ) from aphrodite.distributed.eplb.eplb_utils import override_envs_for_eplb from aphrodite.distributed.kv_transfer import ( @@ -473,6 +475,9 @@ def load_model(self, *, load_dummy_weights: bool = False) -> None: ): self.model_runner.load_model(load_dummy_weights=load_dummy_weights) + if has_ec_transfer(): + get_ec_transfer().start_worker_services() + if self.aphrodite_config.weight_transfer_config is not None: self.weight_transfer_engine = WeightTransferEngineFactory.create_engine( self.aphrodite_config.weight_transfer_config, diff --git a/examples/disaggregated/disaggregated_encoder/disagg_epd_proxy.py b/examples/disaggregated/disaggregated_encoder/disagg_epd_proxy.py new file mode 100644 index 0000000000..57e7159e3c --- /dev/null +++ b/examples/disaggregated/disaggregated_encoder/disagg_epd_proxy.py @@ -0,0 +1,990 @@ +#!/usr/bin/env python3 +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +""" +disagg_encoder_proxy.py + +Proxy that routes OpenAI-compatible “/v1/chat/completions” requests to two +clusters: + • encode (multimodal feature extraction) + • decode (language-model inference) + +For MM input we: + 1. Extract *every* image/audio/video item. + 2. Fire N concurrent requests to the encoder cluster + (one request per item, with **all text removed**). + 3. Wait for all of them to succeed. + 4. Forward the *original* request to a decode server. +""" + +from __future__ import annotations + +import argparse +import asyncio +import hashlib +import io +import itertools +import json +import logging +import os +import random +import time +import uuid +from collections.abc import AsyncIterator +from typing import Any + +import aiohttp +import pybase64 as base64 +import uvicorn +from fastapi import FastAPI, HTTPException, Request +from fastapi.responses import JSONResponse, StreamingResponse + +############################################################################### +# FastAPI app & global state +############################################################################### + +logging.basicConfig(level=logging.INFO, format="%(asctime)s %(levelname)s: %(message)s") +logger = logging.getLogger("proxy") + +app = FastAPI() +encode_session: aiohttp.ClientSession | None = None +prefill_session: aiohttp.ClientSession | None = None +decode_session: aiohttp.ClientSession | None = None + +# Cursor for round-robin encoder assignment, shared across requests so the +# fan-out doesn't restart from e_urls[0] every time. +encoder_rr_idx = 0 +encoder_rr_lock = asyncio.Lock() + +############################################################################### +# Utils +############################################################################### + + +MM_TYPES = {"image_url", "audio_url", "input_audio", "video_url"} + +# The embeds content type each MM item is rewritten to once the encoder has +# published its embedding out of band. +EMBEDS_TYPES = { + "image_url": "image_embeds", + "audio_url": "audio_embeds", + "input_audio": "audio_embeds", + "video_url": "video_embeds", +} + + +def encoder_rr_assignment(e_urls: list[str], start: int, count: int) -> tuple[list[str], int]: + """Assign `count` items to encoder URLs starting from cursor `start`. + + Returns the per-item URL list and the cursor value the next call should + start from, so the assignment is contiguous across calls instead of + restarting at e_urls[0] every time. + """ + urls = [e_urls[(start + i) % len(e_urls)] for i in range(count)] + next_start = (start + count) % len(e_urls) + return urls, next_start + + +def validate_ec_consumer_routing(prefill_urls: list[str], consumer_addrs: list[str]) -> None: + """Reject the topology whose EC destination cannot be routed safely.""" + if prefill_urls and consumer_addrs: + raise ValueError( + "Mooncake EC consumer routing supports E+PD only; disable independent " + "prefill or omit --ec-consumer-zmq-addrs." + ) + + +# Diagnostic switch: forward the original request to the decoder so the +# only difference from the rewrite path is the rewrite itself. +NO_REWRITE = False +# Decode-side retries for a retryable internal error (`finish_reason="error"`, +# e.g. an encoder embedding the connector could not deliver). Re-issuing runs +# the encode again, which produces a fresh transfer. +DECODE_RETRIES = 1 + + +# Grid metadata reported by the encoder instance, keyed by item index. +# Empty when the encoder did not report any (then nothing is rewritten). +def content_uuid(item: dict) -> str: + """Cache key for a multimodal item, derived from its content. + + Must be content-derived, not request-derived: the EC cache is keyed by this + value, so a per-request key (a request id, say) would make every request a + miss and throw away cross-request reuse of already-encoded media -- while + the unmodified path, which hashes the content, would keep it. That asymmetry + silently biases any comparison between the two. + """ + url = (item.get("image_url") or item.get("audio_url") or item.get("video_url") or {}).get("url") or "" + payload = url or json.dumps(item, sort_keys=True) + return hashlib.sha256(payload.encode()).hexdigest() + + +def _b64_tensor(values: list) -> str: + import torch + + buf = io.BytesIO() + flat = [v for item in values for v in (item if isinstance(item, list) else [item])] + # Floats stay float64 so timestamp strings format exactly as the + # encoder computed them. + dtype = torch.float64 if any(isinstance(v, float) for v in flat) else None + # Downstream stacks per item, so hand over a flat vector. + torch.save(torch.tensor(flat, dtype=dtype), buf) + return base64.b64encode(buf.getvalue()).decode() + + +def rewrite_for_decode(req_data: dict, item_meta: dict[int, dict]) -> dict: + """Replace each media item with a metadata-only reference for the decoder. + + The decoder does not need the pixels: the encoder instance already produced + the embedding and published it through the EC connector under the same uuid. + Sending only the grid lets the decoder size the placeholder range without + re-running the media transform. + + `item_meta` holds what the encoder reported for each item (its cache key and + the grid its processor actually produced), so the grid is never re-derived + here -- a second derivation could disagree with the encoder's. + """ + rewritten = 0 + transfer_items = [] + idx = 0 + new_messages = [] + for msg in req_data.get("messages", []): + content = msg.get("content") + if not isinstance(content, list): + new_messages.append(msg) + continue + new_content = [] + for item in content: + if item.get("type") not in MM_TYPES: + new_content.append(item) + continue + meta = dict(item_meta.get(idx) or {}) + idx += 1 + item_uuid = meta.pop("mm_hash", None) + ec_mm_hash = meta.pop("ec_mm_hash", None) or item_uuid + transfer_id = meta.pop("transfer_id", None) + # Whatever keys the encoder reported are the metadata its model + # declared as needed to size the placeholder range; the proxy does + # not need to know their names. + metadata = {k: _b64_tensor(v) for k, v in meta.items()} + if not metadata or not item_uuid: + # Nothing to size the placeholder range with. A processor cache + # hit is not a cause on its own: with the default `lru` type the + # engine restores the item before the scheduler reports it. It + # goes missing when the encode request failed, or under + # `--mm-processor-cache-type shm`, where a hit replaces the item + # with its shared-memory address and only the worker restores + # it. Send the media so the decoder can derive the grid itself. + new_content.append(item) + continue + embeds_type = EMBEDS_TYPES[item["type"]] + new_content.append({"type": embeds_type, embeds_type: metadata, "uuid": item_uuid}) + if transfer_id is not None: + transfer_items.append({"mm_hash": ec_mm_hash, "transfer_id": transfer_id}) + rewritten += 1 + new_messages.append({**msg, "content": new_content}) + + if not rewritten: + return req_data + logger.info("Rewrote %d media item(s) as metadata references", rewritten) + rewritten_request = {**req_data, "messages": new_messages} + if transfer_items: + ec_transfer_params = dict(req_data.get("ec_transfer_params") or {}) + ec_transfer_params["ec_items"] = transfer_items + rewritten_request["ec_transfer_params"] = ec_transfer_params + return rewritten_request + + +def extract_mm_items(request_data: dict) -> list[dict]: + """ + Return *all* image/audio/video items that appear anywhere in `messages`. + + Each returned dict looks like: + { "type": "image_url", "image_url": {...} } + """ + items: list[dict] = [] + for msg in request_data.get("messages", []): + content = msg.get("content") + if not isinstance(content, list): + continue + + for item in content: + if item.get("type") in MM_TYPES: + items.append(item) + return items + + +async def fanout_encoder_primer( + orig_request: dict, + e_urls: list[str], + req_id: str, + consumer_zmq: str | None = None, +) -> tuple[dict[int, dict], dict[str, Any]]: + """ + 1. Build one request *per MM item* with all text removed. + 2. Send them concurrently to the encode cluster. + 3. Raise if any of them fails. + + Returns, per item index, the metadata the encoder reported in + `ec_transfer_params`: its EC cache key and the grid its processor produced. + The proxy still supplies the uuid so both sides key the cache the same way; + the grid can only come from the encoder, which is the side that computed it. + + Also returns the connector handles to put on the decode body, as a fresh + mapping. `orig_request` is left untouched so a retry re-encodes from the + original request instead of carrying the previous attempt's handles. + """ + logger.info("[%s] Processing multimodal items...", req_id) + + mm_items = extract_mm_items(orig_request) + if not mm_items: + logger.info("[%s] No multimodal items, skipping encoder", req_id) + return {}, {} # nothing to do + + logger.info("[%s] got %d multimodal items...", req_id, len(mm_items)) + + tasks = [] + item_uuids: dict[int, str] = {} + item_transfer_ids: dict[int, str] = {} + item_meta: dict[int, dict] = {} + ec_params: dict[str, Any] = {} + + # Round-robin over encode servers to distribute load a bit. The cursor + # persists across requests so fan-out doesn't restart at e_urls[0] every + # time (which would hot-spot the first encoder for single-item requests). + global encoder_rr_idx + async with encoder_rr_lock: + url_cycle, encoder_rr_idx = encoder_rr_assignment(e_urls, encoder_rr_idx, len(mm_items)) + + for idx, (item, target_url) in enumerate(zip(mm_items, url_cycle)): + # Derive a *child* request id: :: + child_req_id = f"{req_id}:{idx}:{uuid.uuid4().hex[:6]}" + headers = {"x-request-id": child_req_id} + + # With --no-rewrite the decoder still receives the raw image and derives + # the cache key by hashing it, so the encoder must do the same -- passing + # a uuid here would make the two disagree and silently defeat the EC + # transfer, leaving the decoder to encode the image itself. + item_uuid = None if NO_REWRITE else content_uuid(item) + if item_uuid is not None: + item_uuids[idx] = item_uuid + transfer_id = uuid.uuid4().hex + item_transfer_ids[idx] = transfer_id + + encoder_req = { + # You *may* need to keep additional fields + "model": orig_request.get("model"), + "messages": [ + { + "role": "user", + "content": [item if item_uuid is None else {**item, "uuid": item_uuid}], + }, + ], + # No max_tokens cap: the encoder instance never samples, it finishes + # once the prompt is encoded and its embeddings are published. + "stream": False, + } + if consumer_zmq is not None: + # No mm_hash here on purpose. The encoder's own + # `mm_features[i].identifier` is derived from the uuid *and* the + # engine's media_io_kwargs / mm_processor_kwargs, so this proxy + # cannot know it before the encoder runs. Sending the bare uuid + # would fail the connector's hash match and make the producer + # invent its own transfer id, which the consumer could then never + # cancel. Omitting it lets the connector match by position, which + # is exact: one encoder request carries exactly one item. + encoder_req["ec_transfer_params"] = { + "consumer_zmq": consumer_zmq, + "ec_items": [{"transfer_id": transfer_id}], + } + tasks.append( + encode_session.post( + f"{target_url}/v1/chat/completions", + json=encoder_req, + headers=headers, + ) + ) + + results = await asyncio.gather(*tasks, return_exceptions=True) + + # Fail fast if any sub-request failed + for idx, r in enumerate(results): + if isinstance(r, Exception): + logger.error( + "[%s] Encoder request #%d raised exception: %s", + req_id, + idx, + r, + exc_info=r, + ) + raise HTTPException(status_code=502, detail=f"Encoder request failed: {str(r)}") + if r.status != 200: + try: + detail = await r.text() + except Exception: + detail = "" + logger.error( + "[%s] Encoder request #%d returned status %s: %s", + req_id, + idx, + r.status, + detail, + ) + raise HTTPException( + status_code=r.status, + detail=f"Encoder request failed: {detail}", + ) + + # The encoder reports each mm_hash's metadata (e.g. the grid) here, + # keyed by the same uuid this proxy assigned above. + try: + params = (await r.json()).get("ec_transfer_params") or {} + except Exception: + logger.warning("[%s] Could not read encoder metadata #%d", req_id, idx) + params = {} + if params: + # One encoder request carries exactly one item, so there is a + # single reported entry. Do not key it by this proxy's uuid: when + # media_io_kwargs or mm_processor_kwargs are set the engine + # re-hashes the uuid together with them, so the encoder's own + # `mm_features[i].identifier` is a derived value this proxy cannot + # predict. Fall back to the sole entry, and carry the key the + # encoder actually used through as `ec_mm_hash`. + ec_mm_hash = item_uuids.get(idx) + reported = params.get(ec_mm_hash) + if reported is None and len(params) == 1: + ((ec_mm_hash, reported),) = params.items() + if reported: + metadata = reported.get("metadata") or {} + if metadata: + item_meta[idx] = { + **metadata, + "mm_hash": item_uuids.get(idx, ec_mm_hash), + "ec_mm_hash": ec_mm_hash, + } + if idx in item_transfer_ids: + item_meta[idx]["transfer_id"] = item_transfer_ids[idx] + # Whatever the encoder reported alongside `metadata` is the + # connector's own handle on the published embedding (for NIXL, + # peer_host/peer_port/size_bytes). The decoder's connector + # looks it up by mm_hash on the request, so carry it through. + ec_params[item_uuids.get(idx, ec_mm_hash)] = reported + if NO_REWRITE and consumer_zmq is not None: + ec_params.setdefault("ec_items", []).append( + {"mm_hash": ec_mm_hash, "transfer_id": item_transfer_ids[idx]} + ) + + logger.info("[%s] All %d encoder requests completed successfully", req_id, len(mm_items)) + return item_meta, ec_params + + +async def maybe_prefill( + req_data: dict, + p_url: str, + req_id: str, +) -> dict: + """ + - Do prefill-only task if p_url exist; + - Return a new body carrying kv transfer params (for nixl connector) + - Else, skip and return the original request data for decode + + `req_data` is never mutated: a decode retry re-enters this function with the + same body, and one attempt's `remote_block_ids` must not reach the next. + """ + if p_url: + logger.info("[%s] Processing through prefill: %s", req_id, p_url) + + prefill_response = await process_prefill_stage(req_data, p_url, req_id) + # for nixl connector to facilitate kv transfer... + prefill_response_json = await prefill_response.json() + kv_transfer_params = prefill_response_json.get("kv_transfer_params", {}) + if kv_transfer_params: + return {**req_data, "kv_transfer_params": kv_transfer_params} + + return req_data + + +async def process_prefill_stage( + req_data: dict, + p_url: str, + req_id: str, +) -> dict: + """Process request through Prefill stage and return kv_transfer_params""" + logger.info("[%s] Sending prefill request to: %s", req_id, p_url) + + prefill_request = req_data.copy() + prefill_request["kv_transfer_params"] = { + "do_remote_decode": True, + "do_remote_prefill": False, + "remote_engine_id": None, + "remote_block_ids": None, + "remote_host": None, + "remote_port": None, + } + prefill_request["stream"] = False + prefill_request["max_tokens"] = 1 + if "max_completion_tokens" in prefill_request: + prefill_request["max_completion_tokens"] = 1 + if "stream_options" in prefill_request: + del prefill_request["stream_options"] + + headers = {"x-request-id": req_id} + try: + prefill_response = await prefill_session.post( + f"{p_url}/v1/chat/completions", json=prefill_request, headers=headers + ) + prefill_response.raise_for_status() + + if prefill_response.status != 200: + error_text = await prefill_response.text() + logger.error( + "[%s] Prefill request failed with status %d: %s", + req_id, + prefill_response.status, + error_text, + ) + raise HTTPException( + status_code=prefill_response.status, + detail={"error": "Prefill request failed", "message": error_text}, + ) + logger.info("[%s] Prefill request completed successfully", req_id) + + return prefill_response + + except Exception as e: + logger.error("Prefill processing failed: %s", str(e)) + raise HTTPException( + status_code=500, + detail={"error": "Prefill processing error", "message": str(e)}, + ) from e + + +############################################################################### +# Middleware for request/response logging +############################################################################### + + +async def log_requests(request: Request, call_next): + """Middleware to log all incoming requests and responses""" + req_id = request.headers.get("x-request-id", str(uuid.uuid4())) + + # Log incoming request + logger.info( + ">>> [%s] %s %s from %s", + req_id, + request.method, + request.url.path, + request.client.host if request.client else "unknown", + ) + + try: + # Process request + response = await call_next(request) + + # Log response + logger.info( + "<<< [%s] %s %s completed with status %d", + req_id, + request.method, + request.url.path, + response.status_code, + ) + + return response + except Exception as e: + # Log errors + logger.exception( + "!!! [%s] %s %s failed with error: %s", + req_id, + request.method, + request.url.path, + str(e), + ) + raise + + +############################################################################### +# FastAPI lifecycle +############################################################################### + + +@app.on_event("startup") +async def on_startup() -> None: + global encode_session, prefill_session, decode_session + timeout = aiohttp.ClientTimeout(total=100_000) + # Aphrodite closes an idle keep-alive connection after + # APHRODITE_HTTP_TIMEOUT_KEEP_ALIVE seconds (5 by default), while aiohttp keeps + # pooling it for 15. Reusing one it has already closed fails the request + # with ServerDisconnectedError, and the server logs nothing at all: it + # closed the socket before the request arrived. Retire ours first. + server_keep_alive = float(os.getenv("APHRODITE_HTTP_TIMEOUT_KEEP_ALIVE", "5")) + connector = aiohttp.TCPConnector( + limit=0, + **({"keepalive_timeout": server_keep_alive / 2} if server_keep_alive > 0 else {"force_close": True}), + ) + encode_session = aiohttp.ClientSession(timeout=timeout, connector=connector) + if app.state.p_urls: + # only setup if prefill instance(s) exist + prefill_session = aiohttp.ClientSession(timeout=timeout, connector=connector) + decode_session = aiohttp.ClientSession(timeout=timeout, connector=connector) + + +@app.on_event("shutdown") +async def on_shutdown() -> None: + global encode_session, prefill_session, decode_session + if encode_session: + await encode_session.close() + if prefill_session: + await prefill_session.close() + if decode_session: + await decode_session.close() + + +############################################################################### +# Core forwarding +############################################################################### + + +async def prepare_for_decode( + req_data: dict, + req_id: str, + e_urls: list[str], + p_url: str, + consumer_zmq: str | None, +) -> tuple[dict, float, float]: + """Encode, rewrite and prefill, returning the body to send to decode. + + `req_data` is left untouched so a retry starts from the original media + rather than from a body whose images are already metadata references. + """ + _t0 = time.perf_counter() + item_meta, ec_params = await fanout_encoder_primer(req_data, e_urls, req_id, consumer_zmq) + _t1 = time.perf_counter() + prepared = req_data if NO_REWRITE else rewrite_for_decode(req_data, item_meta) + if ec_params: + # A fresh body every time: `rewrite_for_decode` hands back `req_data` + # itself when it rewrote nothing, and this attempt's handles must not + # outlive it into a retry. + handles = dict(prepared.get("ec_transfer_params") or {}) + handles.update(ec_params) + prepared = {**prepared, "ec_transfer_params": handles} + _t2 = time.perf_counter() + prepared = await maybe_prefill(prepared, p_url, req_id) + return prepared, _t1 - _t0, _t2 - _t1 + + +async def forward_non_stream( + req_data: dict, + req_id: str, + e_urls: list[str], + p_url: str, + d_url: str, + consumer_zmq: str | None, + dp_rank: int | None = None, +) -> dict: + try: + for attempt in range(DECODE_RETRIES + 1): + _t0 = time.perf_counter() + prepared, encode_s, rewrite_s = await prepare_for_decode(req_data, req_id, e_urls, p_url, consumer_zmq) + _t2 = time.perf_counter() + + logger.info("[%s] Forwarding to decode: %s", req_id, d_url) + headers = {"x-request-id": req_id} + if dp_rank is not None: + headers["X-data-parallel-rank"] = str(dp_rank) + + async with decode_session.post(f"{d_url}/v1/chat/completions", json=prepared, headers=headers) as resp: + if resp.status >= 400: + detail = await resp.text() + # 500 is the decoder's retryable internal error, which + # includes an encoder embedding it could not obtain. Redoing + # the encode publishes the item again. + if resp.status == 500 and attempt < DECODE_RETRIES: + logger.warning( + "[%s] Decode returned 500, re-encoding and retrying (attempt %d/%d): %s", + req_id, + attempt + 1, + DECODE_RETRIES, + detail[:200], + ) + continue + logger.error( + "[%s] Decode request returned status %s: %s", + req_id, + resp.status, + detail, + ) + raise HTTPException(status_code=resp.status, detail=detail) + out = await resp.json() + _t3 = time.perf_counter() + logger.info( + "STAGE %s encode=%.1f rewrite=%.1f decode=%.1f total=%.1f attempt=%d", + "no-rewrite" if NO_REWRITE else "rewrite", + encode_s * 1e3, + rewrite_s * 1e3, + (_t3 - _t2) * 1e3, + (_t3 - _t0) * 1e3, + attempt, + ) + return out + raise HTTPException(status_code=500, detail="Decode failed after re-encoding") + + except HTTPException: + raise + except Exception as e: + logger.exception("[%s] Error in forward_non_stream: %s", req_id, str(e)) + raise HTTPException(status_code=500, detail=f"Proxy error: {str(e)}") from e + + +async def forward_stream( + req_data: dict, + req_id: str, + e_urls: list[str], + p_url: str, + d_url: str, + consumer_zmq: str | None, + dp_rank: int | None = None, +) -> AsyncIterator[str]: + try: + for attempt in range(DECODE_RETRIES + 1): + _t0 = time.perf_counter() + prepared, encode_s, rewrite_s = await prepare_for_decode(req_data, req_id, e_urls, p_url, consumer_zmq) + _t2 = time.perf_counter() + + logger.info("[%s] Starting streaming from decode: %s", req_id, d_url) + headers = {"x-request-id": req_id} + if dp_rank is not None: + headers["X-data-parallel-rank"] = str(dp_rank) + + _first = None + async with decode_session.post( + f"{d_url}/v1/chat/completions", + json=prepared, + headers=headers, + ) as resp: + # Retry only before the first chunk: once anything reached the + # client the response cannot be replaced. + if resp.status == 500 and attempt < DECODE_RETRIES: + detail = await resp.text() + logger.warning( + "[%s] Decode returned 500 before streaming, re-encoding and retrying (attempt %d/%d): %s", + req_id, + attempt + 1, + DECODE_RETRIES, + detail[:200], + ) + continue + resp.raise_for_status() + async for chunk in resp.content.iter_chunked(1024): + if chunk: + if _first is None: + _first = time.perf_counter() + yield chunk.decode("utf-8", errors="ignore") + _t3 = time.perf_counter() + + logger.info( + "STAGE %s encode=%.1f rewrite=%.2f decode_ttfb=%.1f decode_total=%.1f attempt=%d", + "no-rewrite" if NO_REWRITE else "rewrite", + encode_s * 1e3, + rewrite_s * 1e3, + ((_first or _t3) - _t2) * 1e3, + (_t3 - _t2) * 1e3, + attempt, + ) + logger.info("[%s] Streaming completed", req_id) + return + + except HTTPException: + logger.exception("[%s] HTTPException in forward_stream", req_id) + raise + except Exception as e: + logger.exception("[%s] Error in forward_stream: %s", req_id, str(e)) + raise HTTPException(status_code=500, detail=f"Proxy streaming error: {str(e)}") from e + + +############################################################################### +# Public routes +############################################################################### + + +@app.post("/v1/chat/completions") +async def chat_completions(request: Request): + try: + req_data = await request.json() + req_id = request.headers.get("x-request-id", str(uuid.uuid4())) + + e_urls = app.state.e_urls # we want the full list for fan-out + p_url = random.choice(app.state.p_urls) if app.state.p_urls else None + decode_index = random.randrange(len(app.state.d_urls)) + d_url = app.state.d_urls[decode_index] + dp_size = app.state.ec_consumer_dp_size + # Round-robin the replica, then name it to both halves: the decoder + # honours the rank header instead of its own balancer, and the encoder + # pushes to that replica's control channel. Choosing once here means a + # decode retry re-encodes to the same replica. + dp_rank = next(app.state.replica_counter) % dp_size if dp_size > 1 else None + ec_index = decode_index * dp_size + (dp_rank or 0) + consumer_zmq = app.state.d_ec_urls[ec_index] if app.state.d_ec_urls else None + + is_streaming = req_data.get("stream", False) + + if is_streaming: + return StreamingResponse( + forward_stream(req_data, req_id, e_urls, p_url, d_url, consumer_zmq, dp_rank), + media_type="text/event-stream", + ) + result = await forward_non_stream(req_data, req_id, e_urls, p_url, d_url, consumer_zmq, dp_rank) + return JSONResponse(content=result) + + except HTTPException: + raise + except Exception as e: + logger.exception("Error in chat_completions endpoint: %s", str(e)) + raise HTTPException(status_code=500, detail=f"Request processing error: {str(e)}") from e + + +@app.get("/v1/models") +async def list_models(): + async with decode_session.get(f"{app.state.d_urls[0]}/v1/models") as resp: + resp.raise_for_status() + return await resp.json() + + +@app.get("/health") +async def health_check(): + async def healthy(urls): + if not urls: + return "empty" + for u in urls: + try: + async with encode_session.get(f"{u}/health") as resp: + resp.raise_for_status() + except Exception: + return "unhealthy" + return "healthy" + + e_status, p_status, d_status = await asyncio.gather( + healthy(app.state.e_urls), healthy(app.state.p_urls), healthy(app.state.d_urls) + ) + + overall_healthy = all(status != "unhealthy" for status in (e_status, p_status, d_status)) + + status_code = 200 if overall_healthy else 503 + + return JSONResponse( + { + "proxy": "healthy", + "encode_cluster": e_status, + "prefill_cluster": p_status, + "decode_cluster": d_status, + }, + status_code=status_code, + ) + + +############################################################################### +# Simple profiler fan-out (unchanged except for sessions) +############################################################################### + + +async def _post_if_available( + session: aiohttp.ClientSession, + url: str, + payload: dict, + headers: dict, +) -> dict | None: + """ + POST `payload` to `url`. + + Returns + ------- + • The decoded JSON body on success (2xx) + • None if the endpoint does not exist (404) + • Raises for anything else. + """ + try: + resp = await session.post(url, json=payload, headers=headers) + if resp.status == 404: # profiling disabled on that server + logger.warning("Profiling endpoint missing on %s", url) + return None + resp.raise_for_status() + return await resp.json(content_type=None) + except aiohttp.ClientResponseError as exc: + # Pass 404 through the branch above, re-raise everything else + if exc.status == 404: + logger.warning("Profiling endpoint missing on %s", url) + return None + raise + except Exception: + # Network errors etc.: propagate + raise + + +async def _profile_cmd(cmd: str, payload: dict, e_url: str, p_url: str, d_url: str): + """ + Fire & forget to both clusters, tolerate 404. + """ + headers = {"Authorization": f"Bearer {os.getenv('OPENAI_API_KEY', '')}"} + + encode_task = _post_if_available(encode_session, f"{e_url}/{cmd}_profile", payload, headers) + prefill_task = ( + _post_if_available(prefill_session, f"{p_url}/{cmd}_profile", payload, headers) + if p_url is not None + else asyncio.sleep(0) + ) + decode_task = _post_if_available(decode_session, f"{d_url}/{cmd}_profile", payload, headers) + + encode_res, prefill_res, decode_res = await asyncio.gather(encode_task, prefill_task, decode_task) + + # If *all* clusters said “I don’t have that route”, surface an error + if encode_res is prefill_res is decode_res is None: + raise HTTPException( + status_code=503, + detail="Profiling endpoints are disabled on all clusters", + ) + + return { + "encode": encode_res, # may be None + "prefill": prefill_res, # may be None + "decode": decode_res, # may be None + } + + +@app.post("/start_profile") +async def start_profile(request: Request): + body = await request.json() + # TODO: handle multi urls properly + e_url = random.choice(app.state.e_urls) + p_url = random.choice(app.state.p_urls) if app.state.p_urls else None + d_url = random.choice(app.state.d_urls) + return await _profile_cmd("start", body, e_url, p_url, d_url) + + +@app.post("/stop_profile") +async def stop_profile(request: Request): + body = await request.json() + # TODO: handle multi urls properly + e_url = random.choice(app.state.e_urls) + p_url = random.choice(app.state.p_urls) if app.state.p_urls else None + d_url = random.choice(app.state.d_urls) + return await _profile_cmd("stop", body, e_url, p_url, d_url) + + +if __name__ == "__main__": + parser = argparse.ArgumentParser() + parser.add_argument("--host", default="0.0.0.0") + parser.add_argument("--port", type=int, default=8000) + parser.add_argument( + "--log-requests", + action="store_true", + help=( + "Log every request in and out, and raise the log level to DEBUG. " + "Off by default: the proxy is on the request path." + ), + ) + parser.add_argument( + "--no-rewrite", + action="store_true", + help="Forward images to the decoder unchanged (for stage-timing A/B).", + ) + parser.add_argument( + "--encode-servers-urls", + required=True, + help='Comma-separated encode URLs ("http://e1:8001,http://e2:8001")', + ) + parser.add_argument( + "--prefill-servers-urls", + required=True, + help=( + 'Comma-separated prefill URLs ("http://p1:8003,http://p2:8004") ' + 'to enable E->P->D, set "disable" or "none" to enable E->PD' + ), + ) + parser.add_argument( + "--decode-servers-urls", + required=True, + help='Comma-separated decode URLs ("http://d1:8005,http://d2:8006")', + ) + parser.add_argument( + "--decode-retries", + type=int, + default=1, + help=( + "Re-encode and re-send when decode returns 500, which is its " + "retryable internal error (an undeliverable encoder embedding " + "among them). 0 disables." + ), + ) + parser.add_argument( + "--ec-consumer-zmq-addrs", + default="", + help=( + "Comma-separated Mooncake EC consumer control addresses, aligned " + "with --decode-servers-urls. Required for Mooncake EC consumers and " + "supported only in E+PD mode. With --ec-consumer-dp-size > 1, list " + "each server's replicas consecutively: s0r0,s0r1,s1r0,s1r1." + ), + ) + parser.add_argument( + "--ec-consumer-dp-size", + type=int, + default=1, + help=( + "Data-parallel replicas per EC consumer. The proxy picks a replica " + "round-robin and names it to both halves of the request, because an " + "encoder push has to land where the request will run." + ), + ) + + args = parser.parse_args() + if args.log_requests: + logging.getLogger().setLevel(logging.DEBUG) + app.middleware("http")(log_requests) + NO_REWRITE = args.no_rewrite + DECODE_RETRIES = max(0, args.decode_retries) + app.state.e_urls = [u.strip() for u in args.encode_servers_urls.split(",") if u.strip()] + app.state.d_urls = [u.strip() for u in args.decode_servers_urls.split(",") if u.strip()] + app.state.d_ec_urls = [u.strip() for u in args.ec_consumer_zmq_addrs.split(",") if u.strip()] + if args.ec_consumer_dp_size < 1: + parser.error("--ec-consumer-dp-size must be at least 1") + app.state.ec_consumer_dp_size = args.ec_consumer_dp_size + app.state.replica_counter = itertools.count() + expected = len(app.state.d_urls) * args.ec_consumer_dp_size + if app.state.d_ec_urls and len(app.state.d_ec_urls) != expected: + parser.error( + "--ec-consumer-zmq-addrs must contain one address per consumer " + f"replica: expected {expected} " + f"({len(app.state.d_urls)} servers x {args.ec_consumer_dp_size} replicas), " + f"got {len(app.state.d_ec_urls)}" + ) + # handle prefill instances + if args.prefill_servers_urls.lower() in ("disable", "none", ""): + app.state.p_urls = [] + logger.info("Disaggregated prefill phase explicitly disabled by user. Running E + PD...") + else: + app.state.p_urls = [u.strip() for u in args.prefill_servers_urls.split(",") if u.strip()] + logger.info("Disaggregated prefill phase is enabled. Running E + P + D...") + try: + validate_ec_consumer_routing(app.state.p_urls, app.state.d_ec_urls) + except ValueError as exc: + parser.error(str(exc)) + + logger.info("Proxy listening on %s:%s", args.host, args.port) + logger.info("Encode servers: %s", app.state.e_urls) + logger.info("Prefill instances %s", app.state.p_urls) + logger.info("Decode servers: %s", app.state.d_urls) + if app.state.ec_consumer_dp_size > 1: + logger.info( + "EC consumer replicas per server: %d (control addresses: %s)", + app.state.ec_consumer_dp_size, + app.state.d_ec_urls, + ) + + uvicorn.run( + app, + host=args.host, + port=args.port, + log_level="info", + loop="uvloop", + access_log=True, + ) diff --git a/tests/v1/core/test_encoder_cache_manager.py b/tests/v1/core/test_encoder_cache_manager.py index 3fe88bfbb9..74544b7bc2 100644 --- a/tests/v1/core/test_encoder_cache_manager.py +++ b/tests/v1/core/test_encoder_cache_manager.py @@ -167,6 +167,38 @@ def test_get_freed_mm_hashes_clears_freed_list(): assert manager.get_freed_mm_hashes() == [] +def test_referencing_a_freeable_entry_protects_it_from_eviction(): + """A waiting request can hold an entry it is not scheduled for yet. + + Entries with no referent are evictable. A request that is deferred (e.g. + by an EC connector waiting on another item) still needs the items it + already has, and a connector that hands out one transfer per request + cannot always produce the same item a second time. + """ + manager = EncoderCacheManager(cache_size=10) + owner = MockRequest("owner", ["a"], [5]) + waiter = MockRequest("waiter", ["a"], [5]) + newcomer = MockRequest("newcomer", ["b"], [6]) + + manager.allocate(owner, 0) + manager.free_encoder_input(owner, 0) + assert "a" in manager.freeable + + # The waiting request references it; it is no longer reclaimable. + assert manager.check_and_update_cache(waiter, 0) + assert "a" not in manager.freeable + + # 'b' no longer fits, and 'a' must not be taken away to make room. + assert not manager.can_allocate(newcomer, 0, int(1e9), 0) + assert "a" in manager.cached + assert manager.get_freed_mm_hashes() == [] + + # Once the waiter is done with it, the entry is reclaimable again. + manager.free_encoder_input(waiter, 0) + assert manager.can_allocate(newcomer, 0, int(1e9), 0) + assert manager.get_freed_mm_hashes() == ["a"] + + def test_reallocated_hash_is_not_reported_as_freed(): manager = EncoderCacheManager(cache_size=8) req_a = MockRequest("reqA", ["a"], [4]) diff --git a/tests/v1/core/test_scheduler.py b/tests/v1/core/test_scheduler.py index 68d4dc787b..e3cd73afa8 100644 --- a/tests/v1/core/test_scheduler.py +++ b/tests/v1/core/test_scheduler.py @@ -5427,6 +5427,68 @@ def test_free_encoder_inputs_respects_unconfirmed_placeholders(): assert manager.get_cached_input_ids(request) == set() +def test_free_encoder_inputs_notifies_the_ec_connector(): + """The connector learns the item is consumed, not just that the request ended. + + Connectors hold per-item transfer state (remote buffers, reservations); + waiting for `request_finished` pins it for the whole generation. + """ + scheduler = create_scheduler(model="llava-hf/llava-1.5-7b-hf") + mm_start_pos, mm_length = 50, 100 + request = create_requests( + num_requests=1, + num_tokens=mm_start_pos + mm_length + 10, + mm_positions=[[PlaceholderRange(offset=mm_start_pos, length=mm_length)]], + )[0] + scheduler.encoder_cache_manager.allocate(request, 0) + scheduler.ec_connector = Mock() + + request.num_computed_tokens = mm_start_pos + mm_length - 1 + scheduler._free_encoder_inputs(request) + scheduler.ec_connector.update_state_after_free.assert_not_called() + + request.num_computed_tokens = mm_start_pos + mm_length + scheduler._free_encoder_inputs(request) + scheduler.ec_connector.update_state_after_free.assert_called_once_with(request, 0) + + +def test_unavailable_encoder_input_fails_the_request_as_retryable(): + """An encoder input the connector gave up on must end the request. + + `FinishReason.ERROR` is the retryable channel the KV connector already uses + for load failures, so the caller can re-issue; deferring instead parked the + request until the client timed out. + """ + scheduler = create_scheduler(model="llava-hf/llava-1.5-7b-hf") + request = create_requests( + num_requests=1, + num_tokens=160, + mm_positions=[[PlaceholderRange(offset=50, length=100)]], + )[0] + scheduler.add_request(request) + scheduler.ec_connector = Mock() + scheduler.ec_connector.take_unavailable_requests.return_value = {request.request_id} + # `request_finished` reports (delay_free, params) for the request teardown. + scheduler.ec_connector.request_finished.return_value = (False, None) + + scheduler_output = scheduler.schedule() + outputs = scheduler.update_from_output( + scheduler_output, + ModelRunnerOutput( + req_ids=[request.request_id], + req_id_to_index={request.request_id: 0}, + sampled_token_ids=[[]], + logprobs=None, + prompt_logprobs_dict={}, + pooler_output=[], + ), + ) + + assert request.status == RequestStatus.FINISHED_ERROR + engine_outputs = outputs[0].outputs + assert [o.finish_reason for o in engine_outputs] == [FinishReason.ERROR] + + def test_free_encoder_inputs_defers_for_eagle_lookahead(): """With EAGLE speculative decoding, the encoder input is retained one extra position so the drafter's +1 look-ahead mm-embedding gather (which reads one @@ -5622,7 +5684,8 @@ def test_ec_connector_ensure_cache_available_defers_request(use_kv_connector): # ensure_cache_available must have been called with (request, num_computed_tokens=0) # for a brand-new request that has no cached tokens yet. - scheduler.ec_connector.ensure_cache_available.assert_called_once_with(request_deferred, 0) + ensure_call = scheduler.ec_connector.ensure_cache_available.call_args + assert ensure_call.args == (request_deferred, 0) # Deferred request must NOT be scheduled assert request_deferred.request_id not in output.num_scheduled_tokens _assert_right_encoder_cache_allocated(scheduler, expected_total_allocated=0) @@ -5652,6 +5715,31 @@ def test_ec_connector_ensure_cache_available_defers_request(use_kv_connector): _assert_right_encoder_inputs(output, expected_total_reqs=0) +def test_ec_connector_defers_running_request_for_async_reload(): + scheduler = create_scheduler( + model="llava-hf/llava-1.5-7b-hf", + max_num_batched_tokens=32, + use_ec_connector=True, + ec_role="ec_consumer", + ) + request = create_requests( + num_requests=1, + num_tokens=128, + mm_positions=[[PlaceholderRange(offset=48, length=32)]], + req_ids=["request"], + )[0] + scheduler.ec_connector.ensure_cache_available = Mock(side_effect=[True, False]) + + scheduler.add_request(request) + first_output = scheduler.schedule() + assert first_output.num_scheduled_tokens[request.request_id] == 32 + + second_output = scheduler.schedule() + assert request.request_id not in second_output.num_scheduled_tokens + ensure_call = scheduler.ec_connector.ensure_cache_available.call_args + assert ensure_call.args[:2] == (request, 32) + + def test_ec_connector_pending_prefetch_only_checks_future_mm_features(): """Test that future mm feature filtering only yields features beyond the computed token frontier. diff --git a/tests/v1/ec_connector/integration/README.md b/tests/v1/ec_connector/integration/README.md index 44a55ce558..c292ad361f 100644 --- a/tests/v1/ec_connector/integration/README.md +++ b/tests/v1/ec_connector/integration/README.md @@ -4,7 +4,7 @@ This test verifies that EPD (Encoder-Prefill-Decode) disaggregation produces ide ## What It Tests -- **Baseline**: Single Sonar instance serving a multimodal model +- **Baseline**: Single Aphrodite instance serving a multimodal model - **EPD (1E+1PD)**: 1 Encoder + 1 Prefill-Decode instance - **Baseline (1P+1D)**: 1 Prefill + 1 Decode instance - **EPD (1E+1P+1D)**: 1 Encoder + 1 Prefill + 1 Decode instance @@ -58,9 +58,48 @@ EC_SHARED_STORAGE_PATH="/tmp/my_ec_cache" bash ./tests/v1/ec_connector/integrati ## How It Works +### Mooncake EC (1E + 1PD) + +```bash +PYTHON_BIN="$PWD/.venv/bin/python" \ + bash tests/v1/ec_connector/integration/run_epd_mooncake_ec_full_pipeline.sh +``` + +Requires two GPUs and Mooncake TransferEngine. TCP is the default transport; +no RDMA-capable network hardware is required. +The script sets `MC_FORCE_TCP=1` in TCP mode so Mooncake cannot auto-select RDMA. +The baseline runs first on GPU 0, followed by E on GPU 0 and PD on GPU 1. +`MOONCAKE_EC_PROTOCOL=rdma` selects RDMA instead; host-specific transport +environment variables should be set by the caller. + +Three black-box cases check fixed short answers and compare with the baseline: + +- One image: read the STOP sign. +- Two different images (including a local file): identify flowers and birds. +- The same image twice: read both STOP signs. + +By default, all three requests run concurrently for two rounds, exercising +shared hashes across requests and reuse after completion. Every response is +compared, not just the final round. Set `CONCURRENCY` and `REPEAT` to override. +Prefix caching is disabled and CUDA graphs remain enabled. This is a small +correctness suite, not a performance benchmark or failure-injection suite. + +`LOG_PATH` and `BASELINE_FILE` select the log directory and reference output. +`SKIP_BASELINE=1` reuses a reference generated with the same model and test +configuration. `USE_MM_PROMPTS=0` only checks text routing, not EC transfer. + +This fork does not include the upstream Mooncake Buildkite job. Run the script +on a host with two GPUs and a CUDA-compatible Mooncake TransferEngine wheel +installed in the selected Python environment. The default configuration uses +loopback networking and TCP, without RDMA devices or peer-memory setup. + +The script fails on startup errors, request errors, or answer mismatches. On +failure, the script prints the last 100 lines of each server/proxy log to +the CI job log; full files remain under `LOG_PATH` while the container exists. + ### Step 1: Baseline -1. Start a single Sonar instance on the GPU. +1. Start a single Aphrodite instance on the GPU. 2. Run test prompts (multimodal or text-only) 3. Save outputs to `.aphrodite_epd_baseline.txt` 4. Shutdown instance diff --git a/tests/v1/ec_connector/integration/run_epd_mooncake_ec_full_pipeline.sh b/tests/v1/ec_connector/integration/run_epd_mooncake_ec_full_pipeline.sh new file mode 100755 index 0000000000..e5afb9cb2d --- /dev/null +++ b/tests/v1/ec_connector/integration/run_epd_mooncake_ec_full_pipeline.sh @@ -0,0 +1,289 @@ +#!/usr/bin/env bash +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +# +# Full-stack EPD validation with ECMooncakeConnector (Mooncake TransferEngine): +# 1) Single-GPU baseline (multimodal) -> saves baseline JSON +# 2) 1 Encoder + 1 PD with Mooncake EC -> compare outputs to baseline +# +# Usage (from repo root): +# ./tests/v1/ec_connector/integration/run_epd_mooncake_ec_full_pipeline.sh +# +# Env: +# MODEL HF model id (default: Qwen/Qwen2.5-VL-3B-Instruct) +# GPU_SINGLE / GPU_E / GPU_PD GPU ids (defaults 0 / 0 / 1) +# ENDPOINT_PORT, ENCODE_PORT, PREFILL_DECODE_PORT +# EC_MOONCAKE_RESERVATION_HOST consumer control host (default 127.0.0.1) +# MOONCAKE_EC_PROTOCOL tcp | rdma (default tcp) +# USE_MM_PROMPTS 1 (default) or 0 for text-only quick sanity +# TIMEOUT_SECONDS wait_for_server timeout (default 1200) +# SKIP_BASELINE set to 1 to reuse existing BASELINE_FILE +# CONCURRENCY / REPEAT concurrent requests / rounds (defaults 3 / 2) +# MAX_MODEL_LEN context length (default 16384) + +set -euo pipefail + +GIT_ROOT=$(cd "$(dirname "${BASH_SOURCE[0]}")/../../../.." && pwd) +cd "$GIT_ROOT" || exit 1 +export PYTHONPATH="${GIT_ROOT}:${PYTHONPATH:-}" +PYTHON_BIN="${PYTHON_BIN:-${GIT_ROOT}/.venv/bin/python}" + +MODEL="${MODEL:-Qwen/Qwen2.5-VL-3B-Instruct}" +USE_MM_PROMPTS="${USE_MM_PROMPTS:-1}" +TEST_ARGS=(--concurrency "${CONCURRENCY:-3}" --repeat "${REPEAT:-2}") +if [[ "$USE_MM_PROMPTS" == "1" ]]; then + TEST_ARGS+=(--use_mm_prompts --mm_smoke_test) +fi +MAX_MODEL_LEN="${MAX_MODEL_LEN:-16384}" + +GPU_SINGLE="${GPU_SINGLE:-0}" +GPU_E="${GPU_E:-0}" +GPU_PD="${GPU_PD:-1}" + +ENCODE_PORT="${ENCODE_PORT:-19534}" +PREFILL_DECODE_PORT="${PREFILL_DECODE_PORT:-19537}" +ENDPOINT_PORT="${ENDPOINT_PORT:-10002}" +BASELINE_PORT="${BASELINE_PORT:-10003}" + +EC_MOONCAKE_RESERVATION_HOST="${EC_MOONCAKE_RESERVATION_HOST:-127.0.0.1}" +EC_MOONCAKE_RESERVATION_PORT="${EC_MOONCAKE_RESERVATION_PORT:-19019}" +MOONCAKE_EC_PROTOCOL="${MOONCAKE_EC_PROTOCOL:-tcp}" +export EC_MOONCAKE_RESERVATION_HOST +export EC_MOONCAKE_RESERVATION_PORT +export MOONCAKE_EC_PROTOCOL +if [[ "$MOONCAKE_EC_PROTOCOL" == "tcp" ]]; then + # TransferEngine may otherwise auto-select RDMA on hosts with an HCA. + export MC_FORCE_TCP=1 +else + unset MC_FORCE_TCP +fi + +LOG_PATH="${LOG_PATH:-/tmp}" +BASELINE_FILE="${BASELINE_FILE:-/tmp/aphrodite_epd_mooncake_baseline.txt}" +TIMEOUT_SECONDS="${TIMEOUT_SECONDS:-1200}" +PIDS=() + +mkdir -p "$LOG_PATH" + +APHRODITE_SERVE=("$PYTHON_BIN" -m aphrodite.entrypoints.cli.main serve) + +ENC_EC_JSON=$("$PYTHON_BIN" </dev/null || return 1 + if curl --max-time 2 -fsS "http://localhost:${port}/health" >/dev/null 2>&1; then + return 0 + fi + sleep 2 + done + return 1 +} + +cleanup_instances() { + if ((${#PIDS[@]} == 0)); then + return + fi + echo "Cleaning up tracked Aphrodite / proxy processes..." + for pid in "${PIDS[@]}"; do + kill "$pid" 2>/dev/null || true + done + for _ in {1..10}; do + local alive=0 + for pid in "${PIDS[@]}"; do + if kill -0 "$pid" 2>/dev/null; then + alive=1 + fi + done + ((alive == 0)) && break + sleep 1 + done + for pid in "${PIDS[@]}"; do + if kill -0 "$pid" 2>/dev/null; then + kill -KILL "$pid" 2>/dev/null || true + fi + wait "$pid" 2>/dev/null || true + done + PIDS=() +} + +finish() { + local status=$? + cleanup_instances + if ((status != 0)); then + for log in "${LOG_PATH}"/mooncake_epd_*.log; do + [[ -f "$log" ]] || continue + echo "=== $log (last 100 lines) ===" + tail -100 "$log" + done + fi + exit "$status" +} + +trap finish EXIT +trap 'exit 130' INT +trap 'exit 143' TERM + +run_baseline() { + echo "================================" + echo "BASELINE (single Aphrodite, MM if enabled)" + echo "================================" + cleanup_instances + local PORT=$BASELINE_PORT + echo "Starting baseline on GPU $GPU_SINGLE port $PORT" + CUDA_VISIBLE_DEVICES="$GPU_SINGLE" "${APHRODITE_SERVE[@]}" "$MODEL" \ + --port "$PORT" \ + --gpu-memory-utilization 0.75 \ + --max-num-seqs 32 \ + --max-model-len "$MAX_MODEL_LEN" \ + --no-enable-prefix-caching \ + --limit-mm-per-prompt '{"image":2,"video":0}' \ + --allowed-local-media-path "${GIT_ROOT}/tests/v1/ec_connector/integration" \ + >"${LOG_PATH}/mooncake_epd_baseline.log" 2>&1 & + local BASELINE_PID=$! + PIDS+=("$BASELINE_PID") + echo "Waiting for baseline..." + wait_for_server "$PORT" "$BASELINE_PID" || { echo "Baseline failed to start"; return 1; } + curl -s "http://127.0.0.1:${PORT}/v1/models" | head -c 200 || true + echo "" + "$PYTHON_BIN" "${GIT_ROOT}/tests/v1/ec_connector/integration/test_epd_correctness.py" \ + --service_url "http://localhost:$PORT" \ + --model_name "$MODEL" \ + --mode baseline \ + --baseline_file "$BASELINE_FILE" \ + "${TEST_ARGS[@]}" + cleanup_instances +} + +run_epd_mooncake() { + echo "================================" + echo "EPD 1E + 1PD with ECMooncakeConnector" + echo "Reservation port on consumer: $EC_MOONCAKE_RESERVATION_PORT" + echo "Mooncake protocol: $MOONCAKE_EC_PROTOCOL" + echo "================================" + cleanup_instances + + echo "Starting ENCODER on GPU $GPU_E port $ENCODE_PORT" + CUDA_VISIBLE_DEVICES="$GPU_E" "${APHRODITE_SERVE[@]}" "$MODEL" \ + --port "$ENCODE_PORT" \ + --gpu-memory-utilization 0.35 \ + --mm-tensor-ipc torch_shm \ + --enable-request-id-headers \ + --no-enable-prefix-caching \ + --max-num-batched-tokens "$MAX_MODEL_LEN" \ + --max-num-seqs 32 \ + --max-model-len "$MAX_MODEL_LEN" \ + --limit-mm-per-prompt '{"image":2,"video":0}' \ + --allowed-local-media-path "${GIT_ROOT}/tests/v1/ec_connector/integration" \ + --ec-transfer-config "$ENC_EC_JSON" \ + >"${LOG_PATH}/mooncake_epd_encoder.log" 2>&1 & + local ENCODER_PID=$! + PIDS+=("$ENCODER_PID") + + echo "Starting PD on GPU $GPU_PD port $PREFILL_DECODE_PORT" + CUDA_VISIBLE_DEVICES="$GPU_PD" "${APHRODITE_SERVE[@]}" "$MODEL" \ + --port "$PREFILL_DECODE_PORT" \ + --gpu-memory-utilization 0.75 \ + --mm-tensor-ipc torch_shm \ + --enable-mm-embeds \ + --enable-request-id-headers \ + --max-num-seqs 32 \ + --max-model-len "$MAX_MODEL_LEN" \ + --no-enable-prefix-caching \ + --limit-mm-per-prompt '{"image":2,"video":0}' \ + --allowed-local-media-path "${GIT_ROOT}/tests/v1/ec_connector/integration" \ + --ec-transfer-config "$PD_EC_JSON" \ + >"${LOG_PATH}/mooncake_epd_pd.log" 2>&1 & + local PD_PID=$! + PIDS+=("$PD_PID") + + echo "Waiting for encoder..." + wait_for_server "$ENCODE_PORT" "$ENCODER_PID" || { echo "Encoder failed to start"; return 1; } + echo "Waiting for PD..." + wait_for_server "$PREFILL_DECODE_PORT" "$PD_PID" || { echo "PD failed to start"; return 1; } + + echo "Starting EPD proxy on $ENDPOINT_PORT" + "$PYTHON_BIN" "${GIT_ROOT}/examples/disaggregated/disaggregated_encoder/disagg_epd_proxy.py" \ + --host "0.0.0.0" \ + --port "$ENDPOINT_PORT" \ + --encode-servers-urls "http://localhost:$ENCODE_PORT" \ + --prefill-servers-urls "disable" \ + --decode-servers-urls "http://localhost:$PREFILL_DECODE_PORT" \ + --ec-consumer-zmq-addrs \ + "$EC_MOONCAKE_RESERVATION_ADDR" \ + >"${LOG_PATH}/mooncake_epd_proxy.log" 2>&1 & + local PROXY_PID=$! + PIDS+=("$PROXY_PID") + + echo "Waiting for proxy..." + wait_for_server "$ENDPOINT_PORT" "$PROXY_PID" || { echo "Proxy failed to start"; return 1; } + curl -s "http://127.0.0.1:${ENDPOINT_PORT}/health" || true + echo "" + + "$PYTHON_BIN" "${GIT_ROOT}/tests/v1/ec_connector/integration/test_epd_correctness.py" \ + --service_url "http://localhost:$ENDPOINT_PORT" \ + --model_name "$MODEL" \ + --mode disagg \ + --baseline_file "$BASELINE_FILE" \ + "${TEST_ARGS[@]}" + + cleanup_instances +} + +echo "================================" +echo "1E + 1PD ECMooncake end-to-end correctness" +echo "MODEL=$MODEL" +echo "================================" + +if [[ "${SKIP_BASELINE:-0}" != "1" ]]; then + run_baseline +else + echo "SKIP_BASELINE=1 -> using existing $BASELINE_FILE" + [[ -f "$BASELINE_FILE" ]] || { echo "Missing baseline file"; exit 1; } +fi + +run_epd_mooncake + +echo "================================" +echo "PASS: Mooncake EC EPD matches baseline" +echo "Logs: ${LOG_PATH}/mooncake_epd_*.log" +echo "================================" diff --git a/tests/v1/ec_connector/integration/test_epd_correctness.py b/tests/v1/ec_connector/integration/test_epd_correctness.py index 11527316b0..600ad0bc20 100644 --- a/tests/v1/ec_connector/integration/test_epd_correctness.py +++ b/tests/v1/ec_connector/integration/test_epd_correctness.py @@ -26,6 +26,7 @@ import json import os import time +from concurrent.futures import ThreadPoolExecutor import openai import requests @@ -133,17 +134,18 @@ def run_chat_completion( Returns: Generated text content """ - client = openai.OpenAI(api_key="EMPTY", base_url=base_url) - - completion = client.chat.completions.create( - model=model_name, - messages=messages, - max_tokens=max_tokens, - temperature=0.0, - seed=42, - ) + with openai.OpenAI(api_key="EMPTY", base_url=base_url, timeout=120, max_retries=0) as client: + completion = client.chat.completions.create( + model=model_name, + messages=messages, + max_tokens=max_tokens, + temperature=0.0, + seed=42, + ) - return completion.choices[0].message.content + content = completion.choices[0].message.content + assert content, "Expected a nonempty completion" + return content def main(): @@ -191,7 +193,19 @@ def main(): help="Skip the two-image multimodal prompt", ) + parser.add_argument( + "--mm_smoke_test", + action="store_true", + help="Use three short, fixed-answer image cases, including duplicate images", + ) + parser.add_argument("--concurrency", type=int, default=1) + parser.add_argument("--repeat", type=int, default=1) + args = parser.parse_args() + if args.concurrency < 1 or args.repeat < 1: + parser.error("--concurrency and --repeat must be positive") + if args.mm_smoke_test and (not args.use_mm_prompts or args.skip_two_image_prompt): + parser.error("--mm_smoke_test requires all multimodal prompts") print(f"Service URL: {args.service_url}") print(f"Model: {args.model_name}") @@ -226,24 +240,73 @@ def main(): test_prompts = SAMPLE_PROMPTS_TEXT print("Using text-only prompts for quick testing") + if args.mm_smoke_test: + image = SAMPLE_PROMPTS_MM[0]["messages"][0]["content"][0] + image_pair = SAMPLE_PROMPTS_MM[1]["messages"][0]["content"][:2] + cases = [ + ( + "Single image", + [image], + "What word is on the red road sign? Reply with only that word in uppercase.", + "STOP", + ), + ( + "Two different images", + image_pair, + "Do these pictures contain both flowers and birds? Reply with only YES or NO.", + "YES", + ), + ( + "Same image twice", + [image, image], + "Read the red road sign in each image, in image order. Reply " + "with only the two uppercase words separated by a comma and a space.", + "STOP, STOP", + ), + ] + test_prompts = [ + { + "description": description, + "expected": expected, + "messages": [ + { + "role": "user", + "content": [*images, {"type": "text", "text": question}], + } + ], + } + for description, images, question, expected in cases + ] + # Run completions service_url = f"{args.service_url}/v1" output_strs = {} - for i, prompt_data in enumerate(test_prompts): - print(f"\nRunning prompt {i + 1}/{len(test_prompts)}: {prompt_data['description']}") - - output_str = run_chat_completion( + def complete(prompt_data): + output = run_chat_completion( base_url=service_url, model_name=args.model_name, messages=prompt_data["messages"], - max_tokens=MAX_OUTPUT_LEN, + max_tokens=16 if args.mm_smoke_test else MAX_OUTPUT_LEN, ) - - # Use description as key for comparison - key = prompt_data["description"] - output_strs[key] = output_str - print(f"Output: {output_str}") + if args.mm_smoke_test: + output = output.strip() + assert output == prompt_data["expected"], ( + f"{prompt_data['description']}: expected {prompt_data['expected']!r}, got {output!r}" + ) + return output + + # Each round includes concurrent requests sharing image hashes; later rounds + # exercise reuse after previous requests have finished. + with ThreadPoolExecutor(max_workers=args.concurrency) as executor: + for repeat in range(args.repeat): + outputs = executor.map(complete, test_prompts) + for prompt_data, output_str in zip(test_prompts, outputs): + key = prompt_data["description"] + if args.repeat > 1: + key = f"{key} (round {repeat + 1})" + output_strs[key] = output_str + print(f"{key}: {output_str}") if args.mode in ("baseline", "baseline_pd"): # Baseline mode: Save outputs diff --git a/tests/v1/ec_connector/unit/test_epd_proxy_retry.py b/tests/v1/ec_connector/unit/test_epd_proxy_retry.py new file mode 100644 index 0000000000..423b2ac45b --- /dev/null +++ b/tests/v1/ec_connector/unit/test_epd_proxy_retry.py @@ -0,0 +1,164 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +"""Body handling across the EPD proxy's decode retries. + +Exercises the REAL helpers loaded from the ``examples/`` proxy, so a future +change to them is what these tests catch. +""" + +import asyncio +import importlib.util +from pathlib import Path + +import pytest + +PROXY_REL = "examples/disaggregated/disaggregated_encoder/disagg_epd_proxy.py" + + +@pytest.fixture(scope="module") +def proxy(): + path = Path(__file__).parents[4] / PROXY_REL + spec = importlib.util.spec_from_file_location("disagg_epd_proxy_retry", path) + assert spec is not None and spec.loader is not None + module = importlib.util.module_from_spec(spec) + spec.loader.exec_module(module) + return module + + +class _Response: + def __init__(self, params): + self._params = params + + async def json(self): + return {"kv_transfer_params": self._params} + + +def test_maybe_prefill_leaves_the_caller_body_untouched(proxy, monkeypatch): + """A decode retry re-enters this function with the body it was given. + + Mutating that body in place let one attempt's `remote_block_ids` survive + into the next, so a retry whose prefill returns nothing sent decode blocks + the prefiller may already have freed. + """ + served = [{"remote_block_ids": [1, 2]}, {}] + + async def _stage(req_data, p_url, req_id): + assert "kv_transfer_params" not in req_data + return _Response(served.pop(0)) + + monkeypatch.setattr(proxy, "process_prefill_stage", _stage) + + body = {"messages": [], "stream": False} + first = asyncio.run(proxy.maybe_prefill(body, "http://prefill", "r1")) + assert first["kv_transfer_params"] == {"remote_block_ids": [1, 2]} + assert "kv_transfer_params" not in body + + second = asyncio.run(proxy.maybe_prefill(body, "http://prefill", "r1")) + assert "kv_transfer_params" not in second + + +class _EncoderResponse: + def __init__(self, params): + self.status = 200 + self._params = params + + async def json(self): + return {"ec_transfer_params": self._params} + + async def text(self): + return "" + + +class _EncoderSession: + """Serve one canned encoder reply per attempt.""" + + def __init__(self, replies): + self._replies = list(replies) + + async def post(self, url, json=None, headers=None): + return _EncoderResponse(self._replies.pop(0)) + + +def test_a_decode_retry_does_not_inherit_the_previous_handles(proxy, monkeypatch): + """The retry loop re-enters `prepare_for_decode` with the same body. + + Recording the encoder's connector handles on that body in place let + attempt 1's handle survive into attempt 2, so a second encode that + reported nothing still sent decode a handle on an embedding the encoder + no longer publishes -- the exact state the retry exists to leave behind. + """ + handle = {"metadata": {"image_grid_thw": [1, 2, 2]}, "peer_port": 1234} + monkeypatch.setattr( + proxy, + "encode_session", + _EncoderSession([{"encoder-side-hash": handle}, {}]), + ) + + async def _no_prefill(req_data, p_url, req_id): + return req_data + + monkeypatch.setattr(proxy, "maybe_prefill", _no_prefill) + + body = { + "messages": [ + { + "role": "user", + "content": [{"type": "image_url", "image_url": {"url": "image"}}], + } + ], + "stream": False, + } + args = ("r1", ["http://encoder"], "http://prefill", None) + + first, _, _ = asyncio.run(proxy.prepare_for_decode(body, *args)) + reported = first["ec_transfer_params"] + assert [handle] == [value for key, value in reported.items() if key != "ec_items"] + assert "ec_transfer_params" not in body + + # Attempt 2's encode reports nothing: decode must be told nothing. + second, _, _ = asyncio.run(proxy.prepare_for_decode(body, *args)) + assert "ec_transfer_params" not in second + + +def test_raw_media_keeps_encoder_transfer_identity(proxy, monkeypatch): + handle = {"metadata": {"image_grid_thw": [1, 2, 2]}} + monkeypatch.setattr(proxy, "NO_REWRITE", True) + monkeypatch.setattr(proxy, "encode_session", _EncoderSession([{"encoded-hash": handle}])) + body = { + "messages": [ + { + "role": "user", + "content": [{"type": "image_url", "image_url": {"url": "image"}}], + } + ] + } + prepared, _, _ = asyncio.run(proxy.prepare_for_decode(body, "request", ["http://encoder"], "", "tcp://consumer:1")) + assert prepared["messages"] == body["messages"] + assert "ec_transfer_params" not in body + params = prepared["ec_transfer_params"] + assert params["encoded-hash"] == handle + assert params["ec_items"][0]["mm_hash"] == "encoded-hash" + assert params["ec_items"][0]["transfer_id"] + + +@pytest.mark.parametrize("server_keep_alive", ["0.1", "1", "5", "2", "30"]) +def test_pooled_connections_are_retired_before_the_server_closes_them(proxy, monkeypatch, server_keep_alive): + """The proxy must not hand a request a connection the server has dropped. + + Aphrodite closes idle keep-alive connections at `APHRODITE_HTTP_TIMEOUT_KEEP_ALIVE` + seconds while aiohttp pools them for 15, so a slow hop leaves a dead + connection in the pool. The next request fails with + ServerDisconnectedError and the server logs nothing, because it closed the + socket before the request arrived. + """ + monkeypatch.setenv("APHRODITE_HTTP_TIMEOUT_KEEP_ALIVE", server_keep_alive) + monkeypatch.setattr(proxy.app.state, "p_urls", [], raising=False) + + asyncio.run(proxy.on_startup()) + try: + # No public accessor for the pool's idle timeout. + pooled_for = proxy.encode_session.connector._keepalive_timeout + assert pooled_for < float(server_keep_alive) + assert proxy.decode_session.connector._keepalive_timeout == pooled_for + finally: + asyncio.run(proxy.on_shutdown()) diff --git a/tests/v1/ec_connector/unit/test_epd_proxy_round_robin.py b/tests/v1/ec_connector/unit/test_epd_proxy_round_robin.py index 48123235a7..c7278ad55b 100644 --- a/tests/v1/ec_connector/unit/test_epd_proxy_round_robin.py +++ b/tests/v1/ec_connector/unit/test_epd_proxy_round_robin.py @@ -32,8 +32,13 @@ def _load_proxy_module(): @pytest.fixture(scope="module") -def assign(): - return _load_proxy_module().encoder_rr_assignment +def proxy(): + return _load_proxy_module() + + +@pytest.fixture(scope="module") +def assign(proxy): + return proxy.encoder_rr_assignment def _drive(assign, e_urls, counts): @@ -80,3 +85,44 @@ def test_single_encoder_always_resolves_to_it(assign): urls, next_cursor = assign(["E0"], 0, 4) assert urls == ["E0"] * 4 assert next_cursor == 0 + + +def test_generic_e_p_d_routing_without_mooncake_is_supported(proxy): + proxy.validate_ec_consumer_routing(["http://prefill"], []) + + +def test_mooncake_e_pd_routing_is_supported(proxy): + proxy.validate_ec_consumer_routing([], ["tcp://decode:19019"]) + + +def test_mooncake_independent_prefill_routing_fails_fast(proxy): + with pytest.raises(ValueError, match=r"supports E\+PD only"): + proxy.validate_ec_consumer_routing(["http://prefill"], ["tcp://decode:19019"]) + + +def test_decode_rewrite_preserves_engine_reported_ec_hash(proxy): + request = { + "messages": [ + { + "role": "user", + "content": [{"type": "image_url", "image_url": {"url": "image"}}], + } + ] + } + + rewritten = proxy.rewrite_for_decode( + request, + { + 0: { + "mm_hash": "proxy-uuid", + "ec_mm_hash": "engine-derived-hash", + "transfer_id": "transfer", + "image_grid_thw": [1, 2, 3], + } + }, + ) + + assert rewritten["messages"][0]["content"][0]["uuid"] == "proxy-uuid" + assert rewritten["ec_transfer_params"]["ec_items"] == [ + {"mm_hash": "engine-derived-hash", "transfer_id": "transfer"} + ] diff --git a/tests/v1/ec_connector/unit/test_worker_ec_connector.py b/tests/v1/ec_connector/unit/test_worker_ec_connector.py index 4c6ed772f7..c4ce55a75e 100644 --- a/tests/v1/ec_connector/unit/test_worker_ec_connector.py +++ b/tests/v1/ec_connector/unit/test_worker_ec_connector.py @@ -12,6 +12,9 @@ ECConnectorMetadata, ) from aphrodite.v1.outputs import EMPTY_MODEL_RUNNER_OUTPUT +from aphrodite.v1.worker.ec_connector_model_runner_mixin import ( + ECConnectorModelRunnerMixin, +) from aphrodite.v1.worker.gpu.ec_connector import NO_OP_EC_CONNECTOR, ActiveECConnector pytestmark = pytest.mark.cpu_test @@ -48,6 +51,9 @@ def test_saves_newly_added_caches_for_every_producer(is_producer, is_consumer): saved = [call.kwargs["mm_hash"] for call in fake.save_caches.call_args_list] assert saved == (["mm_new"] if is_producer else []) + assert fake.start_save_caches.called == is_producer + if is_producer: + assert fake.start_save_caches.call_args.kwargs["encoder_cache"] is encoder_cache assert fake.start_load_caches.called == is_consumer @@ -62,6 +68,24 @@ def test_worker_meta_is_reported_on_context_exit(): assert fake.clear_connector_metadata.called +def test_v1_producer_receives_existing_encoder_cache(): + encoder_cache = {"mm_cached": None} + fake = MagicMock(spec=ECConnectorBase) + fake.is_producer = True + fake.is_consumer = False + fake.get_finished.return_value = (None, None) + + module = "aphrodite.v1.worker.ec_connector_model_runner_mixin" + with ( + patch(f"{module}.has_ec_transfer", return_value=True), + patch(f"{module}.get_ec_transfer", return_value=fake), + ECConnectorModelRunnerMixin.maybe_get_ec_connector_output(_scheduler_output(), encoder_cache), + ): + pass + + assert fake.start_save_caches.call_args.kwargs["encoder_cache"] is encoder_cache + + def test_no_forward_reports_without_running_the_model(): connector, _ = _connector() From b424bff5dd8d845bd257722e464dfa1c854fe65c Mon Sep 17 00:00:00 2001 From: AlpinDale Date: Mon, 7 Sep 2026 19:58:31 +0000 Subject: [PATCH 14/36] [sync] [CI][ROCm] Temporarily skip unsupported HY-V4 initialization (#55653) Upstream-vLLM: 195bc9c4a1b480c7d2a6131cd8dbb49e9be72762 Co-authored-by: Andreas Karatzas --- .sync/vllm-sha | 2 +- tests/models/test_initialization.py | 6 ++++++ 2 files changed, 7 insertions(+), 1 deletion(-) diff --git a/.sync/vllm-sha b/.sync/vllm-sha index 223e536eb2..09572c4bdc 100644 --- a/.sync/vllm-sha +++ b/.sync/vllm-sha @@ -1 +1 @@ -6fbb00b18874e27ba7d7adc0a3b8e93fee763ab1 +195bc9c4a1b480c7d2a6131cd8dbb49e9be72762 diff --git a/tests/models/test_initialization.py b/tests/models/test_initialization.py index c9bf3aed99..de14bfc259 100644 --- a/tests/models/test_initialization.py +++ b/tests/models/test_initialization.py @@ -206,6 +206,12 @@ def test_can_initialize_large_subset(model_arch: str, monkeypatch: pytest.Monkey This test covers the complement of the tests covered in the "small subset" test. """ + if model_arch in ("HYV4ForCausalLM", "HYV4MTPModel"): + from aphrodite.platforms import current_platform + + if current_platform.is_rocm(): + pytest.skip("HY V4 ROCm initialization requires #54405") + can_initialize(model_arch, monkeypatch, HF_EXAMPLE_MODELS) From 25df714523a01cdd69f8039339e9d359e90d0d91 Mon Sep 17 00:00:00 2001 From: AlpinDale Date: Mon, 7 Sep 2026 20:01:05 +0000 Subject: [PATCH 15/36] [sync] [Bugfix][MooncakeStore] Fix finish-time save crash on hybrid models (#54643) Upstream-vLLM: 34b1e9f7a61e802b870c7d7f78bf90b000b691cc Co-authored-by: Zhewen Li --- .sync/vllm-sha | 2 +- .../v1/mooncake/store/scheduler.py | 7 ++++- .../unit/test_mooncake_store_scheduler.py | 27 +++++++++++++++++++ 3 files changed, 34 insertions(+), 2 deletions(-) diff --git a/.sync/vllm-sha b/.sync/vllm-sha index 09572c4bdc..700864a871 100644 --- a/.sync/vllm-sha +++ b/.sync/vllm-sha @@ -1 +1 @@ -195bc9c4a1b480c7d2a6131cd8dbb49e9be72762 +34b1e9f7a61e802b870c7d7f78bf90b000b691cc diff --git a/aphrodite/distributed/kv_transfer/kv_connector/v1/mooncake/store/scheduler.py b/aphrodite/distributed/kv_transfer/kv_connector/v1/mooncake/store/scheduler.py index e3f84c5024..3d22151dd7 100644 --- a/aphrodite/distributed/kv_transfer/kv_connector/v1/mooncake/store/scheduler.py +++ b/aphrodite/distributed/kv_transfer/kv_connector/v1/mooncake/store/scheduler.py @@ -342,7 +342,12 @@ def build_connector_meta(self, scheduler_output: SchedulerOutput) -> KVConnector ) in self._unfinished_requests.items(): if request_id not in request_ids and request_id not in cached_reqs.req_ids: load_spec = self.load_specs.pop(request_id, None) - if not load_spec: + # A load spec may have been proposed by this connector's + # lookup but rejected by MultiConnector in favor of another + # connector. Only the chosen connector may issue the pending + # load; the normal store path gets its blocks later from + # SchedulerOutput once the request is actually scheduled. + if load_spec is None or not load_spec.can_load: continue num_tokens_to_compute = load_spec.kvpool_cached_tokens request_tracker = RequestTracker( diff --git a/tests/v1/kv_connector/unit/test_mooncake_store_scheduler.py b/tests/v1/kv_connector/unit/test_mooncake_store_scheduler.py index 8584a67926..7afcea1922 100644 --- a/tests/v1/kv_connector/unit/test_mooncake_store_scheduler.py +++ b/tests/v1/kv_connector/unit/test_mooncake_store_scheduler.py @@ -187,6 +187,33 @@ def _add_unfinished_request( ) +def test_pending_load_for_non_chosen_connector_is_dropped(): + """A MultiConnector loser must not turn its proposed load into a save.""" + scheduler = _make_bare_scheduler() + request = SimpleNamespace( + request_id="req-0", + block_hashes=[b"h0", b"h1", b"h2"], + ) + blocks = SimpleNamespace(get_block_ids=lambda: ([1, 2], [9])) + scheduler.load_specs["req-0"] = LoadSpec( + aphrodite_cached_tokens=0, + kvpool_cached_tokens=48, + can_load=False, + ) + + scheduler.update_state_after_alloc(request, blocks, num_external_tokens=0) + meta = scheduler.build_connector_meta(_make_pending_load_scheduler_output()) + + # MultiConnector exposes the real allocation to every child, but only the + # chosen child receives external tokens. The losing store connector must + # neither retain those blocks for a pending load nor enqueue a save from + # the rejected speculative LoadSpec. + assert scheduler._unfinished_requests["req-0"][1] == () + assert meta.requests == [] + assert "req-0" not in scheduler.load_specs + assert "req-0" not in scheduler._request_trackers + + def test_update_state_excludes_nontransfer_groups(): """Store metadata must match the worker's registered cache groups.""" scheduler = _make_bare_scheduler() From 79efeca83ebfe8de3f3af6ffa7cac592ca20aa79 Mon Sep 17 00:00:00 2001 From: AlpinDale Date: Mon, 7 Sep 2026 20:04:26 +0000 Subject: [PATCH 16/36] [sync] [Bugfix] DSv4 MXFP4 selector: stop narrowing explicit aliases to their BF16 variant (#53586) Upstream-vLLM: 5893426b88f7b3cd21101d194eb1c6f0a6f0e27b Co-authored-by: Gabriel Wu <13583761+lucifer1004@users.noreply.github.com> --- .sync/vllm-sha | 2 +- .../layers/fused_moe/oracle/mxfp4.py | 8 +++++- tests/kernels/moe/test_b12x.py | 27 +++++++++++++++++++ 3 files changed, 35 insertions(+), 2 deletions(-) diff --git a/.sync/vllm-sha b/.sync/vllm-sha index 700864a871..2a5567fa0b 100644 --- a/.sync/vllm-sha +++ b/.sync/vllm-sha @@ -1 +1 @@ -34b1e9f7a61e802b870c7d7f78bf90b000b691cc +5893426b88f7b3cd21101d194eb1c6f0a6f0e27b diff --git a/aphrodite/model_executor/layers/fused_moe/oracle/mxfp4.py b/aphrodite/model_executor/layers/fused_moe/oracle/mxfp4.py index 530f25707d..238e8cbdce 100644 --- a/aphrodite/model_executor/layers/fused_moe/oracle/mxfp4.py +++ b/aphrodite/model_executor/layers/fused_moe/oracle/mxfp4.py @@ -643,7 +643,13 @@ def select_deepseek_v4_mxfp4_moe_backend( # falling back to the auto priority list. runner_backend = config.moe_backend if runner_backend != "auto": - requested_backends = _get_requested_backends(runner_backend, None) + if runner_backend == "b12x": + requested_backends = _get_requested_backends(runner_backend, None) + else: + # Try every variant of the alias in priority order. Narrowing to + # the BF16 variant would drop SM100+ W4A8 variants on devices where + # the BF16 variant is gated to SM90. + requested_backends = map_mxfp4_backend(runner_backend) if activation_format == mk.FusedMoEActivationFormat.BatchedExperts: requested_backends = [ Mxfp4MoeBackend.BATCHED_MARLIN if b == Mxfp4MoeBackend.MARLIN else b for b in requested_backends diff --git a/tests/kernels/moe/test_b12x.py b/tests/kernels/moe/test_b12x.py index e47a8cdbdb..7150f29092 100644 --- a/tests/kernels/moe/test_b12x.py +++ b/tests/kernels/moe/test_b12x.py @@ -444,6 +444,33 @@ def test_deepseek_v4_b12x_activation_selection( assert experts_cls is B12xExperts +def test_deepseek_v4_flashinfer_cutlass_falls_through_to_w4a8( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """Explicit flashinfer_cutlass must try the W4A8 variant when the BF16 + variant is unsupported (the BF16 variant is gated to SM90).""" + from aphrodite.model_executor.layers.fused_moe.experts.flashinfer_cutlass_moe import ( + FlashInferExperts, + ) + + monkeypatch.setattr(FlashInferExperts, "_supports_current_device", lambda: True) + + def sm120_quant_gate(weight_key, activation_key): + return (weight_key, activation_key) == ( + mxfp4_oracle.kMxfp4Static, + mxfp4_oracle.kMxfp8Dynamic, + ) + + monkeypatch.setattr(FlashInferExperts, "_supports_quant_scheme", staticmethod(sm120_quant_gate)) + config = make_dummy_moe_config(hidden_dim=256, intermediate_size=64) + config.moe_backend = "flashinfer_cutlass" + + backend, experts_cls = select_deepseek_v4_mxfp4_moe_backend(config) + + assert backend == Mxfp4MoeBackend.FLASHINFER_CUTLASS_MXFP4_MXFP8 + assert experts_cls is FlashInferExperts + + def test_compressed_tensors_mxfp4_preserves_checkpoint_packing( monkeypatch: pytest.MonkeyPatch, ) -> None: From 8d61033cdf45bfe5992a024dfaff5febf8d84670 Mon Sep 17 00:00:00 2001 From: AlpinDale Date: Mon, 7 Sep 2026 20:07:00 +0000 Subject: [PATCH 17/36] [sync] [KVConnector] Guard lmcache_mp_connector state transition with num_external_tokens (#47505) Upstream-vLLM: 49eb2accf0610d220f451f2f3c1dfe93fe395c79 Co-authored-by: AlexHuang --- .sync/vllm-sha | 2 +- .../kv_transfer/kv_connector/v1/lmcache_mp_connector.py | 8 ++++++-- 2 files changed, 7 insertions(+), 3 deletions(-) diff --git a/.sync/vllm-sha b/.sync/vllm-sha index 2a5567fa0b..cea62426f8 100644 --- a/.sync/vllm-sha +++ b/.sync/vllm-sha @@ -1 +1 @@ -5893426b88f7b3cd21101d194eb1c6f0a6f0e27b +49eb2accf0610d220f451f2f3c1dfe93fe395c79 diff --git a/aphrodite/distributed/kv_transfer/kv_connector/v1/lmcache_mp_connector.py b/aphrodite/distributed/kv_transfer/kv_connector/v1/lmcache_mp_connector.py index 828d90b359..cb6eeeb15c 100644 --- a/aphrodite/distributed/kv_transfer/kv_connector/v1/lmcache_mp_connector.py +++ b/aphrodite/distributed/kv_transfer/kv_connector/v1/lmcache_mp_connector.py @@ -799,8 +799,12 @@ def update_state_after_alloc(self, request: "Request", blocks: "KVCacheBlocks", if new_block_ids: tracker.append_block_ids(new_block_ids) - # Update the state of the tracker - condition = tracker.needs_retrieve() + # num_external_tokens > 0 means this connector is chosen for loading. + # When non-chosen (num_external_tokens == 0, e.g. in MultiConnector), + # force condition=False so the state goes to READY and all lookup locks + # are released — end_session does NOT release lookup locks, so without + # this they would leak until TTL expiry (10 minutes). + condition = tracker.needs_retrieve() and num_external_tokens > 0 if tracker.state == LMCacheMPRequestState.PREFETCHING: # If need to retrieve, change to WAITING_FOR_LOAD # Otherwise, change to READY From a0647f34fbfdc053d9d959157c9bac1cf8aca8bb Mon Sep 17 00:00:00 2001 From: AlpinDale Date: Mon, 7 Sep 2026 20:10:36 +0000 Subject: [PATCH 18/36] [sync] [CI] [Test] skip test_wna16_cuda_high_bit_skips_humming on non-CUDA platforms (#55660) Upstream-vLLM: 392db567b260cb4d1e5f91547b65d527038fe641 Co-authored-by: Chaojun Zhang --- .sync/vllm-sha | 2 +- tests/quantization/test_auto_round.py | 4 ++++ 2 files changed, 5 insertions(+), 1 deletion(-) diff --git a/.sync/vllm-sha b/.sync/vllm-sha index cea62426f8..717498c9ef 100644 --- a/.sync/vllm-sha +++ b/.sync/vllm-sha @@ -1 +1 @@ -49eb2accf0610d220f451f2f3c1dfe93fe395c79 +392db567b260cb4d1e5f91547b65d527038fe641 diff --git a/tests/quantization/test_auto_round.py b/tests/quantization/test_auto_round.py index cd9397bc5f..358dea1032 100644 --- a/tests/quantization/test_auto_round.py +++ b/tests/quantization/test_auto_round.py @@ -734,6 +734,10 @@ class DummyMoeConfig: assert captured["layer_config"] is layer_config +@pytest.mark.skipif( + not current_platform.is_cuda(), + reason="This test only exercises the CUDA Marlin path.", +) @pytest.mark.parametrize("bits", [4, 8]) def test_wna16_cuda_high_bit_skips_humming(monkeypatch, bits) -> None: """4/8-bit int stays on the Marlin/GPTQ/AWQ path even on CUDA so a single From 1b116b9a2efe575ff1838023730979e8f49b23fa Mon Sep 17 00:00:00 2001 From: AlpinDale Date: Mon, 7 Sep 2026 20:15:49 +0000 Subject: [PATCH 19/36] [sync] [ROCm][CI] Add attention-sink support to ROCm AITER sparse MLA (#54404) Upstream-vLLM: e476556189506df612218c0888dc675e6ecbfc06 Co-authored-by: Andreas Karatzas --- .sync/vllm-sha | 2 +- aphrodite/_aiter_ops.py | 136 +++++ aphrodite/v1/attention/backend.py | 9 + .../backends/mla/rocm_aiter_mla_sparse.py | 238 +++++++-- .../v1/attention/ops/rocm_aiter_mla_sparse.py | 6 +- .../test_rocm_aiter_mla_op_registration.py | 56 ++- .../attention/test_rocm_aiter_mla_sink.py | 475 ++++++++++++++++++ ...est_rocm_aiter_mla_sparse_metadata_sync.py | 19 + .../v1/attention/test_rocm_glm5next_sparse.py | 48 ++ 9 files changed, 920 insertions(+), 69 deletions(-) create mode 100644 tests/kernels/attention/test_rocm_aiter_mla_sink.py diff --git a/.sync/vllm-sha b/.sync/vllm-sha index 717498c9ef..c5363f21df 100644 --- a/.sync/vllm-sha +++ b/.sync/vllm-sha @@ -1 +1 @@ -392db567b260cb4d1e5f91547b65d527038fe641 +e476556189506df612218c0888dc675e6ecbfc06 diff --git a/aphrodite/_aiter_ops.py b/aphrodite/_aiter_ops.py index 3fd69c7c7e..03f2b3f502 100644 --- a/aphrodite/_aiter_ops.py +++ b/aphrodite/_aiter_ops.py @@ -592,6 +592,104 @@ def _rocm_aiter_mla_decode_fwd_impl( ) +def _rocm_aiter_mla_decode_fwd_lse_impl( + q: torch.Tensor, + kv_buffer: torch.Tensor, + o: torch.Tensor, + qo_indptr: torch.Tensor, + max_seqlen_qo: int, + kv_indptr: torch.Tensor | None = None, + kv_indices: torch.Tensor | None = None, + kv_last_page_lens: torch.Tensor | None = None, + sm_scale: float = 1.0, + logit_cap: float = 0.0, + q_scale: torch.Tensor | None = None, + kv_scale: torch.Tensor | None = None, +) -> torch.Tensor: + """Run non-persistent AITER MLA decode and return natural-log LSE. + + gfx942 persistent MLA kernels do not provide an LSE code object. Keeping + this as a separate op makes that constraint structural: callers that need + LSE cannot accidentally forward persistent work metadata. + """ + from aiter.mla import get_meta_param, mla_decode_fwd + + kwargs: dict[str, float | int | torch.Tensor | None | bool] = { + "sm_scale": sm_scale, + "logit_cap": logit_cap, + "return_lse": True, + } + if _check_aiter_mla_fp8_support(): + kwargs["q_scale"] = q_scale + kwargs["kv_scale"] = kv_scale + + if q.dtype == torch.bfloat16 and kv_buffer.dtype == torch.bfloat16 and q.shape[1] == 32: + from aphrodite.platforms.rocm import on_gfx950 + + if on_gfx950(): + # The gfx950 H32 BF16 kernel mishandles ragged tails with split KV. + # Its single-split path returns correct output and natural-log LSE. + kwargs["num_kv_splits"] = 1 + kwargs["num_kv_splits_indptr"] = torch.arange(qo_indptr.shape[0], dtype=torch.int32, device=q.device) + + if q.dtype == FP8_DTYPE: + assert kv_indices is not None + num_kv_splits, num_kv_splits_indptr = get_meta_param( + None, + qo_indptr.shape[0] - 1, + kv_indices.numel(), + q.shape[1], + max_seqlen_qo, + q.dtype, + ) + if num_kv_splits == 1: + # gfx942's one-split FP8 asm writes the final output directly but + # does not write either LSE buffer. Force the normal split reducer, + # which produces both the same output and an accurate natural LSE. + num_kv_splits = 2 + num_kv_splits_indptr = torch.arange( + 0, + (qo_indptr.shape[0]) * num_kv_splits, + num_kv_splits, + dtype=torch.int32, + device=q.device, + ) + kwargs["num_kv_splits"] = num_kv_splits + kwargs["num_kv_splits_indptr"] = num_kv_splits_indptr + + _, final_lse = mla_decode_fwd( + q, + kv_buffer.view(-1, 1, 1, q.shape[-1]), + o, + qo_indptr, + kv_indptr, + kv_indices, + kv_last_page_lens, + max_seqlen_qo, + **kwargs, + ) + if final_lse is None: + raise RuntimeError("AITER MLA decode did not return the requested LSE") + return final_lse + + +def _rocm_aiter_mla_decode_fwd_lse_fake( + q: torch.Tensor, + kv_buffer: torch.Tensor, + o: torch.Tensor, + qo_indptr: torch.Tensor, + max_seqlen_qo: int, + kv_indptr: torch.Tensor | None = None, + kv_indices: torch.Tensor | None = None, + kv_last_page_lens: torch.Tensor | None = None, + sm_scale: float = 1.0, + logit_cap: float = 0.0, + q_scale: torch.Tensor | None = None, + kv_scale: torch.Tensor | None = None, +) -> torch.Tensor: + return torch.empty(q.shape[:-1], dtype=torch.float32, device=q.device) + + def _rocm_aiter_w8a8_gemm_impl( A: torch.Tensor, B: torch.Tensor, @@ -2117,6 +2215,13 @@ def register_ops_once() -> None: mutates_args=["o"], ) + direct_register_custom_op( + op_name="rocm_aiter_mla_decode_fwd_lse", + op_func=_rocm_aiter_mla_decode_fwd_lse_impl, + mutates_args=["o"], + fake_impl=_rocm_aiter_mla_decode_fwd_lse_fake, + ) + direct_register_custom_op( op_name="rocm_aiter_w8a8_gemm", op_func=_rocm_aiter_w8a8_gemm_impl, @@ -2693,6 +2798,37 @@ def mla_decode_fwd( reduce_partial_map=reduce_partial_map, ) + @staticmethod + def mla_decode_fwd_lse( + q: torch.Tensor, + kv_buffer: torch.Tensor, + o: torch.Tensor, + sm_scale: float, + qo_indptr: torch.Tensor, + max_seqlen_qo: int, + kv_indptr: torch.Tensor | None = None, + kv_indices: torch.Tensor | None = None, + kv_last_page_lens: torch.Tensor | None = None, + logit_cap: float = 0.0, + q_scale: torch.Tensor | None = None, + kv_scale: torch.Tensor | None = None, + ) -> torch.Tensor: + """Run MLA decode without persistent metadata and return its LSE.""" + return torch.ops.aphrodite.rocm_aiter_mla_decode_fwd_lse( + q, + kv_buffer.view(-1, 1, 1, q.shape[-1]), + o, + qo_indptr, + max_seqlen_qo, + kv_indptr, + kv_indices, + kv_last_page_lens, + sm_scale=sm_scale, + logit_cap=logit_cap, + q_scale=q_scale, + kv_scale=kv_scale, + ) + @staticmethod def per_tensor_quant( x: torch.Tensor, diff --git a/aphrodite/v1/attention/backend.py b/aphrodite/v1/attention/backend.py index 647709c1ed..7a2d76b26f 100644 --- a/aphrodite/v1/attention/backend.py +++ b/aphrodite/v1/attention/backend.py @@ -224,6 +224,13 @@ def supports_pcp(cls) -> bool: except NotImplementedError: return False + @classmethod + def supports_dcp(cls) -> bool: + try: + return getattr(cls.get_impl_cls(), "supports_dcp", True) + except NotImplementedError: + return False + @classmethod def supports_non_causal_dcp(cls) -> bool: builder_cls = cls.get_builder_cls() @@ -320,6 +327,8 @@ def validate_configuration( invalid_reasons.append("KV connector not supported") if use_pcp and not cls.supports_pcp(): invalid_reasons.append("PCP not supported") + if use_dcp and not cls.supports_dcp(): + invalid_reasons.append("DCP not supported") if use_adaptive_verification and not cls.supports_device_cpu_query_lens_mismatch(): invalid_reasons.append( "device-cpu query lens mismatch not supported, this is needed for adaptive verification" diff --git a/aphrodite/v1/attention/backends/mla/rocm_aiter_mla_sparse.py b/aphrodite/v1/attention/backends/mla/rocm_aiter_mla_sparse.py index c9d8251aeb..4669be9058 100644 --- a/aphrodite/v1/attention/backends/mla/rocm_aiter_mla_sparse.py +++ b/aphrodite/v1/attention/backends/mla/rocm_aiter_mla_sparse.py @@ -35,7 +35,7 @@ from aphrodite.v1.attention.ops.rocm_aiter_mla_sparse import ( rocm_sparse_attn_prefill, ) -from aphrodite.v1.kv_cache_interface import AttentionSpec +from aphrodite.v1.kv_cache_interface import AttentionSpec, KVCacheLayout from aphrodite.v1.worker.workspace import current_workspace_manager if TYPE_CHECKING: @@ -336,6 +336,17 @@ def is_mla(cls) -> bool: def is_sparse(cls) -> bool: return True + @classmethod + def supported_kv_cache_layouts(cls) -> tuple[KVCacheLayout, ...]: + # Global index conversion assumes contiguous pages within each layer. + return (KVCacheLayout.LBNHC, KVCacheLayout.LBHNC) + + @classmethod + def supports_sink(cls) -> bool: + from aphrodite.platforms.rocm import on_mi3xx + + return on_mi3xx() + @dataclass class ROCMAiterMLASparseMetadata(AttentionMetadata): @@ -402,6 +413,13 @@ def __init__( self.num_heads = self.model_config.get_num_attention_heads(parallel_config) self.mla_dims = get_mla_dims(self.model_config) self.topk_tokens = aphrodite_config.model_config.hf_config.index_topk + attention_context = aphrodite_config.compilation_config.static_forward_context + # Sink decode must use AITER's nonpersistent path. In particular, + # gfx942 has no persistent+LSE kernel, and its metadata heuristic + # terminates for HY-V4's TP1 H64 shape. + self._use_persistent_metadata = all( + getattr(attention_context[name].impl, "sinks", None) is None for name in layer_names + ) # Bounds the KV-split heuristic (see `_sparse_decode_max_split`). self._num_compute_units = current_platform.num_compute_units() self.max_model_len_tensor = torch.tensor([self.model_config.max_model_len], device=device, dtype=torch.int32) @@ -540,14 +558,13 @@ def build( paged_kv_indices = self.paged_kv_indices[: num_tokens * self.topk_tokens] # ----- Compute persistent MLA metadata ----- - # The aiter sparse decode kernel uses qseqlen=1 (each query token is - # treated as its own batch entry), so persistent metadata can always - # be precomputed here. The kernel switches to the persistent - # work-stealing path automatically when work_meta_data is non-None. - # The output is a deterministic function of the per-request query and - # 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. + # The AITER sparse decode kernel uses qseqlen=1 (each query token is + # treated as its own batch entry). Build its persistent work metadata + # only when AITER is selected and no layer needs the nonpersistent LSE + # path for attention sinks. The output is a deterministic function of + # the per-request query and 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. 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, @@ -564,7 +581,7 @@ def build( reduce_indptr = None reduce_final_map = None reduce_partial_map = None - if not use_triton_sparse: + if self._use_persistent_metadata and 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(), @@ -676,6 +693,7 @@ def log2sumexp2(a: torch.Tensor, dim: int) -> torch.Tensor: class ROCMAiterMLASparseImpl(MLAAttentionImpl[ROCMAiterMLASparseMetadata]): is_sparse = True supports_dense_mha_prefill = False + supports_dcp = False def __init__( self, @@ -700,6 +718,15 @@ def __init__( self.head_size = head_size self.scale = float(scale) self.num_kv_heads = num_kv_heads + sinks = mla_args.pop("sinks", None) + if sinks is not None: + if sinks.dtype != torch.float32: + raise ValueError(f"ROCm AITER MLA sinks must be float32, got {sinks.dtype}") + if sinks.ndim != 1 or sinks.numel() != num_heads: + raise ValueError(f"ROCm AITER MLA sinks must have shape ({num_heads},), got {tuple(sinks.shape)}") + if not sinks.is_contiguous(): + raise ValueError("ROCm AITER MLA sinks must be contiguous") + self.sinks: torch.Tensor | None = sinks self.kv_cache_dtype = kv_cache_dtype self.kv_lora_rank: int = mla_args["kv_lora_rank"] self.softmax_scale = scale @@ -723,16 +750,25 @@ def _forward_mla( q: torch.Tensor, # [sq, heads, d_qk] kv_c_and_k_pe_cache: torch.Tensor, # [blocks, heads, d_qk] attn_metadata: ROCMAiterMLASparseMetadata, - ) -> torch.Tensor: + ) -> tuple[torch.Tensor, torch.Tensor | None]: num_tokens = q.shape[0] - mla_num_heads = AiterMLAHelper.get_actual_mla_num_heads(self.num_heads) - output = torch.empty( - [num_tokens, mla_num_heads, self.kv_lora_rank], - dtype=attn_metadata.attn_out_dtype, - device=q.device, + base_mla_num_heads = AiterMLAHelper.get_actual_mla_num_heads(self.num_heads) + mla_num_heads = base_mla_num_heads + need_lse = self.sinks is not None + from aphrodite.platforms.rocm import on_gfx942 + + # Keep sink attention available for dtypes/head shapes without an + # AITER return-LSE kernel. Sink layers never need persistent metadata. + triton_sink_fallback = ( + need_lse + and q.dtype == kv_c_and_k_pe_cache.dtype + and ( + q.dtype == torch.float16 + or (q.dtype == torch.bfloat16 and (mla_num_heads > 128 or (on_gfx942() and 32 < mla_num_heads <= 64))) + ) ) - if _use_rocm_sparse_triton( + if triton_sink_fallback or _use_rocm_sparse_triton( kv_cache_dtype=self.kv_cache_dtype, head_size=q.shape[-1], kv_lora_rank=self.kv_lora_rank, @@ -741,6 +777,18 @@ def _forward_mla( num_decode_tokens=attn_metadata.num_decode_tokens, max_query_len=attn_metadata.max_query_len, ): + output = torch.empty( + [num_tokens, q.shape[1], self.kv_lora_rank], + dtype=attn_metadata.attn_out_dtype, + device=q.device, + ) + triton_sinks = None + if self.sinks is not None: + triton_sinks = AiterMLAHelper.get_mla_padded_q( + self.num_heads, + self.sinks.reshape(1, self.num_heads, 1), + q.shape[1], + ).reshape(-1) rocm_sparse_attn_prefill( q=q, kv=kv_c_and_k_pe_cache.view(-1, 1, q.shape[-1]), @@ -750,44 +798,133 @@ def _forward_mla( 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, + attn_sink=triton_sinks, 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. - mla_kwargs: dict = dict( - q_scale=layer._q_scale, - kv_scale=layer._k_scale, - ) - if attn_metadata.work_meta_data is not None: - mla_kwargs.update( - work_meta_data=attn_metadata.work_meta_data, - work_indptr=attn_metadata.work_indptr, - work_info_set=attn_metadata.work_info_set, - reduce_indptr=attn_metadata.reduce_indptr, - reduce_final_map=attn_metadata.reduce_final_map, - reduce_partial_map=attn_metadata.reduce_partial_map, - ) + output = AiterMLAHelper.get_mla_unpadded_o(self.num_heads, output) + return output, None + + # AITER's nonpersistent return-LSE dispatch has discrete head kernels. + supported_head_buckets: tuple[int, ...] | None = None + head_dtype_name = "" + if need_lse: + if q.dtype == torch.bfloat16 and kv_c_and_k_pe_cache.dtype == torch.bfloat16: + supported_head_buckets = (16, 32, 64, 128) + head_dtype_name = "BF16" + elif q.dtype == current_platform.fp8_dtype() and kv_c_and_k_pe_cache.dtype == current_platform.fp8_dtype(): + supported_head_buckets = (16, 128) + head_dtype_name = "FP8" + else: + raise ValueError( + "ROCm AITER MLA attention sinks require query and KV to " + "both use BF16 or both use FP8, got " + f"query={q.dtype}, KV={kv_c_and_k_pe_cache.dtype}" + ) - rocm_aiter_ops.mla_decode_fwd( - q, - kv_c_and_k_pe_cache, - output, - self.scale, - attn_metadata.qo_indptr, - 1, - attn_metadata.paged_kv_indptr, - attn_metadata.paged_kv_indices, - attn_metadata.paged_kv_last_page_len, - **mla_kwargs, + if supported_head_buckets is not None: + from aphrodite.platforms.rocm import on_mi3xx + + if on_mi3xx(): + supported_heads = next( + (heads for heads in supported_head_buckets if heads >= mla_num_heads), + None, + ) + if supported_heads is None: + raise ValueError( + "ROCm AITER MLA attention sinks support at most 128 " + f"padded local {head_dtype_name} heads; increase " + "tensor_parallel_size" + ) + if supported_heads != mla_num_heads: + q = AiterMLAHelper.get_mla_padded_q(mla_num_heads, q, supported_heads) + mla_num_heads = supported_heads + output = torch.empty( + [num_tokens, mla_num_heads, self.kv_lora_rank], + dtype=attn_metadata.attn_out_dtype, + device=q.device, ) - return AiterMLAHelper.get_mla_unpadded_o(self.num_heads, output) + if need_lse: + # gfx942 has no persistent MLA code object that writes final LSE. + # The split-KV path consumes the same ragged indices and is exact. + lse = rocm_aiter_ops.mla_decode_fwd_lse( + q, + kv_c_and_k_pe_cache, + output, + self.scale, + attn_metadata.qo_indptr, + 1, + attn_metadata.paged_kv_indptr, + attn_metadata.paged_kv_indices, + attn_metadata.paged_kv_last_page_len, + q_scale=layer._q_scale, + kv_scale=layer._k_scale, + ) + else: + # Preserve the persistent work-stealing fast path for models that + # do not use sinks. + mla_kwargs: dict = dict( + q_scale=layer._q_scale, + kv_scale=layer._k_scale, + ) + if attn_metadata.work_meta_data is not None: + mla_kwargs.update( + work_meta_data=attn_metadata.work_meta_data, + work_indptr=attn_metadata.work_indptr, + work_info_set=attn_metadata.work_info_set, + reduce_indptr=attn_metadata.reduce_indptr, + reduce_final_map=attn_metadata.reduce_final_map, + reduce_partial_map=attn_metadata.reduce_partial_map, + ) + + rocm_aiter_ops.mla_decode_fwd( + q, + kv_c_and_k_pe_cache, + output, + self.scale, + attn_metadata.qo_indptr, + 1, + attn_metadata.paged_kv_indptr, + attn_metadata.paged_kv_indices, + attn_metadata.paged_kv_last_page_len, + **mla_kwargs, + ) + lse = None + + if mla_num_heads != base_mla_num_heads: + if mla_num_heads % base_mla_num_heads == 0: + head_stride = mla_num_heads // base_mla_num_heads + output = output[:, ::head_stride] + if lse is not None: + lse = lse[:, ::head_stride] + else: + output = output[:, :base_mla_num_heads] + if lse is not None: + lse = lse[:, :base_mla_num_heads] + + output = AiterMLAHelper.get_mla_unpadded_o(self.num_heads, output) + if lse is not None: + lse = AiterMLAHelper.get_mla_unpadded_o(self.num_heads, lse.unsqueeze(-1)).squeeze(-1) + + if self.sinks is not None: + assert lse is not None + # Empty ragged rows have only sink mass and no value contribution. + # AITER can return NaN output/LSE for those rows; do not multiply it + # by a zero normalization factor and propagate the NaN. + has_keys = (attn_metadata.paged_kv_indptr[1:] > attn_metadata.paged_kv_indptr[:-1]).unsqueeze(-1) + lse = torch.where(has_keys, lse, float("-inf")) + sink_lse = torch.logaddexp(lse, self.sinks) + sink_scale = torch.exp(lse - sink_lse) + output = torch.where( + has_keys.unsqueeze(-1), + output.float() * sink_scale.unsqueeze(-1), + 0.0, + ).to(output.dtype) + lse = sink_lse + + return output, lse def forward_mqa( self, @@ -810,6 +947,8 @@ def forward_mqa( q = self.q_concat_buffer[: ql_nope.shape[0]] if q_pe.shape[-1] == 0: q.copy_(ql_nope) + elif q.dtype == torch.float16: + torch.cat((ql_nope, q_pe), dim=-1, out=q) else: ops.concat_mla_q(ql_nope, q_pe, q) @@ -837,5 +976,4 @@ def forward_mqa( q, _ = ops.scaled_fp8_quant(q.view(q.shape[0], -1), layer._q_scale) q = q.view(original_q_shape) mla_padded_q = AiterMLAHelper.get_mla_padded_q(self.num_heads, q) - attn_out = self._forward_mla(layer, mla_padded_q, kv_c_and_k_pe_cache, attn_metadata) - return attn_out, None + return self._forward_mla(layer, mla_padded_q, kv_c_and_k_pe_cache, attn_metadata) diff --git a/aphrodite/v1/attention/ops/rocm_aiter_mla_sparse.py b/aphrodite/v1/attention/ops/rocm_aiter_mla_sparse.py index ed6cfd1b57..542b1279c5 100644 --- a/aphrodite/v1/attention/ops/rocm_aiter_mla_sparse.py +++ b/aphrodite/v1/attention/ops/rocm_aiter_mla_sparse.py @@ -2484,10 +2484,10 @@ def _rocm_sparse_attn_prefill_ragged_triton( block_k = 16 if head_dim >= 256 else 32 num_warps = 4 if out is None: - out = torch.empty_like(q, dtype=torch.bfloat16) + out = torch.empty_like(q) else: assert out.shape == q.shape, f"expected out shape {q.shape}, got {out.shape}" - assert out.dtype == torch.bfloat16, f"expected bf16 out, got {out.dtype}" + assert out.dtype == q.dtype, f"expected {q.dtype} out, got {out.dtype}" assert out.device == q.device, f"expected out on {q.device}, got {out.device}" _sparse_attn_prefill_ragged_kernel[(num_queries, triton.cdiv(num_heads, block_h))]( q, @@ -3017,7 +3017,7 @@ def rocm_sparse_attn_prefill( rope_head_dim=rope_head_dim, topk_length=topk_length, ) - output.copy_(output_chunk.to(output.dtype)) + output.copy_(output_chunk[..., : output.shape[-1]].to(output.dtype)) def rocm_sparse_attn_decode( diff --git a/tests/kernels/attention/test_rocm_aiter_mla_op_registration.py b/tests/kernels/attention/test_rocm_aiter_mla_op_registration.py index 3f911936d8..4ebfd7f154 100644 --- a/tests/kernels/attention/test_rocm_aiter_mla_op_registration.py +++ b/tests/kernels/attention/test_rocm_aiter_mla_op_registration.py @@ -1,11 +1,11 @@ # SPDX-License-Identifier: Apache-2.0 # SPDX-FileCopyrightText: Copyright contributors to the vLLM project -"""ROCm custom op schema test for AITER MLA decode. +"""ROCm custom op schema tests for AITER MLA decode. -A single ``opcheck`` call on ``torch.ops.aphrodite.rocm_aiter_mla_decode_fwd`` -verifies that the custom op is registered and that its schema and fake -implementation are consistent with the real kernel: fake-tensor support for -torch.compile tracing and the ``mutates_args=["o"]`` in-place output aliasing. +``opcheck`` verifies that the decode ops are registered and that their schemas +and fake implementations are consistent with the real kernels: fake-tensor +support for torch.compile tracing and ``mutates_args=["o"]`` in-place output +aliasing. """ import pytest @@ -14,16 +14,7 @@ from aphrodite.platforms import current_platform from tests.kernels.utils import opcheck -_SKIP_NON_MI3XX = True -if current_platform.is_rocm(): - from aphrodite.platforms.rocm import on_mi3xx - - _SKIP_NON_MI3XX = not on_mi3xx() - -pytestmark = [ - pytest.mark.skipif(not current_platform.is_rocm(), reason="ROCm-specific tests"), - pytest.mark.skipif(_SKIP_NON_MI3XX, reason="MI300/MI350 ROCm only"), -] +pytestmark = pytest.mark.skipif(not current_platform.is_rocm(), reason="ROCm-specific tests") Q_HEAD_DIM = 576 # kv_lora_rank + qk_rope_head_dim V_HEAD_DIM = 512 # kv_lora_rank @@ -31,6 +22,10 @@ def _require_aiter(): from aphrodite._aiter_ops import is_aiter_found_and_supported + from aphrodite.platforms.rocm import get_cdna_version + + if get_cdna_version() not in (3, 4): + pytest.skip("AITER MLA requires CDNA 3 or 4") if not is_aiter_found_and_supported(): pytest.skip("aiter is required on supported ROCm hardware for this test") @@ -77,3 +72,34 @@ def test_mla_decode_fwd_op_schema() -> None: "reduce_partial_map": None, }, ) + + +@torch.inference_mode() +def test_mla_decode_fwd_lse_op_schema() -> None: + """Validate graph registration and mutation schema for LSE decode.""" + _require_aiter() + # Import ensures the custom op is registered. + from aphrodite._aiter_ops import rocm_aiter_ops # noqa: F401 + + batch_size, nhead = 2, 16 + q = torch.randn(batch_size, nhead, Q_HEAD_DIM, dtype=torch.bfloat16, device="cuda") + kv_buffer = torch.randn(32, 1, 1, Q_HEAD_DIM, dtype=torch.bfloat16, device="cuda") + o = torch.zeros(batch_size, nhead, V_HEAD_DIM, dtype=torch.bfloat16, device="cuda") + qo_indptr = torch.arange(batch_size + 1, dtype=torch.int32, device="cuda") + kv_indptr = torch.arange(batch_size + 1, dtype=torch.int32, device="cuda") * 16 + kv_indices = torch.arange(32, dtype=torch.int32, device="cuda") + kv_last_page_lens = torch.ones(batch_size, dtype=torch.int32, device="cuda") + + opcheck( + torch.ops.aphrodite.rocm_aiter_mla_decode_fwd_lse, + (q, kv_buffer, o, qo_indptr, 1), + { + "kv_indptr": kv_indptr, + "kv_indices": kv_indices, + "kv_last_page_lens": kv_last_page_lens, + "sm_scale": Q_HEAD_DIM**-0.5, + "logit_cap": 0.0, + "q_scale": None, + "kv_scale": None, + }, + ) diff --git a/tests/kernels/attention/test_rocm_aiter_mla_sink.py b/tests/kernels/attention/test_rocm_aiter_mla_sink.py new file mode 100644 index 0000000000..a4229538b8 --- /dev/null +++ b/tests/kernels/attention/test_rocm_aiter_mla_sink.py @@ -0,0 +1,475 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +"""Correctness tests for ROCm AITER sparse MLA attention sinks.""" + +from types import SimpleNamespace + +import pytest +import torch + +from aphrodite.platforms import current_platform +from aphrodite.utils.torch_utils import set_random_seed + +pytestmark = pytest.mark.skipif(not current_platform.is_rocm(), reason="ROCm-specific tests") + +Q_HEAD_DIM = 576 +V_HEAD_DIM = 512 +SM_SCALE = Q_HEAD_DIM**-0.5 + + +def _gamma_fp32(operations: int) -> float: + u = torch.finfo(torch.float32).eps / 2 + return operations * u / (1 - operations * u) + + +def _sink_reference(q, keys, sinks, scale): + # Reference the stored operands, so input quantization is not kernel error. + q, keys, sinks = q.double(), keys.double(), sinks.double() + scores = q @ keys.T * scale + logits = torch.cat((scores, sinks[:, None]), dim=-1) + probabilities = logits.softmax(dim=-1)[:, :-1] + values = keys[:, :V_HEAD_DIM] + expected = probabilities @ values + value_scale = probabilities @ values.abs() + # FP32 dot accumulation and conversion/multiplication of the scale. + score_error = _gamma_fp32(q.shape[-1] + 2) * (q.abs() @ keys.abs().T) * abs(scale) + score_error = score_error.amax(dim=-1) if keys.shape[0] else torch.zeros_like(sinks) + return expected, logits.logsumexp(dim=-1), value_scale, score_error + + +def _assert_sink_output_close(actual, expected, value_scale, score_error, probability_dtype, native): + assert actual.shape == expected.shape + # Budget the identifiable rounding stages against the absolute value + # contributions: output-relative ULPs become unbounded when values cancel. + # Rounding P before P@V costs u(P) * sum(P*abs(V)). + # AITER rounds output before and after sink scaling; Triton rounds once. + u_probability = torch.finfo(probability_dtype).eps / 2 + u_output = torch.finfo(actual.dtype).eps / 2 + output_rounds = 2 if native else 1 + amplification = (1 + u_output) ** output_rounds + relative_error = (1 + u_probability) * torch.exp(2 * score_error) - 1 + allowance = ( + relative_error.unsqueeze(-1) * value_scale * amplification + + (amplification - 1) * expected.abs() + + output_rounds * torch.finfo(actual.dtype).tiny * u_output + ) + actual = actual.to(device=expected.device, dtype=torch.float64) + assert torch.isfinite(actual).all() + assert torch.count_nonzero(actual[value_scale == 0]) == 0 + error = (actual - expected).abs() + max_ratio = (error / allowance.clamp_min(1e-300)).max().item() + assert max_ratio <= 1, f"exceeded rounding budget by {max_ratio:.3f}x" + + +def _require_aiter() -> None: + from aphrodite._aiter_ops import is_aiter_found_and_supported + from aphrodite.platforms.rocm import get_cdna_version + + if get_cdna_version() not in (3, 4): + pytest.skip("AITER MLA requires CDNA 3 or 4") + + if not is_aiter_found_and_supported(): + pytest.skip("aiter is required on supported ROCm hardware for this test") + + +@pytest.mark.parametrize( + ("real_heads", "cache_kind"), + [ + (1, "bf16"), + (4, "bf16"), + (8, "bf16"), + (8, "fp8"), + (32, "fp8"), + (64, "fp8"), + (12, "bf16"), + (16, "bf16"), + (24, "bf16"), + (32, "bf16"), + (40, "bf16"), + (48, "bf16"), + (64, "bf16"), + (80, "bf16"), + ], +) +@torch.inference_mode() +def test_sparse_mla_sink_matches_ragged_reference(real_heads: int, cache_kind: str) -> None: + """Exercise ragged sink decode, head padding, and FP8 scale forwarding.""" + _require_aiter() + from aphrodite.v1.attention.backends.mla.rocm_aiter_mla import AiterMLAHelper + from aphrodite.v1.attention.backends.mla.rocm_aiter_mla_sparse import ( + ROCMAiterMLASparseImpl, + ) + + set_random_seed(real_heads * 17 + (cache_kind == "fp8")) + device = torch.device("cuda") + seq_lens = [0, 1, 5, 23, 17] + batch_size = len(seq_lens) + + q_source = torch.randn(batch_size, real_heads, Q_HEAD_DIM, device=device) * 0.2 + + pool_size = sum(seq_lens) + 52 + kv_source = torch.randn(pool_size, 1, Q_HEAD_DIM, device=device) * 0.2 + if cache_kind == "fp8": + fp8_dtype = current_platform.fp8_dtype() + q_scale = torch.tensor(0.5, dtype=torch.float32, device=device) + kv_scale = torch.tensor(0.25, dtype=torch.float32, device=device) + q_real = (q_source / q_scale).to(fp8_dtype) + kv = (kv_source / kv_scale).to(fp8_dtype) + q_ref = q_real.float() * q_scale + kv_ref = kv.float() * kv_scale + else: + q_scale = kv_scale = None + q_real = q_source.to(torch.bfloat16) + kv = kv_source.to(torch.bfloat16) + q_ref = q_real.float() + kv_ref = kv.float() + q = AiterMLAHelper.get_mla_padded_q(real_heads, q_real) + + indices = torch.randperm(pool_size, device=device)[: sum(seq_lens)].to(torch.int32) + kv_indptr = torch.tensor( + [0] + [sum(seq_lens[:i]) for i in range(1, len(seq_lens) + 1)], + dtype=torch.int32, + device=device, + ) + + # Non-None garbage proves the sink path cannot accidentally select the + # gfx942 persistent kernel, which has no return-LSE code object. + metadata = SimpleNamespace( + attn_out_dtype=torch.bfloat16, + qo_indptr=torch.arange(batch_size + 1, dtype=torch.int32, device=device), + paged_kv_indptr=kv_indptr, + paged_kv_indices=indices, + paged_kv_last_page_len=torch.ones(batch_size, dtype=torch.int32, device=device), + work_meta_data=torch.tensor([123], dtype=torch.int32), + work_indptr=None, + work_info_set=None, + reduce_indptr=None, + reduce_final_map=None, + reduce_partial_map=None, + num_prefills=0, + num_decodes=batch_size, + num_decode_tokens=batch_size, + max_query_len=1, + ) + sinks = torch.linspace(-2.0, 6.0, real_heads, device=device) + + impl = object.__new__(ROCMAiterMLASparseImpl) + impl.num_heads = real_heads + impl.kv_lora_rank = V_HEAD_DIM + impl.kv_cache_dtype = cache_kind + impl.scale = SM_SCALE + impl.sinks = sinks + + output, lse = impl._forward_mla(SimpleNamespace(_q_scale=q_scale, _k_scale=kv_scale), q, kv, metadata) + kv_flat = kv_ref[:, 0] + references = [] + start = 0 + for batch_idx, seq_len in enumerate(seq_lens): + rows = kv_flat[indices[start : start + seq_len].long()] + references.append(_sink_reference(q_ref[batch_idx], rows, sinks, SM_SCALE)) + start += seq_len + expected, expected_lse, value_scale, score_error = (torch.stack(values) for values in zip(*references)) + + assert output.dtype == torch.bfloat16 + if lse is not None: + assert lse.dtype == torch.float32 + assert lse.shape == expected_lse.shape + lse_error = (lse.double() - expected_lse).abs() + rounded_lse = expected_lse.float() + positive_inf = torch.full_like(rounded_lse, float("inf")) + ulp = torch.maximum( + torch.nextafter(rounded_lse, positive_inf) - rounded_lse, + rounded_lse - torch.nextafter(rounded_lse, -positive_inf), + ).double() + lse_allowance = score_error + 2 * ulp + assert torch.all(lse_error <= lse_allowance) + torch.testing.assert_close(lse[0].double(), expected_lse[0], atol=0, rtol=0) + else: + from aphrodite.platforms.rocm import on_gfx942 + + assert cache_kind == "bf16" and real_heads in (40, 48, 64) and on_gfx942() + _assert_sink_output_close( + output, + expected, + value_scale, + score_error, + q_real.dtype, + native=lse is not None, + ) + + +@pytest.mark.parametrize( + ("q_dtype", "kv_dtype"), + [ + (torch.float16, torch.bfloat16), + (torch.bfloat16, torch.float16), + ], +) +def test_sparse_mla_sink_rejects_unsupported_aiter_dtypes(q_dtype: torch.dtype, kv_dtype: torch.dtype) -> None: + from aphrodite.v1.attention.backends.mla.rocm_aiter_mla_sparse import ( + ROCMAiterMLASparseImpl, + ) + + impl = object.__new__(ROCMAiterMLASparseImpl) + impl.num_heads = 16 + impl.kv_lora_rank = V_HEAD_DIM + impl.kv_cache_dtype = "auto" + impl.scale = SM_SCALE + impl.sinks = torch.zeros(16, dtype=torch.float32) + metadata = SimpleNamespace( + attn_out_dtype=torch.bfloat16, + num_prefills=0, + num_decodes=1, + num_decode_tokens=1, + max_query_len=1, + ) + + with pytest.raises(ValueError, match="both use BF16 or both use FP8"): + impl._forward_mla( + SimpleNamespace(_q_scale=None, _k_scale=None), + torch.empty(1, 16, Q_HEAD_DIM, dtype=q_dtype), + torch.empty(1, 1, Q_HEAD_DIM, dtype=kv_dtype), + metadata, + ) + + +def _make_noncontiguous_sink() -> torch.Tensor: + return torch.empty(8, dtype=torch.float32)[::2] + + +@pytest.mark.parametrize( + ("sinks", "match"), + [ + (torch.empty(4, dtype=torch.bfloat16), "must be float32"), + (torch.empty(2, 2, dtype=torch.float32), "must have shape"), + (_make_noncontiguous_sink(), "must be contiguous"), + ], +) +def test_sparse_mla_sink_validation(sinks: torch.Tensor, match: str) -> None: + from aphrodite.v1.attention.backends.mla.rocm_aiter_mla_sparse import ( + ROCMAiterMLASparseImpl, + ) + + with pytest.raises(ValueError, match=match): + ROCMAiterMLASparseImpl( + num_heads=4, + head_size=Q_HEAD_DIM, + scale=SM_SCALE, + num_kv_heads=1, + alibi_slopes=None, + sliding_window=None, + kv_cache_dtype="bfloat16", + logits_soft_cap=None, + attn_type="decoder", + kv_sharing_target_layer_name=None, + sinks=sinks, + kv_lora_rank=V_HEAD_DIM, + ) + + +def test_sparse_mla_backend_reports_sink_support_for_current_hardware() -> None: + from aphrodite.platforms.rocm import get_cdna_version + from aphrodite.v1.attention.backends.mla.rocm_aiter_mla_sparse import ( + ROCMAiterMLASparseBackend, + ) + + assert ROCMAiterMLASparseBackend.supports_sink() == (get_cdna_version() in (3, 4)) + + +def test_sparse_mla_backend_rejects_dcp() -> None: + from aphrodite.platforms.rocm import RocmPlatform + from aphrodite.v1.attention.backends.registry import AttentionBackendEnum + from aphrodite.v1.attention.selector import AttentionSelectorConfig + + selector_config = AttentionSelectorConfig( + head_size=Q_HEAD_DIM, + dtype=torch.bfloat16, + kv_cache_dtype="bfloat16", + block_size=16, + use_mla=True, + has_sink=True, + use_sparse=True, + use_mm_prefix=False, + use_per_head_quant_scales=False, + attn_type="decoder", + use_dcp=True, + ) + + with pytest.raises(ValueError, match="DCP not supported"): + RocmPlatform.get_attn_backend_cls( + selected_backend=AttentionBackendEnum.ROCM_AITER_MLA_SPARSE, + attn_selector_config=selector_config, + ) + + +@pytest.mark.parametrize("num_heads,block_size", [(8, 16), (20, 32)]) +@pytest.mark.parametrize("interleaved_pages", [False, True]) +@pytest.mark.parametrize("dtype", [torch.bfloat16, torch.float16]) +@torch.inference_mode() +def test_sparse_mla_sink_matches_dense_attention_with_empty_rows_and_paged_cache( + num_heads: int, + block_size: int, + interleaved_pages: bool, + dtype: torch.dtype, +) -> None: + """Preserve head-specific sinks and latent values across ragged cache pages.""" + _require_aiter() + from aphrodite.v1.attention.backends.mla.rocm_aiter_mla import AiterMLAHelper + from aphrodite.v1.attention.backends.mla.rocm_aiter_mla_sparse import ( + ROCMAiterMLASparseImpl, + ) + from aphrodite.v1.attention.backends.mla.sparse_utils import flat_kv_row_view + + head_dim, value_dim = 576, 512 + generator = torch.Generator().manual_seed(421) + # The empty row exercises sink-only attention; the last row crosses both + # page boundaries and the kernel's 16-key reduction tile. + selected_rows = [ + [], + [(2, 7)], + [(3, 1), (0, block_size - 1), (1, 0), (2, 4)], + [(i % 4, (i * 7) % block_size) for i in range(23)], + ] + q_cpu = torch.randn(len(selected_rows), num_heads, head_dim, generator=generator).to(dtype) + cache_cpu = torch.randn(4, block_size, head_dim, generator=generator).to(dtype) + sinks_cpu = torch.linspace(-4.0, 6.0, num_heads) + + # Layers can share a backing allocation, leaving unused rows between pages. + num_layers = 3 if interleaved_pages else 1 + backing = torch.full( + (4, num_layers, block_size, head_dim), + float("nan"), + dtype=dtype, + device="cuda", + ) + cache = backing[:, num_layers - 1] + cache.copy_(cache_cpu) + indices = [page * num_layers * block_size + offset for rows in selected_rows for page, offset in rows] + lengths = torch.tensor([0, *(len(rows) for rows in selected_rows)]) + metadata = SimpleNamespace( + block_size=block_size, + num_prefills=0, + num_decodes=len(selected_rows), + num_decode_tokens=len(selected_rows), + max_query_len=1, + qo_indptr=torch.arange(len(selected_rows) + 1, dtype=torch.int32, device="cuda"), + paged_kv_last_page_len=torch.ones(len(selected_rows), dtype=torch.int32, device="cuda"), + work_meta_data=None, + attn_out_dtype=dtype, + paged_kv_indices=torch.tensor(indices, dtype=torch.int32, device="cuda"), + paged_kv_indptr=lengths.cumsum(0).to(device="cuda", dtype=torch.int32), + ) + impl = ROCMAiterMLASparseImpl.__new__(ROCMAiterMLASparseImpl) + impl.num_heads = num_heads + impl.head_size = head_dim + impl.kv_lora_rank = value_dim + impl.scale = head_dim**-0.5 + impl.kv_cache_dtype = "auto" + impl.sinks = sinks_cpu.cuda() + padded_q = AiterMLAHelper.get_mla_padded_q(num_heads, q_cpu.cuda()) + + kv_rows, _ = flat_kv_row_view(cache, block_size) + actual, _ = impl._forward_mla( + SimpleNamespace(_q_scale=None, _k_scale=None), + padded_q, + kv_rows.unsqueeze(1), + metadata, + ) + + references = [] + for query_idx, rows in enumerate(selected_rows): + keys = ( + torch.stack([cache_cpu[page, offset] for page, offset in rows]) + if rows + else cache_cpu.new_empty((0, head_dim)) + ) + references.append(_sink_reference(q_cpu[query_idx], keys, sinks_cpu, impl.scale)) + expected, _, value_scale, score_error = (torch.stack(values) for values in zip(*references)) + assert actual.shape == (len(selected_rows), num_heads, value_dim) + assert actual.dtype == dtype + _assert_sink_output_close( + actual, + expected, + value_scale, + score_error, + dtype, + native=dtype == torch.bfloat16, + ) + + +@pytest.mark.parametrize("layout", ["LBNHC", "LBHNC", "BLNHC"]) +def test_sparse_mla_backend_resolves_only_contiguous_layer_layouts(monkeypatch, layout): + from aphrodite.config import CacheConfig + from aphrodite.v1.attention.backends.mla.rocm_aiter_mla_sparse import ( + ROCMAiterMLASparseBackend, + ) + from aphrodite.v1.attention.backends.utils import resolve_kv_cache_layout + + monkeypatch.setenv("APHRODITE_KV_CACHE_LAYOUT", layout) + config = SimpleNamespace(cache_config=CacheConfig()) + supported = [[x.name for x in ROCMAiterMLASparseBackend.supported_kv_cache_layouts()]] + if layout == "BLNHC": + with pytest.raises(ValueError, match="does not satisfy"): + resolve_kv_cache_layout(config, supported) + else: + assert resolve_kv_cache_layout(config, supported).name == layout + + +@pytest.mark.parametrize("dtype", [torch.bfloat16, torch.float16]) +@torch.inference_mode() +def test_sparse_mla_sink_forward_mqa_preserves_split_query(dtype): + """The public sparse forward joins latent/RoPE queries in the model dtype.""" + _require_aiter() + from aphrodite.v1.attention.backends.mla.rocm_aiter_mla_sparse import ( + ROCMAiterMLASparseImpl, + ) + + set_random_seed(412) + num_tokens, num_heads, block_size = 2, 8, 16 + q = torch.randn(num_tokens, num_heads, Q_HEAD_DIM, device="cuda").to(dtype) + kv = torch.randn(2 * block_size, 1, Q_HEAD_DIM, device="cuda").to(dtype) + selected = torch.tensor([[1, 7, 18], [0, 3, 20]], dtype=torch.int32, device="cuda") + sinks = torch.linspace(-3.0, 5.0, num_heads, device="cuda") + impl = object.__new__(ROCMAiterMLASparseImpl) + impl.num_heads = num_heads + impl.kv_lora_rank = V_HEAD_DIM + impl.kv_cache_dtype = "auto" + impl.scale = SM_SCALE + impl.sinks = sinks + impl.topk_indices_buffer = torch.full((num_tokens, 128), -1, dtype=torch.int32, device="cuda") + impl.topk_indices_buffer[:, : selected.shape[1]] = selected + impl.q_concat_buffer = torch.empty_like(q) + metadata = SimpleNamespace( + attn_out_dtype=dtype, + num_actual_tokens=num_tokens, + num_prefills=0, + num_decodes=num_tokens, + num_decode_tokens=num_tokens, + max_query_len=1, + block_size=block_size, + topk_tokens=impl.topk_indices_buffer.shape[1], + req_id_per_token=torch.zeros(num_tokens, dtype=torch.int32, device="cuda"), + block_table=torch.tensor([[0, 1]], dtype=torch.int32, device="cuda"), + qo_indptr=torch.arange(num_tokens + 1, dtype=torch.int32, device="cuda"), + paged_kv_indptr=torch.tensor([0, 3, 6], dtype=torch.int32, device="cuda"), + paged_kv_indices=torch.empty(selected.numel(), dtype=torch.int32, device="cuda"), + paged_kv_last_page_len=torch.ones(num_tokens, dtype=torch.int32, device="cuda"), + work_meta_data=None, + ) + actual, _ = impl.forward_mqa( + (q[..., :V_HEAD_DIM], q[..., V_HEAD_DIM:]), + kv, + metadata, + SimpleNamespace(_q_scale=None, _k_scale=None), + ) + references = [_sink_reference(q[i], kv[:, 0][selected[i].long()], sinks, SM_SCALE) for i in range(num_tokens)] + expected, _, value_scale, score_error = (torch.stack(values) for values in zip(*references)) + assert actual.dtype == dtype + _assert_sink_output_close( + actual, + expected, + value_scale, + score_error, + dtype, + native=dtype == torch.bfloat16, + ) 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 3053fc3c1e..cf71f25386 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 @@ -48,6 +48,7 @@ def _make_builder(): builder.paged_kv_last_page_len = torch.ones(max_num_batched_tokens, dtype=torch.int32, device="cpu") builder.paged_kv_indices = torch.zeros(max_num_batched_tokens * topk_tokens, dtype=torch.int32, device="cpu") builder.paged_kv_indptr = torch.zeros(max_num_batched_tokens + 1, dtype=torch.int32, device="cpu") + builder._use_persistent_metadata = True builder._num_attention_heads = 16 builder._num_compute_units = current_platform.num_compute_units() builder._mla_work_meta_data = torch.empty(1, dtype=torch.int32, device="cpu") @@ -115,6 +116,7 @@ def fake_generate_sparse_seqlen_triton(query_lens, seq_lens, cu_query_lens, topk "current_stream", lambda device=None: SimpleNamespace(synchronize=lambda: events.append("sync") if events is not None else None), ) + return fake_aiter def test_build_populates_decode_only_split_fields(monkeypatch): @@ -146,6 +148,23 @@ def test_build_populates_mixed_split_fields(monkeypatch): assert md.prefill is None +def test_sink_build_skips_persistent_metadata(monkeypatch): + builder = _make_builder() + builder._use_persistent_metadata = False + fake_aiter = _patch_build_deps(monkeypatch) + + md = builder.build(common_prefix_len=0, common_attn_metadata=_make_common_metadata()) + + fake_aiter.get_mla_metadata_v1.assert_not_called() + assert md.work_meta_data is None + assert md.work_indptr is None + assert md.work_info_set is None + assert md.reduce_indptr is None + assert md.reduce_final_map is None + assert md.reduce_partial_map is None + assert builder._prev_metadata_key is None + + def test_sparse_persistent_metadata_syncs_only_after_recompute(monkeypatch): builder = _make_builder() common_metadata = _make_common_metadata() diff --git a/tests/v1/attention/test_rocm_glm5next_sparse.py b/tests/v1/attention/test_rocm_glm5next_sparse.py index e0aec7522a..3d99525e6a 100644 --- a/tests/v1/attention/test_rocm_glm5next_sparse.py +++ b/tests/v1/attention/test_rocm_glm5next_sparse.py @@ -1,11 +1,14 @@ # SPDX-License-Identifier: Apache-2.0 # SPDX-FileCopyrightText: Copyright contributors to the vLLM project +from types import SimpleNamespace + import pytest import torch from aphrodite.platforms import current_platform from aphrodite.triton_utils import tl, triton +from aphrodite.v1.attention.backends.mla import rocm_aiter_mla_sparse as sparse_mod from aphrodite.v1.attention.backends.mla.rocm_aiter_mla_sparse import ( _use_rocm_sparse_triton, fit_kpool_indices_to_aiter, @@ -97,6 +100,51 @@ def test_rocm_sparse_triton_route( ) +@pytest.mark.parametrize("num_heads", [8, 12]) +def test_rocm_sparse_triton_route_preserves_padded_sinks(monkeypatch, num_heads): + captured = {} + + def fake_rocm_sparse_attn_prefill(**kwargs): + output = kwargs["output"] + captured["attn_sink"] = kwargs["attn_sink"] + output.copy_(captured["attn_sink"].to(output.dtype).view(1, -1, 1).expand_as(output)) + + monkeypatch.setattr(sparse_mod, "rocm_sparse_attn_prefill", fake_rocm_sparse_attn_prefill) + + impl = object.__new__(sparse_mod.ROCMAiterMLASparseImpl) + impl.num_heads = num_heads + impl.kv_lora_rank = 512 + impl.kv_cache_dtype = "auto" + impl.scale = 512**-0.5 + impl.sinks = torch.arange(num_heads, dtype=torch.float32) + + q = torch.zeros(2, 16, 512, dtype=torch.bfloat16) + kv = torch.zeros(4, 1, 512, dtype=torch.bfloat16) + metadata = SimpleNamespace( + attn_out_dtype=torch.bfloat16, + num_prefills=1, + num_decodes=0, + num_decode_tokens=0, + max_query_len=2, + paged_kv_indices=torch.empty(0, dtype=torch.int32), + paged_kv_indptr=torch.zeros(3, dtype=torch.int32), + ) + + output, lse = impl._forward_mla(SimpleNamespace(), q, kv, metadata) + + if num_heads == 8: + expected_sinks = impl.sinks.repeat_interleave(2) + else: + expected_sinks = torch.cat((impl.sinks, impl.sinks[:4])) + torch.testing.assert_close(captured["attn_sink"], expected_sinks) + assert output.shape == (2, num_heads, 512) + torch.testing.assert_close( + output[:, :, 0].float(), + impl.sinks.expand(2, -1), + ) + assert lse is None + + def test_rocm_sparse_attention_accepts_glm_nope_dimensions(): _validate_sparse_dims(512, 512, 0, "test") From 34662884281dd9f1a9d0a8fc1fdb4a5a1ca6255b Mon Sep 17 00:00:00 2001 From: AlpinDale Date: Mon, 7 Sep 2026 20:19:05 +0000 Subject: [PATCH 20/36] [sync] [Frontend] Migrate Responses API validation errors to VLLMValidationError (#50257) Upstream-vLLM: 7dbe38683325126de96d442d6b494f4266001bec Co-authored-by: Adababy --- .sync/vllm-sha | 2 +- .../entrypoints/openai/parser/harmony_utils.py | 16 ++++++++++++---- aphrodite/entrypoints/openai/responses/utils.py | 6 +++++- .../openai/parser/test_harmony_utils.py | 6 ++++-- .../openai/responses/test_responses_utils.py | 9 ++++++--- 5 files changed, 28 insertions(+), 11 deletions(-) diff --git a/.sync/vllm-sha b/.sync/vllm-sha index c5363f21df..c6903f20b3 100644 --- a/.sync/vllm-sha +++ b/.sync/vllm-sha @@ -1 +1 @@ -e476556189506df612218c0888dc675e6ecbfc06 +7dbe38683325126de96d442d6b494f4266001bec diff --git a/aphrodite/entrypoints/openai/parser/harmony_utils.py b/aphrodite/entrypoints/openai/parser/harmony_utils.py index 9566c0107b..1b85dfc265 100644 --- a/aphrodite/entrypoints/openai/parser/harmony_utils.py +++ b/aphrodite/entrypoints/openai/parser/harmony_utils.py @@ -24,6 +24,7 @@ from aphrodite import envs from aphrodite.entrypoints.openai.chat_completion.protocol import ChatCompletionToolsParam +from aphrodite.exceptions import AphroditeValidationError from aphrodite.logger import init_logger logger = init_logger(__name__) @@ -121,9 +122,10 @@ def get_system_message( if reasoning_effort is not None: if reasoning_effort not in REASONING_EFFORT: supported_values = ", ".join(REASONING_EFFORT) - raise ValueError( + raise AphroditeValidationError( f"reasoning_effort={reasoning_effort!r} is not supported by " - f"Harmony. Supported values are: {supported_values}." + f"Harmony. Supported values are: {supported_values}.", + parameter="reasoning_effort", ) sys_msg_content = sys_msg_content.with_reasoning_effort(REASONING_EFFORT[reasoning_effort]) if start_date is None: @@ -175,7 +177,10 @@ def get_developer_message( elif tool.type == "function": function_tools.append(tool) else: - raise ValueError(f"tool type {tool.type} not supported") + raise AphroditeValidationError( + f"tool type {tool.type!r} is not supported.", + parameter="tools", + ) if function_tools: function_tool_descriptions = [create_tool_definition(tool) for tool in function_tools] dev_msg_content = dev_msg_content.with_function_tools(function_tool_descriptions) @@ -285,7 +290,10 @@ def extract_instructions_from_messages( elif hasattr(first_message, "model_dump"): first_message = first_message.model_dump(exclude_none=True) else: - raise ValueError(f"Unknown message type: {type(first_message)}") + raise AphroditeValidationError( + f"Unknown message type: {type(first_message)}", + parameter="input", + ) if first_message.get("role") not in ( "system", diff --git a/aphrodite/entrypoints/openai/responses/utils.py b/aphrodite/entrypoints/openai/responses/utils.py index c942a57702..46c7a44d5a 100644 --- a/aphrodite/entrypoints/openai/responses/utils.py +++ b/aphrodite/entrypoints/openai/responses/utils.py @@ -37,6 +37,7 @@ ChatCompletionToolsParam, ) from aphrodite.entrypoints.openai.responses.protocol import ResponseInputOutputItem +from aphrodite.exceptions import AphroditeValidationError from aphrodite.logger import init_logger from aphrodite.tool_parsers.utils import ( build_responses_tool_call_name_map, @@ -262,7 +263,10 @@ def _construct_message_from_response_item( elif isinstance(item, ResponseReasoningItem): reasoning = "" if item.encrypted_content: - raise ValueError("Encrypted content is not supported.") + raise AphroditeValidationError( + "Encrypted content is not supported.", + parameter="input", + ) elif item.content and len(item.content) >= 1: reasoning = item.content[0].text elif len(item.summary) >= 1: diff --git a/tests/entrypoints/openai/parser/test_harmony_utils.py b/tests/entrypoints/openai/parser/test_harmony_utils.py index f8af632434..c621ae6aa8 100644 --- a/tests/entrypoints/openai/parser/test_harmony_utils.py +++ b/tests/entrypoints/openai/parser/test_harmony_utils.py @@ -19,6 +19,7 @@ response_input_to_harmony, response_previous_input_to_harmony, ) +from aphrodite.exceptions import AphroditeValidationError from tests.entrypoints.openai.utils import verify_harmony_messages _TOOL_PARAMETERS = { @@ -929,10 +930,11 @@ def test_all_standard_channels_present(self) -> None: def test_unsupported_reasoning_effort_raises_clear_error(self) -> None: with pytest.raises( - ValueError, + AphroditeValidationError, match="reasoning_effort='max' is not supported by Harmony", - ): + ) as exc_info: get_system_message(reasoning_effort="max") + assert exc_info.value.parameter == "reasoning_effort" class TestResponseInputToHarmonyReasoningItem: diff --git a/tests/entrypoints/openai/responses/test_responses_utils.py b/tests/entrypoints/openai/responses/test_responses_utils.py index 2ba9250e1a..1221b61327 100644 --- a/tests/entrypoints/openai/responses/test_responses_utils.py +++ b/tests/entrypoints/openai/responses/test_responses_utils.py @@ -22,6 +22,7 @@ construct_input_messages, should_continue_final_message, ) +from aphrodite.exceptions import AphroditeValidationError def _single_chat_message(item): @@ -218,8 +219,9 @@ def test_construct_chat_messages_preserves_single_item_conversions(self): encrypted_content="TOP_SECRET_MESSAGE", status=None, ) - with pytest.raises(ValueError): + with pytest.raises(AphroditeValidationError) as exc_info: construct_chat_messages_with_tool_call([item]) + assert exc_info.value.parameter == "input" output_item = ResponseOutputMessage( id="msg_bf585bbbe3d500e0", @@ -341,7 +343,7 @@ def test_neither_content_nor_summary(self): assert formatted["reasoning"] == "" def test_encrypted_content_raises(self): - """Encrypted content should still raise ValueError.""" + """Encrypted content should raise AphroditeValidationError.""" item = ResponseReasoningItem( id="reasoning_6", summary=[ @@ -360,8 +362,9 @@ def test_encrypted_content_raises(self): encrypted_content="ENCRYPTED", status=None, ) - with pytest.raises(ValueError): + with pytest.raises(AphroditeValidationError) as exc_info: construct_chat_messages_with_tool_call([item]) + assert exc_info.value.parameter == "input" @patch("aphrodite.entrypoints.openai.responses.utils.logger") def test_summary_with_multiple_entries_uses_first(self, mock_logger): From 252358692787c7bd6d96ef4948ba495605397018 Mon Sep 17 00:00:00 2001 From: AlpinDale Date: Mon, 7 Sep 2026 20:22:13 +0000 Subject: [PATCH 21/36] [sync] [XPU] Route grouped_topk to the fused _moe_C kernel on XPU (#53580) Upstream-vLLM: 9b85112e130d8beffe9928ec0c86a4029d6ba7a8 Co-authored-by: Marceli Fylcek --- .sync/vllm-sha | 2 +- aphrodite/_custom_ops.py | 24 +++++++++++++++++-- .../fused_moe/router/grouped_topk_router.py | 2 +- 3 files changed, 24 insertions(+), 4 deletions(-) diff --git a/.sync/vllm-sha b/.sync/vllm-sha index c6903f20b3..46feb39f86 100644 --- a/.sync/vllm-sha +++ b/.sync/vllm-sha @@ -1 +1 @@ -7dbe38683325126de96d442d6b494f4266001bec +9b85112e130d8beffe9928ec0c86a4029d6ba7a8 diff --git a/aphrodite/_custom_ops.py b/aphrodite/_custom_ops.py index dda48ddf59..d7a177b45a 100644 --- a/aphrodite/_custom_ops.py +++ b/aphrodite/_custom_ops.py @@ -3052,8 +3052,8 @@ def grouped_topk( bias: Bias tensor (e_score_correction_bias). Always fused in kernel. scoring_func: 0=none (no activation), 1=sigmoid """ - if not current_platform.is_cuda(): - raise NotImplementedError("The fused grouped_topk kernel is only available on CUDA platforms") + if not (current_platform.is_cuda() or current_platform.is_xpu()): + raise NotImplementedError("The fused grouped_topk kernel is only available on CUDA and XPU platforms") return torch.ops._moe_C.grouped_topk( scores, num_expert_group, @@ -3066,6 +3066,26 @@ def grouped_topk( ) +if hasattr(torch.ops, "_moe_C") and hasattr(torch.ops._moe_C, "grouped_topk"): + + @register_fake("_moe_C::grouped_topk") + def _grouped_topk_fake( + scores: torch.Tensor, + num_expert_group: int, + topk_group: int, + topk: int, + renormalize: bool, + routed_scaling_factor: float, + bias: torch.Tensor, + scoring_func: int = 0, + ) -> tuple[torch.Tensor, torch.Tensor]: + num_tokens = scores.size(0) + return ( + scores.new_empty((num_tokens, topk), dtype=torch.float32), + scores.new_empty((num_tokens, topk), dtype=torch.int32), + ) + + def moe_wna16_marlin_gemm( input: torch.Tensor, output: torch.Tensor | None, diff --git a/aphrodite/model_executor/layers/fused_moe/router/grouped_topk_router.py b/aphrodite/model_executor/layers/fused_moe/router/grouped_topk_router.py index 867d4b9ae5..3560dfc9e2 100644 --- a/aphrodite/model_executor/layers/fused_moe/router/grouped_topk_router.py +++ b/aphrodite/model_executor/layers/fused_moe/router/grouped_topk_router.py @@ -90,7 +90,7 @@ def grouped_topk( ) -> tuple[torch.Tensor, torch.Tensor]: if ( envs.APHRODITE_USE_FUSED_MOE_GROUPED_TOPK - and current_platform.is_cuda() + and (current_platform.is_cuda() or current_platform.is_xpu()) and num_expert_group <= 32 and topk <= 32 and e_score_correction_bias is not None From 95935b49476f39dc316c3343d946d8e615a2742c Mon Sep 17 00:00:00 2001 From: AlpinDale Date: Mon, 7 Sep 2026 20:26:18 +0000 Subject: [PATCH 22/36] [sync] [XPU][LoRA] Support LoRA for DeepSeek V4 on XPU (#53689) Upstream-vLLM: 86484465153f08f4523d8be964fd6309146eab09 Co-authored-by: Chaojun Zhang --- .sync/vllm-sha | 2 +- aphrodite/models/deepseek_v4/xpu/model.py | 28 ++++++++++++++++++++++- 2 files changed, 28 insertions(+), 2 deletions(-) diff --git a/.sync/vllm-sha b/.sync/vllm-sha index 46feb39f86..225269d7f7 100644 --- a/.sync/vllm-sha +++ b/.sync/vllm-sha @@ -1 +1 @@ -9b85112e130d8beffe9928ec0c86a4029d6ba7a8 +86484465153f08f4523d8be964fd6309146eab09 diff --git a/aphrodite/models/deepseek_v4/xpu/model.py b/aphrodite/models/deepseek_v4/xpu/model.py index 041cfc8f5e..da95ce4fa6 100644 --- a/aphrodite/models/deepseek_v4/xpu/model.py +++ b/aphrodite/models/deepseek_v4/xpu/model.py @@ -46,6 +46,7 @@ from aphrodite.model_executor.models.interfaces import ( EagleModelMixin, SupportsEagle3, + SupportsLoRA, SupportsPP, ) from aphrodite.model_executor.models.utils import ( @@ -1136,6 +1137,11 @@ def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]: if is_pp_missing_parameter(name, self): break + if name not in params_dict: + head, _, leaf = name.rpartition(".") + suffixed = f"{head}.base_layer.{leaf}" + if suffixed in params_dict: + name = suffixed param = params_dict[name] weight_loader = param.weight_loader weight_loader(param, loaded_weight, shard_id) @@ -1185,6 +1191,13 @@ def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]: else: if is_pp_missing_parameter(name, self): continue + # Non-LoRA params on a LoRA-wrapped module live at + # ``.base_layer.``; the checkpoint is plain. + if name not in params_dict: + head, _, leaf = name.rpartition(".") + suffixed = f"{head}.base_layer.{leaf}" + if suffixed in params_dict: + name = suffixed param = params_dict[name] weight_loader = getattr(param, "weight_loader", default_weight_loader) weight_loader(param, loaded_weight) @@ -1219,6 +1232,8 @@ def _make_deepseek_v4_weights_mapper(expert_dtype: str) -> WeightsMapper: # shared experts use Fp8LinearMethod's block scales, which # register as ``weight_scale_inv``. scale_regex = { + # ``.base_layer.``-namespace variant (LoRA-wrapped experts). + re.compile(r"(\.experts\.\d+\.w[123]\.base_layer)\.scale$"): r"\1.weight_scale", re.compile(r"(\.experts\.\d+\.w[123])\.scale$"): r"\1.weight_scale", re.compile(r"\.scale$"): ".weight_scale_inv", } @@ -1227,6 +1242,8 @@ def _make_deepseek_v4_weights_mapper(expert_dtype: str) -> WeightsMapper: # scales as ``w{13,2}_weight_scale_inv``. Map all ``.scale`` keys # there. scale_regex = { + # ``.base_layer.``-namespace variant of the above. + re.compile(r"(\.experts\.\d+\.w[123]\.base_layer)\.scale$"): r"\1.weight_scale_inv", re.compile(r"\.scale$"): ".weight_scale_inv", } return WeightsMapper( @@ -1250,13 +1267,22 @@ def _make_deepseek_v4_weights_mapper(expert_dtype: str) -> WeightsMapper: ) -class DeepseekV4ForCausalLM(nn.Module, SupportsPP, SupportsEagle3): +class DeepseekV4ForCausalLM(nn.Module, SupportsPP, SupportsEagle3, SupportsLoRA): model_cls = DeepseekV4Model # Default mapper assumes the original FP4-expert checkpoint layout. # Overridden per-instance in __init__ when expert_dtype != "fp4". hf_to_aphrodite_mapper = _make_deepseek_v4_weights_mapper("fp4") + packed_modules_mapping = { + "gate_up_proj": ["w1", "w3"], + "fused_wqa_wkv": ["wq_a", "wkv"], + "fused_wkv_wgate": ["wkv", "wgate"], + } + + # The MTP draft head is not LoRA-adapted. + lora_skip_prefixes = ["mtp."] + def __init__(self, *, aphrodite_config: AphroditeConfig, prefix: str = ""): super().__init__() From 08c3f63297eacca49a11934d5737aaa491144dd2 Mon Sep 17 00:00:00 2001 From: AlpinDale Date: Mon, 7 Sep 2026 20:32:09 +0000 Subject: [PATCH 23/36] [sync] [Qwen3.8-Flash-Next] Support FP8 indexer cache for QSA (#54890) Upstream-vLLM: 94e26dd3dd7d6363b524c96521c78feb501bdc61 Co-authored-by: Thien Tran --- .sync/vllm-sha | 2 +- .../models/qwen4_exp/common/qsa_cache.py | 20 +++++-- .../models/qwen4_exp/nvidia/indexer_qsa.py | 17 +++++- .../qwen4_exp/nvidia/ops/qsa_indexer.py | 23 ++++++-- .../models/qwen4_exp/test_qsa_pre_indexer.py | 31 ++++++++-- tests/models/qwen4_exp/test_qsa_reference.py | 57 +++++++++++++++++-- 6 files changed, 127 insertions(+), 23 deletions(-) diff --git a/.sync/vllm-sha b/.sync/vllm-sha index 225269d7f7..c8f632a4c9 100644 --- a/.sync/vllm-sha +++ b/.sync/vllm-sha @@ -1 +1 @@ -86484465153f08f4523d8be964fd6309146eab09 +94e26dd3dd7d6363b524c96521c78feb501bdc61 diff --git a/aphrodite/models/qwen4_exp/common/qsa_cache.py b/aphrodite/models/qwen4_exp/common/qsa_cache.py index bfc3d97431..081d09e75a 100644 --- a/aphrodite/models/qwen4_exp/common/qsa_cache.py +++ b/aphrodite/models/qwen4_exp/common/qsa_cache.py @@ -649,10 +649,16 @@ def build( class QSAStateBackend(AttentionBackend): - """Key-only dummy backend for out-of-band BF16 QSA side-cache operations.""" + """Key-only dummy backend for out-of-band QSA side-cache operations.""" supported_dtypes: ClassVar[list[torch.dtype]] = [torch.bfloat16] - supported_kv_cache_dtypes: ClassVar[list[CacheDType]] = ["auto", "bfloat16"] + # fp8 entries allow the optional e4m3 compressed indexer cache. + supported_kv_cache_dtypes: ClassVar[list[CacheDType]] = [ + "auto", + "bfloat16", + "fp8", + "fp8_e4m3", + ] @staticmethod def get_name() -> str: @@ -710,8 +716,12 @@ def bind_kv_cache(self, kv_cache: torch.Tensor) -> None: """Adapt the unified [B, H, N, C] view to QSA's [B, N, H, C].""" if kv_cache.ndim != 4 or kv_cache.shape[1] != 1: raise ValueError("QSA state cache must be [blocks, 1, states, width]") - if kv_cache.dtype != torch.bfloat16 or kv_cache.shape[3] != self.head_size: - raise ValueError("QSA state cache does not match its packed BF16 spec") + if kv_cache.dtype != self.dtype or kv_cache.shape[3] != self.head_size: + raise ValueError( + f"QSA state cache does not match its spec " + f"(dtype {kv_cache.dtype} != {self.dtype} or width " + f"{kv_cache.shape[3]} != {self.head_size})" + ) super().bind_kv_cache(kv_cache.transpose(1, 2)) def get_attn_backend(self) -> type[AttentionBackend]: @@ -767,7 +777,7 @@ def get_kv_cache_spec(self, aphrodite_config: AphroditeConfig) -> KVCacheSpec: class QSACompressedKeyCache(_QSAStateCache): - """Normalized, group-first-RoPE BF16 key at one row per complete group.""" + """Normed, group-first-RoPE key at one row per complete group.""" def get_kv_cache_spec(self, aphrodite_config: AphroditeConfig) -> KVCacheSpec: del aphrodite_config diff --git a/aphrodite/models/qwen4_exp/nvidia/indexer_qsa.py b/aphrodite/models/qwen4_exp/nvidia/indexer_qsa.py index c3fa9823af..ed58638d30 100644 --- a/aphrodite/models/qwen4_exp/nvidia/indexer_qsa.py +++ b/aphrodite/models/qwen4_exp/nvidia/indexer_qsa.py @@ -147,6 +147,19 @@ def __init__( cache_config = aphrodite_config.cache_config cache_prefix = f"{prefix}." if prefix else "" + # Plain e4m3 without scales: Q and the compressed K are RMSNormed + # before quantization, and the logits kernels dot fp8 x fp8 directly. + self.indexer_kv_dtype = aphrodite_config.attention_config.resolve_indexer_kv_dtype("bf16") + if self.indexer_kv_dtype == "fp8": + indexer_dtype = torch.float8_e4m3fn + elif self.indexer_kv_dtype == "bf16": + indexer_dtype = torch.bfloat16 + else: + raise NotImplementedError( + f"indexer_kv_dtype={self.indexer_kv_dtype!r} is not supported " + "by the Qwen4Exp QSA indexer (only 'bf16' or 'fp8')." + ) + self.indexer_dtype = indexer_dtype self.raw_key_cache = QSAKeyStateCache( head_size=self.index_head_dim, dtype=torch.bfloat16, @@ -158,7 +171,7 @@ def __init__( ) self.compressed_key_cache = QSACompressedKeyCache( head_size=self.index_head_dim, - dtype=torch.bfloat16, + dtype=indexer_dtype, compress_ratio=self.compress_ratio, prefix=f"{cache_prefix}compressed_key_cache", cache_config=cache_config, @@ -272,6 +285,7 @@ def forward( num_tokens, self.index_n_heads, self.index_head_dim, + dtype=self.indexer_dtype, ) qsa_pre_indexer( projected_q, @@ -309,6 +323,7 @@ def forward( self.q_layernorm.variance_epsilon, ).reshape_as(q) q = apply_qsa_rope(self.rotary_emb, positions, q) + q = q.to(self.indexer_dtype) raw_key_cache = raw_key_state_cache.key_cache rope_position_cache = raw_key_state_cache.rope_position_cache diff --git a/aphrodite/models/qwen4_exp/nvidia/ops/qsa_indexer.py b/aphrodite/models/qwen4_exp/nvidia/ops/qsa_indexer.py index 11f87d1c74..519d9e6aef 100644 --- a/aphrodite/models/qwen4_exp/nvidia/ops/qsa_indexer.py +++ b/aphrodite/models/qwen4_exp/nvidia/ops/qsa_indexer.py @@ -193,6 +193,9 @@ def _qsa_mqa_paged_prefill_kernel( other=0.0, eviction_policy="evict_first", ) + # With an fp8 cache, SM90 lowers this dot to reduced-precision wgmma + # accumulation; SM100 tcgen05 accumulates in fp32. The SM90 error + # (~1e-4 relative) is far below the e4m3 quantization noise floor. scores = tl.dot(keys, query, out_dtype=tl.float32) scores = tl.reshape(scores, (BLOCK_N, TILE_R, NUM_HEADS_PADDED)) score = tl.sum(tl.maximum(scores, 0.0), axis=2) @@ -320,7 +323,7 @@ def warmup_qsa_mqa_paged_decode( for decode_query_len, num_requests in profiles: num_rows = decode_query_len * num_requests q_ptr = TritonWarmupTensor( - torch.bfloat16, + k_cache.dtype, shape=(num_rows, num_heads, head_dim), ) logits_ptr = TritonWarmupTensor( @@ -348,7 +351,8 @@ def warmup_qsa_mqa_paged_decode( BLOCK_N=_DECODE_BLOCK_N, TILES_PER_PROG=tiles_per_program, STAGES=2, - num_warps=2, + # tuned on GB300 + num_warps=1 if k_cache.dtype == torch.float8_e4m3fn else 2, grid=( num_requests, triton.cdiv(columns, _DECODE_BLOCK_N * tiles_per_program), @@ -374,7 +378,11 @@ def _prefill_logits( assert 0 < logits_width <= page_table.shape[1] * k_cache.shape[1] logits = torch.empty((num_queries, logits_width), dtype=torch.float32, device=q.device) - TILE_R = 64 + # tuned on GB300 + if k_cache.dtype == torch.float8_e4m3fn: + TILE_R, STAGES, num_warps = 32, 2, 8 + else: + TILE_R, STAGES, num_warps = 64, 2, 4 BLOCK_N = 64 K_TILES = 16 grid = ( @@ -402,8 +410,8 @@ def _prefill_logits( TILE_R=TILE_R, BLOCK_N=BLOCK_N, K_TILES=K_TILES, - STAGES=2, - num_warps=4, + STAGES=STAGES, + num_warps=num_warps, ) return logits @@ -498,6 +506,7 @@ def qsa_select_paged_decode( assert token_topk % compress_ratio == 0 assert block_indices.shape == (q.shape[0], token_topk // compress_ratio) assert decode_query_len > 0 and q.shape[0] % decode_query_len == 0 + assert q.dtype == k_cache.dtype, "Q and the compressed K cache must match" num_requests = q.shape[0] // decode_query_len assert page_table.shape[0] == num_requests assert visible_blocks.shape == (q.shape[0],) @@ -527,7 +536,8 @@ def qsa_select_paged_decode( BLOCK_N=_DECODE_BLOCK_N, TILES_PER_PROG=tiles_per_program, STAGES=2, - num_warps=2, + # tuned on GB300 + num_warps=1 if k_cache.dtype == torch.float8_e4m3fn else 2, ) _topk( logits, @@ -570,6 +580,7 @@ def qsa_select_paged_prefill( assert token_topk % compress_ratio == 0 assert block_indices.shape == (q.shape[0], token_topk // compress_ratio) + assert q.dtype == k_cache.dtype, "Q and the compressed K cache must match" rows = q.shape[0] # No row scores beyond cdiv(max_seq_len, compress_ratio) compressed # columns. Round up to 64 to keep the logits row stride diff --git a/tests/models/qwen4_exp/test_qsa_pre_indexer.py b/tests/models/qwen4_exp/test_qsa_pre_indexer.py index 3e384d2111..fabebe0377 100644 --- a/tests/models/qwen4_exp/test_qsa_pre_indexer.py +++ b/tests/models/qwen4_exp/test_qsa_pre_indexer.py @@ -49,8 +49,19 @@ def _make_block_table(block_counts): return table, num_blocks +def assert_fp8_within_one_ulp(actual: torch.Tensor, expected: torch.Tensor) -> None: + # e4m3 is sign-magnitude, so within a sign the uint8 code order matches the + # value order and one ulp is one code step. The two paths' intermediates + # differ in the pooling accumulation order, which at denormal magnitudes + # (absolute grid step 2^-9) shows up as up to 2 code steps. + code_diff = (actual.view(torch.uint8).int() - expected.view(torch.uint8).int()).abs() + abs_diff = (actual.float() - expected.float()).abs() + assert bool(((code_diff <= 1) | (abs_diff <= 2**-8)).all()) + + @requires_qsa_kernels @pytest.mark.usefixtures("default_aphrodite_config") +@pytest.mark.parametrize("indexer_dtype", [torch.bfloat16, torch.float8_e4m3fn]) @pytest.mark.parametrize( "mrope,is_2d_positions,cache_rope_positions,state_size,seq_lens,query_lens,history_lens", [ @@ -72,6 +83,7 @@ def _make_block_table(block_counts): ], ) def test_qsa_fused_pre_indexer_matches_unfused( + indexer_dtype, mrope, is_2d_positions, cache_rope_positions, @@ -199,7 +211,7 @@ def test_qsa_fused_pre_indexer_matches_unfused( fused_compressed_storage = torch.zeros( num_compressed_blocks, compressed_page_elements + 16, - dtype=torch.bfloat16, + dtype=indexer_dtype, device=device, ) fused_compressed = torch.as_strided( @@ -213,7 +225,7 @@ def test_qsa_fused_pre_indexer_matches_unfused( q_weight = torch.randn(D, dtype=torch.bfloat16, device=device) * 0.2 k_weight = torch.randn(D, dtype=torch.bfloat16, device=device) * 0.2 - fused_query = torch.empty(num_tokens, HQ, D, dtype=torch.bfloat16, device=device) + fused_query = torch.empty(num_tokens, HQ, D, dtype=indexer_dtype, device=device) qsa_pre_indexer( projected_qk[:, : HQ * D], projected_qk[:, HQ * D :], @@ -262,6 +274,17 @@ def test_qsa_fused_pre_indexer_matches_unfused( if rope_positions is not None: qsa_store_cache_rows(rope_positions, raw_slots, position_rows) - torch.testing.assert_close(fused_query, unfused_query, rtol=RTOL, atol=ATOL) + if indexer_dtype == torch.float8_e4m3fn: + # Both paths round the same ~bf16 intermediates to e4m3. Bitwise + # equality (the dsv4 indexer test's bar) does not hold here: the 1D + # reference rope (forward_cuda) differs from _norm_rope by up to 1 + # bf16 ulp, and even the MRoPE pair flips a code occasionally. + unfused_query = unfused_query.to(indexer_dtype) + assert_fp8_within_one_ulp(fused_query, unfused_query) + else: + torch.testing.assert_close(fused_query, unfused_query, rtol=RTOL, atol=ATOL) assert torch.equal(fused_raw.view(torch.int16), unfused_raw.view(torch.int16)) - torch.testing.assert_close(fused_compressed, unfused_compressed, rtol=RTOL, atol=ATOL) + if indexer_dtype == torch.float8_e4m3fn: + assert_fp8_within_one_ulp(fused_compressed, unfused_compressed) + else: + torch.testing.assert_close(fused_compressed, unfused_compressed, rtol=RTOL, atol=ATOL) diff --git a/tests/models/qwen4_exp/test_qsa_reference.py b/tests/models/qwen4_exp/test_qsa_reference.py index 09bb1072e6..120572d5c2 100644 --- a/tests/models/qwen4_exp/test_qsa_reference.py +++ b/tests/models/qwen4_exp/test_qsa_reference.py @@ -48,6 +48,7 @@ def test_qsa_mtp_index_share_updates_cache_but_skips_selection( index_n_heads=1, index_kv_heads=1, index_head_dim=1, + indexer_dtype=torch.bfloat16, raw_key_cache=SimpleNamespace( kv_cache=torch.empty(0), rope_position_cache=None, @@ -468,6 +469,7 @@ def test_qsa_unfused_cache_update_ignores_padded_qk() -> None: use_fused_pre_indexer=False, index_n_heads=1, index_head_dim=64, + indexer_dtype=torch.bfloat16, q_layernorm=norm, k_layernorm=norm, rotary_emb=rope, @@ -624,11 +626,12 @@ def make_buffers() -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tens (4, 33), ], ) -def test_qsa_decode_selection_correctness(decode_query_len: int, num_requests: int) -> None: +@pytest.mark.parametrize("dtype", [torch.bfloat16, torch.float8_e4m3fn]) +def test_qsa_decode_selection_correctness(decode_query_len: int, num_requests: int, dtype: torch.dtype) -> None: torch.manual_seed(1) heads, head_dim = 4, 128 rows = num_requests * decode_query_len - q = torch.randn(rows, heads, head_dim, device="cuda", dtype=torch.bfloat16) + q = torch.randn(rows, heads, head_dim, device="cuda", dtype=torch.bfloat16).to(dtype) page_size, pages_per_request, max_sequence_length = (16, 40, 2560) if num_requests > 32 else (4, 20, 320) num_pages = num_requests * pages_per_request cache = torch.randn( @@ -638,7 +641,7 @@ def test_qsa_decode_selection_correctness(decode_query_len: int, num_requests: i head_dim, device="cuda", dtype=torch.bfloat16, - ) + ).to(dtype) page_table = torch.randperm(num_pages, device="cuda", dtype=torch.int32).reshape(num_requests, pages_per_request) token_to_req = torch.repeat_interleave( torch.arange(num_requests, device="cuda", dtype=torch.int32), @@ -684,14 +687,37 @@ def test_qsa_decode_selection_correctness(decode_query_len: int, num_requests: i compress_ratio, ) + if dtype == torch.float8_e4m3fn: + # fp8 logits tie at the top-k boundary more often than bf16, so index + # identity is not stable; compare the selected value multisets. + # SM90 wgmma accumulates fp8 in reduced precision (~3e-4 abs + # observed); SM100 tcgen05 is exact fp32. + rtol = atol = 1e-3 if current_platform.is_device_capability(90) else None + logits = _qsa_mqa_paged_reference(q, cache, page_table, token_to_req, visible_blocks) + for row in range(rows): + selected = actual[row][actual[row] >= 0] + wanted = expected[row][expected[row] >= 0] + assert selected.numel() == wanted.numel() + torch.testing.assert_close( + logits[row, selected.long()].sort().values, + logits[row, wanted.long()].sort().values, + rtol=rtol, + atol=atol, + ) + return + torch.testing.assert_close(actual.sort().values, expected.sort().values) @requires_qsa_kernels +@pytest.mark.parametrize("dtype", [torch.bfloat16, torch.float8_e4m3fn]) @pytest.mark.parametrize("seq_len_slack", [0, 1792]) @pytest.mark.parametrize("force_chunk", [False, True]) def test_qsa_prefill_selection_correctness( - monkeypatch: pytest.MonkeyPatch, seq_len_slack: int, force_chunk: bool + monkeypatch: pytest.MonkeyPatch, + seq_len_slack: int, + force_chunk: bool, + dtype: torch.dtype, ) -> None: # page_size=24 (does not divide the 64-aligned clipped width) and an # oversized page table, so the clipped logits width comes from @@ -703,8 +729,8 @@ def test_qsa_prefill_selection_correctness( torch.manual_seed(2) query_lens = [3, 33] rows, heads, head_dim = sum(query_lens), 4, 128 - q = torch.randn(rows, heads, head_dim, device="cuda", dtype=torch.bfloat16) - cache = torch.randn(128, 24, 1, head_dim, device="cuda", dtype=torch.bfloat16) + q = torch.randn(rows, heads, head_dim, device="cuda", dtype=torch.bfloat16).to(dtype) + cache = torch.randn(128, 24, 1, head_dim, device="cuda", dtype=torch.bfloat16).to(dtype) page_table = torch.randperm(128, device="cuda", dtype=torch.int32).reshape(2, 64) token_to_req = torch.repeat_interleave( torch.arange(2, device="cuda", dtype=torch.int32), @@ -748,6 +774,25 @@ def test_qsa_prefill_selection_correctness( compress_ratio, ) + if dtype == torch.float8_e4m3fn: + # fp8 logits tie at the top-k boundary more often than bf16, so index + # identity is not stable; compare the selected value multisets. + # SM90 wgmma accumulates fp8 in reduced precision (~3e-4 abs + # observed); SM100 tcgen05 is exact fp32. + rtol = atol = 1e-3 if current_platform.is_device_capability(90) else None + logits = _qsa_mqa_paged_reference(q, cache, page_table, token_to_req, visible_blocks) + for row in range(rows): + selected = actual[row][actual[row] >= 0] + wanted = expected[row][expected[row] >= 0] + assert selected.numel() == wanted.numel() + torch.testing.assert_close( + logits[row, selected.long()].sort().values, + logits[row, wanted.long()].sort().values, + rtol=rtol, + atol=atol, + ) + return + torch.testing.assert_close(actual.sort().values, expected.sort().values) From caee22bba05d8f22972a3f78f307fc0a1e32e988 Mon Sep 17 00:00:00 2001 From: AlpinDale Date: Mon, 7 Sep 2026 20:36:56 +0000 Subject: [PATCH 24/36] [sync] [Bugfix] OffloadingConnector: stop zeroing offload hits under MTP/EAGLE spec decode (#52771) Upstream-vLLM: 4a806d08ee94b35356dd749bf814ee91ee93f8d0 Co-authored-by: Kam Basra --- .sync/vllm-sha | 2 +- .../kv_connector/v1/offloading/scheduler.py | 42 ++- .../offloading_connector/test_scheduler.py | 257 +++++++++++++++++- .../unit/offloading_connector/utils.py | 5 + 4 files changed, 285 insertions(+), 21 deletions(-) diff --git a/.sync/vllm-sha b/.sync/vllm-sha index c8f632a4c9..34544b7e70 100644 --- a/.sync/vllm-sha +++ b/.sync/vllm-sha @@ -1 +1 @@ -94e26dd3dd7d6363b524c96521c78feb501bdc61 +4a806d08ee94b35356dd749bf814ee91ee93f8d0 diff --git a/aphrodite/distributed/kv_transfer/kv_connector/v1/offloading/scheduler.py b/aphrodite/distributed/kv_transfer/kv_connector/v1/offloading/scheduler.py index 8907666cab..f9f8e1b3b0 100644 --- a/aphrodite/distributed/kv_transfer/kv_connector/v1/offloading/scheduler.py +++ b/aphrodite/distributed/kv_transfer/kv_connector/v1/offloading/scheduler.py @@ -196,7 +196,21 @@ def from_spec( and aphrodite_config.speculative_config.use_eagle_block_drop() ) if use_eagle_block_drop and not eagle_groups: - eagle_groups = set(range(len(kv_cache_config.kv_cache_groups))) + # No group is annotated as holding drafter layers. This happens + # for shared-group MTP models (e.g. Qwen3.5-style), whose drafter + # is a regular decoder layer merged into the target's + # full-attention group, and for separate-draft models that miss + # annotation. Treating every group as a drafter group here makes + # the volatile-tail pop apply to ALL groups, which can zero the + # whole request's offload hit whenever any group's servable + # window is a single chunk (issue #52735). Drafter KV served one + # chunk stale can only lower speculative acceptance rates — the + # target model verifies every draft — so fail toward serving. + logger.info_once( + "KV offloading: speculative decoding is enabled but no " + "KV-cache group is annotated as a drafter group; treating " + "all groups as non-draft for offloading." + ) if eagle_groups: logger.info( @@ -357,15 +371,19 @@ def storable_chunks( accepted position may be rewritten after spec-token rejection. During prefill the trailing chunk is stable (the draft input for a chunk's last position is the next prompt token), so it is stored immediately. - The exclusion must be applied consistently everywhere - ``next_stored_chunk_idx`` is derived: otherwise the trailing chunk of - each step is skipped on collection but jumped over by - ``next_stored_chunk_idx``, so it is never re-considered and a - permanent hole breaks prefix-reuse lookup. + Once the request has finished, no further spec-token rejection can + rewrite the tail, so the exclusion is lifted and the final chunk + becomes storable (issue #52735). The exclusion must be applied + consistently everywhere ``next_stored_chunk_idx`` is derived: + otherwise the trailing chunk of each step is skipped on collection but + jumped over by ``next_stored_chunk_idx``, so it is never re-considered + and a permanent hole breaks prefix-reuse lookup. ``is_finished`` is + monotonic, so the finish-time calls all see the lifted exclusion. """ num_chunks = num_offloadable_tokens // group_config.tokens_per_chunk is_decoding = num_offloadable_tokens > self.req.num_prompt_tokens - if group_config.is_eagle_group and is_decoding: + # Finished requests have no pending speculation. + if group_config.is_eagle_group and is_decoding and not self.req.is_finished(): num_chunks = max(0, num_chunks - 1) num_allocated_chunks = len(group_state.block_ids) // self.config.blocks_per_chunk num_keyed_chunks = len(group_state.offload_keys) @@ -698,10 +716,14 @@ def _lookup_complete_chunks(self, req_status: RequestOffloadState) -> int | None sliding_window_size_in_chunks = group_config.sliding_window_size_in_chunks - # For eagle groups, query one extra block that will be popped. - # We only need to increase the query size for sliding window groups. + # For eagle groups, query one extra chunk that will be popped. + # Widening applies to every group type: without it, the pop + # below shrinks max_hit_size_tokens past what was queried, + # which can push the confirmed boundary under a coarser + # sibling group's chunk granularity and zero the whole + # request's hit (issue #52735). query_max = max_hit_size_tokens - if is_eagle_unverified and sliding_window_size_in_chunks is not None: + if is_eagle_unverified: query_max = min( max_hit_size_tokens + tokens_per_chunk, len(offload_keys) * tokens_per_chunk, diff --git a/tests/v1/kv_connector/unit/offloading_connector/test_scheduler.py b/tests/v1/kv_connector/unit/offloading_connector/test_scheduler.py index d5c5dfe64f..00b80625e1 100644 --- a/tests/v1/kv_connector/unit/offloading_connector/test_scheduler.py +++ b/tests/v1/kv_connector/unit/offloading_connector/test_scheduler.py @@ -10,6 +10,7 @@ from aphrodite.distributed.kv_events import BlockStored from aphrodite.distributed.kv_transfer.kv_connector.v1.offloading.common import ( OffloadingConnectorMetadata, + OffloadingWorkerMetadata, ) from aphrodite.distributed.kv_transfer.kv_connector.v1.offloading.config import ( build_offloading_config, @@ -39,6 +40,7 @@ FullAttentionSpec, KVCacheConfig, KVCacheGroupSpec, + MambaSpec, SlidingWindowSpec, ) from aphrodite.v1.kv_cache_spec_registry import KVCacheSpecRegistry @@ -55,6 +57,7 @@ get_offload_group_idx, make_offload_key, ) +from aphrodite.v1.outputs import KVConnectorOutput from aphrodite.v1.request import RequestStatus from tests.v1.kv_connector.unit.offloading_connector.test_config import ( _make_aphrodite_config, @@ -2619,8 +2622,9 @@ def test_non_eagle_tighten_clears_eagle_verified(self, request_runner): Groups: 0=non-eagle full-attn, 1=eagle full-attn. Group 0 has only 1 hit (out of 3 keys) → max_hit tightens to 4. - This clears eagle_verified. Group 1 runs with max_hit=4 → only 1 - key queried, 1 hit, pop to 0 → returns 0. + This clears eagle_verified. Group 1 runs with max_hit=4 → widened + query of 2 keys, 2 hits, pop to 1 → the confirmed boundary holds + at 4 tokens. """ block_size = 4 groups = [ @@ -2663,9 +2667,12 @@ def test_non_eagle_tighten_clears_eagle_verified(self, request_runner): offload_keys_per_group=[[10, 11, 12], [1, 2, 3]], ) # Group 0 (non-eagle FA): prefix finds 1 hit → max_hit=4, num_hit=4 - # Group 1 (eagle FA): max_hit=4 → num_blocks=1, keys=[1]. - # Finds 1 hit, pop to 0 → new_num_hit = 0 < block_size → return 0 - assert sched._lookup(req_status) == 0 + # Group 1 (eagle FA): max_hit=4, widened query → keys=[1, 2]. + # Finds 2 hits, pop to 1 → boundary stays at 4 tokens. Before the + # #52735 fix the query was not widened for full-attention groups, + # so the pop landed on the only queried chunk and zeroed the + # whole request. + assert sched._lookup(req_status) == 4 def test_eagle_verified_survives_eagle_tighten(self, request_runner): """Eagle group tightening does NOT clear eagle_verified. @@ -2775,10 +2782,12 @@ def test_full_attn_store_excludes_trailing_decode_block(self, request_runner, as runner.new_request(token_ids=[0] * tokens_per_chunk * 3) runner.manager.prepare_store.side_effect = lambda keys, req_context: (generate_store_output(keys)) - # 4 decoded tokens fill block 3 entirely with decode tokens (one - # extra token so the block is stored under async scheduling too). + # 4 decoded tokens fill block 3; the extra decode steps let its store + # job complete within this run under both scheduling modes. The + # eagle group holds back its volatile trailing block (1, 3) while the + # request is still decoding, while the normal group stores (0, 3). runner.run( - decoded_tokens=[1, 1, 1, 1, 1, EOS_TOKEN_ID], + decoded_tokens=[1, 1, 1, 1, 1, 1, 1], expected_stored=( (0, 0), (0, 1), @@ -2790,6 +2799,13 @@ def test_full_attn_store_excludes_trailing_decode_block(self, request_runner, as ), ) + # Once the request finishes, no spec rejection can rewrite the tail, + # so the held-back eagle chunk becomes storable. + runner.run( + decoded_tokens=[EOS_TOKEN_ID], + expected_stored=((1, 3),), + ) + @pytest.mark.parametrize("async_scheduling", [True, False]) def test_sw_store_excludes_trailing_decode_block(self, request_runner, async_scheduling: bool): """Eagle sliding-window group stores all prompt blocks but excludes @@ -2826,12 +2842,19 @@ def test_sw_store_excludes_trailing_decode_block(self, request_runner, async_sch runner.new_request(token_ids=[0] * block_size * 3) runner.manager.prepare_store.side_effect = lambda keys, req_context: (generate_store_output(keys)) - # 4 decoded tokens fill block 3 entirely with decode tokens. + # 4 decoded tokens fill block 3 entirely with decode tokens. The + # eagle group holds back its volatile trailing block while decoding. runner.run( - decoded_tokens=[1, 1, 1, 1, EOS_TOKEN_ID], + decoded_tokens=[1, 1, 1, 1], expected_stored=((0, 0), (0, 1), (0, 2)), ) + # Finish lifts the exclusion; the tail block is stored (issue #52735). + runner.run( + decoded_tokens=[EOS_TOKEN_ID], + expected_stored=((0, 3),), + ) + @pytest.mark.parametrize("async_scheduling", [True, False]) def test_single_block_stored_at_end_of_prefill(self, request_runner, async_scheduling: bool): """An eagle group with a single-block prompt stores it at the end of @@ -3647,3 +3670,217 @@ def test_retention_interval_zero_stores_only_replay_boundary(request_runner, asy (1, 11), ), ) + + +def _shared_kv_mtp_config(): + """Speculative config for a shared-group MTP model: eagle-family method + whose drafter layer merges into a target KV-cache group, so no group + self-identifies as a drafter group (issue #52735).""" + spec = MagicMock(name="shared_kv_mtp_spec") + spec.use_eagle.return_value = True + spec.use_eagle_block_drop.return_value = True + spec.use_multi_module_mtp.return_value = False + spec.num_speculative_tokens_per_batch_size = None + spec.max_num_new_slots_for_drafting = 0 + spec.num_speculative_tokens = 1 + return spec + + +class TestSharedGroupMTPOffload: + """Regression tests for issue #52735: OffloadingConnector must keep + serving when speculative decoding is enabled but no KV-cache group is + annotated as a drafter group (shared-group MTP models).""" + + def test_no_annotation_marks_no_groups(self, request_runner): + """Spec decode on + zero annotated groups must NOT mark every group + as a drafter group; the full store->load roundtrip must match the + non-speculative behavior of test_two_groups_full_and_sliding_window.""" + block_size = 4 + kv_cache_groups = [ + KVCacheGroupSpec( + ["layer0"], + FullAttentionSpec( + block_size=block_size, + num_kv_heads=1, + head_size=1, + dtype=torch.float32, + ), + ), + KVCacheGroupSpec( + ["layer1"], + SlidingWindowSpec( + block_size=block_size, + num_kv_heads=1, + head_size=1, + dtype=torch.float32, + sliding_window=8, + ), + ), + ] + runner = request_runner( + block_size=block_size, + num_gpu_blocks=100, + async_scheduling=True, + kv_cache_groups=kv_cache_groups, + speculative_config=_shared_kv_mtp_config(), + ) + kv_group_configs = runner.connector_scheduler.config.kv_group_configs + assert [c.is_eagle_group for c in kv_group_configs] == [False, False] + + runner.new_request(token_ids=[0] * block_size * 3) + runner.manager.prepare_store.side_effect = lambda keys, req_context: (generate_store_output(keys)) + runner.run(decoded_tokens=[0]) + runner.run( + decoded_tokens=[0] * (block_size * 3 + 2), + expected_stored=(0, 1, 2, 3, 4, 5), + ) + runner.run(decoded_tokens=[EOS_TOKEN_ID]) + + runner.scheduler.reset_prefix_cache() + + runner.new_request(token_ids=[0] * (block_size * 3 + 1)) + runner.manager.lookup.return_value = LookupResult.HIT + runner.run( + decoded_tokens=[EOS_TOKEN_ID], + expected_loaded=((0, 0), (0, 1), (0, 2), (1, 1), (1, 2)), + ) + + +class TestMambaHybridOffloadServing: + """Store->finish->lookup flows on a full-attention + mamba-align hybrid, + with a manager that only HITs keys that were actually stored. Guards the + two collapse routes of issue #52735.""" + + BLOCK = 4 + MAMBA_BLOCK = 16 + PROMPT_TOKENS = 28 # 7 hash blocks; 1 full mamba chunk + + def _make_scheduler(self, speculative_config, mamba_eagle=False): + aphrodite_config = _make_aphrodite_config(extra_config={"self_describing_kv_events": True}) + aphrodite_config.cache_config.block_size = self.BLOCK + aphrodite_config.cache_config.prefix_match_unit = self.BLOCK + aphrodite_config.speculative_config = speculative_config + aphrodite_config.kv_events_config = KVEventsConfig(enable_kv_cache_events=True, publisher="null") + kv_cache_config = KVCacheConfig( + num_blocks=64, + kv_cache_tensors=[], + kv_cache_groups=[ + KVCacheGroupSpec( + ["full_layer"], + FullAttentionSpec( + block_size=self.BLOCK, + num_kv_heads=1, + head_size=1, + dtype=torch.float32, + ), + is_eagle_group=mamba_eagle, + ), + KVCacheGroupSpec( + ["mamba_layer"], + MambaSpec( + block_size=self.MAMBA_BLOCK, + shapes=((1, 1),), + dtypes=(torch.float32,), + mamba_cache_mode="align", + ), + ), + ], + ) + spec = MockOffloadingSpec(build_offloading_config(aphrodite_config, kv_cache_config)) + return OffloadingConnectorScheduler(spec, aphrodite_config, kv_cache_config) + + def _roundtrip_served_tokens(self, scheduler) -> int: + """Request A: prompt-only store, finish; request A2: served tokens.""" + stored: set = set() + scheduler.manager.lookup.side_effect = lambda key, ctx: ( + LookupResult.HIT if key in stored else LookupResult.MISS + ) + scheduler.manager.prepare_store.side_effect = lambda keys, ctx: (generate_store_output(list(keys))) + + def complete_all_jobs(): + job_ids = list(scheduler._jobs.keys()) + for jid in job_ids: + stored.update(scheduler._jobs[jid].keys) + if job_ids: + scheduler.update_connector_output( + KVConnectorOutput( + kv_connector_worker_meta=OffloadingWorkerMetadata(completed_jobs={jid: 1 for jid in job_ids}) + ) + ) + + def make_request(req_id): + request = MagicMock() + request.request_id = req_id + request.kv_transfer_params = None + request.num_prompt_tokens = self.PROMPT_TOKENS + request.num_tokens = self.PROMPT_TOKENS + request.num_computed_tokens = 0 + request.block_hashes = [BlockHash(f"h{i}".encode()) for i in range(self.PROMPT_TOKENS // self.BLOCK)] + request.all_token_ids = list(range(self.PROMPT_TOKENS)) + request.lora_request = None + request.is_finished.return_value = False + request.status = None + scheduler.on_new_request(request) + return request + + req = make_request("A") + req_status = scheduler._req_status["A"] + req_status.num_locally_computed_tokens = 0 + req_status.update_offload_keys() + assert scheduler._lookup(req_status) == 0 # cold + + n_full = self.PROMPT_TOKENS // self.BLOCK + n_mamba = self.PROMPT_TOKENS // self.MAMBA_BLOCK + req_status.group_states[0].block_ids[:] = list(range(1, n_full + 1)) + req_status.group_states[1].block_ids[:] = list(range(101, 101 + n_mamba + 1)) + + out = SimpleNamespace( + num_scheduled_tokens={"A": self.PROMPT_TOKENS}, + finished_req_ids=set(), + kv_connector_block_state=KVConnectorBlockState( + block_ids={}, + boundary_state_offloads={"A": [(1, 101, self.MAMBA_BLOCK)]}, + ), + ) + scheduler._build_partial_tail_store_jobs(out) + scheduler._build_store_jobs(out) + complete_all_jobs() + + req.num_computed_tokens = self.PROMPT_TOKENS + req.num_tokens = self.PROMPT_TOKENS + 1 + req.is_finished.return_value = True + req_status.update_offload_keys() + scheduler.request_finished(req) + out2 = SimpleNamespace(num_scheduled_tokens={}, finished_req_ids={"A"}) + scheduler._build_store_jobs(out2) + complete_all_jobs() + + make_request("A2") + rs2 = scheduler._req_status["A2"] + rs2.num_locally_computed_tokens = 0 + rs2.update_offload_keys() + served = scheduler._lookup(rs2) + return served if served is not None else 0 + + def test_shared_group_mtp_serves_like_non_speculative(self): + """Route 1 of #52735: with spec decode on and no drafter annotation, + serving must equal the non-speculative baseline (was 0 before the + fix: the all-groups fallback marked the mamba group as a drafter and + the volatile-tail pop consumed its only servable chunk).""" + baseline = self._roundtrip_served_tokens(self._make_scheduler(None)) + assert baseline == 16 + served = self._roundtrip_served_tokens(self._make_scheduler(_shared_kv_mtp_config())) + assert served == baseline + + def test_annotated_eagle_group_does_not_starve_coarse_sibling(self): + """Route 2 of #52735 (F1): a genuinely annotated eagle full-attention + group must not drag the confirmed boundary below a coarser sibling's + chunk granularity. The widened query makes the volatile-tail pop + land on an extra queried chunk, holding the boundary at 16 tokens + (was 0 before the fix: 4 hits -> pop -> 12 tokens < mamba chunk).""" + scheduler = self._make_scheduler(None, mamba_eagle=True) + assert [c.is_eagle_group for c in scheduler.config.kv_group_configs] == [ + True, + False, + ] + assert self._roundtrip_served_tokens(scheduler) == 16 diff --git a/tests/v1/kv_connector/unit/offloading_connector/utils.py b/tests/v1/kv_connector/unit/offloading_connector/utils.py index f099ec5fe1..620def5d62 100644 --- a/tests/v1/kv_connector/unit/offloading_connector/utils.py +++ b/tests/v1/kv_connector/unit/offloading_connector/utils.py @@ -172,6 +172,7 @@ def __init__( extra_config_overrides: dict[str, Any] | None = None, worker_count: int = 1, retention_interval: int | None = None, + speculative_config: Any | None = None, ): assert blocks_per_chunk == 1 or kv_cache_groups is None, ( "blocks_per_chunk > 1 requires all groups to have the same " @@ -193,6 +194,8 @@ def __init__( aphrodite_config.scheduler_config.async_scheduling = async_scheduling aphrodite_config.parallel_config.world_size = worker_count aphrodite_config.cache_config.prefix_cache_retention_interval = retention_interval + if speculative_config is not None: + aphrodite_config.speculative_config = speculative_config extra_config: dict[str, Any] = { "spec_name": "MockOffloadingSpec", @@ -628,6 +631,7 @@ def runner_factory( extra_config_overrides=None, worker_count=1, retention_interval=None, + speculative_config=None, ): runner = RequestRunner( block_size=block_size, @@ -638,6 +642,7 @@ def runner_factory( extra_config_overrides=extra_config_overrides, worker_count=worker_count, retention_interval=retention_interval, + speculative_config=speculative_config, ) runners.append(runner) return runner From 78f7f7711e23764d50625e8b40e2d68392977d27 Mon Sep 17 00:00:00 2001 From: AlpinDale Date: Mon, 7 Sep 2026 20:39:22 +0000 Subject: [PATCH 25/36] [sync] [Bugfix][Quantization][XPU] Fix moe_wna16 linear weight loading (#52651) Upstream-vLLM: 5e6f6a8ed42fd420ef206b938e3bb48739a95d6c Co-authored-by: Artur Fierka --- .sync/vllm-sha | 2 +- .../layers/quantization/moe_wna16.py | 13 ++++- aphrodite/platforms/xpu.py | 1 + tests/quantization/test_moe_wna16.py | 51 +++++++++++++++++++ 4 files changed, 64 insertions(+), 3 deletions(-) diff --git a/.sync/vllm-sha b/.sync/vllm-sha index 34544b7e70..34049c0ed1 100644 --- a/.sync/vllm-sha +++ b/.sync/vllm-sha @@ -1 +1 @@ -4a806d08ee94b35356dd749bf814ee91ee93f8d0 +5e6f6a8ed42fd420ef206b938e3bb48739a95d6c diff --git a/aphrodite/model_executor/layers/quantization/moe_wna16.py b/aphrodite/model_executor/layers/quantization/moe_wna16.py index dae45f65d5..5d8ac47280 100644 --- a/aphrodite/model_executor/layers/quantization/moe_wna16.py +++ b/aphrodite/model_executor/layers/quantization/moe_wna16.py @@ -168,12 +168,21 @@ def get_quant_method(self, layer: torch.nn.Module, prefix: str) -> "QuantizeMeth AutoGPTQConfig, ) + linear_config: QuantizationConfig if self.linear_quant_method == "gptq": - return AutoGPTQConfig.from_config(self.full_config).get_quant_method(layer, prefix) + linear_config = AutoGPTQConfig.from_config(self.full_config) elif self.linear_quant_method in ("awq", "awq_marlin"): - return AutoAWQConfig.from_config(self.full_config).get_quant_method(layer, prefix) + linear_config = AutoAWQConfig.from_config(self.full_config) else: raise ValueError("moe_wna16 only support gptq and awq.") + + # The delegate is rebuilt from the raw HF quantization dict, which + # lists shard names and never fused ones. Without the mapping it + # resolves fused layers (qkv_proj, gate_up_proj) to + # UnquantizedLinearMethod and the checkpoint's qweight has nowhere + # to load. + linear_config.packed_modules_mapping = self.packed_modules_mapping + return linear_config.get_quant_method(layer, prefix) elif isinstance(layer, RoutedExperts): return MoeWNA16Method(self, layer.moe_config) return None diff --git a/aphrodite/platforms/xpu.py b/aphrodite/platforms/xpu.py index 7006a86691..984e6052c9 100644 --- a/aphrodite/platforms/xpu.py +++ b/aphrodite/platforms/xpu.py @@ -104,6 +104,7 @@ class XPUPlatform(Platform): supported_quantization: list[str] = [ "awq", "gptq", + "moe_wna16", "auto_awq", "auto_gptq", "inc", diff --git a/tests/quantization/test_moe_wna16.py b/tests/quantization/test_moe_wna16.py index 7d431b2093..d7cbcdc56a 100644 --- a/tests/quantization/test_moe_wna16.py +++ b/tests/quantization/test_moe_wna16.py @@ -305,3 +305,54 @@ def test_compressed_tensors_wna16_moe_converts_and_sets_up_humming_kernel(): assert not hasattr(layer, "w2_weight_packed") assert layer.w13_weight.dtype is torch.int32 assert layer.w2_weight.dtype is torch.int32 + + +def test_moe_wna16_forwards_packed_modules_mapping_to_linear_delegate(monkeypatch): + """The linear delegate must receive packed_modules_mapping. + + It is rebuilt from the raw HF quantization dict, which lists shard names and + never fused ones, so without the mapping a fused layer resolves to + `UnquantizedLinearMethod` and the checkpoint's qweight has nowhere to load. + """ + from aphrodite.model_executor.layers.linear import ColumnParallelLinear + from aphrodite.model_executor.layers.quantization.auto_gptq import AutoGPTQConfig + + config = MoeWNA16Config( + linear_quant_method="gptq", + weight_bits=4, + group_size=128, + has_zp=False, + lm_head_quantized=False, + modules_to_not_convert=None, + full_config={ + "bits": 4, + "group_size": 128, + "desc_act": False, + "sym": True, + "quant_method": "gptq", + # As emitted by AutoGPTQ: shard names, never the fused name. + "modules_in_block_to_quantize": [["mlp.gate_proj", "mlp.up_proj"]], + }, + ) + config.packed_modules_mapping = {"gate_up_proj": ["gate_proj", "up_proj"]} + + seen: dict[str, dict[str, list[str]]] = {} + monkeypatch.setattr( + AutoGPTQConfig, + "get_quant_method", + lambda self, layer, prefix: seen.setdefault("mapping", self.packed_modules_mapping), + ) + layer = ColumnParallelLinear.__new__(ColumnParallelLinear) + config.get_quant_method(layer, "model.layers.0.mlp.gate_up_proj") + + assert seen["mapping"] == {"gate_up_proj": ["gate_proj", "up_proj"]} + + +def test_xpu_platform_supports_moe_wna16(): + """Regression guard for the XPU quantization allowlist.""" + try: + from aphrodite.platforms.xpu import XPUPlatform + except ImportError: + pytest.skip("aphrodite_xpu_kernels not importable outside an XPU stack") + + assert "moe_wna16" in XPUPlatform.supported_quantization From 5951e2f627382d75bca76651d6bbae73a05b912f Mon Sep 17 00:00:00 2001 From: AlpinDale Date: Mon, 7 Sep 2026 20:45:03 +0000 Subject: [PATCH 26/36] [sync] [Frontend] Add stateless /v1/responses/render endpoint (#50195) Upstream-vLLM: 6a2a2bb02b563b83f946012959fd3927984d072a Co-authored-by: Francisco Javier Arceo --- .sync/vllm-sha | 2 +- aphrodite/entrypoints/generate/api_router.py | 39 +- .../launchers/api_server/app_state.py | 19 +- aphrodite/entrypoints/launchers/cli_args.py | 10 + .../entrypoints/launchers/render/app_state.py | 10 +- aphrodite/entrypoints/mcp/tool_server.py | 15 + .../entrypoints/openai/responses/serving.py | 302 +++--------- aphrodite/entrypoints/scale_out/factories.py | 1 + .../scale_out/render/api_router.py | 23 + .../entrypoints/scale_out/render/serving.py | 63 +++ aphrodite/renderers/online_renderer.py | 251 +++++++++- docs/src/content/docs/deployment/security.md | 5 + docs/src/content/docs/serving/other-apis.md | 29 +- tests/entrypoints/launchers/test_cli_args.py | 22 + .../openai/responses/test_mcp_tools.py | 24 + .../responses/test_serving_responses.py | 460 +++++++++++++++++- .../entrypoints/openai/test_render_parity.py | 3 +- .../scale_out/render/test_render.py | 184 +++++++ tests/entrypoints/scale_out/test_factories.py | 1 + 19 files changed, 1182 insertions(+), 281 deletions(-) diff --git a/.sync/vllm-sha b/.sync/vllm-sha index 34049c0ed1..45322d2b8d 100644 --- a/.sync/vllm-sha +++ b/.sync/vllm-sha @@ -1 +1 @@ -5e6f6a8ed42fd420ef206b938e3bb48739a95d6c +6a2a2bb02b563b83f946012959fd3927984d072a diff --git a/aphrodite/entrypoints/generate/api_router.py b/aphrodite/entrypoints/generate/api_router.py index 48447a45ac..55d33a5d73 100644 --- a/aphrodite/entrypoints/generate/api_router.py +++ b/aphrodite/entrypoints/generate/api_router.py @@ -1,6 +1,6 @@ # SPDX-License-Identifier: Apache-2.0 # SPDX-FileCopyrightText: Copyright contributors to the vLLM project -from typing import TYPE_CHECKING +from typing import TYPE_CHECKING, Any from fastapi import FastAPI @@ -66,6 +66,7 @@ async def init_generate_state( args: "Namespace", request_logger: RequestLogger | None, supported_tasks: tuple["SupportedTask", ...], + default_chat_template_kwargs: dict[str, Any], ): from aphrodite.entrypoints.anthropic.serving import AnthropicServingMessages from aphrodite.entrypoints.chat_utils import load_chat_template @@ -85,12 +86,6 @@ async def init_generate_state( CohereServingChatV2 = None # type: ignore[assignment,misc] else: CohereServingChatV2 = None # type: ignore[assignment,misc] - - from aphrodite.entrypoints.mcp.tool_server import ( - DemoToolServer, - MCPToolServer, - ToolServer, - ) from aphrodite.entrypoints.openai.chat_completion.batch_serving import ( OpenAIServingChatBatch, ) @@ -106,33 +101,11 @@ async def init_generate_state( getattr(args, "fingerprint_value", None), ) - if args.tool_server == "demo": - tool_server: ToolServer | None = DemoToolServer() - assert isinstance(tool_server, DemoToolServer) - await tool_server.init_and_validate() - elif args.tool_server: - tool_server = MCPToolServer() - await tool_server.add_tool_server(args.tool_server) - else: - tool_server = None resolved_chat_template = load_chat_template(args.chat_template) - # Fold the dedicated ``--cohere-format`` CLI flag into the renderer's - # default chat-template kwargs. The cohere renderer reads - # ``chat_template_kwargs["cohere_format"]`` to pick cmd3 vs cmd4 - # rendering; making this a first-class flag keeps the right format - # discoverable for ``aphrodite serve --tokenizer-mode cohere`` users - # without forcing them to hand-construct a JSON dict for - # ``--default-chat-template-kwargs``. Per-request overrides still - # take precedence (see ``merge_kwargs`` in - # ``ChatCompletionRequest.build_chat_params``). - default_chat_template_kwargs = dict(args.default_chat_template_kwargs or {}) - if getattr(args, "cohere_format", None): - default_chat_template_kwargs.setdefault("cohere_format", args.cohere_format) - - # Render endpoints are always backed by OnlineRenderer so that - # /v1/chat/completions/render and /v1/completions/render work on both - # generate-mode and render-only servers. Created in init_app_state. + # Render endpoints are always backed by OnlineRenderer so that chat, + # completion, and Responses rendering work on both generate-mode and + # render-only servers. Created in init_app_state. state.openai_serving_responses = ( OpenAIServingResponses( @@ -145,7 +118,7 @@ async def init_generate_state( return_tokens_as_token_ids=args.return_tokens_as_token_ids, enable_auto_tools=args.enable_auto_tool_choice, tool_parser=args.tool_call_parser, - tool_server=tool_server, + tool_server=state.tool_server, reasoning_parser=args.structured_outputs_config.reasoning_parser, enable_prompt_tokens_details=args.enable_prompt_tokens_details, enable_force_include_usage=args.enable_force_include_usage, diff --git a/aphrodite/entrypoints/launchers/api_server/app_state.py b/aphrodite/entrypoints/launchers/api_server/app_state.py index 3c0d616093..271abf450c 100644 --- a/aphrodite/entrypoints/launchers/api_server/app_state.py +++ b/aphrodite/entrypoints/launchers/api_server/app_state.py @@ -8,6 +8,8 @@ from aphrodite.engine.protocol import EngineClient from aphrodite.entrypoints.chat_utils import load_chat_template +from aphrodite.entrypoints.launchers.cli_args import resolve_default_chat_template_kwargs +from aphrodite.entrypoints.mcp.tool_server import init_tool_server from aphrodite.entrypoints.openai.models.protocol import BaseModelPath from aphrodite.entrypoints.openai.models.serving import OpenAIServingModels from aphrodite.entrypoints.serve.tokenize.serving import ServingTokenization @@ -59,6 +61,8 @@ async def init_app_state( state.aphrodite_config = aphrodite_config state.args = args resolved_chat_template = load_chat_template(args.chat_template) + default_chat_template_kwargs = resolve_default_chat_template_kwargs(args) + state.tool_server = await init_tool_server(args) if "generate" in supported_tasks else None # Merge default_mm_loras into the static lora_modules default_mm_loras = aphrodite_config.lora_config.default_mm_loras if aphrodite_config.lora_config is not None else {} @@ -82,7 +86,7 @@ async def init_app_state( exclude_tools_when_tool_choice_none=args.exclude_tools_when_tool_choice_none, tool_parser=args.tool_call_parser, reasoning_parser=args.structured_outputs_config.reasoning_parser, - default_chat_template_kwargs=args.default_chat_template_kwargs, + default_chat_template_kwargs=default_chat_template_kwargs, log_error_stack=args.log_error_stack, ) state.online_renderer.warmup() @@ -98,7 +102,7 @@ async def init_app_state( exclude_tools_when_tool_choice_none=args.exclude_tools_when_tool_choice_none, tool_parser=args.tool_call_parser, reasoning_parser=args.structured_outputs_config.reasoning_parser, - default_chat_template_kwargs=args.default_chat_template_kwargs, + default_chat_template_kwargs=default_chat_template_kwargs, log_error_stack=args.log_error_stack, ) @@ -108,14 +112,21 @@ async def init_app_state( request_logger=request_logger, chat_template=resolved_chat_template, chat_template_content_format=args.chat_template_content_format, - default_chat_template_kwargs=args.default_chat_template_kwargs, + default_chat_template_kwargs=default_chat_template_kwargs, trust_request_chat_template=args.trust_request_chat_template, ) if "generate" in supported_tasks: from aphrodite.entrypoints.generate.api_router import init_generate_state - await init_generate_state(engine_client, state, args, request_logger, supported_tasks) + await init_generate_state( + engine_client, + state, + args, + request_logger, + supported_tasks, + default_chat_template_kwargs, + ) from aphrodite.entrypoints.scale_out.factories import init_scale_out_state diff --git a/aphrodite/entrypoints/launchers/cli_args.py b/aphrodite/entrypoints/launchers/cli_args.py index 37f108336a..482bf0efcb 100644 --- a/aphrodite/entrypoints/launchers/cli_args.py +++ b/aphrodite/entrypoints/launchers/cli_args.py @@ -362,6 +362,16 @@ def _customize_cli_kwargs( return frontend_kwargs +def resolve_default_chat_template_kwargs( + args: argparse.Namespace, +) -> dict[str, Any]: + """Resolve renderer defaults, including the dedicated Cohere format flag.""" + kwargs = dict(args.default_chat_template_kwargs or {}) + if getattr(args, "cohere_format", None): + kwargs.setdefault("cohere_format", args.cohere_format) + return kwargs + + def make_arg_parser(parser: FlexibleArgumentParser) -> FlexibleArgumentParser: """Create the CLI argument parser used by the OpenAI API server. diff --git a/aphrodite/entrypoints/launchers/render/app_state.py b/aphrodite/entrypoints/launchers/render/app_state.py index a9a59cf873..ac0b63884d 100644 --- a/aphrodite/entrypoints/launchers/render/app_state.py +++ b/aphrodite/entrypoints/launchers/render/app_state.py @@ -6,6 +6,8 @@ from aphrodite.config import AphroditeConfig from aphrodite.entrypoints.chat_utils import load_chat_template +from aphrodite.entrypoints.launchers.cli_args import resolve_default_chat_template_kwargs +from aphrodite.entrypoints.mcp.tool_server import init_tool_server from aphrodite.entrypoints.openai.models.protocol import BaseModelPath from aphrodite.entrypoints.openai.models.serving import OpenAIModelRegistry from aphrodite.entrypoints.scale_out.factories import init_render_state @@ -43,6 +45,8 @@ async def init_render_app_state( renderer = renderer_from_config(aphrodite_config) resolved_chat_template = load_chat_template(args.chat_template) + default_chat_template_kwargs = resolve_default_chat_template_kwargs(args) + state.tool_server = await init_tool_server(args) state.online_renderer = OnlineRenderer( model_config=aphrodite_config.model_config, @@ -55,7 +59,7 @@ async def init_render_app_state( exclude_tools_when_tool_choice_none=args.exclude_tools_when_tool_choice_none, tool_parser=args.tool_call_parser, reasoning_parser=args.reasoning_parser, - default_chat_template_kwargs=args.default_chat_template_kwargs, + default_chat_template_kwargs=default_chat_template_kwargs, log_error_stack=args.log_error_stack, ) state.online_renderer.warmup() @@ -71,7 +75,7 @@ async def init_render_app_state( exclude_tools_when_tool_choice_none=args.exclude_tools_when_tool_choice_none, tool_parser=args.tool_call_parser, reasoning_parser=args.reasoning_parser, - default_chat_template_kwargs=args.default_chat_template_kwargs, + default_chat_template_kwargs=default_chat_template_kwargs, log_error_stack=args.log_error_stack, ) @@ -82,7 +86,7 @@ async def init_render_app_state( request_logger=request_logger, chat_template=resolved_chat_template, chat_template_content_format=args.chat_template_content_format, - default_chat_template_kwargs=args.default_chat_template_kwargs, + default_chat_template_kwargs=default_chat_template_kwargs, trust_request_chat_template=args.trust_request_chat_template, ) diff --git a/aphrodite/entrypoints/mcp/tool_server.py b/aphrodite/entrypoints/mcp/tool_server.py index 13604e3102..c5fcc0d856 100644 --- a/aphrodite/entrypoints/mcp/tool_server.py +++ b/aphrodite/entrypoints/mcp/tool_server.py @@ -12,6 +12,8 @@ logger = init_logger(__name__) if TYPE_CHECKING: + from argparse import Namespace + from mcp.types import ListToolsResult @@ -215,3 +217,16 @@ async def new_session(self, tool_name: str, session_id: str, headers: dict[str, if tool_name not in self.tools: raise KeyError(f"Tool '{tool_name}' is not supported") yield self.tools[tool_name] + + +async def init_tool_server(args: "Namespace") -> ToolServer | None: + tool_server = getattr(args, "tool_server", None) + if tool_server == "demo": + demo_server = DemoToolServer() + await demo_server.init_and_validate() + return demo_server + if tool_server: + mcp_server = MCPToolServer() + await mcp_server.add_tool_server(tool_server) + return mcp_server + return None diff --git a/aphrodite/entrypoints/openai/responses/serving.py b/aphrodite/entrypoints/openai/responses/serving.py index e2fc438aa7..d0d25a0e63 100644 --- a/aphrodite/entrypoints/openai/responses/serving.py +++ b/aphrodite/entrypoints/openai/responses/serving.py @@ -6,13 +6,11 @@ from collections import deque from collections.abc import AsyncGenerator, AsyncIterator, Callable, Mapping, Sequence from contextlib import AsyncExitStack -from copy import copy from http import HTTPStatus -from typing import Any, Final +from typing import Any, Final, cast from fastapi import Request from openai.types.responses import ( - ResponseFunctionToolCall, ResponseOutputItem, ResponseOutputMessage, ResponseOutputText, @@ -20,15 +18,12 @@ response_text_delta_event, ) from openai.types.responses.response_output_text import Logprob, LogprobTopLogprob -from openai.types.responses.tool import Mcp, Tool -from openai_harmony import Message as OpenAIHarmonyMessage from pydantic import TypeAdapter from aphrodite import envs from aphrodite.config.utils import replace from aphrodite.engine.protocol import EngineClient from aphrodite.entrypoints.chat_utils import ( - ChatCompletionMessageParam, ChatTemplateContentFormatOption, ) from aphrodite.entrypoints.generate.base.protocol import ( @@ -38,24 +33,13 @@ from aphrodite.entrypoints.generate.base.serving import GenerateBaseServing from aphrodite.entrypoints.mcp.tool_server import ToolServer from aphrodite.entrypoints.openai.models.serving import OpenAIServingModels -from aphrodite.entrypoints.openai.parser.harmony_utils import ( - build_harmony_preamble, - extract_instructions_from_messages, - get_user_message, - has_custom_tools, - render_for_completion, -) from aphrodite.entrypoints.openai.responses.context import ( ConversationContext, HarmonyContext, ParsableContext, SimpleContext, ) -from aphrodite.entrypoints.openai.responses.harmony import ( - construct_harmony_previous_input_messages, - harmony_to_response_output, - response_input_to_harmony, -) +from aphrodite.entrypoints.openai.responses.harmony import harmony_to_response_output from aphrodite.entrypoints.openai.responses.protocol import ( InputTokensDetails, OutputTokensDetails, @@ -80,8 +64,6 @@ ) from aphrodite.entrypoints.openai.responses.utils import ( build_response_output_items, - construct_input_messages, - construct_tool_dicts, extract_function_tool_names, extract_tool_types, ) @@ -89,14 +71,19 @@ from aphrodite.entrypoints.serve.utils.api_utils import get_max_tokens from aphrodite.entrypoints.serve.utils.request_logger import RequestLogger from aphrodite.exceptions import APHRODITEValidationError, GenerationError -from aphrodite.inputs import EngineInput, tokens_input +from aphrodite.inputs import EngineInput from aphrodite.logger import init_logger from aphrodite.logprobs import Logprob as SampleLogprob from aphrodite.logprobs import SampleLogprobs from aphrodite.lora.request import LoRARequest from aphrodite.outputs import CompletionOutput from aphrodite.parser import Parser, ParserManager -from aphrodite.renderers.online_renderer import OnlineRenderer +from aphrodite.renderers import TokenizeParams +from aphrodite.renderers.online_renderer import ( + OnlineRenderer, + ResponsesPreviousMessages, + ResponsesRenderResult, +) from aphrodite.sampling_params import SamplingParams, StructuredOutputsParams from aphrodite.tokenizers import TokenizerLike from aphrodite.utils import random_uuid @@ -105,45 +92,6 @@ logger = init_logger(__name__) -def _extract_allowed_tools_from_mcp_requests( - tools: list[Tool], -) -> dict[str, list[str] | None]: - """ - Extract allowed_tools mapping from MCP tool requests. - - Returns a dictionary mapping server_label to allowed_tools list. - Handles both list format and McpAllowedToolsMcpToolFilter object format. - - Special handling: - - If allowed_tools is None, returns None (allows all tools) - - If allowed_tools contains "*", returns None (allows all tools) - - Otherwise, returns the list of specific tool names - - This function can be reused for both harmony and non-harmony MCP calls. - """ - allowed_tools_map: dict[str, list[str] | None] = {} - for tool in tools: - if not isinstance(tool, Mcp): - continue - - # allowed_tools can be a list or an object with tool_names - # Extract the actual list of tool names - allowed_tools_val = None - if tool.allowed_tools is not None: - if isinstance(tool.allowed_tools, list): - allowed_tools_val = tool.allowed_tools - elif hasattr(tool.allowed_tools, "tool_names"): - # It's an McpAllowedToolsMcpToolFilter object - allowed_tools_val = tool.allowed_tools.tool_names - - # Normalize "*" to None (both mean "allow all tools") - if allowed_tools_val is not None and "*" in allowed_tools_val: - allowed_tools_val = None - - allowed_tools_map[tool.server_label] = allowed_tools_val - return allowed_tools_map - - class OpenAIServingResponses(GenerateBaseServing): def __init__( self, @@ -223,7 +171,7 @@ def __init__( # HACK(woosuk): This is a hack. We should use a better store. # FIXME: If enable_store=True, this may cause a memory leak since we # never remove messages from the store. - self.msg_store: dict[str, list[ChatCompletionMessageParam]] = {} + self.msg_store: dict[str, ResponsesPreviousMessages] = {} # HACK(wuhang): This is a hack. We should use a better store. # FIXME: If enable_store=True, this may cause a memory leak since we @@ -312,6 +260,31 @@ def _validate_create_responses_input(self, request: ResponsesRequest) -> ErrorRe ) return None + async def _render_resolved_response_inputs( + self, + request: ResponsesRequest, + prev_response: ResponsesResponse | None, + *, + skip_mm_cache: bool = False, + ) -> ResponsesRenderResult | ErrorResponse: + previous_messages: ResponsesPreviousMessages | None = None + previous_response_outputs: list[ResponseOutputItem] | None = None + if prev_response is not None: + stored_messages = ( + self.msg_store[prev_response.id] if self.use_harmony else self.msg_store.get(prev_response.id) + ) + if stored_messages is not None: + previous_messages = stored_messages.copy() + previous_response_outputs = list(prev_response.output) + + return await self.online_renderer.render_responses( + request, + previous_messages=previous_messages, + previous_response_outputs=previous_response_outputs, + tool_server=self.tool_server, + skip_mm_cache=skip_mm_cache, + ) + async def create_responses( self, request: ResponsesRequest, @@ -356,10 +329,14 @@ async def _create_responses( lora_request = self._maybe_get_adapters(request) model_name = self.models.model_name(lora_request) - if self.use_harmony: - messages, engine_inputs = self._make_request_with_harmony(request, prev_response) - else: - messages, engine_inputs = await self._make_request(request, prev_response) + render_result = await self._render_resolved_response_inputs( + request, + prev_response, + ) + if isinstance(render_result, ErrorResponse): + return render_result + messages = render_result.messages + engine_inputs = [render_result.engine_input] request_metadata = RequestResponseMetadata(request_id=request.request_id) if raw_request: @@ -466,6 +443,9 @@ async def _create_responses( engine_input=engine_input, sampling_params=sampling_params, context=context, + tok_params=( + request.build_tok_params(self.model_config) if isinstance(context, HarmonyContext) else None + ), lora_request=lora_request, priority=self._get_priority(request, raw_request), trace_headers=trace_headers, @@ -555,58 +535,28 @@ async def _create_responses( request_metadata, ) - async def _make_request( - self, - request: ResponsesRequest, - prev_response: ResponsesResponse | None, - ): - tool_dicts = construct_tool_dicts( - request.tools, - request.tool_choice, - exclude_tools_when_tool_choice_none=(self.online_renderer.exclude_tools_when_tool_choice_none), - ) - # Construct the input messages. - messages = construct_input_messages( - request_instructions=request.instructions, - request_input=request.input, - prev_msg=self.msg_store.get(prev_response.id) if prev_response else None, - prev_response_output=prev_response.output if prev_response else None, - ) - chat_template_kwargs = self._effective_chat_template_kwargs(request) - _, engine_inputs = await self.online_renderer.preprocess_chat( - request, - messages, - default_template=self.chat_template, - default_template_content_format=self.chat_template_content_format, - default_template_kwargs=chat_template_kwargs, - tool_dicts=tool_dicts, - parser=self.parser, - ) - return messages, engine_inputs - async def _render_next_turn( self, request: ResponsesRequest, messages: list[ResponseInputOutputItem], - tool_dicts: list[dict[str, Any]] | None, - parser: type[Parser] | None, - chat_template: str | None, - chat_template_content_format: ChatTemplateContentFormatOption, - ): - new_messages = construct_input_messages( - request_input=messages, + ) -> list[EngineInput]: + render_request = request.model_copy( + update={ + "input": list(messages), + "instructions": None, + "previous_input_messages": None, + "previous_response_id": None, + } ) - chat_template_kwargs = self._effective_chat_template_kwargs(request) - _, engine_inputs = await self.online_renderer.preprocess_chat( - request, - new_messages, - default_template=chat_template, - default_template_content_format=chat_template_content_format, - default_template_kwargs=chat_template_kwargs, - tool_dicts=tool_dicts, - parser=parser, + render_result = await self.online_renderer.render_responses( + render_request, + previous_messages=None, + previous_response_outputs=None, + tool_server=self.tool_server, ) - return engine_inputs + if isinstance(render_result, ErrorResponse): + raise APHRODITEValidationError(render_result.error.message) + return [render_result.engine_input] async def _generate_with_builtin_tools( self, @@ -614,6 +564,7 @@ async def _generate_with_builtin_tools( engine_input: EngineInput, sampling_params: SamplingParams, context: ConversationContext, + tok_params: TokenizeParams | None = None, lora_request: LoRARequest | None = None, priority: int = 0, trace_headers: Mapping[str, str] | None = None, @@ -621,6 +572,7 @@ async def _generate_with_builtin_tools( reasoning_parser_kwargs: dict[str, Any] | None = None, ): max_model_len = self.model_config.max_model_len + cache_salt = cast(str | None, engine_input.get("cache_salt")) orig_priority = priority sub_request = 0 @@ -665,18 +617,17 @@ async def _generate_with_builtin_tools( # Create inputs for the next turn. # Render the next prompt token ids and update sampling_params. if isinstance(context, HarmonyContext): - token_ids = context.render_for_completion() - engine_input = tokens_input(token_ids) + engine_input = self.online_renderer.render_responses_harmony_messages( + context.messages, + cache_salt=cache_salt, + tok_params=tok_params, + ) - sampling_params.max_tokens = max_model_len - len(token_ids) + sampling_params.max_tokens = max_model_len - self._extract_prompt_len(engine_input) elif isinstance(context, ParsableContext): (engine_input,) = await self._render_next_turn( context.request, context.response_messages, - context.tool_dicts, - context.parser_cls, - context.chat_template, - context.chat_template_content_format, ) sampling_params.max_tokens = get_max_tokens( @@ -692,28 +643,6 @@ async def _generate_with_builtin_tools( priority = orig_priority - 1 sub_request += 1 - def _make_request_with_harmony( - self, - request: ResponsesRequest, - prev_response: ResponsesResponse | None, - ): - if self.parser is not None: - # HarmonyParser doesn't need chat_template_kwargs - # TODO: Unify adjust_request() call with non-harmony branch - self.parser( - self.renderer.get_tokenizer(), - request.tools, - model_config=self.model_config, - ).adjust_request(request=request) - - arrival_time = time.time() - messages = self._construct_input_messages_with_harmony(request, prev_response) - prompt_token_ids = render_for_completion(messages) - engine_input = tokens_input(prompt_token_ids, cache_salt=request.cache_salt) - engine_input["arrival_time"] = arrival_time - - return messages, [engine_input] - async def _initialize_tool_sessions( self, request: ResponsesRequest, @@ -1036,97 +965,6 @@ def _make_response_output_items( ) ] - def _get_harmony_builtin_tool_descriptions( - self, request: ResponsesRequest, tool_types: set[str] - ) -> dict[str, str | None]: - # Extract allowed_tools from MCP tool requests - allowed_tools_map = _extract_allowed_tools_from_mcp_requests(request.tools) - - # Get filtered tool descriptions first. - # If get_tool_description returns None (due to filtering), the tool is disabled. - browser_description = ( - self.tool_server.get_tool_description("browser", allowed_tools_map.get("web_search_preview")) - if "web_search_preview" in tool_types - and self.tool_server is not None - and self.tool_server.has_tool("browser") - else None - ) - python_description = ( - self.tool_server.get_tool_description("python", allowed_tools_map.get("code_interpreter")) - if "code_interpreter" in tool_types and self.tool_server is not None and self.tool_server.has_tool("python") - else None - ) - container_description = ( - self.tool_server.get_tool_description("container", allowed_tools_map.get("container")) - if "container" in tool_types and self.tool_server is not None and self.tool_server.has_tool("container") - else None - ) - return { - "browser_description": browser_description, - "python_description": python_description, - "container_description": container_description, - } - - def _construct_input_messages_with_harmony( - self, - request: ResponsesRequest, - prev_response: ResponsesResponse | None, - ) -> list[OpenAIHarmonyMessage]: - messages: list[OpenAIHarmonyMessage] = [] - request_input = request.input - if prev_response is None: - # New conversation. - tool_types = extract_tool_types(request.tools) - with_custom_tools = has_custom_tools(tool_types) - instructions = request.instructions - if instructions is None and isinstance(request_input, list): - instructions, request_input = extract_instructions_from_messages(request_input) - tool_descriptions = self._get_harmony_builtin_tool_descriptions(request, tool_types) - tools = request.tools if with_custom_tools else None - messages.extend( - build_harmony_preamble( - instructions=instructions, - tools=tools, - reasoning_effort=(request.reasoning.effort if request.reasoning else None), - with_custom_tools=with_custom_tools, - **tool_descriptions, - ) - ) - messages += construct_harmony_previous_input_messages(request) - - else: - # Continue the previous conversation. - # FIXME(woosuk): Currently, request params like reasoning and - # instructions are ignored. - prev_msgs = self.msg_store[prev_response.id] - - messages.extend(prev_msgs) - # Append the new input. - # Responses API supports simple text inputs without chat format. - if isinstance(request_input, str): - # Skip empty string input when previous_input_messages supplies - # the full conversation history --- an empty trailing user message - # confuses the model into thinking nothing was sent. - if request_input or not request.previous_input_messages: - messages.append(get_user_message(request_input)) - else: - if prev_response is not None: - prev_outputs = copy(prev_response.output) - else: - prev_outputs = [] - for response_msg in request_input: - new_msg = response_input_to_harmony(response_msg, prev_outputs) - if new_msg is not None: - messages.append(new_msg) - - # User passes in a tool call request and its output. We need - # to add the tool call request to prev_outputs so that - # response_input_to_harmony can find the tool call request when - # parsing the tool call output. - if isinstance(response_msg, ResponseFunctionToolCall): - prev_outputs.append(response_msg) - return messages - async def _run_background_request_stream( self, request: ResponsesRequest, diff --git a/aphrodite/entrypoints/scale_out/factories.py b/aphrodite/entrypoints/scale_out/factories.py index b880ec15d2..efbb9d24a4 100644 --- a/aphrodite/entrypoints/scale_out/factories.py +++ b/aphrodite/entrypoints/scale_out/factories.py @@ -31,6 +31,7 @@ def init_render_state( state.openai_serving_models, state.online_renderer, request_logger=request_logger, + tool_server=state.tool_server, ) state.serving_derender = ServingDerender( diff --git a/aphrodite/entrypoints/scale_out/render/api_router.py b/aphrodite/entrypoints/scale_out/render/api_router.py index 1ab59e0a2c..ac1cf98469 100644 --- a/aphrodite/entrypoints/scale_out/render/api_router.py +++ b/aphrodite/entrypoints/scale_out/render/api_router.py @@ -8,6 +8,7 @@ from aphrodite.entrypoints.anthropic.protocol import AnthropicMessagesRequest from aphrodite.entrypoints.openai.chat_completion.protocol import ChatCompletionRequest from aphrodite.entrypoints.openai.completion.protocol import CompletionRequest +from aphrodite.entrypoints.openai.responses.protocol import ResponsesRequest from aphrodite.entrypoints.serve.engine.protocol import ErrorResponse from aphrodite.entrypoints.serve.utils.api_utils import validate_json_request from aphrodite.logger import init_logger @@ -93,3 +94,25 @@ async def render_completion(request: CompletionRequest, raw_request: Request): return JSONResponse(content=result.model_dump(), status_code=result.error.code) return JSONResponse(content=[item.model_dump() for item in result]) + + +@router.post( + "/v1/responses/render", + dependencies=[Depends(validate_json_request)], + response_model=GenerateRequest, + responses={ + HTTPStatus.BAD_REQUEST.value: {"model": ErrorResponse}, + HTTPStatus.NOT_FOUND.value: {"model": ErrorResponse}, + HTTPStatus.NOT_IMPLEMENTED.value: {"model": ErrorResponse}, + HTTPStatus.INTERNAL_SERVER_ERROR.value: {"model": ErrorResponse}, + }, +) +async def render_responses(request: ResponsesRequest, raw_request: Request): + handler = render(raw_request) + if handler is None: + raise NotImplementedError("The model does not support Responses Render API") + + result = await handler.render_responses_request(request) + if isinstance(result, ErrorResponse): + return JSONResponse(content=result.model_dump(), status_code=result.error.code) + return JSONResponse(content=result.model_dump()) diff --git a/aphrodite/entrypoints/scale_out/render/serving.py b/aphrodite/entrypoints/scale_out/render/serving.py index a65b867980..df86cc60c6 100644 --- a/aphrodite/entrypoints/scale_out/render/serving.py +++ b/aphrodite/entrypoints/scale_out/render/serving.py @@ -1,5 +1,7 @@ # SPDX-License-Identifier: Apache-2.0 # SPDX-FileCopyrightText: Copyright contributors to the vLLM project +from typing import TYPE_CHECKING + from aphrodite.entrypoints.anthropic.protocol import AnthropicMessagesRequest from aphrodite.entrypoints.anthropic.serving import AnthropicServingMessages from aphrodite.entrypoints.openai.chat_completion.protocol import ChatCompletionRequest @@ -8,6 +10,7 @@ OpenAIModelRegistry, OpenAIServingModels, ) +from aphrodite.entrypoints.openai.responses.protocol import ResponsesRequest from aphrodite.entrypoints.scale_out.token_in_token_out.mm_features import ( extract_mm_features, ) @@ -31,6 +34,9 @@ logger = init_logger(__name__) +if TYPE_CHECKING: + from aphrodite.entrypoints.mcp.tool_server import ToolServer + class ServingRender(BaseServing): def __init__( @@ -39,6 +45,7 @@ def __init__( online_renderer: "OnlineRenderer", *, request_logger: RequestLogger | None = None, + tool_server: "ToolServer | None" = None, ) -> None: super().__init__( models=models, @@ -47,6 +54,7 @@ def __init__( ) self.online_renderer = online_renderer + self.tool_server = tool_server self._merge_inline_system = AnthropicServingMessages._detect_merge_inline_system(online_renderer.chat_template) @@ -207,6 +215,61 @@ async def render_completion_request( return generate_requests + async def render_responses_request( + self, + request: ResponsesRequest, + ) -> GenerateRequest | ErrorResponse: + error_check_ret = await self._check_model(request) + if error_check_ret is not None: + return error_check_ret + if request.previous_response_id is not None: + return self.create_error_response( + message=("previous_response_id is not supported by the stateless render endpoint."), + err_type="invalid_request_error", + param="previous_response_id", + ) + + result = await self.online_renderer.render_responses( + request, + previous_messages=None, + previous_response_outputs=None, + tool_server=self.tool_server, + skip_mm_cache=True, + ) + if isinstance(result, ErrorResponse): + return result + + engine_input = result.engine_input + prompt_components = extract_prompt_components(self.model_config, engine_input) + token_ids = prompt_components.token_ids + if not token_ids: + return self.create_error_response("No token_ids rendered") + + input_length = extract_prompt_len(self.model_config, engine_input) + max_tokens = get_max_tokens( + self.model_config.max_model_len, + request.max_output_tokens, + input_length, + self.default_sampling_params, + self.override_max_tokens, + truncate_prompt_tokens=(-1 if request.truncation != "disabled" else None), + ) + params = request.to_sampling_params(max_tokens, self.default_sampling_params) + + return GenerateRequest( + request_id=request.request_id, + token_ids=list(token_ids), + features=self._extract_mm_features(engine_input), + sampling_params=params, + model=request.model, + stream=bool(request.stream), + cache_salt=request.cache_salt, + priority=request.priority, + kv_transfer_params=request.kv_transfer_params, + ec_transfer_params=request.ec_transfer_params, + token_offsets=engine_input.get("prompt_token_offsets"), + ) + def _placeholder_metadata_fields(self, modality: str) -> set[str]: """Fields the EC consumer still requires after embeddings transfer.""" if self._placeholder_metadata_parser_failed: diff --git a/aphrodite/renderers/online_renderer.py b/aphrodite/renderers/online_renderer.py index c5235a677e..af418ca7ee 100644 --- a/aphrodite/renderers/online_renderer.py +++ b/aphrodite/renderers/online_renderer.py @@ -1,13 +1,19 @@ # SPDX-License-Identifier: Apache-2.0 # SPDX-FileCopyrightText: Copyright contributors to the vLLM project +import time from collections.abc import Sequence +from dataclasses import dataclass from http import HTTPStatus -from typing import Any +from typing import TYPE_CHECKING, Any, TypeAlias +from openai.types.responses import ResponseFunctionToolCall, ResponseOutputItem +from openai.types.responses.tool import Mcp, Tool from openai_harmony import Message as OpenAIMessage +from openai_harmony import ToolNamespaceConfig from aphrodite.config import ModelConfig from aphrodite.entrypoints.chat_utils import ( + ChatCompletionMessageParam, ChatTemplateContentFormatOption, ConversationMessage, ) @@ -19,12 +25,24 @@ CompletionRequest, ) from aphrodite.entrypoints.openai.parser.harmony_utils import ( + BUILTIN_TOOL_TO_MCP_SERVER_LABEL, build_harmony_preamble, extract_instructions_from_messages, + get_user_message, + has_custom_tools, parse_chat_inputs_to_harmony_messages, render_for_completion, ) +from aphrodite.entrypoints.openai.responses.harmony import ( + construct_harmony_previous_input_messages, + response_input_to_harmony, +) from aphrodite.entrypoints.openai.responses.protocol import ResponsesRequest +from aphrodite.entrypoints.openai.responses.utils import ( + construct_input_messages, + construct_tool_dicts, + extract_tool_types, +) from aphrodite.entrypoints.serve import create_error_response from aphrodite.entrypoints.serve.engine.protocol import ErrorResponse from aphrodite.entrypoints.serve.utils.request_logger import RequestLogger @@ -32,11 +50,12 @@ EngineInput, PromptType, SingletonPrompt, + TokensPrompt, tokens_input, ) from aphrodite.logger import init_logger from aphrodite.parser import Parser, ParserManager -from aphrodite.renderers import BaseRenderer, ChatParams, merge_kwargs +from aphrodite.renderers import BaseRenderer, ChatParams, TokenizeParams, merge_kwargs from aphrodite.renderers.inputs.preprocess import ( parse_model_prompt, prompt_to_seq, @@ -46,6 +65,37 @@ logger = init_logger(__name__) +if TYPE_CHECKING: + from aphrodite.entrypoints.mcp.tool_server import ToolServer + + +ResponsesPreviousMessages: TypeAlias = list[ChatCompletionMessageParam] | list[OpenAIMessage] + + +@dataclass(frozen=True) +class ResponsesRenderResult: + messages: ResponsesPreviousMessages + engine_input: EngineInput + + +def _extract_allowed_tools_from_mcp_requests( + tools: list[Tool], +) -> dict[str, list[str] | None]: + allowed_tools_map: dict[str, list[str] | None] = {} + for tool in tools: + if not isinstance(tool, Mcp): + continue + + allowed_tools = None + if isinstance(tool.allowed_tools, list): + allowed_tools = tool.allowed_tools + elif tool.allowed_tools is not None: + allowed_tools = tool.allowed_tools.tool_names + if allowed_tools is not None and "*" in allowed_tools: + allowed_tools = None + allowed_tools_map[tool.server_label] = allowed_tools + return allowed_tools_map + def _reused_prompt_token_ids(request: Any) -> list[int] | None: """Pop prompt token ids forwarded for decode-side reuse, if any. @@ -203,6 +253,203 @@ async def render_chat( return conversation, engine_inputs + async def render_responses( + self, + request: ResponsesRequest, + *, + previous_messages: ResponsesPreviousMessages | None = None, + previous_response_outputs: list[ResponseOutputItem] | None = None, + tool_server: "ToolServer | None" = None, + skip_mm_cache: bool = False, + ) -> ResponsesRenderResult | ErrorResponse: + """Render a Responses request using only explicitly supplied history.""" + template_error = self.validate_chat_template( + request_chat_template=None, + chat_template_kwargs=request.chat_template_kwargs, + trust_request_chat_template=self.trust_request_chat_template, + ) + if template_error is not None: + return template_error + + if self.use_harmony: + return self._render_responses_with_harmony( + request, + previous_messages=previous_messages, + previous_response_outputs=previous_response_outputs, + tool_server=tool_server, + ) + + if previous_messages is not None and any(isinstance(message, OpenAIMessage) for message in previous_messages): + return self.create_error_response( + "Non-Harmony Responses history must use chat messages.", + err_type="invalid_request_error", + param="previous_response_id", + ) + + tool_dicts = construct_tool_dicts( + request.tools, + request.tool_choice, + exclude_tools_when_tool_choice_none=(self.exclude_tools_when_tool_choice_none), + ) + messages = construct_input_messages( + request_instructions=request.instructions, + request_input=request.input, + prev_msg=list(previous_messages) if previous_messages is not None else None, + prev_response_output=(list(previous_response_outputs) if previous_response_outputs is not None else None), + ) + chat_template_kwargs = ( + request.build_chat_params( + self.chat_template, + self.chat_template_content_format, + ) + .with_defaults(self.default_chat_template_kwargs) + .chat_template_kwargs + ) + _, engine_inputs = await self.preprocess_chat( + request, + messages, + default_template=self.chat_template, + default_template_content_format=self.chat_template_content_format, + default_template_kwargs=chat_template_kwargs, + tool_dicts=tool_dicts, + parser=self.parser, + skip_mm_cache=skip_mm_cache, + ) + return self._responses_render_result(messages, engine_inputs) + + def _render_responses_with_harmony( + self, + request: ResponsesRequest, + *, + previous_messages: ResponsesPreviousMessages | None, + previous_response_outputs: list[ResponseOutputItem] | None, + tool_server: "ToolServer | None", + ) -> ResponsesRenderResult | ErrorResponse: + if self.parser is not None: + # HarmonyParser doesn't need chat_template_kwargs + # TODO: Unify adjust_request() call with non-harmony branch + self.parser( + self.renderer.get_tokenizer(), + request.tools, + model_config=self.model_config, + ).adjust_request(request=request) + + if previous_messages is not None and any( + not isinstance(message, OpenAIMessage) for message in previous_messages + ): + return self.create_error_response( + "Harmony Responses history must use Harmony messages.", + err_type="invalid_request_error", + param="previous_response_id", + ) + + messages: list[OpenAIMessage] = [] + request_input = request.input + if previous_messages is None: + tool_types = extract_tool_types(request.tools) + with_custom_tools = has_custom_tools(tool_types) + instructions = request.instructions + if instructions is None and isinstance(request_input, list): + instructions, request_input = extract_instructions_from_messages(request_input) + descriptions = self._get_harmony_builtin_tool_descriptions( + request.tools, + tool_types, + tool_server, + ) + messages.extend( + build_harmony_preamble( + instructions=instructions, + tools=request.tools if with_custom_tools else None, + reasoning_effort=(request.reasoning.effort if request.reasoning else None), + with_custom_tools=with_custom_tools, + **descriptions, + ) + ) + messages.extend(construct_harmony_previous_input_messages(request)) + else: + messages.extend(previous_messages) + + previous_outputs = list(previous_response_outputs or ()) + try: + if isinstance(request_input, str): + if request_input or not request.previous_input_messages: + messages.append(get_user_message(request_input)) + else: + for response_message in request_input: + new_message = response_input_to_harmony( + response_message, + previous_outputs, + ) + if new_message is not None: + messages.append(new_message) + if isinstance(response_message, ResponseFunctionToolCall): + previous_outputs.append(response_message) + except ValueError as exc: + return self.create_error_response( + str(exc), + err_type="invalid_request_error", + param="input", + ) + + return ResponsesRenderResult( + messages=messages, + engine_input=self.render_responses_harmony_messages( + messages, + cache_salt=request.cache_salt, + tok_params=request.build_tok_params(self.model_config), + ), + ) + + def _get_harmony_builtin_tool_descriptions( + self, + tools: list[Tool], + tool_types: set[str], + tool_server: "ToolServer | None", + ) -> dict[str, ToolNamespaceConfig | None]: + allowed_tools = _extract_allowed_tools_from_mcp_requests(tools) + descriptions: dict[str, ToolNamespaceConfig | None] = {} + for server_name, request_name in BUILTIN_TOOL_TO_MCP_SERVER_LABEL.items(): + description = ( + tool_server.get_tool_description( + server_name, + allowed_tools.get(request_name), + ) + if request_name in tool_types and tool_server is not None and tool_server.has_tool(server_name) + else None + ) + descriptions[f"{server_name}_description"] = description + return descriptions + + def render_responses_harmony_messages( + self, + messages: list[OpenAIMessage], + *, + cache_salt: str | None, + tok_params: TokenizeParams | None = None, + ) -> EngineInput: + arrival_time = time.time() + prompt = TokensPrompt(prompt_token_ids=render_for_completion(messages)) + if tok_params is not None: + tok_params.apply_post_tokenization( + self.renderer.tokenizer, + prompt, + ) + engine_input = tokens_input(prompt["prompt_token_ids"], cache_salt=cache_salt) + engine_input["arrival_time"] = arrival_time + return engine_input + + def _responses_render_result( + self, + messages: list[ChatCompletionMessageParam], + engine_inputs: list[EngineInput], + ) -> ResponsesRenderResult | ErrorResponse: + if len(engine_inputs) != 1: + return self.create_error_response(f"Expected exactly 1 engine prompt, got {len(engine_inputs)}") + return ResponsesRenderResult( + messages=messages, + engine_input=engine_inputs[0], + ) + def _make_request_with_harmony( self, request: ChatCompletionRequest, diff --git a/docs/src/content/docs/deployment/security.md b/docs/src/content/docs/deployment/security.md index b4d8ff6ea1..d5999c7092 100644 --- a/docs/src/content/docs/deployment/security.md +++ b/docs/src/content/docs/deployment/security.md @@ -15,6 +15,11 @@ Do not expose diagnostic or administrative routes to the public network. Restrict `/metrics`, profiling routes, adapter management, cache reset, sleep, weight update, and collective RPC routes. +`/v1/responses/render` requires bearer authentication when `--api-key` is +configured. It is available on `aphrodite serve` with +`APHRODITE_ENABLE_SCALE_OUT_ENDPOINTS=1`, or on `aphrodite launch render` +unless explicitly disabled. + ## Pin model revisions Use `--revision` and `--tokenizer-revision` with immutable commit identifiers. diff --git a/docs/src/content/docs/serving/other-apis.md b/docs/src/content/docs/serving/other-apis.md index a13f3d3ac7..72205606b0 100644 --- a/docs/src/content/docs/serving/other-apis.md +++ b/docs/src/content/docs/serving/other-apis.md @@ -36,8 +36,8 @@ The render API performs request preprocessing without model execution. The derender API converts generated token IDs into a response. These endpoints help separate CPU preprocessing from GPU inference. -OpenAI-compatible render routes include `/v1/chat/completions/render` and -`/v1/completions/render`. Sonar also provides +OpenAI-compatible render routes include `/v1/chat/completions/render`, +`/v1/completions/render`, and `/v1/responses/render`. Sonar also provides `/inference/v1/generate` for token-input and token-output workflows. These APIs are useful when a gateway performs tokenization or when CPU preprocessing runs separately from GPU workers. @@ -45,6 +45,31 @@ separately from GPU workers. Do not assume token IDs are portable between models. The tokenizer, added tokens, chat template, and model revision must match. +### Responses rendering + +Enable render endpoints with `APHRODITE_ENABLE_SCALE_OUT_ENDPOINTS=1` on +`aphrodite serve`, or use `aphrodite launch render`. + +`/v1/responses/render` uses the same prompt construction as `/v1/responses` +and returns one token-input `GenerateRequest`. It is stateless: inline history +is supported, but `previous_response_id` is not. Resolve stored response state +and include the resulting history in the request before rendering. + +For multimodal requests, the result includes the model-processed payload, +which can be substantially larger than the source image or video. Forward it +unchanged to the generation service and provision transport limits and memory +accordingly. + +```bash +curl http://localhost:2242/v1/responses/render \ + -H "Content-Type: application/json" \ + -d '{ + "model": "meta-llama/Llama-3.1-8B-Instruct", + "input": "Explain prefix caching in one sentence.", + "max_output_tokens": 32 + }' +``` + ### Multimodal Render Features Multimodal render responses include a `features` object with per-modality diff --git a/tests/entrypoints/launchers/test_cli_args.py b/tests/entrypoints/launchers/test_cli_args.py index 7f3020c757..9b8a6978c8 100644 --- a/tests/entrypoints/launchers/test_cli_args.py +++ b/tests/entrypoints/launchers/test_cli_args.py @@ -2,9 +2,11 @@ # SPDX-FileCopyrightText: Copyright contributors to the vLLM project import json +from types import SimpleNamespace import pytest +import aphrodite.entrypoints.launchers.cli_args as cli_args_module from aphrodite.entrypoints.launchers.cli_args import ( make_arg_parser, validate_parsed_serve_args, @@ -316,6 +318,26 @@ def test_default_chat_template_kwargs_default_none(serve_parser): assert args.default_chat_template_kwargs is None +@pytest.mark.parametrize( + "default_kwargs,cohere_format,expected", + [ + (None, "cmd3", {"cohere_format": "cmd3"}), + ( + {"cohere_format": "cmd4", "custom": True}, + "cmd3", + {"cohere_format": "cmd4", "custom": True}, + ), + ], +) +def test_resolve_default_chat_template_kwargs(default_kwargs, cohere_format, expected): + args = SimpleNamespace( + default_chat_template_kwargs=default_kwargs, + cohere_format=cohere_format, + ) + + assert cli_args_module.resolve_default_chat_template_kwargs(args) == expected + + def test_default_chat_template_kwargs_invalid_json(serve_parser): """Ensure invalid JSON raises an error""" with pytest.raises(SystemExit): diff --git a/tests/entrypoints/openai/responses/test_mcp_tools.py b/tests/entrypoints/openai/responses/test_mcp_tools.py index ba356f517d..c0bc3048f7 100644 --- a/tests/entrypoints/openai/responses/test_mcp_tools.py +++ b/tests/entrypoints/openai/responses/test_mcp_tools.py @@ -4,13 +4,18 @@ from __future__ import annotations +from types import SimpleNamespace +from unittest.mock import AsyncMock, MagicMock + import pytest import pytest_asyncio from openai import OpenAI from openai_harmony import Message, ToolDescription, ToolNamespaceConfig +import aphrodite.entrypoints.mcp.tool_server as tool_server_module from aphrodite.entrypoints.mcp.tool_server import ( MCPToolServer, + init_tool_server, post_process_tools_description, ) from tests.utils import RemoteOpenAIServer @@ -37,6 +42,25 @@ _PYTHON_TOOL_INSTRUCTION = "You must use the Python tool to execute code. Never simulate execution." +@pytest.mark.asyncio +async def test_init_tool_server_without_configuration(): + args = SimpleNamespace(tool_server=None) + + assert await init_tool_server(args) is None + + +@pytest.mark.asyncio +async def test_init_tool_server_with_mcp_url(monkeypatch: pytest.MonkeyPatch): + server = MagicMock() + server.add_tool_server = AsyncMock() + monkeypatch.setattr(tool_server_module, "MCPToolServer", lambda: server) + + result = await init_tool_server(SimpleNamespace(tool_server="http://mcp.test")) + + assert result is server + server.add_tool_server.assert_awaited_once_with("http://mcp.test") + + @pytest.mark.skip_global_cleanup class TestMCPToolServerUnit: """Test MCP tool description processing and filtering. diff --git a/tests/entrypoints/openai/responses/test_serving_responses.py b/tests/entrypoints/openai/responses/test_serving_responses.py index f595c7704f..96d6f25b99 100644 --- a/tests/entrypoints/openai/responses/test_serving_responses.py +++ b/tests/entrypoints/openai/responses/test_serving_responses.py @@ -2,12 +2,16 @@ # SPDX-FileCopyrightText: Copyright contributors to the vLLM project from contextlib import AsyncExitStack -from unittest.mock import MagicMock +from types import SimpleNamespace +from unittest.mock import AsyncMock, MagicMock import pytest import pytest_asyncio from openai.types.responses import ( + ResponseFunctionToolCall, ResponseOutputItemDoneEvent, + ResponseOutputMessage, + ResponseOutputText, ResponseReasoningItem, ResponseReasoningTextDeltaEvent, ResponseReasoningTextDoneEvent, @@ -23,6 +27,8 @@ Mcp, Tool, ) +from openai_harmony import Message as OpenAIHarmonyMessage +from openai_harmony import Role import aphrodite.envs as envs from aphrodite.entrypoints.generate.base.protocol import ( @@ -32,7 +38,11 @@ RequestResponseMetadata, ) from aphrodite.entrypoints.mcp.tool_server import ToolServer -from aphrodite.entrypoints.openai.responses.context import ConversationContext, SimpleContext +from aphrodite.entrypoints.openai.responses.context import ( + ConversationContext, + HarmonyContext, + SimpleContext, +) from aphrodite.entrypoints.openai.responses.protocol import ( ResponseCreatedEvent, ResponseRawMessageAndToken, @@ -42,7 +52,6 @@ ) from aphrodite.entrypoints.openai.responses.serving import ( OpenAIServingResponses, - _extract_allowed_tools_from_mcp_requests, extract_tool_types, ) from aphrodite.entrypoints.openai.responses.streaming_events import ( @@ -52,8 +61,23 @@ from aphrodite.inputs import tokens_input from aphrodite.outputs import CompletionOutput, RequestOutput from aphrodite.parser.harmony import Segment +from aphrodite.renderers import TokenizeParams +from aphrodite.renderers.online_renderer import ( + OnlineRenderer, + _extract_allowed_tools_from_mcp_requests, +) from aphrodite.sampling_params import SamplingParams +pytestmark = pytest.mark.skip_global_cleanup + + +def _new_online_renderer() -> OnlineRenderer: + renderer = OnlineRenderer.__new__(OnlineRenderer) + renderer.exclude_tools_when_tool_choice_none = False + renderer.parser = None + renderer.trust_request_chat_template = True + return renderer + class MockConversationContext(ConversationContext): """Mock conversation context for testing""" @@ -203,6 +227,317 @@ def test_response_created_event_uses_public_json_schema_alias() -> None: assert event.response.text.format.model_dump(by_alias=True)["schema"] == schema +@pytest.mark.asyncio +async def test_online_renderer_renders_non_harmony_responses_with_explicit_history(): + renderer = _new_online_renderer() + renderer.use_harmony = False + renderer.chat_template = "server-template" + renderer.chat_template_content_format = "string" + renderer.default_chat_template_kwargs = {"server_default": "kept"} + renderer.parser = None + renderer.preprocess_chat = AsyncMock( + return_value=( + [], + [tokens_input([11, 12, 13], cache_salt="request-salt")], + ) + ) + previous_messages = [ + {"role": "system", "content": "old instructions"}, + {"role": "user", "content": "first"}, + ] + previous_outputs = [ + ResponseOutputMessage( + id="msg_1", + content=[ + ResponseOutputText( + annotations=[], + text="prior answer", + type="output_text", + ) + ], + role="assistant", + status="completed", + type="message", + ) + ] + request = ResponsesRequest( + input="next", + instructions="new instructions", + tools=[], + cache_salt="request-salt", + chat_template_kwargs={"request_default": "kept"}, + ) + + result = await renderer.render_responses( + request, + previous_messages=previous_messages, + previous_response_outputs=previous_outputs, + ) + + assert not isinstance(result, ErrorResponse) + assert result.messages == [ + {"role": "system", "content": "new instructions"}, + {"role": "user", "content": "first"}, + {"role": "assistant", "content": "prior answer"}, + {"role": "user", "content": "next"}, + ] + assert result.engine_input["prompt_token_ids"] == [11, 12, 13] + assert result.engine_input["cache_salt"] == "request-salt" + assert previous_messages[0]["content"] == "old instructions" + assert previous_outputs[0].content[0].text == "prior answer" + + preprocess_kwargs = renderer.preprocess_chat.call_args.kwargs + assert preprocess_kwargs["default_template"] == "server-template" + assert preprocess_kwargs["default_template_content_format"] == "string" + assert preprocess_kwargs["default_template_kwargs"]["server_default"] == "kept" + assert preprocess_kwargs["default_template_kwargs"]["request_default"] == "kept" + + +@pytest.mark.asyncio +async def test_online_renderer_rejects_untrusted_responses_chat_template(): + renderer = _new_online_renderer() + renderer.use_harmony = False + renderer.chat_template = "server-template" + renderer.chat_template_content_format = "string" + renderer.default_chat_template_kwargs = {} + renderer.exclude_tools_when_tool_choice_none = False + renderer.trust_request_chat_template = False + renderer.parser = None + renderer.preprocess_chat = AsyncMock( + return_value=([], [tokens_input([1])]), + ) + request = ResponsesRequest( + input="hello", + chat_template_kwargs={"chat_template": "{{ messages }}"}, + ) + + result = await renderer.render_responses(request) + + assert isinstance(result, ErrorResponse) + assert result.error.code == 400 + assert "untrusted chat template" in result.error.message.lower() + renderer.preprocess_chat.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_online_renderer_rejects_harmony_history_for_non_harmony_model(): + renderer = _new_online_renderer() + renderer.use_harmony = False + request = ResponsesRequest(input="next") + + result = await renderer.render_responses( + request, + previous_messages=[OpenAIHarmonyMessage.from_role_and_content(Role.USER, "first")], + ) + + assert isinstance(result, ErrorResponse) + assert result.error.param == "previous_response_id" + + +@pytest.mark.asyncio +async def test_online_renderer_rejects_chat_history_for_harmony_model(): + renderer = _new_online_renderer() + renderer.use_harmony = True + request = ResponsesRequest(input="next") + + result = await renderer.render_responses( + request, + previous_messages=[{"role": "user", "content": "first"}], + ) + + assert isinstance(result, ErrorResponse) + assert result.error.param == "previous_response_id" + + +@pytest.mark.asyncio +async def test_online_renderer_adjusts_harmony_tool_choice( + monkeypatch: pytest.MonkeyPatch, +): + class StubParser: + def __init__(self, tokenizer, tools, *, model_config): + pass + + def adjust_request(self, request): + request.cache_salt = "adjusted" + return request + + renderer = _new_online_renderer() + renderer.use_harmony = True + renderer.parser = StubParser + renderer.model_config = SimpleNamespace(max_model_len=100) + tokenizer = SimpleNamespace(truncation_side="left") + renderer.renderer = SimpleNamespace( + tokenizer=tokenizer, + get_tokenizer=lambda: tokenizer, + ) + monkeypatch.setattr( + "aphrodite.renderers.online_renderer.render_for_completion", + lambda messages: [1], + ) + request = ResponsesRequest( + input="next", + tools=[ + { + "type": "function", + "name": "lookup", + "parameters": {"type": "object", "properties": {}}, + } + ], + tool_choice={"type": "function", "name": "lookup"}, + ) + + result = await renderer.render_responses(request) + + assert not isinstance(result, ErrorResponse) + assert result.engine_input["cache_salt"] == "adjusted" + + +@pytest.mark.asyncio +async def test_online_renderer_applies_responses_token_budget_to_harmony_prompt( + monkeypatch: pytest.MonkeyPatch, +): + renderer = _new_online_renderer() + renderer.use_harmony = True + renderer.model_config = SimpleNamespace(max_model_len=5) + renderer.renderer = SimpleNamespace(tokenizer=SimpleNamespace(truncation_side="left")) + monkeypatch.setattr( + "aphrodite.renderers.online_renderer.render_for_completion", + lambda messages: [1, 2, 3, 4, 5, 6], + ) + request = ResponsesRequest( + input="hello", + max_output_tokens=2, + truncation="auto", + ) + + result = await renderer.render_responses(request) + + assert not isinstance(result, ErrorResponse) + assert result.engine_input["prompt_token_ids"] == [4, 5, 6] + + +@pytest.mark.asyncio +async def test_online_renderer_renders_harmony_continuation_with_call_linkage( + monkeypatch: pytest.MonkeyPatch, +): + renderer = _new_online_renderer() + renderer.use_harmony = True + renderer.model_config = SimpleNamespace(max_model_len=100) + renderer.renderer = SimpleNamespace(tokenizer=SimpleNamespace(truncation_side="left")) + monkeypatch.setattr( + "aphrodite.renderers.online_renderer.render_for_completion", + lambda messages: [1], + ) + previous_messages = [ + OpenAIHarmonyMessage.from_role_and_content(Role.USER, "first"), + ] + previous_outputs = [ + ResponseFunctionToolCall( + arguments='{"city":"Paris"}', + call_id="call_1", + name="lookup_weather", + type="function_call", + ) + ] + request = ResponsesRequest( + input=[ + { + "type": "function_call_output", + "call_id": "call_1", + "output": "sunny", + } + ], + instructions="replacement instructions are ignored on continuations", + reasoning={"effort": "high"}, + tools=[ + { + "type": "function", + "name": "lookup_weather", + "parameters": {"type": "object", "properties": {}}, + } + ], + cache_salt="request-salt", + ) + + result = await renderer.render_responses( + request, + previous_messages=previous_messages, + previous_response_outputs=previous_outputs, + ) + + assert not isinstance(result, ErrorResponse) + assert result.messages[0] is previous_messages[0] + tool_output = result.messages[-1] + assert tool_output.author.role == "tool" + assert tool_output.author.name == "functions.lookup_weather" + assert tool_output.channel == "commentary" + assert tool_output.recipient == "assistant" + assert result.engine_input["cache_salt"] == "request-salt" + assert "arrival_time" in result.engine_input + assert previous_messages == [ + OpenAIHarmonyMessage.from_role_and_content(Role.USER, "first"), + ] + assert previous_outputs[0].call_id == "call_1" + + +@pytest.mark.asyncio +async def test_online_renderer_returns_typed_error_for_bad_harmony_call_reference(): + renderer = _new_online_renderer() + renderer.use_harmony = True + request = ResponsesRequest( + input=[ + { + "type": "function_call_output", + "call_id": "missing", + "output": "orphaned", + } + ], + tools=[], + ) + + result = await renderer.render_responses(request) + + assert isinstance(result, ErrorResponse) + assert result.error.type == "invalid_request_error" + assert result.error.param == "input" + assert "No call message found for missing" in result.error.message + + +@pytest.mark.asyncio +async def test_online_renderer_preserves_missing_harmony_builtin_tool_behavior( + monkeypatch: pytest.MonkeyPatch, +): + renderer = _new_online_renderer() + renderer.use_harmony = True + renderer.model_config = SimpleNamespace(max_model_len=100) + renderer.renderer = SimpleNamespace(tokenizer=SimpleNamespace(truncation_side="left")) + monkeypatch.setattr( + "aphrodite.renderers.online_renderer.render_for_completion", + lambda messages: [1], + ) + request = ResponsesRequest( + input="search", + tools=[{"type": "web_search_preview"}], + ) + + result = await renderer.render_responses(request, tool_server=None) + + assert not isinstance(result, ErrorResponse) + + +@pytest.mark.parametrize( + "engine_inputs", + [[], [tokens_input([1]), tokens_input([2])]], +) +def test_responses_render_result_rejects_non_single_prompt(engine_inputs): + renderer = _new_online_renderer() + + result = renderer._responses_render_result([], engine_inputs) + + assert isinstance(result, ErrorResponse) + assert f"got {len(engine_inputs)}" in result.error.message + + class TestInitializeToolSessions: """Test class for _initialize_tool_sessions method""" @@ -238,6 +573,125 @@ async def serving_responses_instance(self): return instance + @pytest.mark.asyncio + async def test_stateful_harmony_delegates_named_tool_choice_to_renderer( + self, + serving_responses_instance, + ): + serving_responses_instance.use_harmony = True + serving_responses_instance.online_renderer.render_responses = AsyncMock( + return_value=SimpleNamespace( + messages=[], + engine_input=tokens_input([1]), + ) + ) + request = ResponsesRequest( + input="hello", + tool_choice={"type": "function", "name": "lookup"}, + tools=[ + { + "type": "function", + "name": "lookup", + "parameters": {"type": "object", "properties": {}}, + } + ], + ) + + result = await serving_responses_instance._render_resolved_response_inputs( + request, + None, + ) + + assert result.engine_input["prompt_token_ids"] == [1] + serving_responses_instance.online_renderer.render_responses.assert_awaited_once() + + @pytest.mark.asyncio + async def test_render_next_turn_reuses_shared_responses_renderer(self, serving_responses_instance): + serving_responses_instance.online_renderer.render_responses = AsyncMock( + return_value=SimpleNamespace( + messages=[], + engine_input=tokens_input([8, 9], cache_salt="request-salt"), + ) + ) + serving_responses_instance.online_renderer.preprocess_chat = AsyncMock(return_value=([], [tokens_input([1])])) + request = ResponsesRequest( + input="first", + instructions="initial instructions", + cache_salt="request-salt", + ) + turn_messages = [ + {"role": "user", "content": "first"}, + {"role": "assistant", "content": "tool call"}, + {"role": "tool", "content": "tool output"}, + ] + + engine_inputs = await serving_responses_instance._render_next_turn( + request, + turn_messages, + ) + + assert engine_inputs[0]["prompt_token_ids"] == [8, 9] + render_request = serving_responses_instance.online_renderer.render_responses.call_args.args[0] + assert render_request.input == turn_messages + assert render_request.instructions is None + assert render_request.cache_salt == "request-salt" + + @pytest.mark.asyncio + async def test_harmony_tool_followup_preserves_render_params( + self, serving_responses_instance, monkeypatch: pytest.MonkeyPatch + ): + class ToolCallingHarmonyContext(HarmonyContext): + def __init__(self): + self._messages = [] + self._needs_tool_call = True + + def append_output(self, output) -> None: + pass + + def need_builtin_tool_call(self) -> bool: + return self._needs_tool_call + + async def call_tool(self): + self._needs_tool_call = False + return [] + + def append_tool_output(self, output) -> None: + self._messages.extend(output) + + async def generate_output(): + yield MagicMock() + + context = ToolCallingHarmonyContext() + serving_responses_instance.engine_client.generate.side_effect = lambda *args, **kwargs: generate_output() + monkeypatch.setattr( + "aphrodite.renderers.online_renderer.render_for_completion", + lambda messages: [1, 2, 3, 4], + ) + online_renderer = serving_responses_instance.online_renderer + online_renderer.renderer = SimpleNamespace(tokenizer=SimpleNamespace(truncation_side="left")) + serving_responses_instance.online_renderer.render_responses_harmony_messages = ( + OnlineRenderer.render_responses_harmony_messages.__get__(online_renderer) + ) + serving_responses_instance._extract_prompt_len = MagicMock(return_value=2) + + async for _ in serving_responses_instance._generate_with_builtin_tools( + request_id="req", + engine_input=tokens_input([1], cache_salt="request-salt"), + sampling_params=MagicMock(), + context=context, + tok_params=TokenizeParams( + max_total_tokens=3, + max_output_tokens=1, + truncate_prompt_tokens=-1, + ), + ): + pass + + assert serving_responses_instance.engine_client.generate.call_count == 2 + followup_engine_input = serving_responses_instance.engine_client.generate.call_args_list[1].args[0] + assert followup_engine_input.get("cache_salt") == "request-salt" + assert followup_engine_input["prompt_token_ids"] == [3, 4] + @pytest.mark.asyncio async def test_initialize_tool_sessions(self, serving_responses_instance, mock_context, mock_exit_stack): """Test that method works correctly with only MCP tools""" diff --git a/tests/entrypoints/openai/test_render_parity.py b/tests/entrypoints/openai/test_render_parity.py index 395516ccf3..74db78b97b 100644 --- a/tests/entrypoints/openai/test_render_parity.py +++ b/tests/entrypoints/openai/test_render_parity.py @@ -151,7 +151,8 @@ async def _capture_responses( request: ResponsesRequest, ) -> CapturedRenderInputs: capture = RenderCapture(serving.online_renderer) - await serving._make_request(request, prev_response=None) + result = await serving.online_renderer.render_responses(request) + assert not isinstance(result, ErrorResponse), result return capture.take() diff --git a/tests/entrypoints/scale_out/render/test_render.py b/tests/entrypoints/scale_out/render/test_render.py index 8ed76c5a5a..e99d57eec1 100644 --- a/tests/entrypoints/scale_out/render/test_render.py +++ b/tests/entrypoints/scale_out/render/test_render.py @@ -3,15 +3,134 @@ """Tests for the /render endpoints that expose prompt preprocessing.""" +from http import HTTPStatus +from types import SimpleNamespace +from unittest.mock import AsyncMock, MagicMock + import httpx import pytest import pytest_asyncio +from aphrodite.entrypoints.openai.responses.protocol import ResponsesRequest +from aphrodite.entrypoints.scale_out.render.api_router import router +from aphrodite.entrypoints.scale_out.render.serving import ServingRender +from aphrodite.entrypoints.serve.engine.protocol import ErrorResponse +from aphrodite.renderers.online_renderer import OnlineRenderer from tests.utils import RemoteLaunchRenderServer MODEL_NAME = "hmellor/tiny-random-LlamaForCausalLM" +def _build_responses_serving_render() -> ServingRender: + serving = ServingRender.__new__(ServingRender) + serving.model_config = SimpleNamespace( + max_model_len=100, + is_encoder_decoder=False, + ) + serving.default_sampling_params = {} + serving.override_max_tokens = None + serving.tool_server = MagicMock() + serving.online_renderer = MagicMock() + serving.online_renderer.create_error_response = OnlineRenderer.create_error_response.__get__( + serving.online_renderer + ) + serving.online_renderer.validate_chat_template = OnlineRenderer.validate_chat_template.__get__( + serving.online_renderer + ) + serving.online_renderer.trust_request_chat_template = True + serving._check_model = AsyncMock(return_value=None) + return serving + + +@pytest.mark.skip_global_cleanup +def test_responses_render_route_is_registered(): + assert any(route.path == "/v1/responses/render" for route in router.routes) + + +@pytest.mark.skip_global_cleanup +def test_responses_render_route_documents_not_implemented(): + route = next(route for route in router.routes if route.path == "/v1/responses/render") + + assert HTTPStatus.NOT_IMPLEMENTED.value in route.responses + + +@pytest.mark.asyncio +@pytest.mark.skip_global_cleanup +async def test_render_responses_returns_generate_request_without_stored_state(): + serving = _build_responses_serving_render() + serving.online_renderer.render_responses = AsyncMock( + return_value=MagicMock( + messages=[{"role": "user", "content": "Test prompt"}], + engine_input={"prompt_token_ids": [7, 8, 9]}, + ) + ) + request = ResponsesRequest( + model=MODEL_NAME, + input="Test prompt", + max_output_tokens=12, + stream=True, + cache_salt="request-salt", + priority=3, + kv_transfer_params={"do_remote_prefill": True}, + ec_transfer_params={"remote": True}, + ) + + response = await serving.render_responses_request(request) + + assert response.token_ids == [7, 8, 9] + assert response.request_id == request.request_id + assert response.sampling_params.max_tokens == 12 + assert response.model == MODEL_NAME + assert response.stream is True + assert response.cache_salt == "request-salt" + assert response.priority == 3 + assert response.kv_transfer_params == {"do_remote_prefill": True} + assert response.ec_transfer_params == {"remote": True} + serving.online_renderer.render_responses.assert_awaited_once_with( + request, + previous_messages=None, + previous_response_outputs=None, + tool_server=serving.tool_server, + skip_mm_cache=True, + ) + + +@pytest.mark.asyncio +@pytest.mark.skip_global_cleanup +async def test_render_responses_rejects_previous_response_id(): + serving = _build_responses_serving_render() + serving.online_renderer.render_responses = AsyncMock() + request = ResponsesRequest( + model=MODEL_NAME, + input="Test prompt", + previous_response_id="resp_previous", + ) + + response = await serving.render_responses_request(request) + + assert isinstance(response, ErrorResponse) + assert response.error.code == 400 + assert response.error.param == "previous_response_id" + serving.online_renderer.render_responses.assert_not_awaited() + + +@pytest.mark.asyncio +@pytest.mark.skip_global_cleanup +async def test_render_responses_rejects_empty_token_ids(): + serving = _build_responses_serving_render() + serving.online_renderer.render_responses = AsyncMock( + return_value=MagicMock( + messages=[], + engine_input={"prompt_token_ids": []}, + ) + ) + + response = await serving.render_responses_request(ResponsesRequest(model=MODEL_NAME, input="Test prompt")) + + assert isinstance(response, ErrorResponse) + assert response.error.message == "No token_ids rendered" + + @pytest.fixture(scope="module") def server(): args: list[str] = ["--trust-request-chat-template"] @@ -26,6 +145,71 @@ async def client(server): yield http_client +@pytest.mark.asyncio +async def test_responses_render_basic(client): + response = await client.post( + "/v1/responses/render", + json={ + "model": MODEL_NAME, + "input": "When should a Responses handler return an empty string?", + "max_output_tokens": 7, + }, + ) + + assert response.status_code == 200 + data = response.json() + assert data["model"] == MODEL_NAME + assert data["request_id"].startswith("resp_") + assert data["sampling_params"]["max_tokens"] == 7 + assert data["token_ids"] + + +@pytest.mark.asyncio +async def test_responses_render_includes_prior_input_items(client): + current_turn = { + "role": "user", + "content": "Which color should I remember?", + } + current_turn_response = await client.post( + "/v1/responses/render", + json={"model": MODEL_NAME, "input": [current_turn]}, + ) + multi_turn_response = await client.post( + "/v1/responses/render", + json={ + "model": MODEL_NAME, + "input": [ + {"role": "user", "content": "Remember the color cobalt."}, + {"role": "assistant", "content": "I will remember cobalt."}, + current_turn, + ], + }, + ) + + assert current_turn_response.status_code == 200 + assert multi_turn_response.status_code == 200 + current_turn_token_ids = current_turn_response.json()["token_ids"] + multi_turn_token_ids = multi_turn_response.json()["token_ids"] + assert len(multi_turn_token_ids) > len(current_turn_token_ids) + + +@pytest.mark.asyncio +async def test_responses_render_rejects_previous_response_id_over_http(client): + response = await client.post( + "/v1/responses/render", + json={ + "model": MODEL_NAME, + "input": "Continue the previous response.", + "previous_response_id": "resp_previous", + }, + ) + + assert response.status_code == 400 + error = response.json()["error"] + assert error["type"] == "invalid_request_error" + assert error["param"] == "previous_response_id" + + @pytest.mark.asyncio async def test_completion_render_basic(client): """Test basic completion render endpoint.""" diff --git a/tests/entrypoints/scale_out/test_factories.py b/tests/entrypoints/scale_out/test_factories.py index 9c39638ee2..d737f58c0d 100644 --- a/tests/entrypoints/scale_out/test_factories.py +++ b/tests/entrypoints/scale_out/test_factories.py @@ -15,6 +15,7 @@ "/v1/completions/render", "/v1/completions/derender", "/v1/messages/render", + "/v1/responses/render", } GENERATE_PATH = "/inference/v1/generate" SCALE_OUT_PATHS = RENDER_PATHS | {GENERATE_PATH} From b753ef5ae46e842cc0f6defa435448296f0353ec Mon Sep 17 00:00:00 2001 From: AlpinDale Date: Mon, 7 Sep 2026 20:47:59 +0000 Subject: [PATCH 27/36] [sync] [Bugfix][Spec Decode] Resolve n_predict from text_config for Qwen3.5 multimodal MTP (#55369) Upstream-vLLM: b339d75a410e4e889049dff59042f385909bad67 Co-authored-by: Soumyajit Ghosh --- .sync/vllm-sha | 2 +- aphrodite/config/speculative.py | 3 + tests/models/test_qwen3_5_mtp_config.py | 92 ++++++++++++++++++++++--- 3 files changed, 88 insertions(+), 9 deletions(-) diff --git a/.sync/vllm-sha b/.sync/vllm-sha index 45322d2b8d..65a3baa455 100644 --- a/.sync/vllm-sha +++ b/.sync/vllm-sha @@ -1 +1 @@ -6a2a2bb02b563b83f946012959fd3927984d072a +b339d75a410e4e889049dff59042f385909bad67 diff --git a/aphrodite/config/speculative.py b/aphrodite/config/speculative.py index f06b5aa0b3..e184bce6e3 100644 --- a/aphrodite/config/speculative.py +++ b/aphrodite/config/speculative.py @@ -838,6 +838,9 @@ def hf_config_override(hf_config: PretrainedConfig) -> PretrainedConfig: is_moe = hf_config.model_type in ("qwen3_5_moe", "qwen3_5_moe_text") hf_config.model_type = "qwen3_5_mtp" n_predict = getattr(hf_config, "mtp_num_hidden_layers", None) + if n_predict is None: + text_config = get_hf_text_config(hf_config) + n_predict = getattr(text_config, "mtp_num_hidden_layers", None) hf_config.update( { "n_predict": n_predict, diff --git a/tests/models/test_qwen3_5_mtp_config.py b/tests/models/test_qwen3_5_mtp_config.py index 4fde0d3da0..d4a4efe2e4 100644 --- a/tests/models/test_qwen3_5_mtp_config.py +++ b/tests/models/test_qwen3_5_mtp_config.py @@ -1,19 +1,43 @@ # SPDX-License-Identifier: Apache-2.0 # SPDX-FileCopyrightText: Copyright contributors to the vLLM project -"""CPU-only tests for Qwen3.5 text-only MTP speculative decoding.""" +"""CPU-only tests for Qwen3.5 MTP speculative decoding config overrides.""" + +from typing import Any import pytest -from transformers import PretrainedConfig +from transformers import AutoConfig, PretrainedConfig from aphrodite.config.speculative import SpeculativeConfig +_CHECKPOINTS = { + "qwen3_5": "Qwen/Qwen3.8-27B", + "qwen3_5_moe": "Qwen/Qwen3.6-35B-A3B", +} + def _mtp_config(model_type: str) -> PretrainedConfig: - return PretrainedConfig( - model_type=model_type, - architectures=["SomeArch"], - mtp_num_hidden_layers=1, - ) + """Create a top-level MTP configuration with mtp_num_hidden_layers.""" + kwargs: dict[str, Any] = { + "model_type": model_type, + "architectures": ["SomeArch"], + "mtp_num_hidden_layers": 1, + } + return PretrainedConfig(**kwargs) + + +def _multimodal_wrapper_mtp_config(model_type: str, mtp_layers: int = 1) -> PretrainedConfig: + """Download a multimodal wrapper config via AutoConfig.from_pretrained + and configure mtp_num_hidden_layers in text_config. + + Uses real-world Hugging Face Hub checkpoints: + - Dense (qwen3_5): Qwen/Qwen3.8-27B + - MoE (qwen3_5_moe): Qwen/Qwen3.6-35B-A3B + """ + repo = _CHECKPOINTS[model_type] + config: PretrainedConfig = AutoConfig.from_pretrained(repo) + text_config = config.get_text_config() + text_config.mtp_num_hidden_layers = mtp_layers + return config @pytest.mark.parametrize( @@ -26,8 +50,60 @@ def _mtp_config(model_type: str) -> PretrainedConfig: ("qwen3_5_moe_text", "Qwen3_5MoeMTP"), ], ) -def test_mtp_override_recognizes_text_only_types(model_type, expected_arch): +def test_mtp_override_recognizes_text_only_types(model_type: str, expected_arch: str) -> None: + """Verify that text-only config variants map to the expected MTP architectures.""" cfg = SpeculativeConfig.hf_config_override(_mtp_config(model_type)) assert cfg.model_type == "qwen3_5_mtp" assert cfg.architectures == [expected_arch] assert cfg.n_predict == 1 + + +@pytest.mark.parametrize( + "model_type,expected_arch", + [ + ("qwen3_5", "Qwen3_5MTP"), + ("qwen3_5_moe", "Qwen3_5MoeMTP"), + ], +) +def test_mtp_override_extracts_n_predict_from_multimodal_wrapper(model_type: str, expected_arch: str) -> None: + """Verify that multimodal wrapper checkpoints with mtp_num_hidden_layers + in text_config resolve n_predict and architecture correctly.""" + cfg = SpeculativeConfig.hf_config_override(_multimodal_wrapper_mtp_config(model_type, mtp_layers=2)) + assert cfg.model_type == "qwen3_5_mtp" + assert cfg.architectures == [expected_arch] + assert cfg.n_predict == 2 + + +@pytest.mark.parametrize( + "model_type,expected_arch", + [ + ("qwen3_5", "Qwen3_5MTP"), + ("qwen3_5_moe", "Qwen3_5MoeMTP"), + ], +) +def test_mtp_override_top_level_precedence_over_nested_text_config(model_type: str, expected_arch: str) -> None: + """Verify that an explicit top-level mtp_num_hidden_layers takes precedence + over a nested text_config value.""" + cfg = _multimodal_wrapper_mtp_config(model_type, mtp_layers=2) + cfg.mtp_num_hidden_layers = 3 + overridden = SpeculativeConfig.hf_config_override(cfg) + assert overridden.model_type == "qwen3_5_mtp" + assert overridden.architectures == [expected_arch] + assert overridden.n_predict == 3 + + +@pytest.mark.parametrize( + "model_id,expected_arch", + [ + ("Qwen/Qwen3.8-27B", "Qwen3_5MTP"), + ("Qwen/Qwen3.6-35B-A3B", "Qwen3_5MoeMTP"), + ], +) +def test_mtp_override_downloads_real_hf_hub_configs(model_id: str, expected_arch: str) -> None: + """Verify that unmodified real-world checkpoints downloaded via + AutoConfig.from_pretrained resolve n_predict=1 from text_config.""" + hf_config = AutoConfig.from_pretrained(model_id) + cfg = SpeculativeConfig.hf_config_override(hf_config) + assert cfg.model_type == "qwen3_5_mtp" + assert cfg.architectures == [expected_arch] + assert cfg.n_predict == 1 From ca199e46d3133c6dd70e3df53d4a009e7020146c Mon Sep 17 00:00:00 2001 From: AlpinDale Date: Mon, 7 Sep 2026 20:51:12 +0000 Subject: [PATCH 28/36] [sync] [Frontend] Migrate Responses harmony input validation to VLLMValidationError (#55701) Upstream-vLLM: 167858ea174efd0655ae821d3f3a50b51d721525 Co-authored-by: Adababy --- .sync/vllm-sha | 2 +- aphrodite/entrypoints/openai/responses/harmony.py | 7 ++++--- .../openai/responses/test_harmony_utils.py | 9 +++++++++ .../responses/test_response_input_to_harmony.py | 12 ++++++++---- 4 files changed, 22 insertions(+), 8 deletions(-) diff --git a/.sync/vllm-sha b/.sync/vllm-sha index 65a3baa455..3430b92ebd 100644 --- a/.sync/vllm-sha +++ b/.sync/vllm-sha @@ -1 +1 @@ -b339d75a410e4e889049dff59042f385909bad67 +167858ea174efd0655ae821d3f3a50b51d721525 diff --git a/aphrodite/entrypoints/openai/responses/harmony.py b/aphrodite/entrypoints/openai/responses/harmony.py index 2d8d763cad..f27ae6703d 100644 --- a/aphrodite/entrypoints/openai/responses/harmony.py +++ b/aphrodite/entrypoints/openai/responses/harmony.py @@ -40,6 +40,7 @@ ResponseInputOutputItem, ResponsesRequest, ) +from aphrodite.exceptions import AphroditeValidationError from aphrodite.logger import init_logger from aphrodite.utils import random_uuid @@ -88,7 +89,7 @@ def _parse_chat_format_message(chat_msg: dict) -> list[Message]: """Parse an OpenAI chat-format dict into Harmony messages.""" role = chat_msg.get("role") if role is None: - raise ValueError(f"Message has no 'role' key: {chat_msg}") + raise AphroditeValidationError(f"Message has no 'role' key: {chat_msg}", parameter="input") # Assistant message with tool calls tool_calls = chat_msg.get("tool_calls") @@ -181,7 +182,7 @@ def response_input_to_harmony( call_response = prev_response break if call_response is None: - raise ValueError(f"No call message found for {call_id}") + raise AphroditeValidationError(f"No call message found for {call_id}", parameter="input") msg = Message.from_author_and_content( Author.new(Role.TOOL, f"functions.{call_response.name}"), response_msg["output"], @@ -202,7 +203,7 @@ def response_input_to_harmony( msg = msg.with_recipient(f"functions.{response_msg['name']}") msg = msg.with_content_type("json") else: - raise ValueError(f"Unknown input type: {response_msg['type']}") + raise AphroditeValidationError(f"Unknown input type: {response_msg['type']}", parameter="input") return msg diff --git a/tests/entrypoints/openai/responses/test_harmony_utils.py b/tests/entrypoints/openai/responses/test_harmony_utils.py index ce02ba2297..8d41125eba 100644 --- a/tests/entrypoints/openai/responses/test_harmony_utils.py +++ b/tests/entrypoints/openai/responses/test_harmony_utils.py @@ -16,6 +16,7 @@ harmony_to_response_output, response_previous_input_to_harmony, ) +from aphrodite.exceptions import AphroditeValidationError class TestResponsePreviousInputToHarmony: @@ -90,6 +91,14 @@ def test_tool_message_with_empty_content(self): assert messages[0].author.name == "functions.empty_tool" assert messages[0].content[0].text == "" + def test_chat_message_without_role_raises_validation_error(self): + """A chat-format message missing the 'role' key is invalid input.""" + chat_msg = {"content": "hello, no role here"} + + with pytest.raises(AphroditeValidationError, match="Message has no 'role' key") as exc_info: + response_previous_input_to_harmony(chat_msg) + assert exc_info.value.parameter == "input" + class TestHarmonyToResponseOutput: """Tests for harmony_to_response_output function.""" diff --git a/tests/entrypoints/openai/responses/test_response_input_to_harmony.py b/tests/entrypoints/openai/responses/test_response_input_to_harmony.py index 6aaa226d16..ab3332a406 100644 --- a/tests/entrypoints/openai/responses/test_response_input_to_harmony.py +++ b/tests/entrypoints/openai/responses/test_response_input_to_harmony.py @@ -15,6 +15,7 @@ from openai_harmony import DeveloperContent, Role from aphrodite.entrypoints.openai.responses.harmony import response_input_to_harmony +from aphrodite.exceptions import AphroditeValidationError # --------------------------------------------------------------------------- # Shared fixtures @@ -250,7 +251,7 @@ def test_function_call_output_skips_non_function_call_items_in_prev_responses( assert msg.author.name == "functions.get_weather" def test_function_call_output_raises_if_no_matching_call(self): - with pytest.raises(ValueError, match="No call message found for"): + with pytest.raises(AphroditeValidationError, match="No call message found for") as exc_info: response_input_to_harmony( { "type": "function_call_output", @@ -259,21 +260,24 @@ def test_function_call_output_raises_if_no_matching_call(self): }, prev_responses=[_PREV_CALL], ) + assert exc_info.value.parameter == "input" def test_function_call_output_raises_on_empty_prev_responses(self): - with pytest.raises(ValueError, match="No call message found for"): + with pytest.raises(AphroditeValidationError, match="No call message found for") as exc_info: response_input_to_harmony( {"type": "function_call_output", "call_id": "call_test", "output": "x"}, prev_responses=[], ) + assert exc_info.value.parameter == "input" # ----------------------------------------------------------------------- # Error cases # ----------------------------------------------------------------------- - def test_unknown_type_raises_value_error(self): - with pytest.raises(ValueError, match="Unknown input type"): + def test_unknown_type_raises_validation_error(self): + with pytest.raises(AphroditeValidationError, match="Unknown input type") as exc_info: response_input_to_harmony( {"type": "image_url", "url": "https://example.com/img.png"}, prev_responses=[], ) + assert exc_info.value.parameter == "input" From 92df6125f1f12db2422221cb149b0be2c760fe04 Mon Sep 17 00:00:00 2001 From: AlpinDale Date: Mon, 7 Sep 2026 20:53:24 +0000 Subject: [PATCH 29/36] [sync] [Pooling] Honor max_embed_len for chunked embeddings (#55551) Upstream-vLLM: ecd600d91e373d939ede33253a75d86bc5bf7338 Co-authored-by: Taneem Ibrahim --- .sync/vllm-sha | 2 +- aphrodite/entrypoints/pooling/base/protocol.py | 2 +- 2 files changed, 2 insertions(+), 2 deletions(-) diff --git a/.sync/vllm-sha b/.sync/vllm-sha index 3430b92ebd..7c95fadaea 100644 --- a/.sync/vllm-sha +++ b/.sync/vllm-sha @@ -1 +1 @@ -167858ea174efd0655ae821d3f3a50b51d721525 +ecd600d91e373d939ede33253a75d86bc5bf7338 diff --git a/aphrodite/entrypoints/pooling/base/protocol.py b/aphrodite/entrypoints/pooling/base/protocol.py index 942a0b9c8c..c5fb261027 100644 --- a/aphrodite/entrypoints/pooling/base/protocol.py +++ b/aphrodite/entrypoints/pooling/base/protocol.py @@ -179,7 +179,7 @@ def build_tok_params(self, model_config: ModelConfig) -> TokenizeParams: pooler_config = model_config.pooler_config if pooler_config is not None: if pooler_config.enable_chunked_processing: - max_total_tokens = None + max_total_tokens = pooler_config.max_embed_len else: max_embed_len = pooler_config.max_embed_len or default_max_total_tokens max_output_tokens = default_max_total_tokens - max_embed_len From b2b5570063523c2e364669e508606315ae050152 Mon Sep 17 00:00:00 2001 From: AlpinDale Date: Mon, 7 Sep 2026 21:00:40 +0000 Subject: [PATCH 30/36] [sync] [2/N][warmup][DSv4] Migrate sequence and DCP kernels (#53564) Upstream-vLLM: 1713b9866ae5c380d97da5796217064ac7ab57b8 Co-authored-by: Roberto L. Castro <38211239+LopezCastroRoberto@users.noreply.github.com> --- .sync/vllm-sha | 2 +- .../attention/dsa/dcp_indexer_cutedsl.py | 842 +++++++++++------- .../layers/sparse_attn_indexer.py | 25 +- aphrodite/model_executor/warmup/jit_warmup.py | 38 +- .../warmup/jit_warmup_cutedsl_helper.py | 80 ++ .../warmup/jit_warmup_triton_helper.py | 19 +- aphrodite/utils/cpu_triton_utils.py | 2 + .../v1/attention/backends/mla/indexer.py | 5 +- aphrodite/v1/attention/ops/common.py | 384 +++++--- aphrodite/v1/attention/ops/dcp.py | 350 ++++++-- aphrodite/v1/worker/block_table.py | 72 +- tests/model_executor/test_jit_warmup.py | 4 +- .../test_jit_warmup_cutedsl_launcher.py | 76 ++ .../test_jit_warmup_triton_launcher.py | 163 ++++ 14 files changed, 1468 insertions(+), 594 deletions(-) create mode 100644 aphrodite/model_executor/warmup/jit_warmup_cutedsl_helper.py create mode 100644 tests/model_executor/test_jit_warmup_cutedsl_launcher.py create mode 100644 tests/model_executor/test_jit_warmup_triton_launcher.py diff --git a/.sync/vllm-sha b/.sync/vllm-sha index 7c95fadaea..f1dd3b8b83 100644 --- a/.sync/vllm-sha +++ b/.sync/vllm-sha @@ -1 +1 @@ -ecd600d91e373d939ede33253a75d86bc5bf7338 +1713b9866ae5c380d97da5796217064ac7ab57b8 diff --git a/aphrodite/model_executor/kernels/attention/dsa/dcp_indexer_cutedsl.py b/aphrodite/model_executor/kernels/attention/dsa/dcp_indexer_cutedsl.py index ce75c252ba..ae86f30545 100644 --- a/aphrodite/model_executor/kernels/attention/dsa/dcp_indexer_cutedsl.py +++ b/aphrodite/model_executor/kernels/attention/dsa/dcp_indexer_cutedsl.py @@ -1,7 +1,8 @@ # SPDX-License-Identifier: Apache-2.0 # SPDX-FileCopyrightText: Copyright contributors to the vLLM project -from functools import cache +from dataclasses import dataclass +from typing import Any import cutlass import cutlass.cute as cute @@ -11,6 +12,17 @@ from quack.compile_utils import make_fake_tensor from aphrodite.cute_utils import recast_val +from aphrodite.model_executor.warmup.jit_warmup_cutedsl_helper import ( + AphroditeCuTeDSLJitKernel, + CuTeDSLLaunchSpec, + cutedsl_kernel_launcher, +) +from aphrodite.model_executor.warmup.jit_warmup_triton_helper import ( + AphroditeTritonJitKernel, + LaunchSpec, + TritonWarmupTensor, + kernel_launcher, +) from aphrodite.triton_utils import tl, triton @@ -25,7 +37,7 @@ def stable_topk_from_gathered_candidates_cutedsl( dtype=torch.int32, device=gathered.device, ) - StableTopKFromGatheredCandidatesKernel.compile(topk, gathered.shape[1])(gathered, out) + _STABLE_TOPK_FROM_GATHERED_CANDIDATES_KERNEL(gathered, out, topk=topk) return out @@ -39,393 +51,537 @@ def pack_dcp_topk_candidates_cutedsl( row_starts: torch.Tensor | None, ) -> None: topk = topk_indices.shape[1] - grid = (topk_indices.shape[0], triton.cdiv(topk, 512)) row_starts_arg = row_starts if row_starts is not None else topk_indices - _pack_dcp_topk_candidates_triton_kernel[grid]( + _PACK_DCP_TOPK_CANDIDATES_KERNEL( logits, topk_indices, packed, row_starts_arg, - logits.stride(0), - logits.stride(1), - topk_indices.stride(0), - topk_indices.stride(1), - packed.stride(0), - packed.stride(1), - packed.stride(2), - logits.shape[1], - DCP_RANK=dcp_rank, - DCP_WORLD_SIZE=dcp_world_size, - CP_INTERLEAVE=cp_interleave, - HAS_ROW_STARTS=row_starts is not None, - TOPK=topk, - BLOCK_SIZE=512, - num_warps=8, + logits_stride0=logits.stride(0), + logits_stride1=logits.stride(1), + topk_stride0=topk_indices.stride(0), + topk_stride1=topk_indices.stride(1), + packed_stride0=packed.stride(0), + packed_stride1=packed.stride(1), + packed_stride2=packed.stride(2), + num_cols=logits.shape[1], + dcp_rank=dcp_rank, + dcp_world_size=dcp_world_size, + cp_interleave=cp_interleave, + has_row_starts=row_starts is not None, + topk=topk, + block_size=512, ) -# Strides and num_cols are runtime scalars excluded from specialization. -# logits_stride0 is the gathered-K row length, which varies with every batch's -# prompt/context length, so folding it into the compile key (constexpr or -# divisibility specialization) would re-JIT this kernel per prefill shape. -@triton.jit( - do_not_specialize=[ - "logits_stride0", - "logits_stride1", - "topk_stride0", - "topk_stride1", - "packed_stride0", - "packed_stride1", - "packed_stride2", - "num_cols", - ] -) -def _pack_dcp_topk_candidates_triton_kernel( - logits, - topk_indices, - packed, - row_starts, - logits_stride0, - logits_stride1, - topk_stride0, - topk_stride1, - packed_stride0, - packed_stride1, - packed_stride2, - num_cols, - DCP_RANK: tl.constexpr, - DCP_WORLD_SIZE: tl.constexpr, - CP_INTERLEAVE: tl.constexpr, - HAS_ROW_STARTS: tl.constexpr, - TOPK: tl.constexpr, - BLOCK_SIZE: tl.constexpr, -): - row = tl.program_id(0) - tile = tl.program_id(1) - cols = tile * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE) - mask = cols < TOPK - - local_idx = tl.load( - topk_indices + row * topk_stride0 + cols * topk_stride1, - mask=mask, - other=-1, - ) - valid = local_idx >= 0 - safe_local_idx = tl.maximum(local_idx, 0) - - row_start = 0 - if HAS_ROW_STARTS: - row_start = tl.load(row_starts + row) - - score_col = safe_local_idx + row_start - score_col = tl.minimum(score_col, tl.maximum(num_cols - 1, 0)) - score = tl.load( - logits + row * logits_stride0 + score_col * logits_stride1, - mask=mask & valid, - other=-float("inf"), - ) +class PackDCPTopkCandidatesKernel(AphroditeTritonJitKernel["PackDCPTopkCandidatesKernel.CompileKey"]): + @dataclass(frozen=True) + class CompileKey: + dcp_rank: int + dcp_world_size: int + cp_interleave: int + has_row_starts: bool + topk: int + block_size: int - global_id = ( - (safe_local_idx // CP_INTERLEAVE) * (DCP_WORLD_SIZE * CP_INTERLEAVE) - + DCP_RANK * CP_INTERLEAVE - + safe_local_idx % CP_INTERLEAVE + # These scalars only describe runtime layouts and bounds; the constexpr + # fields below own the launch geometry and algorithmic specialization. + @staticmethod + @triton.jit( + do_not_specialize=[ + "logits_stride0", + "logits_stride1", + "topk_stride0", + "topk_stride1", + "packed_stride0", + "packed_stride1", + "packed_stride2", + "num_cols", + ] ) - global_id = tl.where(valid, global_id, -1) - - packed_base = packed + row * packed_stride0 + cols * packed_stride1 - tl.store(packed_base, score, mask=mask) - tl.store(packed_base + packed_stride2, global_id.to(tl.float32), mask=mask) - - -@cute.jit -def _warp_scan_inclusive_i32(val: Int32, lane: Int32) -> Int32: - for i in cutlass.range_constexpr(cute.arch.WARP_SIZE.bit_length() - 1): - offset = 1 << i - partial = cute.arch.shuffle_sync_up(val, offset=offset, mask_and_clamp=0) - if lane >= offset: - val += partial - return val - - -@cute.jit -def _block_scan_inclusive_i32( - val: Int32, - lane: Int32, - warp_id: Int32, - warp_scratch: cute.Tensor, - warps_per_block: int, -) -> Int32: - prefix = _warp_scan_inclusive_i32(val, lane) - if lane == Int32(cute.arch.WARP_SIZE - 1): - warp_scratch[0, warp_id] = prefix - cute.arch.sync_threads() - - if warp_id == Int32(0): - warp_total = Int32(0) - if lane < Int32(warps_per_block): - warp_total = warp_scratch[0, lane] - warp_prefix = _warp_scan_inclusive_i32(warp_total, lane) - if lane < Int32(warps_per_block): - warp_scratch[0, lane] = warp_prefix - warp_total - cute.arch.sync_threads() - - return prefix + warp_scratch[0, warp_id] - - -class StableTopKFromGatheredCandidatesKernel: + def kernel( + logits, + topk_indices, + packed, + row_starts, + logits_stride0, + logits_stride1, + topk_stride0, + topk_stride1, + packed_stride0, + packed_stride1, + packed_stride2, + num_cols, + DCP_RANK: tl.constexpr, + DCP_WORLD_SIZE: tl.constexpr, + CP_INTERLEAVE: tl.constexpr, + HAS_ROW_STARTS: tl.constexpr, + TOPK: tl.constexpr, + BLOCK_SIZE: tl.constexpr, + ): + row = tl.program_id(0) + tile = tl.program_id(1) + cols = tile * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE) + mask = cols < TOPK + + local_idx = tl.load( + topk_indices + row * topk_stride0 + cols * topk_stride1, + mask=mask, + other=-1, + ) + valid = local_idx >= 0 + safe_local_idx = tl.maximum(local_idx, 0) + + row_start = 0 + if HAS_ROW_STARTS: + row_start = tl.load(row_starts + row) + + score_col = safe_local_idx + row_start + score_col = tl.minimum(score_col, tl.maximum(num_cols - 1, 0)) + score = tl.load( + logits + row * logits_stride0 + score_col * logits_stride1, + mask=mask & valid, + other=-float("inf"), + ) + + global_id = ( + (safe_local_idx // CP_INTERLEAVE) * (DCP_WORLD_SIZE * CP_INTERLEAVE) + + DCP_RANK * CP_INTERLEAVE + + safe_local_idx % CP_INTERLEAVE + ) + global_id = tl.where(valid, global_id, -1) + + packed_base = packed + row * packed_stride0 + cols * packed_stride1 + tl.store(packed_base, score, mask=mask) + tl.store(packed_base + packed_stride2, global_id.to(tl.float32), mask=mask) + + def dispatch( # type: ignore[override] + self, + *, + has_row_starts: bool, + **compile_key_fields: int, + ) -> CompileKey: + return self.CompileKey( + **compile_key_fields, + has_row_starts=has_row_starts, + ) + + def get_warmup_keys(self, aphrodite_config: Any) -> list[CompileKey]: + dcp_world_size = aphrodite_config.parallel_config.decode_context_parallel_size + if dcp_world_size <= 1: + return [] + cp_interleave = aphrodite_config.parallel_config.cp_kv_cache_interleave_size + topk = aphrodite_config.model_config.hf_config.index_topk + if topk <= 0: + return [] + + try: + from aphrodite.distributed.parallel_state import get_dcp_group + + dcp_rank = get_dcp_group().rank_in_group + except Exception: + dcp_rank = 0 + + return self._trace_dispatch(self.dispatch)( + dcp_rank=dcp_rank, + dcp_world_size=dcp_world_size, + cp_interleave=cp_interleave, + has_row_starts=(False, True), + topk=topk, + block_size=512, + ) + + def warmup_inputs(self, compile_key: CompileKey) -> dict[str, Any]: + fp32_ptr = TritonWarmupTensor(torch.float32) + int32_ptr = TritonWarmupTensor(torch.int32) + return dict( + logits=fp32_ptr, + topk_indices=TritonWarmupTensor(torch.int32, shape=(1, compile_key.topk)), + packed=fp32_ptr, + row_starts_arg=int32_ptr, + logits_stride0=1, + logits_stride1=1, + topk_stride0=1, + topk_stride1=1, + packed_stride0=1, + packed_stride1=1, + packed_stride2=1, + num_cols=1, + dcp_rank=compile_key.dcp_rank, + dcp_world_size=compile_key.dcp_world_size, + cp_interleave=compile_key.cp_interleave, + has_row_starts=compile_key.has_row_starts, + topk=compile_key.topk, + block_size=compile_key.block_size, + ) + + @kernel_launcher + def __call__( + self, + logits: torch.Tensor, + topk_indices: torch.Tensor, + packed: torch.Tensor, + row_starts_arg: torch.Tensor, + *, + logits_stride0: int, + logits_stride1: int, + topk_stride0: int, + topk_stride1: int, + packed_stride0: int, + packed_stride1: int, + packed_stride2: int, + num_cols: int, + dcp_rank: int, + dcp_world_size: int, + cp_interleave: int, + has_row_starts: bool, + topk: int, + block_size: int, + ) -> LaunchSpec: + grid = (topk_indices.shape[0], triton.cdiv(topk, block_size)) + return grid, dict( + row_starts=row_starts_arg, + DCP_RANK=dcp_rank, + DCP_WORLD_SIZE=dcp_world_size, + CP_INTERLEAVE=cp_interleave, + HAS_ROW_STARTS=has_row_starts, + TOPK=topk, + BLOCK_SIZE=block_size, + num_warps=8, + ) + + +class StableTopKFromGatheredCandidatesKernel( + AphroditeCuTeDSLJitKernel["StableTopKFromGatheredCandidatesKernel.CompileKey"] +): tb_size = 512 hist_bins = 2048 radix_bits = (hist_bins - 1).bit_length() - assert hist_bins == 1 << radix_bits key_bits = Uint64.width radix_passes = (key_bits + radix_bits - 1) // radix_bits final_radix_bits = key_bits - radix_bits * (radix_passes - 1) hist_chunks = (hist_bins + tb_size - 1) // tb_size warps_per_block = tb_size // cute.arch.WARP_SIZE - def __init__(self, topk: int, num_candidates: int): - assert num_candidates % self.tb_size == 0, ( + @dataclass(frozen=True) + class CompileKey: + topk: int + num_candidates: int + + @staticmethod + def kernel(compile_key: CompileKey) -> Any: + tb_size = StableTopKFromGatheredCandidatesKernel.tb_size + hist_bins = StableTopKFromGatheredCandidatesKernel.hist_bins + radix_bits = StableTopKFromGatheredCandidatesKernel.radix_bits + key_bits = StableTopKFromGatheredCandidatesKernel.key_bits + radix_passes = StableTopKFromGatheredCandidatesKernel.radix_passes + final_radix_bits = StableTopKFromGatheredCandidatesKernel.final_radix_bits + hist_chunks = StableTopKFromGatheredCandidatesKernel.hist_chunks + warps_per_block = StableTopKFromGatheredCandidatesKernel.warps_per_block + topk = compile_key.topk + assert hist_bins == 1 << radix_bits + assert compile_key.num_candidates % tb_size == 0, ( "StableTopKFromGatheredCandidatesKernel requires candidate count " - f"to be a multiple of {self.tb_size}, got {num_candidates}" + f"to be a multiple of {tb_size}, got {compile_key.num_candidates}" ) - self.topk = topk - self.keys_per_thread = num_candidates // self.tb_size + keys_per_thread = compile_key.num_candidates // tb_size @cute.struct class SharedStorage: - hist: cute.struct.MemRange[Int32, self.hist_bins] + hist: cute.struct.MemRange[Int32, hist_bins] committed_count: cute.struct.MemRange[Int32, 1] running_count: cute.struct.MemRange[Int32, 1] threshold_bin: cute.struct.MemRange[Int32, 1] threshold_found: cute.struct.MemRange[Int32, 1] include_threshold_bin: cute.struct.MemRange[Int32, 1] prefix_s: cute.struct.Align[cute.struct.MemRange[Uint64, 1], 8] - warp_totals: cute.struct.MemRange[Int32, self.warps_per_block] - - self.shared_storage = SharedStorage - - @cute.jit - def __call__( - self, - gathered: cute.Tensor, - out: cute.Tensor, - stream: CUstream, - ): - grid = (gathered.shape[0], 1, 1) - self.kernel(gathered, out).launch( - grid=grid, - block=(self.tb_size, 1, 1), - stream=stream, - ) + warp_totals: cute.struct.MemRange[Int32, warps_per_block] + + shared_storage = SharedStorage + + @cute.jit + def warp_scan_inclusive_i32(val: Int32, lane: Int32) -> Int32: + for i in cutlass.range_constexpr(cute.arch.WARP_SIZE.bit_length() - 1): + offset = 1 << i + partial = cute.arch.shuffle_sync_up(val, offset=offset, mask_and_clamp=0) + if lane >= offset: + val += partial + return val + + @cute.jit + def block_scan_inclusive_i32( + val: Int32, + lane: Int32, + warp_id: Int32, + warp_scratch: cute.Tensor, + warps_per_block: int, + ) -> Int32: + prefix = warp_scan_inclusive_i32(val, lane) + if lane == Int32(cute.arch.WARP_SIZE - 1): + warp_scratch[0, warp_id] = prefix + cute.arch.sync_threads() - @cute.jit - def _stable_key(self, score: Float32, token_id: Int32) -> Uint64: - bits = recast_val(score, Uint32) - mask = Uint32(0x80000000) - if (bits & Uint32(0x80000000)) != Uint32(0): - mask = Uint32(0xFFFFFFFF) - score_key = Uint64(bits ^ mask) << Uint64(32) - id_key = Uint64(~Uint32(token_id)) - key = score_key | id_key - if token_id < Int32(0): - key = Uint64(0) - return key - - @cute.jit - def _prefix_matches( - self, - key: Uint64, - prefix: Uint64, - prefix_bits: Int32, - ): - matches = prefix_bits == Int32(0) - if prefix_bits != Int32(0): - shift = Int32(self.key_bits) - prefix_bits - matches = (key >> Uint64(shift)) == (prefix >> Uint64(shift)) - return matches - - @cute.jit - def _radix_pass( - self, - keys: cute.Tensor, - output: cute.Tensor, - storage, - tid: Int32, - step: Int32, - bits: int, - is_final_pass: bool, - ): - hist_smem = storage.hist.get_tensor(cute.make_layout((self.hist_bins,))) - committed_count_smem = storage.committed_count.data_ptr() - running_count_smem = storage.running_count.data_ptr() - threshold_bin_smem = storage.threshold_bin.data_ptr() - threshold_found_smem = storage.threshold_found.data_ptr() - include_threshold_bin_smem = storage.include_threshold_bin.data_ptr() - prefix_smem = storage.prefix_s.data_ptr() - warp_totals_smem = storage.warp_totals.get_tensor(cute.make_layout((1, self.warps_per_block))) - - prefix_bits = step * Int32(self.radix_bits) - num_bins = 1 << bits - block_scan_iterations = (num_bins + self.tb_size - 1) // self.tb_size - shift = Int32(self.key_bits) - prefix_bits - Int32(bits) - bin_mask = Uint64(num_bins - 1) - prefix = prefix_smem.load() - - for chunk in cutlass.range_constexpr(self.hist_chunks): - hist_smem[tid + Int32(chunk * self.tb_size)] = Int32(0) - if tid == Int32(0): - running_count_smem.store(committed_count_smem.load()) - include_threshold_bin_smem.store(Int32(0)) - threshold_found_smem.store(Int32(0)) - cute.arch.sync_threads() - - for key_idx in cutlass.range_constexpr(self.keys_per_thread): - key = keys[key_idx] - if self._prefix_matches(key, prefix, prefix_bits): - bin_idx = Int32((key >> Uint64(shift)) & bin_mask) - cute.arch.atomic_add( - hist_smem.iterator + bin_idx, - Int32(1), - sem="relaxed", - scope="cta", - ) - cute.arch.sync_threads() - - lane = cute.arch.lane_idx() - warp_id = cute.arch.warp_idx() - # Each iteration scans one tb_size-wide slice of bins, high to low. - iter = Int32(0) - threshold_found = threshold_found_smem.load() - while threshold_found == Int32(0) and iter < Int32(block_scan_iterations): - bin_idx = Int32(num_bins - 1) - (iter * Int32(self.tb_size) + tid) - count = hist_smem[bin_idx] - chunk_inclusive = _block_scan_inclusive_i32( - count, - lane, - warp_id, - warp_totals_smem, - self.warps_per_block, - ) - running_count = running_count_smem.load() - prior_in_scan_slice = chunk_inclusive - count - remaining = Int32(self.topk) - running_count - prior_in_scan_slice - if count > Int32(0) and remaining > Int32(0) and remaining <= count: - threshold_bin_smem.store(bin_idx) - if count <= remaining or cutlass.const_expr(is_final_pass): - include_threshold_bin_smem.store(Int32(1)) - threshold_found_smem.store(Int32(1)) - # Barrier: every thread must finish reading running_count for this - # slice before tb_size-1 advances it, else a warp racing ahead to - # the store makes a lagging thread double-count the slice total - # (-> remaining too small -> threshold too high -> under-fill). + if warp_id == Int32(0): + warp_total = Int32(0) + if lane < Int32(warps_per_block): + warp_total = warp_scratch[0, lane] + warp_prefix = warp_scan_inclusive_i32(warp_total, lane) + if lane < Int32(warps_per_block): + warp_scratch[0, lane] = warp_prefix - warp_total cute.arch.sync_threads() - if tid == Int32(self.tb_size - 1): - running_count_smem.store(running_count + chunk_inclusive) + + return prefix + warp_scratch[0, warp_id] + + @cute.jit + def stable_key(score: Float32, token_id: Int32) -> Uint64: + bits = recast_val(score, Uint32) + mask = Uint32(0x80000000) + if (bits & Uint32(0x80000000)) != Uint32(0): + mask = Uint32(0xFFFFFFFF) + score_key = Uint64(bits ^ mask) << Uint64(32) + id_key = Uint64(~Uint32(token_id)) + key = score_key | id_key + if token_id < Int32(0): + key = Uint64(0) + return key + + @cute.jit + def prefix_matches( + key: Uint64, + prefix: Uint64, + prefix_bits: Int32, + ): + matches = prefix_bits == Int32(0) + if prefix_bits != Int32(0): + shift = Int32(key_bits) - prefix_bits + matches = (key >> Uint64(shift)) == (prefix >> Uint64(shift)) + return matches + + @cute.jit + def radix_pass( + keys: cute.Tensor, + output: cute.Tensor, + storage, + tid: Int32, + step: Int32, + bits: int, + is_final_pass: bool, + ): + hist_smem = storage.hist.get_tensor(cute.make_layout((hist_bins,))) + committed_count_smem = storage.committed_count.data_ptr() + running_count_smem = storage.running_count.data_ptr() + threshold_bin_smem = storage.threshold_bin.data_ptr() + threshold_found_smem = storage.threshold_found.data_ptr() + include_threshold_bin_smem = storage.include_threshold_bin.data_ptr() + prefix_smem = storage.prefix_s.data_ptr() + warp_totals_smem = storage.warp_totals.get_tensor(cute.make_layout((1, warps_per_block))) + + prefix_bits = step * Int32(radix_bits) + num_bins = 1 << bits + block_scan_iterations = (num_bins + tb_size - 1) // tb_size + shift = Int32(key_bits) - prefix_bits - Int32(bits) + bin_mask = Uint64(num_bins - 1) + prefix = prefix_smem.load() + + for chunk in cutlass.range_constexpr(hist_chunks): + hist_smem[tid + Int32(chunk * tb_size)] = Int32(0) + if tid == Int32(0): + running_count_smem.store(committed_count_smem.load()) + include_threshold_bin_smem.store(Int32(0)) + threshold_found_smem.store(Int32(0)) cute.arch.sync_threads() - threshold_found = threshold_found_smem.load() - iter += Int32(1) - - threshold = threshold_bin_smem.load() - should_include_threshold = include_threshold_bin_smem.load() != Int32(0) - for key_idx in cutlass.range_constexpr(self.keys_per_thread): - key = keys[key_idx] - if self._prefix_matches(key, prefix, prefix_bits): - bin_idx = Int32((key >> Uint64(shift)) & bin_mask) - selected = bin_idx > threshold - if should_include_threshold: - selected = selected or bin_idx == threshold - if selected: - dst = cute.arch.atomic_add( - committed_count_smem, + for key_idx in cutlass.range_constexpr(keys_per_thread): + key = keys[key_idx] + if prefix_matches(key, prefix, prefix_bits): + bin_idx = Int32((key >> Uint64(shift)) & bin_mask) + cute.arch.atomic_add( + hist_smem.iterator + bin_idx, Int32(1), sem="relaxed", scope="cta", ) - if dst < Int32(self.topk): - output[dst] = recast_val(~Uint32(key), Int32) - cute.arch.sync_threads() + cute.arch.sync_threads() - pass_finished = include_threshold_bin_smem.load() - if tid == Int32(0) and pass_finished == Int32(0): - prefix_smem.store(prefix | (Uint64(threshold) << Uint64(shift))) - cute.arch.sync_threads() - return pass_finished + lane = cute.arch.lane_idx() + warp_id = cute.arch.warp_idx() + # Each iteration scans one tb_size-wide slice of bins, high to low. + iter = Int32(0) + threshold_found = threshold_found_smem.load() + while threshold_found == Int32(0) and iter < Int32(block_scan_iterations): + bin_idx = Int32(num_bins - 1) - (iter * Int32(tb_size) + tid) + count = hist_smem[bin_idx] + chunk_inclusive = block_scan_inclusive_i32( + count, + lane, + warp_id, + warp_totals_smem, + warps_per_block, + ) + running_count = running_count_smem.load() + prior_in_scan_slice = chunk_inclusive - count + remaining = Int32(topk) - running_count - prior_in_scan_slice + if count > Int32(0) and remaining > Int32(0) and remaining <= count: + threshold_bin_smem.store(bin_idx) + if count <= remaining or cutlass.const_expr(is_final_pass): + include_threshold_bin_smem.store(Int32(1)) + threshold_found_smem.store(Int32(1)) + # Barrier: every thread must finish reading running_count for this + # slice before tb_size-1 advances it, else a warp racing ahead to + # the store makes a lagging thread double-count the slice total + # (-> remaining too small -> threshold too high -> under-fill). + cute.arch.sync_threads() + if tid == Int32(tb_size - 1): + running_count_smem.store(running_count + chunk_inclusive) + cute.arch.sync_threads() + + threshold_found = threshold_found_smem.load() + iter += Int32(1) + + threshold = threshold_bin_smem.load() + should_include_threshold = include_threshold_bin_smem.load() != Int32(0) + for key_idx in cutlass.range_constexpr(keys_per_thread): + key = keys[key_idx] + if prefix_matches(key, prefix, prefix_bits): + bin_idx = Int32((key >> Uint64(shift)) & bin_mask) + selected = bin_idx > threshold + if should_include_threshold: + selected = selected or bin_idx == threshold + if selected: + dst = cute.arch.atomic_add( + committed_count_smem, + Int32(1), + sem="relaxed", + scope="cta", + ) + if dst < Int32(topk): + output[dst] = recast_val(~Uint32(key), Int32) + cute.arch.sync_threads() - @cute.kernel - def kernel( - self, - input: cute.Tensor, - out: cute.Tensor, - ): - row, _, _ = cute.arch.block_idx() - tid, _, _ = cute.arch.thread_idx() - input_row = input[row, None, None] - output_row = out[row, None] - keys = cute.make_rmem_tensor((self.keys_per_thread,), Uint64) - - smem = cutlass.utils.SmemAllocator() - storage = smem.allocate(self.shared_storage, 8) - committed_count_smem = storage.committed_count.data_ptr() - prefix_smem = storage.prefix_s.data_ptr() - for i in range(tid, self.topk, self.tb_size): - output_row[i] = Int32(-1) - - for key_idx in cutlass.range_constexpr(self.keys_per_thread): - col = tid + Int32(key_idx * self.tb_size) - score = Float32(input_row[col, 0]) - token_id = Int32(input_row[col, 1]) - keys[key_idx] = self._stable_key(score, token_id) - - if tid == Int32(0): - committed_count_smem.store(Int32(0)) - prefix_smem.store(Uint64(0)) - cute.arch.sync_threads() - - step = Int32(0) - finished = Int32(0) - while finished == Int32(0) and step < Int32(self.radix_passes - 1): - finished = self._radix_pass( - keys, - output_row, - storage, - tid, - step, - self.radix_bits, - False, - ) - step += Int32(1) - - if finished == Int32(0): - self._radix_pass( - keys, - output_row, - storage, - tid, - Int32(self.radix_passes - 1), - self.final_radix_bits, - True, + pass_finished = include_threshold_bin_smem.load() + if tid == Int32(0) and pass_finished == Int32(0): + prefix_smem.store(prefix | (Uint64(threshold) << Uint64(shift))) + cute.arch.sync_threads() + return pass_finished + + @cute.kernel + def device_kernel(input: cute.Tensor, out: cute.Tensor): + row, _, _ = cute.arch.block_idx() + tid, _, _ = cute.arch.thread_idx() + input_row = input[row, None, None] + output_row = out[row, None] + keys = cute.make_rmem_tensor((keys_per_thread,), Uint64) + + smem = cutlass.utils.SmemAllocator() + storage = smem.allocate(shared_storage, 8) + committed_count_smem = storage.committed_count.data_ptr() + prefix_smem = storage.prefix_s.data_ptr() + for i in range(tid, topk, tb_size): + output_row[i] = Int32(-1) + + for key_idx in cutlass.range_constexpr(keys_per_thread): + col = tid + Int32(key_idx * tb_size) + score = Float32(input_row[col, 0]) + token_id = Int32(input_row[col, 1]) + keys[key_idx] = stable_key(score, token_id) + + if tid == Int32(0): + committed_count_smem.store(Int32(0)) + prefix_smem.store(Uint64(0)) + cute.arch.sync_threads() + + step = Int32(0) + finished = Int32(0) + while finished == Int32(0) and step < Int32(radix_passes - 1): + finished = radix_pass( + keys, + output_row, + storage, + tid, + step, + radix_bits, + False, + ) + step += Int32(1) + + if finished == Int32(0): + radix_pass( + keys, + output_row, + storage, + tid, + Int32(radix_passes - 1), + final_radix_bits, + True, + ) + + @cute.jit + def host_entrypoint( + gathered: cute.Tensor, + out: cute.Tensor, + stream: CUstream, + ): + grid = (gathered.shape[0], 1, 1) + device_kernel(gathered, out).launch( + grid=grid, + block=(tb_size, 1, 1), + stream=stream, ) - @cache - @staticmethod - def compile(topk: int, num_candidates: int): - num_rows = cute.sym_int() + return host_entrypoint + + def dispatch( # type: ignore[override] + self, + **compile_key_fields: int, + ) -> CompileKey: + return self.CompileKey(**compile_key_fields) + + def get_warmup_keys(self, aphrodite_config: Any) -> list[CompileKey]: + dcp_world_size = aphrodite_config.parallel_config.decode_context_parallel_size + if dcp_world_size <= 1: + return [] + topk = aphrodite_config.model_config.hf_config.index_topk + if topk <= 0: + return [] + return self._trace_dispatch(self.dispatch)( + topk=topk, + num_candidates=topk * dcp_world_size, + ) + def warmup_inputs(self, compile_key: CompileKey) -> tuple[Any, ...]: + num_rows = cute.sym_int() gathered = cute.runtime.make_fake_tensor( Float32, - (num_rows, num_candidates, 2), + (num_rows, compile_key.num_candidates, 2), stride=(cute.sym_int64(divisibility=2), 2, 1), assumed_align=8, ) - out = make_fake_tensor(Int32, (num_rows, topk), divisibility=1) - - kernel = StableTopKFromGatheredCandidatesKernel(topk, num_candidates) - stream = cute.runtime.make_fake_stream(use_tvm_ffi_env_stream=True) - return cute.compile( - kernel, - gathered, - out, - stream, - options="--enable-tvm-ffi", + out = make_fake_tensor( + Int32, + (num_rows, compile_key.topk), + divisibility=1, ) + return gathered, out + + @cutedsl_kernel_launcher + def __call__( + self, + gathered: torch.Tensor, + out: torch.Tensor, + *, + topk: int, + ) -> CuTeDSLLaunchSpec[CompileKey]: + compile_key = self.dispatch(topk=topk, num_candidates=gathered.shape[1]) + return ( + compile_key, + (gathered, out), + { + "gathered_shape": tuple(gathered.shape), + "out_shape": tuple(out.shape), + "topk": topk, + }, + ) + + +_PACK_DCP_TOPK_CANDIDATES_KERNEL = PackDCPTopkCandidatesKernel() +_STABLE_TOPK_FROM_GATHERED_CANDIDATES_KERNEL = StableTopKFromGatheredCandidatesKernel() diff --git a/aphrodite/model_executor/layers/sparse_attn_indexer.py b/aphrodite/model_executor/layers/sparse_attn_indexer.py index 63cddc2009..c10be3a1d8 100644 --- a/aphrodite/model_executor/layers/sparse_attn_indexer.py +++ b/aphrodite/model_executor/layers/sparse_attn_indexer.py @@ -871,7 +871,8 @@ def __init__( # DCP scalars are constant for the run; resolve them here (config is set # during model construction) and pass them into the custom op, rather # than threading them through per-step metadata. - parallel_config = get_current_aphrodite_config().parallel_config + aphrodite_config = get_current_aphrodite_config() + parallel_config = aphrodite_config.parallel_config self._parallel_config = parallel_config self.dcp_world_size = parallel_config.decode_context_parallel_size self.dcp_rank = get_dcp_group().rank_in_group if self.dcp_world_size > 1 else 0 @@ -892,6 +893,28 @@ def __init__( "Sparse Attention Indexer CUDA op requires DeepGEMM support in the current Aphrodite environment." ) + if aphrodite_config.kernel_config.enable_jit_warmup: + from aphrodite.v1.attention.ops.common import ( + _PACK_SEQ_TRITON_KERNEL, + _UNPACK_SEQ_TRITON_KERNEL, + ) + + pack_dtype = torch.uint8 if use_fp4_cache else current_platform.fp8_dtype() + _PACK_SEQ_TRITON_KERNEL.register_warmup( + dtype=pack_dtype, + pad_value=0 if use_fp4_cache else -float("inf"), + ) + _UNPACK_SEQ_TRITON_KERNEL.register_warmup() + + if self.dcp_world_size > 1 and current_platform.is_cuda() and has_cutedsl(): + from aphrodite.model_executor.kernels.attention.dsa.dcp_indexer_cutedsl import ( # noqa: E501 + _PACK_DCP_TOPK_CANDIDATES_KERNEL, + _STABLE_TOPK_FROM_GATHERED_CANDIDATES_KERNEL, + ) + + _PACK_DCP_TOPK_CANDIDATES_KERNEL.register_warmup() + _STABLE_TOPK_FROM_GATHERED_CANDIDATES_KERNEL.register_warmup() + @property def cp_kv_cache_interleave_size(self) -> int: """With PD+DCP, the real value isn't known until block_size is finalized, diff --git a/aphrodite/model_executor/warmup/jit_warmup.py b/aphrodite/model_executor/warmup/jit_warmup.py index 6c90dc39cd..f5455c9f43 100644 --- a/aphrodite/model_executor/warmup/jit_warmup.py +++ b/aphrodite/model_executor/warmup/jit_warmup.py @@ -612,6 +612,34 @@ def warmup(self, *args: Any, **kwargs: Any) -> None: self.compile(compile_key) +def _same_value(left: Any, right: Any) -> bool: + """True if two registration values are the same object or compare equal. + + Identity is checked first so shared singletons (e.g. ``aphrodite_config``, torch + dtypes) short-circuit before any potentially deep or non-boolean ``__eq__``. + """ + if left is right: + return True + try: + return bool(left == right) + except (TypeError, ValueError, RuntimeError): + return False + + +def _same_registration( + left: tuple[tuple[Any, ...], dict[str, Any]], + right: tuple[tuple[Any, ...], dict[str, Any]], +) -> bool: + """True if two ``(args, kwargs)`` registrations are equivalent.""" + left_args, left_kwargs = left + right_args, right_kwargs = right + if len(left_args) != len(right_args) or left_kwargs.keys() != right_kwargs.keys(): + return False + return all(_same_value(x, y) for x, y in zip(left_args, right_args, strict=True)) and all( + _same_value(left_kwargs[k], right_kwargs[k]) for k in left_kwargs + ) + + class JitWarmupRegistry: """Collect and compile JIT kernels selected during runner setup.""" @@ -655,13 +683,9 @@ def _add( kwargs: dict[str, Any], ) -> None: registrations = self._registrations.setdefault(kernel, []) - if ( - not args - and not kwargs - and any( - not registered_args and not registered_kwargs for registered_args, registered_kwargs in registrations - ) - ): + # Layers can share the same configuration object. Deduplicate before + # tracing warmup keys, checking identity before potentially deep equality. + if any(_same_registration(registered, (args, kwargs)) for registered in registrations): return registrations.append((args, kwargs)) diff --git a/aphrodite/model_executor/warmup/jit_warmup_cutedsl_helper.py b/aphrodite/model_executor/warmup/jit_warmup_cutedsl_helper.py new file mode 100644 index 0000000000..796042a8c5 --- /dev/null +++ b/aphrodite/model_executor/warmup/jit_warmup_cutedsl_helper.py @@ -0,0 +1,80 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project + +from abc import abstractmethod +from collections.abc import Callable, Mapping +from functools import wraps +from typing import Any, ClassVar, Generic, TypeAlias, TypeVar + +from aphrodite.model_executor.warmup.jit_warmup import AphroditeJitKernel + +DEFAULT_CUTEDSL_COMPILE_OPTIONS = "--enable-tvm-ffi" + +CompileKeyT = TypeVar("CompileKeyT") +CuTeDSLLaunchSpec: TypeAlias = tuple[ + CompileKeyT, + tuple[Any, ...], + Mapping[str, Any] | None, +] + + +def cutedsl_fake_stream(*, use_tvm_ffi_env_stream: bool = True) -> Any: + import cutlass.cute as cute + + return cute.runtime.make_fake_stream(use_tvm_ffi_env_stream=use_tvm_ffi_env_stream) + + +def compile_cutedsl( + entry: Callable[..., Any], + *args: Any, + options: str = DEFAULT_CUTEDSL_COMPILE_OPTIONS, + use_tvm_ffi_env_stream: bool = True, +) -> Any: + import cutlass.cute as cute + + return cute.compile( + entry, + *args, + cutedsl_fake_stream(use_tvm_ffi_env_stream=use_tvm_ffi_env_stream), + options=options, + ) + + +class AphroditeCuTeDSLJitKernel(AphroditeJitKernel[CompileKeyT], Generic[CompileKeyT]): + """CuTeDSL owner whose compiled executor is shared by warmup and runtime.""" + + kernel: ClassVar[Any] + + @abstractmethod + def warmup_inputs(self, compile_key: CompileKeyT) -> tuple[Any, ...]: + """Return fake arguments that compile one executor specialization.""" + raise NotImplementedError + + def compile(self, compile_key: CompileKeyT) -> None: + if compile_key in self._compiled_cache: + return + self._compiled_cache[compile_key] = compile_cutedsl( + self.kernel(compile_key), + *self.warmup_inputs(compile_key), + ) + + +def cutedsl_kernel_launcher( + call_fn: Callable[..., CuTeDSLLaunchSpec[CompileKeyT]], +) -> Callable[..., Any]: + """Invoke a cached CuTeDSL executor from a declarative ``__call__``.""" + + @wraps(call_fn) + def wrapper( + self: AphroditeCuTeDSLJitKernel[CompileKeyT], + *args: Any, + **kwargs: Any, + ) -> Any: + compile_key, launch_args, runtime_context = call_fn(self, *args, **kwargs) + executor = self._get_or_compile( + compile_key, + runtime_context=runtime_context, + ) + return executor(*launch_args) + + return wrapper diff --git a/aphrodite/model_executor/warmup/jit_warmup_triton_helper.py b/aphrodite/model_executor/warmup/jit_warmup_triton_helper.py index f4860e6e56..ead0a11482 100644 --- a/aphrodite/model_executor/warmup/jit_warmup_triton_helper.py +++ b/aphrodite/model_executor/warmup/jit_warmup_triton_helper.py @@ -107,8 +107,14 @@ def compile(self, compile_key: CompileKeyT) -> None: self._warming = False @cached_property - def _kernel_param_names(self) -> frozenset[str]: - return frozenset(self.kernel.arg_names) + def _kernel_arg_names(self) -> tuple[str, ...]: + arg_names = getattr(self.kernel, "arg_names", None) + if arg_names is not None: + return tuple(arg_names) + wrapped = getattr(self.kernel, "func", None) + if wrapped is not None: + return tuple(inspect.signature(wrapped).parameters) + raise TypeError(f"Cannot inspect kernel parameters for {type(self.kernel).__name__}") def launch( self, @@ -117,14 +123,19 @@ def launch( /, **kwargs: Any, ) -> Any: + runtime_launcher = kwargs.pop("_runtime_launcher", None) + runtime_launcher_arg_count = kwargs.pop("_runtime_launcher_arg_count", 0) for name, value in inputs.items(): - target = name if name in self._kernel_param_names else f"{name}_ptr" - if target in self._kernel_param_names and target not in kwargs: + target = name if name in self._kernel_arg_names else f"{name}_ptr" + if target in self._kernel_arg_names and target not in kwargs: kwargs[target] = value if self._warming: warmup = getattr(self.kernel, "warmup", None) assert warmup is not None return warmup(grid=(1,), **kwargs) + if runtime_launcher is not None: + regular_args = [kwargs.pop(name) for name in self._kernel_arg_names[:runtime_launcher_arg_count]] + return runtime_launcher(self.kernel, grid, *regular_args, **kwargs) return self.kernel[grid](**kwargs) diff --git a/aphrodite/utils/cpu_triton_utils.py b/aphrodite/utils/cpu_triton_utils.py index 2e99dbc4d6..0761dde422 100644 --- a/aphrodite/utils/cpu_triton_utils.py +++ b/aphrodite/utils/cpu_triton_utils.py @@ -29,6 +29,8 @@ def _compute_slot_mapping_kernel_impl( block_table_stride: int, # max_num_blocks_per_req block_size: int, slot_mapping: torch.Tensor, # [max_num_tokens], int64 + KV_CACHE_BLOCK_SIZE: int, + BLOCKS_PER_KV_BLOCK: int, TOTAL_CP_WORLD_SIZE: int, TOTAL_CP_RANK: int, CP_KV_CACHE_INTERLEAVE_SIZE: int, diff --git a/aphrodite/v1/attention/backends/mla/indexer.py b/aphrodite/v1/attention/backends/mla/indexer.py index 7ddb2475bb..5ede1fe12e 100644 --- a/aphrodite/v1/attention/backends/mla/indexer.py +++ b/aphrodite/v1/attention/backends/mla/indexer.py @@ -873,7 +873,7 @@ def _warmup_dcp_kernels(self) -> None: index_topk = getattr(self.aphrodite_config.model_config.hf_config, "index_topk", None) if has_cutedsl() and index_topk in (512, 1024, 2048): from aphrodite.model_executor.kernels.attention.dsa.dcp_indexer_cutedsl import ( - StableTopKFromGatheredCandidatesKernel, + _STABLE_TOPK_FROM_GATHERED_CANDIDATES_KERNEL, pack_dcp_topk_candidates_cutedsl, ) @@ -896,7 +896,8 @@ def _warmup_dcp_kernels(self) -> None: # Keyed on (topk, num_candidates) with a symbolic row count, # so one compile covers every batch. This is the seconds-scale # CuteDSL compile. - StableTopKFromGatheredCandidatesKernel.compile(index_topk, self.dcp_world_size * index_topk) + kernel = _STABLE_TOPK_FROM_GATHERED_CANDIDATES_KERNEL + kernel.compile(kernel.dispatch(topk=index_topk, num_candidates=self.dcp_world_size * index_topk)) except Exception: logger.warning( "DCP indexer kernel warmup failed; kernels will be JIT-compiled on first use instead.", diff --git a/aphrodite/v1/attention/ops/common.py b/aphrodite/v1/attention/ops/common.py index 5e0bbf06ba..61b5d7d7e0 100644 --- a/aphrodite/v1/attention/ops/common.py +++ b/aphrodite/v1/attention/ops/common.py @@ -1,63 +1,157 @@ # SPDX-License-Identifier: Apache-2.0 # SPDX-FileCopyrightText: Copyright contributors to the vLLM project +from dataclasses import dataclass +from typing import Any + import torch +from aphrodite.model_executor.warmup.jit_warmup_triton_helper import ( + AphroditeTritonJitKernel, + LaunchSpec, + TritonWarmupTensor, + kernel_launcher, +) from aphrodite.triton_utils import tl, triton -@triton.jit -def _pack_seq_kernel( - x_ptr, # [N, D] - out_ptr, # [B, Lmax, D] - lengths_ptr, # *i32, [B] - N: tl.constexpr, - D: tl.constexpr, - Lmax: tl.constexpr, - PAD_VALUE: tl.constexpr, - PAD_IS_UINT8: tl.constexpr, - BLOCK_T: tl.constexpr, # timesteps per program - BLOCK_D: tl.constexpr, # features per program -): - pid_b = tl.program_id(0) # batch id - pid_t = tl.program_id(1) # block over time dimension - pid_d = tl.program_id(2) # block over feature dimension - off_t = pid_t * BLOCK_T + tl.arange(0, BLOCK_T) # [BLOCK_T] - off_d = pid_d * BLOCK_D + tl.arange(0, BLOCK_D) # [BLOCK_D] - - # Compute start index and sequence length from cumulative lengths - in_start = 0 - for i in range(pid_b): - in_start += tl.load(lengths_ptr + i) - seq_len = tl.load(lengths_ptr + pid_b) - - # valid time positions for this block - t_mask = off_t < Lmax - - # compute input row indices for valid (b, t) - in_row = in_start + off_t - valid_row = (off_t < seq_len) & t_mask - - # Pointers - # x_ptr: row-major [N, D] - x_row_ptr = x_ptr + in_row[:, None] * D + off_d[None, :] - - # out_ptr: row-major [B, Lmax, D] - out_row_ptr = out_ptr + (pid_b * Lmax + off_t)[:, None] * D + off_d[None, :] - - # Initialize with PAD. PAD_IS_UINT8 selects the pad tensor's dtype so - # integer-typed outputs (e.g. MXFP4 packed nibbles, ue8m0 scale bytes) - # get an exact-byte pad rather than going through an fp32→uint8 cast - # that's implementation-defined outside of value 0. - d_mask = off_d[None, :] < D - if PAD_IS_UINT8: - pad_vals = tl.full([BLOCK_T, BLOCK_D], PAD_VALUE, tl.uint8) - else: - pad_vals = tl.full([BLOCK_T, BLOCK_D], PAD_VALUE, tl.float32) - tl.store(out_row_ptr, pad_vals, mask=t_mask[:, None] & d_mask) +class PackSeqTritonKernel(AphroditeTritonJitKernel["PackSeqTritonKernel.CompileKey"]): + @dataclass(frozen=True) + class CompileKey: + dtype: torch.dtype + pad_value: float | int + pad_is_uint8: bool + block_t: int + block_d: int + + # N is unused; D and Lmax only affect runtime addressing and masks. The + # compile-time tile shape remains controlled by BLOCK_T and BLOCK_D. + @staticmethod + @triton.jit(do_not_specialize=["N", "D", "Lmax"]) + def kernel( + x_ptr, # [N, D] + out_ptr, # [B, Lmax, D] + lengths_ptr, # *i32, [B] + N, + D, + Lmax, + PAD_VALUE: tl.constexpr, + PAD_IS_UINT8: tl.constexpr, + BLOCK_T: tl.constexpr, # timesteps per program + BLOCK_D: tl.constexpr, # features per program + ): + pid_b = tl.program_id(0) # batch id + pid_t = tl.program_id(1) # block over time dimension + pid_d = tl.program_id(2) # block over feature dimension + off_t = pid_t * BLOCK_T + tl.arange(0, BLOCK_T) # [BLOCK_T] + off_d = pid_d * BLOCK_D + tl.arange(0, BLOCK_D) # [BLOCK_D] + + # Compute start index and sequence length from cumulative lengths + in_start = 0 + for i in range(pid_b): + in_start += tl.load(lengths_ptr + i) + seq_len = tl.load(lengths_ptr + pid_b) + + # valid time positions for this block + t_mask = off_t < Lmax + + # compute input row indices for valid (b, t) + in_row = in_start + off_t + valid_row = (off_t < seq_len) & t_mask + + # Pointers + # x_ptr: row-major [N, D] + x_row_ptr = x_ptr + in_row[:, None] * D + off_d[None, :] + + # out_ptr: row-major [B, Lmax, D] + out_row_ptr = out_ptr + (pid_b * Lmax + off_t)[:, None] * D + off_d[None, :] + + # Initialize with PAD. PAD_IS_UINT8 selects the pad tensor's dtype so + # integer-typed outputs (e.g. MXFP4 packed nibbles, ue8m0 scale bytes) + # get an exact-byte pad rather than going through an fp32→uint8 cast + # that's implementation-defined outside of value 0. + d_mask = off_d[None, :] < D + if PAD_IS_UINT8: + pad_vals = tl.full([BLOCK_T, BLOCK_D], PAD_VALUE, tl.uint8) + else: + pad_vals = tl.full([BLOCK_T, BLOCK_D], PAD_VALUE, tl.float32) + tl.store(out_row_ptr, pad_vals, mask=t_mask[:, None] & d_mask) + + # Load & write only where within seq_len + x_vals = tl.load(x_row_ptr, mask=valid_row[:, None] & d_mask) + tl.store(out_row_ptr, x_vals, mask=valid_row[:, None] & d_mask) + + def dispatch( # type: ignore[override] + self, + *, + dtype: torch.dtype, + pad_value: float | int, + **compile_key_fields: int, + ) -> CompileKey: + is_uint8 = dtype == torch.uint8 + return self.CompileKey( + **compile_key_fields, + dtype=dtype, + pad_value=int(pad_value) if is_uint8 else float(pad_value), + pad_is_uint8=is_uint8, + ) + + def get_warmup_keys( + self, + *, + dtype: torch.dtype, + pad_value: float | int, + ) -> list[CompileKey]: + return self._trace_dispatch(self.dispatch)( + dtype=dtype, + pad_value=pad_value, + block_t=64, + block_d=64, + ) - # Load & write only where within seq_len - x_vals = tl.load(x_row_ptr, mask=valid_row[:, None] & d_mask) - tl.store(out_row_ptr, x_vals, mask=valid_row[:, None] & d_mask) + def warmup_inputs(self, compile_key: CompileKey) -> dict[str, Any]: + return dict( + x_reshaped=TritonWarmupTensor(compile_key.dtype), + out=TritonWarmupTensor(compile_key.dtype), + lengths=TritonWarmupTensor(torch.int32), + N=1, + D=1, + Lmax=1, + pad_value=compile_key.pad_value, + pad_is_uint8=compile_key.pad_is_uint8, + block_t=compile_key.block_t, + block_d=compile_key.block_d, + ) + + @kernel_launcher + def __call__( + self, + x_reshaped: torch.Tensor, + out: torch.Tensor, + lengths: torch.Tensor, + *, + N: int, + D: int, + Lmax: int, + pad_value: float | int, + pad_is_uint8: bool, + block_t: int, + block_d: int, + ) -> LaunchSpec: + grid = ( + lengths.shape[0], + triton.cdiv(Lmax, block_t), + triton.cdiv(D, block_d), + ) + return grid, dict( + x_ptr=x_reshaped, + lengths_ptr=lengths, + PAD_VALUE=pad_value, + PAD_IS_UINT8=pad_is_uint8, + BLOCK_T=block_t, + BLOCK_D=block_d, + num_warps=4, + num_stages=2, + ) def pack_seq_triton( @@ -109,20 +203,18 @@ def pack_seq_triton( out = torch.empty((B, Lmax, D), device=x.device, dtype=x.dtype) - grid = (B, triton.cdiv(Lmax, block_t), triton.cdiv(D, block_d)) - _pack_seq_kernel[grid]( + lengths = lengths.int() + _PACK_SEQ_TRITON_KERNEL( x_reshaped, out, - lengths.int(), - N, - D, - Lmax, - PAD_VALUE=pad_constexpr, - PAD_IS_UINT8=is_uint8, - BLOCK_T=block_t, - BLOCK_D=block_d, - num_warps=4, - num_stages=2, + lengths, + N=N, + D=D, + Lmax=Lmax, + pad_value=pad_constexpr, + pad_is_uint8=is_uint8, + block_t=block_t, + block_d=block_d, ) if len(original_shape) > 2: @@ -131,47 +223,107 @@ def pack_seq_triton( return out -@triton.jit -def _unpack_seq_triton_kernel( - packed_ptr, # [B, Lmax, D] - out_ptr, # [N, D] - lengths_ptr, # *i32, [B] - B: tl.constexpr, - Lmax: tl.constexpr, - D: tl.constexpr, - BLOCK_T: tl.constexpr, # timesteps per program - BLOCK_D: tl.constexpr, # features per program -): - pid_b = tl.program_id(0) # batch id - pid_t = tl.program_id(1) # block over time dimension - pid_d = tl.program_id(2) # block over feature dimension - off_t = pid_t * BLOCK_T + tl.arange(0, BLOCK_T) # [BLOCK_T] - off_d = pid_d * BLOCK_D + tl.arange(0, BLOCK_D) # [BLOCK_D] - - # bounds: compute start from cumulative lengths - in_start = 0 - for i in range(pid_b): - in_start += tl.load(lengths_ptr + i) - seq_len = tl.load(lengths_ptr + pid_b) - - # valid time positions for this block - t_mask = off_t < Lmax - valid_row = (off_t < seq_len) & t_mask - - # compute output row indices for valid (b, t) - out_row = in_start + off_t - - # Pointers - # packed_ptr: row-major [B, Lmax, D] - packed_row_ptr = packed_ptr + (pid_b * Lmax + off_t)[:, None] * D + off_d[None, :] - - # out_ptr: row-major [N, D] - out_row_ptr = out_ptr + out_row[:, None] * D + off_d[None, :] - - # Load from packed tensor and store to output - d_mask = off_d[None, :] < D - packed_vals = tl.load(packed_row_ptr, mask=valid_row[:, None] & d_mask) - tl.store(out_row_ptr, packed_vals, mask=valid_row[:, None] & d_mask) +class UnpackSeqTritonKernel(AphroditeTritonJitKernel["UnpackSeqTritonKernel.CompileKey"]): + @dataclass(frozen=True) + class CompileKey: + dtype: torch.dtype + block_t: int + block_d: int + + # B, Lmax, and D only affect runtime addressing and masks. The compile-time + # tile shape remains controlled by BLOCK_T and BLOCK_D. + @staticmethod + @triton.jit(do_not_specialize=["B", "Lmax", "D"]) + def kernel( + packed_ptr, # [B, Lmax, D] + out_ptr, # [N, D] + lengths_ptr, # *i32, [B] + B, + Lmax, + D, + BLOCK_T: tl.constexpr, # timesteps per program + BLOCK_D: tl.constexpr, # features per program + ): + pid_b = tl.program_id(0) # batch id + pid_t = tl.program_id(1) # block over time dimension + pid_d = tl.program_id(2) # block over feature dimension + off_t = pid_t * BLOCK_T + tl.arange(0, BLOCK_T) # [BLOCK_T] + off_d = pid_d * BLOCK_D + tl.arange(0, BLOCK_D) # [BLOCK_D] + + # bounds: compute start from cumulative lengths + in_start = 0 + for i in range(pid_b): + in_start += tl.load(lengths_ptr + i) + seq_len = tl.load(lengths_ptr + pid_b) + + # valid time positions for this block + t_mask = off_t < Lmax + valid_row = (off_t < seq_len) & t_mask + + # compute output row indices for valid (b, t) + out_row = in_start + off_t + + # Pointers + # packed_ptr: row-major [B, Lmax, D] + packed_row_ptr = packed_ptr + (pid_b * Lmax + off_t)[:, None] * D + off_d[None, :] + + # out_ptr: row-major [N, D] + out_row_ptr = out_ptr + out_row[:, None] * D + off_d[None, :] + + # Load from packed tensor and store to output + d_mask = off_d[None, :] < D + packed_vals = tl.load(packed_row_ptr, mask=valid_row[:, None] & d_mask) + tl.store(out_row_ptr, packed_vals, mask=valid_row[:, None] & d_mask) + + def dispatch( # type: ignore[override] + self, + *, + dtype: torch.dtype, + **compile_key_fields: int, + ) -> CompileKey: + return self.CompileKey(**compile_key_fields, dtype=dtype) + + def get_warmup_keys(self) -> list[CompileKey]: + return self._trace_dispatch(self.dispatch)( + dtype=torch.int32, + block_t=64, + block_d=64, + ) + + def warmup_inputs(self, compile_key: CompileKey) -> dict[str, Any]: + return dict( + packed_reshaped=TritonWarmupTensor(compile_key.dtype), + out=TritonWarmupTensor(compile_key.dtype), + lengths=TritonWarmupTensor(torch.int32), + B=1, + Lmax=1, + D=1, + block_t=compile_key.block_t, + block_d=compile_key.block_d, + ) + + @kernel_launcher + def __call__( + self, + packed_reshaped: torch.Tensor, + out: torch.Tensor, + lengths: torch.Tensor, + *, + B: int, + Lmax: int, + D: int, + block_t: int, + block_d: int, + ) -> LaunchSpec: + grid = (B, triton.cdiv(Lmax, block_t), triton.cdiv(D, block_d)) + return grid, dict( + packed_ptr=packed_reshaped, + lengths_ptr=lengths, + BLOCK_T=block_t, + BLOCK_D=block_d, + num_warps=4, + num_stages=2, + ) def unpack_seq_triton( @@ -209,18 +361,16 @@ def unpack_seq_triton( out = torch.empty((N, D), device=packed_tensor.device, dtype=packed_tensor.dtype) - grid = (B, triton.cdiv(Lmax, block_t), triton.cdiv(D, block_d)) - _unpack_seq_triton_kernel[grid]( + lengths = lengths.int() + _UNPACK_SEQ_TRITON_KERNEL( packed_reshaped, out, - lengths.int(), - B, - Lmax, - D, - BLOCK_T=block_t, - BLOCK_D=block_d, - num_warps=4, - num_stages=2, + lengths, + B=B, + Lmax=Lmax, + D=D, + block_t=block_t, + block_d=block_d, ) # Reshape output back to original dimensions (except first dimension) @@ -229,3 +379,7 @@ def unpack_seq_triton( out = out.reshape(output_shape) return out + + +_PACK_SEQ_TRITON_KERNEL = PackSeqTritonKernel() +_UNPACK_SEQ_TRITON_KERNEL = UnpackSeqTritonKernel() diff --git a/aphrodite/v1/attention/ops/dcp.py b/aphrodite/v1/attention/ops/dcp.py index 27c33d0102..a93fb336c2 100644 --- a/aphrodite/v1/attention/ops/dcp.py +++ b/aphrodite/v1/attention/ops/dcp.py @@ -6,7 +6,8 @@ import functools from collections.abc import Callable -from typing import TYPE_CHECKING, Protocol +from dataclasses import dataclass +from typing import TYPE_CHECKING, Any, Protocol import torch import torch.distributed as dist @@ -15,6 +16,16 @@ from aphrodite.config import AphroditeConfig from aphrodite.distributed import get_dcp_group from aphrodite.logger import init_logger +from aphrodite.model_executor.warmup.jit_warmup import ( + WarmupIntRange, +) +from aphrodite.model_executor.warmup.jit_warmup_triton_helper import ( + AphroditeTritonJitKernel, + LaunchSpec, + TritonWarmupTensor, + kernel_launcher, + triton_scalar_specialization_rep, +) from aphrodite.triton_utils import tl, triton from aphrodite.v1.attention.ops.cp_common import ( DirectCPWorkspace, @@ -55,83 +66,238 @@ def mask_dcp_empty_shards_( # AG + RS/AR implementation -@triton.jit -def _correct_attn_cp_out_kernel( - outputs_ptr, - new_output_ptr, - lses_ptr, - vlse_ptr, - outputs_stride_B, - outputs_stride_H, - outputs_stride_D, - lses_stride_N, - lses_stride_B, - lses_stride_H, - lse_idx, - HEAD_DIM: tl.constexpr, - N_ROUNDED: tl.constexpr, - IS_BASE_E: tl.constexpr, -): - """ - Apply the all-gathered lses to correct each local rank's attention - output. we still need perform a cross-rank reduction to obtain the - final attention output. +class CorrectAttnCPOutKernel(AphroditeTritonJitKernel["CorrectAttnCPOutKernel.CompileKey"]): + @dataclass(frozen=True) + class CompileKey: + output_dtype: torch.dtype + lse_dtype: torch.dtype + outputs_stride_b: int + outputs_stride_h: int + outputs_stride_d: int + lses_stride_n: int + lses_stride_b: int + lses_stride_h: int + lse_idx: int + head_dim: int + n_rounded: int + is_base_e: bool + + @staticmethod + @triton.jit + def kernel( + outputs_ptr, + new_output_ptr, + lses_ptr, + vlse_ptr, + outputs_stride_B, + outputs_stride_H, + outputs_stride_D, + lses_stride_N, + lses_stride_B, + lses_stride_H, + lse_idx, + HEAD_DIM: tl.constexpr, + N_ROUNDED: tl.constexpr, + IS_BASE_E: tl.constexpr, + ): + """ + Apply the all-gathered lses to correct each local rank's attention + output. we still need perform a cross-rank reduction to obtain the + final attention output. + + Args: + outputs_ptr (triton.PointerType): + Pointer to input tensor of shape [ B, H, D ] + lses_ptr (triton.PointerType): + Pointer to input tensor of shape [ N, B, H ] + new_output_ptr (triton.PointerType): + Pointer to output tensor of shape [ B, H, D ] + vlse_ptr (triton.PointerType): + Pointer to output tensor of shape [ B, H ] + """ + batch_idx = tl.program_id(axis=0).to(tl.int64) + head_idx = tl.program_id(axis=1).to(tl.int64) + d_offsets = tl.arange(0, HEAD_DIM) + num_n_offsets = tl.arange(0, N_ROUNDED) + + # shape = [N] + lse_offsets = num_n_offsets * lses_stride_N + batch_idx * lses_stride_B + head_idx * lses_stride_H + + # calc final lse + lse = tl.load(lses_ptr + lse_offsets).to(tl.float32) + lse = tl.where((lse != lse) | (lse == float("inf")), -float("inf"), lse) + lse_max = tl.max(lse, axis=0) + lse_max = tl.where(lse_max == -float("inf"), 0, lse_max) + lse -= lse_max + if IS_BASE_E: + lse_exp = tl.exp(lse) + lse_acc = tl.sum(lse_exp, axis=0) + lse = tl.log(lse_acc) + else: + lse_exp = tl.exp2(lse) + lse_acc = tl.sum(lse_exp, axis=0) + lse = tl.log2(lse_acc) + lse += lse_max + + lse_offsets = batch_idx * lses_stride_B + head_idx * lses_stride_H + tl.store(vlse_ptr + lse_offsets, lse) + + # shape = [D] + output_offsets = batch_idx * outputs_stride_B + head_idx * outputs_stride_H + d_offsets * outputs_stride_D + + # correct output + lse_offset = lse_idx * lses_stride_N + batch_idx * lses_stride_B + head_idx * lses_stride_H + lse_tmp = tl.load(lses_ptr + lse_offset).to(tl.float32) + lse_finally = lse_tmp - lse + lse_finally = tl.where( + (lse_finally != lse_finally) | (lse_finally == float("inf")), + -float("inf"), + lse_finally, + ) + factor = tl.exp(lse_finally) if IS_BASE_E else tl.exp2(lse_finally) + output = tl.load(outputs_ptr + output_offsets) + output = output * factor + output = tl.where(factor == 0.0, 0.0, output) - Args: - outputs_ptr (triton.PointerType): - Pointer to input tensor of shape [ B, H, D ] - lses_ptr (triton.PointerType): - Pointer to input tensor of shape [ N, B, H ] - new_output_ptr (triton.PointerType): - Pointer to output tensor of shape [ B, H, D ] - vlse_ptr (triton.PointerType): - Pointer to output tensor of shape [ B, H ] - """ - batch_idx = tl.program_id(axis=0).to(tl.int64) - head_idx = tl.program_id(axis=1).to(tl.int64) - d_offsets = tl.arange(0, HEAD_DIM) - num_n_offsets = tl.arange(0, N_ROUNDED) - - # shape = [N] - lse_offsets = num_n_offsets * lses_stride_N + batch_idx * lses_stride_B + head_idx * lses_stride_H - - # calc final lse - lse = tl.load(lses_ptr + lse_offsets).to(tl.float32) - lse = tl.where((lse != lse) | (lse == float("inf")), -float("inf"), lse) - lse_max = tl.max(lse, axis=0) - lse_max = tl.where(lse_max == -float("inf"), 0, lse_max) - lse -= lse_max - if IS_BASE_E: - lse_exp = tl.exp(lse) - lse_acc = tl.sum(lse_exp, axis=0) - lse = tl.log(lse_acc) - else: - lse_exp = tl.exp2(lse) - lse_acc = tl.sum(lse_exp, axis=0) - lse = tl.log2(lse_acc) - lse += lse_max - - lse_offsets = batch_idx * lses_stride_B + head_idx * lses_stride_H - tl.store(vlse_ptr + lse_offsets, lse) - - # shape = [D] - output_offsets = batch_idx * outputs_stride_B + head_idx * outputs_stride_H + d_offsets * outputs_stride_D - - # correct output - lse_offset = lse_idx * lses_stride_N + batch_idx * lses_stride_B + head_idx * lses_stride_H - lse_tmp = tl.load(lses_ptr + lse_offset).to(tl.float32) - lse_finally = lse_tmp - lse - lse_finally = tl.where( - (lse_finally != lse_finally) | (lse_finally == float("inf")), - -float("inf"), - lse_finally, - ) - factor = tl.exp(lse_finally) if IS_BASE_E else tl.exp2(lse_finally) - output = tl.load(outputs_ptr + output_offsets) - output = output * factor - output = tl.where(factor == 0.0, 0.0, output) + tl.store(new_output_ptr + output_offsets, output) + + def dispatch( # type: ignore[override] + self, + *, + output_dtype: torch.dtype, + lse_dtype: torch.dtype, + num_tokens: int, + num_heads: int, + head_dim: int, + n_rounded: int, + lse_idx: int, + is_base_e: bool, + ) -> CompileKey: + # Warm the contiguous [B, H, D] output and [N, B, H] LSE layouts used + # by the production collective path. Runtime still forwards real + # strides so unsupported layouts remain correct, but may compile on + # first use under the JIT monitor. + outputs_stride_d = 1 + outputs_stride_h = head_dim + outputs_stride_b = num_heads * head_dim + lses_stride_h = 1 + lses_stride_b = num_heads + lses_stride_n = num_tokens * num_heads + return self.CompileKey( + output_dtype=output_dtype, + lse_dtype=lse_dtype, + outputs_stride_b=triton_scalar_specialization_rep(outputs_stride_b), + outputs_stride_h=triton_scalar_specialization_rep(outputs_stride_h), + outputs_stride_d=triton_scalar_specialization_rep(outputs_stride_d), + lses_stride_n=triton_scalar_specialization_rep(lses_stride_n), + lses_stride_b=triton_scalar_specialization_rep(lses_stride_b), + lses_stride_h=triton_scalar_specialization_rep(lses_stride_h), + lse_idx=triton_scalar_specialization_rep(lse_idx), + head_dim=head_dim, + n_rounded=n_rounded, + is_base_e=is_base_e, + ) + + def get_warmup_keys( + self, + aphrodite_config: Any, + *, + output_dtype: torch.dtype | None = None, + num_heads: int | None = None, + head_dim: int | None = None, + is_base_e: bool | tuple[bool, ...] = (False, True), + ) -> list[CompileKey]: + from aphrodite.model_executor.layers.attention.mla_attention import ( + get_mla_dims, + ) - tl.store(new_output_ptr + output_offsets, output) + dcp_world_size = aphrodite_config.parallel_config.decode_context_parallel_size + max_tokens = min( + 16, + aphrodite_config.scheduler_config.max_num_batched_tokens, + ) + hf_config = aphrodite_config.model_config.hf_config + if num_heads is None: + num_heads = hf_config.num_attention_heads // aphrodite_config.parallel_config.tensor_parallel_size + if head_dim is None: + head_dim = get_mla_dims(aphrodite_config.model_config).v_head_dim + if output_dtype is None: + output_dtype = aphrodite_config.model_config.dtype + if dcp_world_size <= 1 or max_tokens <= 0 or num_heads <= 0: + return [] + return self._trace_dispatch(self.dispatch)( + output_dtype=output_dtype, + lse_dtype=torch.float32, + num_tokens=WarmupIntRange(1, max_tokens + 1), + num_heads=num_heads, + head_dim=head_dim, + n_rounded=dcp_world_size, + lse_idx=WarmupIntRange(0, dcp_world_size), + is_base_e=is_base_e, + ) + + def warmup_inputs(self, compile_key: CompileKey) -> dict[str, Any]: + output_ptr = TritonWarmupTensor( + compile_key.output_dtype, + shape=(1, 1, compile_key.head_dim), + strides=( + compile_key.outputs_stride_b, + compile_key.outputs_stride_h, + compile_key.outputs_stride_d, + ), + ) + lse_ptr = TritonWarmupTensor( + compile_key.lse_dtype, + shape=(compile_key.n_rounded, 1, 1), + strides=( + compile_key.lses_stride_n, + compile_key.lses_stride_b, + compile_key.lses_stride_h, + ), + ) + return dict( + outputs=output_ptr, + new_output=output_ptr, + lses=lse_ptr, + vlse=lse_ptr, + lse_idx=compile_key.lse_idx, + ctx=None, + is_base_e=compile_key.is_base_e, + ) + + @kernel_launcher + def __call__( + self, + outputs: torch.Tensor, + new_output: torch.Tensor, + lses: torch.Tensor, + vlse: torch.Tensor, + lse_idx: int, + ctx: Any, + *, + is_base_e: bool, + ) -> LaunchSpec: + num_tokens, num_heads, head_dim = outputs.shape + n_rounded = lses.shape[0] + outputs_stride_b, outputs_stride_h, outputs_stride_d = outputs.stride() + lses_stride_n, lses_stride_b, lses_stride_h = lses.stride() + grid = (num_tokens, num_heads, 1) + return grid, dict( + outputs_stride_B=outputs_stride_b, + outputs_stride_H=outputs_stride_h, + outputs_stride_D=outputs_stride_d, + lses_stride_N=lses_stride_n, + lses_stride_B=lses_stride_b, + lses_stride_H=lses_stride_h, + HEAD_DIM=head_dim, + N_ROUNDED=n_rounded, + IS_BASE_E=is_base_e, + _runtime_launcher=None if self._warming else ctx.call_kernel, + # CPTritonContext caches the non-constexpr positional prefix; derive + # its length so adding/reordering a kernel arg cannot silently + # misalign the cached replay path. + _runtime_launcher_arg_count=len(self.kernel.arg_names) - len(self.kernel.constexprs), + ) class CPTritonContext: @@ -179,37 +345,26 @@ def correct_attn_out( lses = lses.squeeze(1) assert lses.ndim == 3, f"expected lses [N,B,H] (optionally with a 1-sized extra dim), got {tuple(lses.shape)}" - B, H, D = out.shape - N = lses.shape[0] + B, H, _ = out.shape # Strides after we normalized shapes to 3-D views. The kernel computes # offsets for `vlse_ptr` using lses_stride_B/H, so the output buffer must # have the same B/H stride layout as a slice of `lses`. - o_sB, o_sH, o_sD = out.stride() - l_sN, l_sB, l_sH = lses.stride() + _, l_sB, l_sH = lses.stride() # Allocate LSE with the same B/H strides as `lses` so writes land correctly # even when `lses` is a non-contiguous view (e.g., 4-D to 3-D squeeze). lse = torch.empty_strided((B, H), (l_sB, l_sH), device=lses.device, dtype=lses.dtype) - # Kernel launch config - grid = (B, H, 1) - - regular_args = ( + _CORRECT_ATTN_CP_OUT_KERNEL( out, out, lses, lse, - o_sB, - o_sH, - o_sD, - l_sN, - l_sB, - l_sH, cp_rank, + ctx, + is_base_e=is_lse_base_on_e, ) - const_args = {"HEAD_DIM": D, "N_ROUNDED": N, "IS_BASE_E": is_lse_base_on_e} - ctx.call_kernel(_correct_attn_cp_out_kernel, grid, *regular_args, **const_args) return out, lse @@ -1167,6 +1322,14 @@ def __init__( is_lse_base_on_e, use_pcp, ) + if aphrodite_config.kernel_config.enable_jit_warmup and not self.use_a2a: + _CORRECT_ATTN_CP_OUT_KERNEL.register_warmup( + aphrodite_config, + output_dtype=output_dtype, + num_heads=num_heads, + head_dim=output_head_dim, + is_base_e=is_lse_base_on_e, + ) self.query_gather = ( None if use_pcp @@ -1318,3 +1481,6 @@ def kv_gather( local_kv: torch.Tensor, ) -> object: return self._kv_gather(gathered_kv, local_kv) + + +_CORRECT_ATTN_CP_OUT_KERNEL = CorrectAttnCPOutKernel() diff --git a/aphrodite/v1/worker/block_table.py b/aphrodite/v1/worker/block_table.py index 4025bf4071..d633a86f79 100644 --- a/aphrodite/v1/worker/block_table.py +++ b/aphrodite/v1/worker/block_table.py @@ -11,11 +11,11 @@ from aphrodite.distributed import get_dcp_group, get_pcp_group from aphrodite.logger import init_logger -from aphrodite.model_executor.warmup.jit_warmup import ( - AphroditeJitKernel, -) from aphrodite.model_executor.warmup.jit_warmup_triton_helper import ( + AphroditeTritonJitKernel, + LaunchSpec, TritonWarmupTensor, + kernel_launcher, triton_scalar_specialization_rep, ) from aphrodite.triton_utils import tl, triton @@ -372,7 +372,7 @@ def __getitem__(self, idx: int) -> "BlockTable": return self.block_tables[idx] -class ComputeSlotMappingKernel(AphroditeJitKernel["ComputeSlotMappingKernel.CompileKey"]): +class ComputeSlotMappingKernel(AphroditeTritonJitKernel["ComputeSlotMappingKernel.CompileKey"]): triton_block_size = 1024 @dataclass(frozen=True) @@ -460,37 +460,53 @@ def dispatch( # type: ignore[override] def get_warmup_keys(self, **dispatch_kwargs: int) -> list[CompileKey]: return self._trace_dispatch(self.dispatch)(**dispatch_kwargs) - def compile(self, compile_key: CompileKey) -> None: - warmup = getattr(self.kernel, "warmup", None) - assert warmup is not None + def warmup_inputs(self, compile_key: CompileKey) -> dict[str, Any]: int32_ptr = TritonWarmupTensor(torch.int32) int64_ptr = TritonWarmupTensor(torch.int64) - warmup( - 2, # arbitrary, num_tokens in do_not_specialize - 2, # arbitrary, max_num_tokens in do_not_specialize - int32_ptr, - int64_ptr, - int32_ptr, - compile_key.block_table_stride, - compile_key.block_size, - int64_ptr, - KV_CACHE_BLOCK_SIZE=compile_key.kv_cache_block_size, - BLOCKS_PER_KV_BLOCK=compile_key.blocks_per_kv_block, - TOTAL_CP_WORLD_SIZE=compile_key.total_cp_world_size, - TOTAL_CP_RANK=compile_key.total_cp_rank, - CP_KV_CACHE_INTERLEAVE_SIZE=compile_key.cp_kv_cache_interleave_size, - PAD_ID=PAD_SLOT_ID, - BLOCK_SIZE=self.triton_block_size, - grid=(2,), + return dict( + num_reqs=1, + num_tokens=2, # arbitrary, in do_not_specialize + max_num_tokens=2, # arbitrary, in do_not_specialize + query_start_loc=int32_ptr, + positions=int64_ptr, + block_table=TritonWarmupTensor( + torch.int32, + shape=(1, compile_key.block_table_stride), + ), + block_table_stride=compile_key.block_table_stride, + block_size=compile_key.block_size, + slot_mapping=int64_ptr, + kv_cache_block_size=compile_key.kv_cache_block_size, + blocks_per_kv_block=compile_key.blocks_per_kv_block, + total_cp_world_size=compile_key.total_cp_world_size, + total_cp_rank=compile_key.total_cp_rank, + cp_kv_cache_interleave_size=compile_key.cp_kv_cache_interleave_size, ) + @kernel_launcher def __call__( self, num_reqs: int, - *args: Any, - ) -> None: - self.kernel[(num_reqs + 1,)]( - *args, + num_tokens: int, + max_num_tokens: int, + query_start_loc: torch.Tensor, + positions: torch.Tensor, + block_table: torch.Tensor, + block_table_stride: int, + block_size: int, + slot_mapping: torch.Tensor, + kv_cache_block_size: int, + blocks_per_kv_block: int, + total_cp_world_size: int, + total_cp_rank: int, + cp_kv_cache_interleave_size: int, + ) -> LaunchSpec: + return (num_reqs + 1,), dict( + KV_CACHE_BLOCK_SIZE=kv_cache_block_size, + BLOCKS_PER_KV_BLOCK=blocks_per_kv_block, + TOTAL_CP_WORLD_SIZE=total_cp_world_size, + TOTAL_CP_RANK=total_cp_rank, + CP_KV_CACHE_INTERLEAVE_SIZE=cp_kv_cache_interleave_size, PAD_ID=PAD_SLOT_ID, BLOCK_SIZE=self.triton_block_size, ) diff --git a/tests/model_executor/test_jit_warmup.py b/tests/model_executor/test_jit_warmup.py index b3f4d81ca8..c28e44f2d8 100644 --- a/tests/model_executor/test_jit_warmup.py +++ b/tests/model_executor/test_jit_warmup.py @@ -528,10 +528,12 @@ def __call__(self, value: int) -> None: def test_registry_records_only_inside_model_setup_context() -> None: registry = JitWarmupRegistry(_config()) kernel = RecordingToyKernel() + config = _config() kernel.register_warmup(3, _config()) with registry.activate(): - kernel.register_warmup(3, _config()) + kernel.register_warmup(3, config) + kernel.register_warmup(3, config) assert len(registry) == 1 assert kernel.compiled == [] diff --git a/tests/model_executor/test_jit_warmup_cutedsl_launcher.py b/tests/model_executor/test_jit_warmup_cutedsl_launcher.py new file mode 100644 index 0000000000..b07949a07f --- /dev/null +++ b/tests/model_executor/test_jit_warmup_cutedsl_launcher.py @@ -0,0 +1,76 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project + +from dataclasses import dataclass +from typing import Any + +from aphrodite.model_executor.warmup import jit_warmup_cutedsl_helper +from aphrodite.model_executor.warmup.jit_warmup_cutedsl_helper import ( + AphroditeCuTeDSLJitKernel, + CuTeDSLLaunchSpec, + cutedsl_kernel_launcher, +) + + +class _TestCuTeDSLKernel(AphroditeCuTeDSLJitKernel["_TestCuTeDSLKernel.CompileKey"]): + @dataclass(frozen=True) + class CompileKey: + variant: int + + @staticmethod + def kernel(compile_key: CompileKey) -> tuple[str, int]: + return "entry", compile_key.variant + + def dispatch(self, *, variant: int) -> CompileKey: + return self.CompileKey(variant=variant) + + def get_warmup_keys(self) -> list[CompileKey]: + return self._trace_dispatch(self.dispatch)(variant=(1, 2)) + + def warmup_inputs(self, compile_key: CompileKey) -> tuple[Any, ...]: + return (f"fake-{compile_key.variant}",) + + @cutedsl_kernel_launcher + def __call__( + self, + payload: str, + *, + variant: int, + ) -> CuTeDSLLaunchSpec[CompileKey]: + compile_key = self.dispatch(variant=variant) + return compile_key, (payload,), {"variant": variant} + + +def test_cutedsl_launcher_reuses_compiled_executor(monkeypatch) -> None: + compile_calls: list[tuple[Any, tuple[Any, ...]]] = [] + launch_calls: list[tuple[int, tuple[Any, ...]]] = [] + + def compile_cutedsl(entry: tuple[str, int], *args: Any) -> Any: + compile_calls.append((entry, args)) + variant = entry[1] + + def executor(*runtime_args: Any) -> tuple[int, tuple[Any, ...]]: + launch_calls.append((variant, runtime_args)) + return variant, runtime_args + + return executor + + monkeypatch.setattr( + jit_warmup_cutedsl_helper, + "compile_cutedsl", + compile_cutedsl, + ) + kernel = _TestCuTeDSLKernel() + compile_key = kernel.CompileKey(variant=1) + + kernel.compile(compile_key) + kernel.compile(compile_key) + assert compile_calls == [(("entry", 1), ("fake-1",))] + + assert kernel("runtime", variant=1) == (1, ("runtime",)) + assert kernel("other", variant=2) == (2, ("other",)) + assert compile_calls == [ + (("entry", 1), ("fake-1",)), + (("entry", 2), ("fake-2",)), + ] + assert launch_calls == [(1, ("runtime",)), (2, ("other",))] diff --git a/tests/model_executor/test_jit_warmup_triton_launcher.py b/tests/model_executor/test_jit_warmup_triton_launcher.py new file mode 100644 index 0000000000..a58799e12d --- /dev/null +++ b/tests/model_executor/test_jit_warmup_triton_launcher.py @@ -0,0 +1,163 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project + +from dataclasses import dataclass +from typing import Any + +from aphrodite.model_executor.warmup.jit_warmup_triton_helper import ( + AphroditeTritonJitKernel, + LaunchSpec, + kernel_launcher, +) + + +class _FakeTritonKernel: + arg_names = ("first", "second", "CONST") + + def __init__(self) -> None: + self.warmup_calls: list[dict[str, Any]] = [] + + def warmup(self, **kwargs: Any) -> None: + self.warmup_calls.append(kwargs) + + +class _TestTritonKernel(AphroditeTritonJitKernel["_TestTritonKernel.CompileKey"]): + kernel = _FakeTritonKernel() + + @dataclass(frozen=True) + class CompileKey: + value: int + + def dispatch(self, *, value: int) -> CompileKey: + return self.CompileKey(value=value) + + def get_warmup_keys(self) -> list[CompileKey]: + return self._trace_dispatch(self.dispatch)(value=1) + + def warmup_inputs(self, compile_key: CompileKey) -> dict[str, Any]: + return dict(first="warmup", second=compile_key.value, runtime_launcher=None) + + @kernel_launcher + def __call__( + self, + first: str, + second: int, + runtime_launcher: Any, + ) -> LaunchSpec: + return (2,), dict( + CONST=7, + _runtime_launcher=runtime_launcher, + _runtime_launcher_arg_count=2, + ) + + +def test_triton_launcher_supports_compile_and_runtime_adapters() -> None: + owner = _TestTritonKernel() + owner.kernel.warmup_calls.clear() + + owner.compile(owner.CompileKey(value=1)) + assert owner.kernel.warmup_calls == [{"grid": (1,), "first": "warmup", "second": 1, "CONST": 7}] + + runtime_calls: list[tuple[Any, ...]] = [] + + def runtime_launcher(kernel: Any, grid: Any, *args: Any, **kwargs: Any) -> None: + runtime_calls.append((kernel, grid, args, kwargs)) + + owner("runtime", 2, runtime_launcher) + assert runtime_calls == [(owner.kernel, (2,), ("runtime", 2), {"CONST": 7})] + + +def test_triton_launcher_supports_cpu_function_wrappers() -> None: + calls: list[tuple[Any, ...]] = [] + + def kernel(first: str, second: int, CONST: int) -> None: + calls.append((first, second, CONST)) + + class FuncWrapper: + def __init__(self) -> None: + self.func = kernel + + def __getitem__(self, _grid: Any) -> Any: + return self.func + + class TestCpuKernel(_TestTritonKernel): + kernel: Any = FuncWrapper() + + owner = TestCpuKernel() + + owner("runtime", 2, None) + assert calls == [("runtime", 2, 7)] + + +def test_compute_slot_mapping_uses_named_launcher_inputs(monkeypatch) -> None: + from aphrodite.v1.attention.backends.utils import PAD_SLOT_ID + from aphrodite.v1.worker.block_table import ComputeSlotMappingKernel + + owner = ComputeSlotMappingKernel() + compile_key = owner.CompileKey( + kv_cache_block_size=16, + blocks_per_kv_block=1, + total_cp_world_size=2, + total_cp_rank=1, + cp_kv_cache_interleave_size=1, + block_table_stride=128, + block_size=16, + ) + launches: list[tuple[Any, ...]] = [] + + def launch(grid: Any, inputs: Any, **kwargs: Any) -> None: + launches.append((grid, inputs, kwargs)) + + monkeypatch.setattr(owner, "launch", launch) + owner.compile(compile_key) + + grid, inputs, kwargs = launches[0] + assert grid == (2,) + assert inputs["num_tokens"] == 2 + assert inputs["block_table_stride"] == 128 + assert kwargs == { + "KV_CACHE_BLOCK_SIZE": 16, + "BLOCKS_PER_KV_BLOCK": 1, + "TOTAL_CP_WORLD_SIZE": 2, + "TOTAL_CP_RANK": 1, + "CP_KV_CACHE_INTERLEAVE_SIZE": 1, + "PAD_ID": PAD_SLOT_ID, + "BLOCK_SIZE": owner.triton_block_size, + } + + +def test_compute_slot_mapping_supports_cpu_adapter(monkeypatch) -> None: + from types import SimpleNamespace + + from aphrodite.utils import cpu_triton_utils + from aphrodite.v1.worker.block_table import ComputeSlotMappingKernel + + calls: list[tuple[Any, ...]] = [] + monkeypatch.setattr( + cpu_triton_utils, + "torch", + SimpleNamespace( + ops=SimpleNamespace( + _C=SimpleNamespace(compute_slot_mapping_kernel_impl=lambda *args: calls.append(args)), + ), + ), + ) + owner = ComputeSlotMappingKernel() + owner.kernel = cpu_triton_utils.compute_slot_mapping_kernel + owner( + num_reqs=1, + num_tokens=2, + max_num_tokens=8, + query_start_loc="query", + positions="positions", + block_table="blocks", + block_table_stride=128, + block_size=16, + slot_mapping="slots", + kv_cache_block_size=16, + blocks_per_kv_block=1, + total_cp_world_size=1, + total_cp_rank=0, + cp_kv_cache_interleave_size=1, + ) + assert calls == [("query", "positions", "blocks", "slots", 16)] From 1a85cdab22dc469921b5b1eea6067894ecbd1f1b Mon Sep 17 00:00:00 2001 From: AlpinDale Date: Mon, 7 Sep 2026 21:04:37 +0000 Subject: [PATCH 31/36] [sync] [3/N][warmup][DSv4] Migrate FA4 MLA and shared CuTeDSL kernels (#53565) Upstream-vLLM: 42801b3a6b3bb92397a04117449f4699b09db1df Co-authored-by: Roberto L. Castro <38211239+LopezCastroRoberto@users.noreply.github.com> --- .sync/vllm-sha | 2 +- .../warmup/fa4_cutedsl_warmup.py | 53 ------------------- .../model_executor/warmup/kernel_warmup.py | 4 -- aphrodite/models/inkling/nvidia/attention.py | 7 ++- .../models/inkling/nvidia/ops/__init__.py | 4 +- .../inkling/nvidia/ops/fa4_rel_attention.py | 2 +- .../backends/mla/prefill/flash_attn.py | 7 ++- .../models/inkling/test_fa4_rel_attention.py | 6 +-- .../test_mla_prefill_quant_output.py | 31 +++++++++++ 9 files changed, 48 insertions(+), 68 deletions(-) delete mode 100644 aphrodite/model_executor/warmup/fa4_cutedsl_warmup.py diff --git a/.sync/vllm-sha b/.sync/vllm-sha index f1dd3b8b83..c48ea8b058 100644 --- a/.sync/vllm-sha +++ b/.sync/vllm-sha @@ -1 +1 @@ -1713b9866ae5c380d97da5796217064ac7ab57b8 +42801b3a6b3bb92397a04117449f4699b09db1df diff --git a/aphrodite/model_executor/warmup/fa4_cutedsl_warmup.py b/aphrodite/model_executor/warmup/fa4_cutedsl_warmup.py deleted file mode 100644 index 7249989454..0000000000 --- a/aphrodite/model_executor/warmup/fa4_cutedsl_warmup.py +++ /dev/null @@ -1,53 +0,0 @@ -# SPDX-License-Identifier: Apache-2.0 -# SPDX-FileCopyrightText: Copyright contributors to the vLLM project -"""Warm up FA4 CuTeDSL kernels.""" - -from __future__ import annotations - -from typing import TYPE_CHECKING - -from aphrodite.v1.attention.backends.mla.prefill import get_mla_prefill_backend - -if TYPE_CHECKING: - from aphrodite.v1.worker.gpu_worker import Worker - - -def _warm_fa4_mla_prefill(worker: Worker) -> None: - runner = worker.model_runner - if runner.is_pooling_model: - return - - aphrodite_config = runner.aphrodite_config - if not aphrodite_config.model_config.use_mla: - return - - try: - backend_cls = get_mla_prefill_backend(aphrodite_config) - except ValueError: - # fall back to top-k MQA prefill path. - return - if backend_cls.get_name() != "FLASH_ATTN": - return - - from aphrodite.v1.attention.backends.mla.prefill import flash_attn - - flash_attn.FA4_MLA_PREFILL_KERNEL.warmup(aphrodite_config) - - -def _warm_inkling_fa4_rel_attention(worker: Worker) -> None: - from aphrodite.models.inkling.configs import InklingMMConfig, InklingModelConfig - from aphrodite.models.inkling.nvidia.ops.fa4_rel_attention import ( - INKLING_FA4_REL_ATTENTION_KERNEL, - ) - - aphrodite_config = worker.aphrodite_config - hf_config = aphrodite_config.model_config.hf_config - if not isinstance(hf_config, (InklingMMConfig, InklingModelConfig)): - return - - INKLING_FA4_REL_ATTENTION_KERNEL.warmup(aphrodite_config) - - -def fa4_cutedsl_warmup(worker: Worker) -> None: - _warm_fa4_mla_prefill(worker) - _warm_inkling_fa4_rel_attention(worker) diff --git a/aphrodite/model_executor/warmup/kernel_warmup.py b/aphrodite/model_executor/warmup/kernel_warmup.py index f5db4a6176..1f082fabb9 100644 --- a/aphrodite/model_executor/warmup/kernel_warmup.py +++ b/aphrodite/model_executor/warmup/kernel_warmup.py @@ -20,9 +20,6 @@ from aphrodite.model_executor.warmup.deepseek_v4_mhc_warmup import ( deepseek_v4_mhc_warmup, ) -from aphrodite.model_executor.warmup.fa4_cutedsl_warmup import ( - fa4_cutedsl_warmup, -) from aphrodite.model_executor.warmup.flashinfer_autotune_cache import ( resolve_flashinfer_autotune_file, write_flashinfer_autotune_cache, @@ -173,7 +170,6 @@ def kernel_warmup(worker: "Worker", *, process_local_only: bool = False): # Run next so input-prep kernels JIT against pristine runner state. if worker.aphrodite_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) diff --git a/aphrodite/models/inkling/nvidia/attention.py b/aphrodite/models/inkling/nvidia/attention.py index 848bb8b0d2..8dbf1baba0 100644 --- a/aphrodite/models/inkling/nvidia/attention.py +++ b/aphrodite/models/inkling/nvidia/attention.py @@ -35,7 +35,7 @@ from ..configs import InklingModelConfig from .layernorm import InklingRMSNorm from .ops.fa4_rel_attention import ( - INKLING_FA4_REL_ATTENTION_KERNEL, + _INKLING_FA4_REL_ATTENTION_KERNEL, bucket_max_seqlen_q, inkling_fa4_num_splits, ) @@ -160,6 +160,9 @@ def __init__( compilation_config.static_forward_context[prefix] = self self.kv_cache = torch.tensor([]) # replaced by bind_kv_cache + if aphrodite_config.kernel_config.enable_jit_warmup: + _INKLING_FA4_REL_ATTENTION_KERNEL.register_warmup() + def get_attn_backend(self) -> type[AttentionBackend]: return FlashAttentionBackend @@ -273,7 +276,7 @@ def _attention( num_kv_heads=self.num_kv_heads, max_kv_len=self._max_kv_len, ) - INKLING_FA4_REL_ATTENTION_KERNEL( + _INKLING_FA4_REL_ATTENTION_KERNEL( q[:nt], key_cache, value_cache, diff --git a/aphrodite/models/inkling/nvidia/ops/__init__.py b/aphrodite/models/inkling/nvidia/ops/__init__.py index bb7f066ddb..fe6bc87e30 100644 --- a/aphrodite/models/inkling/nvidia/ops/__init__.py +++ b/aphrodite/models/inkling/nvidia/ops/__init__.py @@ -15,12 +15,12 @@ _LAZY_EXPORTS = { "silu_and_mul_triton": "silu_and_mul", "sink_silu_mul_epilogue": "silu_and_mul", - "INKLING_FA4_REL_ATTENTION_KERNEL": "fa4_rel_attention", + "_INKLING_FA4_REL_ATTENTION_KERNEL": "fa4_rel_attention", } if TYPE_CHECKING: from .fa4_rel_attention import ( # noqa: F401 - INKLING_FA4_REL_ATTENTION_KERNEL, + _INKLING_FA4_REL_ATTENTION_KERNEL, ) from .silu_and_mul import ( # noqa: F401 silu_and_mul_triton, diff --git a/aphrodite/models/inkling/nvidia/ops/fa4_rel_attention.py b/aphrodite/models/inkling/nvidia/ops/fa4_rel_attention.py index b0d3e06a35..bde129e005 100644 --- a/aphrodite/models/inkling/nvidia/ops/fa4_rel_attention.py +++ b/aphrodite/models/inkling/nvidia/ops/fa4_rel_attention.py @@ -425,4 +425,4 @@ def __call__( ) -INKLING_FA4_REL_ATTENTION_KERNEL = InklingFA4RelAttentionKernel() +_INKLING_FA4_REL_ATTENTION_KERNEL = InklingFA4RelAttentionKernel() diff --git a/aphrodite/v1/attention/backends/mla/prefill/flash_attn.py b/aphrodite/v1/attention/backends/mla/prefill/flash_attn.py index 8d0fea4a98..56915dc750 100644 --- a/aphrodite/v1/attention/backends/mla/prefill/flash_attn.py +++ b/aphrodite/v1/attention/backends/mla/prefill/flash_attn.py @@ -326,6 +326,9 @@ def __init__( if self.flash_attn_version is not None: self.flash_attn_varlen_func = functools.partial(flash_attn_varlen_func, fa_version=self.flash_attn_version) + if aphrodite_config.kernel_config.enable_jit_warmup and self.flash_attn_version == 4: + _FA4_MLA_PREFILL_KERNEL.register_warmup() + # Determine if we need to pad V # For MLA the v head dim is smaller than qk head dim so we pad out # v with 0s to match the qk head dim for attention backends that do @@ -377,7 +380,7 @@ def _flash_attn_varlen_diff_headdims( if envs.APHRODITE_BATCH_INVARIANT: kwargs["num_splits"] = 1 - attn_out = FA4_MLA_PREFILL_KERNEL( + attn_out = _FA4_MLA_PREFILL_KERNEL( q=q, k=k, v=maybe_padded_v, @@ -449,4 +452,4 @@ def run_prefill_context_chunk( ) -FA4_MLA_PREFILL_KERNEL = FA4MLAPrefillKernel() +_FA4_MLA_PREFILL_KERNEL = FA4MLAPrefillKernel() diff --git a/tests/models/inkling/test_fa4_rel_attention.py b/tests/models/inkling/test_fa4_rel_attention.py index 56d9747145..26db6d8788 100644 --- a/tests/models/inkling/test_fa4_rel_attention.py +++ b/tests/models/inkling/test_fa4_rel_attention.py @@ -2,7 +2,7 @@ # SPDX-FileCopyrightText: Copyright contributors to the vLLM project """Correctness test for the Inkling FA4 relative-attention score-mod kernel. -Checks ``INKLING_FA4_REL_ATTENTION_KERNEL`` against a pure-PyTorch reference +Checks ``_INKLING_FA4_REL_ATTENTION_KERNEL`` against a pure-PyTorch reference that implements the relative bias exactly as documented in the Inkling architecture guide:: @@ -23,7 +23,7 @@ compute_log_scaling_tau, ) from aphrodite.models.inkling.nvidia.ops.fa4_rel_attention import ( - INKLING_FA4_REL_ATTENTION_KERNEL, + _INKLING_FA4_REL_ATTENTION_KERNEL, _use_sheared_bias, inkling_fa4_num_splits, ) @@ -271,7 +271,7 @@ def _run_case(seq_lens, num_heads, num_kv_heads, rel_extent, window_left, seed=0 num_kv_heads=num_kv_heads, max_kv_len=max(kv_lens), ) - out = INKLING_FA4_REL_ATTENTION_KERNEL( + out = _INKLING_FA4_REL_ATTENTION_KERNEL( q, key_cache, value_cache, diff --git a/tests/v1/attention/test_mla_prefill_quant_output.py b/tests/v1/attention/test_mla_prefill_quant_output.py index 3cc76ec25d..ee12376856 100644 --- a/tests/v1/attention/test_mla_prefill_quant_output.py +++ b/tests/v1/attention/test_mla_prefill_quant_output.py @@ -10,6 +10,7 @@ + standalone static-FP8-quant path it replaces (GPU-only, SM100/SM110). """ +from types import SimpleNamespace from unittest.mock import patch import pytest @@ -93,6 +94,36 @@ def test_flash_attn_supports_quant_output_unknown_device(): assert backend.supports_quant_output(kFp8StaticTensorSym) is False +@pytest.mark.parametrize( + ("version", "enable_jit_warmup", "expected_calls"), + [(4, True, 1), (4, False, 0), (3, True, 0), (2, True, 0), (None, True, 0)], +) +def test_flash_attn_registers_warmup_only_for_fa4( + version: int | None, + enable_jit_warmup: bool, + expected_calls: int, +): + aphrodite_config = SimpleNamespace(kernel_config=SimpleNamespace(enable_jit_warmup=enable_jit_warmup)) + with ( + patch(f"{_FA_MODULE}.flash_attn_varlen_func"), + patch(f"{_FA_MODULE}.get_flash_attn_version", return_value=version), + patch(f"{_FA_MODULE}._FA4_MLA_PREFILL_KERNEL.register_warmup") as register, + patch(f"{_FA_MODULE}.current_platform") as platform, + ): + platform.get_device_capability.return_value = DeviceCapability(10, 0) + FlashAttnPrefillBackend( + num_heads=16, + scale=1.0, + kv_lora_rank=512, + qk_nope_head_dim=128, + qk_rope_head_dim=64, + v_head_dim=128, + aphrodite_config=aphrodite_config, + ) + + assert register.call_count == expected_calls + + def test_flash_attn_prefill_backend_signature_accepts_fused_kwargs(): """run_prefill_new_tokens must accept out/output_scale so the direct (non-**kwargs) call in forward_mha type- and runtime-checks.""" From c68266c6e73c6693c6df5415daadff3f2a9b7156 Mon Sep 17 00:00:00 2001 From: AlpinDale Date: Mon, 7 Sep 2026 21:08:46 +0000 Subject: [PATCH 32/36] [sync] [Bugfix][Gemma] Conditionally create KV projections/norms on KV-shared layers (#54917) Upstream-vLLM: f2d45f26bd6a2c841ffbd0030ebaefe87c274e58 Co-authored-by: Asaf Gardin <39553475+Josephasafg@users.noreply.github.com> --- .sync/vllm-sha | 2 +- aphrodite/model_executor/models/gemma3n.py | 108 ++++++++++----- aphrodite/model_executor/models/gemma4.py | 129 ++++++++++-------- .../models/language/generation/test_gemma.py | 59 ++++++++ 4 files changed, 204 insertions(+), 94 deletions(-) diff --git a/.sync/vllm-sha b/.sync/vllm-sha index c48ea8b058..364ed64c64 100644 --- a/.sync/vllm-sha +++ b/.sync/vllm-sha @@ -1 +1 @@ -42801b3a6b3bb92397a04117449f4699b09db1df +f2d45f26bd6a2c841ffbd0030ebaefe87c274e58 diff --git a/aphrodite/model_executor/models/gemma3n.py b/aphrodite/model_executor/models/gemma3n.py index 71c3e530ee..6263ce5c16 100644 --- a/aphrodite/model_executor/models/gemma3n.py +++ b/aphrodite/model_executor/models/gemma3n.py @@ -297,15 +297,32 @@ def __init__( self.q_size = self.num_heads * self.head_dim self.kv_size = self.num_kv_heads * self.head_dim - self.qkv_proj = QKVParallelLinear( - hidden_size, - self.head_dim, - self.total_num_heads, - self.total_num_kv_heads, - bias=config.attention_bias, - quant_config=quant_config, - prefix=f"{prefix}.qkv_proj", - ) + layer_idx = extract_layer_index(prefix) + first_kv_shared_layer_idx = config.num_hidden_layers - config.num_kv_shared_layers + self.is_kv_shared_layer = layer_idx >= first_kv_shared_layer_idx + + if self.is_kv_shared_layer: + # K/V come from the target layer's cache, so like HF this layer + # only has q_proj (and no k_norm/v_norm). + self.q_proj = ColumnParallelLinear( + hidden_size, + self.total_num_heads * self.head_dim, + bias=config.attention_bias, + quant_config=quant_config, + prefix=f"{prefix}.q_proj", + ) + else: + self.qkv_proj = QKVParallelLinear( + hidden_size, + self.head_dim, + self.total_num_heads, + self.total_num_kv_heads, + bias=config.attention_bias, + quant_config=quant_config, + prefix=f"{prefix}.qkv_proj", + ) + self.k_norm = RMSNorm(hidden_size=self.head_dim, eps=config.rms_norm_eps) + self.v_norm = RMSNorm(hidden_size=self.head_dim, eps=config.rms_norm_eps, has_weight=False) self.o_proj = RowParallelLinear( self.total_num_heads * self.head_dim, hidden_size, @@ -314,10 +331,7 @@ def __init__( prefix=f"{prefix}.o_proj", ) self.q_norm = RMSNorm(hidden_size=self.head_dim, eps=config.rms_norm_eps) - self.k_norm = RMSNorm(hidden_size=self.head_dim, eps=config.rms_norm_eps) - self.v_norm = RMSNorm(hidden_size=self.head_dim, eps=config.rms_norm_eps, has_weight=False) - layer_idx = extract_layer_index(prefix) layer_type = config.layer_types[layer_idx] is_sliding = layer_type == "sliding_attention" self.sliding_window = config.sliding_window if is_sliding else None @@ -334,11 +348,8 @@ def __init__( if is_sliding: rope_parameters["rope_theta"] = config.rope_local_base_freq - first_kv_shared_layer_idx = config.num_hidden_layers - config.num_kv_shared_layers - self.is_kv_shared = layer_idx >= first_kv_shared_layer_idx - kv_sharing_target_layer_name = None - if self.is_kv_shared: + if self.is_kv_shared_layer: # Last full attention layer is 1 before sharing # Last sliding attention layer is 2 before sharing offset = 2 if self.sliding_window is not None else 1 @@ -393,21 +404,30 @@ def forward( hidden_states: torch.Tensor, **kwargs, ) -> torch.Tensor: - qkv, _ = self.qkv_proj(hidden_states) - q, k, v = qkv.split([self.q_size, self.kv_size, self.kv_size], dim=-1) - - q = q.unflatten(-1, (self.num_heads, self.head_dim)) - q = self.q_norm(q) - q = q.flatten(-2, -1) - k = k.unflatten(-1, (self.num_kv_heads, self.head_dim)) - k = self.k_norm(k) - k = k.flatten(-2, -1) - v = v.unflatten(-1, (self.num_kv_heads, self.head_dim)) - v = self.v_norm(v) - v = v.flatten(-2, -1) - - q, k = self.rotary_emb(positions, q, k) - attn_output = self.attn(q, k, v) + if self.is_kv_shared_layer: + # Shared KV: only Q is projected; K/V come from the target layer. + q, _ = self.q_proj(hidden_states) + q = q.unflatten(-1, (self.num_heads, self.head_dim)) + q = self.q_norm(q) + q = q.flatten(-2, -1) + q, _ = self.rotary_emb(positions, q, None) + attn_output = self.attn(q, None, None) + else: + qkv, _ = self.qkv_proj(hidden_states) + q, k, v = qkv.split([self.q_size, self.kv_size, self.kv_size], dim=-1) + + q = q.unflatten(-1, (self.num_heads, self.head_dim)) + q = self.q_norm(q) + q = q.flatten(-2, -1) + k = k.unflatten(-1, (self.num_kv_heads, self.head_dim)) + k = self.k_norm(k) + k = k.flatten(-2, -1) + v = v.unflatten(-1, (self.num_kv_heads, self.head_dim)) + v = self.v_norm(v) + v = v.flatten(-2, -1) + + q, k = self.rotary_emb(positions, q, k) + attn_output = self.attn(q, k, v) output, _ = self.o_proj(attn_output) return output @@ -766,6 +786,28 @@ def forward( return hidden_states +def _kv_sharing_weights_mapper(config: Gemma3nTextConfig) -> WeightsMapper: + """KV-shared layers only have q_proj, so qkv_proj packing applies to the + other layers. Original checkpoints still ship K/V tensors for the shared + layers (fine-tuned ones omit them); those are dropped.""" + first_kv_shared_layer_idx = config.num_hidden_layers - config.num_kv_shared_layers + return WeightsMapper( + orig_to_new_substr={ + f"layers.{i}.self_attn.{name}.": None + for i in range(first_kv_shared_layer_idx, config.num_hidden_layers) + for name in ("k_proj", "v_proj", "k_norm") + }, + orig_to_new_stacked={ + f"layers.{i}.self_attn.{shard}_proj.": ( + f"layers.{i}.self_attn.qkv_proj.", + shard, + ) + for i in range(first_kv_shared_layer_idx) + for shard in ("q", "k", "v") + }, + ) + + # This disables torch.compile if --kv-sharing-fast-prefill passed @support_torch_compile(enable_if=lambda aphrodite_config: not aphrodite_config.cache_config.kv_sharing_fast_prefill) class Gemma3nTextModel(nn.Module, SupportsQuant): @@ -778,9 +820,6 @@ class Gemma3nTextModel(nn.Module, SupportsQuant): "altup_projections.": "self_decoder.altup_projections.", }, orig_to_new_stacked={ - ".q_proj": (".qkv_proj", "q"), - ".k_proj": (".qkv_proj", "k"), - ".v_proj": (".qkv_proj", "v"), ".gate_proj": (".gate_up_proj", 0), ".up_proj": (".gate_up_proj", 1), }, @@ -793,6 +832,7 @@ def __init__(self, *, aphrodite_config: AphroditeConfig, prefix: str = ""): quant_config = aphrodite_config.quant_config self.config = config self.quant_config = quant_config + self.hf_to_aphrodite_mapper = self.hf_to_aphrodite_mapper | _kv_sharing_weights_mapper(config) self.altup_unembed_projections = nn.ModuleList( [ diff --git a/aphrodite/model_executor/models/gemma4.py b/aphrodite/model_executor/models/gemma4.py index ef816b15da..bbca97ff1b 100644 --- a/aphrodite/model_executor/models/gemma4.py +++ b/aphrodite/model_executor/models/gemma4.py @@ -408,18 +408,37 @@ def __init__( # Q/K norms with learnable weights handle scaling implicitly. self.scaling = 1.0 - # QKVParallelLinear handles GQA correctly for all layer types. - # k_eq_v layers load K weights into both K and V slots via - # _weight_iterator remapping — no structural difference needed. - self.qkv_proj = QKVParallelLinear( - hidden_size, - self.head_dim, - self.total_num_heads, - self.total_num_kv_heads, - bias=config.attention_bias, - quant_config=quant_config, - prefix=f"{prefix}.qkv_proj", - ) + layer_idx = extract_layer_index(prefix) + num_kv_shared_layers = getattr(config, "num_kv_shared_layers", 0) + first_kv_shared_layer_idx = config.num_hidden_layers - num_kv_shared_layers + self.is_kv_shared_layer = num_kv_shared_layers > 0 and layer_idx >= first_kv_shared_layer_idx + + if self.is_kv_shared_layer: + # K/V come from the target layer's cache, so like HF this layer + # only has q_proj (and no k_norm/v_norm). + self.q_proj = ColumnParallelLinear( + hidden_size, + self.total_num_heads * self.head_dim, + bias=config.attention_bias, + quant_config=quant_config, + prefix=f"{prefix}.q_proj", + ) + else: + # QKVParallelLinear handles GQA correctly for all layer types. + # k_eq_v layers load K weights into both K and V slots via + # _weight_iterator remapping — no structural difference needed. + self.qkv_proj = QKVParallelLinear( + hidden_size, + self.head_dim, + self.total_num_heads, + self.total_num_kv_heads, + bias=config.attention_bias, + quant_config=quant_config, + prefix=f"{prefix}.qkv_proj", + ) + self.k_norm = RMSNorm(self.head_dim, eps=config.rms_norm_eps) + # V norm: no learnable scale (pure normalization only) + self.v_norm = RMSNorm(self.head_dim, eps=config.rms_norm_eps, has_weight=False) self.o_proj = RowParallelLinear( self.total_num_heads * self.head_dim, hidden_size, @@ -428,14 +447,10 @@ def __init__( prefix=f"{prefix}.o_proj", ) - # Q/K norms: output = norm(x) * weight (learnable per-head scale) + # Q norm: output = norm(x) * weight (learnable per-head scale) self.q_norm = RMSNorm(self.head_dim, eps=config.rms_norm_eps) - self.k_norm = RMSNorm(self.head_dim, eps=config.rms_norm_eps) - # V norm: no learnable scale (pure normalization only) - self.v_norm = RMSNorm(self.head_dim, eps=config.rms_norm_eps, has_weight=False) # Determine layer type and sliding window - layer_idx = extract_layer_index(prefix) layer_type = config.layer_types[layer_idx] self.is_sliding = layer_type == "sliding_attention" sliding_window = config.sliding_window if self.is_sliding else None @@ -459,26 +474,21 @@ def __init__( # KV sharing: layers in the last `num_kv_shared_layers` share KV # cache with earlier layers of the same type. kv_sharing_target_layer_name = None - self.is_kv_shared_layer = False - num_kv_shared_layers = getattr(config, "num_kv_shared_layers", 0) - if num_kv_shared_layers > 0: - first_kv_shared_layer_idx = config.num_hidden_layers - num_kv_shared_layers - if layer_idx >= first_kv_shared_layer_idx: - self.is_kv_shared_layer = True - # Find the last non-shared layer of the same attention type - prev_layers = config.layer_types[:first_kv_shared_layer_idx] - current_layer_type = config.layer_types[layer_idx] - kv_shared_layer_index = len(prev_layers) - 1 - prev_layers[::-1].index(current_layer_type) - if kv_shared_layer_index >= 0: - if ".layers." in prefix: - param_name_before_layers = prefix.split(".layers.")[0] - else: - raise ValueError( - f"Unexpected prefix format for Gemma4Attention: '{prefix}'. Expected to contain '.layers.'." - ) - kv_sharing_target_layer_name = ( - f"{param_name_before_layers}.layers.{kv_shared_layer_index}.self_attn.attn" + if self.is_kv_shared_layer: + # Find the last non-shared layer of the same attention type + prev_layers = config.layer_types[:first_kv_shared_layer_idx] + current_layer_type = config.layer_types[layer_idx] + kv_shared_layer_index = len(prev_layers) - 1 - prev_layers[::-1].index(current_layer_type) + if kv_shared_layer_index >= 0: + if ".layers." in prefix: + param_name_before_layers = prefix.split(".layers.")[0] + else: + raise ValueError( + f"Unexpected prefix format for Gemma4Attention: '{prefix}'. Expected to contain '.layers.'." ) + kv_sharing_target_layer_name = ( + f"{param_name_before_layers}.layers.{kv_shared_layer_index}.self_attn.attn" + ) self.rotary_emb = get_rope( self.head_dim, @@ -513,19 +523,24 @@ def forward( hidden_states: torch.Tensor, **kwargs, ) -> torch.Tensor: - # Unified QKV path (works for both k_eq_v and standard layers). - # For k_eq_v, K weights are loaded into both K and V slots of - # qkv_proj, so V == K automatically. - qkv, _ = self.qkv_proj(hidden_states) - q, k, v = qkv.split([self.q_size, self.kv_size, self.kv_size], dim=-1) - - # Q norm (always applied) - q = q.unflatten(-1, (self.num_heads, self.head_dim)) - q = self.q_norm(q) - q = q.flatten(-2, -1) - - if not self.is_kv_shared_layer: - # Non-shared: apply K norm + RoPE, V norm + if self.is_kv_shared_layer: + # Shared KV: only Q is projected; K/V come from the target layer. + q, _ = self.q_proj(hidden_states) + q = q.unflatten(-1, (self.num_heads, self.head_dim)) + q = self.q_norm(q) + q = q.flatten(-2, -1) + q, _ = self.rotary_emb(positions, q, None) + attn_output = self.attn(q, None, None) + else: + # For k_eq_v, K weights are loaded into both K and V slots of + # qkv_proj, so V == K automatically. + qkv, _ = self.qkv_proj(hidden_states) + q, k, v = qkv.split([self.q_size, self.kv_size, self.kv_size], dim=-1) + + q = q.unflatten(-1, (self.num_heads, self.head_dim)) + q = self.q_norm(q) + q = q.flatten(-2, -1) + k = k.unflatten(-1, (self.num_kv_heads, self.head_dim)) k = self.k_norm(k) k = k.flatten(-2, -1) @@ -534,11 +549,8 @@ def forward( v = v.unflatten(-1, (self.num_kv_heads, self.head_dim)) v = self.v_norm(v) v = v.flatten(-2, -1) - else: - # Shared: only apply RoPE to Q - q = self.rotary_emb(positions, q, k)[0] + attn_output = self.attn(q, k, v) - attn_output = self.attn(q, k, v) output, _ = self.o_proj(attn_output) return output @@ -1332,9 +1344,9 @@ def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]: if shard_name not in name: continue stacked_name = name.replace(shard_name, param_name) - # k_eq_v layers use separate q_proj/k_proj instead of - # packed qkv_proj. If the stacked param doesn't exist, - # skip this mapping and fall through to direct load. + # KV-shared layers have a plain q_proj instead of a packed + # qkv_proj: fall through to the direct load, which also + # skips the redundant K/V tensors original checkpoints ship. if stacked_name not in params_dict: continue if is_pp_missing_parameter(stacked_name, self): @@ -1422,9 +1434,8 @@ class Gemma4ForCausalLM(nn.Module, SupportsLoRA, SupportsPP, MixtureOfExperts, S ".moe.experts.down_proj": ".moe.down_proj", }, ) - # Note: qkv_proj packing applies to non-k_eq_v layers (sliding - # attention and full attention without k_eq_v). k_eq_v layers use - # separate q_proj + k_proj without packing. + # KV-shared layers only have q_proj, so qkv_proj packing applies to + # the other layers. packed_modules_mapping = { "qkv_proj": [ "q_proj", diff --git a/tests/models/language/generation/test_gemma.py b/tests/models/language/generation/test_gemma.py index 2aa21d848c..873707e9bd 100644 --- a/tests/models/language/generation/test_gemma.py +++ b/tests/models/language/generation/test_gemma.py @@ -10,6 +10,11 @@ from aphrodite.config import AphroditeConfig from aphrodite.model_executor.layers.vocab_parallel_embedding import VocabParallelEmbedding from aphrodite.model_executor.models import gemma +from aphrodite.model_executor.models.gemma3n import ( + Gemma3nTextModel, + _kv_sharing_weights_mapper, +) +from aphrodite.model_executor.models.gemma4 import Gemma4Model MODELS = ["google/gemma-2b", "google/gemma-2-2b", "google/gemma-3-4b-it"] @@ -51,6 +56,60 @@ def __init__(self, *, aphrodite_config, prefix): assert torch.equal(model.lm_head.weight[:4], lm_head_weight) +@pytest.mark.cpu_test +def test_gemma4_kv_shared_layer_loads_plain_q_proj() -> None: + """KV-shared layers have q_proj instead of a packed qkv_proj; their + redundant K/V tensors in original checkpoints have no parameter and are + skipped rather than failing the load.""" + model = torch.nn.Module() + model.config = SimpleNamespace(num_experts=0) + model.start_layer, model.end_layer = 0, 2 + model.layers = torch.nn.ModuleList([torch.nn.Module(), torch.nn.Module()]) + for layer in model.layers: + layer.self_attn = torch.nn.Module() + model.layers[0].self_attn.qkv_proj = torch.nn.Linear(2, 6, bias=False) + model.layers[1].self_attn.q_proj = torch.nn.Linear(2, 2, bias=False) + shards: list[str] = [] + model.layers[0].self_attn.qkv_proj.weight.weight_loader = lambda param, weight, shard_id: shards.append(shard_id) + + q_weight = torch.full((2, 2), 2.0) + weights = [ + (f"layers.{i}.self_attn.{tensor}.weight", q_weight) for i in (0, 1) for tensor in ("q_proj", "k_proj", "v_proj") + ] + [("layers.1.self_attn.k_norm.weight", torch.ones(2))] + loaded = Gemma4Model.load_weights(cast(Gemma4Model, model), weights) + + assert shards == ["q", "k", "v"] + assert "layers.1.self_attn.q_proj.weight" in loaded + assert not any(name.startswith("layers.1.self_attn.k") for name in loaded) + assert not any(name.startswith("layers.1.self_attn.v") for name in loaded) + assert torch.equal(model.layers[1].self_attn.q_proj.weight, q_weight) + + +@pytest.mark.cpu_test +def test_gemma3n_kv_shared_layer_mapper() -> None: + """Only non-shared layers pack q/k/v into qkv_proj; KV-shared layers keep + q_proj and drop the redundant K/V tensors original checkpoints ship.""" + config = SimpleNamespace(num_hidden_layers=4, num_kv_shared_layers=2) + mapper = Gemma3nTextModel.hf_to_aphrodite_mapper | _kv_sharing_weights_mapper(config) + weights = [ + (f"layers.{i}.self_attn.{tensor}.weight", torch.empty(0)) + for i in (1, 3) + for tensor in ("q_proj", "k_proj", "v_proj", "k_norm", "o_proj") + ] + + mapped = [(name, getattr(weight, "shard_id", None)) for name, weight in mapper.apply(weights)] + + assert mapped == [ + ("layers.1.self_attn.qkv_proj.weight", "q"), + ("layers.1.self_attn.qkv_proj.weight", "k"), + ("layers.1.self_attn.qkv_proj.weight", "v"), + ("layers.1.self_attn.k_norm.weight", None), + ("layers.1.self_attn.o_proj.weight", None), + ("layers.3.self_attn.q_proj.weight", None), + ("layers.3.self_attn.o_proj.weight", None), + ] + + @pytest.mark.parametrize("model", MODELS) def test_dummy_loader(aphrodite_runner, monkeypatch, model: str) -> None: with monkeypatch.context() as m: From 8347c70090db35fe1d3d7609cd9e7b1037f85357 Mon Sep 17 00:00:00 2001 From: AlpinDale Date: Mon, 7 Sep 2026 21:11:19 +0000 Subject: [PATCH 33/36] [sync] [Docs] Clarify admission control limits apply server-wide, not per DP rank (#55124) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Upstream-vLLM: cd64c2dea9c72af333de7ec05293d54bbf1bd128 Co-authored-by: Nicolò Lucchesi --- .sync/vllm-sha | 2 +- aphrodite/config/scheduler.py | 13 ++++++++++++- docs/src/content/docs/deployment/parallelism.md | 10 ++++++++++ 3 files changed, 23 insertions(+), 2 deletions(-) diff --git a/.sync/vllm-sha b/.sync/vllm-sha index 364ed64c64..1459001f30 100644 --- a/.sync/vllm-sha +++ b/.sync/vllm-sha @@ -1 +1 @@ -f2d45f26bd6a2c841ffbd0030ebaefe87c274e58 +cd64c2dea9c72af333de7ec05293d54bbf1bd128 diff --git a/aphrodite/config/scheduler.py b/aphrodite/config/scheduler.py index b4a8f45a89..e58238c116 100644 --- a/aphrodite/config/scheduler.py +++ b/aphrodite/config/scheduler.py @@ -76,7 +76,13 @@ class SchedulerConfig: at the same time, or None for no limit. When the limit is reached, new requests are rejected with HTTP 503 so the client can retry on another instance. This bounds Aphrodite's otherwise unbounded request queue and is - primarily a coarse capacity valve.""" + primarily a coarse capacity valve. + + Unlike ``max_num_seqs``, which applies per data-parallel rank, this + limit is enforced in the API server process and counts in-flight + requests across all DP ranks it routes to. Size it as roughly + ``data_parallel_size * max_num_seqs`` plus the desired queue depth if + it should not bind before per-rank admission does.""" max_num_queued_tokens: int | None = Field(default=None, ge=0) """Maximum total prompt tokens of requests currently in the prefill @@ -89,6 +95,11 @@ class SchedulerConfig: prefill-decode setup this maps directly to the prefill pool's capacity. + Like ``max_num_queued_reqs``, this limit is enforced in the API + server process and covers the prefill backlog across all DP ranks it + routes to, so ``prefill_throughput`` in the formula above is the + aggregate throughput of the deployment. + Note: the count is conservative. A partially prefilled request still contributes its full ``prompt_len`` until it transitions out of the prefill phase, because the scheduler's per-iteration diff --git a/docs/src/content/docs/deployment/parallelism.md b/docs/src/content/docs/deployment/parallelism.md index dbcf8ef4c0..e9a02484d8 100644 --- a/docs/src/content/docs/deployment/parallelism.md +++ b/docs/src/content/docs/deployment/parallelism.md @@ -85,6 +85,16 @@ aphrodite serve MODEL --data-parallel-size 4 Use data parallelism when one replica fits and request volume is high. A data parallel size of four needs enough memory for four copies of the model. +When sizing deployments, `--max-num-seqs` applies per data-parallel rank. +The admission control limits `--max-num-queued-reqs` and +`--max-num-queued-tokens` apply across all DP ranks that each API server +process routes to: each process counts in-flight requests and prefill backlog +across those ranks. For example, with `--data-parallel-size=4 +--max-num-seqs=256 --max-num-queued-reqs=256`, an API server process rejects new +requests once 256 are in-flight in total, even though the ranks could jointly +run 1024. To keep ranks saturated, size the request cap as roughly +`data-parallel-size * max-num-seqs` plus the desired queue depth. + Data parallelism can give better throughput isolation than one large replica. The router must distribute requests across the replicas. Use session affinity only when an application requires it. Prefix cache entries are local to a From f239831572d25e3928f72d2e38fbb0c81304d66c Mon Sep 17 00:00:00 2001 From: AlpinDale Date: Mon, 7 Sep 2026 21:13:48 +0000 Subject: [PATCH 34/36] [sync] [Tests] Update SarvamMLA transformers v5 compatibility reason to hf (#55728) Upstream-vLLM: 7cc89a7ddabd6be3ff64535139bea20548062ca6 Co-authored-by: Abhi <108084481+Aj2280@users.noreply.github.com> --- .sync/vllm-sha | 2 +- tests/models/registry.py | 5 +++-- 2 files changed, 4 insertions(+), 3 deletions(-) diff --git a/.sync/vllm-sha b/.sync/vllm-sha index 1459001f30..45a682a3ab 100644 --- a/.sync/vllm-sha +++ b/.sync/vllm-sha @@ -1 +1 @@ -cd64c2dea9c72af333de7ec05293d54bbf1bd128 +7cc89a7ddabd6be3ff64535139bea20548062ca6 diff --git a/tests/models/registry.py b/tests/models/registry.py index 3c6238d817..6b93135597 100644 --- a/tests/models/registry.py +++ b/tests/models/registry.py @@ -471,8 +471,9 @@ def check_available_online( is_available_online=True, max_transformers_version="5.3", transformers_version_reason={ - "aphrodite": ( - "aphrodite upgraded transformers above v5.4 where validate_rope() no longer accepts ignore_keys param" + "hf": ( + "Transformers v5.4 removed the ignore_keys param from " + "validate_rope(); Aphrodite patches validate_rope() and is unaffected" ) }, ), From 8a459bfa2b41b0ce2e28ee30cf5b73f8eb718c39 Mon Sep 17 00:00:00 2001 From: AlpinDale Date: Mon, 7 Sep 2026 21:18:53 +0000 Subject: [PATCH 35/36] [sync] [Model] Add Cohere Compass model (#54774) Upstream-vLLM: 70584f69b1bf866224604029fae038864a368019 Co-authored-by: am-cohere <312513976+am-cohere@users.noreply.github.com> --- .sync/vllm-sha | 2 +- .../model_executor/models/cohere_compass.py | 2229 +++++++++++++++++ aphrodite/model_executor/models/registry.py | 4 + aphrodite/transformers_utils/config.py | 4 +- tests/models/registry.py | 4 + 5 files changed, 2241 insertions(+), 2 deletions(-) create mode 100644 aphrodite/model_executor/models/cohere_compass.py diff --git a/.sync/vllm-sha b/.sync/vllm-sha index 45a682a3ab..1a13682dd0 100644 --- a/.sync/vllm-sha +++ b/.sync/vllm-sha @@ -1 +1 @@ -7cc89a7ddabd6be3ff64535139bea20548062ca6 +70584f69b1bf866224604029fae038864a368019 diff --git a/aphrodite/model_executor/models/cohere_compass.py b/aphrodite/model_executor/models/cohere_compass.py new file mode 100644 index 0000000000..4936e387a4 --- /dev/null +++ b/aphrodite/model_executor/models/cohere_compass.py @@ -0,0 +1,2229 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project + +# Copyright 2025 The vLLM team. +# Copyright 2026 The Cohere Team. +# All rights reserved. +# +# This code is based on EleutherAI's GPT-NeoX library and the GPT-NeoX +# and OPT implementations in this library. It has been modified from its +# original forms to accommodate minor architectural differences compared +# to GPT-NeoX and OPT used by the Meta AI team that trained the model. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +"""Inference-only Cohere Compass model compatible with HuggingFace weights.""" + +import math +from collections.abc import Callable, Iterable, Iterator, Mapping, Sequence +from functools import lru_cache, partial +from itertools import islice +from typing import Annotated, Any, Literal, TypeAlias + +import einops +import numpy as np +import torch +import torch.nn as nn +import torch.nn.functional as F +from transformers import BatchFeature +from transformers.models.cohere_compass import ( + CohereCompassConfig, + CohereCompassProcessor, +) +from transformers.models.cohere_compass.configuration_cohere_compass import ( + CohereCompassVisionConfig, +) +from transformers.models.cohere_compass.image_processing_cohere_compass import ( + CohereCompassImageProcessor, +) +from transformers.models.cohere_compass.image_processing_cohere_compass import ( + smart_resize as image_smart_resize, +) + +from aphrodite.compilation.decorators import support_torch_compile +from aphrodite.config import AphroditeConfig, CacheConfig +from aphrodite.config.multimodal import BaseDummyOptions +from aphrodite.distributed import ( + get_pp_group, + get_tensor_model_parallel_world_size, + parallel_state, +) +from aphrodite.distributed import utils as dist_utils +from aphrodite.inputs import ModalityData, MultiModalDataDict +from aphrodite.logger import init_logger +from aphrodite.model_executor.layers.activation import _ACTIVATION_REGISTRY +from aphrodite.model_executor.layers.attention import Attention +from aphrodite.model_executor.layers.attention.mm_encoder_attention import ( + MMEncoderAttention, +) +from aphrodite.model_executor.layers.conv import Conv3dLayer +from aphrodite.model_executor.layers.linear import ( + ColumnParallelLinear, + QKVParallelLinear, + RowParallelLinear, +) +from aphrodite.model_executor.layers.logits_processor import LogitsProcessor +from aphrodite.model_executor.layers.quantization import QuantizationConfig +from aphrodite.model_executor.layers.rotary_embedding import get_rope +from aphrodite.model_executor.layers.rotary_embedding.common import ApplyRotaryEmb +from aphrodite.model_executor.layers.vocab_parallel_embedding import VocabParallelEmbedding +from aphrodite.model_executor.models.module_mapping import MultiModelKeys +from aphrodite.multimodal import MULTIMODAL_REGISTRY +from aphrodite.multimodal.inputs import ( + ImageItem, + MultiModalFeatureSpec, + MultiModalFieldConfig, + MultiModalKwargsItems, +) +from aphrodite.multimodal.parse import ( + DictEmbeddingItems, + ImageSize, + ModalityDataItems, + MultiModalDataItems, + MultiModalDataParser, +) +from aphrodite.multimodal.processing import ( + BaseDummyInputsBuilder, + BaseMultiModalProcessor, + BaseProcessingInfo, + PromptReplacement, + PromptUpdate, +) +from aphrodite.multimodal.video_prune.evs import ( + compute_mrope_for_media, + recompute_mrope_positions, +) +from aphrodite.sequence import IntermediateTensors +from aphrodite.tokenizers.registry import cached_tokenizer_from_config +from aphrodite.triton_utils import HAS_TRITON, tl, triton +from aphrodite.utils.math_utils import round_up +from aphrodite.utils.tensor_schema import TensorSchema, TensorShape +from aphrodite.v1.worker.encoder_cudagraph_defs import EncoderCudaGraphReplayBuffers + +from ...utils.torch_utils import async_tensor_h2d +from .commandr import CohereForCausalLM, CohereMLP, LayerNorm +from .interfaces import ( + MultiModalEmbeddings, + SupportsEagle, + SupportsEagle3, + SupportsEncoderCudaGraph, + SupportsLoRA, + SupportsMRoPE, + SupportsMultiModal, + SupportsMultiModalPruning, + SupportsPP, + _require_is_multimodal, +) +from .utils import ( + AutoWeightsLoader, + WeightsMapper, + _merge_multimodal_embeddings, + extract_layer_index, + make_empty_intermediate_tensors_factory, + make_layers, + maybe_prefix, +) +from .vision import ( + get_fp8_padded_hidden_size, + get_vit_attn_backend, + is_vit_use_data_parallel, + run_dp_sharded_mrope_vision_model, +) + +logger = init_logger(__name__) + +# --------------------------------------------------------------------------- +# Triton kernel: fused bilinear position-embedding interpolation +# --------------------------------------------------------------------------- +# Replaces many small eager-mode CUDA kernels with a single launch. +# The spatial-merge reorder is baked into the index math so the output +# is ready to be added to the patch embeddings directly. +# --------------------------------------------------------------------------- + +if HAS_TRITON: + + @triton.jit + def _bilinear_pos_embed_kernel( + embed_ptr, + output_ptr, + H, + W, + h_scale, + w_scale, + NUM_GRID: tl.constexpr, + M_SIZE: tl.constexpr, + HIDDEN_DIM: tl.constexpr, + BLOCK_D: tl.constexpr, + ): + """Fused bilinear pos-embed interpolation with spatial-merge reorder.""" + pid = tl.program_id(0) + total_spatial = H * W + spatial_idx = pid % total_spatial + + num_blocks_w = W // M_SIZE + block_idx = spatial_idx // (M_SIZE * M_SIZE) + local_idx = spatial_idx % (M_SIZE * M_SIZE) + br = block_idx // num_blocks_w + bc = block_idx % num_blocks_w + lr = local_idx // M_SIZE + lc = local_idx % M_SIZE + row = br * M_SIZE + lr + col = bc * M_SIZE + lc + + h_frac = row.to(tl.float32) * h_scale + w_frac = col.to(tl.float32) * w_scale + + hf = tl.math.floor(h_frac).to(tl.int32) + wf = tl.math.floor(w_frac).to(tl.int32) + hc = tl.minimum(hf + 1, NUM_GRID - 1) + wc = tl.minimum(wf + 1, NUM_GRID - 1) + + dh = h_frac - hf.to(tl.float32) + dw = w_frac - wf.to(tl.float32) + w11 = dh * dw + w10 = dh - w11 + w01 = dw - w11 + w00 = 1.0 - dh - w01 + + off00 = (hf * NUM_GRID + wf) * HIDDEN_DIM + off01 = (hf * NUM_GRID + wc) * HIDDEN_DIM + off10 = (hc * NUM_GRID + wf) * HIDDEN_DIM + off11 = (hc * NUM_GRID + wc) * HIDDEN_DIM + out_off = pid * HIDDEN_DIM + + # Cast weights to output dtype so the multiply-accumulate stays + # in the same precision as the native PyTorch implementation. + out_dtype = output_ptr.dtype.element_ty + w00_c = w00.to(out_dtype) + w01_c = w01.to(out_dtype) + w10_c = w10.to(out_dtype) + w11_c = w11.to(out_dtype) + + for d in tl.range(0, HIDDEN_DIM, BLOCK_D): + cols = d + tl.arange(0, BLOCK_D) + mask = cols < HIDDEN_DIM + + e00 = tl.load(embed_ptr + off00 + cols, mask=mask) + e01 = tl.load(embed_ptr + off01 + cols, mask=mask) + e10 = tl.load(embed_ptr + off10 + cols, mask=mask) + e11 = tl.load(embed_ptr + off11 + cols, mask=mask) + + val = w00_c * e00 + w01_c * e01 + w10_c * e10 + w11_c * e11 + + tl.store(output_ptr + out_off + cols, val, mask=mask) + + def triton_pos_embed_interpolate( + embed_weight: torch.Tensor, + t: int, + h: int, + w: int, + num_grid_per_side: int, + m_size: int, + dtype: torch.dtype, + ) -> torch.Tensor: + """Launch the fused Triton kernel for one (t,h,w) grid. + + Returns a tensor of shape ``(t * h * w, hidden_dim)`` with the + bilinearly-interpolated position embeddings in spatial-merge order. + """ + assert h % m_size == 0 and w % m_size == 0, f"h={h} and w={w} must be divisible by m_size={m_size}" + hidden_dim = embed_weight.shape[1] + total_out = t * h * w + output = torch.empty( + total_out, + hidden_dim, + device=embed_weight.device, + dtype=dtype, + ) + + h_scale = float(num_grid_per_side - 1) / float(h - 1) if h > 1 else 0.0 + w_scale = float(num_grid_per_side - 1) / float(w - 1) if w > 1 else 0.0 + + BLOCK_D = triton.next_power_of_2(hidden_dim) + + _bilinear_pos_embed_kernel[(total_out,)]( + embed_weight, + output, + h, + w, + h_scale, + w_scale, + num_grid_per_side, + m_size, + hidden_dim, + BLOCK_D, + ) + return output + + +def pos_embed_interpolate_native( + embed_weight: torch.Tensor, + t: int, + h: int, + w: int, + num_grid_per_side: int, + m_size: int, + dtype: torch.dtype, +) -> torch.Tensor: + """Eager PyTorch bilinear position-embedding interpolation. + + Returns a tensor of shape ``(t * h * w, hidden_dim)`` with the + bilinearly-interpolated position embeddings in spatial-merge order. + """ + assert h % m_size == 0 and w % m_size == 0, f"h={h} and w={w} must be divisible by m_size={m_size}" + hidden_dim = embed_weight.shape[1] + device = embed_weight.device + + h_idxs = torch.linspace( + 0, + num_grid_per_side - 1, + h, + dtype=torch.float32, + device=device, + ) + w_idxs = torch.linspace( + 0, + num_grid_per_side - 1, + w, + dtype=torch.float32, + device=device, + ) + + h_floor = h_idxs.to(torch.long) + w_floor = w_idxs.to(torch.long) + h_ceil = torch.clamp(h_floor + 1, max=num_grid_per_side - 1) + w_ceil = torch.clamp(w_floor + 1, max=num_grid_per_side - 1) + + dh = h_idxs - h_floor + dw = w_idxs - w_floor + + dh_grid, dw_grid = torch.meshgrid(dh, dw, indexing="ij") + h_floor_grid, w_floor_grid = torch.meshgrid(h_floor, w_floor, indexing="ij") + h_ceil_grid, w_ceil_grid = torch.meshgrid(h_ceil, w_ceil, indexing="ij") + + w11 = dh_grid * dw_grid + w10 = dh_grid - w11 + w01 = dw_grid - w11 + w00 = 1 - dh_grid - w01 + + h_grid = torch.stack([h_floor_grid, h_floor_grid, h_ceil_grid, h_ceil_grid]) + w_grid = torch.stack([w_floor_grid, w_ceil_grid, w_floor_grid, w_ceil_grid]) + h_grid_idx = h_grid * num_grid_per_side + + indices = (h_grid_idx + w_grid).reshape(4, -1) + weights = torch.stack([w00, w01, w10, w11], dim=0).reshape(4, -1, 1) + weights = weights.to(dtype=dtype) + + embeds = embed_weight[indices] + embeds *= weights + combined = embeds.sum(dim=0) + + combined = combined.reshape(h // m_size, m_size, w // m_size, m_size, hidden_dim) + combined = combined.permute(0, 2, 1, 3, 4).reshape(1, -1, hidden_dim) + repeated = combined.expand(t, -1, -1).reshape(-1, hidden_dim) + return repeated.to(dtype=dtype) + + +class Cohere_VLImagePixelInputs(TensorSchema): + """ + Dimensions: + - np: Number of patches + - ni: Number of images + - cps: Number of channels * patch_size * patch_size + + Historical context: + - pixel_values shape: (num_patches, num_channels * patch_size * + patch_size) + - image_grid_thw shape: (num_images, 3) in (grid_t, grid_h, grid_w) + format. + """ + + type: Literal["pixel_values"] + + pixel_values: Annotated[ + torch.Tensor, + TensorShape("np", "cps"), + ] + + image_grid_thw: Annotated[ + torch.Tensor, + TensorShape("ni", 3), + ] + + +class Cohere_VLImageEmbeddingInputs(TensorSchema): + """ + Dimensions: + - nf: Number of image features + - hs: Hidden size + - ni: Number of images + + Historical context: + - image_embeds shape: (num_image_features, hidden_size) + - num_image_features varies based on the number and resolution of the + images. + - hidden_size must match the hidden size of language model backbone. + - image_grid_thw shape: (num_images, 3) in (grid_t, grid_h, grid_w) + format + """ + + type: Literal["image_embeds"] + + image_embeds: Annotated[ + torch.Tensor, + TensorShape("nf", "hs"), + ] + + image_grid_thw: Annotated[ + torch.Tensor, + TensorShape("ni", 3), + ] + + +Cohere_VLImageInputs: TypeAlias = Cohere_VLImagePixelInputs | Cohere_VLImageEmbeddingInputs + + +class Cohere_VisionAttention(nn.Module): + def __init__( + self, + embed_dim: int, + num_heads: int, + projection_size: int, + quant_config: QuantizationConfig | None = None, + prefix: str = "", + ) -> None: + super().__init__() + # Per attention head and per partition values. + use_data_parallel = is_vit_use_data_parallel() + self.tp_size = 1 if use_data_parallel else parallel_state.get_tensor_model_parallel_world_size() + self.tp_rank = 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.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", + disable_tp=use_data_parallel, + ) + + self.proj = RowParallelLinear( + input_size=projection_size, + output_size=embed_dim, + quant_config=quant_config, + prefix=f"{prefix}.proj", + 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 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, # Only used for Flash Attention + # Only used for FlashInfer CuDNN backend. + sequence_lengths: torch.Tensor | None, + ) -> torch.Tensor: + # [s, b, c] --> [s, b, head * 3 * head_dim] + x, _ = self.qkv(x) + seq_len, batch_size, _ = x.shape + + qkv = einops.rearrange( + x, + "s b (three head head_dim) -> b s three head head_dim", + three=3, + head=self.num_attention_heads_per_partition, + ) + + if rotary_pos_emb_cos is not None and rotary_pos_emb_sin is not None: + qk, v = qkv[:, :, :2], qkv[:, :, 2] + + qk_reshaped = einops.rearrange(qk, "b s two head head_dim -> (two b) s head head_dim", two=2) + qk_reshaped = qk_reshaped.contiguous() + qk_rotated = self.apply_rotary_emb( + qk_reshaped, + rotary_pos_emb_cos, + rotary_pos_emb_sin, + ) + qk_rotated = qk_rotated.view( + 2, + batch_size, + seq_len, + self.num_attention_heads_per_partition, + self.hidden_size_per_attention_head, + ) + q, k = qk_rotated.unbind(dim=0) + else: + q, k, v = qkv.unbind(dim=2) + + context_layer = self.attn( + query=q, + key=k, + value=v, + cu_seqlens=cu_seqlens, + max_seqlen=max_seqlen, + sequence_lengths=sequence_lengths, + ) + + context_layer = einops.rearrange(context_layer, "b s h d -> s b (h d)", b=batch_size).contiguous() + + output, _ = self.proj(context_layer) + return output + + +class Cohere_VisionPatchEmbed(nn.Module): + def __init__( + self, + patch_size: int = 14, + temporal_patch_size: int = 2, + in_channels: int = 3, + hidden_size: int = 1152, + ) -> 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 Cohere_VisionMLP(nn.Module): + def __init__( + self, + in_features: int, + hidden_features: int, + bias: bool = False, + act_fn: Callable[[torch.Tensor], torch.Tensor] = F.silu, + quant_config: QuantizationConfig | None = None, + prefix: str = "", + ): + super().__init__() + use_data_parallel = is_vit_use_data_parallel() + self.linear_fc1 = ColumnParallelLinear( + in_features, + hidden_features, + bias=bias, + quant_config=quant_config, + return_bias=False, + prefix=f"{prefix}.linear_fc1", + disable_tp=use_data_parallel, + ) + self.linear_fc2 = RowParallelLinear( + hidden_features, + in_features, + bias=bias, + quant_config=quant_config, + return_bias=False, + prefix=f"{prefix}.linear_fc2", + disable_tp=use_data_parallel, + ) + self.act_fn = act_fn + + def forward(self, x: torch.Tensor): + mlp_output = self.linear_fc2(self.act_fn(self.linear_fc1(x))) + return mlp_output + + +class Cohere_VisionBlock(nn.Module): + def __init__( + self, + dim: int, + num_heads: int, + mlp_hidden_dim: int, + act_fn: Callable[[torch.Tensor], torch.Tensor] = F.silu, + norm_layer: Callable[[int], 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 = Cohere_VisionAttention( + embed_dim=dim, + num_heads=num_heads, + projection_size=dim, + quant_config=quant_config, + prefix=f"{prefix}.attn", + ) + self.mlp = Cohere_VisionMLP( + dim, + mlp_hidden_dim, + act_fn=act_fn, + 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: torch.Tensor, # Only used for Flash Attention + sequence_lengths: torch.Tensor, # Only used for FlashInfer CuDNN backend + ) -> torch.Tensor: + x = x + 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, + sequence_lengths=sequence_lengths, + ) + + x = x + self.mlp(self.norm2(x)) + return x + + +class Cohere_VisionPatchMerger(nn.Module): + def __init__( + self, + d_model: int, + context_dim: int, + norm_layer: Callable[[int], nn.Module] | None = None, + spatial_merge_size: int = 2, + use_postshuffle_norm: bool = False, + quant_config: QuantizationConfig | None = None, + prefix: str = "", + ) -> None: + super().__init__() + use_data_parallel = is_vit_use_data_parallel() + self.hidden_size = context_dim * (spatial_merge_size**2) + + self.use_postshuffle_norm = use_postshuffle_norm + if self.use_postshuffle_norm: + context_dim = self.hidden_size + + if norm_layer is None: + norm_layer = partial(nn.LayerNorm, eps=1e-6) + self.norm = norm_layer(context_dim) + self.linear_fc1 = ColumnParallelLinear( + self.hidden_size, + self.hidden_size, + bias=True, + quant_config=quant_config, + prefix=f"{prefix}.linear_fc1", + disable_tp=use_data_parallel, + ) + self.act_fn = nn.GELU() + self.linear_fc2 = RowParallelLinear( + self.hidden_size, + d_model, + bias=True, + quant_config=quant_config, + prefix=f"{prefix}.linear_fc2", + disable_tp=use_data_parallel, + ) + + def forward(self, x: torch.Tensor) -> torch.Tensor: + if self.use_postshuffle_norm: + x = self.norm(x.view(-1, self.hidden_size)) + else: + x = self.norm(x).view(-1, self.hidden_size) + + x_parallel, _ = self.linear_fc1(x) + x_parallel = self.act_fn(x_parallel) + out, _ = self.linear_fc2(x_parallel) + return out + + +class Cohere_VisionTransformer(nn.Module): + hf_to_aphrodite_mapper = WeightsMapper( + orig_to_new_stacked={ + ".q.": (".qkv.", "q"), + ".k.": (".qkv.", "k"), + ".v.": (".qkv.", "v"), + } + ) + + def __init__( + self, + vision_config: CohereCompassVisionConfig, + norm_eps: float = 1e-6, + quant_config: QuantizationConfig | None = None, + prefix: str = "", + ) -> None: + super().__init__() + self.hidden_size = vision_config.hidden_size + self.num_heads = vision_config.num_heads + self.num_position_embeddings = vision_config.num_position_embeddings + self.patch_size = vision_config.patch_size + self.spatial_merge_size = vision_config.spatial_merge_size + self.spatial_merge_unit = self.spatial_merge_size**2 + self.temporal_patch_size = vision_config.temporal_patch_size + self.deepstack_visual_indexes = ( + vision_config.deepstack_visual_indexes if hasattr(vision_config, "deepstack_visual_indexes") else [] + ) + self.num_grid_per_side = int(self.num_position_embeddings**0.5) + + use_data_parallel = is_vit_use_data_parallel() + self.tp_size = 1 if use_data_parallel else parallel_state.get_tensor_model_parallel_world_size() + + # NOTE: This is used for creating empty tensor for all_gather for + # DP ViT. Here out_hidden_size is enlarged due to deepstack + self.out_hidden_size = vision_config.out_hidden_size * (1 + len(self.deepstack_visual_indexes)) + + self.patch_embed = Cohere_VisionPatchEmbed( + patch_size=self.patch_size, + temporal_patch_size=self.temporal_patch_size, + in_channels=vision_config.in_channels, + hidden_size=self.hidden_size, + ) + + self.pos_embed = nn.Embedding(self.num_position_embeddings, self.hidden_size) + + norm_layer = partial(nn.LayerNorm, eps=norm_eps) + head_dim = self.hidden_size // self.num_heads + + # FP8 attention: Q/K/V become independent contiguous tensors + # after quantization, so cu_seqlens uses uniform stride (no 3x V). + self.fp8_padded_hidden_size = get_fp8_padded_hidden_size(self.num_heads, head_dim) + + 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.merger = Cohere_VisionPatchMerger( + d_model=vision_config.out_hidden_size, + context_dim=self.hidden_size, + norm_layer=norm_layer, + spatial_merge_size=self.spatial_merge_size, + quant_config=quant_config, + prefix=f"{prefix}.merger", + ) + + self.deepstack_merger_list = nn.ModuleList( + [ + Cohere_VisionPatchMerger( + d_model=vision_config.out_hidden_size, + context_dim=self.hidden_size, + spatial_merge_size=self.spatial_merge_size, + use_postshuffle_norm=True, + norm_layer=norm_layer, + quant_config=quant_config, + prefix=f"{prefix}.deepstack_merger_list.{layer_idx}", + ) + for layer_idx in range(len(self.deepstack_visual_indexes)) + ] + ) + + self.attn_backend = get_vit_attn_backend( + head_size=head_dim, + dtype=torch.get_default_dtype(), + ) + + self.blocks = nn.ModuleList( + [ + Cohere_VisionBlock( + dim=self.hidden_size, + num_heads=self.num_heads, + mlp_hidden_dim=vision_config.intermediate_size, + act_fn=_ACTIVATION_REGISTRY[vision_config.hidden_act], + norm_layer=norm_layer, + quant_config=quant_config, + prefix=f"{prefix}.blocks.{layer_idx}", + ) + for layer_idx in range(vision_config.depth) + ] + ) + + @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 + + @staticmethod + @lru_cache(maxsize=1024) + def rot_pos_ids(h: int, w: int, spatial_merge_size: int) -> torch.Tensor: + hpos_ids = np.broadcast_to(np.arange(h).reshape(h, 1), (h, w)) + h_div = h // spatial_merge_size + w_div = w // spatial_merge_size + hpos_ids = hpos_ids.reshape( + h_div, + spatial_merge_size, + w_div, + spatial_merge_size, + ) + hpos_ids = hpos_ids.transpose(0, 2, 1, 3) + hpos_ids = hpos_ids.flatten() + + wpos_ids = np.broadcast_to(np.arange(w).reshape(1, w), (h, w)) + wpos_ids = wpos_ids.reshape( + h_div, + spatial_merge_size, + w_div, + spatial_merge_size, + ) + wpos_ids = wpos_ids.transpose(0, 2, 1, 3) + wpos_ids = wpos_ids.flatten() + + return torch.from_numpy(np.stack([hpos_ids, wpos_ids], axis=-1)) + + def rot_pos_emb(self, grid_thw: list[list[int]]): + max_grid_size = max(max(h, w) for _, h, w in grid_thw) + pos_ids = [ + self.rot_pos_ids(h, w, self.spatial_merge_size) + if t == 1 + else self.rot_pos_ids(h, w, self.spatial_merge_size).repeat(t, 1) + for t, h, w in grid_thw + ] + pos_ids = torch.cat(pos_ids, dim=0).to(self.device, non_blocking=True) + + # Use pre-computed cos_sin_cache from RotaryEmbedding + cos, sin = self.rotary_pos_emb.get_cos_sin(max_grid_size) + + cos_combined = cos[pos_ids].flatten(1) + sin_combined = sin[pos_ids].flatten(1) + + return cos_combined, sin_combined + + def fast_pos_embed_interpolate(self, grid_thw: list[list[int]]) -> torch.Tensor: + interpolate_fn = triton_pos_embed_interpolate if HAS_TRITON else pos_embed_interpolate_native + outputs = [] + for t, h, w in grid_thw: + outputs.append( + interpolate_fn( + self.pos_embed.weight, + t, + h, + w, + self.num_grid_per_side, + self.spatial_merge_size, + self.dtype, + ) + ) + return torch.cat(outputs, dim=0) + + 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 from grid_thw_list. + + Shared by the eager forward path, CUDA graph capture, and + CUDA graph replay to avoid duplicated implementation. + + Args: + grid_thw_list: Grid configurations as list of [t, h, w]. + max_batch_size: If set, pad cu_seqlens to this size + (needed for CUDA graph capture/replay). + max_frames_per_batch: If set, overrides max_batch_size for + cu_seqlens padding. For video inputs each item contributes + T attention sequences (frames); this sizes the buffer to + the total frame budget so video replays never overflow. + max_seqlen_override: If set, use this value for max_seqlen + instead of computing from cu_seqlens (needed for CUDA + graph capture to cover worst-case replay scenarios). + device: Device to place tensors on. Defaults to self.device. + """ + if device is None: + device = self.device + + metadata: dict[str, torch.Tensor | None] = {} + + # Positional embeddings + metadata["pos_embeds"] = self.fast_pos_embed_interpolate(grid_thw_list) + 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 + + # cu_seqlens from grid_thw + 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 cu_seqlens to the required number of sequences. + # For videos each item contributes T frames = T attention sequences, + # so the total can exceed max_batch_size. max_frames_per_batch + # overrides the pad target when set. + 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, + ), + ] + ) + + # sequence_lengths (backend-specific) + metadata["sequence_lengths"] = MMEncoderAttention.maybe_compute_seq_lens(self.attn_backend, cu_seqlens, device) + + # max_seqlen + 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) + # Keep max_seqlen on CPU: attention wrappers call .item() on it, + # and having it on GPU would capture a wasteful D2H copy in CUDA + # graphs without changing behavior (the scalar is baked at capture). + metadata["max_seqlen"] = torch.tensor(max_seqlen_val, dtype=torch.int32) + + # Recompute cu_seqlens (backend-specific transformation) + metadata["cu_seqlens"] = MMEncoderAttention.maybe_recompute_cu_seqlens( + self.attn_backend, + cu_seqlens, + self.hidden_size, + self.tp_size, + device, + fp8_padded_hidden_size=self.fp8_padded_hidden_size, + ) + + 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: + hidden_states = x.to(device=self.device, dtype=self.dtype, non_blocking=True) + hidden_states = self.patch_embed(hidden_states) + + if encoder_metadata is None: + if isinstance(grid_thw, list): + grid_thw_list = grid_thw + else: + grid_thw_list = grid_thw.tolist() + encoder_metadata = self.prepare_encoder_metadata(grid_thw_list) + + pos_embeds = encoder_metadata["pos_embeds"] + hidden_states = hidden_states + pos_embeds + hidden_states = hidden_states.unsqueeze(1) + + deepstack_feature_lists = [] + for layer_num, blk in enumerate(self.blocks): + hidden_states = blk( + hidden_states, + cu_seqlens=encoder_metadata["cu_seqlens"], + rotary_pos_emb_cos=encoder_metadata["rotary_pos_emb_cos"], + rotary_pos_emb_sin=encoder_metadata["rotary_pos_emb_sin"], + max_seqlen=encoder_metadata["max_seqlen"], + sequence_lengths=encoder_metadata.get("sequence_lengths"), + ) + if layer_num in self.deepstack_visual_indexes: + deepstack_merger_idx = self.deepstack_visual_indexes.index(layer_num) + deepstack_feature = self.deepstack_merger_list[deepstack_merger_idx](hidden_states) + deepstack_feature_lists.append(deepstack_feature) + hidden_states = self.merger(hidden_states) + hidden_states = torch.cat( + [hidden_states] + deepstack_feature_lists, dim=1 + ) # [seq_len, hidden_size * (1 + depth_of_deepstack)] + return hidden_states + + def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]: + loader = AutoWeightsLoader(self) + return loader.load_weights(weights, mapper=self.hf_to_aphrodite_mapper) + + +def _create_cohere_compass_field_factory( + spatial_merge_size: int, +) -> Callable[ + [Mapping[str, torch.Tensor]], + Mapping[str, MultiModalFieldConfig], +]: + def _field_config( + hf_inputs: Mapping[str, torch.Tensor], + ) -> Mapping[str, MultiModalFieldConfig]: + image_grid_thw = hf_inputs.get("image_grid_thw", torch.empty((0, 3))) + image_pixel_grid_sizes = image_grid_thw.prod(-1) + image_embed_grid_sizes = image_pixel_grid_sizes // spatial_merge_size // spatial_merge_size + + return { + "pixel_values": MultiModalFieldConfig.flat_from_sizes("image", image_pixel_grid_sizes), + "image_embeds": MultiModalFieldConfig.flat_from_sizes("image", image_embed_grid_sizes), + "image_grid_thw": MultiModalFieldConfig.batched("image", keep_on_cpu=True), + } + + return _field_config + + +class CohereCompassMultiModalDataParser(MultiModalDataParser): + def __init__(self, spatial_merge_size: int, *args, **kwargs): + self._spatial_merge_size = spatial_merge_size + super().__init__(*args, **kwargs) + + def _parse_image_data( + self, + data: dict[str, torch.Tensor] | ModalityData[ImageItem], + ) -> ModalityDataItems[Any, Any] | None: + if isinstance(data, dict): + return DictEmbeddingItems( + data, + modality="image", + required_fields={"image_embeds", "image_grid_thw"}, + fields_factory=_create_cohere_compass_field_factory(self._spatial_merge_size), + ) + + return super()._parse_image_data(data) + + +class CohereCompassProcessingInfo(BaseProcessingInfo): + def get_supported_mm_limits(self) -> Mapping[str, int | None]: + return {"image": None} + + def get_mm_max_tokens_per_item( + self, + seq_len: int, + mm_counts: Mapping[str, int], + ) -> Mapping[str, int]: + max_image_tokens = self.get_max_image_tokens() + return {"image": max_image_tokens} + + def get_num_image_tokens( + self, + *, + image_width: int, + image_height: int, + image_processor: CohereCompassImageProcessor, + mm_kwargs: Mapping[str, object], + ) -> int: + _, num_image_tokens = self._get_vision_info( + image_width=image_width, + image_height=image_height, + num_frames=1, + image_processor=image_processor, + mm_kwargs=mm_kwargs, + ) + return num_image_tokens + + def get_image_size_with_most_features(self, max_pixels: int | None = None) -> ImageSize: + # NOTE: Simply processing a huge size with _get_vision_info might not give a + # size that maximizes the number of features, i.e., the number of (merged) + # patches. This is because the number of patches limits the allowed aspect + # ratios. For example, suppose the maximum number of patches is 1280. A square + # image cannot be broken down into 1280 patches, so feeding a giant square image + # into _get_vision_info will not yield a size that maximizes the number of + # patches. Therefore, we directly factorize the maximum number of patches into + # height and width. The tricky part is to avoid extreme aspect ratios (>200). + # If we can't find a suitable aspect ratio, we decrease the number of + # patches and retry. This is safe because the processor does not accept extreme + # aspect ratios, so there is no valid post-resize image with the number of + # patches that yields extreme aspect ratios. + + hf_config = self.get_hf_config() + vision_config = hf_config.vision_config + patch_size = vision_config.patch_size + merge_size = vision_config.spatial_merge_size + + if max_pixels is None: + image_processor = self.get_image_processor() + + mm_kwargs = self.ctx.get_merged_mm_kwargs({}) + size = image_processor.size + if override_size := mm_kwargs.get("size"): + size = size | override_size + if (override_min_pixels := mm_kwargs.get("min_pixels")) is not None: + size = size | {"shortest_edge": override_min_pixels} + if (override_max_pixels := mm_kwargs.get("max_pixels")) is not None: + size = size | {"longest_edge": override_max_pixels} + + max_pixels = size["longest_edge"] + + unit = patch_size * merge_size + max_seq_len = max_pixels // (unit * unit) + + def closest_factor_pair(n: int) -> tuple[int, int]: + # left <= right + for d in range(math.isqrt(n), 0, -1): + if n % d == 0: + return d, n // d + return 1, n + + height_factor, width_factor = 1, max_seq_len + for seq_len in range(max_seq_len, 0, -1): + height_factor, width_factor = closest_factor_pair(seq_len) + if width_factor / height_factor <= 200: + break + + return ImageSize(width=unit * width_factor, height=unit * height_factor) + + def get_max_image_tokens(self) -> int: + image_processor = self.get_image_processor() + target_width, target_height = self.get_image_size_with_most_features() + + return self.get_num_image_tokens( + image_width=target_width, + image_height=target_height, + image_processor=image_processor, + mm_kwargs={}, + ) + + def get_hf_config(self) -> CohereCompassConfig: + return self.ctx.get_hf_config(CohereCompassConfig) + + def get_hf_processor(self, **kwargs: object) -> CohereCompassProcessor: + return self.ctx.get_hf_processor( + CohereCompassProcessor, + use_fast=kwargs.pop("use_fast", True), + **kwargs, + ) + + def get_image_processor(self, **kwargs: object) -> CohereCompassImageProcessor: + return self.get_hf_processor(**kwargs).image_processor + + def get_data_parser(self): + return CohereCompassMultiModalDataParser( + self.get_hf_config().vision_config.spatial_merge_size, + expected_hidden_size=self._get_expected_hidden_size(), + ) + + def _get_vision_info( + self, + *, + image_width: int, + image_height: int, + num_frames: int = 2, + do_resize: bool = True, + image_processor: CohereCompassImageProcessor, + mm_kwargs: Mapping[str, object], + ) -> tuple[ImageSize, int]: + hf_config = self.get_hf_config() + vision_config = 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 + + mm_kwargs = self.ctx.get_merged_mm_kwargs(mm_kwargs) + size = image_processor.size + if override_size := mm_kwargs.get("size"): + size = size | override_size + if (override_min_pixels := mm_kwargs.get("min_pixels")) is not None: + size = size | {"shortest_edge": override_min_pixels} + if (override_max_pixels := mm_kwargs.get("max_pixels")) is not None: + size = size | {"longest_edge": override_max_pixels} + + if do_resize: + smart_resize = image_smart_resize + extra_kwargs: dict[str, Any] = {} + + resized_height, resized_width = smart_resize( + height=image_height, + width=image_width, + factor=patch_size * merge_size, + min_pixels=size["shortest_edge"], + max_pixels=size["longest_edge"], + **extra_kwargs, + ) + preprocessed_size = ImageSize(width=resized_width, height=resized_height) + else: + preprocessed_size = ImageSize(width=image_width, height=image_height) + + padded_num_frames = round_up(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 CohereCompassDummyInputsBuilder(BaseDummyInputsBuilder[CohereCompassProcessingInfo]): + def get_dummy_text(self, mm_counts: Mapping[str, int]) -> str: + num_images = mm_counts.get("image", 0) + return "<|VISION_START|><|IMAGE_PAD|><|VISION_END|>" * num_images + + def get_dummy_mm_data( + self, + seq_len: int, + mm_counts: Mapping[str, int], + mm_options: Mapping[str, BaseDummyOptions], + ) -> MultiModalDataDict: + num_images = mm_counts.get("image", 0) + image_overrides = mm_options.get("image") + + target_image_width, target_image_height = self.info.get_image_size_with_most_features() + + return { + "image": self._get_dummy_images( + width=target_image_width, + height=target_image_height, + num_images=num_images, + overrides=image_overrides, + ), + } + + +class CohereCompassMultiModalProcessor(BaseMultiModalProcessor[CohereCompassProcessingInfo]): + def _get_mm_fields_config( + self, + hf_inputs: BatchFeature, + hf_processor_mm_kwargs: Mapping[str, object], + ) -> Mapping[str, MultiModalFieldConfig]: + return _create_cohere_compass_field_factory(self.info.get_hf_config().vision_config.spatial_merge_size)( + hf_inputs + ) + + def _get_prompt_updates( + self, + mm_items: MultiModalDataItems, + hf_processor_mm_kwargs: Mapping[str, Any], + out_mm_kwargs: MultiModalKwargsItems, + ) -> Sequence[PromptUpdate]: + hf_processor = self.info.get_hf_processor(**hf_processor_mm_kwargs) + image_processor = self.info.get_image_processor(**hf_processor_mm_kwargs) + merge_length = image_processor.merge_size**2 + + def get_image_replacement_cohere_compass(item_idx: int): + out_item = out_mm_kwargs["image"][item_idx] + grid_thw = out_item["image_grid_thw"].data + assert isinstance(grid_thw, torch.Tensor) + + num_tokens = int(grid_thw.prod()) // merge_length + return [hf_processor.image_token_id] * num_tokens + + return [ + PromptReplacement( + modality="image", + target=[hf_processor.image_token_id], + replacement=get_image_replacement_cohere_compass, + ), + ] + + +@MULTIMODAL_REGISTRY.register_processor( + CohereCompassMultiModalProcessor, + info=CohereCompassProcessingInfo, + dummy_inputs=CohereCompassDummyInputsBuilder, +) +class CohereCompassForConditionalGeneration( + nn.Module, + SupportsMultiModal, + SupportsEncoderCudaGraph, + SupportsLoRA, + SupportsPP, + SupportsMRoPE, + SupportsEagle, + SupportsEagle3, + SupportsMultiModalPruning, +): + packed_modules_mapping = { + "qkv_proj": [ + "q_proj", + "k_proj", + "v_proj", + ], + "gate_up_proj": [ + "gate_proj", + "up_proj", + ], + "qkv": ["qkv"], # For vision tower's already-packed QKV + } + + supports_encoder_tp_data = True + + # To ensure correct weight loading and mapping. + hf_to_aphrodite_mapper = WeightsMapper( + orig_to_new_prefix={ + "model.visual.": "visual.", + "lm_head.": "language_model.lm_head.", + "model.language_model.": "language_model.model.", + } + ) + + @classmethod + def get_placeholder_str(cls, modality: str, i: int) -> str | None: + if modality.startswith("image"): + return "<|VISION_START|><|IMAGE_PAD|><|VISION_END|>" + raise ValueError("Only image modality is supported") + + def __init__(self, *, aphrodite_config: AphroditeConfig, prefix: str = "model"): + nn.Module.__init__(self) + config = aphrodite_config.model_config.hf_config + quant_config = aphrodite_config.quant_config + multimodal_config = aphrodite_config.model_config.multimodal_config + + self.config = config + self.model_config = aphrodite_config.model_config + self._tokenizer = cached_tokenizer_from_config(aphrodite_config.model_config) + self.multimodal_config = multimodal_config + assert multimodal_config is not None + self.use_data_parallel = multimodal_config.mm_encoder_tp_mode == "data" + self.is_multimodal_pruning_enabled = multimodal_config.is_multimodal_pruning_enabled() + + self.use_deepstack = bool(config.vision_config.deepstack_visual_indexes) + self.deepstack_num_level = len(config.vision_config.deepstack_visual_indexes) + self.visual_dim = config.vision_config.out_hidden_size + self.multiscale_dim = self.visual_dim * self.deepstack_num_level + + with self._mark_tower_model(aphrodite_config, {"image"}): + self.visual = Cohere_VisionTransformer( + config.vision_config, + norm_eps=1e-6, + quant_config=quant_config, + prefix=maybe_prefix(prefix, "visual"), + ) + if self.use_deepstack: + self.deepstack_input_embeds = [ + torch.zeros( + aphrodite_config.scheduler_config.max_num_batched_tokens, + config.text_config.hidden_size, + ) + for _ in range(self.deepstack_num_level) + ] + self.deepstack_input_embeds_num_tokens = 0 + + with self._mark_language_model(aphrodite_config): + self.language_model = CohereCompassForCausalLM( + aphrodite_config=aphrodite_config.with_hf_config(config.text_config), + prefix=maybe_prefix(prefix, "language_model"), + ) + + if not get_pp_group().is_first_rank and self.use_deepstack: + assert self.language_model.model.start_layer >= self.deepstack_num_level + + self.make_empty_intermediate_tensors = self.language_model.make_empty_intermediate_tensors + + def iter_mm_grid_thw(self, mm_features: list[MultiModalFeatureSpec]) -> Iterator[tuple[int, int, int, int, float]]: + spatial_merge_size = self.config.vision_config.spatial_merge_size + for mm_feature in sorted(mm_features, key=lambda f: f.mm_position.offset): + if mm_feature.modality != "image": + raise ValueError(f"Unsupported modality: {mm_feature.modality}") + + offset = mm_feature.mm_position.offset + feature_data = mm_feature.data + assert feature_data is not None + grid_item = feature_data.get("image_grid_thw") + assert grid_item is not None + grid_data = grid_item.data + assert isinstance(grid_data, torch.Tensor) + t, h, w = grid_data.tolist() + assert t == 1, f"Image must have 1 frame, got {t}" + yield offset, 1, h // spatial_merge_size, w // spatial_merge_size, 1.0 + + def _get_deepstack_input_embeds( + self, + num_tokens: int, + ) -> IntermediateTensors | None: + if not getattr(self, "deepstack_input_embeds", None): + return None # If vision tower is skipped + if num_tokens > self.deepstack_input_embeds[0].size(0): + self._resize_deepstack_input_embeds(num_tokens) + + # get deepstack_input_embeds from buffer, and clear the buffer + return IntermediateTensors( + { + f"deepstack_input_embeds_{idx}": self.deepstack_input_embeds[idx][:num_tokens] + for idx in range(self.deepstack_num_level) + } + ) + + def _resize_deepstack_input_embeds(self, num_tokens: int) -> None: + self.deepstack_input_embeds = [ + torch.zeros( + num_tokens, + self.config.text_config.hidden_size, + device=self.deepstack_input_embeds[0].device, + dtype=self.deepstack_input_embeds[0].dtype, + ) + for _ in range(self.deepstack_num_level) + ] + + def _set_deepstack_input_embeds(self, deepstack_input_embeds: torch.Tensor) -> None: + if not getattr(self, "deepstack_input_embeds", None): + return + + # set deepstack_input_embeds to buffer + num_tokens = deepstack_input_embeds.size(1) + if num_tokens > self.deepstack_input_embeds[0].size(0): + self._resize_deepstack_input_embeds(num_tokens) + for idx in range(self.deepstack_num_level): + self.deepstack_input_embeds[idx][:num_tokens].copy_(deepstack_input_embeds[idx]) + self.deepstack_input_embeds_num_tokens = num_tokens + + def _clear_deepstack_input_embeds(self, num_tokens: int) -> None: + if not getattr(self, "deepstack_input_embeds", None): + return + if getattr(self, "deepstack_input_embeds_num_tokens", 0) == 0: + return + + # clear deepstack_input_embeds in buffer + if num_tokens > 0: + for idx in range(self.deepstack_num_level): + self.deepstack_input_embeds[idx][:num_tokens].zero_() + self.deepstack_input_embeds_num_tokens = 0 + + # -- SupportsEncoderCudaGraph protocol methods -- + + def get_encoder_cudagraph_config(self): + from aphrodite.v1.worker.encoder_cudagraph_defs import ( + EncoderCudaGraphConfig, + ) + + # When EVS pruning is enabled, embed_multimodal post-processes both + # image and video embeddings (mrope positions are appended for image, + # prune+append for video). The encoder CUDA graph path bypasses that + # post-process, producing inconsistent embedding formats vs eager. So + # disable CUDA graph for all modalities when pruning is on. + modalities = [] if self.is_multimodal_pruning_enabled else ["image"] + + # Compute max_frames_per_video for budget sizing. + max_frames = 1 + + return EncoderCudaGraphConfig( + modalities=modalities, + buffer_keys=[ + "pixel_values", + "pos_embeds", + "rotary_pos_emb_cos", + "rotary_pos_emb_sin", + "cu_seqlens", + "max_seqlen", + "sequence_lengths", + ], + out_hidden_size=self.visual.out_hidden_size, + max_frames_per_video=max_frames, + ) + + def get_input_modality( + self, + mm_kwargs: dict[str, Any], + ) -> str: + if "image_grid_thw" in mm_kwargs: + return "image" + raise AssertionError("This line should be unreachable.") + + def get_encoder_cudagraph_budget_range( + self, + aphrodite_config, + ) -> tuple[int, int]: + # Min: estimated smallest possible encoder input. + # 224x224 image → 16x16 patches (patch_size=14) + # spatial_merge_size=2 → 8x8 = 64 tokens + min_budget = 64 + # Max: capped by max_num_batched_tokens + # TODO(shen-shanshan): the max_budget auto-infer needs to be optimized later. + max_budget = min( + aphrodite_config.scheduler_config.max_num_batched_tokens, + self.model_config.max_model_len, + ) + return (min_budget, max_budget) + + def _get_pixel_values_by_modality( + self, + mm_kwargs: dict[str, Any], + ) -> torch.Tensor: + modality = self.get_input_modality(mm_kwargs) + if modality == "image": + return mm_kwargs["pixel_values"] + raise AssertionError("This line should be unreachable.") + + def _get_grid_thw_by_modality( + self, + mm_kwargs: dict[str, Any], + ) -> list[list[int]]: + grid_thw_key = f"{self.get_input_modality(mm_kwargs)}_grid_thw" + grid_thw = mm_kwargs[grid_thw_key] + if not isinstance(grid_thw, list): + grid_thw = grid_thw.tolist() + return grid_thw + + def get_encoder_cudagraph_item_specs( + self, + mm_kwargs: dict[str, Any], + ): + from aphrodite.v1.worker.encoder_cudagraph_defs import EncoderItemSpec + + m = self.visual.spatial_merge_size + grid_thw = self._get_grid_thw_by_modality(mm_kwargs) + return [ + EncoderItemSpec( + input_size=t * h * w, + output_tokens=t * (h // m) * (w // m), + ) + for t, h, w in grid_thw + ] + + def select_encoder_cudagraph_items( + self, + mm_kwargs: dict[str, Any], + indices: list[int], + ) -> dict[str, Any]: + grid_thw = self._get_grid_thw_by_modality(mm_kwargs) + pixel_values = self._get_pixel_values_by_modality(mm_kwargs) + + if len(indices) == 0: + return { + "pixel_values": pixel_values[:0], + "image_grid_thw": [], + } + + # Compute cumulative patch offsets for slicing pixel_values + patches_per_item = [t * h * w for t, h, w in grid_thw] + cum_patches = [0] + for p in patches_per_item: + cum_patches.append(cum_patches[-1] + p) + + selected_pv = torch.cat([pixel_values[cum_patches[i] : cum_patches[i + 1]] for i in indices]) + selected_grid = [grid_thw[i] for i in indices] + + return { + "pixel_values": selected_pv, + "image_grid_thw": selected_grid, + } + + def prepare_encoder_cudagraph_capture_inputs( + self, + token_budget: int, + max_batch_size: int, + max_frames_per_batch: int, + device: torch.device, + dtype: torch.dtype, + path: str = "default", + ): + from aphrodite.v1.worker.encoder_cudagraph_defs import ( + EncoderCudaGraphCaptureInputs, + ) + + spatial_merge_size = self.visual.spatial_merge_size + # Ceil so the buffer fits the worst case of one item using the full + # budget. Floor under-allocates when budget is not a multiple of + # max_batch_size. + per_mm_item_output = (token_budget + max_batch_size - 1) // max_batch_size + + # Image-format grid_config (T=1). + grid_config = [[1, spatial_merge_size, per_mm_item_output * spatial_merge_size] for _ in range(max_batch_size)] + + # Create dummy pixel_values + patch_embed = self.visual.patch_embed + in_channels = patch_embed.proj.in_channels + patch_size = patch_embed.patch_size + temporal_patch_size = patch_embed.temporal_patch_size + total_patches = sum(t * h * w for t, h, w in grid_config) + flattened_patch_size = in_channels * temporal_patch_size * patch_size * patch_size + dummy_pixel_values = torch.randn(total_patches, flattened_patch_size, device=device, dtype=dtype) + + # Override max_seqlen with a safe upper bound for capture. + # max_seqlen.item() gets baked into the CUDA graph (not replayed), + # so the capture value must cover any replay scenario. + # Worst case: 1 item consuming the full budget -> + # seq_len = token_budget * spatial_merge_size^2. + metadata = self.visual.prepare_encoder_metadata( + grid_config, + max_batch_size=max_batch_size, + max_frames_per_batch=max_frames_per_batch, + max_seqlen_override=token_budget * (spatial_merge_size**2), + device=device, + ) + + # Just use image-modality dummy input_buffer for capturing + values = metadata | { + "pixel_values": dummy_pixel_values, + } + + return EncoderCudaGraphCaptureInputs( + values=values, + ) + + def prepare_encoder_cudagraph_replay_buffers( + self, + mm_kwargs: dict[str, Any], + max_batch_size: int, + max_frames_per_batch: int, + path: str = "default", + ) -> EncoderCudaGraphReplayBuffers: + modality = self.get_input_modality(mm_kwargs) + grid_thw_list = self._get_grid_thw_by_modality(mm_kwargs) + + if modality == "image": + metadata = self.visual.prepare_encoder_metadata( + grid_thw_list, + max_batch_size=max_batch_size, + ) + else: + raise AssertionError("This line should be unreachable.") + + values = metadata | { + "pixel_values": self._get_pixel_values_by_modality(mm_kwargs), + } + return EncoderCudaGraphReplayBuffers(values=values) + + def encoder_cudagraph_forward( + self, + values: dict[str, torch.Tensor], + path: str = "default", + ) -> torch.Tensor: + pixel_values = values.pop("pixel_values") + metadata = values + return self.visual(pixel_values, None, encoder_metadata=metadata) + + def encoder_eager_forward( + self, + mm_kwargs: dict[str, Any], + path: str = "default", + ) -> torch.Tensor: + pixel_values = self._get_pixel_values_by_modality(mm_kwargs) + grid_thw = self._get_grid_thw_by_modality(mm_kwargs) + return self.visual(pixel_values, grid_thw) + + def _parse_and_validate_image_input(self, **kwargs: object) -> Cohere_VLImageInputs | None: + pixel_values = kwargs.pop("pixel_values", None) + image_embeds = kwargs.pop("image_embeds", None) + image_grid_thw = kwargs.pop("image_grid_thw", None) + + if pixel_values is None and image_embeds is None: + return None + + if pixel_values is not None: + return Cohere_VLImagePixelInputs( + type="pixel_values", + pixel_values=pixel_values, + image_grid_thw=image_grid_thw, + ) + + if image_embeds is not None: + return Cohere_VLImageEmbeddingInputs( + type="image_embeds", + image_embeds=image_embeds, + image_grid_thw=image_grid_thw, + ) + raise AssertionError("image input must contain pixels or embeddings") + + def _process_image_input(self, image_input: Cohere_VLImageInputs) -> tuple[torch.Tensor, ...]: + grid_thw = image_input["image_grid_thw"] + assert grid_thw.ndim == 2 + + if image_input["type"] == "image_embeds": + image_embeds = image_input["image_embeds"].type(self.visual.dtype) + else: + pixel_values = image_input["pixel_values"].type(self.visual.dtype) + if self.use_data_parallel: + return run_dp_sharded_mrope_vision_model( + self.visual, pixel_values, grid_thw.tolist(), rope_type="rope_3d" + ) + else: + image_embeds = self.visual(pixel_values, grid_thw=grid_thw) + + # Split concatenated embeddings for each image item. + merge_size = self.visual.spatial_merge_size + sizes = (grid_thw.prod(-1) // merge_size // merge_size).tolist() + return image_embeds.split(sizes) + + def _postprocess_image_embeds_evs( + self, + image_embeds_split: tuple[torch.Tensor, ...], + image_input: Cohere_VLImageInputs, + ) -> tuple[torch.Tensor, ...]: + """ + Append mrope positions for each for images. + This is necessary to recover correct mrope positions + + Args: + image_embeds_split: Tuple of image embeddings for + each image item. + image_input: Image input data. + + Returns: + Tuple of image embeddings for each image item. + Resulting embeddings will have extra 5 channels for + computed mrope positions, consistent with video embeddings. + """ + if self.is_multimodal_pruning_enabled: + merge_size = self.visual.spatial_merge_size + grid_thw = image_input["image_grid_thw"] + grid_thw_list = grid_thw.tolist() + image_embeds_out = [] + for emb, size in zip(image_embeds_split, grid_thw_list): + positions = compute_mrope_for_media(size, merge_size).to(emb.device, non_blocking=True) + positions = torch.cat( + [ + positions, + torch.zeros_like(positions[:, 0:1]), # Dummy extra fifth channel + ], + dim=1, + ) + emb = torch.cat([emb, positions], dim=1) + image_embeds_out.append(emb) + image_embeds_split = tuple(image_embeds_out) + return image_embeds_split + + def _parse_and_validate_multimodal_inputs(self, **kwargs: object) -> dict: + mm_input_by_modality = {} + for input_key in kwargs: + if input_key in ("pixel_values", "image_embeds") and "image" not in mm_input_by_modality: + mm_input_by_modality["image"] = self._parse_and_validate_image_input(**kwargs) + return mm_input_by_modality + + def get_mrope_input_positions( + self, + input_tokens: list[int], + mm_features: list[MultiModalFeatureSpec], + ) -> tuple[torch.Tensor, int]: + llm_pos_ids_list: list[np.ndarray] = [] + st = 0 + + for ( + offset, + llm_grid_t, + llm_grid_h, + llm_grid_w, + _, + ) in self.iter_mm_grid_thw(mm_features): + text_len = offset - st + st_idx = llm_pos_ids_list[-1].max() + 1 if llm_pos_ids_list else 0 + llm_pos_ids_list.append(np.broadcast_to(np.arange(text_len), (3, text_len)) + st_idx) + grid_indices = np.indices((llm_grid_t, llm_grid_h, llm_grid_w)).reshape(3, -1) + llm_pos_ids_list.append(grid_indices + text_len + st_idx) + st = offset + llm_grid_t * llm_grid_h * llm_grid_w + + if st < len(input_tokens): + st_idx = llm_pos_ids_list[-1].max() + 1 if llm_pos_ids_list else 0 + text_len = len(input_tokens) - st + llm_pos_ids_list.append(np.broadcast_to(np.arange(text_len), (3, text_len)) + st_idx) + + llm_positions = np.concatenate(llm_pos_ids_list, axis=1).reshape(3, -1) + mrope_position_delta = (llm_positions.max() + 1 - len(input_tokens)).item() + return torch.from_numpy(llm_positions), mrope_position_delta + + def recompute_mrope_positions( + self, + input_ids: list[int], + multimodal_embeddings: MultiModalEmbeddings, + mrope_positions: torch.LongTensor, + num_computed_tokens: int, + ) -> tuple[MultiModalEmbeddings, torch.Tensor, int]: + """ + Update part of input mrope positions (starting with + num_computed_tokens index). Original mrope_positions are computed + for unpruned sequence and becomes incorrect once pruning occurs, + so once we prune media tokens we should reflect this in the + mrope_positions before we feed it to LLM. + + Args: + input_ids: (N,) All input tokens of the prompt containing + entire sequence. + multimodal_embeddings: Tuple of multimodal embeddings that + fits into the prefill chunk that is being processed. + mrope_positions: Existing mrope positions (3, N) for entire + sequence + num_computed_tokens: A number of computed tokens so far. + + Returns: + Tuple of (multimodal_embeddings, mrope_positions, + mrope_position_delta). + """ + return self._recompute_mrope_positions( + input_ids=input_ids, + multimodal_embeddings=multimodal_embeddings, + mrope_positions=mrope_positions, + num_computed_tokens=num_computed_tokens, + image_token_id=self.config.image_token_id, + video_token_id=self.config.video_token_id, + vision_start_token_id=self.config.vision_start_token_id, + ) + + @staticmethod + def _recompute_mrope_positions( + input_ids: list[int], + multimodal_embeddings: MultiModalEmbeddings, + mrope_positions: torch.LongTensor, + num_computed_tokens: int, + vision_start_token_id: int, + image_token_id: int, + video_token_id: int, + ) -> tuple[MultiModalEmbeddings, torch.Tensor, int]: + # Device + device = multimodal_embeddings[0].device if len(multimodal_embeddings) else mrope_positions.device + + # Tensors + input_ids_t = async_tensor_h2d(input_ids, device=device, dtype=torch.long) + + mm_embeddings_out = [] + mm_embeddings_pos = [] + # Strip position information from embeddings (last 5 channels) + # For Cohere VL, handle potentially empty frames (from unpacking) + for mm in multimodal_embeddings: + if mm.shape[0] > 0: # Only process non-empty frames + mm_embeddings_out.append(mm[:, :-5]) + mm_embeddings_pos.append(mm[:, -5:].permute(1, 0).long()) + else: + # Empty frame - keep as is + mm_embeddings_out.append(mm) + # Create empty position tensor with correct shape + mm_embeddings_pos.append(torch.empty(5, 0, device=device, dtype=torch.long)) + + positions, mrope_positions_delta = recompute_mrope_positions( + input_ids_t, + mm_embeddings_pos, + mrope_positions, + num_computed_tokens, + vision_start_token_id, + image_token_id, + video_token_id, + ) + + return tuple(mm_embeddings_out), positions, mrope_positions_delta + + def embed_multimodal(self, **kwargs: object) -> MultiModalEmbeddings | None: + mm_input_by_modality = self._parse_and_validate_multimodal_inputs(**kwargs) + if not mm_input_by_modality: + return None + + # The result multimodal_embeddings is tuple of tensors, with each + # tensor corresponding to a multimodal data item (image). + multimodal_embeddings: list[torch.Tensor] = [] + + # NOTE: It is important to iterate over the keys in this dictionary + # to preserve the order of the modalities. + for modality in mm_input_by_modality: + multimodal_input = mm_input_by_modality[modality] + if modality == "image": + image_embeddings = self._process_image_input(multimodal_input) + image_embeddings = self._postprocess_image_embeds_evs(image_embeddings, multimodal_input) + multimodal_embeddings.extend(image_embeddings) + + embeddings_tuple = tuple(multimodal_embeddings) + return embeddings_tuple + + def _compute_deepstack_embeds( + self, + inputs_embeds: torch.Tensor, + multimodal_embeddings: MultiModalEmbeddings, + is_multimodal: torch.Tensor, + ) -> tuple[torch.Tensor, MultiModalEmbeddings]: + visual_lens = [len(x) for x in multimodal_embeddings] + multimodal_embeddings_cat = torch.cat(multimodal_embeddings, dim=0) + + ( + multimodal_embeddings_main, + multimodal_embeddings_multiscale, + ) = torch.split( + multimodal_embeddings_cat, + [self.visual_dim, self.multiscale_dim], + dim=-1, + ) + + multimodal_embeddings = torch.split(multimodal_embeddings_main, visual_lens, dim=0) + multimodal_embeddings_multiscale = torch.split(multimodal_embeddings_multiscale, visual_lens, dim=0) + + deepstack_input_embeds = inputs_embeds.new_zeros( + inputs_embeds.size(0), self.deepstack_num_level * inputs_embeds.size(1) + ) + + deepstack_input_embeds = _merge_multimodal_embeddings( + inputs_embeds=deepstack_input_embeds, + multimodal_embeddings=multimodal_embeddings_multiscale, + is_multimodal=is_multimodal, + ) + deepstack_input_embeds = deepstack_input_embeds.view( + inputs_embeds.shape[0], self.deepstack_num_level, self.visual_dim + ) + deepstack_input_embeds = deepstack_input_embeds.permute(1, 0, 2) + + return deepstack_input_embeds, multimodal_embeddings + + def embed_input_ids( + self, + input_ids: torch.Tensor, + multimodal_embeddings: MultiModalEmbeddings | None = None, + *, + is_multimodal: torch.Tensor | None = None, + ) -> torch.Tensor: + inputs_embeds = self._embed_text_input_ids( + input_ids, + self.language_model.embed_input_ids, + is_multimodal=is_multimodal, + ) + + if multimodal_embeddings is None or len(multimodal_embeddings) == 0: + return inputs_embeds + + is_multimodal = _require_is_multimodal(is_multimodal) + + if self.use_deepstack: + ( + deepstack_input_embeds, + multimodal_embeddings, + ) = self._compute_deepstack_embeds( + inputs_embeds=inputs_embeds, + multimodal_embeddings=multimodal_embeddings, + is_multimodal=is_multimodal, + ) + else: + deepstack_input_embeds = None + + inputs_embeds = _merge_multimodal_embeddings( + inputs_embeds=inputs_embeds, + multimodal_embeddings=multimodal_embeddings, + is_multimodal=is_multimodal, + ) + + if deepstack_input_embeds is not None: + self._set_deepstack_input_embeds(deepstack_input_embeds) + + return inputs_embeds + + def forward( + self, + input_ids: torch.Tensor | None, + positions: torch.Tensor, + intermediate_tensors: IntermediateTensors | None = None, + inputs_embeds: torch.Tensor | None = None, + **kwargs: object, + ) -> torch.Tensor | IntermediateTensors: + """Run forward pass for Cohere Compass. + + Args: + input_ids: Flattened (concatenated) input_ids corresponding to a + batch. + positions: Flattened (concatenated) position ids corresponding to a + batch. + **NOTE**: If mrope is enabled (default setting for Cohere Compass + opensource models), the shape will be `(3, seq_len)`, + otherwise it will be `(seq_len,). + intermediate_tensors: Intermediate tensors from previous pipeline + stages. + inputs_embeds: Pre-computed input embeddings. + **kwargs: Additional keyword arguments including: + - pixel_values: Pixel values to be fed to a model. + `None` if no images are passed. + - image_grid_thw: Tensor `(n_images, 3)` of image 3D grid in + LLM. `None` if no images are passed. + """ + + if intermediate_tensors is not None: + inputs_embeds = None + + if inputs_embeds is not None and get_pp_group().is_first_rank: + deepstack_input_embeds = self._get_deepstack_input_embeds(inputs_embeds.size(0)) + else: + deepstack_input_embeds = None + + hidden_states = self.language_model.model( + input_ids=input_ids, + positions=positions, + intermediate_tensors=intermediate_tensors, + inputs_embeds=inputs_embeds, + # args for deepstack + deepstack_input_embeds=deepstack_input_embeds, + ) + + if inputs_embeds is not None and get_pp_group().is_first_rank: + self._clear_deepstack_input_embeds(inputs_embeds.size(0)) + + return hidden_states + + def compute_logits( + self, + hidden_states: torch.Tensor, + ) -> torch.Tensor | None: + return self.language_model.compute_logits(hidden_states) + + def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]: + loader = AutoWeightsLoader(self) + return loader.load_weights(weights, mapper=self.hf_to_aphrodite_mapper) + + def get_mm_mapping(self) -> MultiModelKeys: + """ + Get the module prefix in multimodal models + """ + return MultiModelKeys.from_string_field( + language_model="language_model", + connector=["visual.merger", "visual.deepstack_merger_list"], + tower_model="visual.", + ) + + def get_num_mm_encoder_tokens( + self, + num_image_tokens: int, + ) -> int: + hf_config = self.config + vision_config = hf_config.vision_config + merge_size = vision_config.spatial_merge_size + + return num_image_tokens * merge_size**2 + + def get_num_mm_connector_tokens( + self, + num_vision_tokens: int, + ) -> int: + hf_config = self.config + vision_config = hf_config.vision_config + merge_size = vision_config.spatial_merge_size + return num_vision_tokens // merge_size**2 + + +@lru_cache +def _cached_tensor(x, device) -> torch.Tensor: + return torch.tensor(x, device=device) + + +class CohereCompassAttention(nn.Module): + def __init__( + self, + config, + cache_config: CacheConfig | None = None, + quant_config: QuantizationConfig | None = None, + max_position_embeddings: int | None = None, + prefix: str = "", + ): + super().__init__() + tp_size = get_tensor_model_parallel_world_size() + layer_idx = extract_layer_index(prefix) + layer_type = config.layer_types[layer_idx] + + self.head_dim = config.head_dim + self.total_num_heads = config.num_attention_heads + self.num_heads = self.total_num_heads // tp_size + self.total_num_kv_heads = config.num_key_value_heads + if self.total_num_kv_heads >= tp_size: + assert self.total_num_kv_heads % tp_size == 0 + else: + assert tp_size % self.total_num_kv_heads == 0 + self.num_kv_heads = max(1, self.total_num_kv_heads // tp_size) + self.q_size = self.num_heads * self.head_dim + self.kv_size = self.num_kv_heads * self.head_dim + self.scaling = self.head_dim**-0.5 + + self.qkv_proj = QKVParallelLinear( + config.hidden_size, + self.head_dim, + self.total_num_heads, + self.total_num_kv_heads, + bias=config.attention_bias, + quant_config=quant_config, + prefix=f"{prefix}.qkv_proj", + ) + self.o_proj = RowParallelLinear( + self.total_num_heads * self.head_dim, + config.hidden_size, + bias=config.attention_bias, + quant_config=quant_config, + prefix=f"{prefix}.o_proj", + ) + + rope_parameters = config.rope_parameters.get(layer_type) + self.rotary_emb = ( + get_rope( + self.head_dim, + max_position=max_position_embeddings or config.max_position_embeddings, + rope_parameters=rope_parameters, + is_neox_style=True, + ) + if rope_parameters is not None + else None + ) + sliding_window = config.sliding_window if layer_type == "sliding_attention" else None + self.attn = Attention( + self.num_heads, + self.head_dim, + self.scaling, + num_kv_heads=self.num_kv_heads, + cache_config=cache_config, + quant_config=quant_config, + per_layer_sliding_window=sliding_window, + prefix=f"{prefix}.attn", + ) + + def forward( + self, + positions: torch.Tensor, + hidden_states: torch.Tensor, + ) -> torch.Tensor: + qkv, _ = self.qkv_proj(hidden_states) + q, k, v = qkv.split([self.q_size, self.kv_size, self.kv_size], dim=-1) + if self.rotary_emb is not None: + q, k = self.rotary_emb(positions, q, k) + attn_output = self.attn(q, k, v) + output, _ = self.o_proj(attn_output) + return output + + +class CohereCompassDecoderLayer(nn.Module): + def __init__( + self, + config, + cache_config: CacheConfig | None = None, + quant_config: QuantizationConfig | None = None, + max_position_embeddings: int | None = None, + prefix: str = "", + ): + super().__init__() + self.self_attn = CohereCompassAttention( + config, + cache_config, + quant_config=quant_config, + max_position_embeddings=max_position_embeddings, + prefix=f"{prefix}.self_attn", + ) + self.mlp = CohereMLP( + config, + quant_config=quant_config, + reduce_results=True, + prefix=f"{prefix}.mlp", + ) + self.input_layernorm = LayerNorm( + param_shape=config.hidden_size, + eps=config.layer_norm_eps, + ) + + def forward( + self, + positions: torch.Tensor, + hidden_states: torch.Tensor, + residual: torch.Tensor | None, + ) -> tuple[torch.Tensor, torch.Tensor]: + residual = hidden_states + hidden_states, residual = self.input_layernorm(hidden_states, residual) + hidden_states_attention = self.self_attn( + positions=positions, + hidden_states=hidden_states, + ) + hidden_states_mlp = self.mlp(hidden_states) + hidden_states = residual + hidden_states_attention + hidden_states_mlp + return hidden_states, residual + + +@support_torch_compile( + dynamic_arg_dims={ + "input_ids": 0, + "positions": -1, + "intermediate_tensors": 0, + "inputs_embeds": 0, + "deepstack_input_embeds": 0, + } +) +class CohereCompassTextModel(nn.Module): + def __init__(self, *, aphrodite_config: AphroditeConfig, prefix: str = ""): + super().__init__() + config = aphrodite_config.model_config.hf_config + cache_config = aphrodite_config.cache_config + quant_config = aphrodite_config.quant_config + + self.config = config + self.vocab_size = config.vocab_size + self.embed_tokens = VocabParallelEmbedding( + config.vocab_size, + config.hidden_size, + ) + self.start_layer, self.end_layer, self.layers = make_layers( + config.num_hidden_layers, + lambda prefix: CohereCompassDecoderLayer( + config, + cache_config, + quant_config=quant_config, + max_position_embeddings=aphrodite_config.model_config.max_model_len, + prefix=prefix, + ), + prefix=f"{prefix}.layers", + ) + self.norm = LayerNorm( + param_shape=config.hidden_size, + eps=config.layer_norm_eps, + ) + self.make_empty_intermediate_tensors = make_empty_intermediate_tensors_factory( + ["hidden_states", "residual"], config.hidden_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 = None, + inputs_embeds: torch.Tensor | None = None, + deepstack_input_embeds: IntermediateTensors | None = None, + ) -> torch.Tensor | IntermediateTensors: + if get_pp_group().is_first_rank: + hidden_states = inputs_embeds if inputs_embeds is not None else self.embed_input_ids(input_ids) + residual = None + else: + assert intermediate_tensors is not None + hidden_states = intermediate_tensors["hidden_states"] + residual = intermediate_tensors["residual"] + + for layer_idx, layer in islice(enumerate(self.layers), self.start_layer, self.end_layer): + hidden_states, residual = layer( + positions, + hidden_states, + residual, + ) + if deepstack_input_embeds is not None and layer_idx < len(deepstack_input_embeds): + hidden_states = hidden_states + deepstack_input_embeds[f"deepstack_input_embeds_{layer_idx}"] + + if not get_pp_group().is_last_rank: + return IntermediateTensors({"hidden_states": hidden_states, "residual": residual}) + hidden_states, _ = self.norm(hidden_states, residual) + return hidden_states + + +class CohereCompassForCausalLM(CohereForCausalLM): + def __init__(self, *, aphrodite_config: AphroditeConfig, prefix: str = ""): + nn.Module.__init__(self) + config = aphrodite_config.model_config.hf_config + assert config.tie_word_embeddings + + self.config = config + self.quant_config = aphrodite_config.quant_config + logit_scale = 1.0 if config.logit_scale is None else config.logit_scale + self.logits_processor = LogitsProcessor(config.vocab_size, scale=logit_scale) + self.model = CohereCompassTextModel( + aphrodite_config=aphrodite_config, + prefix=maybe_prefix(prefix, "model"), + ) + self.make_empty_intermediate_tensors = self.model.make_empty_intermediate_tensors + + +__all__ = ["CohereCompassForConditionalGeneration"] diff --git a/aphrodite/model_executor/models/registry.py b/aphrodite/model_executor/models/registry.py index eded571fc1..8166b1aefb 100644 --- a/aphrodite/model_executor/models/registry.py +++ b/aphrodite/model_executor/models/registry.py @@ -357,6 +357,10 @@ "cohere2_vision", "Cohere2VisionForConditionalGeneration", ), + "CohereCompassForConditionalGeneration": ( + "cohere_compass", + "CohereCompassForConditionalGeneration", + ), "Cosmos3ForConditionalGeneration": ("cosmos3", "Cosmos3ForConditionalGeneration"), "Cosmos3EdgeForConditionalGeneration": ( "cosmos3_edge", diff --git a/aphrodite/transformers_utils/config.py b/aphrodite/transformers_utils/config.py index e02aa069c4..270215e4ee 100644 --- a/aphrodite/transformers_utils/config.py +++ b/aphrodite/transformers_utils/config.py @@ -547,7 +547,9 @@ def _uses_mrope(config: PretrainedConfig) -> bool: if rope_parameters is None: return False - return "mrope_section" in rope_parameters + return "mrope_section" in rope_parameters or any( + isinstance(params, dict) and "mrope_section" in params for params in rope_parameters.values() + ) def uses_mrope(config: PretrainedConfig) -> bool: diff --git a/tests/models/registry.py b/tests/models/registry.py index 6b93135597..4f6dbe18f8 100644 --- a/tests/models/registry.py +++ b/tests/models/registry.py @@ -688,6 +688,10 @@ def check_available_online( extras={"command-a-plus": "CohereLabs/command-a-plus-05-2026-bf16"}, min_transformers_version="5.9.0", ), + "CohereCompassForConditionalGeneration": _HfExamplesInfo( + "CohereLabs/North-Micro-Vision-Instruct", + min_transformers_version="5.16.0", + ), "Cosmos3ForConditionalGeneration": _HfExamplesInfo( "nvidia/Cosmos3-Nano", extras={"super": "nvidia/Cosmos3-Super"}, From d65368b22fed27471386d2db23a5cf1cfef24bfd Mon Sep 17 00:00:00 2001 From: AlpinDale Date: Mon, 7 Sep 2026 21:23:18 +0000 Subject: [PATCH 36/36] [sync] [CI/Build] Unskip ColQwen3 multimodal pooling tests on Transformers v5 (#55588) Upstream-vLLM: 51da0ca66c8065619c79e35dff97aa99aeaf5644 Co-authored-by: Sahil Patel <91423311+Sip4818@users.noreply.github.com> --- .sync/vllm-sha | 2 +- tests/models/multimodal/pooling/test_colqwen3.py | 4 ---- 2 files changed, 1 insertion(+), 5 deletions(-) diff --git a/.sync/vllm-sha b/.sync/vllm-sha index 1a13682dd0..c27e5a6a0d 100644 --- a/.sync/vllm-sha +++ b/.sync/vllm-sha @@ -1 +1 @@ -70584f69b1bf866224604029fae038864a368019 +51da0ca66c8065619c79e35dff97aa99aeaf5644 diff --git a/tests/models/multimodal/pooling/test_colqwen3.py b/tests/models/multimodal/pooling/test_colqwen3.py index fd1a9d6895..850d1f08ef 100644 --- a/tests/models/multimodal/pooling/test_colqwen3.py +++ b/tests/models/multimodal/pooling/test_colqwen3.py @@ -22,10 +22,6 @@ from ....conftest import AphroditeRunner -pytestmark = pytest.mark.skip( - reason="ColQwen3 model's weight tying is incompatible with transformers v5 (missing all_tied_weights_keys)" -) - MODELS = [ "TomoroAI/tomoro-colqwen3-embed-4b", "OpenSearch-AI/Ops-Colqwen3-4B",