acdir-llada-math500 / scripts /batch /eval_math500.batch
NYCU-MLLab's picture
Upload folder using huggingface_hub
4a28d4d verified
Raw
History Blame Contribute Delete
8.97 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/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[@]}"