Fix CUTLASS FMHA BiasLoader alignment for unaligned kernel path - #28369
Conversation
BiasLoader hardcoded 128-bit (8 fp16 element) vectorized loads via ElementsPerAccess = 128 / sizeof_bits<scalar_t> regardless of the isAligned template parameter. When attention bias stride (total_sequence_length) was not a multiple of 8, the unaligned kernel was selected but BiasLoader still used 128-bit loads on the bias, causing cudaErrorMisalignedAddress. Fix: Use kAlignmentA (kMinimumAlignment=4 for unaligned path, kAlignmentA=8 for aligned path) as BiasLoader's ElementsPerAccess. This allows the unaligned kernel to use 64-bit loads for the bias. Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> Signed-off-by: Justin Chu <justinchu@microsoft.com>
Test the BiasLoader alignment fix with total_kv_length values that are not divisible by 8 (the fp16 vectorized load width). Before the fix, these would cause wrong results or crashes in the MEA kernel path. 16 test cases across 3 categories: - MHA decode: 8 lengths (5, 7, 9, 13, 27 unaligned + 8, 16, 32 aligned) - MHA prompt: 5 lengths (5, 7, 13 unaligned + 8, 16 aligned) - GQA decode: 3 lengths (5, 9 unaligned + 16 aligned) All use float16 with 4D additive attention mask on CUDA EP, comparing against PyTorch reference. Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> Signed-off-by: Justin Chu <justinchu@microsoft.com>
Tianlei Wu (tianleiwu)
left a comment
There was a problem hiding this comment.
Summary
Clean, minimal fix for a real crash caused by the CUTLASS BiasLoader hardcoding 128-bit vectorized loads regardless of the kernel's alignment path. The one-line kernel change is correct and consistent with how kAlignmentA is already used for Q/K iterators. The test coverage is thorough in terms of length combinations but has one high-priority gap: it does not disable Flash Attention, so on SM80+ GPUs the CUTLASS MEA path may never be exercised.
Kernel Fix
Positive: The fix is minimal and exactly right — replacing 128 / cutlass::sizeof_bits<scalar_t>::value (always 8 for fp16) with kAlignmentA, which already adapts via kIsAligned ? DefaultConfig::kAlignmentA : GemmType::kMinimumAlignment (8 aligned, 4 unaligned). This makes BiasLoader consistent with Q/K iterators and with the existing check_supported() validation and dispatch-side alignment check in fmha_launch_template.h.
Tests
Positive: Good test matrix covering decode, prompt, and GQA paths with both aligned and unaligned total KV lengths.
See inline comment for one concern.
…mha.py Co-authored-by: Tianlei Wu <tlwu@microsoft.com>
|
Copilot fix lint |
BiasLoader hardcoded 128-bit vectorized loads (
ElementsPerAccess = 128/sizeof_bits = 8for fp16) regardless of theisAlignedtemplate flag. When bias stride was not a multiple of 8, the unaligned kernel was selected but BiasLoader still used 128-bit loads →cudaErrorMisalignedAddress.Fix: Use
kAlignmentA(4 for unaligned, 8 for aligned) instead of hardcoded 8.Tested with Gemma4 Attention + mask at all seq lengths 1–32.