#!/usr/bin/env bash
# README example: LoRA fine-tuning of Pythia-410M on Alpaca2k (2 epochs), then ARC-Easy evaluation.
# Metrics are parsed from litgpt's output and the results.json it writes.
set -o pipefail
ROW=/var/tmp/kv-repro/015-Lightning-AI-litgpt
C=$ROW/data/checkpoints/EleutherAI/pythia-410m
OUT=$ROW/data/run
rm -rf "$OUT"
echo "KVERITAS_PHASE name=finetune"
litgpt finetune_lora "$C" --data Alpaca2k --data.download_dir "$ROW/data/alpaca2k" --train.epochs 2 --out_dir "$OUT" 2>&1 \
  | tee "$ROW/finetune.out" || exit 1
python - "$ROW/finetune.out" <<'PY'
import re, sys
text = open(sys.argv[1]).read()
for e, i, loss in re.findall(r"Epoch (\d+) \| iter (\d+) step \d+ \| loss train: ([0-9.]+)", text):
    if int(i) % 50 == 0:
        print(f"KVERITAS_METRIC name=train_loss value={loss} step={i}")
for i, v in re.findall(r"iter (\d+): val loss ([0-9.]+)", text):
    print(f"KVERITAS_METRIC name=val_loss value={v} step={i}")
m = re.search(r"Final evaluation \| val loss: ([0-9.]+) \| val ppl: ([0-9.]+)", text)
print(f"KVERITAS_CLAIM metric=final_val_loss value={m.group(1)}")
print(f"KVERITAS_CLAIM metric=final_val_ppl value={m.group(2)}")
PY
echo "KVERITAS_PHASE name=evaluate"
litgpt evaluate "$OUT/final" --tasks arc_easy --out_dir "$OUT/eval" 2>&1 | tail -20 || exit 1
python - "$OUT/eval/results.json" <<'PY'
import json, sys
r = json.load(open(sys.argv[1]))["results"]["arc_easy"]
for k in ("acc,none", "acc_norm,none"):
    name = "arc_easy_" + k.split(",")[0]
    print(f"KVERITAS_METRIC name={name} value={r[k]:.6g}")
    print(f"KVERITAS_CLAIM metric={name} value={r[k]:.6g}")
PY
