Skip to content

fix(inference): remove dynamic device scalar from multi-control attention - #271

Open
sunil-srinivasa wants to merge 2 commits into
NVIDIA:mainfrom
sunil-srinivasa:fix/multi-control-cudnn-symint
Open

sunil-srinivasa wants to merge 2 commits into
NVIDIA:mainfrom
sunil-srinivasa:fix/multi-control-cudnn-symint

Conversation

@sunil-srinivasa

Copy link
Copy Markdown
Contributor

Summary

  • fix multi-control transfer inference failing during torch.compile on Blackwell before a GPU kernel launches
  • read the real text/GEN token counts from the sequence pack's host-side metadata instead of converting CUDA offset scalars to Python integers
  • preserve the cuDNN path and its existing GB200 performance advantage over the varlen NATTEN fallback
  • add a CPU dynamic full-graph compile regression test that fails before the fix

Root cause

multi_control_two_way_attention used:

n_text = int(causal_k_offsets[-1])
n_full = int(full_q_offsets[-1])

Under torch.compile(dynamic=True), those CUDA scalar reads become unbacked symbolic integers. The per-control KV length then reaches cuDNN's Inductor lowering as a symbolic stride expression such as:

128 * Max(1, u2 + 3720)

The lowering cannot resolve its stride comparison against the 128-element head dimension. Head dimension 128 is supported; the failure is the data-dependent scalar read.

This presents on Blackwell because cuDNN leads the attention backend list for sm_100/103, while Hopper normally selects FlashAttention 3 first.

The pack already stores the same token counts as host-side integers. _compute_mode_indices_and_offsets constructs offsets and indices from the same split_lens, so offsets[-1] == len(indices) == num_<mode>_tokens.

Validation

  • GB200 hardware: PASS
    • NVIDIA GB200, sm_100, driver 580.126.20
    • PyTorch 2.12.0+cu130
    • forced I4_ATTN_BACKENDS=cudnn
    • torch.compile(fullgraph=True, dynamic=True)
    • compile + first execution: 16.715s
    • cached execution: 1ms
    • finite output, shape [2049, 2048]
  • Regression test: PASS
    • fails on the parent source commit with an unbacked-scalar runtime assertion
    • passes after this change
  • python -m py_compile on both changed public files

The public files were generated through the cosmos-framework-release mapping/rewrite stages from internal source MR !13565; they exactly match the release output for those mapped paths.

sunil-srinivasa and others added 2 commits September 22, 2026 11:01
…tion

Signed-off-by: Sunil Srinivasa <106262814+sunil-srinivasa@users.noreply.github.com>
@lfengad

lfengad commented Sep 23, 2026

Copy link
Copy Markdown
Collaborator

Is this change also updated in internal I4 repo? If so, it will be synced to this OSS repo. Thanks!

@chrisvoncsefalvay

Copy link
Copy Markdown
Contributor

This change also fixes an eager-mode multi-control transfer failure with --no-use-torch-compile -- we encountered it running Cosmos3-Nano with simultaneous depth and segmentation controls on an RTX PRO 6000 Blackwell, using PyTorch 2.10.0+cu130.

In multi_control_two_way_attention, the final offset includes the trailing padding segment:

full_q_offsets = [0, 85560, 85561]
causal_q_offsets = [0, 1422, 1423]
noisy_token_range = (57040, 85560)

So n_full becomes 85561 and torch._check(n_full == noisy_e) fails at attention.py:661. The causal token count also includes a padding row. Thanks for this PR -- we tested it with 1, 3 and 16 padding rows and passes. Could we include these padded eager-mode cases explicitly? This would document that the change fixes ordinary multi-control inference as well as the compilation failure.

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.

4 participants