faster-diffbloch 0.1.0__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
@@ -0,0 +1,701 @@
1
+ #include <stdint.h>
2
+ #include <stdbool.h>
3
+ #include <stdio.h>
4
+ #include <stdlib.h>
5
+ #include <string.h>
6
+
7
+ /* Flow runtime helpers */
8
+ typedef struct flow_temp_node { struct flow_temp_node* next; } flow_temp_node;
9
+ static flow_temp_node* flow_temp_head = NULL;
10
+ static int flow_temp_atexit_set = 0;
11
+ __attribute__((unused)) static void flow_temp_free_all(void) {
12
+ while (flow_temp_head) {
13
+ flow_temp_node* n = flow_temp_head;
14
+ flow_temp_head = n->next;
15
+ free(n);
16
+ }
17
+ }
18
+ __attribute__((unused)) static void* flow_temp_alloc(size_t nbytes) {
19
+ flow_temp_node* node = (flow_temp_node*)malloc(sizeof(flow_temp_node) + nbytes);
20
+ if (!node) return NULL;
21
+ node->next = flow_temp_head;
22
+ flow_temp_head = node;
23
+ if (!flow_temp_atexit_set) {
24
+ flow_temp_atexit_set = 1;
25
+ atexit(flow_temp_free_all);
26
+ }
27
+ return (void*)(node + 1);
28
+ }
29
+ #ifndef FLOW_DIAG
30
+ #define FLOW_DIAG(msg) fprintf(stderr, "%s", (msg))
31
+ #endif
32
+ #ifndef FLOW_LOG
33
+ #define FLOW_LOG(fmt, ...) printf(fmt, __VA_ARGS__)
34
+ #endif
35
+ #ifndef FLOW_LOG_EMPTY
36
+ #define FLOW_LOG_EMPTY(fmt) printf(fmt)
37
+ #endif
38
+ static char* flow_strcat(const char* a, const char* b) {
39
+ size_t la = strlen(a ? a : ""), lb = strlen(b ? b : "");
40
+ char* r = (char*)flow_temp_alloc(la + lb + 1);
41
+ if (!r) return NULL;
42
+ if (la) memcpy(r, a, la);
43
+ if (lb) memcpy(r + la, b, lb);
44
+ r[la + lb] = '\0';
45
+ return r;
46
+ }
47
+
48
+ #define __flow_in_arr(arr, val) __extension__ ({ \
49
+ int _found = 0; \
50
+ size_t _n = sizeof(arr)/sizeof((arr)[0]); \
51
+ for (size_t _i = 0; _i < _n; _i++) { \
52
+ if ((arr)[_i] == (val)) { _found = 1; break; } \
53
+ } _found; })
54
+
55
+ #include <math.h>
56
+ #include <complex.h>
57
+
58
+ void* _ui_state = NULL;
59
+
60
+ static inline float i32_to_f32(int32_t v) { return (float)v; }
61
+
62
+ /* Host stub for @gpu kernels (device codegen replaces this). */
63
+ static inline int32_t gpu_thread_id(void) { return 0; }
64
+
65
+ typedef enum {
66
+ Adjoint_Blocked,
67
+ Adjoint_Dense
68
+ } Adjoint_Tag;
69
+
70
+ typedef struct {
71
+ Adjoint_Tag tag;
72
+ } Adjoint;
73
+
74
+ typedef struct CMat CMat;
75
+ typedef struct Scalars Scalars;
76
+ typedef struct c128 c128;
77
+ typedef struct c64 c64;
78
+
79
+ /* Spans: borrowed {pointer, length} views */
80
+ typedef struct { double *data; int64_t len; } flow_span_mut_f64;
81
+ typedef struct { float complex *data; int64_t len; } flow_span_mut_c64;
82
+ typedef struct { double complex *data; int64_t len; } flow_span_mut_c128;
83
+ typedef struct { const double complex *data; int64_t len; } flow_span_const_c128;
84
+ typedef struct { const int32_t *data; int64_t len; } flow_span_const_i32;
85
+ typedef struct { const double *data; int64_t len; } flow_span_const_f64;
86
+ typedef struct { int32_t *data; int64_t len; } flow_span_mut_i32;
87
+
88
+ struct CMat {
89
+ flow_span_mut_c64 data;
90
+ int32_t n;
91
+ };
92
+
93
+ struct Scalars {
94
+ float complex* one;
95
+ float complex* zero;
96
+ float complex* beta_one;
97
+ };
98
+
99
+ struct c128 {
100
+ };
101
+
102
+ struct c64 {
103
+ };
104
+
105
+ void cblas_cgemm(int32_t order, int32_t ta, int32_t tb, int32_t m, int32_t n, int32_t k, float complex* alpha, float complex* A, int32_t lda, float complex* B, int32_t ldb, float complex* beta, float complex* C, int32_t ldc);
106
+ double theta18(void);
107
+ double ln2(void);
108
+ CMat cmat_new_i32(int32_t n);
109
+ void cmat_free_CMat(CMat m);
110
+ void cmat_copy_into_CMat_CMat(CMat dst, CMat src);
111
+ Scalars scalars_new(void);
112
+ void scalars_free_Scalars(Scalars s);
113
+ void gemm_CMat_CMat_CMat_Scalars(CMat a, CMat b, CMat out, Scalars s);
114
+ double norm1_CMat(CMat m);
115
+ int32_t squarings_for_f64(double norm);
116
+ float halving_factor_i32(int32_t s);
117
+ flow_span_mut_f64 taylor_coefficients(void);
118
+ void expm_CMat_CMat(CMat m, CMat out);
119
+ void conjugate_transpose_CMat_CMat(CMat m, CMat out);
120
+ void expm_backward_dense_CMat_CMat_CMat(CMat m, CMat ebar, CMat mbar);
121
+ void pair_mul_CMat_CMat_CMat_CMat_CMat_CMat_Scalars(CMat y1, CMat l1, CMat y2, CMat l2, CMat yo, CMat lo, Scalars s);
122
+ void pair_sq_CMat_CMat_CMat_CMat_Scalars(CMat y, CMat l, CMat yo, CMat lo, Scalars s);
123
+ double pair_norm1_CMat_CMat(CMat y, CMat l);
124
+ void expm_backward_blocked_CMat_CMat_CMat(CMat m, CMat ebar, CMat mbar);
125
+ void expm_backward_Adjoint_CMat_CMat_CMat(Adjoint method, CMat m, CMat ebar, CMat mbar);
126
+ bool shapes_ok_i32_i32(int32_t n, int32_t count);
127
+ int32_t metal_matrix_exp_c64(float complex* a, int32_t n, int32_t count, float complex* out);
128
+ int32_t metal_matrix_exp_backward_c64(float complex* m, float complex* ebar, int32_t n, int32_t count, int32_t dense, float complex* mbar);
129
+ int32_t bridge_abi_version(void);
130
+ int32_t bridge_matrix_exp_gpu_ptr_c64_i32_i32_ptr_c64(float complex* a, int32_t n, int32_t count, float complex* out);
131
+ int32_t bridge_matrix_exp_backward_gpu_ptr_c64_ptr_c64_i32_i32_i32_ptr_c64(float complex* m, float complex* ebar, int32_t n, int32_t count, int32_t dense, float complex* mbar);
132
+ int32_t bridge_matrix_exp_ptr_c64_i32_i32_ptr_c64(float complex* a, int32_t n, int32_t count, float complex* out);
133
+ int32_t bridge_matrix_exp_backward_ptr_c64_ptr_c64_i32_i32_i32_ptr_c64(float complex* m, float complex* ebar, int32_t n, int32_t count, int32_t dense, float complex* mbar);
134
+ double bridge_norm1_ptr_c64_i32(float complex* a, int32_t n);
135
+ flow_span_mut_c128 zeroed_c128_i32(int32_t n);
136
+ void gather_into_span_const_c128_span_const_i32_i32_span_const_i32_span_mut_c128(flow_span_const_c128 fgb, flow_span_const_i32 source, int32_t buffer_size, flow_span_const_i32 destination, flow_span_mut_c128 gathered);
137
+ int32_t bridge_structure_matrix_ptr_c128_i32_ptr_i32_i32_ptr_i32_i32_ptr_f64_ptr_f64_i32_f64_i32_ptr_c128(double complex* fgb, int32_t n_grid, int32_t* source, int32_t buffer_size, int32_t* destination, int32_t n, double* mii, double* diagonal, int32_t n_batch, double prefactor, int32_t absorption, double complex* out);
138
+ int32_t bridge_structure_matrix_backward_ptr_c128_ptr_i32_i32_ptr_f64_i32_f64_i32_ptr_i32_i32_i32_ptr_c128(double complex* abar, int32_t* destination, int32_t n, double* mii, int32_t n_batch, double prefactor, int32_t absorption, int32_t* source, int32_t buffer_size, int32_t n_grid, double complex* fbar);
139
+
140
+ static const int32_t ROW_MAJOR = 101;
141
+ static const int32_t NO_TRANS = 111;
142
+
143
+
144
+
145
+
146
+
147
+
148
+
149
+
150
+
151
+ double theta18(void) {
152
+ return 3.010066362817634;
153
+ }
154
+
155
+ double ln2(void) {
156
+ return 0.6931471805599453;
157
+ }
158
+
159
+ CMat cmat_new_i32(int32_t n) {
160
+ int64_t bytes = ((((int64_t)(n)) * ((int64_t)(n))) * 8);
161
+ float complex* raw = (float complex*)(malloc(bytes));
162
+ memset(raw, 0, bytes);
163
+ return (CMat){ .data = ((flow_span_mut_c64){ .data = (((raw)) + (0)), .len = (int64_t)(((n * n)) - (0)) }), .n = n };
164
+ }
165
+
166
+ void cmat_free_CMat(CMat m) {
167
+ free(m.data.data);
168
+ }
169
+
170
+ void cmat_copy_into_CMat_CMat(CMat dst, CMat src) {
171
+ memcpy(dst.data.data, src.data.data, ((((int64_t)(dst.n)) * ((int64_t)(dst.n))) * 8));
172
+ }
173
+
174
+ Scalars scalars_new(void) {
175
+ float complex* one = (float complex*)(malloc(8));
176
+ float complex* zero = (float complex*)(malloc(8));
177
+ float complex* beta_one = (float complex*)(malloc(8));
178
+ one[0] = ((float)(1.0) + (float)(0.0) * I);
179
+ zero[0] = ((float)(0.0) + (float)(0.0) * I);
180
+ beta_one[0] = ((float)(1.0) + (float)(0.0) * I);
181
+ return (Scalars){ .one = one, .zero = zero, .beta_one = beta_one };
182
+ }
183
+
184
+ void scalars_free_Scalars(Scalars s) {
185
+ free(s.one);
186
+ free(s.zero);
187
+ free(s.beta_one);
188
+ }
189
+
190
+ void gemm_CMat_CMat_CMat_Scalars(CMat a, CMat b, CMat out, Scalars s) {
191
+ cblas_cgemm(ROW_MAJOR, NO_TRANS, NO_TRANS, a.n, a.n, a.n, s.one, a.data.data, a.n, b.data.data, b.n, s.zero, out.data.data, out.n);
192
+ }
193
+
194
+ double norm1_CMat(CMat m) {
195
+ double worst = 0.0;
196
+ int32_t __flow_step_1 = 1;
197
+ for (int32_t j = 0; (0 <= m.n) ? j < m.n : j > m.n; j += (0 <= m.n) ? 1 : -1) {
198
+ double column = 0.0;
199
+ int32_t __flow_step_2 = 1;
200
+ for (int32_t i = 0; (0 <= m.n) ? i < m.n : i > m.n; i += (0 <= m.n) ? 1 : -1) {
201
+ column = (column + ((double)(cabs((m.data).data[((i * m.n) + j)]))));
202
+ }
203
+ if (column > worst) {
204
+ worst = column;
205
+ }
206
+ }
207
+ return worst;
208
+ }
209
+
210
+ int32_t squarings_for_f64(double norm) {
211
+ double threshold = theta18();
212
+ if (norm <= threshold) {
213
+ return 0;
214
+ }
215
+ return ((int32_t)(ceil(log2((norm / threshold)))));
216
+ }
217
+
218
+ float halving_factor_i32(int32_t s) {
219
+ return ((float)(exp(((-((double)(s))) * ln2()))));
220
+ }
221
+
222
+ flow_span_mut_f64 taylor_coefficients(void) {
223
+ double* raw = (double*)(malloc((19 * 8)));
224
+ raw[0] = 1.0;
225
+ int32_t __flow_step_3 = 1;
226
+ for (int32_t k = 1; (1 <= 19) ? k < 19 : k > 19; k += (1 <= 19) ? 1 : -1) {
227
+ raw[k] = (raw[(k - 1)] / ((double)(k)));
228
+ }
229
+ return ((flow_span_mut_f64){ .data = (((raw)) + (0)), .len = (int64_t)((19) - (0)) });
230
+ }
231
+
232
+ void expm_CMat_CMat(CMat m, CMat out) {
233
+ int32_t n = m.n;
234
+ int32_t count = (n * n);
235
+ Scalars s = scalars_new();
236
+ int32_t squarings = squarings_for_f64(norm1_CMat(m));
237
+ if (squarings > 0) {
238
+ float factor = halving_factor_i32(squarings);
239
+ int32_t __flow_step_4 = 1;
240
+ for (int32_t i = 0; (0 <= count) ? i < count : i > count; i += (0 <= count) ? 1 : -1) {
241
+ (m.data).data[i] = ((m.data).data[i] * ((float)(factor) + (float)(0.0) * I));
242
+ }
243
+ }
244
+ flow_span_mut_f64 __flow_span_init_1 = taylor_coefficients();
245
+ flow_span_mut_f64 coefficients = __flow_span_init_1;
246
+ CMat m2 = cmat_new_i32(n);
247
+ CMat m3 = cmat_new_i32(n);
248
+ gemm_CMat_CMat_CMat_Scalars(m, m, m2, s);
249
+ gemm_CMat_CMat_CMat_Scalars(m2, m, m3, s);
250
+ CMat acc = cmat_new_i32(n);
251
+ CMat tmp = cmat_new_i32(n);
252
+ int32_t j = 5;
253
+ while (j >= 0) {
254
+ if (j == 5) {
255
+ int32_t __flow_step_5 = 1;
256
+ for (int32_t i = 0; (0 <= count) ? i < count : i > count; i += (0 <= count) ? 1 : -1) {
257
+ (acc.data).data[i] = ((float)(0.0) + (float)(0.0) * I);
258
+ }
259
+ } else {
260
+ gemm_CMat_CMat_CMat_Scalars(acc, m3, tmp, s);
261
+ cmat_copy_into_CMat_CMat(acc, tmp);
262
+ }
263
+ float c1 = ((float)((coefficients).data[((j * 3) + 1)]));
264
+ float c2 = ((float)((coefficients).data[((j * 3) + 2)]));
265
+ int32_t __flow_step_6 = 1;
266
+ for (int32_t i = 0; (0 <= count) ? i < count : i > count; i += (0 <= count) ? 1 : -1) {
267
+ (acc.data).data[i] = (((acc.data).data[i] + ((m.data).data[i] * ((float)(c1) + (float)(0.0) * I))) + ((m2.data).data[i] * ((float)(c2) + (float)(0.0) * I)));
268
+ }
269
+ float c0 = ((float)((coefficients).data[(j * 3)]));
270
+ int32_t __flow_step_7 = 1;
271
+ for (int32_t i = 0; (0 <= n) ? i < n : i > n; i += (0 <= n) ? 1 : -1) {
272
+ (acc.data).data[((i * n) + i)] = ((acc.data).data[((i * n) + i)] + ((float)(c0) + (float)(0.0) * I));
273
+ }
274
+ j = (j - 1);
275
+ }
276
+ CMat src = acc;
277
+ CMat dst = tmp;
278
+ int32_t __flow_step_8 = 1;
279
+ for (int32_t k = 0; (0 <= squarings) ? k < squarings : k > squarings; k += (0 <= squarings) ? 1 : -1) {
280
+ gemm_CMat_CMat_CMat_Scalars(src, src, dst, s);
281
+ CMat swap = src;
282
+ src = dst;
283
+ dst = swap;
284
+ }
285
+ cmat_copy_into_CMat_CMat(out, src);
286
+ cmat_free_CMat(m2);
287
+ cmat_free_CMat(m3);
288
+ cmat_free_CMat(acc);
289
+ cmat_free_CMat(tmp);
290
+ free(coefficients.data);
291
+ scalars_free_Scalars(s);
292
+ }
293
+
294
+ void conjugate_transpose_CMat_CMat(CMat m, CMat out) {
295
+ int32_t __flow_step_9 = 1;
296
+ for (int32_t i = 0; (0 <= m.n) ? i < m.n : i > m.n; i += (0 <= m.n) ? 1 : -1) {
297
+ int32_t __flow_step_10 = 1;
298
+ for (int32_t j = 0; (0 <= m.n) ? j < m.n : j > m.n; j += (0 <= m.n) ? 1 : -1) {
299
+ float complex z = (m.data).data[((j * m.n) + i)];
300
+ (out.data).data[((i * m.n) + j)] = ((float)(creal(z)) + (float)((-cimag(z))) * I);
301
+ }
302
+ }
303
+ }
304
+
305
+ void expm_backward_dense_CMat_CMat_CMat(CMat m, CMat ebar, CMat mbar) {
306
+ int32_t n = m.n;
307
+ int32_t wide = (2 * n);
308
+ CMat block = cmat_new_i32(wide);
309
+ CMat result = cmat_new_i32(wide);
310
+ int32_t __flow_step_11 = 1;
311
+ for (int32_t i = 0; (0 <= n) ? i < n : i > n; i += (0 <= n) ? 1 : -1) {
312
+ int32_t __flow_step_12 = 1;
313
+ for (int32_t j = 0; (0 <= n) ? j < n : j > n; j += (0 <= n) ? 1 : -1) {
314
+ float complex z = (m.data).data[((j * n) + i)];
315
+ float complex adjointed = ((float)(creal(z)) + (float)((-cimag(z))) * I);
316
+ (block.data).data[((i * wide) + j)] = adjointed;
317
+ (block.data).data[(((i + n) * wide) + (j + n))] = adjointed;
318
+ (block.data).data[((i * wide) + (j + n))] = (ebar.data).data[((i * n) + j)];
319
+ }
320
+ }
321
+ expm_CMat_CMat(block, result);
322
+ int32_t __flow_step_13 = 1;
323
+ for (int32_t i = 0; (0 <= n) ? i < n : i > n; i += (0 <= n) ? 1 : -1) {
324
+ int32_t __flow_step_14 = 1;
325
+ for (int32_t j = 0; (0 <= n) ? j < n : j > n; j += (0 <= n) ? 1 : -1) {
326
+ (mbar.data).data[((i * n) + j)] = (result.data).data[((i * wide) + (j + n))];
327
+ }
328
+ }
329
+ cmat_free_CMat(block);
330
+ cmat_free_CMat(result);
331
+ }
332
+
333
+ void pair_mul_CMat_CMat_CMat_CMat_CMat_CMat_Scalars(CMat y1, CMat l1, CMat y2, CMat l2, CMat yo, CMat lo, Scalars s) {
334
+ gemm_CMat_CMat_CMat_Scalars(y1, y2, yo, s);
335
+ gemm_CMat_CMat_CMat_Scalars(y1, l2, lo, s);
336
+ cblas_cgemm(ROW_MAJOR, NO_TRANS, NO_TRANS, l1.n, l1.n, l1.n, s.one, l1.data.data, l1.n, y2.data.data, y2.n, s.beta_one, lo.data.data, lo.n);
337
+ }
338
+
339
+ void pair_sq_CMat_CMat_CMat_CMat_Scalars(CMat y, CMat l, CMat yo, CMat lo, Scalars s) {
340
+ gemm_CMat_CMat_CMat_Scalars(y, y, yo, s);
341
+ gemm_CMat_CMat_CMat_Scalars(y, l, lo, s);
342
+ cblas_cgemm(ROW_MAJOR, NO_TRANS, NO_TRANS, l.n, l.n, l.n, s.one, l.data.data, l.n, y.data.data, y.n, s.beta_one, lo.data.data, lo.n);
343
+ }
344
+
345
+ double pair_norm1_CMat_CMat(CMat y, CMat l) {
346
+ double worst = 0.0;
347
+ int32_t __flow_step_15 = 1;
348
+ for (int32_t j = 0; (0 <= y.n) ? j < y.n : j > y.n; j += (0 <= y.n) ? 1 : -1) {
349
+ double column_y = 0.0;
350
+ double column_l = 0.0;
351
+ int32_t __flow_step_16 = 1;
352
+ for (int32_t i = 0; (0 <= y.n) ? i < y.n : i > y.n; i += (0 <= y.n) ? 1 : -1) {
353
+ column_y = (column_y + ((double)(cabs((y.data).data[((i * y.n) + j)]))));
354
+ column_l = (column_l + ((double)(cabs((l.data).data[((i * y.n) + j)]))));
355
+ }
356
+ if (column_y > worst) {
357
+ worst = column_y;
358
+ }
359
+ if ((column_y + column_l) > worst) {
360
+ worst = (column_y + column_l);
361
+ }
362
+ }
363
+ return worst;
364
+ }
365
+
366
+ void expm_backward_blocked_CMat_CMat_CMat(CMat m, CMat ebar, CMat mbar) {
367
+ int32_t n = m.n;
368
+ int32_t count = (n * n);
369
+ Scalars s = scalars_new();
370
+ CMat y = cmat_new_i32(n);
371
+ CMat l = cmat_new_i32(n);
372
+ conjugate_transpose_CMat_CMat(m, y);
373
+ cmat_copy_into_CMat_CMat(l, ebar);
374
+ int32_t squarings = squarings_for_f64(pair_norm1_CMat_CMat(y, l));
375
+ if (squarings > 0) {
376
+ float factor = halving_factor_i32(squarings);
377
+ int32_t __flow_step_17 = 1;
378
+ for (int32_t i = 0; (0 <= count) ? i < count : i > count; i += (0 <= count) ? 1 : -1) {
379
+ (y.data).data[i] = ((y.data).data[i] * ((float)(factor) + (float)(0.0) * I));
380
+ (l.data).data[i] = ((l.data).data[i] * ((float)(factor) + (float)(0.0) * I));
381
+ }
382
+ }
383
+ flow_span_mut_f64 __flow_span_init_2 = taylor_coefficients();
384
+ flow_span_mut_f64 coefficients = __flow_span_init_2;
385
+ CMat y2 = cmat_new_i32(n);
386
+ CMat l2 = cmat_new_i32(n);
387
+ CMat y3 = cmat_new_i32(n);
388
+ CMat l3 = cmat_new_i32(n);
389
+ pair_sq_CMat_CMat_CMat_CMat_Scalars(y, l, y2, l2, s);
390
+ pair_mul_CMat_CMat_CMat_CMat_CMat_CMat_Scalars(y2, l2, y, l, y3, l3, s);
391
+ CMat acc_y = cmat_new_i32(n);
392
+ CMat acc_l = cmat_new_i32(n);
393
+ CMat tmp_y = cmat_new_i32(n);
394
+ CMat tmp_l = cmat_new_i32(n);
395
+ int32_t j = 5;
396
+ while (j >= 0) {
397
+ if (j == 5) {
398
+ int32_t __flow_step_18 = 1;
399
+ for (int32_t i = 0; (0 <= count) ? i < count : i > count; i += (0 <= count) ? 1 : -1) {
400
+ (acc_y.data).data[i] = ((float)(0.0) + (float)(0.0) * I);
401
+ (acc_l.data).data[i] = ((float)(0.0) + (float)(0.0) * I);
402
+ }
403
+ } else {
404
+ pair_mul_CMat_CMat_CMat_CMat_CMat_CMat_Scalars(acc_y, acc_l, y3, l3, tmp_y, tmp_l, s);
405
+ cmat_copy_into_CMat_CMat(acc_y, tmp_y);
406
+ cmat_copy_into_CMat_CMat(acc_l, tmp_l);
407
+ }
408
+ float c1 = ((float)((coefficients).data[((j * 3) + 1)]));
409
+ float c2 = ((float)((coefficients).data[((j * 3) + 2)]));
410
+ int32_t __flow_step_19 = 1;
411
+ for (int32_t i = 0; (0 <= count) ? i < count : i > count; i += (0 <= count) ? 1 : -1) {
412
+ (acc_y.data).data[i] = (((acc_y.data).data[i] + ((y.data).data[i] * ((float)(c1) + (float)(0.0) * I))) + ((y2.data).data[i] * ((float)(c2) + (float)(0.0) * I)));
413
+ (acc_l.data).data[i] = (((acc_l.data).data[i] + ((l.data).data[i] * ((float)(c1) + (float)(0.0) * I))) + ((l2.data).data[i] * ((float)(c2) + (float)(0.0) * I)));
414
+ }
415
+ float c0 = ((float)((coefficients).data[(j * 3)]));
416
+ int32_t __flow_step_20 = 1;
417
+ for (int32_t i = 0; (0 <= n) ? i < n : i > n; i += (0 <= n) ? 1 : -1) {
418
+ (acc_y.data).data[((i * n) + i)] = ((acc_y.data).data[((i * n) + i)] + ((float)(c0) + (float)(0.0) * I));
419
+ }
420
+ j = (j - 1);
421
+ }
422
+ int32_t __flow_step_21 = 1;
423
+ for (int32_t k = 0; (0 <= squarings) ? k < squarings : k > squarings; k += (0 <= squarings) ? 1 : -1) {
424
+ pair_sq_CMat_CMat_CMat_CMat_Scalars(acc_y, acc_l, tmp_y, tmp_l, s);
425
+ cmat_copy_into_CMat_CMat(acc_y, tmp_y);
426
+ cmat_copy_into_CMat_CMat(acc_l, tmp_l);
427
+ }
428
+ cmat_copy_into_CMat_CMat(mbar, acc_l);
429
+ cmat_free_CMat(y);
430
+ cmat_free_CMat(l);
431
+ cmat_free_CMat(y2);
432
+ cmat_free_CMat(l2);
433
+ cmat_free_CMat(y3);
434
+ cmat_free_CMat(l3);
435
+ cmat_free_CMat(acc_y);
436
+ cmat_free_CMat(acc_l);
437
+ cmat_free_CMat(tmp_y);
438
+ cmat_free_CMat(tmp_l);
439
+ free(coefficients.data);
440
+ scalars_free_Scalars(s);
441
+ }
442
+
443
+ void expm_backward_Adjoint_CMat_CMat_CMat(Adjoint method, CMat m, CMat ebar, CMat mbar) {
444
+ { // match block
445
+ if ((method.tag) == Adjoint_Blocked) {
446
+ expm_backward_blocked_CMat_CMat_CMat(m, ebar, mbar);
447
+ } else { // exhaustive
448
+ expm_backward_dense_CMat_CMat_CMat(m, ebar, mbar);
449
+ }
450
+ } // end match
451
+ }
452
+
453
+ bool shapes_ok_i32_i32(int32_t n, int32_t count) {
454
+ if (n <= 0) {
455
+ return 0;
456
+ }
457
+ if (count <= 0) {
458
+ return 0;
459
+ }
460
+ return 1;
461
+ }
462
+
463
+
464
+
465
+ int32_t bridge_abi_version(void) {
466
+ return 3;
467
+ }
468
+
469
+ int32_t bridge_matrix_exp_gpu_ptr_c64_i32_i32_ptr_c64(float complex* a, int32_t n, int32_t count, float complex* out) {
470
+ if ((!(shapes_ok_i32_i32(n, count)))) {
471
+ return (-1);
472
+ }
473
+ return metal_matrix_exp_c64(a, n, count, out);
474
+ }
475
+
476
+ int32_t bridge_matrix_exp_backward_gpu_ptr_c64_ptr_c64_i32_i32_i32_ptr_c64(float complex* m, float complex* ebar, int32_t n, int32_t count, int32_t dense, float complex* mbar) {
477
+ if ((!(shapes_ok_i32_i32(n, count)))) {
478
+ return (-1);
479
+ }
480
+ return metal_matrix_exp_backward_c64(m, ebar, n, count, dense, mbar);
481
+ }
482
+
483
+ int32_t bridge_matrix_exp_ptr_c64_i32_i32_ptr_c64(float complex* a, int32_t n, int32_t count, float complex* out) {
484
+ if ((!(shapes_ok_i32_i32(n, count)))) {
485
+ return (-1);
486
+ }
487
+ int32_t block = (n * n);
488
+ int32_t total = (block * count);
489
+ flow_span_mut_c64 __flow_span_init_3 = ((flow_span_mut_c64){ .data = (((a)) + (0)), .len = (int64_t)((total) - (0)) });
490
+ flow_span_mut_c64 src = __flow_span_init_3;
491
+ flow_span_mut_c64 __flow_span_init_4 = ((flow_span_mut_c64){ .data = (((out)) + (0)), .len = (int64_t)((total) - (0)) });
492
+ flow_span_mut_c64 dst = __flow_span_init_4;
493
+ CMat work = cmat_new_i32(n);
494
+ CMat result = cmat_new_i32(n);
495
+ int32_t __flow_step_22 = 1;
496
+ for (int32_t k = 0; (0 <= count) ? k < count : k > count; k += (0 <= count) ? 1 : -1) {
497
+ int32_t base = (k * block);
498
+ int32_t __flow_step_23 = 1;
499
+ for (int32_t i = 0; (0 <= block) ? i < block : i > block; i += (0 <= block) ? 1 : -1) {
500
+ (work.data).data[i] = (src).data[(base + i)];
501
+ }
502
+ expm_CMat_CMat(work, result);
503
+ int32_t __flow_step_24 = 1;
504
+ for (int32_t i = 0; (0 <= block) ? i < block : i > block; i += (0 <= block) ? 1 : -1) {
505
+ (dst).data[(base + i)] = (result.data).data[i];
506
+ }
507
+ }
508
+ cmat_free_CMat(work);
509
+ cmat_free_CMat(result);
510
+ return 0;
511
+ }
512
+
513
+ int32_t bridge_matrix_exp_backward_ptr_c64_ptr_c64_i32_i32_i32_ptr_c64(float complex* m, float complex* ebar, int32_t n, int32_t count, int32_t dense, float complex* mbar) {
514
+ if ((!(shapes_ok_i32_i32(n, count)))) {
515
+ return (-1);
516
+ }
517
+ int32_t block = (n * n);
518
+ int32_t total = (block * count);
519
+ flow_span_mut_c64 __flow_span_init_5 = ((flow_span_mut_c64){ .data = (((m)) + (0)), .len = (int64_t)((total) - (0)) });
520
+ flow_span_mut_c64 m_all = __flow_span_init_5;
521
+ flow_span_mut_c64 __flow_span_init_6 = ((flow_span_mut_c64){ .data = (((ebar)) + (0)), .len = (int64_t)((total) - (0)) });
522
+ flow_span_mut_c64 ebar_all = __flow_span_init_6;
523
+ flow_span_mut_c64 __flow_span_init_7 = ((flow_span_mut_c64){ .data = (((mbar)) + (0)), .len = (int64_t)((total) - (0)) });
524
+ flow_span_mut_c64 mbar_all = __flow_span_init_7;
525
+ Adjoint method = (Adjoint){ .tag = Adjoint_Blocked };
526
+ if (dense != 0) {
527
+ method = (Adjoint){ .tag = Adjoint_Dense };
528
+ }
529
+ CMat m_k = cmat_new_i32(n);
530
+ CMat ebar_k = cmat_new_i32(n);
531
+ CMat mbar_k = cmat_new_i32(n);
532
+ int32_t __flow_step_25 = 1;
533
+ for (int32_t k = 0; (0 <= count) ? k < count : k > count; k += (0 <= count) ? 1 : -1) {
534
+ int32_t base = (k * block);
535
+ int32_t __flow_step_26 = 1;
536
+ for (int32_t i = 0; (0 <= block) ? i < block : i > block; i += (0 <= block) ? 1 : -1) {
537
+ (m_k.data).data[i] = (m_all).data[(base + i)];
538
+ (ebar_k.data).data[i] = (ebar_all).data[(base + i)];
539
+ }
540
+ expm_backward_Adjoint_CMat_CMat_CMat(method, m_k, ebar_k, mbar_k);
541
+ int32_t __flow_step_27 = 1;
542
+ for (int32_t i = 0; (0 <= block) ? i < block : i > block; i += (0 <= block) ? 1 : -1) {
543
+ (mbar_all).data[(base + i)] = (mbar_k.data).data[i];
544
+ }
545
+ }
546
+ cmat_free_CMat(m_k);
547
+ cmat_free_CMat(ebar_k);
548
+ cmat_free_CMat(mbar_k);
549
+ return 0;
550
+ }
551
+
552
+ double bridge_norm1_ptr_c64_i32(float complex* a, int32_t n) {
553
+ if ((!(shapes_ok_i32_i32(n, 1)))) {
554
+ return 0.0;
555
+ }
556
+ CMat m = (CMat){ .data = ((flow_span_mut_c64){ .data = (((a)) + (0)), .len = (int64_t)(((n * n)) - (0)) }), .n = n };
557
+ return norm1_CMat(m);
558
+ }
559
+
560
+
561
+
562
+
563
+ flow_span_mut_c128 zeroed_c128_i32(int32_t n) {
564
+ int64_t bytes = (((int64_t)(n)) * 16);
565
+ double complex* raw = (double complex*)(malloc(bytes));
566
+ memset(raw, 0, bytes);
567
+ return ((flow_span_mut_c128){ .data = (((raw)) + (0)), .len = (int64_t)((n) - (0)) });
568
+ }
569
+
570
+ void gather_into_span_const_c128_span_const_i32_i32_span_const_i32_span_mut_c128(flow_span_const_c128 fgb, flow_span_const_i32 source, int32_t buffer_size, flow_span_const_i32 destination, flow_span_mut_c128 gathered) {
571
+ flow_span_mut_c128 __flow_span_init_8 = zeroed_c128_i32(buffer_size);
572
+ flow_span_mut_c128 buffer = __flow_span_init_8;
573
+ int32_t __flow_step_28 = 1;
574
+ for (int32_t k = 0; (0 <= ((int32_t)(fgb.len))) ? k < ((int32_t)(fgb.len)) : k > ((int32_t)(fgb.len)); k += (0 <= ((int32_t)(fgb.len))) ? 1 : -1) {
575
+ int32_t slot = (source).data[k];
576
+ (buffer).data[slot] = ((buffer).data[slot] + (fgb).data[k]);
577
+ }
578
+ int32_t __flow_step_29 = 1;
579
+ for (int32_t p = 0; (0 <= ((int32_t)(gathered.len))) ? p < ((int32_t)(gathered.len)) : p > ((int32_t)(gathered.len)); p += (0 <= ((int32_t)(gathered.len))) ? 1 : -1) {
580
+ (gathered).data[p] = (buffer).data[(destination).data[p]];
581
+ }
582
+ free(buffer.data);
583
+ }
584
+
585
+ int32_t bridge_structure_matrix_ptr_c128_i32_ptr_i32_i32_ptr_i32_i32_ptr_f64_ptr_f64_i32_f64_i32_ptr_c128(double complex* fgb, int32_t n_grid, int32_t* source, int32_t buffer_size, int32_t* destination, int32_t n, double* mii, double* diagonal, int32_t n_batch, double prefactor, int32_t absorption, double complex* out) {
586
+ if ((!(shapes_ok_i32_i32(n, n_batch)))) {
587
+ return (-1);
588
+ }
589
+ if (n_grid <= 0) {
590
+ return (-2);
591
+ }
592
+ if (buffer_size <= 0) {
593
+ return (-3);
594
+ }
595
+ int32_t pairs = (n * n);
596
+ flow_span_mut_c128 __flow_span_init_9 = ((flow_span_mut_c128){ .data = (((fgb)) + (0)), .len = (int64_t)((n_grid) - (0)) });
597
+ flow_span_const_c128 f = ((flow_span_const_c128){ .data = (const double complex*)(__flow_span_init_9).data, .len = (__flow_span_init_9).len });
598
+ flow_span_mut_i32 __flow_span_init_10 = ((flow_span_mut_i32){ .data = (((source)) + (0)), .len = (int64_t)((n_grid) - (0)) });
599
+ flow_span_const_i32 src = ((flow_span_const_i32){ .data = (const int32_t*)(__flow_span_init_10).data, .len = (__flow_span_init_10).len });
600
+ flow_span_mut_i32 __flow_span_init_11 = ((flow_span_mut_i32){ .data = (((destination)) + (0)), .len = (int64_t)((pairs) - (0)) });
601
+ flow_span_const_i32 dst = ((flow_span_const_i32){ .data = (const int32_t*)(__flow_span_init_11).data, .len = (__flow_span_init_11).len });
602
+ flow_span_mut_f64 __flow_span_init_12 = ((flow_span_mut_f64){ .data = (((mii)) + (0)), .len = (int64_t)(((n_batch * n)) - (0)) });
603
+ flow_span_const_f64 m = ((flow_span_const_f64){ .data = (const double*)(__flow_span_init_12).data, .len = (__flow_span_init_12).len });
604
+ flow_span_mut_f64 __flow_span_init_13 = ((flow_span_mut_f64){ .data = (((diagonal)) + (0)), .len = (int64_t)(((n_batch * n)) - (0)) });
605
+ flow_span_const_f64 d = ((flow_span_const_f64){ .data = (const double*)(__flow_span_init_13).data, .len = (__flow_span_init_13).len });
606
+ flow_span_mut_c128 __flow_span_init_14 = ((flow_span_mut_c128){ .data = (((out)) + (0)), .len = (int64_t)(((n_batch * pairs)) - (0)) });
607
+ flow_span_mut_c128 a = __flow_span_init_14;
608
+ flow_span_mut_c128 __flow_span_init_15 = zeroed_c128_i32(pairs);
609
+ flow_span_mut_c128 gathered = __flow_span_init_15;
610
+ gather_into_span_const_c128_span_const_i32_i32_span_const_i32_span_mut_c128(f, src, buffer_size, dst, gathered);
611
+ double u0_prime = 0.0;
612
+ if (absorption != 0) {
613
+ double total = 0.0;
614
+ int32_t __flow_step_30 = 1;
615
+ for (int32_t i = 0; (0 <= n) ? i < n : i > n; i += (0 <= n) ? 1 : -1) {
616
+ total = (total + cimag((gathered).data[((i * n) + i)]));
617
+ }
618
+ u0_prime = (prefactor * (total / ((double)(n))));
619
+ }
620
+ int32_t __flow_step_31 = 1;
621
+ for (int32_t b = 0; (0 <= n_batch) ? b < n_batch : b > n_batch; b += (0 <= n_batch) ? 1 : -1) {
622
+ int32_t base_m = (b * n);
623
+ int32_t base_a = (b * pairs);
624
+ int32_t __flow_step_32 = 1;
625
+ for (int32_t i = 0; (0 <= n) ? i < n : i > n; i += (0 <= n) ? 1 : -1) {
626
+ int32_t __flow_step_33 = 1;
627
+ for (int32_t j = 0; (0 <= n) ? j < n : j > n; j += (0 <= n) ? 1 : -1) {
628
+ double scale = ((prefactor * (m).data[(base_m + j)]) * (m).data[(base_m + i)]);
629
+ double complex z = (gathered).data[((i * n) + j)];
630
+ (a).data[((base_a + (i * n)) + j)] = ((double)((creal(z) * scale)) + (double)((cimag(z) * scale)) * I);
631
+ }
632
+ (a).data[((base_a + (i * n)) + i)] = ((double)((d).data[(base_m + i)]) + (double)((u0_prime * (m).data[(base_m + i)])) * I);
633
+ }
634
+ }
635
+ free(gathered.data);
636
+ return 0;
637
+ }
638
+
639
+ int32_t bridge_structure_matrix_backward_ptr_c128_ptr_i32_i32_ptr_f64_i32_f64_i32_ptr_i32_i32_i32_ptr_c128(double complex* abar, int32_t* destination, int32_t n, double* mii, int32_t n_batch, double prefactor, int32_t absorption, int32_t* source, int32_t buffer_size, int32_t n_grid, double complex* fbar) {
640
+ if ((!(shapes_ok_i32_i32(n, n_batch)))) {
641
+ return (-1);
642
+ }
643
+ if (n_grid <= 0) {
644
+ return (-2);
645
+ }
646
+ if (buffer_size <= 0) {
647
+ return (-3);
648
+ }
649
+ int32_t pairs = (n * n);
650
+ flow_span_mut_c128 __flow_span_init_16 = ((flow_span_mut_c128){ .data = (((abar)) + (0)), .len = (int64_t)(((n_batch * pairs)) - (0)) });
651
+ flow_span_const_c128 ab = ((flow_span_const_c128){ .data = (const double complex*)(__flow_span_init_16).data, .len = (__flow_span_init_16).len });
652
+ flow_span_mut_i32 __flow_span_init_17 = ((flow_span_mut_i32){ .data = (((destination)) + (0)), .len = (int64_t)((pairs) - (0)) });
653
+ flow_span_const_i32 dst = ((flow_span_const_i32){ .data = (const int32_t*)(__flow_span_init_17).data, .len = (__flow_span_init_17).len });
654
+ flow_span_mut_f64 __flow_span_init_18 = ((flow_span_mut_f64){ .data = (((mii)) + (0)), .len = (int64_t)(((n_batch * n)) - (0)) });
655
+ flow_span_const_f64 m = ((flow_span_const_f64){ .data = (const double*)(__flow_span_init_18).data, .len = (__flow_span_init_18).len });
656
+ flow_span_mut_i32 __flow_span_init_19 = ((flow_span_mut_i32){ .data = (((source)) + (0)), .len = (int64_t)((n_grid) - (0)) });
657
+ flow_span_const_i32 src = ((flow_span_const_i32){ .data = (const int32_t*)(__flow_span_init_19).data, .len = (__flow_span_init_19).len });
658
+ flow_span_mut_c128 __flow_span_init_20 = ((flow_span_mut_c128){ .data = (((fbar)) + (0)), .len = (int64_t)((n_grid) - (0)) });
659
+ flow_span_mut_c128 fb = __flow_span_init_20;
660
+ flow_span_mut_c128 __flow_span_init_21 = zeroed_c128_i32(buffer_size);
661
+ flow_span_mut_c128 buffer = __flow_span_init_21;
662
+ int32_t __flow_step_34 = 1;
663
+ for (int32_t b = 0; (0 <= n_batch) ? b < n_batch : b > n_batch; b += (0 <= n_batch) ? 1 : -1) {
664
+ int32_t base_m = (b * n);
665
+ int32_t base_a = (b * pairs);
666
+ int32_t __flow_step_35 = 1;
667
+ for (int32_t i = 0; (0 <= n) ? i < n : i > n; i += (0 <= n) ? 1 : -1) {
668
+ int32_t __flow_step_36 = 1;
669
+ for (int32_t j = 0; (0 <= n) ? j < n : j > n; j += (0 <= n) ? 1 : -1) {
670
+ if (i != j) {
671
+ double scale = ((prefactor * (m).data[(base_m + j)]) * (m).data[(base_m + i)]);
672
+ int32_t slot = (dst).data[((i * n) + j)];
673
+ double complex z = (ab).data[((base_a + (i * n)) + j)];
674
+ (buffer).data[slot] = ((buffer).data[slot] + ((double)((creal(z) * scale)) + (double)((cimag(z) * scale)) * I));
675
+ }
676
+ }
677
+ }
678
+ }
679
+ if (absorption != 0) {
680
+ double total = 0.0;
681
+ int32_t __flow_step_37 = 1;
682
+ for (int32_t b = 0; (0 <= n_batch) ? b < n_batch : b > n_batch; b += (0 <= n_batch) ? 1 : -1) {
683
+ int32_t __flow_step_38 = 1;
684
+ for (int32_t i = 0; (0 <= n) ? i < n : i > n; i += (0 <= n) ? 1 : -1) {
685
+ total = (total + (cimag((ab).data[(((b * pairs) + (i * n)) + i)]) * (m).data[((b * n) + i)]));
686
+ }
687
+ }
688
+ double share = ((prefactor * total) / ((double)(n)));
689
+ int32_t __flow_step_39 = 1;
690
+ for (int32_t i = 0; (0 <= n) ? i < n : i > n; i += (0 <= n) ? 1 : -1) {
691
+ int32_t slot = (dst).data[((i * n) + i)];
692
+ (buffer).data[slot] = ((buffer).data[slot] + ((double)(0.0) + (double)(share) * I));
693
+ }
694
+ }
695
+ int32_t __flow_step_40 = 1;
696
+ for (int32_t k = 0; (0 <= n_grid) ? k < n_grid : k > n_grid; k += (0 <= n_grid) ? 1 : -1) {
697
+ (fb).data[k] = (buffer).data[(src).data[k]];
698
+ }
699
+ free(buffer.data);
700
+ return 0;
701
+ }