#!/usr/bin/env bash
# README GaLore pre-training command for LLaMA-60M on streamed C4 (allenai/c4), 2000 steps instead of 10k,
# micro-batch 64 with total batch 512 (the README's 256 does not fit 16 GB), eval every 500 steps.
# Metrics are parsed from the script's log lines.
set -o pipefail
ROW=/var/tmp/kv-repro/033-jiaweizzhao-GaLore
export HOME=$ROW/data/home WANDB_MODE=offline WANDB_DIR=$ROW/data
echo "KVERITAS_INPUT src=seed:0"
echo "KVERITAS_PHASE name=pretrain"
torchrun --standalone --nproc_per_node 1 torchrun_main.py --model_config configs/llama_60m.json --lr 0.01 \
  --galore_scale 0.25 --rank 128 --update_proj_gap 200 --batch_size 64 --total_batch_size 512 \
  --num_training_steps 2000 --warmup_steps 200 --dtype bfloat16 --optimizer galore_adamw \
  --eval_every 500 --save_dir "$ROW/data/checkpoints" 2>&1 | tr '\r' '\n' | grep -E "Eval loss|Final eval|Saving|Total params" | tee ../train.out || exit 1
python - ../train.out <<'PY'
import math, re, sys
text = open(sys.argv[1]).read()
for step, loss in re.findall(r"Eval loss at step (\d+): ([0-9.]+)", text):
    print(f"KVERITAS_METRIC name=eval_loss value={loss} step={step}")
    print(f"KVERITAS_METRIC name=eval_perplexity value={math.exp(float(loss)):.6g} step={step}")
final = float(re.search(r"Final eval loss: ([0-9.]+)", text).group(1))
print(f"KVERITAS_CLAIM metric=final_eval_loss value={final:.6g}")
print(f"KVERITAS_CLAIM metric=final_eval_perplexity value={math.exp(final):.6g}")
PY
