File size: 1,636 Bytes
994182c | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 | #!/usr/bin/env bash
# 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."
|