File size: 8,970 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 | #!/bin/bash
# Portable Slurm helper. Supply site-specific account/partition flags to
# sbatch, for example: sbatch -A <account> -p <partition> ...
#
# Usage:
# sbatch scripts/batch/eval_math500.batch
#
# Small smoke test:
# sbatch --export=ALL,MAX_EVAL_SAMPLES=8 scripts/batch/eval_math500.batch
#
# Default settings match the reported MATH500 run: lmdeploy, full_window,
# varlen flash on, CUDA graph off, batch size 1.
#SBATCH -J acdir_math500
#SBATCH --nodes=1
#SBATCH --ntasks-per-node=1
#SBATCH --gpus-per-node=1
#SBATCH --cpus-per-task=12
#SBATCH --time=1-00:00:00
#SBATCH -o batch_eval_math500.out
#SBATCH -e batch_eval_math500.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}"
CONDA_ENV_PATH="${CONDA_ENV_PATH:-}"
PYTHON_BIN="${PYTHON_BIN:-}"
CONDA_MODULE="${CONDA_MODULE:-}"
CONDA_ROOT="${CONDA_ROOT:-}"
CUDA_MODULE="${CUDA_MODULE:-}"
DEFAULT_CUDA_HOME="${DEFAULT_CUDA_HOME:-}"
BASE_MODEL="${BASE_MODEL:-${HF_BASE_MODEL:-GSAI-ML/LLaDA-8B-Instruct}}"
EVAL_BATCH_SIZE="${EVAL_BATCH_SIZE:-1}"
MAX_EVAL_SAMPLES="${MAX_EVAL_SAMPLES:-0}"
EVAL_CRITIC_CKPT="${EVAL_CRITIC_CKPT:-}"
EVAL_COMPARE_WITH_BASELINE="${EVAL_COMPARE_WITH_BASELINE:-}"
EVAL_LOOKBACK_BLOCKS="${EVAL_LOOKBACK_BLOCKS:-}"
EVAL_REMASK_MIN_AGE_CURRENT="${EVAL_REMASK_MIN_AGE_CURRENT:-}"
EVAL_REMASK_MAX_AGE_LOOKBACK="${EVAL_REMASK_MAX_AGE_LOOKBACK:-}"
EVAL_MAX_TOTAL_REMASK_PER_SAMPLE="${EVAL_MAX_TOTAL_REMASK_PER_SAMPLE:-}"
EVAL_FORCE_REMASK_WINDOW="${EVAL_FORCE_REMASK_WINDOW:-}"
EVAL_REFORWARD_AFTER_REMASK="${EVAL_REFORWARD_AFTER_REMASK:-}"
EVAL_DETERMINISTIC_JOINT_ARGMAX="${EVAL_DETERMINISTIC_JOINT_ARGMAX:-}"
EVAL_SAMPLE_REMASK="${EVAL_SAMPLE_REMASK:-}"
EVAL_REMASK_TEMPERATURE="${EVAL_REMASK_TEMPERATURE:-}"
EVAL_REMASK_TIMING="${EVAL_REMASK_TIMING:-}"
EVAL_COUNT_LOGIT_BIAS="${EVAL_COUNT_LOGIT_BIAS:-}"
EVAL_SWEEP_PRESET="${EVAL_SWEEP_PRESET:-}"
EVAL_CLEAN_OUTPUT="${EVAL_CLEAN_OUTPUT:-True}"
EVAL_PROGRESS_EVERY="${EVAL_PROGRESS_EVERY:-50}"
MASTER_PORT="${MASTER_PORT:-$((29500 + (${SLURM_JOB_ID:-0} % 1000)))}"
RESULT_DIR="${RESULT_DIR:-outputs/math500_${SLURM_JOB_ID:-manual}}"
case "${EVAL_SWEEP_PRESET}" in
""|none)
;;
baseline|k4be_baseline)
EVAL_COUNT_LOGIT_BIAS="${EVAL_COUNT_LOGIT_BIAS:-}"
;;
mild|k4be_mild)
EVAL_COUNT_LOGIT_BIAS="${EVAL_COUNT_LOGIT_BIAS:--0.25,0.05,0.25,0.35,0.35}"
;;
medium|k4be_medium)
EVAL_COUNT_LOGIT_BIAS="${EVAL_COUNT_LOGIT_BIAS:--0.50,0.10,0.45,0.60,0.60}"
;;
mediumcap|k4be_mediumcap)
EVAL_COUNT_LOGIT_BIAS="${EVAL_COUNT_LOGIT_BIAS:--0.50,0.10,0.45,0.60,0.60}"
EVAL_MAX_TOTAL_REMASK_PER_SAMPLE="${EVAL_MAX_TOTAL_REMASK_PER_SAMPLE:-4}"
;;
strong|k4be_strong)
EVAL_COUNT_LOGIT_BIAS="${EVAL_COUNT_LOGIT_BIAS:--1.00,0.00,1.00,1.25,1.25}"
;;
strongcap|k4be_strongcap)
EVAL_COUNT_LOGIT_BIAS="${EVAL_COUNT_LOGIT_BIAS:--1.00,0.00,1.00,1.25,1.25}"
EVAL_MAX_TOTAL_REMASK_PER_SAMPLE="${EVAL_MAX_TOTAL_REMASK_PER_SAMPLE:-4}"
;;
*)
echo "[error] unknown EVAL_SWEEP_PRESET: ${EVAL_SWEEP_PRESET}" >&2
exit 1
;;
esac
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 "${CUDA_MODULE}" ]; then module load "${CUDA_MODULE}" || true; fi
fi
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. Run scripts/batch/setup_env.batch first or set CONDA_MODULE/CONDA_ROOT." >&2
exit 1
fi
if [ -z "${PYTHON_BIN}" ] && [ -n "${CONDA_ENV_PATH}" ] && [ -x "${CONDA_ENV_PATH}/bin/python" ]; then
PYTHON_BIN="${CONDA_ENV_PATH}/bin/python"
fi
if [ -n "${PYTHON_BIN}" ] && [ -x "${PYTHON_BIN}" ]; then
export PATH="$(dirname "${PYTHON_BIN}"):${PATH}"
elif conda activate "${CONDA_ENV_NAME}"; then
PYTHON_BIN="$(command -v python)"
:
elif [ -n "${CONDA_ENV_PATH}" ] && [ -d "${CONDA_ENV_PATH}" ]; then
conda activate "${CONDA_ENV_PATH}"
PYTHON_BIN="$(command -v python)"
else
echo "[error] conda environment not found: ${CONDA_ENV_NAME}" >&2
if [ -n "${CONDA_ENV_PATH}" ]; then
echo "[error] fallback path missing: ${CONDA_ENV_PATH}" >&2
fi
conda info --envs >&2 || true
exit 1
fi
gpu_count="${SLURM_GPUS_ON_NODE:-${GPUS_PER_NODE:-1}}"
gpu_count="$(printf '%s\n' "${gpu_count}" | grep -o '[0-9][0-9]*' | tail -n 1 || true)"
NPROC_PER_NODE="${NPROC_PER_NODE:-${gpu_count:-1}}"
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 PYTHONPATH="${RELEASE_DIR}:${RELEASE_DIR}/lmdeploy:${PYTHONPATH:-}"
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 LLADA_EXACT_BACKEND="${LLADA_EXACT_BACKEND:-lmdeploy}"
export LLADA_LMDEPLOY_FAST_MODE="${LLADA_LMDEPLOY_FAST_MODE:-full_window}"
export LLADA_LMDEPLOY_CUDAGRAPH="${LLADA_LMDEPLOY_CUDAGRAPH:-0}"
export LLADA_LMDEPLOY_VARLEN_FLASH="${LLADA_LMDEPLOY_VARLEN_FLASH:-1}"
export LLADA_FORCE_MATH_SDPA="${LLADA_FORCE_MATH_SDPA:-0}"
export ACDIR_DIST_TIMEOUT_MIN="${ACDIR_DIST_TIMEOUT_MIN:-120}"
export TOKENIZERS_PARALLELISM="${TOKENIZERS_PARALLELISM:-false}"
export PYTHONDONTWRITEBYTECODE="${PYTHONDONTWRITEBYTECODE:-1}"
mkdir -p "${HF_HOME}" "${HF_HUB_CACHE}" "${HF_DATASETS_CACHE}" "${TRITON_CACHE_DIR}" "${RESULT_DIR}"
COMPARE_ARGS=()
if [ -n "${EVAL_COMPARE_WITH_BASELINE}" ]; then
COMPARE_ARGS=(--compare_with_baseline "${EVAL_COMPARE_WITH_BASELINE}")
fi
CRITIC_ARGS=()
if [ -n "${EVAL_CRITIC_CKPT}" ]; then
CRITIC_ARGS=(--critic_ckpt "${EVAL_CRITIC_CKPT}")
fi
echo "================ ACDiR-LLaDA MATH500 eval ================"
echo "release dir = ${RELEASE_DIR}"
echo "base model = ${BASE_MODEL}"
echo "critic ckpt = ${EVAL_CRITIC_CKPT:-<config>}"
echo "conda env = ${CONDA_ENV_NAME}"
echo "python = ${PYTHON_BIN}"
echo "batch size = ${EVAL_BATCH_SIZE}"
echo "nproc = ${NPROC_PER_NODE}"
echo "max samples = ${MAX_EVAL_SAMPLES}"
echo "result dir = ${RESULT_DIR}"
echo "compare(base) = ${EVAL_COMPARE_WITH_BASELINE:-<config>}"
echo "lookback blocks= ${EVAL_LOOKBACK_BLOCKS:-<config>}"
echo "age current/lb = ${EVAL_REMASK_MIN_AGE_CURRENT:-<config>}/${EVAL_REMASK_MAX_AGE_LOOKBACK:-<config>}"
echo "max remask cap = ${EVAL_MAX_TOTAL_REMASK_PER_SAMPLE:-<config>}"
echo "force window = ${EVAL_FORCE_REMASK_WINDOW:-<config>}"
echo "reforward = ${EVAL_REFORWARD_AFTER_REMASK:-<config>}"
echo "joint argmax = ${EVAL_DETERMINISTIC_JOINT_ARGMAX:-<config>}"
echo "sample remask = ${EVAL_SAMPLE_REMASK:-<config>}"
echo "remask temp = ${EVAL_REMASK_TEMPERATURE:-<config>}"
echo "remask timing = ${EVAL_REMASK_TIMING:-<config>}"
echo "count bias = ${EVAL_COUNT_LOGIT_BIAS:-<none>}"
echo "sweep preset = ${EVAL_SWEEP_PRESET:-<none>}"
echo "clean output = ${EVAL_CLEAN_OUTPUT}"
echo "progress every = ${EVAL_PROGRESS_EVERY}"
echo "fast mode = ${LLADA_LMDEPLOY_FAST_MODE}"
echo "varlen flash = ${LLADA_LMDEPLOY_VARLEN_FLASH}"
echo "math sdpa = ${LLADA_FORCE_MATH_SDPA}"
echo "cuda graph = ${LLADA_LMDEPLOY_CUDAGRAPH}"
echo "==========================================================="
"${PYTHON_BIN}" eval_math500.py \
--base_model "${BASE_MODEL}" \
--batch_size "${EVAL_BATCH_SIZE}" \
"${CRITIC_ARGS[@]}" \
--nproc_per_node "${NPROC_PER_NODE}" \
--master_port "${MASTER_PORT}" \
--max_eval_samples "${MAX_EVAL_SAMPLES}" \
--lookback_blocks "${EVAL_LOOKBACK_BLOCKS}" \
--remask_min_age_current "${EVAL_REMASK_MIN_AGE_CURRENT}" \
--remask_max_age_lookback "${EVAL_REMASK_MAX_AGE_LOOKBACK}" \
--max_total_remask_per_sample "${EVAL_MAX_TOTAL_REMASK_PER_SAMPLE}" \
--force_remask_window "${EVAL_FORCE_REMASK_WINDOW}" \
--reforward_after_remask "${EVAL_REFORWARD_AFTER_REMASK}" \
--deterministic_joint_argmax "${EVAL_DETERMINISTIC_JOINT_ARGMAX}" \
--sample_remask "${EVAL_SAMPLE_REMASK}" \
--remask_temperature "${EVAL_REMASK_TEMPERATURE}" \
--remask_timing "${EVAL_REMASK_TIMING}" \
--count_logit_bias="${EVAL_COUNT_LOGIT_BIAS}" \
--clean_output "${EVAL_CLEAN_OUTPUT}" \
--progress_every "${EVAL_PROGRESS_EVERY}" \
--result_dir "${RESULT_DIR}" \
"${COMPARE_ARGS[@]}"
|