#!/usr/bin/env bash
# tinygrad examples: hlb_cifar10.py (default 1000 steps, BS 512), then beautiful_mnist.py (default 70 steps).
set -o pipefail
cd "$(dirname "$0")"
PY=../env/bin/python
export PYTHONPATH=.
export LD_LIBRARY_PATH="$(cd .. && pwd)/env/lib"  # conda-forge cuda-nvrtc for the NV backend
SEED=$(sed -n "s/^ *'seed' *: *\([0-9]*\).*/\1/p" examples/hlb_cifar10.py)
echo "KVERITAS_INPUT src=seed:$SEED"
echo "KVERITAS_PHASE name=hlb_cifar10"
$PY -u examples/hlb_cifar10.py 2>&1 | awk '
{ print; fflush() }
/^ *[0-9]+ +[0-9.]+ ms run,/ {
  if ($1 % 50 == 0) {
    match($0, /[0-9.]+ loss/); split(substr($0,RSTART,RLENGTH), l, " ")
    printf "KVERITAS_METRIC name=cifar_train_loss value=%s step=%s\n", l[1], $1; fflush()
  }
}
/^eval +[0-9]+\/[0-9]+ / {
  match($0, /[0-9.]+%/); a=substr($0,RSTART,RLENGTH-1)
  match($0, /STEP=[0-9]+/); s=substr($0,RSTART+5,RLENGTH-5)
  printf "KVERITAS_METRIC name=cifar_eval_acc_pct value=%s step=%s\nKVERITAS_CLAIM metric=cifar_eval_acc_pct value=%s\n", a, s, a; fflush()
}' || exit 1
echo "KVERITAS_PHASE name=beautiful_mnist"
$PY -u examples/beautiful_mnist.py 2>&1 | tr '\r' '\n' | awk '
{ print; fflush() }
/test_accuracy: +[0-9.]+%/ {
  match($0, /test_accuracy: +[0-9.]+/); split(substr($0,RSTART,RLENGTH), a, " "); acc=a[2]
  match($0, /loss: +[0-9.]+/); split(substr($0,RSTART,RLENGTH), l, " "); loss=l[2]
}
END { if (acc!="") printf "KVERITAS_METRIC name=mnist_final_loss value=%s\nKVERITAS_CLAIM metric=mnist_test_acc_pct value=%s\n", loss, acc }'
