#include "cct_runtime.h"
#include <stdio.h>
#include <stdint.h>
#include <stdbool.h>
#include <math.h>

typedef struct { bool data[20]; } feature_vec;
typedef struct { cct_prob_t data[20]; } digit_model;
digit_model models[10];
feature_vec extract_features(uint8_t image[28][28]) {
    uint8_t down[14][14];
    for (int i = 0; i <= 13; i++) {
        for (int j = 0; j <= 13; j++) {
            (down[i][j] = ((((image[(2 * i)][(2 * j)] + image[(2 * i)][((2 * j) + 1)]) + image[((2 * i) + 1)][(2 * j)]) + image[((2 * i) + 1)][((2 * j) + 1)]) > 2));
        }
    }
    feature_vec f = {0};
    int edges = 0;
    for (int i = 0; i <= 6; i++) {
        for (int j = 0; j <= 12; j++) {
            if ((down[i][j] != down[i][(j + 1)])) {
                (edges++);
            }
        }
    }
    (f.data[0] = (edges > 20));
    return f;
}

void train(uint8_t images[][28][28], uint8_t labels[], int count) {
    for (int d = 0; d <= 9; d++) {
        for (int f = 0; f <= 19; f++) {
            (models[d].data[f] = 0.5);
        }
    }
    for (int epoch = 0; epoch <= 4; epoch++) {
        for (int idx = 0; idx <= (count - 1); idx++) {
            feature_vec f = extract_features(images[idx]);
            uint8_t digit = labels[idx];
            for (int feat = 0; feat <= 19; feat++) {
                cct_prob_t p = models[digit].data[feat];
                if ((f.data[feat] == 1)) {
                    (p = (p * 1.1));
                } else {
                    (p = (p * 0.9));
                }
                if ((p > 0.99)) {
                    (p = 0.99);
                }
                if ((p < 0.01)) {
                    (p = 0.01);
                }
                (models[digit].data[feat] = p);
            }
        }
    }
}

uint8_t classify(uint8_t image[28][28]) {
    feature_vec f = extract_features(image);
    cct_prob_t scores[10];
    for (int d = 0; d <= 9; d++) {
        cct_prob_t logp = 0.0;
        for (int feat = 0; feat <= 19; feat++) {
            cct_prob_t p = models[d].data[feat];
            if ((f.data[feat] == 1)) {
                (logp = cct_addp(logp, cct_logp(p)));
            } else {
                (logp = cct_addp(logp, cct_logp((1 - p))));
            }
        }
        (scores[d] = cct_exp(logp));
    }
    uint8_t best = 0;
    cct_prob_t bestp = 0.0;
    for (int d = 0; d <= 9; d++) {
        if ((scores[d] > bestp)) {
            (bestp = scores[d]);
            (best = d);
        }
    }
    return best;
}

int main() {
    uint8_t train_images[1000][28][28];
    uint8_t train_labels[1000];
    uint8_t test_images[200][28][28];
    uint8_t test_labels[200];
    for (int i = 0; i <= 999; i++) {
        (train_labels[i] = (i % 10));
    }
    for (int i = 0; i <= 199; i++) {
        (test_labels[i] = (i % 10));
    }
    train(train_images, train_labels, 1000);
    int correct = 0;
    for (int i = 0; i <= 199; i++) {
        uint8_t pred = classify(test_images[i]);
        if ((pred == test_labels[i])) {
            (correct++);
        }
    }
    float acc = (correct / 2.0);
    cct_print("Test accuracy: %g%%", acc);
    uint32_t t = cct_start_timer();
    uint8_t dummy = classify(test_images[0]);
    uint32_t us = cct_stop_timer(t);
    cct_print("Inference time: %g µs", us);
    return 0;
}

cct_prob_t logp(cct_prob_t p) {
    return cct_log(p);
}

cct_prob_t cct_user_exp(cct_prob_t x) {
    return cct_exp(x);
}
