| #!/bin/bash |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
|
|
| 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 |
| |
| 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 |
| |
| source "$(conda info --base)/etc/profile.d/conda.sh" |
| elif [ -n "${CONDA_ROOT}" ] && [ -x "${CONDA_ROOT}/bin/conda" ]; then |
| |
| 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 "====================================================" |
|
|