(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.
BONUS : Links for Windows attentions for Torch 2.13, CUDA 130
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.whlFlashAttention
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.whlPython 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.whlThis 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)
Helpful - thanks for making and sharing.
