#!/usr/bin/env bash
set -euo pipefail

: "${STORAGE_ROOT:?Set STORAGE_ROOT to the shared-storage mount path}"
: "${RUN_ID:?Set a unique RUN_ID}"

readonly VERL_COMMIT=ddd86f527a4af75095e4677b02b5aa272913a088
readonly DATASET_REVISION=740312add88f781978c0658806c59bc2815b9866
readonly TARGET_STEPS=2
readonly TRAIN_BATCH_SIZE=4
readonly ROLLOUT_N=2
readonly VERL_ROOT="${STORAGE_ROOT}/code/verl/${VERL_COMMIT}"
readonly MODEL_ROOT="${STORAGE_ROOT}/models/modelscope/deepseek-ai/DeepSeek-R1-Distill-Qwen-1.5B/6fc93244f442ee2b5ab5c8000687ef5f7ffe1d03"
readonly DATA_ROOT="${STORAGE_ROOT}/datasets/huggingface/openai/gsm8k/${DATASET_REVISION}/processed/verl-v0.6.0-v1"
readonly TOOL_ROOT="${STORAGE_ROOT}/tools/verl/single-gpu-grpo"
readonly VERIFY_SCRIPT="${TOOL_ROOT}/verify-single-gpu-grpo-v1.py"
readonly OUTPUT_ROOT="${STORAGE_ROOT}/runs/verl/single-gpu-grpo"
readonly OUTPUT_DIR="${OUTPUT_ROOT}/${RUN_ID}"

if [[ -e "${OUTPUT_DIR}" ]]; then
  echo "Refusing to reuse OUTPUT_DIR: ${OUTPUT_DIR}" >&2
  exit 1
fi
mkdir -p "${OUTPUT_ROOT}"
mkdir "${OUTPUT_DIR}"

readonly DRIVER_LOG="${OUTPUT_DIR}/driver.log"
readonly TRAINING_LOG="${OUTPUT_DIR}/training.log"
readonly CHECKPOINT_ROOT="${OUTPUT_DIR}/checkpoints"
readonly TENSORBOARD_DIR="${OUTPUT_DIR}/tensorboard"
readonly ROLLOUT_DIR="${OUTPUT_DIR}/rollouts"
readonly HYDRA_DIR="${OUTPUT_DIR}/hydra"
exec > >(tee -a "${DRIVER_LOG}") 2>&1

on_error() {
  local status=$?
  if [[ ! -e "${OUTPUT_DIR}/SUCCESS" ]]; then
    printf '%s exit=%s\n' "$(date -u +%Y-%m-%dT%H:%M:%SZ)" "${status}" > "${OUTPUT_DIR}/FAILED"
  fi
  exit "${status}"
}
trap on_error ERR

for path in "${VERL_ROOT}" "${MODEL_ROOT}" "${DATA_ROOT}"; do
  test -d "${path}"
done
for path in "${DATA_ROOT}/train.parquet" "${DATA_ROOT}/test.parquet" "${VERIFY_SCRIPT}"; do
  test -r "${path}"
done

if [[ -n "${PYTHONPATH:-}" ]]; then
  export PYTHONPATH="${VERL_ROOT}:${PYTHONPATH}"
else
  export PYTHONPATH="${VERL_ROOT}"
fi
export PYTHONUNBUFFERED=1
export HYDRA_FULL_ERROR=1
export HF_HUB_OFFLINE=1
export HF_DATASETS_OFFLINE=1
export TRANSFORMERS_OFFLINE=1
export WANDB_MODE=disabled
export TOKENIZERS_PARALLELISM=true
export TENSORBOARD_DIR
unset RAY_ADDRESS || true

echo "RUN_ID: ${RUN_ID}"
echo "STORAGE_ROOT: ${STORAGE_ROOT}"
findmnt -T "${STORAGE_ROOT}" -o TARGET,SOURCE,FSTYPE,OPTIONS
nvidia-smi --query-gpu=driver_version,name,memory.total --format=csv,noheader
echo "VERL_COMMIT: $(git -C "${VERL_ROOT}" rev-parse HEAD)"
test -z "$(git -C "${VERL_ROOT}" status --porcelain)"
(cd "${MODEL_ROOT}" && sha256sum -c SHA256SUMS)
(cd "${DATA_ROOT}" && sha256sum -c SHA256SUMS)

python3 "${VERIFY_SCRIPT}" preflight \
  --run-id "${RUN_ID}" \
  --output-dir "${OUTPUT_DIR}" \
  --target-steps "${TARGET_STEPS}" \
  --train-batch-size "${TRAIN_BATCH_SIZE}" \
  --rollout-n "${ROLLOUT_N}" \
  --storage-root "${STORAGE_ROOT}" \
  --verl-root "${VERL_ROOT}" \
  --model-root "${MODEL_ROOT}" \
  --data-root "${DATA_ROOT}"

training_command=(
  python3 -m verl.trainer.main_ppo
  "algorithm.adv_estimator=grpo"
  "algorithm.use_kl_in_reward=false"
  "data.train_files=${DATA_ROOT}/train.parquet"
  "data.val_files=${DATA_ROOT}/test.parquet"
  "data.train_batch_size=${TRAIN_BATCH_SIZE}"
  "data.max_prompt_length=512"
  "data.max_response_length=256"
  "data.filter_overlong_prompts=true"
  "data.truncation=error"
  "data.shuffle=false"
  "actor_rollout_ref.model.path=${MODEL_ROOT}"
  "actor_rollout_ref.model.use_shm=false"
  "actor_rollout_ref.model.use_remove_padding=true"
  "actor_rollout_ref.model.enable_gradient_checkpointing=true"
  "actor_rollout_ref.actor.strategy=fsdp2"
  "actor_rollout_ref.ref.strategy=fsdp2"
  "critic.strategy=fsdp2"
  "reward_model.strategy=fsdp2"
  "actor_rollout_ref.actor.optim.lr=1e-6"
  "actor_rollout_ref.actor.ppo_mini_batch_size=${TRAIN_BATCH_SIZE}"
  "actor_rollout_ref.actor.ppo_micro_batch_size_per_gpu=2"
  "actor_rollout_ref.actor.use_kl_loss=false"
  "actor_rollout_ref.actor.entropy_coeff=0"
  "actor_rollout_ref.actor.fsdp_config.param_offload=true"
  "actor_rollout_ref.actor.fsdp_config.optimizer_offload=true"
  "actor_rollout_ref.rollout.name=sglang"
  "actor_rollout_ref.rollout.tensor_model_parallel_size=1"
  "actor_rollout_ref.rollout.gpu_memory_utilization=0.3"
  "actor_rollout_ref.rollout.n=${ROLLOUT_N}"
  "actor_rollout_ref.rollout.log_prob_micro_batch_size_per_gpu=2"
  "actor_rollout_ref.rollout.max_num_seqs=16"
  "actor_rollout_ref.rollout.max_model_len=768"
  "actor_rollout_ref.rollout.max_num_batched_tokens=2048"
  "actor_rollout_ref.rollout.enable_chunked_prefill=false"
  "actor_rollout_ref.rollout.enforce_eager=true"
  "trainer.logger=[console,tensorboard]"
  "trainer.project_name=verl_aistudio"
  "trainer.experiment_name=${RUN_ID}"
  "trainer.n_gpus_per_node=1"
  "trainer.nnodes=1"
  "trainer.val_before_train=false"
  "trainer.test_freq=-1"
  "trainer.save_freq=${TARGET_STEPS}"
  "trainer.total_epochs=1"
  "trainer.total_training_steps=${TARGET_STEPS}"
  "trainer.resume_mode=disable"
  "trainer.default_local_dir=${CHECKPOINT_ROOT}"
  "trainer.rollout_data_dir=${ROLLOUT_DIR}"
  "ray_kwargs.ray_init.num_cpus=8"
  "hydra.job.chdir=false"
  "hydra.output_subdir=null"
  "hydra.run.dir=${HYDRA_DIR}"
)

printf '%q ' "${training_command[@]}" > "${OUTPUT_DIR}/resolved-command.txt"
printf '\n' >> "${OUTPUT_DIR}/resolved-command.txt"
printf 'Executing: '
printf '%q ' "${training_command[@]}"
printf '\n'

set +e
"${training_command[@]}" 2>&1 | tee -a "${TRAINING_LOG}"
training_status=${PIPESTATUS[0]}
set -e
if (( training_status != 0 )); then
  echo "verl training exited with status ${training_status}" >&2
  exit "${training_status}"
fi

python3 "${VERIFY_SCRIPT}" postflight \
  --run-id "${RUN_ID}" \
  --output-dir "${OUTPUT_DIR}" \
  --target-steps "${TARGET_STEPS}" \
  --train-batch-size "${TRAIN_BATCH_SIZE}" \
  --rollout-n "${ROLLOUT_N}" \
  --training-log "${TRAINING_LOG}"

trap - ERR
echo "Verification passed: ${OUTPUT_DIR}/verification.json"
