coder-fake / scripts /setup_fast_kernels.sh
halle01's picture
Add files using upload-large-folder tool
78cb49e verified
Raw History Blame Contribute Delete
2.94 kB
#!/bin/bash
# Enable the flash-linear-attention (fla) fused Gated-DeltaNet kernel β€” the
# single biggest training-throughput lever. VERIFIED on a Thunder A6000:
# fixed-4096 fwd+bwd dropped 32.8s -> 9.0s/step (3.6x) with this exact recipe.
#
# THE KEY FINDING (see NEXT_INSTANCE_TRAINING.md): fla's kernel "hangs" on the
# default stack's Triton 3.7.1 (ships with torch 2.13+cu130) β€” it's not a real
# hang, triton 3.7.1 DEADLOCKS compiling the Gated-DeltaNet kernel. Pinning
# **triton 3.6.0** fixes it (torch 2.13 warns about its version pin but runs
# fine). This was proven by vLLM running the same model's GDN kernels fine on
# triton 3.6.0. torch.compile also hangs on 3.7.1 for the same reason.
set -euo pipefail
cd "$(dirname "$0")/.."
source .venv/bin/activate
echo "=== current versions ==="
python -c "import torch,triton;print('torch',torch.__version__,'triton',triton.__version__)"
echo "=== pin triton 3.6.0 (THE fix) + fla 0.5.2 ==="
# 0.5.2, not 0.4.x: works with triton 3.6.0 for TRAINING (the 0.5.0 '!!!!'
# corruption bug is inference/generation-only, doesn't affect loss training).
pip install "triton==3.6.0" "flash-linear-attention==0.5.2"
python -c "import torch,triton,fla;print('now: torch',torch.__version__,'triton',triton.__version__,'fla',fla.__version__)"
echo "=== verify fla kernel WARMS (do NOT judge on the first call β€” it compiles ~2s) ==="
python - <<'PY'
import torch, time
from fla.ops.gated_delta_rule import chunk_gated_delta_rule as chunk
B,T,H,D=1,4096,4,128
q=torch.randn(B,T,H,D,device="cuda",dtype=torch.bfloat16,requires_grad=True)
k=torch.randn(B,T,H,D,device="cuda",dtype=torch.bfloat16); v=torch.randn(B,T,H,D,device="cuda",dtype=torch.bfloat16)
g=torch.rand(B,T,H,device="cuda").log(); beta=torch.rand(B,T,H,device="cuda",dtype=torch.bfloat16)
def step():
o,_=chunk(q,k,v,g,beta,use_qk_l2norm_in_kernel=True); o.sum().backward(); q.grad=None
t0=time.time(); step(); torch.cuda.synchronize(); print(f"cold fwd+bwd (compile): {time.time()-t0:.1f}s")
t0=time.time()
for _ in range(5): step()
torch.cuda.synchronize()
ms=1000*(time.time()-t0)/5
print(f"WARM fwd+bwd: {ms:.1f}ms/call")
if ms > 500:
raise SystemExit("!! fla kernel is slow even warm β€” check triton version (need 3.6.0)")
print("OK β€” fla is active and fast. Re-run the smoke and compare s/step to the torch-fallback baseline.")
PY
echo
echo "=== OPTIONAL: causal-conv1d for the COMPLETE fast path (needs nvcc) ==="
# is_fast_path_available also needs causal-conv1d. Without it the conv falls
# back to torch (why we get 3.6x not ~10x). causal-conv1d is source-only and
# needs nvcc (absent by default on Thunder). To try: install a pip nvcc first
# pip install nvidia-cuda-nvcc-cu12 # then ensure it's on PATH
# pip install causal-conv1d --no-build-isolation
# Skip if it fails β€” fla alone is the bulk of the win.
echo "(skipped by default β€” see NEXT_INSTANCE_TRAINING.md if you want to attempt it)"