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