#!/usr/bin/env bash
# README example: Diffusion Policy on PushT, 20,000 steps, batch 64, seed 1000; simulator evaluation
# (50 episodes) at steps 10,000 and 20,000. Metrics are parsed from the training log.
set -o pipefail
ROW=/var/tmp/kv-repro/018-huggingface-lerobot
OUT=$ROW/data/run
rm -rf "$OUT"
echo "KVERITAS_INPUT src=seed:1000"
echo "KVERITAS_PHASE name=train"
lerobot-train --policy.type=diffusion --dataset.repo_id=lerobot/pusht --env.type=pusht \
  --steps=20000 --batch_size=64 --env_eval_freq=10000 --save_freq=20000 --log_freq=200 \
  --eval.n_episodes=50 --eval.batch_size=10 --eval.use_async_envs=false \
  --policy.push_to_hub=false --output_dir="$OUT" --seed=1000 2>&1 | tee "$ROW/train.out" || exit 1
python - "$ROW/train.out" <<'PY'
import ast, re, sys
text = open(sys.argv[1]).read().replace("\r", "\n")
step = 0
evals = []
for line in text.splitlines():
    m = re.search(r"step:([0-9.]+)([KM]?) .*?loss:([0-9.]+)", line)
    if m:
        step = round(float(m.group(1)) * {"": 1, "K": 1e3, "M": 1e6}[m.group(2)])
        print(f"KVERITAS_METRIC name=train_loss value={m.group(3)} step={step}")
    m = re.search(r"Suite overall aggregated: (\{.*\})", line)
    if m:
        r = ast.literal_eval(m.group(1))
        evals.append(r)
        print(f"KVERITAS_METRIC name=pc_success value={r['pc_success']} step={step}")
        print(f"KVERITAS_METRIC name=avg_max_reward value={r['avg_max_reward']:.6g} step={step}")
final = evals[-1]
print(f"KVERITAS_CLAIM metric=pc_success value={final['pc_success']}")
print(f"KVERITAS_CLAIM metric=avg_max_reward value={final['avg_max_reward']:.6g}")
PY
