SpaceQuester

Untitled

Mar 29th, 2026
72
0
Never
Not a member of Pastebin yet? Sign Up, it unlocks many cool features!
C 55.87 KB | None | 0 0
  1. #define _CRT_SECURE_NO_WARNINGS
  2. #include <stdio.h>
  3. #include <stdlib.h>
  4. #include <string.h>
  5. #include <math.h>
  6. #include <time.h>
  7. #ifdef _WIN32
  8. #include <direct.h>
  9. #define GETCWD _getcwd
  10. #else
  11. #include <unistd.h>
  12. #define GETCWD getcwd
  13. #endif
  14.  
  15. #ifndef M_PI
  16. #define M_PI 3.14159265358979323846
  17. #endif
  18.  
  19.  
  20. #define IMG_ROWS 28
  21. #define IMG_COLS 28
  22. #define INPUT_DIM 784
  23. #define HIDDEN_DIM 128
  24. #define N_CLASSES 10
  25. #define N_QUBITS 5
  26. #define STATE_SIZE (1 << N_QUBITS)
  27.  
  28. #define TRAIN_SUBSET 8000
  29. #define TEST_SUBSET 10000
  30. #define BATCH_SIZE 32
  31. #define EPOCHS 140
  32. #define XI 0.8
  33.  
  34. #define LEARNING_RATE 0.001
  35. #define BETA1 0.90
  36. #define BETA2 0.99
  37. #define WEIGHT_DECAY 0.0
  38. #define EPS_ADAM 1e-8
  39.  
  40. /* Stage-2 implementation parameters: required for discrete simulation of time */
  41. #define DEFAULT_TIME_STEPS 32
  42. #define DT 0.01
  43. #define V_REST 0.0
  44. #define V_RESET 0.0
  45. #define V_THRESHOLD 1.0
  46. #define TAU_M 1.0
  47.  
  48. #define Q_VAR_PARAM_COUNT ((N_QUBITS - 1) * 3)
  49. #define Q_HEAD_IN N_QUBITS
  50.  
  51. typedef struct {
  52.     int num_images;
  53.     int rows;
  54.     int cols;
  55.     unsigned char *images;
  56.     unsigned char *labels;
  57. } MnistSet;
  58.  
  59. typedef struct {
  60.     double real;
  61.     double imag;
  62. } Complex;
  63.  
  64. typedef struct {
  65.     int true_label;
  66.     int pred_label;
  67.     double confidence;
  68.     double loss;
  69. } EvalInfo;
  70.  
  71. static double fc1_w[HIDDEN_DIM][INPUT_DIM];
  72. static double fc1_b[HIDDEN_DIM];
  73.  
  74. static double qc_w[N_CLASSES][HIDDEN_DIM];
  75. static double qc_b[N_CLASSES];
  76.  
  77. static double qq_w[N_CLASSES][Q_HEAD_IN];
  78. static double qq_b[N_CLASSES];
  79.  
  80. /* trainable quantum variation parameters: for each i=1..4, RZ RX RZ */
  81. static double q_theta_var[Q_VAR_PARAM_COUNT];
  82.  
  83. /* Adam state */
  84. static double m_fc1_w[HIDDEN_DIM][INPUT_DIM];
  85. static double v_fc1_w[HIDDEN_DIM][INPUT_DIM];
  86. static double m_fc1_b[HIDDEN_DIM];
  87. static double v_fc1_b[HIDDEN_DIM];
  88.  
  89. static double m_qc_w[N_CLASSES][HIDDEN_DIM];
  90. static double v_qc_w[N_CLASSES][HIDDEN_DIM];
  91. static double m_qc_b[N_CLASSES];
  92. static double v_qc_b[N_CLASSES];
  93.  
  94. static double m_qq_w[N_CLASSES][Q_HEAD_IN];
  95. static double v_qq_w[N_CLASSES][Q_HEAD_IN];
  96. static double m_qq_b[N_CLASSES];
  97. static double v_qq_b[N_CLASSES];
  98.  
  99. static double m_qtheta[Q_VAR_PARAM_COUNT];
  100. static double v_qtheta[Q_VAR_PARAM_COUNT];
  101.  
  102. /* gradients */
  103. static double g_fc1_w[HIDDEN_DIM][INPUT_DIM];
  104. static double g_fc1_b[HIDDEN_DIM];
  105. static double g_qc_w[N_CLASSES][HIDDEN_DIM];
  106. static double g_qc_b[N_CLASSES];
  107. static double g_qq_w[N_CLASSES][Q_HEAD_IN];
  108. static double g_qq_b[N_CLASSES];
  109. static double g_qtheta[Q_VAR_PARAM_COUNT];
  110.  
  111. static double qpre_w[N_QUBITS][HIDDEN_DIM];
  112. static double qpre_b[N_QUBITS];
  113.  
  114. static double m_qpre_w[N_QUBITS][HIDDEN_DIM];
  115. static double v_qpre_w[N_QUBITS][HIDDEN_DIM];
  116. static double m_qpre_b[N_QUBITS];
  117. static double v_qpre_b[N_QUBITS];
  118.  
  119. static double g_qpre_w[N_QUBITS][HIDDEN_DIM];
  120. static double g_qpre_b[N_QUBITS];
  121.  
  122. static Complex q_state[STATE_SIZE];
  123.  
  124. static unsigned int rng_state = 1u;
  125. static int g_time_steps = DEFAULT_TIME_STEPS;
  126. static int g_model_initialized = 0;
  127.  
  128. static int read_be_int(FILE *f) {
  129.     unsigned char b[4];
  130.     if (fread(b, 1, 4, f) != 4) return -1;
  131.     return ((int)b[0] << 24) | ((int)b[1] << 16) | ((int)b[2] << 8) | (int)b[3];
  132. }
  133.  
  134. static double urand01(void) {
  135.     rng_state = 1664525u * rng_state + 1013904223u;
  136.     return (double)(rng_state & 0xFFFFFFFFu) / 4294967295.0;
  137. }
  138.  
  139. static double rand_uniform(double a, double b) {
  140.     return a + (b - a) * urand01();
  141. }
  142.  
  143. static void shuffle_int(int *arr, int n) {
  144.     for (int i = n - 1; i > 0; --i) {
  145.         int j = (int)(urand01() * (i + 1));
  146.         int tmp = arr[i];
  147.         arr[i] = arr[j];
  148.         arr[j] = tmp;
  149.     }
  150. }
  151.  
  152. static int file_exists(const char *path) {
  153.     FILE *f = fopen(path, "rb");
  154.     if (f) {
  155.         fclose(f);
  156.         return 1;
  157.     }
  158.     return 0;
  159. }
  160.  
  161. static void join_path(char *out, size_t out_sz, const char *a, const char *b) {
  162.     size_t la = strlen(a);
  163. #ifdef _WIN32
  164.     const char sep = '\\';
  165. #else
  166.     const char sep = '/';
  167. #endif
  168.  
  169.     if (la > 0 && (a[la - 1] == '\\' || a[la - 1] == '/')) {
  170.         snprintf(out, out_sz, "%s%s", a, b);
  171.     } else {
  172.         snprintf(out, out_sz, "%s%c%s", a, sep, b);
  173.     }
  174. }
  175.  
  176. static int dirname_of(char *path) {
  177.     size_t n = strlen(path);
  178.     if (n == 0) return 0;
  179.     while (n > 0 && path[n - 1] != '\\' && path[n - 1] != '/') --n;
  180.     if (n == 0) return 0;
  181.     path[n - 1] = '\0';
  182.     return 1;
  183. }
  184.  
  185. static int find_mnist_files(char *img_tr, size_t sz1,
  186.                             char *lbl_tr, size_t sz2,
  187.                             char *img_te, size_t sz3,
  188.                             char *lbl_te, size_t sz4) {
  189.     const char *names[4] = {
  190.         "train-images-idx3-ubyte",
  191.         "train-labels-idx1-ubyte",
  192.         "t10k-images-idx3-ubyte",
  193.         "t10k-labels-idx1-ubyte"
  194.     };
  195.  
  196.     char cwd[1024];
  197.     if (!GETCWD(cwd, sizeof(cwd))) {
  198.         return 0;
  199.     }
  200.  
  201.     char dirs[12][1024];
  202.     int dir_count = 0;
  203.     strncpy(dirs[dir_count++], cwd, sizeof(dirs[0]) - 1);
  204.     dirs[dir_count - 1][sizeof(dirs[0]) - 1] = '\0';
  205.  
  206.     if (dir_count < 12) {
  207.         char ds[1024];
  208.         join_path(ds, sizeof(ds), cwd, "Datasets");
  209.         strncpy(dirs[dir_count++], ds, sizeof(dirs[0]) - 1);
  210.         dirs[dir_count - 1][sizeof(dirs[0]) - 1] = '\0';
  211.     }
  212.  
  213.     char exe_dir[1024];
  214.     strncpy(exe_dir, cwd, sizeof(exe_dir) - 1);
  215.     exe_dir[sizeof(exe_dir) - 1] = '\0';
  216.     if (dirname_of(exe_dir)) {
  217.         strncpy(dirs[dir_count++], exe_dir, sizeof(dirs[0]) - 1);
  218.         dirs[dir_count - 1][sizeof(dirs[0]) - 1] = '\0';
  219.     }
  220.  
  221.     /* walk up several parents from cwd */
  222.     char cur[1024];
  223.     strncpy(cur, cwd, sizeof(cur) - 1);
  224.     cur[sizeof(cur) - 1] = '\0';
  225.     for (int k = 0; k < 6; ++k) {
  226.         if (!dirname_of(cur)) break;
  227.         strncpy(dirs[dir_count++], cur, sizeof(dirs[0]) - 1);
  228.         dirs[dir_count - 1][sizeof(dirs[0]) - 1] = '\0';
  229.     }
  230.  
  231. #ifdef _WIN32
  232.     /* also try x64\\Debug and x64\\Release below cwd and parents */
  233.     int base_count = dir_count;
  234.     for (int i = 0; i < base_count && dir_count < 12; ++i) {
  235.         char tmp[1024];
  236.         join_path(tmp, sizeof(tmp), dirs[i], "x64\\Debug");
  237.         strncpy(dirs[dir_count++], tmp, sizeof(dirs[0]) - 1);
  238.         dirs[dir_count - 1][sizeof(dirs[0]) - 1] = '\0';
  239.         if (dir_count >= 12) break;
  240.         join_path(tmp, sizeof(tmp), dirs[i], "x64\\Release");
  241.         strncpy(dirs[dir_count++], tmp, sizeof(dirs[0]) - 1);
  242.         dirs[dir_count - 1][sizeof(dirs[0]) - 1] = '\0';
  243.     }
  244. #endif
  245.  
  246.     printf("Searching MNIST files in:\n");
  247.     for (int i = 0; i < dir_count; ++i) {
  248.         printf("  %s\n", dirs[i]);
  249.     }
  250.     printf("\n");
  251.  
  252.     for (int i = 0; i < dir_count; ++i) {
  253.         char p1[1024], p2[1024], p3[1024], p4[1024];
  254.         join_path(p1, sizeof(p1), dirs[i], names[0]);
  255.         join_path(p2, sizeof(p2), dirs[i], names[1]);
  256.         join_path(p3, sizeof(p3), dirs[i], names[2]);
  257.         join_path(p4, sizeof(p4), dirs[i], names[3]);
  258.         if (file_exists(p1) && file_exists(p2) && file_exists(p3) && file_exists(p4)) {
  259.             strncpy(img_tr, p1, sz1 - 1); img_tr[sz1 - 1] = '\0';
  260.             strncpy(lbl_tr, p2, sz2 - 1); lbl_tr[sz2 - 1] = '\0';
  261.             strncpy(img_te, p3, sz3 - 1); img_te[sz3 - 1] = '\0';
  262.             strncpy(lbl_te, p4, sz4 - 1); lbl_te[sz4 - 1] = '\0';
  263.             return 1;
  264.         }
  265.     }
  266.     return 0;
  267. }
  268.  
  269. static int load_mnist_images(const char *path, int *num, int *rows, int *cols, unsigned char **data) {
  270.     FILE *f = fopen(path, "rb");
  271.     if (!f) return 0;
  272.     int magic = read_be_int(f);
  273.     int n = read_be_int(f);
  274.     int r = read_be_int(f);
  275.     int c = read_be_int(f);
  276.     if (magic != 2051 || n <= 0 || r != IMG_ROWS || c != IMG_COLS) {
  277.         fclose(f);
  278.         return 0;
  279.     }
  280.     size_t sz = (size_t)n * r * c;
  281.     unsigned char *buf = (unsigned char *)malloc(sz);
  282.     if (!buf) {
  283.         fclose(f);
  284.         return 0;
  285.     }
  286.     if (fread(buf, 1, sz, f) != sz) {
  287.         free(buf);
  288.         fclose(f);
  289.         return 0;
  290.     }
  291.     fclose(f);
  292.     *num = n; *rows = r; *cols = c; *data = buf;
  293.     return 1;
  294. }
  295.  
  296. static int load_mnist_labels(const char *path, int *num, unsigned char **labels) {
  297.     FILE *f = fopen(path, "rb");
  298.     if (!f) return 0;
  299.     int magic = read_be_int(f);
  300.     int n = read_be_int(f);
  301.     if (magic != 2049 || n <= 0) {
  302.         fclose(f);
  303.         return 0;
  304.     }
  305.     unsigned char *buf = (unsigned char *)malloc((size_t)n);
  306.     if (!buf) {
  307.         fclose(f);
  308.         return 0;
  309.     }
  310.     if (fread(buf, 1, (size_t)n, f) != (size_t)n) {
  311.         free(buf);
  312.         fclose(f);
  313.         return 0;
  314.     }
  315.     fclose(f);
  316.     *num = n; *labels = buf;
  317.     return 1;
  318. }
  319.  
  320. static int load_mnist_set(const char *img_path, const char *lbl_path, MnistSet *set) {
  321.     int n1, n2, r, c;
  322.     unsigned char *imgs = NULL;
  323.     unsigned char *lbls = NULL;
  324.     if (!load_mnist_images(img_path, &n1, &r, &c, &imgs)) return 0;
  325.     if (!load_mnist_labels(lbl_path, &n2, &lbls)) {
  326.         free(imgs);
  327.         return 0;
  328.     }
  329.     if (n1 != n2) {
  330.         free(imgs); free(lbls);
  331.         return 0;
  332.     }
  333.     set->num_images = n1;
  334.     set->rows = r;
  335.     set->cols = c;
  336.     set->images = imgs;
  337.     set->labels = lbls;
  338.     return 1;
  339. }
  340.  
  341. static void free_mnist_set(MnistSet *set) {
  342.     free(set->images);
  343.     free(set->labels);
  344.     set->images = NULL;
  345.     set->labels = NULL;
  346. }
  347.  
  348. static void get_image_vector(const MnistSet *set, int idx, double x[INPUT_DIM]) {
  349.     const unsigned char *img = set->images + (size_t)idx * INPUT_DIM;
  350.     for (int i = 0; i < INPUT_DIM; ++i) {
  351.         x[i] = ((double)img[i]) / 255.0;
  352.     }
  353. }
  354.  
  355. static void print_ascii_digit(const MnistSet *set, int idx) {
  356.     const unsigned char *img = set->images + (size_t)idx * INPUT_DIM;
  357.     const char *lut = " .:-=+*#%@";
  358.     printf("ASCII preview (label=%d):\n", set->labels[idx]);
  359.     for (int r = 0; r < IMG_ROWS; ++r) {
  360.         for (int c = 0; c < IMG_COLS; ++c) {
  361.             int v = img[r * IMG_COLS + c];
  362.             int k = (v * 9) / 255;
  363.             putchar(lut[k]);
  364.         }
  365.         putchar('\n');
  366.     }
  367.     putchar('\n');
  368. }
  369.  
  370. static void print_ascii_vector28(const double x[INPUT_DIM], int label, const char *title) {
  371.     const char *lut = " .:-=+*#%@";
  372.     printf("%s (label=%d):\n", title, label);
  373.     for (int r = 0; r < IMG_ROWS; ++r) {
  374.         for (int c = 0; c < IMG_COLS; ++c) {
  375.             double v = x[r * IMG_COLS + c];
  376.             if (v < 0.0) v = 0.0;
  377.             if (v > 1.0) v = 1.0;
  378.             int iv = (int)(v * 255.0 + 0.5);
  379.             int k = (iv * 9) / 255;
  380.             putchar(lut[k]);
  381.         }
  382.         putchar('\n');
  383.     }
  384.     putchar('\n');
  385. }
  386.  
  387. static int save_pgm_from_vector(const char *path, const double x[INPUT_DIM]) {
  388.     FILE *f = fopen(path, "wb");
  389.     if (!f) return 0;
  390.     fprintf(f, "P5\n%d %d\n255\n", IMG_COLS, IMG_ROWS);
  391.     for (int i = 0; i < INPUT_DIM; ++i) {
  392.         double v = x[i];
  393.         if (v < 0.0) v = 0.0;
  394.         if (v > 1.0) v = 1.0;
  395.         unsigned char px = (unsigned char)(v * 255.0 + 0.5);
  396.         fwrite(&px, 1, 1, f);
  397.     }
  398.     fclose(f);
  399.     return 1;
  400. }
  401.  
  402. static int save_pgm_from_dataset_image(const char *path, const MnistSet *set, int idx) {
  403.     FILE *f = fopen(path, "wb");
  404.     if (!f) return 0;
  405.     fprintf(f, "P5\n%d %d\n255\n", IMG_COLS, IMG_ROWS);
  406.     const unsigned char *img = set->images + (size_t)idx * INPUT_DIM;
  407.     fwrite(img, 1, INPUT_DIM, f);
  408.     fclose(f);
  409.     return 1;
  410. }
  411.  
  412. static int find_nth_occurrence_of_label(const MnistSet *set, int digit, int occurrence_zero_based) {
  413.     int seen = 0;
  414.     for (int i = 0; i < set->num_images; ++i) {
  415.         if (set->labels[i] == digit) {
  416.             if (seen == occurrence_zero_based) return i;
  417.             ++seen;
  418.         }
  419.     }
  420.     return -1;
  421. }
  422.  
  423. static int count_label_occurrences(const MnistSet *set, int digit) {
  424.     int cnt = 0;
  425.     for (int i = 0; i < set->num_images; ++i) {
  426.         if (set->labels[i] == digit) ++cnt;
  427.     }
  428.     return cnt;
  429. }
  430.  
  431. static int choose_random_index_of_label(const MnistSet *set, int digit) {
  432.     int count = count_label_occurrences(set, digit);
  433.     if (count <= 0) return -1;
  434.     int pick = (int)(urand01() * count);
  435.     if (pick >= count) pick = count - 1;
  436.     return find_nth_occurrence_of_label(set, digit, pick);
  437. }
  438.  
  439. static void add_noise_uniform(double x[INPUT_DIM], double noise_percent) {
  440.     double a = noise_percent / 100.0;
  441.     for (int i = 0; i < INPUT_DIM; ++i) {
  442.         double n = rand_uniform(-a, a);
  443.         x[i] += n;
  444.         if (x[i] < 0.0) x[i] = 0.0;
  445.         if (x[i] > 1.0) x[i] = 1.0;
  446.     }
  447. }
  448.  
  449. static double randn_box_muller(void) {
  450.     double u1 = urand01();
  451.     double u2 = urand01();
  452.     if (u1 < 1e-12) u1 = 1e-12;
  453.     return sqrt(-2.0 * log(u1)) * cos(2.0 * M_PI * u2);
  454. }
  455.  
  456. static void add_noise_gaussian(double x[INPUT_DIM], double noise_percent) {
  457.     double sigma = noise_percent / 100.0;
  458.     for (int i = 0; i < INPUT_DIM; ++i) {
  459.         double n = sigma * randn_box_muller();
  460.         x[i] += n;
  461.         if (x[i] < 0.0) x[i] = 0.0;
  462.         if (x[i] > 1.0) x[i] = 1.0;
  463.     }
  464. }
  465.  
  466. static double clamp01(double x) {
  467.     if (x < 0.0) return 0.0;
  468.     if (x > 1.0) return 1.0;
  469.     return x;
  470. }
  471.  
  472. static void print_ascii_vector(const double x[INPUT_DIM], int label, const char *title) {
  473.     const char *lut = " .:-=+*#%@";
  474.     printf("%s (label=%d):\n", title, label);
  475.     for (int r = 0; r < IMG_ROWS; ++r) {
  476.         for (int c = 0; c < IMG_COLS; ++c) {
  477.             double v = clamp01(x[r * IMG_COLS + c]);
  478.             int k = (int)(v * 9.0);
  479.             if (k < 0) k = 0;
  480.             if (k > 9) k = 9;
  481.             putchar(lut[k]);
  482.         }
  483.         putchar('\n');
  484.     }
  485.     putchar('\n');
  486. }
  487.  
  488. static double rand_normal01(void) {
  489.     double u1 = urand01();
  490.     double u2 = urand01();
  491.     if (u1 < 1e-12) u1 = 1e-12;
  492.     return sqrt(-2.0 * log(u1)) * cos(2.0 * M_PI * u2);
  493. }
  494.  
  495. static void apply_uniform_noise(const double src[INPUT_DIM], double dst[INPUT_DIM], double percent) {
  496.     double amp = percent / 100.0;
  497.     for (int i = 0; i < INPUT_DIM; ++i) {
  498.         double noise = rand_uniform(-amp, amp);
  499.         dst[i] = clamp01(src[i] + noise);
  500.     }
  501. }
  502.  
  503. static void apply_gaussian_noise(const double src[INPUT_DIM], double dst[INPUT_DIM], double percent) {
  504.     double sigma = percent / 100.0;
  505.     for (int i = 0; i < INPUT_DIM; ++i) {
  506.         double noise = rand_normal01() * sigma;
  507.         dst[i] = clamp01(src[i] + noise);
  508.     }
  509. }
  510.  
  511. static int find_nth_label_index(const MnistSet *set, int label, int occurrence) {
  512.     int seen = 0;
  513.     for (int i = 0; i < set->num_images; ++i) {
  514.         if (set->labels[i] == label) {
  515.             if (seen == occurrence) return i;
  516.             ++seen;
  517.         }
  518.     }
  519.     return -1;
  520. }
  521.  
  522. static int count_label_images(const MnistSet *set, int label) {
  523.     int count = 0;
  524.     for (int i = 0; i < set->num_images; ++i) {
  525.         if (set->labels[i] == label) ++count;
  526.     }
  527.     return count;
  528. }
  529.  
  530. static double relu(double x) {
  531.     return x > 0.0 ? x : 0.0;
  532. }
  533.  
  534. static void init_weights(void) {
  535.     for (int j = 0; j < HIDDEN_DIM; ++j) {
  536.         for (int i = 0; i < INPUT_DIM; ++i) {
  537.             fc1_w[j][i] = rand_uniform(-0.05, 0.05);
  538.             m_fc1_w[j][i] = 0.0;
  539.             v_fc1_w[j][i] = 0.0;
  540.         }
  541.         fc1_b[j] = 0.0;
  542.         m_fc1_b[j] = 0.0;
  543.         v_fc1_b[j] = 0.0;
  544.     }
  545.  
  546.     for (int k = 0; k < N_CLASSES; ++k) {
  547.         for (int j = 0; j < HIDDEN_DIM; ++j) {
  548.             qc_w[k][j] = rand_uniform(-0.05, 0.05);
  549.             m_qc_w[k][j] = 0.0;
  550.             v_qc_w[k][j] = 0.0;
  551.         }
  552.         qc_b[k] = 0.0;
  553.         m_qc_b[k] = 0.0;
  554.         v_qc_b[k] = 0.0;
  555.     }
  556.  
  557.     for (int k = 0; k < N_CLASSES; ++k) {
  558.         for (int j = 0; j < Q_HEAD_IN; ++j) {
  559.             qq_w[k][j] = rand_uniform(-0.05, 0.05);
  560.             m_qq_w[k][j] = 0.0;
  561.             v_qq_w[k][j] = 0.0;
  562.         }
  563.         qq_b[k] = 0.0;
  564.         m_qq_b[k] = 0.0;
  565.         v_qq_b[k] = 0.0;
  566.     }
  567.  
  568.     for (int p = 0; p < Q_VAR_PARAM_COUNT; ++p) {
  569.         q_theta_var[p] = rand_uniform(-0.1, 0.1);
  570.         m_qtheta[p] = 0.0;
  571.         v_qtheta[p] = 0.0;
  572.     }
  573.  
  574.     for (int q = 0; q < N_QUBITS; ++q) {
  575.         for (int j = 0; j < HIDDEN_DIM; ++j) {
  576.             qpre_w[q][j] = rand_uniform(-0.05, 0.05);
  577.             m_qpre_w[q][j] = 0.0;
  578.             v_qpre_w[q][j] = 0.0;
  579.         }
  580.         qpre_b[q] = 0.0;
  581.         m_qpre_b[q] = 0.0;
  582.         v_qpre_b[q] = 0.0;
  583.     }
  584. }
  585.  
  586. static void zero_grads(void) {
  587.     memset(g_fc1_w, 0, sizeof(g_fc1_w));
  588.     memset(g_fc1_b, 0, sizeof(g_fc1_b));
  589.     memset(g_qc_w, 0, sizeof(g_qc_w));
  590.     memset(g_qc_b, 0, sizeof(g_qc_b));
  591.     memset(g_qq_w, 0, sizeof(g_qq_w));
  592.     memset(g_qq_b, 0, sizeof(g_qq_b));
  593.     memset(g_qtheta, 0, sizeof(g_qtheta));
  594.     memset(g_qpre_w, 0, sizeof(g_qpre_w));
  595.     memset(g_qpre_b, 0, sizeof(g_qpre_b));
  596. }
  597.  
  598. static void adam_update_scalar(double *param, double grad, double *m, double *v, int t) {
  599.     double g = grad + WEIGHT_DECAY * (*param);
  600.     *m = BETA1 * (*m) + (1.0 - BETA1) * g;
  601.     *v = BETA2 * (*v) + (1.0 - BETA2) * g * g;
  602.     double mhat = (*m) / (1.0 - pow(BETA1, (double)t));
  603.     double vhat = (*v) / (1.0 - pow(BETA2, (double)t));
  604.     *param -= LEARNING_RATE * mhat / (sqrt(vhat) + EPS_ADAM);
  605. }
  606.  
  607. static void apply_adam_all(int t) {
  608.     for (int j = 0; j < HIDDEN_DIM; ++j) {
  609.         for (int i = 0; i < INPUT_DIM; ++i) {
  610.             adam_update_scalar(&fc1_w[j][i], g_fc1_w[j][i], &m_fc1_w[j][i], &v_fc1_w[j][i], t);
  611.         }
  612.         adam_update_scalar(&fc1_b[j], g_fc1_b[j], &m_fc1_b[j], &v_fc1_b[j], t);
  613.     }
  614.     for (int k = 0; k < N_CLASSES; ++k) {
  615.         for (int j = 0; j < HIDDEN_DIM; ++j) {
  616.             adam_update_scalar(&qc_w[k][j], g_qc_w[k][j], &m_qc_w[k][j], &v_qc_w[k][j], t);
  617.         }
  618.         adam_update_scalar(&qc_b[k], g_qc_b[k], &m_qc_b[k], &v_qc_b[k], t);
  619.     }
  620.     for (int k = 0; k < N_CLASSES; ++k) {
  621.         for (int j = 0; j < Q_HEAD_IN; ++j) {
  622.             adam_update_scalar(&qq_w[k][j], g_qq_w[k][j], &m_qq_w[k][j], &v_qq_w[k][j], t);
  623.         }
  624.         adam_update_scalar(&qq_b[k], g_qq_b[k], &m_qq_b[k], &v_qq_b[k], t);
  625.     }
  626.     for (int q = 0; q < N_QUBITS; ++q) {
  627.         for (int j = 0; j < HIDDEN_DIM; ++j) {
  628.             adam_update_scalar(&qpre_w[q][j], g_qpre_w[q][j],
  629.                 &m_qpre_w[q][j], &v_qpre_w[q][j], t);
  630.         }
  631.         adam_update_scalar(&qpre_b[q], g_qpre_b[q],
  632.             &m_qpre_b[q], &v_qpre_b[q], t);
  633.     }
  634.     for (int p = 0; p < Q_VAR_PARAM_COUNT; ++p) {
  635.         adam_update_scalar(&q_theta_var[p], g_qtheta[p], &m_qtheta[p], &v_qtheta[p], t);
  636.     }
  637. }
  638.  
  639. static Complex c_add(Complex a, Complex b) {
  640.     Complex r = {a.real + b.real, a.imag + b.imag};
  641.     return r;
  642. }
  643.  
  644. static Complex c_mul(Complex a, Complex b) {
  645.     Complex r = {a.real * b.real - a.imag * b.imag,
  646.                  a.real * b.imag + a.imag * b.real};
  647.     return r;
  648. }
  649.  
  650. static void q_apply_gate(int qubit, Complex m00, Complex m01, Complex m10, Complex m11) {
  651.     int step = 1 << qubit;
  652.     for (int i = 0; i < STATE_SIZE; i += 2 * step) {
  653.         for (int j = 0; j < step; ++j) {
  654.             int i0 = i + j;
  655.             int i1 = i0 + step;
  656.             Complex a = q_state[i0];
  657.             Complex b = q_state[i1];
  658.             q_state[i0] = c_add(c_mul(m00, a), c_mul(m01, b));
  659.             q_state[i1] = c_add(c_mul(m10, a), c_mul(m11, b));
  660.         }
  661.     }
  662. }
  663.  
  664. static void q_apply_H(int q) {
  665.     double s = 1.0 / sqrt(2.0);
  666.     q_apply_gate(q, (Complex){s,0}, (Complex){s,0}, (Complex){s,0}, (Complex){-s,0});
  667. }
  668.  
  669. static void q_apply_RY(int q, double t) {
  670.     double c = cos(t / 2.0);
  671.     double s = sin(t / 2.0);
  672.     q_apply_gate(q, (Complex){c,0}, (Complex){-s,0}, (Complex){s,0}, (Complex){c,0});
  673. }
  674.  
  675. static void q_apply_RX(int q, double t) {
  676.     double c = cos(t / 2.0);
  677.     double s = sin(t / 2.0);
  678.     q_apply_gate(q, (Complex){c,0}, (Complex){0,-s}, (Complex){0,-s}, (Complex){c,0});
  679. }
  680.  
  681. static void q_apply_RZ(int q, double t) {
  682.     Complex e0 = {cos(-t / 2.0), sin(-t / 2.0)};
  683.     Complex e1 = {cos(t / 2.0), sin(t / 2.0)};
  684.     q_apply_gate(q, e0, (Complex){0,0}, (Complex){0,0}, e1);
  685. }
  686.  
  687. static void q_apply_CNOT(int control, int target) {
  688.     for (int i = 0; i < STATE_SIZE; ++i) {
  689.         if ((i >> control) & 1) {
  690.             int j = i ^ (1 << target);
  691.             if (i < j) {
  692.                 Complex tmp = q_state[i];
  693.                 q_state[i] = q_state[j];
  694.                 q_state[j] = tmp;
  695.             }
  696.         }
  697.     }
  698. }
  699.  
  700. static double q_measure_Z(int qubit) {
  701.     double expv = 0.0;
  702.     for (int i = 0; i < STATE_SIZE; ++i) {
  703.         double prob = q_state[i].real * q_state[i].real + q_state[i].imag * q_state[i].imag;
  704.         int bit = (i >> qubit) & 1;
  705.         expv += bit ? -prob : prob;
  706.     }
  707.     return expv;
  708. }
  709.  
  710. static void quantum_expectation_from_hidden(
  711.     const double hidden[HIDDEN_DIM],
  712.     double qpre_w_local[N_QUBITS][HIDDEN_DIM],
  713.     double qpre_b_local[N_QUBITS],
  714.     const double theta_var[Q_VAR_PARAM_COUNT],
  715.     double q_in_out[N_QUBITS],
  716.     double z_out[N_QUBITS]
  717. ) {
  718.     for (int i = 0; i < STATE_SIZE; ++i) {
  719.         q_state[i].real = 0.0;
  720.         q_state[i].imag = 0.0;
  721.     }
  722.     q_state[0].real = 1.0;
  723.  
  724.     double pre_out[N_QUBITS];
  725.     double q_in[N_QUBITS];
  726.  
  727.     for (int q = 0; q < N_QUBITS; ++q) {
  728.         double s = qpre_b[q];
  729.         for (int j = 0; j < HIDDEN_DIM; ++j) {
  730.             s += qpre_w[q][j] * hidden[j];
  731.         }
  732.         pre_out[q] = s;
  733.         q_in[q] = tanh(pre_out[q]) * (M_PI / 2.0);
  734.     }
  735.  
  736.     for (int q = 0; q < N_QUBITS; ++q) q_apply_H(q);
  737.     for (int q = 0; q < N_QUBITS; ++q) q_apply_RY(q, q_in[q]);
  738.  
  739.     int p = 0;
  740.     for (int i = 1; i < N_QUBITS; ++i) {
  741.         q_apply_CNOT(0, i);
  742.         q_apply_RZ(i, theta_var[p++]);
  743.         q_apply_RX(i, theta_var[p++]);
  744.         q_apply_RZ(i, theta_var[p++]);
  745.         q_apply_CNOT(0, i);
  746.     }
  747.  
  748.     for (int q = 0; q < N_QUBITS; ++q) {
  749.         z_out[q] = q_measure_Z(q);
  750.     }
  751. }
  752.  
  753. static void linear_forward_hidden_to_class(double w[N_CLASSES][HIDDEN_DIM], double b[N_CLASSES], const double x[HIDDEN_DIM], double out[N_CLASSES]) {
  754.     for (int k = 0; k < N_CLASSES; ++k) {
  755.         double s = b[k];
  756.         for (int j = 0; j < HIDDEN_DIM; ++j) s += w[k][j] * x[j];
  757.         out[k] = s;
  758.     }
  759. }
  760.  
  761. static void linear_forward_qhead_to_class(double w[N_CLASSES][Q_HEAD_IN], double b[N_CLASSES], const double x[Q_HEAD_IN], double out[N_CLASSES]) {
  762.     for (int k = 0; k < N_CLASSES; ++k) {
  763.         double s = b[k];
  764.         for (int j = 0; j < Q_HEAD_IN; ++j) s += w[k][j] * x[j];
  765.         out[k] = s;
  766.     }
  767. }
  768.  
  769. static void softmax10(const double in[N_CLASSES], double out[N_CLASSES]) {
  770.     double max_v = in[0];
  771.     for (int i = 1; i < N_CLASSES; ++i) if (in[i] > max_v) max_v = in[i];
  772.     double sum = 0.0;
  773.     for (int i = 0; i < N_CLASSES; ++i) {
  774.         out[i] = exp(in[i] - max_v);
  775.         sum += out[i];
  776.     }
  777.     if (sum <= 0.0) sum = 1.0;
  778.     for (int i = 0; i < N_CLASSES; ++i) out[i] /= sum;
  779. }
  780.  
  781. static int argmax10(const double v[N_CLASSES]) {
  782.     int idx = 0;
  783.     double best = v[0];
  784.     for (int i = 1; i < N_CLASSES; ++i) {
  785.         if (v[i] > best) {
  786.             best = v[i];
  787.             idx = i;
  788.         }
  789.     }
  790.     return idx;
  791. }
  792.  
  793. static double nll_one(const double probs[N_CLASSES], int label) {
  794.     double p = probs[label];
  795.     if (p < 1e-12) p = 1e-12;
  796.     return -log(p);
  797. }
  798.  
  799. static void print_vec10_dbg(const char* name, const double v[N_CLASSES]) {
  800.     printf("%s:\n", name);
  801.     for (int i = 0; i < N_CLASSES; ++i) {
  802.         printf("  [%d] %.10f\n", i, v[i]);
  803.     }
  804.     printf("\n");
  805. }
  806.  
  807. static void preprocess_fc1(const double x[INPUT_DIM], double pre[HIDDEN_DIM], double act[HIDDEN_DIM]) {
  808.     for (int j = 0; j < HIDDEN_DIM; ++j) {
  809.         double s = fc1_b[j];
  810.         for (int i = 0; i < INPUT_DIM; ++i) s += fc1_w[j][i] * x[i];
  811.         pre[j] = s;
  812.         act[j] = relu(s);
  813.     }
  814. }
  815.  
  816. static void snn_forward_poisson_lif(const double act[HIDDEN_DIM], double pooled[HIDDEN_DIM]) {
  817.     /*
  818.        SpikingActivation-style forward:
  819.        - frequency proportional to base activation
  820.        - multiple spikes per timestep allowed
  821.        - spike height = 1/dt
  822.        - TemporalAvgPool = average over time axis
  823.     */
  824.     double voltage[HIDDEN_DIM];
  825.     double sum_spike_output[HIDDEN_DIM];
  826.  
  827.     for (int j = 0; j < HIDDEN_DIM; ++j) {
  828.         /* initial state in [0,1), как у spiking voltage state */
  829.         voltage[j] = urand01();
  830.         sum_spike_output[j] = 0.0;
  831.     }
  832.  
  833.     for (int t = 0; t < g_time_steps; ++t) {
  834.         for (int j = 0; j < HIDDEN_DIM; ++j) {
  835.             double rate = act[j];
  836.  
  837.             /* ReLU already used before this function, but guard anyway */
  838.             if (rate < 0.0) rate = 0.0;
  839.  
  840.             /* integrate rate over one timestep */
  841.             voltage[j] += rate * DT;
  842.  
  843.             /* allow multiple spikes per timestep */
  844.             int n_spikes = (int)floor(voltage[j]);
  845.  
  846.             /* subtract emitted spikes from voltage */
  847.             voltage[j] -= (double)n_spikes;
  848.  
  849.             /* each spike has height 1/dt */
  850.             double spike_output = ((double)n_spikes) / DT;
  851.  
  852.             /* TemporalAvgPool later averages over time axis */
  853.             sum_spike_output[j] += spike_output;
  854.         }
  855.     }
  856.  
  857.     for (int j = 0; j < HIDDEN_DIM; ++j) {
  858.         pooled[j] = sum_spike_output[j] / (double)g_time_steps;
  859.     }
  860. }
  861.  
  862. static void fuse_logits(const double qc[N_CLASSES], const double qq[N_CLASSES], double qh[N_CLASSES]) {
  863.     for (int k = 0; k < N_CLASSES; ++k) qh[k] = XI * qq[k] + (1.0 - XI) * qc[k];
  864. }
  865.  
  866. static void model_forward(const double x[INPUT_DIM],
  867.     double fc1_pre[HIDDEN_DIM],
  868.     double fc1_act[HIDDEN_DIM],
  869.     double pooled[HIDDEN_DIM],
  870.     double q_in[N_QUBITS],
  871.     double z[N_QUBITS],
  872.     double qc[N_CLASSES],
  873.     double qq[N_CLASSES],
  874.     double qh[N_CLASSES],
  875.     double probs[N_CLASSES]) {
  876.     preprocess_fc1(x, fc1_pre, fc1_act);
  877.     snn_forward_poisson_lif(fc1_act, pooled);
  878.  
  879.     /* classical branch */
  880.     linear_forward_hidden_to_class(qc_w, qc_b, pooled, qc);
  881.  
  882.     /* quantum branch with q_pre -> tanh -> pi/2 encoding */
  883.     quantum_expectation_from_hidden(fc1_act, qpre_w, qpre_b, q_theta_var, q_in, z);
  884.  
  885.     /* quantum head */
  886.     linear_forward_qhead_to_class(qq_w, qq_b, z, qq);
  887.  
  888.     /* fusion */
  889.     fuse_logits(qc, qq, qh);
  890.     softmax10(qh, probs);
  891. }
  892.  
  893. static double loss_with_theta_shift(const double x[INPUT_DIM], int label, int param_idx, double shift) {
  894.     double saved = q_theta_var[param_idx];
  895.     q_theta_var[param_idx] = saved + shift;
  896.     double fc1_pre[HIDDEN_DIM], fc1_act[HIDDEN_DIM], pooled[HIDDEN_DIM];
  897.     double q_in[N_QUBITS], z[N_QUBITS];
  898.     double qc[N_CLASSES], qq[N_CLASSES], qh[N_CLASSES], probs[N_CLASSES];
  899.     model_forward(x, fc1_pre, fc1_act, pooled, q_in, z, qc, qq, qh, probs);
  900.     double loss = nll_one(probs, label);
  901.     q_theta_var[param_idx] = saved;
  902.     return loss;
  903. }
  904.  
  905. static double loss_with_qin_shift(const double x[INPUT_DIM], int label, int q_idx, double shift) {
  906.     double fc1_pre[HIDDEN_DIM], fc1_act[HIDDEN_DIM], pooled[HIDDEN_DIM];
  907.     double q_in[N_QUBITS], z[N_QUBITS];
  908.     double qc[N_CLASSES], qq[N_CLASSES], qh[N_CLASSES], probs[N_CLASSES];
  909.  
  910.     model_forward(x, fc1_pre, fc1_act, pooled, q_in, z, qc, qq, qh, probs);
  911.  
  912.     /* recompute only the quantum branch with shifted q_in[q_idx] */
  913.     for (int i = 0; i < STATE_SIZE; ++i) {
  914.         q_state[i].real = 0.0;
  915.         q_state[i].imag = 0.0;
  916.     }
  917.     q_state[0].real = 1.0;
  918.     q_state[0].imag = 0.0;
  919.  
  920.     double q_in_shifted[N_QUBITS];
  921.     for (int q = 0; q < N_QUBITS; ++q) {
  922.         q_in_shifted[q] = q_in[q];
  923.     }
  924.     q_in_shifted[q_idx] += shift;
  925.  
  926.     for (int q = 0; q < N_QUBITS; ++q) {
  927.         q_apply_H(q);
  928.     }
  929.     for (int q = 0; q < N_QUBITS; ++q) {
  930.         q_apply_RY(q, q_in_shifted[q]);
  931.     }
  932.  
  933.     int p = 0;
  934.     for (int i = 1; i < N_QUBITS; ++i) {
  935.         q_apply_CNOT(0, i);
  936.         q_apply_RZ(i, q_theta_var[p++]);
  937.         q_apply_RX(i, q_theta_var[p++]);
  938.         q_apply_RZ(i, q_theta_var[p++]);
  939.         q_apply_CNOT(0, i);
  940.     }
  941.  
  942.     for (int q = 0; q < N_QUBITS; ++q) {
  943.         z[q] = q_measure_Z(q);
  944.     }
  945.  
  946.     linear_forward_qhead_to_class(qq_w, qq_b, z, qq);
  947.     fuse_logits(qc, qq, qh);
  948.     softmax10(qh, probs);
  949.  
  950.     return nll_one(probs, label);
  951. }
  952.  
  953. static void accumulate_sample_gradients(const double x[INPUT_DIM], int label) {
  954.     double fc1_pre[HIDDEN_DIM], fc1_act[HIDDEN_DIM], pooled[HIDDEN_DIM];
  955.     double q_in[N_QUBITS], z[N_QUBITS];
  956.     double qc[N_CLASSES], qq[N_CLASSES], qh[N_CLASSES], probs[N_CLASSES];
  957.  
  958.     model_forward(x, fc1_pre, fc1_act, pooled, q_in, z, qc, qq, qh, probs);
  959.  
  960.     double dL_dqh[N_CLASSES];
  961.     for (int k = 0; k < N_CLASSES; ++k) {
  962.         dL_dqh[k] = probs[k];
  963.     }
  964.     dL_dqh[label] -= 1.0;
  965.  
  966.     double dL_dqc[N_CLASSES], dL_dqq[N_CLASSES];
  967.     for (int k = 0; k < N_CLASSES; ++k) {
  968.         dL_dqq[k] = XI * dL_dqh[k];
  969.         dL_dqc[k] = (1.0 - XI) * dL_dqh[k];
  970.     }
  971.  
  972.     /* qc head grads */
  973.     for (int k = 0; k < N_CLASSES; ++k) {
  974.         for (int j = 0; j < HIDDEN_DIM; ++j) {
  975.             g_qc_w[k][j] += dL_dqc[k] * pooled[j];
  976.         }
  977.         g_qc_b[k] += dL_dqc[k];
  978.     }
  979.  
  980.     /* qq head grads */
  981.     for (int k = 0; k < N_CLASSES; ++k) {
  982.         for (int j = 0; j < Q_HEAD_IN; ++j) {
  983.             g_qq_w[k][j] += dL_dqq[k] * z[j];
  984.         }
  985.         g_qq_b[k] += dL_dqq[k];
  986.     }
  987.  
  988.     /* -------- qpre gradients via dL/dq_in + chain rule -------- */
  989.  
  990.     double dL_dqin[N_QUBITS];
  991.     {
  992.         const double eps_qin = 1e-3;
  993.         for (int q = 0; q < N_QUBITS; ++q) {
  994.             double lp = loss_with_qin_shift(x, label, q, +eps_qin);
  995.             double lm = loss_with_qin_shift(x, label, q, -eps_qin);
  996.             dL_dqin[q] = (lp - lm) / (2.0 * eps_qin);
  997.         }
  998.     }
  999.  
  1000.     for (int q = 0; q < N_QUBITS; ++q) {
  1001.         double pre_out = qpre_b[q];
  1002.         for (int j = 0; j < HIDDEN_DIM; ++j) {
  1003.             pre_out += qpre_w[q][j] * fc1_act[j];
  1004.         }
  1005.  
  1006.         /* q_in = tanh(pre_out) * (pi/2) */
  1007.         double th = tanh(pre_out);
  1008.         double dqindpre = (1.0 - th * th) * (M_PI / 2.0);
  1009.         double dL_dpre = dL_dqin[q] * dqindpre;
  1010.  
  1011.         for (int j = 0; j < HIDDEN_DIM; ++j) {
  1012.             g_qpre_w[q][j] += dL_dpre * fc1_act[j];
  1013.         }
  1014.         g_qpre_b[q] += dL_dpre;
  1015.     }
  1016.  
  1017.     /* -------- classical branch backward -------- */
  1018.  
  1019.     double dL_dpooled[HIDDEN_DIM];
  1020.     for (int j = 0; j < HIDDEN_DIM; ++j) {
  1021.         double s = 0.0;
  1022.         for (int k = 0; k < N_CLASSES; ++k) {
  1023.             s += qc_w[k][j] * dL_dqc[k];
  1024.         }
  1025.         dL_dpooled[j] = s;
  1026.     }
  1027.  
  1028.     /* spiking-aware training surrogate: use base activation backward */
  1029.     double dL_dfc1pre[HIDDEN_DIM];
  1030.     for (int j = 0; j < HIDDEN_DIM; ++j) {
  1031.         double grad = dL_dpooled[j];
  1032.         if (fc1_pre[j] <= 0.0) {
  1033.             grad = 0.0;
  1034.         }
  1035.         dL_dfc1pre[j] = grad;
  1036.     }
  1037.  
  1038.     for (int j = 0; j < HIDDEN_DIM; ++j) {
  1039.         for (int i = 0; i < INPUT_DIM; ++i) {
  1040.             g_fc1_w[j][i] += dL_dfc1pre[j] * x[i];
  1041.         }
  1042.         g_fc1_b[j] += dL_dfc1pre[j];
  1043.     }
  1044.  
  1045.     /* parameter-shift for quantum variational parameters */
  1046.     for (int p = 0; p < Q_VAR_PARAM_COUNT; ++p) {
  1047.         double lp = loss_with_theta_shift(x, label, p, M_PI / 2.0);
  1048.         double lm = loss_with_theta_shift(x, label, p, -M_PI / 2.0);
  1049.         g_qtheta[p] += 0.5 * (lp - lm);
  1050.     }
  1051. }
  1052.  
  1053. static void scale_all_grads(double scale) {
  1054.     for (int j = 0; j < HIDDEN_DIM; ++j) {
  1055.         for (int i = 0; i < INPUT_DIM; ++i) {
  1056.             g_fc1_w[j][i] *= scale;
  1057.         }
  1058.         g_fc1_b[j] *= scale;
  1059.     }
  1060.  
  1061.     for (int k = 0; k < N_CLASSES; ++k) {
  1062.         for (int j = 0; j < HIDDEN_DIM; ++j) {
  1063.             g_qc_w[k][j] *= scale;
  1064.         }
  1065.         g_qc_b[k] *= scale;
  1066.     }
  1067.  
  1068.     for (int k = 0; k < N_CLASSES; ++k) {
  1069.         for (int j = 0; j < Q_HEAD_IN; ++j) {
  1070.             g_qq_w[k][j] *= scale;
  1071.         }
  1072.         g_qq_b[k] *= scale;
  1073.     }
  1074.  
  1075.     for (int q = 0; q < N_QUBITS; ++q) {
  1076.         for (int j = 0; j < HIDDEN_DIM; ++j) {
  1077.             g_qpre_w[q][j] *= scale;
  1078.         }
  1079.         g_qpre_b[q] *= scale;
  1080.     }
  1081.  
  1082.     for (int p = 0; p < Q_VAR_PARAM_COUNT; ++p) {
  1083.         g_qtheta[p] *= scale;
  1084.     }
  1085. }
  1086.  
  1087. static double train_one_epoch(const MnistSet *train_set, const int *train_idx, int epoch, int *adam_step) {
  1088.     double epoch_loss = 0.0;
  1089.     int n_batches = TRAIN_SUBSET / BATCH_SIZE;
  1090.  
  1091.     int order[TRAIN_SUBSET];
  1092.     for (int i = 0; i < TRAIN_SUBSET; ++i) order[i] = train_idx[i];
  1093.     shuffle_int(order, TRAIN_SUBSET);
  1094.  
  1095.     for (int b = 0; b < n_batches; ++b) {
  1096.         zero_grads();
  1097.         double batch_loss = 0.0;
  1098.         int start = b * BATCH_SIZE;
  1099.         int end = start + BATCH_SIZE;
  1100.         for (int s = start; s < end; ++s) {
  1101.             int idx = order[s];
  1102.             double x[INPUT_DIM];
  1103.             get_image_vector(train_set, idx, x);
  1104.             int label = train_set->labels[idx];
  1105.  
  1106.             double fc1_pre[HIDDEN_DIM], fc1_act[HIDDEN_DIM], pooled[HIDDEN_DIM];
  1107.             double q_in[N_QUBITS], z[N_QUBITS];
  1108.             double qc[N_CLASSES], qq[N_CLASSES], qh[N_CLASSES], probs[N_CLASSES];
  1109.             model_forward(x, fc1_pre, fc1_act, pooled, q_in, z, qc, qq, qh, probs);
  1110.             static int debug_forward_done = 0;
  1111.             if (!debug_forward_done) {
  1112.                 print_vec10_dbg("Qc", qc);
  1113.                 print_vec10_dbg("Qq", qq);
  1114.                 print_vec10_dbg("Qh", qh);
  1115.                 print_vec10_dbg("probs", probs);
  1116.  
  1117.                 printf("true label = %d\n", label);
  1118.                 printf("pred label = %d\n", argmax10(probs));
  1119.  
  1120.                 double avg_pooled = 0.0;
  1121.                 for (int j = 0; j < HIDDEN_DIM; ++j) {
  1122.                     avg_pooled += pooled[j];
  1123.                 }
  1124.                 avg_pooled /= (double)HIDDEN_DIM;
  1125.                 printf("avg pooled (SNN output) = %.10f\n", avg_pooled);
  1126.  
  1127.                 for (int i = 0; i < N_QUBITS; ++i) {
  1128.                     printf("z[%d] = %.10f\n", i, z[i]);
  1129.                 }
  1130.                 printf("\n");
  1131.  
  1132.                 debug_forward_done = 1;
  1133.             }
  1134.             batch_loss += nll_one(probs, label);
  1135.  
  1136.             accumulate_sample_gradients(x, label);
  1137.         }
  1138.         scale_all_grads(1.0 / (double)BATCH_SIZE);
  1139.         static int debug_grad_done = 0;
  1140.         if (!debug_grad_done) {
  1141.             printf("grad g_qc_w[0][0]     = %.10e\n", g_qc_w[0][0]);
  1142.             printf("grad g_qq_w[0][0]     = %.10e\n", g_qq_w[0][0]);
  1143.             printf("grad g_fc1_w[0][0]    = %.10e\n", g_fc1_w[0][0]);
  1144.             printf("grad g_qtheta[0]      = %.10e\n", g_qtheta[0]);
  1145.             printf("grad g_qpre_w[0][0]   = %.10e\n", g_qpre_w[0][0]);
  1146.             printf("grad g_qpre_b[0]      = %.10e\n", g_qpre_b[0]);
  1147.             printf("\n");
  1148.  
  1149.             debug_grad_done = 1;
  1150.         }
  1151.         ++(*adam_step);
  1152.         static int debug_before_update_done = 0;
  1153.         if (!debug_before_update_done) {
  1154.             printf("BEFORE update:\n");
  1155.             printf("qc_w[0][0]     = %.10f\n", qc_w[0][0]);
  1156.             printf("qq_w[0][0]     = %.10f\n", qq_w[0][0]);
  1157.             printf("fc1_w[0][0]    = %.10f\n", fc1_w[0][0]);
  1158.             printf("q_theta_var[0] = %.10f\n", q_theta_var[0]);
  1159.             printf("qpre_w[0][0]    = %.10f\n", qpre_w[0][0]);
  1160.             printf("qpre_b[0]       = %.10f\n", qpre_b[0]);
  1161.             printf("\n");
  1162.         }
  1163.         apply_adam_all(*adam_step);
  1164.         static int debug_after_update_done = 0;
  1165.         if (!debug_after_update_done) {
  1166.             printf("AFTER update:\n");
  1167.             printf("qc_w[0][0]     = %.10f\n", qc_w[0][0]);
  1168.             printf("qq_w[0][0]     = %.10f\n", qq_w[0][0]);
  1169.             printf("fc1_w[0][0]    = %.10f\n", fc1_w[0][0]);
  1170.             printf("q_theta_var[0] = %.10f\n", q_theta_var[0]);
  1171.             printf("qpre_w[0][0]    = %.10f\n", qpre_w[0][0]);
  1172.             printf("qpre_b[0]       = %.10f\n", qpre_b[0]);
  1173.             printf("\n");
  1174.  
  1175.             debug_before_update_done = 1;
  1176.             debug_after_update_done = 1;
  1177.         }
  1178.         epoch_loss += batch_loss / (double)BATCH_SIZE;
  1179.  
  1180.         if ((b + 1) % 25 == 0 || b == n_batches - 1) {
  1181.             printf("  epoch %d batch %d/%d loss=%.6f\n", epoch + 1, b + 1, n_batches, batch_loss / (double)BATCH_SIZE);
  1182.         }
  1183.     }
  1184.     return epoch_loss / (double)n_batches;
  1185. }
  1186.  
  1187. static EvalInfo evaluate_one_sample(const MnistSet *set, int idx) {
  1188.     EvalInfo e;
  1189.     double x[INPUT_DIM];
  1190.     get_image_vector(set, idx, x);
  1191.     double fc1_pre[HIDDEN_DIM], fc1_act[HIDDEN_DIM], pooled[HIDDEN_DIM];
  1192.     double q_in[N_QUBITS], z[N_QUBITS];
  1193.     double qc[N_CLASSES], qq[N_CLASSES], qh[N_CLASSES], probs[N_CLASSES];
  1194.     model_forward(x, fc1_pre, fc1_act, pooled, q_in, z, qc, qq, qh, probs);
  1195.     e.true_label = set->labels[idx];
  1196.     e.pred_label = argmax10(probs);
  1197.     e.confidence = probs[e.pred_label];
  1198.     e.loss = nll_one(probs, e.true_label);
  1199.     return e;
  1200. }
  1201.  
  1202. static double evaluate_accuracy(const MnistSet *set, const int *idxs, int n, double *avg_loss) {
  1203.     int correct = 0;
  1204.     double loss = 0.0;
  1205.     for (int t = 0; t < n; ++t) {
  1206.         double x[INPUT_DIM];
  1207.         get_image_vector(set, idxs[t], x);
  1208.         double fc1_pre[HIDDEN_DIM], fc1_act[HIDDEN_DIM], pooled[HIDDEN_DIM];
  1209.         double q_in[N_QUBITS], z[N_QUBITS];
  1210.         double qc[N_CLASSES], qq[N_CLASSES], qh[N_CLASSES], probs[N_CLASSES];
  1211.         model_forward(x, fc1_pre, fc1_act, pooled, q_in, z, qc, qq, qh, probs);
  1212.         int pred = argmax10(probs);
  1213.         int label = set->labels[idxs[t]];
  1214.         if (pred == label) ++correct;
  1215.         loss += nll_one(probs, label);
  1216.     }
  1217.     *avg_loss = loss / (double)n;
  1218.     return (double)correct / (double)n;
  1219. }
  1220.  
  1221.  
  1222. static void wait_enter(void) {
  1223.     int c;
  1224.     printf("\nPress Enter to continue...");
  1225.     fflush(stdout);
  1226.     do {
  1227.         c = getchar();
  1228.     } while (c != '\n' && c != EOF);
  1229. }
  1230.  
  1231. static void clear_stdin_line(void) {
  1232.     int c;
  1233.     do {
  1234.         c = getchar();
  1235.     } while (c != '\n' && c != EOF);
  1236. }
  1237.  
  1238. static int read_int_prompt(const char *prompt, int default_value) {
  1239.     char buf[128];
  1240.     printf("%s", prompt);
  1241.     if (!fgets(buf, (int)sizeof(buf), stdin)) return default_value;
  1242.     if (buf[0] == '\n' || buf[0] == '\0') return default_value;
  1243.     return atoi(buf);
  1244. }
  1245.  
  1246. static double read_double_prompt(const char *prompt, double default_value) {
  1247.     char buf[128];
  1248.     printf("%s", prompt);
  1249.     if (!fgets(buf, (int)sizeof(buf), stdin)) return default_value;
  1250.     if (buf[0] == '\n' || buf[0] == '\0') return default_value;
  1251.     return atof(buf);
  1252. }
  1253.  
  1254. static void save_model_to_file(const char *path) {
  1255.     FILE *f = fopen(path, "wb");
  1256.     if (!f) {
  1257.         printf("Failed to open model file for writing: %s\n", path);
  1258.         return;
  1259.     }
  1260.     const char magic[8] = { 'P','P','F','S','Q','N','N','1' };
  1261.     fwrite(magic, 1, 8, f);
  1262.     fwrite(&g_time_steps, sizeof(g_time_steps), 1, f);
  1263.     fwrite(fc1_w, sizeof(fc1_w), 1, f);
  1264.     fwrite(fc1_b, sizeof(fc1_b), 1, f);
  1265.     fwrite(qc_w, sizeof(qc_w), 1, f);
  1266.     fwrite(qc_b, sizeof(qc_b), 1, f);
  1267.     fwrite(qq_w, sizeof(qq_w), 1, f);
  1268.     fwrite(qq_b, sizeof(qq_b), 1, f);
  1269.     fwrite(q_theta_var, sizeof(q_theta_var), 1, f);
  1270.     fclose(f);
  1271.     printf("Model saved to %s\n", path);
  1272. }
  1273.  
  1274. static int load_model_from_file(const char *path) {
  1275.     FILE *f = fopen(path, "rb");
  1276.     if (!f) {
  1277.         printf("Failed to open model file for reading: %s\n", path);
  1278.         return 0;
  1279.     }
  1280.     char magic[8];
  1281.     if (fread(magic, 1, 8, f) != 8 || memcmp(magic, "PPFSQNN1", 8) != 0) {
  1282.         fclose(f);
  1283.         printf("Invalid model file format: %s\n", path);
  1284.         return 0;
  1285.     }
  1286.     if (fread(&g_time_steps, sizeof(g_time_steps), 1, f) != 1 ||
  1287.         fread(fc1_w, sizeof(fc1_w), 1, f) != 1 ||
  1288.         fread(fc1_b, sizeof(fc1_b), 1, f) != 1 ||
  1289.         fread(qc_w, sizeof(qc_w), 1, f) != 1 ||
  1290.         fread(qc_b, sizeof(qc_b), 1, f) != 1 ||
  1291.         fread(qq_w, sizeof(qq_w), 1, f) != 1 ||
  1292.         fread(qq_b, sizeof(qq_b), 1, f) != 1 ||
  1293.         fread(q_theta_var, sizeof(q_theta_var), 1, f) != 1) {
  1294.         fclose(f);
  1295.         printf("Failed to read full model file: %s\n", path);
  1296.         return 0;
  1297.     }
  1298.     fclose(f);
  1299.     g_model_initialized = 1;
  1300.     printf("Model loaded from %s\n", path);
  1301.     printf("Loaded implementation parameter: time steps = %d\n", g_time_steps);
  1302.     return 1;
  1303. }
  1304.  
  1305. static void save_model_to_text_file(const char *path) {
  1306.     FILE *f = fopen(path, "w");
  1307.     if (!f) {
  1308.         printf("Failed to open text model file for writing: %s\n", path);
  1309.         return;
  1310.     }
  1311.  
  1312.     fprintf(f, "PPF-SQNN model dump\n");
  1313.     fprintf(f, "g_time_steps %d\n", g_time_steps);
  1314.  
  1315.     fprintf(f, "\n[fc1_w]\n");
  1316.     for (int j = 0; j < HIDDEN_DIM; ++j) {
  1317.         for (int i = 0; i < INPUT_DIM; ++i) {
  1318.             fprintf(f, "%.17g", fc1_w[j][i]);
  1319.             if (i + 1 < INPUT_DIM) fprintf(f, " ");
  1320.         }
  1321.         fprintf(f, "\n");
  1322.     }
  1323.  
  1324.     fprintf(f, "\n[fc1_b]\n");
  1325.     for (int j = 0; j < HIDDEN_DIM; ++j) fprintf(f, "%.17g\n", fc1_b[j]);
  1326.  
  1327.     fprintf(f, "\n[qc_w]\n");
  1328.     for (int k = 0; k < N_CLASSES; ++k) {
  1329.         for (int j = 0; j < HIDDEN_DIM; ++j) {
  1330.             fprintf(f, "%.17g", qc_w[k][j]);
  1331.             if (j + 1 < HIDDEN_DIM) fprintf(f, " ");
  1332.         }
  1333.         fprintf(f, "\n");
  1334.     }
  1335.  
  1336.     fprintf(f, "\n[qc_b]\n");
  1337.     for (int k = 0; k < N_CLASSES; ++k) fprintf(f, "%.17g\n", qc_b[k]);
  1338.  
  1339.     fprintf(f, "\n[qq_w]\n");
  1340.     for (int k = 0; k < N_CLASSES; ++k) {
  1341.         for (int j = 0; j < Q_HEAD_IN; ++j) {
  1342.             fprintf(f, "%.17g", qq_w[k][j]);
  1343.             if (j + 1 < Q_HEAD_IN) fprintf(f, " ");
  1344.         }
  1345.         fprintf(f, "\n");
  1346.     }
  1347.  
  1348.     fprintf(f, "\n[qq_b]\n");
  1349.     for (int k = 0; k < N_CLASSES; ++k) fprintf(f, "%.17g\n", qq_b[k]);
  1350.  
  1351.     fprintf(f, "\n[qpre_w]\n");
  1352.     for (int q = 0; q < N_QUBITS; ++q) {
  1353.         for (int j = 0; j < HIDDEN_DIM; ++j) {
  1354.             fprintf(f, "%.17g", qpre_w[q][j]);
  1355.             if (j + 1 < HIDDEN_DIM) fprintf(f, " ");
  1356.         }
  1357.         fprintf(f, "\n");
  1358.     }
  1359.  
  1360.     fprintf(f, "\n[qpre_b]\n");
  1361.     for (int q = 0; q < N_QUBITS; ++q) fprintf(f, "%.17g\n", qpre_b[q]);
  1362.  
  1363.     fprintf(f, "\n[q_theta_var]\n");
  1364.     for (int p = 0; p < Q_VAR_PARAM_COUNT; ++p) fprintf(f, "%.17g\n", q_theta_var[p]);
  1365.  
  1366.     fclose(f);
  1367.     printf("Model saved to %s\n", path);
  1368. }
  1369.  
  1370. static int skip_blank_lines(FILE *f, char *line, size_t sz) {
  1371.     while (fgets(line, (int)sz, f)) {
  1372.         if (line[0] == '\n' || line[0] == '\r') continue;
  1373.         return 1;
  1374.     }
  1375.     return 0;
  1376. }
  1377.  
  1378. static int expect_section(FILE *f, const char *expected) {
  1379.     char line[256];
  1380.     if (!skip_blank_lines(f, line, sizeof(line))) return 0;
  1381.     line[strcspn(line, "\r\n")] = '\0';
  1382.     return strcmp(line, expected) == 0;
  1383. }
  1384.  
  1385. static int load_model_from_text_file(const char *path) {
  1386.     FILE *f = fopen(path, "r");
  1387.     if (!f) {
  1388.         printf("Failed to open text model file for reading: %s\n", path);
  1389.         return 0;
  1390.     }
  1391.  
  1392.     char line[256];
  1393.     if (!fgets(line, sizeof(line), f)) {
  1394.         fclose(f);
  1395.         printf("Empty text model file: %s\n", path);
  1396.         return 0;
  1397.     }
  1398.     line[strcspn(line, "\r\n")] = '\0';
  1399.     if (strcmp(line, "PPF-SQNN model dump") != 0) {
  1400.         fclose(f);
  1401.         printf("Invalid text model file header: %s\n", path);
  1402.         return 0;
  1403.     }
  1404.  
  1405.     if (!fgets(line, sizeof(line), f) || sscanf(line, "g_time_steps %d", &g_time_steps) != 1) {
  1406.         fclose(f);
  1407.         printf("Failed to read g_time_steps from %s\n", path);
  1408.         return 0;
  1409.     }
  1410.  
  1411.     if (!expect_section(f, "[fc1_w]")) { fclose(f); printf("Missing [fc1_w] in %s\n", path); return 0; }
  1412.     for (int j = 0; j < HIDDEN_DIM; ++j)
  1413.         for (int i = 0; i < INPUT_DIM; ++i)
  1414.             if (fscanf(f, "%lf", &fc1_w[j][i]) != 1) { fclose(f); printf("Read error in [fc1_w]\n"); return 0; }
  1415.  
  1416.     if (!expect_section(f, "[fc1_b]")) { fclose(f); printf("Missing [fc1_b] in %s\n", path); return 0; }
  1417.     for (int j = 0; j < HIDDEN_DIM; ++j)
  1418.         if (fscanf(f, "%lf", &fc1_b[j]) != 1) { fclose(f); printf("Read error in [fc1_b]\n"); return 0; }
  1419.  
  1420.     if (!expect_section(f, "[qc_w]")) { fclose(f); printf("Missing [qc_w] in %s\n", path); return 0; }
  1421.     for (int k = 0; k < N_CLASSES; ++k)
  1422.         for (int j = 0; j < HIDDEN_DIM; ++j)
  1423.             if (fscanf(f, "%lf", &qc_w[k][j]) != 1) { fclose(f); printf("Read error in [qc_w]\n"); return 0; }
  1424.  
  1425.     if (!expect_section(f, "[qc_b]")) { fclose(f); printf("Missing [qc_b] in %s\n", path); return 0; }
  1426.     for (int k = 0; k < N_CLASSES; ++k)
  1427.         if (fscanf(f, "%lf", &qc_b[k]) != 1) { fclose(f); printf("Read error in [qc_b]\n"); return 0; }
  1428.  
  1429.     if (!expect_section(f, "[qq_w]")) { fclose(f); printf("Missing [qq_w] in %s\n", path); return 0; }
  1430.     for (int k = 0; k < N_CLASSES; ++k)
  1431.         for (int j = 0; j < Q_HEAD_IN; ++j)
  1432.             if (fscanf(f, "%lf", &qq_w[k][j]) != 1) { fclose(f); printf("Read error in [qq_w]\n"); return 0; }
  1433.  
  1434.     if (!expect_section(f, "[qq_b]")) { fclose(f); printf("Missing [qq_b] in %s\n", path); return 0; }
  1435.     for (int k = 0; k < N_CLASSES; ++k)
  1436.         if (fscanf(f, "%lf", &qq_b[k]) != 1) { fclose(f); printf("Read error in [qq_b]\n"); return 0; }
  1437.  
  1438.     if (!expect_section(f, "[qpre_w]")) { fclose(f); printf("Missing [qpre_w] in %s\n", path); return 0; }
  1439.     for (int q = 0; q < N_QUBITS; ++q)
  1440.         for (int j = 0; j < HIDDEN_DIM; ++j)
  1441.             if (fscanf(f, "%lf", &qpre_w[q][j]) != 1) { fclose(f); printf("Read error in [qpre_w]\n"); return 0; }
  1442.  
  1443.     if (!expect_section(f, "[qpre_b]")) { fclose(f); printf("Missing [qpre_b] in %s\n", path); return 0; }
  1444.     for (int q = 0; q < N_QUBITS; ++q)
  1445.         if (fscanf(f, "%lf", &qpre_b[q]) != 1) { fclose(f); printf("Read error in [qpre_b]\n"); return 0; }
  1446.  
  1447.     if (!expect_section(f, "[q_theta_var]")) { fclose(f); printf("Missing [q_theta_var] in %s\n", path); return 0; }
  1448.     for (int p = 0; p < Q_VAR_PARAM_COUNT; ++p)
  1449.         if (fscanf(f, "%lf", &q_theta_var[p]) != 1) { fclose(f); printf("Read error in [q_theta_var]\n"); return 0; }
  1450.  
  1451.     fclose(f);
  1452.     g_model_initialized = 1;
  1453.     printf("Model loaded from %s\n", path);
  1454.     printf("Loaded implementation parameter: time steps = %d\n", g_time_steps);
  1455.     return 1;
  1456. }
  1457.  
  1458.  
  1459. static void predict_one_image_menu(const MnistSet *train_set, const MnistSet *test_set) {
  1460.     if (!g_model_initialized) {
  1461.         printf("Model is not initialized. Train or load a model first.");
  1462.         return;
  1463.     }
  1464.  
  1465.     int dataset_choice = read_int_prompt("Dataset (1=train, 2=test) [2]: ", 2);
  1466.     const MnistSet *set = (dataset_choice == 1) ? train_set : test_set;
  1467.     const char *dataset_name = (dataset_choice == 1) ? "train" : "test";
  1468.  
  1469.     printf("Select image mode:");
  1470.     printf("  1 - by index");
  1471.     printf("  2 - by true label and occurrence");
  1472.     printf("  3 - random image of a chosen digit");
  1473.     int mode = read_int_prompt("Mode [1]: ", 1);
  1474.  
  1475.     int idx = 0;
  1476.  
  1477.     if (mode == 1) {
  1478.         char prompt[256];
  1479.         snprintf(prompt, sizeof(prompt), "Image index [0..%d] [0]: ", set->num_images - 1);
  1480.         idx = read_int_prompt(prompt, 0);
  1481.         if (idx < 0) idx = 0;
  1482.         if (idx >= set->num_images) idx = set->num_images - 1;
  1483.     } else if (mode == 2) {
  1484.         int digit = read_int_prompt("True digit [0..9] [0]: ", 0);
  1485.         if (digit < 0) digit = 0;
  1486.         if (digit > 9) digit = 9;
  1487.  
  1488.         int total = count_label_occurrences(set, digit);
  1489.         if (total <= 0) {
  1490.             printf("No images with digit %d in %s dataset.", digit, dataset_name);
  1491.             return;
  1492.         }
  1493.  
  1494.         char prompt[256];
  1495.         snprintf(prompt, sizeof(prompt), "Occurrence [1..%d] [1]: ", total);
  1496.         int occ = read_int_prompt(prompt, 1);
  1497.         if (occ < 1) occ = 1;
  1498.         if (occ > total) occ = total;
  1499.  
  1500.         idx = find_nth_occurrence_of_label(set, digit, occ - 1);
  1501.         printf("Chosen occurrence: %d of %d", occ, total);
  1502.     } else if (mode == 3) {
  1503.         int digit = read_int_prompt("True digit [0..9] [0]: ", 0);
  1504.         if (digit < 0) digit = 0;
  1505.         if (digit > 9) digit = 9;
  1506.  
  1507.         int total = count_label_occurrences(set, digit);
  1508.         if (total <= 0) {
  1509.             printf("No images with digit %d in %s dataset.", digit, dataset_name);
  1510.             return;
  1511.         }
  1512.  
  1513.         idx = choose_random_index_of_label(set, digit);
  1514.         int occ = 1;
  1515.         for (int i = 0; i < set->num_images; ++i) {
  1516.             if (set->labels[i] == digit) {
  1517.                 if (i == idx) break;
  1518.                 ++occ;
  1519.             }
  1520.         }
  1521.         printf("Random occurrence chosen: %d of %d", occ, total);
  1522.     } else {
  1523.         printf("Unknown mode.");
  1524.         return;
  1525.     }
  1526.  
  1527.     double noise_percent = read_double_prompt("Noise percent [0..100] [0]: ", 0.0);
  1528.     if (noise_percent < 0.0) noise_percent = 0.0;
  1529.     if (noise_percent > 100.0) noise_percent = 100.0;
  1530.  
  1531.     printf("Select noise type:");
  1532.     printf("  1 - uniform");
  1533.     printf("  2 - gaussian");
  1534.     int noise_type = read_int_prompt("Noise type [1]: ", 1);
  1535.  
  1536.     print_ascii_digit(set, idx);
  1537.  
  1538.     double x_clean[INPUT_DIM];
  1539.     double x_noisy[INPUT_DIM];
  1540.     get_image_vector(set, idx, x_clean);
  1541.     for (int i = 0; i < INPUT_DIM; ++i) {
  1542.         x_noisy[i] = x_clean[i];
  1543.     }
  1544.  
  1545.     if (noise_percent > 0.0) {
  1546.         if (noise_type == 2) {
  1547.             add_noise_gaussian(x_noisy, noise_percent);
  1548.         } else {
  1549.             add_noise_uniform(x_noisy, noise_percent);
  1550.         }
  1551.     }
  1552.  
  1553.     print_ascii_vector28(x_noisy, set->labels[idx], "ASCII preview after noise");
  1554.  
  1555.     double fc1_pre[HIDDEN_DIM], fc1_act[HIDDEN_DIM], pooled[HIDDEN_DIM];
  1556.     double q_in[N_QUBITS], z[N_QUBITS];
  1557.     double qc[N_CLASSES], qq[N_CLASSES], qh[N_CLASSES], probs[N_CLASSES];
  1558.     model_forward(x_noisy, fc1_pre, fc1_act, pooled, q_in, z, qc, qq, qh, probs);
  1559.  
  1560.     int pred = argmax10(probs);
  1561.     int label = set->labels[idx];
  1562.     int pred_ref_idx = choose_random_index_of_label(set, pred);
  1563.  
  1564.     if (save_pgm_from_vector("predict_input_noisy.pgm", x_noisy)) {
  1565.         printf("Saved noisy input image to predict_input_noisy.pgm");
  1566.     } else {
  1567.         printf("Failed to save noisy input image.");
  1568.     }
  1569.  
  1570.     if (pred_ref_idx >= 0 && save_pgm_from_dataset_image("predict_predicted_class_reference.pgm", set, pred_ref_idx)) {
  1571.         printf("Saved predicted-class reference image to predict_predicted_class_reference.pgm");
  1572.         printf("Note: this second file is a clean reference example of the predicted class, not a reconstruction.");
  1573.     } else {
  1574.         printf("Failed to save predicted-class reference image.");
  1575.     }
  1576.  
  1577.     printf("Prediction:");
  1578.     printf("  dataset         : %s", dataset_name);
  1579.     printf("  image index     : %d", idx);
  1580.     printf("  true label      : %d", label);
  1581.     printf("  predicted label : %d", pred);
  1582.     printf("  noise percent   : %.2f", noise_percent);
  1583.     printf("  noise type      : %s", (noise_type == 2) ? "gaussian" : "uniform");
  1584.     printf("  confidence      : %.6f", probs[pred]);
  1585.     printf("  loss            : %.6f", nll_one(probs, label));
  1586.  
  1587.     printf("Class probabilities:");
  1588.     for (int k = 0; k < N_CLASSES; ++k) {
  1589.         printf("  [%d] %.10f", k, probs[k]);
  1590.     }
  1591. }
  1592.  
  1593. static void train_model_menu(const MnistSet *train_set, const MnistSet *test_set, int *train_idx, int *test_idx) {
  1594.     if (!g_model_initialized) {
  1595.         init_weights();
  1596.         g_model_initialized = 1;
  1597.     }
  1598.  
  1599.     int epochs = read_int_prompt("Epochs [140]: ", EPOCHS);
  1600.     if (epochs <= 0) epochs = EPOCHS;
  1601.  
  1602.     int new_steps = read_int_prompt("Time steps for spiking simulation [current value]: ", g_time_steps);
  1603.     if (new_steps > 0) g_time_steps = new_steps;
  1604.  
  1605.     int adam_step = 0;
  1606.     for (int epoch = 0; epoch < epochs; ++epoch) {
  1607.         double train_loss = train_one_epoch(train_set, train_idx, epoch, &adam_step);
  1608.         printf("Epoch %d/%d finished, avg train loss = %.6f\n", epoch + 1, epochs, train_loss);
  1609.  
  1610.         if ((epoch + 1) % 10 == 0 || epoch == epochs - 1) {
  1611.             double test_loss = 0.0;
  1612.             double acc = evaluate_accuracy(test_set, test_idx, TEST_SUBSET, &test_loss);
  1613.             printf("  Test subset accuracy = %.4f, avg loss = %.6f\n", acc, test_loss);
  1614.         }
  1615.         printf("\n");
  1616.     }
  1617.  
  1618.     save_model_to_text_file("ppf_sqnn_model.txt");
  1619. }
  1620.  
  1621. static void print_menu(void) {
  1622.     printf("\n================ PPF-SQNN menu ================\n");
  1623.     printf("1 - Train model\n");
  1624.     printf("2 - Predict one image\n");
  1625.     printf("3 - Save model\n");
  1626.     printf("4 - Load model\n");
  1627.     printf("5 - Set time steps (current: %d)\n", g_time_steps);
  1628.     printf("0 - Exit\n");
  1629.     printf("===============================================\n");
  1630. }
  1631.  
  1632. int main(void) {
  1633.     rng_state = (unsigned int)time(NULL);
  1634.  
  1635.     printf("PPF-SQNN exact-replication project, menu build\n");
  1636.     printf("- Input: 28x28 -> 784\n");
  1637.     printf("- Reduced dimension: 128\n");
  1638.     printf("- Classes: 10\n");
  1639.     printf("- Qubits: 5\n");
  1640.     printf("- Train subset: 8000\n");
  1641.     printf("- Test subset: 10000\n");
  1642.     printf("- Batch size: 32\n");
  1643.     printf("- Default epochs: 140\n");
  1644.     printf("- xi: 0.8\n");
  1645.     printf("- Adam: lr=0.001 beta1=0.90 beta2=0.99 wd=0.0\n");
  1646.     printf("- Stage-2 time steps: runtime implementation parameter (default %d)\n\n", g_time_steps);
  1647.  
  1648.     char train_img[1024], train_lbl[1024], test_img[1024], test_lbl[1024];
  1649.     if (!find_mnist_files(train_img, sizeof(train_img), train_lbl, sizeof(train_lbl), test_img, sizeof(test_img), test_lbl, sizeof(test_lbl))) {
  1650.         printf("Could not locate all four MNIST IDX files.\n");
  1651.         return 1;
  1652.     }
  1653.  
  1654.     printf("Using MNIST files:\n");
  1655.     printf("  %s\n", train_img);
  1656.     printf("  %s\n", train_lbl);
  1657.     printf("  %s\n", test_img);
  1658.     printf("  %s\n\n", test_lbl);
  1659.  
  1660.     MnistSet train_set, test_set;
  1661.     if (!load_mnist_set(train_img, train_lbl, &train_set) || !load_mnist_set(test_img, test_lbl, &test_set)) {
  1662.         printf("Failed to load MNIST data.\n");
  1663.         return 1;
  1664.     }
  1665.  
  1666.     int *train_idx = (int *)malloc((size_t)train_set.num_images * sizeof(int));
  1667.     int *test_idx = (int *)malloc((size_t)test_set.num_images * sizeof(int));
  1668.     if (!train_idx || !test_idx) {
  1669.         printf("Allocation failure.\n");
  1670.         free_mnist_set(&train_set);
  1671.         free_mnist_set(&test_set);
  1672.         free(train_idx);
  1673.         free(test_idx);
  1674.         return 1;
  1675.     }
  1676.  
  1677.     for (int i = 0; i < train_set.num_images; ++i) train_idx[i] = i;
  1678.     for (int i = 0; i < test_set.num_images; ++i) test_idx[i] = i;
  1679.     shuffle_int(train_idx, train_set.num_images);
  1680.     shuffle_int(test_idx, test_set.num_images);
  1681.  
  1682.     if (file_exists("ppf_sqnn_model.txt")) {
  1683.         printf("Found existing model file ppf_sqnn_model.txt\n");
  1684.         printf("Trying to load it...\n");
  1685.         load_model_from_text_file("ppf_sqnn_model.txt");
  1686.     } else {
  1687.         init_weights();
  1688.         g_model_initialized = 1;
  1689.     }
  1690.  
  1691.     for (;;) {
  1692.         print_menu();
  1693.         int choice = read_int_prompt("Choose action: ", -1);
  1694.  
  1695.         if (choice == 0) {
  1696.             break;
  1697.         } else if (choice == 1) {
  1698.             train_model_menu(&train_set, &test_set, train_idx, test_idx);
  1699.             wait_enter();
  1700.         } else if (choice == 2) {
  1701.             predict_one_image_menu(&train_set, &test_set);
  1702.             wait_enter();
  1703.         } else if (choice == 3) {
  1704.             save_model_to_text_file("ppf_sqnn_model.txt");
  1705.             wait_enter();
  1706.         } else if (choice == 4) {
  1707.             load_model_from_text_file("ppf_sqnn_model.txt");
  1708.             wait_enter();
  1709.         } else if (choice == 5) {
  1710.             int new_steps = read_int_prompt("New time steps value: ", g_time_steps);
  1711.             if (new_steps > 0) {
  1712.                 g_time_steps = new_steps;
  1713.                 printf("Time steps updated to %d\n", g_time_steps);
  1714.             } else {
  1715.                 printf("Time steps value must be positive.\n");
  1716.             }
  1717.             wait_enter();
  1718.         } else {
  1719.             printf("Unknown menu item.\n");
  1720.             wait_enter();
  1721.         }
  1722.     }
  1723.  
  1724.     free(train_idx);
  1725.     free(test_idx);
  1726.     free_mnist_set(&train_set);
  1727.     free_mnist_set(&test_set);
  1728.     return 0;
  1729. }
  1730.  
Advertisement
Add Comment
Please, Sign In to add comment