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,374 @@
1
+ #include <metal_stdlib>
2
+ using namespace metal;
3
+
4
+ static inline float2 c_mul(float2 a, float2 b) {
5
+ return float2(a.x * b.x - a.y * b.y, a.x * b.y + a.y * b.x);
6
+ }
7
+
8
+ static inline float2 c_fma(float2 a, float2 b, float2 acc) {
9
+ return float2(acc.x + (a.x * b.x - a.y * b.y), acc.y + (a.x * b.y + a.y * b.x));
10
+ }
11
+
12
+ static inline float2 c_conj(float2 a) {
13
+ return float2(a.x, -a.y);
14
+ }
15
+
16
+ constant int TILE_SIZE = 16;
17
+
18
+ struct BatchGemmParams {
19
+ int M;
20
+ int N;
21
+ int K;
22
+ int lda;
23
+ int ldb;
24
+ int ldc;
25
+ int strideA;
26
+ int strideB;
27
+ int strideC;
28
+ int transA;
29
+ int transB;
30
+ int pad;
31
+ float2 alpha;
32
+ float2 beta;
33
+ int batch_count;
34
+ int pad2;
35
+ };
36
+
37
+ // -----------------------------------------------------------------------------
38
+ // Batched Complex GEMM Kernel (c64)
39
+ // -----------------------------------------------------------------------------
40
+ kernel void batch_cgemm_c64(
41
+ device const float2* A [[buffer(0)]],
42
+ device const float2* B [[buffer(1)]],
43
+ device float2* C [[buffer(2)]],
44
+ constant BatchGemmParams& params [[buffer(3)]],
45
+ uint3 thread_pos_in_grid [[thread_position_in_grid]],
46
+ uint3 thread_pos_in_tg [[thread_position_in_threadgroup]],
47
+ uint3 tg_pos [[threadgroup_position_in_grid]]
48
+ ) {
49
+ int col = thread_pos_in_grid.x;
50
+ int row = thread_pos_in_grid.y;
51
+ int batch_idx = thread_pos_in_grid.z;
52
+
53
+ if (batch_idx >= params.batch_count) {
54
+ return;
55
+ }
56
+
57
+ threadgroup float2 sA[TILE_SIZE][TILE_SIZE];
58
+ threadgroup float2 sB[TILE_SIZE][TILE_SIZE];
59
+
60
+ int tx = thread_pos_in_tg.x;
61
+ int ty = thread_pos_in_tg.y;
62
+
63
+ device const float2* currA = A + (int64_t)batch_idx * params.strideA;
64
+ device const float2* currB = B + (int64_t)batch_idx * params.strideB;
65
+ device float2* currC = C + (int64_t)batch_idx * params.strideC;
66
+
67
+ float2 acc = float2(0.0f, 0.0f);
68
+
69
+ int num_tiles = (params.K + TILE_SIZE - 1) / TILE_SIZE;
70
+
71
+ for (int t = 0; t < num_tiles; t++) {
72
+ int a_k = t * TILE_SIZE + tx;
73
+ if (row < params.M && a_k < params.K) {
74
+ float2 val;
75
+ if (params.transA == 0) {
76
+ val = currA[row * params.lda + a_k];
77
+ } else if (params.transA == 1) {
78
+ val = currA[a_k * params.lda + row];
79
+ } else {
80
+ val = c_conj(currA[a_k * params.lda + row]);
81
+ }
82
+ sA[ty][tx] = val;
83
+ } else {
84
+ sA[ty][tx] = float2(0.0f, 0.0f);
85
+ }
86
+
87
+ int b_k = t * TILE_SIZE + ty;
88
+ if (b_k < params.K && col < params.N) {
89
+ float2 val;
90
+ if (params.transB == 0) {
91
+ val = currB[b_k * params.ldb + col];
92
+ } else if (params.transB == 1) {
93
+ val = currB[col * params.ldb + b_k];
94
+ } else {
95
+ val = c_conj(currB[col * params.ldb + b_k]);
96
+ }
97
+ sB[ty][tx] = val;
98
+ } else {
99
+ sB[ty][tx] = float2(0.0f, 0.0f);
100
+ }
101
+
102
+ threadgroup_barrier(mem_flags::mem_threadgroup);
103
+
104
+ #pragma unroll
105
+ for (int k = 0; k < TILE_SIZE; k++) {
106
+ acc = c_fma(sA[ty][k], sB[k][tx], acc);
107
+ }
108
+
109
+ threadgroup_barrier(mem_flags::mem_threadgroup);
110
+ }
111
+
112
+ if (row < params.M && col < params.N) {
113
+ int c_idx = row * params.ldc + col;
114
+ float2 out_val = c_mul(params.alpha, acc);
115
+ if (params.beta.x != 0.0f || params.beta.y != 0.0f) {
116
+ out_val += c_mul(params.beta, currC[c_idx]);
117
+ }
118
+ currC[c_idx] = out_val;
119
+ }
120
+ }
121
+
122
+ // -----------------------------------------------------------------------------
123
+ // Matrix Operations for GPU-Resident Expm & Adjoint
124
+ // -----------------------------------------------------------------------------
125
+
126
+ struct ElementwiseParams {
127
+ int total_elements;
128
+ int n;
129
+ int batch_count;
130
+ float scalar_a;
131
+ float scalar_b;
132
+ float scalar_c;
133
+ };
134
+
135
+ kernel void cmat_scale_kernel(
136
+ device const float2* src [[buffer(0)]],
137
+ device float2* dst [[buffer(1)]],
138
+ constant ElementwiseParams& p [[buffer(2)]],
139
+ uint id [[thread_position_in_grid]]
140
+ ) {
141
+ if (id < (uint)p.total_elements) {
142
+ dst[id] = src[id] * p.scalar_a;
143
+ }
144
+ }
145
+
146
+ kernel void cmat_scale_adaptive_kernel(
147
+ device const float2* src [[buffer(0)]],
148
+ device float2* dst [[buffer(1)]],
149
+ device const float* scales [[buffer(2)]],
150
+ constant ElementwiseParams& p [[buffer(3)]],
151
+ uint id [[thread_position_in_grid]]
152
+ ) {
153
+ if (id < (uint)p.total_elements) {
154
+ int b = id / (p.n * p.n);
155
+ dst[id] = src[id] * scales[b];
156
+ }
157
+ }
158
+
159
+ kernel void cmat_squaring_gate_kernel(
160
+ device const float2* gemm_res [[buffer(0)]],
161
+ device const float2* acc_prev [[buffer(1)]],
162
+ device float2* dst [[buffer(2)]],
163
+ device const int* squarings [[buffer(3)]],
164
+ constant int& current_k [[buffer(4)]],
165
+ constant ElementwiseParams& p [[buffer(5)]],
166
+ uint id [[thread_position_in_grid]]
167
+ ) {
168
+ if (id < (uint)p.total_elements) {
169
+ int b = id / (p.n * p.n);
170
+ if (current_k < squarings[b]) {
171
+ dst[id] = gemm_res[id];
172
+ } else {
173
+ dst[id] = acc_prev[id];
174
+ }
175
+ }
176
+ }
177
+
178
+ kernel void cmat_copy_kernel(
179
+ device const float2* src [[buffer(0)]],
180
+ device float2* dst [[buffer(1)]],
181
+ constant ElementwiseParams& p [[buffer(2)]],
182
+ uint id [[thread_position_in_grid]]
183
+ ) {
184
+ if (id < (uint)p.total_elements) {
185
+ dst[id] = src[id];
186
+ }
187
+ }
188
+
189
+ kernel void cmat_conjugate_transpose_kernel(
190
+ device const float2* src [[buffer(0)]],
191
+ device float2* dst [[buffer(1)]],
192
+ constant ElementwiseParams& p [[buffer(2)]],
193
+ uint3 pos [[thread_position_in_grid]]
194
+ ) {
195
+ int col = pos.x;
196
+ int row = pos.y;
197
+ int b = pos.z;
198
+ if (col < p.n && row < p.n && b < p.batch_count) {
199
+ int src_idx = b * (p.n * p.n) + row * p.n + col;
200
+ int dst_idx = b * (p.n * p.n) + col * p.n + row;
201
+ float2 val = src[src_idx];
202
+ dst[dst_idx] = float2(val.x, -val.y);
203
+ }
204
+ }
205
+
206
+ kernel void cmat_taylor_horner_combine(
207
+ device float2* acc [[buffer(0)]],
208
+ device const float2* m [[buffer(1)]],
209
+ device const float2* m2 [[buffer(2)]],
210
+ constant ElementwiseParams& p [[buffer(3)]],
211
+ uint3 pos [[thread_position_in_grid]]
212
+ ) {
213
+ int col = pos.x;
214
+ int row = pos.y;
215
+ int b = pos.z;
216
+ if (col < p.n && row < p.n && b < p.batch_count) {
217
+ int idx = b * (p.n * p.n) + row * p.n + col;
218
+ float2 val = acc[idx] + m[idx] * p.scalar_b + m2[idx] * p.scalar_c;
219
+ if (row == col) {
220
+ val.x += p.scalar_a;
221
+ }
222
+ acc[idx] = val;
223
+ }
224
+ }
225
+
226
+ kernel void cmat_taylor_horner_combine_pair(
227
+ device float2* acc_y [[buffer(0)]],
228
+ device float2* acc_l [[buffer(1)]],
229
+ device const float2* y [[buffer(2)]],
230
+ device const float2* l [[buffer(3)]],
231
+ device const float2* y2 [[buffer(4)]],
232
+ device const float2* l2 [[buffer(5)]],
233
+ constant ElementwiseParams& p [[buffer(6)]],
234
+ uint3 pos [[thread_position_in_grid]]
235
+ ) {
236
+ int col = pos.x;
237
+ int row = pos.y;
238
+ int b = pos.z;
239
+ if (col < p.n && row < p.n && b < p.batch_count) {
240
+ int idx = b * (p.n * p.n) + row * p.n + col;
241
+ float2 vy = acc_y[idx] + y[idx] * p.scalar_b + y2[idx] * p.scalar_c;
242
+ if (row == col) {
243
+ vy.x += p.scalar_a;
244
+ }
245
+ acc_y[idx] = vy;
246
+ acc_l[idx] = acc_l[idx] + l[idx] * p.scalar_b + l2[idx] * p.scalar_c;
247
+ }
248
+ }
249
+
250
+ // -----------------------------------------------------------------------------
251
+ // Structure Factors on Metal (c64 / float2)
252
+ // -----------------------------------------------------------------------------
253
+
254
+ struct StructureFactorParams {
255
+ int n_grid;
256
+ int n_atoms;
257
+ float volume;
258
+ float two_pi;
259
+ float minus_two_pi_sq;
260
+ int absorption;
261
+ };
262
+
263
+ static inline float lobato_eval(device const float* a_coeffs, device const float* b_coeffs, int offset, float g2) {
264
+ float sum = 0.0f;
265
+ for (int i = 0; i < 5; i++) {
266
+ float a = a_coeffs[offset + i];
267
+ float b = b_coeffs[offset + i];
268
+ sum += a / (1.0f + b * g2);
269
+ }
270
+ return sum;
271
+ }
272
+
273
+ kernel void structure_factors_c64_kernel(
274
+ device float2* fgb [[buffer(0)]],
275
+ device const float* g2_arr [[buffer(1)]],
276
+ device const float* hkl_arr [[buffer(2)]],
277
+ device const float* lobato_a [[buffer(3)]],
278
+ device const float* lobato_b [[buffer(4)]],
279
+ device const float* uij_arr [[buffer(5)]],
280
+ device const float* pos_arr [[buffer(6)]],
281
+ device const float* occ_arr [[buffer(7)]],
282
+ constant StructureFactorParams& p [[buffer(8)]],
283
+ uint id [[thread_position_in_grid]]
284
+ ) {
285
+ if (id >= (uint)p.n_grid) return;
286
+
287
+ float g2 = g2_arr[id];
288
+ float h0 = hkl_arr[id * 3 + 0];
289
+ float h1 = hkl_arr[id * 3 + 1];
290
+ float h2 = hkl_arr[id * 3 + 2];
291
+
292
+ float re_acc = 0.0f;
293
+ float im_acc = 0.0f;
294
+
295
+ for (int a = 0; a < p.n_atoms; a++) {
296
+ float form = lobato_eval(lobato_a, lobato_b, a * 5, g2);
297
+
298
+ int base_u = a * 9;
299
+ float quad = h0 * (uij_arr[base_u + 0] * h0 + uij_arr[base_u + 1] * h1 + uij_arr[base_u + 2] * h2)
300
+ + h1 * (uij_arr[base_u + 3] * h0 + uij_arr[base_u + 4] * h1 + uij_arr[base_u + 5] * h2)
301
+ + h2 * (uij_arr[base_u + 6] * h0 + uij_arr[base_u + 7] * h1 + uij_arr[base_u + 8] * h2);
302
+ float dwf = metal::exp(p.minus_two_pi_sq * quad);
303
+
304
+ float amplitude = (form * dwf) / p.volume;
305
+ float phase = p.two_pi * (pos_arr[a * 3 + 0] * h0 + pos_arr[a * 3 + 1] * h1 + pos_arr[a * 3 + 2] * h2);
306
+
307
+ float cp = metal::cos(phase);
308
+ float sp = metal::sin(phase);
309
+
310
+ float occ = occ_arr[a];
311
+ re_acc += occ * amplitude * cp;
312
+ im_acc += occ * amplitude * sp;
313
+ }
314
+
315
+ fgb[id] = float2(re_acc, im_acc);
316
+ }
317
+
318
+ // -----------------------------------------------------------------------------
319
+ // Structure Matrix Assembly on Metal (c64 / float2)
320
+ // -----------------------------------------------------------------------------
321
+
322
+ struct StructureMatrixParams {
323
+ int n_beams;
324
+ int n_batch;
325
+ int buffer_size;
326
+ int n_grid;
327
+ float prefactor;
328
+ float u0_prime;
329
+ };
330
+
331
+ kernel void structure_matrix_scatter_kernel(
332
+ device float2* buffer [[buffer(0)]],
333
+ device const float2* fgb [[buffer(1)]],
334
+ device const int* source [[buffer(2)]],
335
+ constant StructureMatrixParams& p [[buffer(3)]],
336
+ uint id [[thread_position_in_grid]]
337
+ ) {
338
+ if (id < (uint)p.n_grid) {
339
+ int slot = source[id];
340
+ buffer[slot] = fgb[id];
341
+ }
342
+ }
343
+
344
+ kernel void structure_matrix_gather_kernel(
345
+ device float2* a [[buffer(0)]],
346
+ device const float2* buffer [[buffer(1)]],
347
+ device const int* destination [[buffer(2)]],
348
+ device const float* mii [[buffer(3)]],
349
+ device const float* diagonal [[buffer(4)]],
350
+ constant StructureMatrixParams& p [[buffer(5)]],
351
+ uint3 pos [[thread_position_in_grid]]
352
+ ) {
353
+ int j = pos.x;
354
+ int i = pos.y;
355
+ int b = pos.z;
356
+
357
+ int n = p.n_beams;
358
+ if (j < n && i < n && b < p.n_batch) {
359
+ int base_m = b * n;
360
+ int base_a = b * (n * n);
361
+ int out_idx = base_a + i * n + j;
362
+
363
+ if (i == j) {
364
+ float diag_re = diagonal[base_m + i];
365
+ float diag_im = p.u0_prime * mii[base_m + i];
366
+ a[out_idx] = float2(diag_re, diag_im);
367
+ } else {
368
+ int slot = destination[i * n + j];
369
+ float2 z = buffer[slot];
370
+ float scale = p.prefactor * mii[base_m + j] * mii[base_m + i];
371
+ a[out_idx] = float2(z.x * scale, z.y * scale);
372
+ }
373
+ }
374
+ }