#!/usr/bin/env bash
# timm train.py: resnet26d from scratch on Imagenette (160px), 40 epochs, AMP, batch 128.
set -o pipefail
cd "$(dirname "$0")"
SEED=42
OUT=../data/train_out
rm -rf "$OUT/kv"
echo "KVERITAS_INPUT src=seed:$SEED"
echo "KVERITAS_PHASE name=train"
../env/bin/python -u train.py --data-dir ../data/imagenette2-160 --model resnet26d --num-classes 10 \
  --epochs 40 --img-size 160 -b 128 --amp --seed "$SEED" -j 8 --output "$OUT" --experiment kv 2>&1
RC=$?
echo "KVERITAS_PHASE name=summary"
../env/bin/python - "$OUT/kv/summary.csv" <<'PY'
import csv, sys
rows = list(csv.DictReader(open(sys.argv[1])))
for r in rows:
    for k in ("train_loss", "eval_loss", "eval_top1", "eval_top5"):
        if r.get(k):
            print(f"KVERITAS_METRIC name={k} value={r[k]} step={r['epoch']}")
last = rows[-1]
best = max(rows, key=lambda r: float(r["eval_top1"]))
print(f"KVERITAS_CLAIM metric=final_eval_top1 value={last['eval_top1']}")
print(f"KVERITAS_CLAIM metric=final_eval_top5 value={last['eval_top5']}")
print(f"KVERITAS_CLAIM metric=best_eval_top1 value={best['eval_top1']}")
PY
exit $RC
