Conversation
…next chunk fused_kernel123's K1 warps read their scanned prefix row and the chunk total from sPartialLast after the scan barrier, and the next chunk's partial sums overwrite sPartialLast before any K1 barrier (racecheck: write at fused_k123.py:1268 against reads at :1293 / :1299). A K1 warp that reached its reads late used the next chunk's sums: with one K1 warp stalled 20 us after the scan barrier, every case of test_fused_prefill_beta_sigmoid_matches_unfused_kernel changed (output max |d| up to 7e14). One K1 barrier after the reads orders them; outputs are bit-identical to before. test_fused_prefill_does_not_depend_on_k1_warp_timing runs a copy of the kernel with that stall and requires the unpatched kernel's bits; it fails without the barrier. Signed-off-by: Vasanth Sabavat <[email protected]>
…ensor on every call) The ssm/kda_prefill entry's test pins the op's output as a new tensor on every call, which NVIDIA#19835 makes true. NVIDIA#19826 (this branch's base) and NVIDIA#19835 touch different code in test_kda_prefill_op.py and otherwise disjoint files. Signed-off-by: Vasanth Sabavat <[email protected]>
|
Navigate logical layers of code changes, visualize relationships, and explore their blast radius. 🧰 Additional context used📚 Code guidelines (1)No actionable comments were generated in the recent review. 🎉 ℹ️ Recent review info⚙️ Run configuration
📒 Files selected for processing (2)
Included review availability: This review used your included allowance. Your plan provides up to 12 included reviews per hour; 8 remain after this review. WalkthroughThe K1 kernel now synchronizes after reading prefix and chunk-total values, before Pass 2a. A parameterized test delays a K1 warp and checks output and state parity for padded and variable-length prefill inputs. ChangesK1 synchronization
Priority: ➖ Normal Estimated code review effort: 3 (Moderate) | ~20 minutes Change: Bug fix Suggested reviewers: Merge Risk: ⚪ Minimal · up to The barrier and timing test have no identified issue that needs resolution before merge. 🚥 Pre-merge checks | ✅ 4 | ❌ 1❌ Failed checks (1 warning)
✅ Passed checks (4 passed)
✨ Finishing Touches🧪 Generate unit tests (beta)
Comment |
Description
What goes wrong.
fused_kernel123(
tensorrt_llm/_torch/cute_dsl_kernels/blackwell/kimi_k3_kda/fused_k123.py) has a race between chunks:sPartialLastafter the scan barrier;sPartialLastbefore any K1 barrier.A K1 warp that reaches its reads late uses the next chunk's sums.
test_fused_prefill_beta_sigmoid_matches_unfused_kernelchanges, by up to 7e14 in the output.test_fused_prefill_beta_sigmoid_matches_unfused_kernel[padded_eqlen]gives a NaN output.Fix. One K1 barrier after the reads orders them before the next chunk's writes. Outputs are bit-identical to
before.
Test Coverage
tests/unittest/_torch/modules/kimi_kda/test_kda_prefill_op.py::test_fused_prefill_does_not_depend_on_k1_warp_timing(SM 100):
The file is already in
l0_b200.yml.Results on GB200:
padded_eqlen,varlen).test_kda_prefill_op.pypasses 26 / 26.modules/kimi_kda/has no new failures: 119 passed and 1 skipped with this PR, 116 passed and 2 skipped onmain. The new test adds 2 passes, and one test that skipped in main's run ran and passed in this PR's run.
PR Checklist
Please review the following before submitting your PR:
PR description clearly explains what and why. If using CodeRabbit's summary, please make sure it makes sense.
PR Follows TRT-LLM CODING GUIDELINES to the best of your knowledge.
Test cases are provided for new code paths (see test instructions)
If PR introduces API changes, an appropriate PR label is added - either
api-compatibleorapi-breaking. Forapi-breaking, includeBREAKINGin the PR title. (No API change.)Any new dependencies have been scanned for license and vulnerabilities (No new dependencies.)
CODEOWNERS updated if ownership changes (No ownership change.)
Documentation updated as needed (No user-facing change.)
Update tava architecture diagram if there is a significant design change in PR. (No design change.)
The reviewers assigned automatically/manually are appropriate for the PR. (To check once the PR is open.)
Please check this after reviewing the above items as appropriate for this PR.
GitHub Bot Help
To see a list of available CI bot commands, please comment
/bot help.Dev Engineer Review
fused_kernel123adds a K1 barrier after each warp reads its prefix and chunk total fromsPartialLast. This prevents the next chunk from overwriting those values before all K1 warps finish reading them.QA Engineer Review
tests/unittest/_torch/modules/kimi_kda/test_kda_prefill_op.pyadds the parameterizedtest_fused_prefill_does_not_depend_on_k1_warp_timingtest for padded equal-length and variable-length inputs. It injects a delay into a K1 warp and checks bit-exact output and state parity against the unmodified kernel.tests/integration/test_lists/test-db/l0_b200.ymlandtests/integration/test_lists/test-db/l0_gb300_multi_gpus.yml. No changed test-list files were reported.test_kda_prefill_op.pyon GB200. These results were not independently verified here.Per-File QA Perspective
tensorrt_llm/_torch/cute_dsl_kernels/blackwell/kimi_k3_kda/fused_k123.py: Verify K1 synchronization prevents reads of overwrittensPartialLastvalues and preserves prefill results. The barrier also adds synchronization overhead.tests/unittest/_torch/modules/kimi_kda/test_kda_prefill_op.py: The timing-stall test checks bit-exact output and state parity for padded equal-length and variable-length inputs. The test file appears in the GB200 CI test lists noted above.