From 648d6f8c169ae38df66ef824a970c3ff45f7e90d Mon Sep 17 00:00:00 2001 From: Mohamed Ghoneim <138689320+moghon92@users.noreply.github.com> Date: Tue, 31 Mar 2026 18:17:55 +0200 Subject: [PATCH] Fix SM120 forward pass crash Fix SM120 forward pass crash: parent __init__ overwrites arch, enabling unsupported TMA path FlashAttentionForwardSm120 sets class variable arch=80 to force CpAsync code paths (no TMA for output). However, FlashAttentionForwardSm80.__init__() calls self.arch = BaseDSL._get_dsl().get_arch_enum(), which returns the real GPU architecture (Arch.sm_120), overwriting the class variable with an instance variable. This causes use_tma_O = (self.arch >= Arch.sm_90) to evaluate True, and the epilogue enters the TMA output path where tma_atom_O is None (never created for SM120), resulting in: AttributeError: 'NoneType' object has no attribute '_trait' in copy_utils.tma_get_copy_fn -> cpasync.tma_partition Fix: override __init__ to reset self.arch = Arch.sm_80 after super().__init__(). Tested on NVIDIA B200 (SM 12.0) with: - torch 2.9.1+cu129 - nvidia-cutlass-dsl 4.4.2 - quack-kernels 0.3.7 --- flash_attn/cute/flash_fwd_sm120.py | 10 ++++++++++ 1 file changed, 10 insertions(+) diff --git a/flash_attn/cute/flash_fwd_sm120.py b/flash_attn/cute/flash_fwd_sm120.py index 08d219acfa8..c5cf5577a74 100644 --- a/flash_attn/cute/flash_fwd_sm120.py +++ b/flash_attn/cute/flash_fwd_sm120.py @@ -7,6 +7,7 @@ import cutlass import cutlass.utils as utils_basic +from cutlass.base_dsl.arch import Arch from flash_attn.cute.flash_fwd import FlashAttentionForwardSm80 @@ -16,6 +17,15 @@ class FlashAttentionForwardSm120(FlashAttentionForwardSm80): # The compilation target is determined by the GPU at compile time, not this field. arch = 80 + def __init__(self, *args, **kwargs): + super().__init__(*args, **kwargs) + # FlashAttentionForwardSm80.__init__ sets self.arch from + # BaseDSL.get_arch_enum() which returns the real GPU arch (sm_120), + # overwriting the class-level arch = 80. This enables TMA code paths + # (use_tma_O) that SM120 doesn't support, since tma_atom_O is never + # created for this subclass. Reset to sm_80 to stay on CpAsync paths. + self.arch = Arch.sm_80 + @staticmethod def can_implement( dtype,