acdir-llada-math500 / scripts /batch /setup_env.batch
NYCU-MLLab's picture
Upload folder using huggingface_hub
4a28d4d verified
Raw
History Blame Contribute Delete
9.28 kB
#!/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 "===================================================="