[Bug] Forward fixed_split_size in SWA / cross-attention paths of FlashInfer backend (#26412)

Signed-off-by: Xingyu Liu <charlotteliu12x@gmail.com>
This commit is contained in:
Xingyu Liu
2026-05-28 00:52:26 -07:00
committed by GitHub
parent 8dca6291c7
commit 770c51b127
@@ -1059,6 +1059,8 @@ class FlashInferIndicesUpdaterDecode:
spec_info,
seq_lens_cpu=seq_lens_cpu_tmp,
use_sliding_window_kv_pool=use_sliding_window_kv_pool,
fixed_split_size=fixed_split_size,
disable_split_kv=disable_split_kv,
)
def update_cross_attention(
@@ -1096,6 +1098,8 @@ class FlashInferIndicesUpdaterDecode:
kv_start_idx,
spec_info,
seq_lens_cpu=kv_lens_cpu,
fixed_split_size=fixed_split_size,
disable_split_kv=disable_split_kv,
)
def call_begin_forward(
@@ -1340,6 +1344,7 @@ class FlashInferIndicesUpdaterPrefill:
use_ragged,
spec_info,
use_sliding_window_kv_pool=use_sliding_window_kv_pool,
fixed_split_size=fixed_split_size,
multi_item_params=multi_item_params,
)
@@ -1383,6 +1388,7 @@ class FlashInferIndicesUpdaterPrefill:
self.qo_indptr[wrapper_id],
use_ragged,
spec_info,
fixed_split_size=fixed_split_size,
multi_item_params=multi_item_params,
cross_attention_custom_mask=(
cross_attention_custom_mask if wrapper_id == 1 else None