CivArchive
    Test your Python setup : Torch-Compile, Flash Attention and Sage Attention - v1.0
    Preview 144449484

    (You're on Windows and) you really struggled to install Torch, a Clang compiler,Visual Studio C++ Compiler, FlashAttention, and SageAttention.
    Yet your brand new ComfyUI setup keeps crashing miserably and Forge shows you're running on your CPU...

    This little script verifies that Torch compiles, FlashAttention is flashing ans SageAttention is behaving hitself.

    Note: Each component should be able to work even if the others do not, it depends on what you want.

    Well... those who work on my machines and probably not on yours (?)

    SageAttn (Python 3.12/3.13)

    source : https://github.com/woct0rdho/SageAttention

    uv pip install triton-windows>3.7
    uv pip install https://github.com/woct0rdho/SageAttention/releases/download/v2.2.0-windows.post6/sageattention-2.2.0+cu130torch2.10.0andhigher.post6-cp310-abi3-win_amd64.whl

    FlashAttention

    source : https://github.com/mjun0812/flash-attention-prebuild-wheels

    Python 3.13

    uv pip install https://github.com/mjun0812/flash-attention-prebuild-wheels/releases/download/v0.9.52/flash_attn-2.8.3+cu130torch2.13-cp313-cp313-win_amd64.whl

    Python 3.12

    uv pip install https://github.com/mjun0812/flash-attention-prebuild-wheels/releases/download/v0.9.52/flash_attn-2.8.3+cu130torch2.13-cp312-cp312-win_amd64.whl

    This script Python code

    import torch
    
    # ANSI color codes (pure Python, no external dependency)
    GREEN = "\033[92m"
    RED = "\033[91m"
    YELLOW = "\033[93m"
    CYAN = "\033[96m"
    BOLD = "\033[1m"
    RESET = "\033[0m"
    
    PASS = f"{GREEN}{BOLD}[PASS]{RESET}"
    FAIL = f"{RED}{BOLD}[FAIL]{RESET}"
    INFO = f"{CYAN}[INFO]{RESET}"
    
    
    def print_header(title: str) -> None:
        """Print a standardized test section header."""
        print(f"\n{BOLD}==== START TEST {title} ===={RESET}")
    
    
    def print_footer(title: str) -> None:
        """Print a standardized test section footer."""
        print(f"{BOLD}===== END TEST {title} ====={RESET}")
    
    
    if __name__ == "__main__":
    
        # ---------------------------------------------------------------
        # Test 1: torch.compile
        # ---------------------------------------------------------------
        print_header("TORCH COMPILE")
        try:
            device = "cpu"  # or "xpu" for XPU
            print(f"{INFO} Device: {device}")
    
            def foo(x, y):
                a = torch.sin(x)
                b = torch.cos(x)
                return a + b
    
            opt_foo1 = torch.compile(foo)
            result = opt_foo1(
                torch.randn(10, 10).to(device),
                torch.randn(10, 10).to(device),
            )
            print(f"{INFO} Output shape: {tuple(result.shape)}")
            print(f"{PASS} torch.compile works.")
        except Exception as e:
            print(f"{FAIL} torch.compile error: {e}")
    
        # Test with fullgraph=True to detect graph breaks
        try:
            opt_foo2 = torch.compile(foo, fullgraph=True)
            result = opt_foo2(
                torch.randn(10, 10).to(device),
                torch.randn(10, 10).to(device),
            )
            print(f"{PASS} torch.compile with fullgraph=True OK.")
        except Exception as e:
            print(f"{FAIL} torch.compile fullgraph error: {e}")
    
    
        print_footer("TORCH COMPILE")
    
        # ---------------------------------------------------------------
        # Test 2: Flash Attention
        # ---------------------------------------------------------------
        print_header("FLASH ATTENTION")
        try:
            import flash_attn
            from flash_attn import flash_attn_func
    
            print(f"{INFO} Flash Attention version : {flash_attn.__version__}")
    
            if not torch.cuda.is_available():
                print(f"{YELLOW}[SKIP]{RESET} CUDA not available, Flash Attention test skipped.")
            else:
                # Test réel avec des tenseurs sur GPU
                q = torch.randn(2, 128, 8, 64, device="cuda", dtype=torch.float16)
                k = torch.randn(2, 128, 8, 64, device="cuda", dtype=torch.float16)
                v = torch.randn(2, 128, 8, 64, device="cuda", dtype=torch.float16)
    
                output = flash_attn_func(q, k, v)
                print(f"{INFO} Output shape: {tuple(output.shape)}")
                print(f"{PASS} Flash Attention forward pass OK.")
        except Exception as e:
            print(f"{FAIL} Flash Attention error: {e}")
    
        print_footer("FLASH ATTENTION")
    
        # ---------------------------------------------------------------
        # Test 3: Sage Attention
        # ---------------------------------------------------------------
        print_header("SAGE ATTENTION")
        try:
            from sageattention import sageattn
    
            # Test tensors
            batch_size = 2
            num_heads = 8
            seq_len = 1024
            head_dim = 64
    
            # SageAttention requires CUDA and fp16 inputs
            if not torch.cuda.is_available():
                print(f"{YELLOW}[SKIP]{RESET} CUDA not available, SageAttention test skipped.")
            else:
                q = torch.randn(batch_size, num_heads, seq_len, head_dim, device="cuda", dtype=torch.float16)
                k = torch.randn(batch_size, num_heads, seq_len, head_dim, device="cuda", dtype=torch.float16)
                v = torch.randn(batch_size, num_heads, seq_len, head_dim, device="cuda", dtype=torch.float16)
    
                output = sageattn(q, k, v)
                print(f"{INFO} Output shape: {tuple(output.shape)}")
    
                import sageattention.core as sage_core
    
                # Chech SM (Streaming Multiprocessor) capacities
                # Check if kernels INT8 and INT4 are available
                print(f"{INFO} SM80 (Ampere) enabled : {sage_core.SM80_ENABLED}")
                print(f"{INFO} SM89 (Ada) enabled    : {sage_core.SM89_ENABLED}")
                print(f"{INFO} SM90 (Hopper) enabled : {sage_core.SM90_ENABLED}")
    
                # Test with tensor_layout='HND' (default) and 'NHD'
                output_hnd = sageattn(q, k, v, tensor_layout='HND', is_causal=False)
                output_nhd = sageattn(q, k, v, tensor_layout='NHD', is_causal=False)
                print(f"{PASS} SageAttention HND layout OK, shape: {tuple(output_hnd.shape)}")
                print(f"{PASS} SageAttention NHD layout OK, shape: {tuple(output_nhd.shape)}")
    
    
                print(f"{PASS} SageAttention test passed.")
        except Exception as e:
            print(f"{FAIL} SageAttention error: {e}")
    
    
    
        print_footer("SAGE ATTENTION")

    Description

    v1

    Comments (1)

    EnragedAntelopeOct 2, 2026
    CivitAI

    Helpful - thanks for making and sharing.

    Other
    Other

    Details

    Downloads
    21
    Platform
    CivitAI
    Platform Status
    Available
    Created
    10/2/2026
    Updated
    10/11/2026
    Deleted
    -

    Files

    testYourPythonSetupTorch_v10.zip

    Mirrors