#!/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)"