| #!/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}" |
| 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 |
| |
| 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 |
| |
| 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. 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[@]}" |
|
|