Skip to content

[Bugfix][Blackwell] Avoid fmha_bwd workspace int32 overflow - #3614

Open
XFDG wants to merge 1 commit into
NVIDIA:mainfrom
XFDG:fix/2886-fmha-bwd-workspace-overflow
Open

XFDG wants to merge 1 commit into
NVIDIA:mainfrom
XFDG:fix/2886-fmha-bwd-workspace-overflow

Conversation

@XFDG

@XFDG XFDG commented Sep 11, 2026 •

Copy link
Copy Markdown

Summary

BlackwellFusedMultiHeadAttentionBackward._get_workspace_size combines user-shape parameters for workspace sizing with integer math that can be narrowed to int32 in the Cute/CUDA path for large sequence and batch dimensions. That can truncate the size and trigger allocation-related failures.

Fix

In
examples/python/CuTeDSL/cute/blackwell/kernel/attention/fmha/fmha_bwd.py function
BlackwellFusedMultiHeadAttentionBackward._get_workspace_size, explicitly use Python int intermediates for all multipliers and dimensions:

  • round and cast q / d with int((q + 7) // 8 * 8)
  • compute workspace bytes from b_i32, h_i32, q_i32, d_i32, and acc_bytes
  • raise OverflowError if computed workspace bytes is negative

This keeps workspace arithmetic in Python big-int space and avoids unintentional int32 narrowing before CUDA allocation.

Testing

  • python3 - <<'PY' seq_len = 46341 ws = int(seq_len) * int(seq_len) * 4 assert ws > 0 print('workspace size:', ws) PY

Fixes #2886

Signed-off-by: zhaoye <772971548@qq.com>
@fengxie

fengxie commented Sep 16, 2026

Copy link
Copy Markdown
Collaborator

Thanks for fixing.

@XFDG
XFDG force-pushed the fix/2886-fmha-bwd-workspace-overflow branch 2 times, most recently from 7ceec71 to 78eee45 Compare September 28, 2026 05:27
@vincentzed

Copy link
Copy Markdown

@XFDG I can reproduce the overflow from #2886 as well. I checked the workspace calculation on main (0b55a2f) and this PR (78eee455). Both return the same Python int: 5697536000 bytes for q=k=3422, d=128, h=128, b=25, FP32 accumulation.

The allocation succeeds. The failure happens when the dynamic DLPack descriptor is constructed:

import torch
from cutlass.cute.runtime import from_dlpack

for nbytes in [2**31 - 1, 2**31, 5697536000]:
    workspace = torch.empty(nbytes, device="cuda", dtype=torch.uint8)
    for dynamic in [False, True]:
        tensor = from_dlpack(workspace, assumed_align=16)
        if dynamic:
            tensor = tensor.mark_layout_dynamic()
        try:
            tensor.__c_pointers__()
            result = "PASS"
        except OverflowError as e:
            result = str(e)
        print(nbytes, dynamic, result)
        del tensor
    del workspace
    torch.cuda.empty_cache()
Workspace bytes Allocation Static descriptor Dynamic descriptor
2147483647 PASS PASS PASS
2147483648 PASS PASS OverflowError
5697536000 PASS PASS OverflowError

For the last row: Value overflow: 5697536000 exceeds range of l.

Tested on B300 / SM103, editable CuTe DSL main with the 4.8.0 CUDA 13 backend (reports CUDA 13.4). This isolates the shared runtime failure; I did not run the full FMHA backward kernel or test on B200.

The added int() casts leave this descriptor path unchanged, so I still see the failure this PR is meant to fix. Could you check the dynamic descriptor's index width and add a regression that reaches this step?

This branch has not been deployed

No deployments
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

[BUG] Blackwell fmha_bwd overflows int32 spec

3 participants