[ROCm][MLA][Bugfix] Reserve FP8 prefill workspace before lock for Kimi-K2.5 (#42978)

Signed-off-by: Xavier Aguilar <xavier.aguilarfruto@amd.com>
Co-authored-by: Douglas Lehr <91553416+dllehr-amd@users.noreply.github.com>
Co-authored-by: Andreas Karatzas <akaratza@amd.com>
This commit is contained in:
xaguilar-amd
2026-06-08 22:25:52 -07:00
committed by GitHub
co-authored by Douglas Lehr Andreas Karatzas
parent d8218b1ee7
commit baacbfcebf
@@ -242,7 +242,12 @@ class AiterMLAMetadataBuilder(MLACommonMetadataBuilder[AiterMLAMetadata]):
vllm_config.model_config.max_model_len,
vllm_config.scheduler_config.max_num_batched_tokens,
)
self._init_fp8_prefill_ps_buffers(max_num_reqs, max_prefill_qlen, device)
self._init_fp8_prefill_ps_buffers(
max_num_reqs,
max_prefill_qlen,
vllm_config.scheduler_config.max_num_batched_tokens,
device,
)
if self.compilation_config.cudagraph_mode.has_full_cudagraphs():
self.paged_kv_indptr = torch.zeros(
@@ -257,21 +262,29 @@ class AiterMLAMetadataBuilder(MLACommonMetadataBuilder[AiterMLAMetadata]):
self,
max_num_reqs: int,
max_prefill_qlen: int,
max_num_batched_tokens: int,
device: torch.device,
) -> None:
"""Pre-allocate persistent buffers for FP8 MLA prefill PS metadata.
Uses ``get_ps_metadata_info_v1`` with max values so the buffers are
large enough for any batch. ``get_ps_metadata_v1`` fills them
per-batch in ``build()``.
per-batch in ``build()``. The FP8 prefill forward path also uses the
global workspace manager for per-call scratch, so reserve its maximum
shape here before the workspace manager is locked after warmup.
Args:
max_num_reqs: Maximum number of concurrent requests.
max_prefill_qlen: Maximum Q-length for a single request in one
prefill batch. Should be ``min(max_model_len,
max_num_batched_tokens)`` — the chunked-prefill scheduler
never emits more than ``max_num_batched_tokens`` new tokens
per batch.
max_num_batched_tokens)`` — a single request never exceeds
``max_model_len`` tokens, nor the per-batch token budget.
max_num_batched_tokens: Maximum number of tokens scheduled in one
batch. The ``final_lse`` scratch is sized by ``total_q`` (the
summed Q-length over all prefill requests in the batch), which
is bounded by this budget rather than by a single request's
``max_prefill_qlen`` — concurrent requests can sum to more than
``max_model_len`` when ``max_model_len < max_num_batched_tokens``.
device: Target device for the buffers.
"""
from aiter import get_ps_metadata_info_v1
@@ -279,6 +292,7 @@ class AiterMLAMetadataBuilder(MLACommonMetadataBuilder[AiterMLAMetadata]):
# After kv_b_proj decompression, K has num_heads heads (same as Q).
# So gqa_ratio=1 and num_head_k=num_heads for the PS kernel.
num_head_k = self.num_heads
v_head_dim = self.mla_dims.v_head_dim
# gqa_ratio = 1
# qlen_granularity = _FP8_PREFILL_TILE_Q // max(gqa_ratio, 1)
qlen_granularity = _FP8_PREFILL_TILE_Q
@@ -318,6 +332,21 @@ class AiterMLAMetadataBuilder(MLACommonMetadataBuilder[AiterMLAMetadata]):
device=device,
)
from vllm.v1.worker.workspace import current_workspace_manager
max_num_partial_tiles = reduce_partial_map_size
current_workspace_manager().get_simultaneous(
(
(max_num_partial_tiles * _FP8_PREFILL_TILE_Q, num_head_k, v_head_dim),
torch.float32,
),
(
(max_num_partial_tiles * _FP8_PREFILL_TILE_Q, num_head_k),
torch.float32,
),
((max_num_batched_tokens, num_head_k), torch.float32),
)
logger.info(
"FP8 MLA prefill PS buffers allocated "
"(max_batch=%d, max_qlen=%d, num_head_k=%d)",