Download scripts/setup_fast_kernels.sh from halle01/coder-fake: direct link, hf CLI and curl.
- Browser
- Download file 2.94 kB
-
https://huggingface.co/halle01/coder-fake/resolve/main/scripts/setup_fast_kernels.sh
- Command line
-
hf download hf://halle01/coder-fake/scripts/setup_fast_kernels.sh
-
curl -L -o setup_fast_kernels.sh https://huggingface.co/halle01/coder-fake/resolve/main/scripts/setup_fast_kernels.sh
2.94 kB
| # 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)" | |