#!/bin/bash # Portable Slurm helper. Supply site-specific account/partition flags to # sbatch, for example: sbatch -A -p ... # # 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 "===================================================="