/*
 * inference.h — OmniPen CNN forward pass
 * ========================================
 * Pure C++ implementation of the trained neural network.
 * No external ML library required — works on any Arduino board.
 *
 * Network architecture (matches train_model.py):
 *   Input:   28×28 float image, values 0.0–1.0
 *   Conv1:   8 filters, 3×3, same padding, ReLU → 28×28×8
 *   Pool1:   MaxPool 2×2 → 14×14×8
 *   Conv2:   16 filters, 3×3, same padding, ReLU → 14×14×16
 *   Pool2:   MaxPool 2×2 → 7×7×16
 *   Conv3:   32 filters, 3×3, same padding, ReLU → 7×7×32
 *   Flatten: → 1568
 *   Dense1:  32 units, ReLU → 32
 *   Dense2:  26 units, Softmax → 26 (one probability per letter A–Z)
 *
 * The weights are in weights.h, which is generated by train_model.py and
 * must be copied into this sketch folder before compiling.
 */

#pragma once
#include <math.h>
#include "weights.h"

// ── Activation buffer sizes ───────────────────────────────────────────────────
// The two ping-pong buffers must hold the largest activation in the network.
// Largest: Conv1 output = 28×28×8 = 6 272 floats.
#define ACT_BUF_SIZE 6272


// ── Conv2D with ReLU, "same" padding ─────────────────────────────────────────
/*
 * Computes one Conv2D layer followed by ReLU activation.
 *
 * Parameters
 *   in      : input  tensor, layout [H][W][Cin]  (row-major, channel-last)
 *   H, W    : spatial size of the input
 *   Cin     : number of input channels
 *   W_kern  : weight array, layout [Cout][Cin][kH][kW]
 *   bias    : bias array,   layout [Cout]
 *   Cout    : number of output filters
 *   K       : kernel size (assumed square, e.g. 3 for 3×3)
 *   out     : output tensor, layout [H][W][Cout]  (same spatial size, "same" padding)
 *
 * "Same" padding means the output has the same H×W as the input.
 * Pixels outside the border are treated as zero (zero-padding).
 *
 * Weight layout [Cout][Cin][kH][kW] differs from Keras ([kH][kW][Cin][Cout]).
 * train_model.py transposes the weights during export so they match this layout.
 */
static void conv2d_relu(
    const float* in, int H, int W, int Cin,
    const float* W_kern, const float* bias,
    int Cout, int K,
    float* out)
{
    int pad = K / 2;
    for (int oh = 0; oh < H; oh++) {
        for (int ow = 0; ow < W; ow++) {
            for (int oc = 0; oc < Cout; oc++) {
                float sum = bias[oc];
                for (int ic = 0; ic < Cin; ic++) {
                    for (int kh = 0; kh < K; kh++) {
                        int ih = oh - pad + kh;
                        if (ih < 0 || ih >= H) continue;
                        for (int kw = 0; kw < K; kw++) {
                            int iw = ow - pad + kw;
                            if (iw < 0 || iw >= W) continue;
                            float px = in[(ih * W + iw) * Cin + ic];
                            float wt = W_kern[((oc * Cin + ic) * K + kh) * K + kw];
                            sum += px * wt;
                        }
                    }
                }
                // ReLU: clamp negative values to zero
                out[(oh * W + ow) * Cout + oc] = (sum > 0.0f) ? sum : 0.0f;
            }
        }
    }
}


// ── MaxPool2D, 2×2 window, stride 2 ──────────────────────────────────────────
/*
 * Slides a 2×2 window over the spatial dimensions with stride 2 and keeps
 * the maximum value in each window.  Output size: (H/2) × (W/2) × C.
 *
 * This halves the spatial resolution, which reduces computation in later
 * layers and makes the network somewhat translation-invariant.
 */
static void maxpool2d(
    const float* in, int H, int W, int C,
    float* out)
{
    int oH = H / 2, oW = W / 2;
    for (int oh = 0; oh < oH; oh++) {
        for (int ow = 0; ow < oW; ow++) {
            for (int c = 0; c < C; c++) {
                float mx = -1e30f;
                for (int ph = 0; ph < 2; ph++) {
                    for (int pw = 0; pw < 2; pw++) {
                        float v = in[((oh*2+ph)*W + (ow*2+pw))*C + c];
                        if (v > mx) mx = v;
                    }
                }
                out[(oh * oW + ow) * C + c] = mx;
            }
        }
    }
}


// ── Fully-connected (Dense) layer with ReLU ───────────────────────────────────
/*
 * Each output neuron computes:  out[o] = ReLU( bias[o] + sum_i( W[o][i] * in[i] ) )
 *
 * Weight layout [Nout][Nin]: row o contains the weights going INTO output neuron o.
 * train_model.py transposes the Keras weight matrix ([Nin][Nout]) during export.
 */
static void dense_relu(
    const float* in, int Nin,
    const float* W_fc, const float* bias,
    int Nout,
    float* out)
{
    for (int o = 0; o < Nout; o++) {
        float sum = bias[o];
        const float* row = W_fc + (long)o * Nin;
        for (int i = 0; i < Nin; i++) sum += row[i] * in[i];
        out[o] = (sum > 0.0f) ? sum : 0.0f;
    }
}


// ── Fully-connected layer with Softmax ────────────────────────────────────────
/*
 * Same matrix multiply as dense_relu but the activation is Softmax instead
 * of ReLU.  Softmax converts raw scores into probabilities that sum to 1.0.
 *
 * Numerically stable form:
 *   1. Subtract the maximum logit before exp() to prevent overflow.
 *   2. Divide each exp() by the sum of all exp() values.
 */
static void dense_softmax(
    const float* in, int Nin,
    const float* W_fc, const float* bias,
    int Nout,
    float* out)
{
    // Compute logits
    float max_logit = -1e30f;
    for (int o = 0; o < Nout; o++) {
        float sum = bias[o];
        const float* row = W_fc + (long)o * Nin;
        for (int i = 0; i < Nin; i++) sum += row[i] * in[i];
        out[o] = sum;
        if (sum > max_logit) max_logit = sum;
    }
    // Softmax
    float sum_exp = 0.0f;
    for (int o = 0; o < Nout; o++) {
        out[o] = expf(out[o] - max_logit);
        sum_exp += out[o];
    }
    for (int o = 0; o < Nout; o++) out[o] /= sum_exp;
}


// ── Full forward pass ─────────────────────────────────────────────────────────
/*
 * Runs the complete CNN on a 28×28 float image and returns 26 probabilities.
 *
 *   image  : 2-D array [28][28], pixel values 0.0 (background) to 1.0 (ink)
 *   probs  : output array [26], probs[0]=P(A) … probs[25]=P(Z)
 *
 * Two static ping-pong buffers (buf_a and buf_b) hold the activations.
 * 'static' puts them in BSS (global RAM), not on the stack, which avoids
 * stack overflow from the large (~25 KB each) allocations.
 *
 * Typical execution time on XIAO nRF52840 @ 64 MHz: 100–300 ms.
 */
static void run_inference(float image[28][28], float probs[26])
{
    static float buf_a[ACT_BUF_SIZE];
    static float buf_b[ACT_BUF_SIZE];

    // Flatten the 28×28 image into buf_a as [28×28×1] (single channel)
    for (int r = 0; r < 28; r++)
        for (int c = 0; c < 28; c++)
            buf_a[r * 28 + c] = image[r][c];

    // Conv1: [28,28,1] → [28,28,8]
    conv2d_relu(buf_a, 28, 28, 1,  CONV1_W, CONV1_B,  8, 3, buf_b);
    // Pool1: [28,28,8] → [14,14,8]
    maxpool2d  (buf_b, 28, 28, 8,  buf_a);

    // Conv2: [14,14,8] → [14,14,16]
    conv2d_relu(buf_a, 14, 14, 8,  CONV2_W, CONV2_B, 16, 3, buf_b);
    // Pool2: [14,14,16] → [7,7,16]
    maxpool2d  (buf_b, 14, 14, 16, buf_a);

    // Conv3: [7,7,16] → [7,7,32]
    conv2d_relu(buf_a,  7,  7, 16, CONV3_W, CONV3_B, 32, 3, buf_b);
    // Flatten: 7×7×32 = 1568 values already contiguous in buf_b

    // Dense1: 1568 → 32, ReLU
    dense_relu   (buf_b, 1568, DENSE1_W, DENSE1_B, 32, buf_a);
    // Dense2: 32 → 26, Softmax → write directly into probs
    dense_softmax(buf_a,   32, DENSE2_W, DENSE2_B, 26, probs);
}
