#!/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 EXPECTED_NODES=2
readonly GPUS_PER_NODE=2
readonly WORLD_SIZE=$((EXPECTED_NODES * GPUS_PER_NODE))
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/two-worker-grpo"
readonly VERIFY_SCRIPT="${TOOL_ROOT}/verify-two-worker-grpo-v1.py"
readonly EXPECTED_VERIFY_SHA256="3a2d6ffc85beef7479b08762f4dc00820957ae21dc153227b981bf5eb1b7a527"
readonly OUTPUT_ROOT="${STORAGE_ROOT}/runs/verl/two-worker-grpo"
readonly OUTPUT_DIR="${OUTPUT_ROOT}/${RUN_ID}"

[[ "${RUN_ID}" =~ ^[A-Za-z0-9][A-Za-z0-9._-]*$ ]] || {
  echo "RUN_ID contains unsupported characters" >&2
  exit 1
}
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_exit() {
  local status=$?
  if (( status != 0 )) && [[ ! -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_exit EXIT

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
echo "${EXPECTED_VERIFY_SHA256}  ${VERIFY_SCRIPT}" | sha256sum -c -

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 NCCL_IB_DISABLE=1
export CUDA_DEVICE_MAX_CONNECTIONS=1
export TENSORBOARD_DIR

echo "RUN_ID: ${RUN_ID}"
echo "STORAGE_ROOT: ${STORAGE_ROOT}"
echo "RAY_ADDRESS: ${RAY_ADDRESS:-auto}"
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}" \
  --expected-nodes "${EXPECTED_NODES}" \
  --expected-gpus-per-node "${GPUS_PER_NODE}"

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=${GPUS_PER_NODE}"
  "trainer.nnodes=${EXPECTED_NODES}"
  "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.address=auto"
  "+ray_kwargs.ray_init.runtime_env.env_vars.TENSORBOARD_DIR=${TENSORBOARD_DIR}"
  "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}" \
  --expected-nodes "${EXPECTED_NODES}" \
  --expected-gpus-per-node "${GPUS_PER_NODE}" \
  --world-size "${WORLD_SIZE}"

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