| # Portable, arch-aware env bootstrap. Run on ANY rented GPU (A100/H100/Blackwell). | |
| # Standard image for all arches: pytorch/pytorch:2.12.1-cuda13.0-cudnn9-devel | |
| # (host drivers are recent enough that cuda13 torch covers sm80/sm90/sm100/sm120). | |
| set -uo pipefail | |
| export PIP_BREAK_SYSTEM_PACKAGES=1 PIP_ROOT_USER_ACTION=ignore | |
| export HF_HOME=/workspace/hf-cache HF_HUB_CACHE=/workspace/hf-cache/hub | |
| mkdir -p /workspace/hf-cache /workspace/tmp | |
| echo "== core deps ==" | |
| python -m pip install -q -U "transformers>=5.12.1" accelerate datasets peft huggingface_hub \ | |
| safetensors tokenizers sentencepiece protobuf einops pyyaml | |
| python -m pip install -q -U causal-conv1d "flash-linear-attention>=0.5.1" || echo "WARN: fla/causal-conv1d (torch fallback ok)" | |
| CAP=$(python -c "import torch;print('%d%d'%torch.cuda.get_device_capability(0))" 2>/dev/null || echo "?") | |
| NAME=$(python -c "import torch;print(torch.cuda.get_device_name(0))" 2>/dev/null || echo "?") | |
| echo "== GPU: $NAME (sm_$CAP) ==" | |
| case "$CAP" in | |
| 80|86|89) echo "Ampere/Ada -> FLA gated-delta backward works natively. Use seq up to ~4k on 80GB.";; | |
| 90) echo "Hopper -> FLA backward needs tilelang (or triton<3.4)."; python -m pip install -q tilelang apache-tvm-ffi || echo " tilelang install/import may fail (see task #19).";; | |
| 100|120) echo "Blackwell -> FLA backward UNVERIFIED. Run run_phase0.sh smoke BEFORE a long run.";; | |
| *) echo "Unknown cap '$CAP' — verify with run_phase0.sh.";; | |
| esac | |
| python -c "import torch;print('torch',torch.__version__,'cuda',torch.version.cuda,'cap',torch.cuda.get_device_capability(0))" | |
| echo "bootstrap done." | |