Sign In

Test your Python setup : Torch-Compile, Flash Attention and Sage Attention

Updated: Oct 2, 2026

tooltorchpython

Download

1 variant available

Archive Other

test_torch_compile_and_attentions.zip

1.49 KB

Verified:

Type
Other
Stats

21

Reviews
Published

Oct 2, 2026

Base Model

Other

Hash
AutoV2
715FCFB0ED
Comfy-20261002_112045_a_pass1_00001__bmod.webp

(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")