From 0e7fe57d96928a64d902bf712304a380796b5490 Mon Sep 17 00:00:00 2001 From: ViperEkura <3081035982@qq.com> Date: Fri, 7 Aug 2026 14:42:53 +0800 Subject: [PATCH] fix: use max_context_len for stable num_splits in paged decode - PagedKV::host_kv_len now returns max_context_len instead of max_seq_len - Eliminates grid-z instability for CUDA graph capture/replay - Restore skip_no_kernel re-export accidentally removed by ruff --fix --- csrc/kernels/attn_kv_source.cuh | 2 +- tests/extension/conftest.py | 1 + 2 files changed, 2 insertions(+), 1 deletion(-) diff --git a/csrc/kernels/attn_kv_source.cuh b/csrc/kernels/attn_kv_source.cuh index 416cd9c..4c181cd 100644 --- a/csrc/kernels/attn_kv_source.cuh +++ b/csrc/kernels/attn_kv_source.cuh @@ -111,7 +111,7 @@ struct PagedKV { return p.max_q_len; } HOST_DEV_FORCEINLINE int host_kv_len(const AttentionParams& p) { - return p.max_seq_len; + return p.max_context_len; } // prefill: Q rows start at qo_indptr[batch] (ragged batch base) diff --git a/tests/extension/conftest.py b/tests/extension/conftest.py index 8e3e373..e0c18ac 100644 --- a/tests/extension/conftest.py +++ b/tests/extension/conftest.py @@ -5,6 +5,7 @@ import torch from astrai.config.model_config import AutoRegressiveLMConfig from astrai.model.transformer import AutoRegressiveLM +from tests.conftest import skip_no_kernel # noqa: F401 re-export for test modules D = 64 CFG = dict(