Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion cmake/external_projects/flashkda.cmake
Original file line number Diff line number Diff line change
Expand Up @@ -13,7 +13,7 @@ else()
FetchContent_Declare(
flashkda
GIT_REPOSITORY https://github.com/vllm-project/FlashKDA.git
GIT_TAG 053de1b716ef3255873e02d2d28f4adf09951978
GIT_TAG ee0be888cd0e972f9409bf53756f8c38c6652173
GIT_PROGRESS TRUE
GIT_SUBMODULES cutlass
)
Expand Down
3 changes: 2 additions & 1 deletion csrc/flashkda_registration.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -10,7 +10,8 @@ STABLE_TORCH_LIBRARY(_flashkda_C, m) {
"Tensor(a!) out, Tensor(c!) workspace, Tensor A_log, Tensor dt_bias, "
"float lower_bound, "
"Tensor? initial_state=None, Tensor(b!)? final_state=None, "
"Tensor? cu_seqlens=None) -> ()");
"Tensor? cu_seqlens=None, Tensor(d!)? checkpoint_state=None, "
"Tensor? checkpoint_offsets=None) -> ()");
}

STABLE_TORCH_LIBRARY_IMPL(_flashkda_C, CompositeExplicitAutograd, m) {
Expand Down
81 changes: 81 additions & 0 deletions tests/models/kimi_k3/test_kda.py
Original file line number Diff line number Diff line change
Expand Up @@ -21,6 +21,8 @@
fused_recurrent_kda_packed_decode as fused_recurrent_kda_packed_decode_amd,
)
from vllm.models.kimi_k3.nvidia.kda import (
_flashkda_prefill,
_store_cache_checkpoints_kernel,
is_flashkda_supported,
is_fused_kda_decode_supported,
)
Expand All @@ -40,6 +42,7 @@
)
from vllm.platforms import current_platform
from vllm.third_party.flash_linear_attention.ops.l2norm import l2norm_fwd
from vllm.v1.attention.backends.utils import NULL_BLOCK_ID

DEVICE = current_platform.device_type

Expand Down Expand Up @@ -1088,6 +1091,17 @@ def test_flashkda_correctness():
expected_states.append(final_state)
expected_out = torch.cat(expected_outputs, dim=1)
expected_state = torch.cat(expected_states).transpose(-1, -2).contiguous()
_, expected_checkpoint = naive_recurrent_kda(
q_norm[:, :16],
k_norm[:, :16],
v[:, :16],
gate[:, :16],
beta[:, :16],
initial_state=initial_state[0:1].transpose(-1, -2),
output_final_state=True,
)
assert expected_checkpoint is not None
expected_checkpoint = expected_checkpoint.transpose(-1, -2).contiguous()

actual_out = torch.empty_like(v)
actual_state = torch.empty_like(initial_state)
Expand Down Expand Up @@ -1115,3 +1129,70 @@ def test_flashkda_correctness():

assert_close("o", expected_out, actual_out, 0.01)
assert_close("ht", expected_state, actual_state, 0.01)

checkpoint_out = torch.empty_like(v)
checkpoint_final_state = torch.empty_like(initial_state)
checkpoint_state = torch.empty_like(initial_state)
checkpoint_offsets = torch.tensor([16, 31], dtype=torch.int32, device=DEVICE)
_flashkda_prefill(
q=q,
k=k,
v=v,
g=raw_g,
beta=beta_logits,
A_log=A_log,
dt_bias=dt_bias,
lower_bound=lower_bound,
initial_state=initial_state,
cu_seqlens=cu_seqlens,
out=checkpoint_out,
final_state=checkpoint_final_state,
workspace=workspace,
checkpoint_state=checkpoint_state,
checkpoint_offsets=checkpoint_offsets,
)

assert_close("checkpoint_o", expected_out, checkpoint_out, 0.01)
assert_close("checkpoint_ht", expected_state, checkpoint_final_state, 0.01)
assert_close("checkpoint", expected_checkpoint, checkpoint_state[:1], 0.01)

conv_state = torch.zeros(2, H * D, 3, dtype=q.dtype, device=DEVICE)
recurrent_storage = torch.zeros(2, H * D * D + 8, device=DEVICE)
recurrent_state = recurrent_storage[:, : H * D * D].view(2, H, D, D)
conv_input = q[0].flatten(1)
checkpoint_state_indices = torch.tensor(
[1, NULL_BLOCK_ID], dtype=torch.int32, device=DEVICE
)
state_len = conv_state.shape[-1]
width = H * D
recurrent_row_size = checkpoint_state[0].numel()
block_size = 256
_store_cache_checkpoints_kernel[
(
checkpoint_state_indices.numel(),
(max(width * state_len, recurrent_row_size) + block_size - 1) // block_size,
)
](
conv_input,
conv_state,
checkpoint_state,
recurrent_state,
cu_seqlens,
checkpoint_offsets,
checkpoint_state_indices,
conv_input.stride(0),
conv_input.stride(1),
conv_state.stride(0),
conv_state.stride(1),
conv_state.stride(2),
checkpoint_state.stride(0),
recurrent_state.stride(0),
checkpoint_offsets.stride(0),
state_len,
width,
recurrent_row_size,
NULL_BLOCK_ID,
block_size,
)
torch.testing.assert_close(conv_state[1], q[0, 13:16].flatten(1).transpose(0, 1))
torch.testing.assert_close(recurrent_state[1], checkpoint_state[0])
50 changes: 50 additions & 0 deletions tests/models/kimi_k3/test_kda_metadata.py
Original file line number Diff line number Diff line change
Expand Up @@ -87,6 +87,7 @@ def _make_builder(
device: torch.device = DEVICE,
mamba_cache_mode: str = "none",
use_recoverssm: bool = False,
num_prefill_checkpoint_blocks: int = 0,
) -> AttentionMetadataBuilder:
vllm_config = create_vllm_config(
model_name="Qwen/Qwen3.5-0.8B",
Expand All @@ -110,6 +111,7 @@ def _make_builder(
dtypes=(torch.float16,),
mamba_cache_mode=mamba_cache_mode,
num_speculative_blocks=(0 if use_recoverssm else num_speculative_tokens),
num_prefill_checkpoint_blocks=num_prefill_checkpoint_blocks,
),
layer_names=["layer.0"],
vllm_config=vllm_config,
Expand All @@ -121,6 +123,54 @@ def _make_builder(
return builder


@pytest.mark.skipif(not torch.cuda.is_available(), reason="requires CUDA")
def test_internal_checkpoint_metadata_targets_last_aligned_boundary():
device = torch.device("cuda")
batch = BatchSpec(seq_lens=[50, 32], query_lens=[50, 16])
common_attn_metadata = create_common_attn_metadata(
batch, BLOCK_SIZE, device, arange_block_indices=True
).replace(is_prefilling=torch.tensor([True, True]))
actual = _make_builder(
KimiK3KDAMetadataBuilder,
num_speculative_tokens=0,
full_cuda_graph=False,
mamba_cache_mode="align",
num_prefill_checkpoint_blocks=1,
device=device,
).build(0, common_attn_metadata)

assert actual.checkpoint is not None
torch.testing.assert_close(
actual.checkpoint.state_indices,
torch.tensor([2, NULL_BLOCK_ID], dtype=torch.int32, device=device),
)
torch.testing.assert_close(
actual.checkpoint.checkpoint_offsets,
torch.tensor([48, 0], dtype=torch.int32, device=device),
)


@pytest.mark.skipif(not torch.cuda.is_available(), reason="requires CUDA")
def test_internal_checkpoint_metadata_skips_unaligned_offset():
device = torch.device("cuda")
# The checkpoint block boundary is 48, but this query starts at token 1,
# making the real checkpoint offset 47, which is not FlashKDA-aligned.
batch = BatchSpec(seq_lens=[50], query_lens=[49])
common_attn_metadata = create_common_attn_metadata(
batch, BLOCK_SIZE, device, arange_block_indices=True
).replace(is_prefilling=torch.tensor([True]))
actual = _make_builder(
KimiK3KDAMetadataBuilder,
num_speculative_tokens=0,
full_cuda_graph=False,
mamba_cache_mode="align",
num_prefill_checkpoint_blocks=1,
device=device,
).build(0, common_attn_metadata)

assert actual.checkpoint is None


@pytest.mark.parametrize(
(
"batch",
Expand Down
85 changes: 85 additions & 0 deletions tests/v1/core/prefix_cache/test_partial_prefix_cache_hits.py
Original file line number Diff line number Diff line change
Expand Up @@ -77,6 +77,7 @@ def make_full_mamba_manager(
mamba_block_size: int = 4,
num_blocks: int = 32,
use_eagle: bool = False,
num_prefill_checkpoint_blocks: int = 0,
):
kv_cache_config = KVCacheConfig(
num_blocks=num_blocks,
Expand All @@ -98,6 +99,7 @@ def make_full_mamba_manager(
shapes=(1, 1),
dtypes=(torch.float32,),
mamba_cache_mode="align",
num_prefill_checkpoint_blocks=num_prefill_checkpoint_blocks,
),
),
],
Expand Down Expand Up @@ -135,6 +137,7 @@ def test_mamba_align_split_partial_tail_schedule(dcp_world_size: int):
dcp_world_size=dcp_world_size,
scheduler_block_size=scheduler_block_size,
mamba_partial_cache_hit=True,
mamba_has_prefill_checkpoint_blocks=False,
)
split = Scheduler._mamba_block_aligned_split

Expand Down Expand Up @@ -180,6 +183,7 @@ def test_mamba_align_split_when_block_exceeds_scheduling_budget():
use_eagle=False,
hash_block_size=32,
mamba_partial_cache_hit=False,
mamba_has_prefill_checkpoint_blocks=False,
)
req = make_request("0", [0] * prompt_length, 32, sha256)
split = Scheduler._mamba_block_aligned_split
Expand Down Expand Up @@ -218,6 +222,7 @@ def test_mamba_align_split_when_block_exceeds_long_prefill_threshold():
use_eagle=False,
hash_block_size=32,
mamba_partial_cache_hit=False,
mamba_has_prefill_checkpoint_blocks=False,
)
req = make_request("0", [0] * prompt_length, 32, sha256)
split = Scheduler._mamba_block_aligned_split
Expand Down Expand Up @@ -462,6 +467,85 @@ def test_hybrid_mamba_partial_tail_owner_uses_cow_on_continue():
assert moved[0].block_hash_num_tokens == 6


def test_partial_hit_then_internal_checkpoint_uses_distinct_mamba_blocks():
hash_block_size = 2
mamba_block_size = 4
manager = make_full_mamba_manager(
dcp_world_size=1,
hash_block_size=hash_block_size,
full_block_size=hash_block_size,
mamba_block_size=mamba_block_size,
num_prefill_checkpoint_blocks=1,
)

owner = make_request("owner", [0, 0, 1, 1, 2, 2], hash_block_size, sha256)
computed_blocks, num_computed, _ = manager.get_computed_blocks(owner)
assert manager.allocate_slots(owner, 6, num_computed, computed_blocks) is not None
manager.free(owner)
manager.new_step_starts()

partial_hash = owner.block_hashes[2]
partial_block = manager.block_pool.get_cached_block(partial_hash, [1])
assert partial_block is not None
partial_block_id = partial_block[0].block_id

replay = make_request(
"replay",
[0, 0, 1, 1, 2, 2, 3, 3, 4, 4, 5, 5, 6, 6],
hash_block_size,
sha256,
)
computed_blocks, num_computed, _ = manager.get_computed_blocks(replay)
assert num_computed == 6

realigned_blocks = manager.allocate_slots(replay, 2, num_computed, computed_blocks)
assert realigned_blocks is not None
mamba_cow_block_id = realigned_blocks.get_block_ids()[1][0]
copies, retained = manager.take_kv_cache_block_copies()
assert KVCacheBlockCopy(partial_block_id, mamba_cow_block_id) in copies
manager.block_pool.free_blocks(retained)

replay.num_computed_tokens = 8
manager.new_step_starts()
final_blocks = manager.allocate_slots(replay, 6)
assert final_blocks is not None

checkpoint_block_id, running_block_id = final_blocks.get_block_ids()[1]
assert len({mamba_cow_block_id, checkpoint_block_id, running_block_id}) == 3
mamba_blocks = manager.get_blocks(replay.request_id).blocks[1]
assert mamba_blocks[2].block_id == checkpoint_block_id
assert mamba_blocks[3].block_id == running_block_id


def test_internal_checkpoint_requires_block_aligned_start():
hash_block_size = 2
mamba_block_size = 16
manager = make_full_mamba_manager(
dcp_world_size=1,
hash_block_size=hash_block_size,
full_block_size=hash_block_size,
mamba_block_size=mamba_block_size,
num_prefill_checkpoint_blocks=1,
)
request = make_request("producer", list(range(50)), hash_block_size, sha256)

# Compute token 0 first, so the next query starts at token 1, which is not
# aligned to the Mamba block size.
assert manager.allocate_slots(request, 1) is not None
request.num_computed_tokens = 1
manager.new_step_starts()

new_blocks = manager.allocate_slots(request, 49)

assert new_blocks is not None
mamba_blocks = manager.get_blocks(request.request_id).blocks[1]
checkpoint_block_idx = 48 // mamba_block_size - 1
assert mamba_blocks[checkpoint_block_idx].is_null
assert not mamba_blocks[-1].is_null
checkpoint_hash = request.block_hashes[48 // hash_block_size - 1]
assert manager.block_pool.get_cached_block(checkpoint_hash, [1]) is None


def test_external_mamba_hit_same_block_uses_running_cow_on_continue():
"""An external mid-block hit must become a running request even when its
first continuation does not need another Mamba block."""
Expand Down Expand Up @@ -560,6 +644,7 @@ def test_take_partial_tail_offloads_returns_cow_target():
shapes=(1, 1),
dtypes=(torch.float32,),
mamba_cache_mode="align",
num_prefill_checkpoint_blocks=1,
),
),
],
Expand Down
Loading
Loading