Skip to content

[CuTe] Make is_fake_mode() torch.compile-safe (skip active_fake_mode under Dynamo) - #2678

Closed
johnnynunez wants to merge 1 commit into
Dao-AILab:mainfrom
johnnynunez:fix/cute-torch-compile-fake-mode
Closed

johnnynunez wants to merge 1 commit into
Dao-AILab:mainfrom
johnnynunez:fix/cute-torch-compile-fake-mode

Conversation

@johnnynunez

Copy link
Copy Markdown
Contributor

What

flash_attn_func (CuTeDSL / FA4) is currently 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
  Explanation: ... the function `active_fake_mode` in file `torch/_guards.py` should not be traced.

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 checks torch.compiler.is_compiling() first (which is safe to trace) and short-circuits to True before touching active_fake_mode().

def is_fake_mode() -> bool:
    try:
        if torch.compiler.is_compiling():
            return True
    except Exception:
        pass
    return active_fake_mode() is not None

Repro

import torch
from flash_attn.cute import flash_attn_func

q = torch.randn(1, 4096, 16, 128, device="cuda", dtype=torch.bfloat16)
k = torch.randn(1, 4096, 16, 128, device="cuda", dtype=torch.bfloat16)
v = torch.randn(1, 4096, 16, 128, device="cuda", dtype=torch.bfloat16)

def f(q, k, v):
    return flash_attn_func(q, k, v, causal=False)

f(q, k, v)                                   # eager: OK
torch.compile(f, fullgraph=True)(q, k, v)    # before: Unsupported; after: OK

Verification

On sm120 (RTX PRO 6000 Blackwell, torch 2.12, CUDA 13.2):

  • torch.compile(flash_attn_func, fullgraph=True) now traces successfully.
  • Compiled output is bitwise-identical to eager (max_diff == 0.0).
  • Both match F.scaled_dot_product_attention at the bf16 noise floor (4.9e-4).

The change only affects the tracing-time guard path; eager numerics are untouched.

@johnnynunez

Copy link
Copy Markdown
Contributor Author

Note for reviewers testing on consumer Blackwell (sm120)

This fix is self-contained and independent — it only touches the tracing-time guard in is_fake_mode().

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 use_tma_O fix from #2655. Without #2655, the kernel epilogue raises an unrelated error during the TMA output copy:

AttributeError: 'NoneType' object has no attribute '_trait'
  ... copy_utils.tma_get_copy_fn -> cpasync.tma_partition

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):

  • torch.compile(flash_attn_func, fullgraph=True) traces
  • compiled output is bitwise-identical to eager (max_diff == 0.0)
  • both match SDPA at the bf16 noise floor (4.9e-4)

On Hopper (sm90) / sm100 the TMA-O path is unaffected, so this PR stands alone there.

@johnnynunez

Copy link
Copy Markdown
Contributor Author

Follow-up: tracing fix is correct, but end-to-end FA4+compile needs more work

To be fully transparent after deeper testing on sm120: this PR makes flash_attn_func traceable under torch.compile, and the standalone check is solid — compiled output is bitwise-identical to eager (max_diff == 0.0) on the shapes in the repro.

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:

  • 256p (seqlen ~4k): compiled latent is finite, output looks sane.
  • 480p (seqlen ~16k): the compiled forward produces NaN (single-forward cosine=nan, mean|val|=nan), while the same FA4 kernel in eager is correct (cosine=0.99995 vs native).

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 active_fake_mode skip was a real, unconditional blocker to tracing), but I wanted to flag that FA4 + torch.compile is not yet numerically safe end-to-end so nobody ships it expecting correctness on large shapes. Eager FA4 is unaffected and correct. Happy to investigate the NaN path separately if useful.

…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).
@johnnynunez
johnnynunez force-pushed the fix/cute-torch-compile-fake-mode branch from f40c873 to fd95611 Compare June 28, 2026 18:08
@johnnynunez johnnynunez closed this Jul 7, 2026
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.

1 participant