#!/usr/bin/env bash
# examples/mnist/main.py as shipped: Schedule-Free AdamW, 14 epochs, lr 0.0025, seed 1.
# Test loss and accuracy are parsed from its per-epoch output.
set -o pipefail
cd examples/mnist || exit 1
echo "KVERITAS_INPUT src=seed:1"
echo "KVERITAS_PHASE name=train"
python -u main.py 2>&1 | grep -v "^Train Epoch" | tee ../../../train.out || exit 1
python - ../../../train.out <<'PY'
import re, sys
rows = re.findall(r"Test set: Average loss: ([0-9.]+), Accuracy: (\d+)/(\d+) \(([0-9.]+)%\)", open(sys.argv[1]).read())
for i, (loss, _, _, acc) in enumerate(rows, 1):
    print(f"KVERITAS_METRIC name=test_loss value={loss} step={i}")
    print(f"KVERITAS_METRIC name=test_accuracy value={acc} step={i}")
print(f"KVERITAS_CLAIM metric=final_test_accuracy value={rows[-1][3]}")
print(f"KVERITAS_CLAIM metric=final_test_loss value={rows[-1][0]}")
PY
