Skip to content
Open
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
Original file line number Diff line number Diff line change
Expand Up @@ -1298,6 +1298,10 @@ def fused_kernel123(
for vi in cutlass.range_constexpr(VEC):
rPrefix[vi] = sPartialLast[warp_row_group - 1, col_base + vi]

# The next chunk's partial sums overwrite sPartialLast, so every
# K1 warp must have read its prefix and the chunk total first.
k1_internal_barrier()

# ---- Pass 2a: ONLY cumsum + write csGcum (critical path, minimal work) ----
for vi in cutlass.range_constexpr(VEC):
rAcc[vi] = rPrefix[vi]
Expand Down
125 changes: 125 additions & 0 deletions tests/unittest/_torch/modules/kimi_kda/test_kda_prefill_op.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,7 +2,10 @@
# SPDX-License-Identifier: Apache-2.0
"""Parity tests for the optimized Kimi K3 KDA prefill op."""

import importlib.util
import sys
from collections.abc import Mapping
from pathlib import Path
from types import SimpleNamespace

import pytest
Expand Down Expand Up @@ -488,6 +491,128 @@ def randn(
_assert_close(state_fused, state_unfused)


# fused_k123's K1 warps meet at this barrier after the cross-warp scan, then
# read the scanned prefix rows and the chunk total from sPartialLast.
_K1_SCAN_BARRIER = " k1_internal_barrier()\n"
_K1_SCAN_READS = (
_K1_SCAN_BARRIER
+ "\n"
+ " for vi in cutlass.range_constexpr(VEC):\n"
+ " rGkLast[vi] = sPartialLast[K1_ROW_GROUPS - 1, col_base + vi]\n"
)
# K1 warp 1 reads the prefix row 0 and the chunk total.
_LATE_K1_WARP = (
" if k1_warp == 1:\n"
" late_t = _test_globaltimer()\n"
" late_end = late_t + cutlass.Int64(20000)\n"
" while late_t < late_end:\n"
" late_t = _test_globaltimer()\n"
)
_GLOBALTIMER = """

@dsl_user_op
def _test_globaltimer(*, loc=None, ip=None):
return cutlass.Int64(
llvm.inline_asm(
T.i64(),
[],
"mov.u64 $0, %globaltimer;",
"=l",
has_side_effects=True,
is_align_stack=False,
asm_dialect=llvm.AsmDialect.AD_ATT,
loc=loc,
ip=ip,
)
)
"""


@torch.no_grad()
@pytest.mark.parametrize(
"sequence_lengths",
[None, [100, 257, 63]],
ids=["padded_eqlen", "varlen"],
)
def test_fused_prefill_does_not_depend_on_k1_warp_timing(
sequence_lengths: list[int] | None, tmp_path: Path, monkeypatch: pytest.MonkeyPatch
) -> None:
"""A K1 warp that reads the scanned partial sums late gets the same bits.

A copy of fused_k123 stalls K1 warp 1 for 20 us right after the scan
barrier, so the other K1 warps run ahead into the next chunk before it
reads sPartialLast. Outputs and states must equal the unpatched kernel's.
"""
from tensorrt_llm._torch.custom_ops import cute_dsl_kimi_k3_custom_ops as kimi_k3_ops
from tensorrt_llm._torch.cute_dsl_kernels.blackwell.kimi_k3_kda import fused_k123

source = Path(fused_k123.__file__).read_text()
assert source.count(_K1_SCAN_READS) == 1, "fused_k123's K1 scan reads moved"
late_source = source.replace(
_K1_SCAN_READS, _K1_SCAN_BARRIER + _LATE_K1_WARP + _K1_SCAN_READS[len(_K1_SCAN_BARRIER) :]
)
late_path = tmp_path / "fused_k123_late_k1_warp.py"
late_path.write_text(late_source + _GLOBALTIMER)
spec = importlib.util.spec_from_file_location(late_path.stem, late_path)
late_module = importlib.util.module_from_spec(spec)
monkeypatch.setitem(sys.modules, late_path.stem, late_module)
spec.loader.exec_module(late_module)

torch.manual_seed(10)
if sequence_lengths is None:
sequence_length = 300
num_sequences = 1
cu_seqlens = None
else:
sequence_length = sum(sequence_lengths)
num_sequences = len(sequence_lengths)
cu_seqlens = torch.tensor(
[0, *torch.tensor(sequence_lengths).cumsum(0).tolist()],
dtype=torch.long,
device="cuda",
)

def randn(*shape: int, dtype: torch.dtype = torch.bfloat16, scale: float = 0.05):
return (torch.randn(*shape, dtype=torch.float32, device="cuda") * scale).to(dtype)

def l2norm(x: torch.Tensor) -> torch.Tensor:
return torch.nn.functional.normalize(x.float(), dim=-1).to(torch.bfloat16)

shape = (1, sequence_length, NUM_HEADS, HEAD_DIM)
inputs = dict(
q=l2norm(randn(*shape)),
k=l2norm(randn(*shape)),
v=randn(*shape),
g=randn(*shape),
beta=randn(1, sequence_length, NUM_HEADS, dtype=torch.float32),
scale=HEAD_DIM**-0.5,
state_indices=torch.arange(num_sequences, dtype=torch.int32, device="cuda"),
safe_gate=True,
lower_bound=-5.0,
use_gate_in_kernel=True,
A_log=torch.randn(NUM_HEADS, dtype=torch.float32, device="cuda") * 0.1,
dt_bias=torch.randn(NUM_HEADS * HEAD_DIM, dtype=torch.float32, device="cuda") * 0.1,
cu_seqlens=cu_seqlens,
use_beta_sigmoid_in_kernel=True,
)
initial_state = randn(
num_sequences, NUM_HEADS, HEAD_DIM, HEAD_DIM, dtype=torch.float32, scale=0.01
)

def run() -> tuple[torch.Tensor, torch.Tensor]:
state = initial_state.clone()
output = torch.ops.trtllm.kda_prefill(state_pool=state, **inputs).clone()
return output, state

expected_output, expected_state = run()
monkeypatch.setattr(kimi_k3_ops, "_fused_make_host", late_module.make_host_function)
monkeypatch.setattr(kimi_k3_ops, "_fused_k123_cache", {})
output, state = run()

torch.testing.assert_close(output, expected_output, rtol=0, atol=0)
torch.testing.assert_close(state, expected_state, rtol=0, atol=0)


@torch.no_grad()
def test_kda_mixer_empty_prefill():
"""An empty token payload still publishes the conv state for its slot."""
Expand Down
Loading