File size: 9,278 Bytes
4a28d4d
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
#!/bin/bash
# Portable Slurm helper. Supply site-specific account/partition flags to
# sbatch, for example: sbatch -A <account> -p <partition> ...
#
# Usage:
#   sbatch scripts/batch/setup_env.batch
#
# Useful overrides:
#   sbatch --export=ALL,CONDA_ENV_NAME=acdir_math500,FORCE_RECREATE_ENV=1 scripts/batch/setup_env.batch
#
# Flash attention is optional. The default skips it for portability; enable it
# with RUN_FLASH_ATTN_BUILD=1 if you want to compare the faster kernel path.
#SBATCH -J acdir_hf_env
#SBATCH --nodes=1
#SBATCH --gpus-per-node=1
#SBATCH --cpus-per-task=12
#SBATCH --mem=160G
#SBATCH --time=12:00:00
#SBATCH -o batch_setup_env_%j.out
#SBATCH -e batch_setup_env_%j.err
#SBATCH --mail-type=END,FAIL

set -euo pipefail

SCRIPT_DIR="$(cd "$(dirname "${BASH_SOURCE[0]}")" && pwd)"
RELEASE_DIR="${RELEASE_DIR:-$(cd "${SCRIPT_DIR}/../.." && pwd)}"
CONDA_ENV_NAME="${CONDA_ENV_NAME:-acdir_math500}"
PYTHON_VERSION="${PYTHON_VERSION:-3.11}"
FORCE_RECREATE_ENV="${FORCE_RECREATE_ENV:-0}"
CONDA_MODULE="${CONDA_MODULE:-}"
CONDA_ROOT="${CONDA_ROOT:-}"
GCC_MODULE="${GCC_MODULE:-}"
CUDA_MODULE="${CUDA_MODULE:-}"
DEFAULT_CUDA_HOME="${DEFAULT_CUDA_HOME:-}"
PYTORCH_INDEX_URL="${PYTORCH_INDEX_URL:-https://download.pytorch.org/whl/cu121}"
TORCH_SPEC="${TORCH_SPEC:-torch==2.5.1+cu121}"
CONDA_CREATE_CHANNEL="${CONDA_CREATE_CHANNEL:-conda-forge}"
RUN_FLASH_ATTN_BUILD="${RUN_FLASH_ATTN_BUILD:-0}"
KEEP_FLASH_BUILD_ARTIFACTS="${KEEP_FLASH_BUILD_ARTIFACTS:-0}"
FLASH_ATTN_VERSION="${FLASH_ATTN_VERSION:-2.8.3}"
FLASH_ATTN_CUDA_ARCHS="${FLASH_ATTN_CUDA_ARCHS:-90}"
TORCH_CUDA_ARCH_LIST="${TORCH_CUDA_ARCH_LIST:-9.0}"
SLURM_MAX_JOBS="${SLURM_MAX_JOBS:-4}"
NVCC_THREADS="${NVCC_THREADS:-1}"

cd "${RELEASE_DIR}"

if ! type module >/dev/null 2>&1 && [ -f /etc/profile.d/modules.sh ]; then
  # shellcheck disable=SC1091
  source /etc/profile.d/modules.sh || true
fi
if type module >/dev/null 2>&1 && [ -n "${CONDA_MODULE}" ]; then
    module purge || true
    module load "${CONDA_MODULE}" || true
    if [ -n "${GCC_MODULE}" ]; then module load "${GCC_MODULE}" || true; fi
    if [ -n "${CUDA_MODULE}" ]; then module load "${CUDA_MODULE}" || true; fi
fi

if [ -z "${CUDA_HOME:-}" ] && [ -n "${DEFAULT_CUDA_HOME}" ]; then
  export CUDA_HOME="${DEFAULT_CUDA_HOME}"
elif [ -z "${CUDA_HOME:-}" ] && command -v nvcc >/dev/null 2>&1; then
  export CUDA_HOME="$(cd "$(dirname "$(command -v nvcc)")/.." && pwd)"
fi
export HF_HOME="${HF_HOME:-${RELEASE_DIR}/.cache/huggingface}"
export HF_HUB_CACHE="${HF_HUB_CACHE:-${HF_HOME}/hub}"
export HF_DATASETS_CACHE="${HF_DATASETS_CACHE:-${RELEASE_DIR}/.cache/datasets}"
export TRITON_CACHE_DIR="${TRITON_CACHE_DIR:-${RELEASE_DIR}/.cache/triton}"
export PYTHONPATH="${RELEASE_DIR}:${RELEASE_DIR}/lmdeploy:${PYTHONPATH:-}"
export PYTHONDONTWRITEBYTECODE="${PYTHONDONTWRITEBYTECODE:-1}"
mkdir -p "${HF_HOME}" "${HF_HUB_CACHE}" "${HF_DATASETS_CACHE}" "${TRITON_CACHE_DIR}"

if command -v conda >/dev/null 2>&1; then
  # shellcheck disable=SC1090
  source "$(conda info --base)/etc/profile.d/conda.sh"
elif [ -n "${CONDA_ROOT}" ] && [ -x "${CONDA_ROOT}/bin/conda" ]; then
  # shellcheck disable=SC1091
  source "${CONDA_ROOT}/etc/profile.d/conda.sh"
else
  echo "[error] conda not found. Load conda first or set CONDA_MODULE/CONDA_ROOT." >&2
  exit 1
fi

if [ "${FORCE_RECREATE_ENV}" = "1" ]; then
  conda env remove -y -n "${CONDA_ENV_NAME}" || true
fi
if conda env list | awk '{print $1}' | grep -qx "${CONDA_ENV_NAME}"; then
  echo "[env] using existing env: ${CONDA_ENV_NAME}"
else
  echo "[env] creating env: ${CONDA_ENV_NAME}"
  conda create -y -n "${CONDA_ENV_NAME}" --override-channels -c "${CONDA_CREATE_CHANNEL}" "python=${PYTHON_VERSION}" pip
fi
conda activate "${CONDA_ENV_NAME}"

python -m pip install -U pip setuptools wheel ninja packaging
python -m pip install --extra-index-url "${PYTORCH_INDEX_URL}" "${TORCH_SPEC}"
python -m pip install --extra-index-url "${PYTORCH_INDEX_URL}" -r requirements.txt

build_flash_attn() {
  if type module >/dev/null 2>&1; then
    if [ -n "${GCC_MODULE}" ]; then module load "${GCC_MODULE}" || true; fi
    if [ -n "${CUDA_MODULE}" ]; then module load "${CUDA_MODULE}" || true; fi
  fi
  hash -r

  if [ -z "${CUDA_HOME:-}" ]; then
    echo "[error] CUDA_HOME is unset and nvcc was not discoverable." >&2
    exit 1
  fi
  export PATH="${CUDA_HOME}/bin:${PATH}"
  export LD_LIBRARY_PATH="${CUDA_HOME}/lib64:${LD_LIBRARY_PATH:-}"
  export CC="${CC:-$(command -v gcc)}"
  export CXX="${CXX:-$(command -v g++)}"
  export CUDAHOSTCXX="${CUDAHOSTCXX:-${CXX}}"
  unset NVCC_PREPEND_FLAGS
  export NVCC_APPEND_FLAGS="${NVCC_APPEND_FLAGS:+${NVCC_APPEND_FLAGS} }--objdir-as-tempdir"
  export FLASH_ATTN_TMP_ROOT="${FLASH_ATTN_TMP_ROOT:-${RELEASE_DIR}/.tmp/flash-attn/${SLURM_JOB_ID:-manual}}"
  export FLASH_ATTN_TMPDIR="${FLASH_ATTN_TMPDIR:-${FLASH_ATTN_TMP_ROOT}/tmp}"
  export TMPDIR="${FLASH_ATTN_TMPDIR}"
  export TMP="${TMPDIR}"
  export TEMP="${TMPDIR}"
  export MAX_JOBS="${SLURM_MAX_JOBS}"
  export NVCC_THREADS
  export FLASH_ATTN_VERSION
  export TORCH_CUDA_ARCH_LIST
  export FLASH_ATTN_CUDA_ARCHS
  export FLASH_ATTENTION_FORCE_BUILD=TRUE
  mkdir -p "${TMPDIR}"

  if ! command -v nvcc >/dev/null 2>&1; then
    echo "[error] nvcc not found. Set CUDA_MODULE/CUDA_HOME before RUN_FLASH_ATTN_BUILD=1." >&2
    exit 1
  fi

  echo "================ Flash-attn build env ================"
  echo "python: $(which python)"
  python -V
  echo "CUDA_HOME: ${CUDA_HOME}"
  echo "TMPDIR: ${TMPDIR}"
  echo "NVCC_APPEND_FLAGS: ${NVCC_APPEND_FLAGS}"
  echo "nvcc: $(nvcc --version | grep release || true)"
  echo "gcc: $(gcc --version | head -n 1)"
  echo "TORCH_CUDA_ARCH_LIST: ${TORCH_CUDA_ARCH_LIST}"
  echo "FLASH_ATTN_CUDA_ARCHS: ${FLASH_ATTN_CUDA_ARCHS}"
  echo "MAX_JOBS: ${MAX_JOBS}"
  python - <<'PY'
import torch
print("torch:", torch.__version__, "cuda:", torch.version.cuda, "cuda_available:", torch.cuda.is_available())
if torch.cuda.is_available():
    print("gpu:", torch.cuda.get_device_name(0))
    print("capability:", torch.cuda.get_device_capability(0))
PY

  python -m pip uninstall -y flash-attn || true
  python -m pip install -U pip setuptools wheel ninja packaging
  export FLASH_ATTN_BUILD_ROOT="${FLASH_ATTN_BUILD_ROOT:-${FLASH_ATTN_TMP_ROOT}/build}"
  export FLASH_ATTN_WHEEL_DIR="${FLASH_ATTN_WHEEL_DIR:-${FLASH_ATTN_BUILD_ROOT}/wheelhouse}"
  export FLASH_ATTN_SRC_PARENT="${FLASH_ATTN_SRC_PARENT:-${FLASH_ATTN_BUILD_ROOT}/src}"
  mkdir -p "${FLASH_ATTN_BUILD_ROOT}" "${FLASH_ATTN_WHEEL_DIR}" "${FLASH_ATTN_SRC_PARENT}"
  rm -rf "${FLASH_ATTN_SRC_PARENT}"/flash?attn-"${FLASH_ATTN_VERSION}" "${FLASH_ATTN_WHEEL_DIR:?}"/*

  python - <<'PY'
import io
import json
import os
import tarfile
import urllib.request

version = os.environ["FLASH_ATTN_VERSION"]
src_parent = os.environ["FLASH_ATTN_SRC_PARENT"]
url = f"https://pypi.org/pypi/flash-attn/{version}/json"
data = json.load(urllib.request.urlopen(url))
sdist_url = [u["url"] for u in data["urls"] if u["packagetype"] == "sdist"][0]
print("Downloading sdist:", sdist_url)
blob = urllib.request.urlopen(sdist_url).read()
with tarfile.open(fileobj=io.BytesIO(blob), mode="r:gz") as tf:
    top_dirs = {m.name.split("/")[0] for m in tf.getmembers() if "/" in m.name}
    tf.extractall(path=src_parent)
candidates = [
    os.path.join(src_parent, d)
    for d in sorted(top_dirs)
    if d.startswith(("flash_attn-", "flash-attn-"))
]
if not candidates:
    raise SystemExit(f"Could not find extracted flash-attn source in {src_parent}")
with open(os.path.join(src_parent, ".src_dir"), "w") as f:
    f.write(candidates[0])
print("Extracted to:", candidates[0])
PY

  export FLASH_ATTN_SRC_DIR="$(cat "${FLASH_ATTN_SRC_PARENT}/.src_dir")"
  echo "Building wheel from ${FLASH_ATTN_SRC_DIR}"
  ( cd "${FLASH_ATTN_SRC_DIR}" && python setup.py bdist_wheel -d "${FLASH_ATTN_WHEEL_DIR}" )
  ls -lh "${FLASH_ATTN_WHEEL_DIR}"
  python -m pip install --force-reinstall --no-deps "${FLASH_ATTN_WHEEL_DIR}"/flash_attn-*.whl

  echo "================ Verify flash-attn ================"
  python - <<'PY'
import importlib.util
import torch
import flash_attn

print("torch:", torch.__version__, "cuda:", torch.version.cuda)
print("flash_attn:", flash_attn.__version__)
if torch.cuda.is_available():
    print("gpu:", torch.cuda.get_device_name(0))
spec = importlib.util.find_spec("flash_attn_2_cuda")
print("flash_attn_2_cuda:", spec.origin if spec else None)
PY

  if [ "${KEEP_FLASH_BUILD_ARTIFACTS}" != "1" ] \
    && [[ "${FLASH_ATTN_TMP_ROOT}" == "${RELEASE_DIR}/.tmp/flash-attn/"* ]]; then
    rm -rf -- "${FLASH_ATTN_TMP_ROOT}"
  fi
}

if [ "${RUN_FLASH_ATTN_BUILD}" = "1" ]; then
  build_flash_attn
fi

echo "================ Environment check ================"
python - <<'PY'
import importlib.util
import sys
import torch
import transformers
import datasets
from lmdeploy.pytorch.llada_exact import LLaDAExactRunner

print("python:", sys.executable)
print("torch:", torch.__version__, "cuda:", torch.version.cuda, "cuda_available:", torch.cuda.is_available())
print("transformers:", transformers.__version__)
print("datasets:", datasets.__version__)
print("vendored LLaDAExactRunner:", LLaDAExactRunner)
print("flash_attn spec:", importlib.util.find_spec("flash_attn"))
PY
echo "===================================================="