#!/usr/bin/env bash
# Purpose: pretrain a small nanochat base model inside a stated wall-clock budget:
#          calibrate the throughput with a short run, turn the budget into an
#          iteration count, train with intermediate checkpoints, then evaluate the
#          result and write the configuration out for the recording script
# Platform: all (CUDA on Tracks S and N, ROCm on Track X, MPS on Track M, CPU anywhere)
# Minimum memory: 8 GB
# Assumes: prepare-data.sh has already been run for this track, so $NANOCHAT holds
#          a synced .venv, $NANOCHAT_BASE_DIR holds corpus shards and a trained
#          tokeniser, and about 2 GB of free disk is available for checkpoints
set -euo pipefail

# ---------------------------------------------------------------- settings ---
TRACK="${TRACK:-}"                 # spark | strix | mac | nvidia | cpu
NANOCHAT="${NANOCHAT:-$HOME/nanochat}"
TARGET_MINUTES="${TARGET_MINUTES:-25}"   # wall clock the real run should take
DEPTH="${DEPTH:-6}"                # the one size dial; width and heads follow
MAX_SEQ_LEN="${MAX_SEQ_LEN:-512}"
# The logits tensor is device_batch_size x max_seq_len x vocab_size, materialised in
# fp32 for the loss, so it dominates peak memory: at 8 x 512 x 32768 x 4 bytes it is
# already about half a gigabyte before the copies the softcap makes. 8 is chosen for
# the 8 GB memory floor; raise it while the peak memory the run prints leaves room.
DEVICE_BATCH_SIZE="${DEVICE_BATCH_SIZE:-8}"
TOTAL_BATCH_SIZE="${TOTAL_BATCH_SIZE:-16384}"  # must be a multiple of batch x seq len
HEAD_DIM="${HEAD_DIM:-64}"
WINDOW_PATTERN="${WINDOW_PATTERN:-L}"   # full context on every layer; see the page
CALIBRATE_STEPS="${CALIBRATE_STEPS:-40}"
MIN_ITERATIONS="${MIN_ITERATIONS:-200}"
EVAL_EVERY="${EVAL_EVERY:-100}"
EVAL_TOKENS="${EVAL_TOKENS:-524288}"
SAMPLE_EVERY="${SAMPLE_EVERY:-100}"
EVAL_SPLIT_TOKENS="${EVAL_SPLIT_TOKENS:-16384}"
EVAL_MAX_PER_TASK="${EVAL_MAX_PER_TASK:-16}"
MODEL_TAG="${MODEL_TAG:-d${DEPTH}-lab}"
ITERATIONS="${ITERATIONS:-0}"      # set this to skip calibration entirely
LOG_DIR="${LOG_DIR:-$PWD/part-12-logs}"
FIELDS="${FIELDS:-$PWD/train-fields.json}"

export NANOCHAT_BASE_DIR="${NANOCHAT_BASE_DIR:-$HOME/.cache/nanochat}"
export OMP_NUM_THREADS="${OMP_NUM_THREADS:-1}"

CAL_TAG="calibration-tmp"

usage() {
    cat <<'USAGE'
Usage: TRACK=<spark|strix|mac|nvidia|cpu> bash train-small.sh

Environment variables (all optional except TRACK):
  TARGET_MINUTES      wall clock the training run should take   (default 25)
  DEPTH               transformer layers; the only size dial    (default 6)
  MAX_SEQ_LEN         context length used for training          (default 512)
  DEVICE_BATCH_SIZE   rows per forward/backward                 (default 8)
  TOTAL_BATCH_SIZE    tokens per optimiser step                 (default 16384)
  ITERATIONS          skip calibration and use this step count  (default 0 = calibrate)
  CALIBRATE_STEPS     steps in the calibration run              (default 40)
  NANOCHAT_DTYPE      float32 | bfloat16 | float16, passed through to nanochat
  LOG_DIR             where logs are written                    (default ./part-12-logs)
  FIELDS              where the run configuration is written    (default ./train-fields.json)
USAGE
}

case "$TRACK" in
    spark|strix|mac|nvidia|cpu) ;;
    *) usage; echo; echo "ERROR: set TRACK to one of spark, strix, mac, nvidia, cpu." >&2; exit 2 ;;
esac

# ------------------------------------------------------------- preflight ----
[ -d "$NANOCHAT/.git" ] || { echo "ERROR: no nanochat repository at $NANOCHAT." >&2; exit 1; }
[ -d "$NANOCHAT/.venv" ] || { echo "ERROR: no .venv in $NANOCHAT; run prepare-data.sh first." >&2; exit 1; }
[ -d "$NANOCHAT_BASE_DIR/tokenizer" ] || { echo "ERROR: no tokeniser in $NANOCHAT_BASE_DIR; run prepare-data.sh first." >&2; exit 1; }

# total_batch_size must divide evenly into whole forward/backward passes
TOKENS_PER_FWDBWD=$(( DEVICE_BATCH_SIZE * MAX_SEQ_LEN ))
if [ $(( TOTAL_BATCH_SIZE % TOKENS_PER_FWDBWD )) -ne 0 ]; then
    echo "ERROR: TOTAL_BATCH_SIZE ($TOTAL_BATCH_SIZE) must be a multiple of" >&2
    echo "       DEVICE_BATCH_SIZE x MAX_SEQ_LEN ($TOKENS_PER_FWDBWD). Keep both powers of two." >&2
    exit 2
fi

mkdir -p "$LOG_DIR"
cd "$NANOCHAT"
# shellcheck disable=SC1091
source .venv/bin/activate
COMMIT="$(git rev-parse --short HEAD)"

COMMON_ARGS=(
    --depth="$DEPTH"
    --head-dim="$HEAD_DIM"
    --window-pattern="$WINDOW_PATTERN"
    --max-seq-len="$MAX_SEQ_LEN"
    --device-batch-size="$DEVICE_BATCH_SIZE"
    --total-batch-size="$TOTAL_BATCH_SIZE"
)

# ----------------------------------------------------------- calibration ----
# A short run at exactly the real run's settings, so the throughput it reports is
# the throughput the real run will get. Everything that is not the training step
# is switched off, and the checkpoint it leaves behind is deleted afterwards.
RATE=""
if [ "$ITERATIONS" -le 0 ]; then
    CAL_LOG="$LOG_DIR/calibrate-$MODEL_TAG.log"
    echo "==> calibrating with $CALIBRATE_STEPS steps (this pays the compile cost once)"
    python -m scripts.base_train \
        "${COMMON_ARGS[@]}" \
        --num-iterations="$CALIBRATE_STEPS" \
        --eval-every=-1 \
        --core-metric-every=-1 \
        --sample-every=-1 \
        --model-tag="$CAL_TAG" 2>&1 | tee "$CAL_LOG"

    # Average the last ten step rates: the first steps include warm-up and compile.
    RATE="$(grep -oE 'tok/sec: [0-9,]+' "$CAL_LOG" \
        | tr -d ',' \
        | awk '{print $2}' \
        | tail -n 10 \
        | awk '{s+=$1; n+=1} END {if (n > 0) printf "%d", s/n}')"

    if [ -z "$RATE" ] || [ "$RATE" -le 0 ]; then
        echo "ERROR: could not read a tok/sec figure from $CAL_LOG." >&2
        echo "       Read the log, then rerun with ITERATIONS=<count> to skip calibration." >&2
        exit 1
    fi

    ITERATIONS="$(awk -v m="$TARGET_MINUTES" -v r="$RATE" -v b="$TOTAL_BATCH_SIZE" \
        'BEGIN { printf "%d", (m * 60 * r) / b }')"
    if [ "$ITERATIONS" -lt "$MIN_ITERATIONS" ]; then
        echo "==> calibration suggests $ITERATIONS steps; raising to the $MIN_ITERATIONS-step floor"
        echo "    (the learning-rate warm-up alone is 40 steps, so shorter runs are not meaningful)"
        ITERATIONS="$MIN_ITERATIONS"
    fi

    # The calibration checkpoint has served its purpose; reclaim the disk.
    CAL_DIR="$NANOCHAT_BASE_DIR/base_checkpoints/$CAL_TAG"
    if [ -d "$CAL_DIR" ]; then
        rm -rf "$CAL_DIR"
        echo "==> removed the calibration checkpoint at $CAL_DIR"
    fi

    echo "==> measured throughput and the budget give $ITERATIONS steps"
    echo "    ${TARGET_MINUTES} min x 60 s x ${RATE} tokens/s / ${TOTAL_BATCH_SIZE} tokens per step"
fi

SAVE_EVERY=$(( ITERATIONS / 4 ))
[ "$SAVE_EVERY" -lt 100 ] && SAVE_EVERY=100

# -------------------------------------------------------------- training ----
TRAIN_LOG="$LOG_DIR/train-$MODEL_TAG.log"
echo "==> training $MODEL_TAG for $ITERATIONS steps, checkpointing every $SAVE_EVERY"
python -m scripts.base_train \
    "${COMMON_ARGS[@]}" \
    --num-iterations="$ITERATIONS" \
    --eval-every="$EVAL_EVERY" \
    --eval-tokens="$EVAL_TOKENS" \
    --core-metric-every=-1 \
    --sample-every="$SAMPLE_EVERY" \
    --save-every="$SAVE_EVERY" \
    --model-tag="$MODEL_TAG" 2>&1 | tee "$TRAIN_LOG"

# ------------------------------------------------------------ evaluation ----
# bpb on both splits, the CORE benchmark on a small sample of each task, and the
# fixed sample prompts. The CORE bundle is downloaded on first use.
EVAL_LOG="$LOG_DIR/eval-$MODEL_TAG.log"
echo "==> evaluating $MODEL_TAG"
python -m scripts.base_eval \
    --eval=core,bpb,sample \
    --model-tag="$MODEL_TAG" \
    --device-batch-size=1 \
    --split-tokens="$EVAL_SPLIT_TOKENS" \
    --max-per-task="$EVAL_MAX_PER_TASK" 2>&1 | tee "$EVAL_LOG"

# ------------------------------------------------ configuration for the log --
cat > "$FIELDS" <<JSON
{
  "lab": "part-12/train-a-small-model",
  "track": "$TRACK",
  "commit": "$COMMIT",
  "model_tag": "$MODEL_TAG",
  "dataset": "nanochat default shards in $NANOCHAT_BASE_DIR",
  "depth": $DEPTH,
  "head_dim": $HEAD_DIM,
  "window_pattern": "$WINDOW_PATTERN",
  "max_seq_len": $MAX_SEQ_LEN,
  "device_batch_size": $DEVICE_BATCH_SIZE,
  "total_batch_size": $TOTAL_BATCH_SIZE,
  "num_iterations": $ITERATIONS,
  "save_every": $SAVE_EVERY,
  "target_minutes": $TARGET_MINUTES,
  "calibrated_tokens_per_second": ${RATE:-null},
  "dtype_override": "${NANOCHAT_DTYPE:-auto}",
  "train_log": "$TRAIN_LOG",
  "eval_log": "$EVAL_LOG"
}
JSON

cat <<SUMMARY

==> done
    model tag      $MODEL_TAG
    steps          $ITERATIONS
    checkpoints    $NANOCHAT_BASE_DIR/base_checkpoints/$MODEL_TAG
    training log   $TRAIN_LOG
    eval log       $EVAL_LOG
    configuration  $FIELDS

    Read these lines out of the training log before you go on:
      "Parameter counts:"                  where the parameters actually are
      "Estimated FLOPs per token:"         the six-per-parameter rule, applied
      "Tokens : Scaling params ratio:"     compare with the compute-optimal 12
      "Total training time:"               against your TARGET_MINUTES
      "Minimum validation bpb:"            the number worth comparing between runs

    Next: python sample-and-record.py --model-tag $MODEL_TAG --fields $FIELDS \\
              --labbook labbook.md
SUMMARY
