coder-fake / scripts /setup_blackwell.sh
halle01's picture
Add files using upload-large-folder tool
78cb49e verified
Raw History Blame Contribute Delete
5.24 kB
#!/usr/bin/env bash
# ============================================================================
# Blackwell (RTX PRO 6000, 96GB, sm_120) setup for the threejs coder fine-tune.
# Adapts our A6000-proven stack to Blackwell. Reference stacks that informed
# this: affine/setup_blackwell.sh and affine/install_finetune_blackwell.sh
# (both use CUDA 12.8 / cu128 β€” nightly torch for sm_120).
#
# THE ONE THING THAT MATTERS: does fla's Gated-DeltaNet Triton kernel COMPILE
# on sm_120? Our fix on the A6000 was triton 3.6.0, but Blackwell is NEWER than
# triton 3.6.0 β€” its sm_120 codegen may be absent. This script installs, then
# RUNS THE 30s fla ISOLATION TEST AS A HARD GATE before you trust anything.
# Do not launch a real run until the gate prints FLA_OK.
#
# Run: bash scripts/setup_blackwell.sh
# ============================================================================
set -euo pipefail
cd "$(dirname "$0")/.."
echo "=== 0. GPU check (expect RTX PRO 6000 / sm_120 / ~96GB) ==="
nvidia-smi --query-gpu=name,memory.total,compute_cap --format=csv,noheader || nvidia-smi --query-gpu=name,memory.total --format=csv,noheader
echo "=== 1. venv ==="
python3 -m venv .venv || python3.12 -m venv .venv
source .venv/bin/activate
pip install --upgrade pip wheel
# --- 2. PyTorch for Blackwell -----------------------------------------------
# Two candidate stacks, try in order. Our A6000 stack was torch 2.13.0+cu130;
# CUDA 13.0 >= 12.8 so it SHOULD cover sm_120 IF the wheel ships sm_120 kernels.
# If torch can't see the GPU (sm_120 not in the wheel), fall back to cu128
# nightly (what affine's script uses for Blackwell).
echo "=== 2. install torch (try cu130 stable, else cu128 nightly) ==="
pip install "torch==2.13.0" torchvision --index-url https://download.pytorch.org/whl/cu130 || \
pip install --pre torch torchvision --index-url https://download.pytorch.org/whl/nightly/cu128
python - <<'PY'
import torch
print("torch", torch.__version__, "| cuda", torch.version.cuda, "| avail", torch.cuda.is_available())
if torch.cuda.is_available():
p=torch.cuda.get_device_properties(0)
print("GPU", p.name, "| sm", f"{p.major}.{p.minor}", "| VRAM", round(p.total_memory/1e9), "GB")
# sanity: a real matmul must run (proves the wheel has sm_120 kernels)
a=torch.randn(1024,1024,device="cuda"); (a@a).sum().item(); print("cuda matmul OK")
else:
raise SystemExit("torch cannot see the Blackwell GPU β€” the wheel lacks sm_120. Re-run and let it fall through to cu128 nightly.")
PY
echo "=== 3. training deps (transformers 5.x REQUIRED for qwen3_5) ==="
# NOTE: affine's ref pins transformers<5.0, but OUR model_type qwen3_5 only
# exists in transformers 5.x β€” do NOT downgrade transformers here.
pip install "transformers>=5.0.0" "trl>=1.9.0" "peft>=0.13.0" "accelerate>=1.0.0" \
"bitsandbytes>=0.48.1" "datasets>=3.0.0" pillow httpx pyyaml wandb
echo "=== 4. fla + triton (the fast Gated-DeltaNet kernel) ==="
# Start with the A6000-proven pin. If the gate below fails on sm_120, the doc
# (NEXT_INSTANCE_TRAINING.md, Blackwell section) lists the escalation: try the
# triton that ships WITH this torch (uninstall the 3.6.0 pin), or fla from git.
pip install "triton==3.6.0" "flash-linear-attention==0.5.2" || \
pip install "flash-linear-attention==0.5.2" # keep torch's bundled triton if 3.6.0 refuses
python -c "import torch,triton,fla;print('torch',torch.__version__,'| triton',triton.__version__,'| fla',fla.__version__)"
echo "=== 5. HARD GATE: does the fla chunk kernel compile+run on sm_120? (30s) ==="
python - <<'PY'
import torch, time, signal, os
def to(s,f):
print("FLA_HANG: chunk kernel did not finish in 90s on sm_120 β€” triton likely lacks sm_120 codegen.", flush=True)
os._exit(2)
signal.signal(signal.SIGALRM, to); signal.alarm(90)
from fla.ops.gated_delta_rule import chunk_gated_delta_rule as chunk
B,T,H,D=1,4096,4,128; dev="cuda"; dt=torch.bfloat16
q=torch.randn(B,T,H,D,device=dev,dtype=dt,requires_grad=True)
k=torch.randn(B,T,H,D,device=dev,dtype=dt); v=torch.randn(B,T,H,D,device=dev,dtype=dt)
g=torch.rand(B,T,H,device=dev,dtype=torch.float32).log(); beta=torch.rand(B,T,H,device=dev,dtype=dt)
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 {time.time()-t0:.1f}s (compile)", flush=True)
t0=time.time()
for _ in range(5): step()
torch.cuda.synchronize(); signal.alarm(0)
ms=1000*(time.time()-t0)/5
print(f"warm fwd+bwd {ms:.1f}ms/layer-call", flush=True)
print("FLA_OK" if ms < 500 else "FLA_SLOW: ran but slow β€” investigate before a long run.", flush=True)
PY
echo ""
echo "=== Done. If you saw FLA_OK above, fla works on this Blackwell β€” proceed. ==="
echo "If FLA_HANG/FLA_SLOW: see NEXT_INSTANCE_TRAINING.md 'Blackwell' section for the"
echo "triton/torch escalation, or train on the torch fallback (Blackwell's silicon is"
echo "fast and 96GB removes the KTO memory limits either way)."
echo ""
echo "Blackwell config deltas (96GB): in configs/kto.yaml set max_seq_length: 4096"
echo "and per_device_train_batch_size: 4 (accum 8); SFT can use batch 2. No"
echo "expandable_segments juggling needed. Keep optim: adamw_torch."