fix(inference): remove dynamic device scalar from multi-control attention - #271
sunil-srinivasa wants to merge 2 commits into
Conversation
…tion Signed-off-by: Sunil Srinivasa <106262814+sunil-srinivasa@users.noreply.github.com>
|
Is this change also updated in internal I4 repo? If so, it will be synced to this OSS repo. Thanks! |
|
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: So |
Summary
torch.compileon Blackwell before a GPU kernel launchesRoot cause
multi_control_two_way_attentionused: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: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_offsetsconstructs offsets and indices from the samesplit_lens, sooffsets[-1] == len(indices) == num_<mode>_tokens.Validation
I4_ATTN_BACKENDS=cudnntorch.compile(fullgraph=True, dynamic=True)[2049, 2048]python -m py_compileon both changed public filesThe public files were generated through the
cosmos-framework-releasemapping/rewrite stages from internal source MR !13565; they exactly match the release output for those mapped paths.