Skip to content
88 changes: 88 additions & 0 deletions tests/models/kimi_k3/test_kda_metadata.py
Original file line number Diff line number Diff line change
Expand Up @@ -934,3 +934,91 @@ def test_aligned_block_table_matches_shared_gdn():
)

torch.testing.assert_close(actual, expected)


def _build_non_spec(batch, is_prefilling, full_cuda_graph=False):
common_attn_metadata = create_common_attn_metadata(
batch, BLOCK_SIZE, DEVICE
).replace(
is_prefilling=None
if is_prefilling is None
else torch.tensor(is_prefilling, dtype=torch.bool)
)
builder = _make_builder(
KimiK3KDAMetadataBuilder,
num_speculative_tokens=0,
full_cuda_graph=full_cuda_graph,
)
return builder, common_attn_metadata, builder.build(0, common_attn_metadata)


def test_one_token_first_chunk_excludes_padding():
"""Neither padding requests nor padding tokens count as prefill work."""
common = create_common_attn_metadata(
BatchSpec(seq_lens=[100, 1, 0, 0], query_lens=[1, 1, 0, 0]),
BLOCK_SIZE,
DEVICE,
).replace(
is_prefilling=torch.tensor([False, True, False, False], dtype=torch.bool),
num_actual_tokens=4,
)
builder = _make_builder(
KimiK3KDAMetadataBuilder, num_speculative_tokens=0, full_cuda_graph=False
)
actual = builder.build(0, common)

assert actual.num_decodes == 1
assert actual.num_prefills == 1
assert actual.num_decode_tokens == 1
assert actual.num_prefill_tokens == 1


@pytest.mark.parametrize(
("seq_len", "query_len", "is_prefilling", "num_prefills"),
[
pytest.param(1, 1, True, 1, id="first-chunk"),
pytest.param(65, 1, True, 0, id="resumed-chunk"),
pytest.param(0, 0, True, 0, id="padding"),
pytest.param(1, 1, None, 0, id="missing-prefill-flag"),
],
)
def test_one_token_chunk_classification(
seq_len, query_len, is_prefilling, num_prefills
):
"""Only a real first chunk with a prefill flag needs state initialization."""
_, _, actual = _build_non_spec(
BatchSpec(seq_lens=[100, seq_len], query_lens=[1, query_len]),
is_prefilling=None if is_prefilling is None else [False, is_prefilling],
)

assert actual.num_prefills == num_prefills
assert actual.num_decodes == 2 - num_prefills
assert actual.num_prefill_tokens == num_prefills
assert actual.num_decode_tokens == 1 + query_len - num_prefills
if num_prefills:
assert actual.has_initial_state is not None
assert actual.has_initial_state.tolist() == [True, False]
else:
assert actual.has_initial_state is None


def test_cudagraph_capture_batch_stays_decode_only():
"""Capture rows have no history, but must still select decode kernels."""
batch = BatchSpec(seq_lens=[1] * 4, query_lens=[1] * 4)
common_attn_metadata = create_common_attn_metadata(
batch, BLOCK_SIZE, DEVICE
).replace(is_prefilling=torch.zeros(4, dtype=torch.bool))
builder = _make_builder(
KimiK3KDAMetadataBuilder,
num_speculative_tokens=0,
full_cuda_graph=True,
)
actual = builder.build_for_cudagraph_capture(common_attn_metadata)

assert actual.num_prefills == 0
assert actual.num_decodes == 4
assert actual.has_initial_state is None
staged = actual.non_spec_state_indices_tensor
assert staged is not None
assert staged.data_ptr() == builder.non_spec_state_indices_tensor.data_ptr()
torch.testing.assert_close(staged, common_attn_metadata.block_table_tensor[:, 0])
27 changes: 25 additions & 2 deletions vllm/models/kimi_k3/nvidia/kda_metadata.py
Original file line number Diff line number Diff line change
Expand Up @@ -427,10 +427,33 @@ def build( # type: ignore[override]
spec_token_start = None
non_spec_token_start = None
if num_spec_decodes == 0:
# The runner orders ordinary decodes before prefills.
# V2 already excludes prefills from full decode graphs via has_prefill.
# Classify first chunks as prefills to mask recycled state;
# resumed one-token chunks can still use the decode kernels.
assert m.seq_lens_cpu_upper_bound is not None
query_lens_cpu = query_start_loc_cpu.diff()
no_prior_state = (query_lens_cpu > 0) & (
m.seq_lens_cpu_upper_bound <= query_lens_cpu
)
# Capture batches also have seq_len == query_len, but are not prefills.
if m.is_prefilling is not None:
no_prior_state &= m.is_prefilling
else:
no_prior_state = torch.zeros_like(no_prior_state)
num_decodes, num_prefills, num_decode_tokens, num_prefill_tokens = (
split_decodes_and_prefills(m, decode_threshold=1)
split_decodes_and_prefills(
m.replace(is_prefilling=no_prior_state),
decode_threshold=1,
treat_short_extends_as_decodes=False,
)
)
# Exclude trailing padding from both prefill counts.
if num_prefills:
num_prefills -= int((query_lens_cpu[num_decodes:] == 0).sum())
num_prefill_tokens = (
int(query_start_loc_cpu[num_decodes + num_prefills])
- num_decode_tokens
)
num_spec_decode_tokens = 0
spec_token_indx = None
non_spec_token_indx = None
Expand Down
Loading