attention / README.md
Dunfan's picture
Publish kernel (build/ + card + benchmarks); drop source/build-tooling files
2c1bc68 verified
|
Raw
History Blame
2.28 kB
metadata
library_name: kernels
license: apache-2.0
tags:
  - kernels
  - helion
  - attention

attention (Helion)

A scaled-dot-product / flash-attention kernel written in Helion and packaged for the kernels library. It is a Triton-backend kernel that ships pre-tuned configs (via helion.aot_kernel), so the first call selects a known-good configuration and compiles it — no autotuning search at run time.

How to use

# make sure `kernels` is installed: `pip install -U kernels`
import torch
from kernels import get_kernel

kernel = get_kernel("HelionDSL/attention", version=1)

# flash-attn layout: (batch, seq, heads, head_dim), fp16/bf16
q = torch.randn(2, 512, 8, 64, device="cuda", dtype=torch.bfloat16)
k = torch.randn_like(q)
v = torch.randn_like(q)

out = kernel.flash_attn_func(q, k, v, causal=False)

get_kernel requires a pinned version or revision — use version=1 (the stable kernel-API major, backed by the repo's v1 branch) or revision="main" / revision="<commit>" for an explicit Hub revision.

HelionDSL is a trusted kernel publisher on the Hub, so trust_remote_code=True is not required. (On older kernels releases that predate publisher trust, pass trust_remote_code=True to get_kernel.)

Available functions

  • flash_attn_func(q, k, v, softmax_scale=None, causal=False) — flash-attn compatible entry point. Inputs and output are (batch, seq, heads, head_dim).
  • attention(q, k, v, causal=False, out=None) — SDPA-layout entry point. Inputs and output are (batch, heads, seq, head_dim).

Both accept torch.float16 / torch.bfloat16 and support causal=True.

Available layers

  • Attention — an nn.Module wrapper (layers.Attention(causal=False)) whose forward(q, k, v) calls the kernel.

Notes

  • Backend: Helion → Triton. Correctness is verified against torch.nn.functional.scaled_dot_product_attention.
  • Pre-tuned: ships _helion_aot_attention_cuda_sm100.py and _helion_aot_attention_cuda_sm90.py; the right one is selected for your GPU (with compute-capability fallback). On a shape/GPU without a shipped config, the kernel falls back to normal Helion autotuning.