"""
OmniPen Training Script
=======================
Trains a compact CNN on the EMNIST Letters dataset and exports it as a
C header file that you embed in your Arduino firmware.

Run this on your computer (NOT on the microcontroller).

What happens, step by step:
  1. Download the EMNIST/Letters dataset (~145 k handwritten letter images)
  2. Build a small CNN that can fit in ~80 KB after compression
  3. Train the CNN (~15 epochs, a few minutes on CPU or ~1 min on GPU)
  4. Convert the trained model to TensorFlow Lite with int8 quantisation
     so it runs efficiently on the Cortex-M4 microcontroller
  5. Write the model as a C byte-array header (model_data.h) that you
     copy into your Arduino sketch folder

A note on EMNIST vs. real pen motion
─────────────────────────────────────
EMNIST contains images of letters drawn on a 2-D canvas. Your pen produces
6-axis IMU time series. The firmware bridges the gap by integrating the IMU
data to reconstruct a 2-D trajectory, then rendering that trajectory as a
28×28 greyscale image — the same format EMNIST uses. This means the CNN
trained here can classify the pen's trajectory directly.

The accuracy you get out of the box will be decent but not perfect, because
real pen-in-air strokes look slightly different from EMNIST pen-on-paper
strokes. See the "Going further" section at the bottom of this file for how
to fine-tune with your own collected data.

Requirements (install once)
───────────────────────────
  pip install tensorflow tensorflow-datasets numpy

Usage
─────
  python train_model.py

  Outputs:
    omnipen_model.tflite   — the quantised model (for reference / PC testing)
    model_data.h           — copy this file into firmware/omnipen/
"""

import os
import numpy as np
import tensorflow as tf
import tensorflow_datasets as tfds

# Suppress TensorFlow info/warning noise in the terminal
os.environ["TF_CPP_MIN_LOG_LEVEL"] = "2"


# ── Configuration ─────────────────────────────────────────────────────────────
# You rarely need to change these unless you want to experiment.

EPOCHS      = 15         # Max training passes over the full dataset.
                         # EarlyStopping will usually stop before this.
BATCH_SIZE  = 128        # Images processed per gradient-descent step.
HEADER_FILE = "weights.h"    # copied to firmware/omnipen/weights.h

# The 26 output classes in order (index 0 = 'A', index 25 = 'Z')
LETTERS = "ABCDEFGHIJKLMNOPQRSTUVWXYZ"


# ── Step 1: Load EMNIST Letters ───────────────────────────────────────────────

def load_data():
    """
    Download and return the EMNIST/Letters dataset as TF datasets.

    EMNIST (Extended MNIST) extends the classic MNIST digit dataset to
    include 145,600 handwritten letter images:
      • Each image is 28×28 pixels, single channel (greyscale)
      • Labels are integers 1–26 mapping to A–Z

    The dataset ships stored in a transposed orientation — we fix that
    during preprocessing below.
    """
    print("Loading EMNIST/Letters dataset …")
    print("(First run downloads ~535 MB; subsequent runs use the cache.)\n")

    (ds_train, ds_test), info = tfds.load(
        "emnist/letters",
        split=["train", "test"],
        as_supervised=True,    # Yields (image, label) pairs
        with_info=True,
    )

    n_train = info.splits["train"].num_examples
    n_test  = info.splits["test"].num_examples
    print(f"  Training images : {n_train:,}")
    print(f"  Test images     : {n_test:,}\n")

    return ds_train, ds_test


def preprocess(image, label):
    """
    Prepare a single (image, label) pair for training.

    Three transformations are applied:

    1. Fix EMNIST orientation
       EMNIST images are stored with height and width swapped relative to
       what you'd expect. Transposing axes 0 and 1 (height ↔ width)
       corrects this so letters appear upright.

    2. Normalise pixel values
       The raw images are uint8 in [0, 255]. Dividing by 255 brings them
       into [0.0, 1.0] as float32. Neural networks train more stably on
       small, normalised values.

    3. Shift label range
       EMNIST labels are 1–26 (A=1, Z=26). Subtracting 1 gives 0–25,
       which is what our 26-output softmax layer expects.
    """
    image = tf.transpose(image, perm=[1, 0, 2])   # fix EMNIST rotation
    image = tf.cast(image, tf.float32) / 255.0    # normalise to [0, 1]
    label = label - 1                              # shift to 0-based index
    return image, label


def make_pipeline(ds, *, training: bool):
    """
    Apply preprocessing, optional shuffling, batching, and prefetching.

    Shuffling during training prevents the model from memorising the order
    of examples (which would look like good training accuracy but fail on
    new data).

    Prefetching prepares the next batch in the background while the GPU/CPU
    is training on the current one, reducing idle time.
    """
    ds = ds.map(preprocess, num_parallel_calls=tf.data.AUTOTUNE)
    if training:
        ds = ds.shuffle(buffer_size=10_000)
    ds = ds.batch(BATCH_SIZE)
    ds = ds.prefetch(tf.data.AUTOTUNE)
    return ds


# ── Step 2: Build the model ───────────────────────────────────────────────────

def build_model() -> tf.keras.Model:
    """
    Build a compact CNN sized to fit in the XIAO nRF52840's 1 MB Flash.

    Architecture walkthrough
    ────────────────────────
    Input (28×28×1)
      │
      ├─ Conv2D(8 filters, 3×3) + ReLU
      │    Learns 8 simple patterns: horizontal edges, diagonal strokes, dots …
      │    Output: 28×28×8
      │
      ├─ MaxPool2D(2×2)   →  14×14×8
      │    Keeps the strongest response in each 2×2 block, halves spatial size.
      │
      ├─ Conv2D(16 filters, 3×3) + ReLU
      │    Combines simple patterns into curved strokes, corners, arcs …
      │    Output: 14×14×16
      │
      ├─ MaxPool2D(2×2)   →  7×7×16
      │
      ├─ Conv2D(32 filters, 3×3) + ReLU
      │    Detects letter-specific shapes by combining the previous features.
      │    Output: 7×7×32
      │
      ├─ Flatten   →  1 568 values
      │
      ├─ Dropout(0.25)
      │    During training, randomly zeros 25 % of activations each step.
      │    Forces the network to learn redundant representations, which
      │    helps it generalise to new handwriting styles.
      │
      ├─ Dense(32) + ReLU
      │    A 32-neuron summary layer.  32 instead of 64 keeps the float
      │    weight array at ~200 KB — comfortable within the 1 MB Flash.
      │
      └─ Dense(26) + Softmax
           One score per letter. Softmax turns raw scores into probabilities
           that sum to 1.0, so we can read off the confidence of each guess.

    Why export float weights instead of a TFLite model?
    ─────────────────────────────────────────────────────
    Every available TFLite Micro Arduino library (Arduino_TensorFlowLite,
    tflm_cortexm, EloquentTinyML) depends on the Arduino mbed core and
    fails to compile against the Seeed/Adafruit nRF52 core used by XIAO.
    Exporting plain float arrays and implementing the forward pass in C++
    requires NO external library and compiles on any board.
    Float inference on Cortex-M4 @ 64 MHz is fast enough (~100 ms per letter).
    """
    model = tf.keras.Sequential(
        [
            tf.keras.layers.Input(shape=(28, 28, 1)),

            # Random rotation: active only during training.
            # The firmware projects the trajectory onto a plane whose azimuth
            # is arbitrary, so the rendered letter may be rotated by any angle.
            # Training with rotation makes the model tolerant of this.
            tf.keras.layers.RandomRotation(
                factor=0.12,
                fill_mode="constant",
                fill_value=0.0,
            ),

            tf.keras.layers.Conv2D(8,  3, padding="same", activation="relu"),
            tf.keras.layers.MaxPooling2D(2),

            tf.keras.layers.Conv2D(16, 3, padding="same", activation="relu"),
            tf.keras.layers.MaxPooling2D(2),

            tf.keras.layers.Conv2D(32, 3, padding="same", activation="relu"),

            tf.keras.layers.Flatten(),
            tf.keras.layers.Dropout(0.25),
            tf.keras.layers.Dense(32, activation="relu"),   # 32, not 64 — keeps weights ~200 KB
            tf.keras.layers.Dense(26, activation="softmax"),
        ],
        name="omnipen_cnn",
    )

    model.compile(
        optimizer="adam",
        # sparse_categorical_crossentropy expects integer labels (0–25),
        # not one-hot vectors — saves a conversion step.
        loss="sparse_categorical_crossentropy",
        metrics=["accuracy"],
    )

    model.summary()
    return model


# ── Step 3: Train ─────────────────────────────────────────────────────────────

def train(model: tf.keras.Model, ds_train, ds_test):
    """
    Fit the model and stop early if validation accuracy plateaus.

    EarlyStopping monitors validation accuracy after each epoch. If it
    does not improve for `patience` consecutive epochs, training stops and
    the best weights are restored. This prevents wasting time and avoids
    overfitting (the model starting to memorise rather than generalise).
    """
    print(f"\nTraining for up to {EPOCHS} epochs …")
    print("(Each epoch processes all ~112 k training images.)\n")

    callbacks = [
        tf.keras.callbacks.EarlyStopping(
            monitor="val_accuracy",
            patience=3,                  # Stop after 3 non-improving epochs
            restore_best_weights=True,   # Rewind to the best checkpoint
            verbose=1,
        )
    ]

    model.fit(
        ds_train,
        epochs=EPOCHS,
        validation_data=ds_test,
        callbacks=callbacks,
    )

    loss, acc = model.evaluate(ds_test, verbose=0)
    print(f"\nFinal test accuracy : {acc:.2%}")
    print(f"Final test loss     : {loss:.4f}")


# ── Step 4: Export weights as a C++ header ───────────────────────────────────

def export_weights(model: tf.keras.Model, path: str = HEADER_FILE):
    """
    Walk every layer, extract its float32 weights, and write them as C arrays.

    Why float32 instead of a TFLite binary?
    ─────────────────────────────────────────
    All TFLite Micro Arduino libraries depend on the Arduino mbed core and
    fail to compile against the Seeed/Adafruit nRF52 core.  Exporting plain
    float arrays and implementing the forward pass in C++ (see inference.h)
    needs no external library and compiles on any board.

    Weight layout conversions
    ─────────────────────────
    Keras stores Conv2D kernels as  [kH, kW, Cin, Cout].
    Our C++ inference code expects  [Cout, Cin, kH, kW]  (output-filter-first
    so the inner loops over Cin/kH/kW are contiguous for each output filter).

    Keras stores Dense kernels as   [Nin, Nout].
    Our C++ code expects            [Nout, Nin]  (row = one output neuron).
    """
    print(f"\nExporting weights to {path!r} …")

    lines = [
        "/*",
        " * OmniPen weights — auto-generated by train_model.py, do NOT edit.",
        " * Included by inference.h.",
        " */",
        "#pragma once",
        "",
    ]

    conv_idx  = 1
    dense_idx = 1
    total_bytes = 0

    for layer in model.layers:
        wts = layer.get_weights()
        if not wts:
            continue

        lname = layer.name.lower()

        if "conv" in lname:
            # kernel shape: [kH, kW, Cin, Cout] → transpose to [Cout, Cin, kH, kW]
            kernel = np.transpose(wts[0], (3, 2, 0, 1)).flatten()
            bias   = wts[1].flatten()
            tag    = f"CONV{conv_idx}"
            lines += [
                f"// {layer.name}: kernel {wts[0].shape} → [{tag}]",
                f"const float {tag}_W[{len(kernel)}] = {{",
                "  " + ", ".join(f"{v:.7g}f" for v in kernel),
                "};",
                f"const float {tag}_B[{len(bias)}] = {{",
                "  " + ", ".join(f"{v:.7g}f" for v in bias),
                "};",
                "",
            ]
            total_bytes += (len(kernel) + len(bias)) * 4
            conv_idx += 1

        elif "dense" in lname:
            # kernel shape: [Nin, Nout] → transpose to [Nout, Nin]
            kernel = np.transpose(wts[0]).flatten()
            bias   = wts[1].flatten()
            tag    = f"DENSE{dense_idx}"
            lines += [
                f"// {layer.name}: kernel {wts[0].shape} → [{tag}]",
                f"const float {tag}_W[{len(kernel)}] = {{",
                "  " + ", ".join(f"{v:.7g}f" for v in kernel),
                "};",
                f"const float {tag}_B[{len(bias)}] = {{",
                "  " + ", ".join(f"{v:.7g}f" for v in bias),
                "};",
                "",
            ]
            total_bytes += (len(kernel) + len(bias)) * 4
            dense_idx += 1

    with open(path, "w") as f:
        f.write("\n".join(lines))

    print(f"  {conv_idx-1} conv layers, {dense_idx-1} dense layers exported.")
    print(f"  Total weight size: {total_bytes/1024:.1f} KB")
    print(f"  Copy {path} → firmware/omnipen/weights.h")


# ── Entry point ───────────────────────────────────────────────────────────────

if __name__ == "__main__":
    print("=" * 56)
    print("  OmniPen — Training Script")
    print("=" * 56)

    ds_train_raw, ds_test_raw = load_data()
    ds_train = make_pipeline(ds_train_raw, training=True)
    ds_test  = make_pipeline(ds_test_raw,  training=False)

    model = build_model()
    train(model, ds_train, ds_test)

    export_weights(model, HEADER_FILE)

    print("\n" + "=" * 56)
    print("  Done!")
    print("  Next steps:")
    print("    1. Copy weights.h  →  firmware/omnipen/weights.h")
    print("    2. Flash omnipen.ino to your XIAO nRF52840 Sense")
    print("    3. Open Serial Monitor at 115200 baud")
    print("    4. Write a letter in the air; watch it get recognised!")
    print("=" * 56)


# ── Going further ─────────────────────────────────────────────────────────────
#
# The model trained here has never seen real pen-in-air IMU trajectories —
# only drawn images. If recognition accuracy is low for certain letters,
# the most effective fix is to collect real data from the pen and fine-tune:
#
#   1. Add a data-collection mode to the firmware that saves IMU samples
#      over Serial (or BLE) while you write each letter many times.
#   2. Feed that data through the same trajectory-to-image pipeline used
#      in the firmware (recreated in Python), producing real 28×28 images.
#   3. Mix those real images with the EMNIST data (or replace EMNIST for
#      the fine-tuning phase) and re-run the training with a lower
#      learning rate (e.g., change optimizer="adam" to
#      optimizer=tf.keras.optimizers.Adam(1e-4)).
#
# Even 50–100 samples per letter is usually enough to noticeably improve
# accuracy on your personal handwriting style.
