|
| 1 | +#!/bin/bash |
| 2 | +# Copyright (c) 2025 BAAI. All rights reserved. |
| 3 | +# Setup script for Hygon DCU CI environment. |
| 4 | +set -euo pipefail |
| 5 | + |
| 6 | +git config --global --add safe.directory "$(pwd)" |
| 7 | + |
| 8 | +export GEMS_VENDOR="${GEMS_VENDOR:-hygon}" |
| 9 | +export DTK_HOME="${DTK_HOME:-/opt/dtk}" |
| 10 | +export ROCM_PATH="${ROCM_PATH:-${DTK_HOME}}" |
| 11 | +export HIP_PATH="${HIP_PATH:-${DTK_HOME}/hip}" |
| 12 | +export HSA_PATH="${HSA_PATH:-${DTK_HOME}/hsa}" |
| 13 | +export HIP_CLANG_PATH="${HIP_CLANG_PATH:-${DTK_HOME}/llvm/bin}" |
| 14 | +export DEVICE_LIB_PATH="${DEVICE_LIB_PATH:-${DTK_HOME}/amdgcn/bitcode}" |
| 15 | + |
| 16 | +DTK_PATH="${DTK_HOME}/bin:${HIP_PATH}/bin:${HIP_CLANG_PATH}" |
| 17 | +DTK_LIBRARY_PATH="/opt/hyhal/lib/criu:/opt/hyhal/lib/rocprofiler:/opt/hyhal/lib:${HIP_PATH}/lib:${DTK_HOME}/lib:${DTK_HOME}/llvm/lib:${DTK_HOME}/dcc/lib:${DTK_HOME}/aillvm/lib:${HSA_PATH}/lib" |
| 18 | +export PATH="${DTK_PATH}:${PATH}" |
| 19 | +export LD_LIBRARY_PATH="${DTK_LIBRARY_PATH}${LD_LIBRARY_PATH:+:${LD_LIBRARY_PATH}}" |
| 20 | + |
| 21 | +if [[ -n "${GITHUB_ENV:-}" ]]; then |
| 22 | + for name in \ |
| 23 | + GEMS_VENDOR \ |
| 24 | + DTK_HOME \ |
| 25 | + ROCM_PATH \ |
| 26 | + HIP_PATH \ |
| 27 | + HSA_PATH \ |
| 28 | + HIP_CLANG_PATH \ |
| 29 | + DEVICE_LIB_PATH \ |
| 30 | + PATH \ |
| 31 | + LD_LIBRARY_PATH; do |
| 32 | + echo "${name}=${!name}" >> "${GITHUB_ENV}" |
| 33 | + done |
| 34 | +fi |
| 35 | + |
| 36 | +echo "DTK_HOME=${DTK_HOME}" |
| 37 | +echo "LD_LIBRARY_PATH=${LD_LIBRARY_PATH}" |
| 38 | +test -e "${HIP_PATH}/lib/libgalaxyhip.so.5" |
| 39 | +test -e "${DTK_HOME}/llvm/lib/libomp.so" |
| 40 | + |
| 41 | +TEST_DEPS=( |
| 42 | + pytest |
| 43 | + pytest-cov |
| 44 | + pytest-timeout |
| 45 | + pytest-json-report |
| 46 | + scikit-build-core==0.11 |
| 47 | + pybind11 |
| 48 | + ninja |
| 49 | + cmake |
| 50 | + numpy |
| 51 | + requests |
| 52 | + openai |
| 53 | + decorator |
| 54 | + pyyaml |
| 55 | + sqlalchemy |
| 56 | +) |
| 57 | + |
| 58 | +FLAGGEMS_REF="8d23621eb1381ae96b315a9287ce3cc555433824" |
| 59 | +FLAGGEMS_SOURCE="${FLAGGEMS_PATH:-/workspace/FlagGems}" |
| 60 | + |
| 61 | +if command -v uv >/dev/null 2>&1; then |
| 62 | + uv pip install --system --upgrade pip |
| 63 | + uv pip install --system --no-build-isolation -e . --no-deps |
| 64 | + uv pip install --system "${TEST_DEPS[@]}" |
| 65 | +else |
| 66 | + python -m pip install --upgrade pip |
| 67 | + python -m pip install --no-build-isolation -e . --no-deps |
| 68 | + python -m pip install "${TEST_DEPS[@]}" |
| 69 | +fi |
| 70 | + |
| 71 | +test -d "${FLAGGEMS_SOURCE}/src/flag_gems" |
| 72 | +git config --global --add safe.directory "${FLAGGEMS_SOURCE}" |
| 73 | + |
| 74 | +FLAGGEMS_HEAD="$(git -C "${FLAGGEMS_SOURCE}" rev-parse HEAD)" |
| 75 | +if [[ "${FLAGGEMS_HEAD}" != "${FLAGGEMS_REF}" ]]; then |
| 76 | + echo "Unexpected FlagGems commit: ${FLAGGEMS_HEAD}; expected ${FLAGGEMS_REF}." |
| 77 | + exit 1 |
| 78 | +fi |
| 79 | + |
| 80 | +if ! grep -q 'kwargs.pop("num_ldmatrixes", None)' \ |
| 81 | + "${FLAGGEMS_SOURCE}/src/flag_gems/runtime/configloader.py"; then |
| 82 | + echo "FlagGems num_ldmatrixes compatibility patch is missing." |
| 83 | + exit 1 |
| 84 | +fi |
| 85 | + |
| 86 | +FLAGGEMS_DIR="$(mktemp -d)/FlagGems" |
| 87 | +cp -a "${FLAGGEMS_SOURCE}" "${FLAGGEMS_DIR}" |
| 88 | + |
| 89 | +if command -v uv >/dev/null 2>&1; then |
| 90 | + uv pip install --system --no-build-isolation -e "${FLAGGEMS_DIR}" --no-deps |
| 91 | +else |
| 92 | + python -m pip install --no-build-isolation -e "${FLAGGEMS_DIR}" --no-deps |
| 93 | +fi |
| 94 | + |
| 95 | +python - <<'PY' |
| 96 | +import flag_gems |
| 97 | +import torch |
| 98 | +
|
| 99 | +print(f"FlagGems import ok: {getattr(flag_gems, '__version__', 'unknown')}") |
| 100 | +print(f"Torch import ok: {torch.__version__}") |
| 101 | +print(f"Accelerator available: {torch.cuda.is_available()}") |
| 102 | +print(f"Accelerator count: {torch.cuda.device_count()}") |
| 103 | +PY |
0 commit comments