#!/bin/bash # Portable Slurm helper. Supply site-specific account/partition flags to # sbatch, for example: sbatch -A -p ... # # 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:-}" 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:-}" echo "lookback blocks= ${EVAL_LOOKBACK_BLOCKS:-}" echo "age current/lb = ${EVAL_REMASK_MIN_AGE_CURRENT:-}/${EVAL_REMASK_MAX_AGE_LOOKBACK:-}" echo "max remask cap = ${EVAL_MAX_TOTAL_REMASK_PER_SAMPLE:-}" echo "force window = ${EVAL_FORCE_REMASK_WINDOW:-}" echo "reforward = ${EVAL_REFORWARD_AFTER_REMASK:-}" echo "joint argmax = ${EVAL_DETERMINISTIC_JOINT_ARGMAX:-}" echo "sample remask = ${EVAL_SAMPLE_REMASK:-}" echo "remask temp = ${EVAL_REMASK_TEMPERATURE:-}" echo "remask timing = ${EVAL_REMASK_TIMING:-}" echo "count bias = ${EVAL_COUNT_LOGIT_BIAS:-}" echo "sweep preset = ${EVAL_SWEEP_PRESET:-}" 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[@]}"