[CuTe] Make is_fake_mode() torch.compile-safe (skip active_fake_mode under Dynamo) - #2678
johnnynunez wants to merge 1 commit into
Conversation
Note for reviewers testing on consumer Blackwell (sm120)This fix is self-contained and independent — it only touches the tracing-time guard in One heads-up if you reproduce on an sm120 GPU (RTX PRO 6000 / RTX 50-series): the end-to-end FA4 forward path on sm120 also needs the That crash is not caused by this PR (it reproduces in plain eager on sm120 too) — it's the pre-existing TMA-O path that #2655 addresses. With both #2655 and this PR applied, the full check passes on sm120 (torch 2.12, CUDA 13.2):
On Hopper (sm90) / sm100 the TMA-O path is unaffected, so this PR stands alone there. |
Follow-up: tracing fix is correct, but end-to-end FA4+compile needs more workTo be fully transparent after deeper testing on sm120: this PR makes However, enabling tracing is not the same as guaranteeing a correct compiled graph end-to-end. Plugging FA4 into a real model (NVIDIA Cosmos3 DiT) and compiling the whole transformer, I see shape-dependent NaNs:
So the composition of the CuTeDSL JIT kernel with the inductor graph appears to break on longer sequences — likely a graph-capture / tile-scheduler interaction, not this guard change. This PR stands on its own merit (the |
…under Dynamo) flash_attn_func is uncompilable under torch.compile: _flash_attn_fwd gates its real-tensor `.is_cuda` assertions behind is_fake_mode(), which calls torch._guards.active_fake_mode(). That symbol is in Dynamo's MOD_SKIPLIST, so tracing into it raises: torch._dynamo.exc.Unsupported: Attempted to call function marked as skipped (active_fake_mode in torch/_guards.py) Under compilation there are no real tensors to validate, so fake mode and compile-time tracing should be treated identically. Check torch.compiler.is_compiling() first (safe to trace) and short-circuit to True before touching active_fake_mode(). Verified on sm120 (RTX PRO 6000, torch 2.12): torch.compile(flash_attn_func, fullgraph=True) now traces; compiled output is bitwise-identical to eager (max_diff 0.0) and matches SDPA at the bf16 noise floor (4.9e-4).
f40c873 to
fd95611
Compare
What
flash_attn_func(CuTeDSL / FA4) is currently uncompilable undertorch.compile._flash_attn_fwdgates its real-tensor.is_cudaassertions behindis_fake_mode(), which callstorch._guards.active_fake_mode(). That symbol is in Dynamo'sMOD_SKIPLIST, so tracing into it raises:Fix
Under compilation there are no real tensors to validate, so compile-time tracing and fake mode should be treated identically.
is_fake_mode()now checkstorch.compiler.is_compiling()first (which is safe to trace) and short-circuits toTruebefore touchingactive_fake_mode().Repro
Verification
On sm120 (RTX PRO 6000 Blackwell, torch 2.12, CUDA 13.2):
torch.compile(flash_attn_func, fullgraph=True)now traces successfully.max_diff == 0.0).F.scaled_dot_product_attentionat the bf16 noise floor (4.9e-4).The change only affects the tracing-time guard path; eager numerics are untouched.