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.
- faster_diffbloch/__init__.py +14 -0
- faster_diffbloch/backend.py +149 -0
- faster_diffbloch/builder.py +58 -0
- faster_diffbloch/cli.py +21 -0
- faster_diffbloch/native/batch_cgemm.c +214 -0
- faster_diffbloch/native/batch_cgemm.h +142 -0
- faster_diffbloch/native/batch_cgemm.metal +374 -0
- faster_diffbloch/native/bridge_lib.c +701 -0
- faster_diffbloch/native/metal_batch_cgemm.m +776 -0
- faster_diffbloch/native/native_scattering.c +233 -0
- faster_diffbloch/native/native_scattering.h +60 -0
- faster_diffbloch-0.1.0.dist-info/METADATA +98 -0
- faster_diffbloch-0.1.0.dist-info/RECORD +16 -0
- faster_diffbloch-0.1.0.dist-info/WHEEL +4 -0
- faster_diffbloch-0.1.0.dist-info/entry_points.txt +3 -0
- faster_diffbloch-0.1.0.dist-info/licenses/LICENSE +21 -0
|
@@ -0,0 +1,776 @@
|
|
|
1
|
+
#import <Foundation/Foundation.h>
|
|
2
|
+
#import <Metal/Metal.h>
|
|
3
|
+
#include "batch_cgemm.h"
|
|
4
|
+
#include <math.h>
|
|
5
|
+
#include <dlfcn.h>
|
|
6
|
+
|
|
7
|
+
struct __attribute__((aligned(8))) BatchGemmParams {
|
|
8
|
+
int32_t M;
|
|
9
|
+
int32_t N;
|
|
10
|
+
int32_t K;
|
|
11
|
+
int32_t lda;
|
|
12
|
+
int32_t ldb;
|
|
13
|
+
int32_t ldc;
|
|
14
|
+
int32_t strideA;
|
|
15
|
+
int32_t strideB;
|
|
16
|
+
int32_t strideC;
|
|
17
|
+
int32_t transA;
|
|
18
|
+
int32_t transB;
|
|
19
|
+
int32_t pad;
|
|
20
|
+
struct { float x; float y; } alpha;
|
|
21
|
+
struct { float x; float y; } beta;
|
|
22
|
+
int32_t batch_count;
|
|
23
|
+
int32_t pad2;
|
|
24
|
+
};
|
|
25
|
+
|
|
26
|
+
struct __attribute__((aligned(8))) ElementwiseParams {
|
|
27
|
+
int32_t total_elements;
|
|
28
|
+
int32_t n;
|
|
29
|
+
int32_t batch_count;
|
|
30
|
+
float scalar_a;
|
|
31
|
+
float scalar_b;
|
|
32
|
+
float scalar_c;
|
|
33
|
+
};
|
|
34
|
+
|
|
35
|
+
struct __attribute__((aligned(8))) StructureFactorParams {
|
|
36
|
+
int32_t n_grid;
|
|
37
|
+
int32_t n_atoms;
|
|
38
|
+
float volume;
|
|
39
|
+
float two_pi;
|
|
40
|
+
float minus_two_pi_sq;
|
|
41
|
+
int32_t absorption;
|
|
42
|
+
};
|
|
43
|
+
|
|
44
|
+
struct __attribute__((aligned(8))) StructureMatrixParams {
|
|
45
|
+
int32_t n_beams;
|
|
46
|
+
int32_t n_batch;
|
|
47
|
+
int32_t buffer_size;
|
|
48
|
+
int32_t n_grid;
|
|
49
|
+
float prefactor;
|
|
50
|
+
float u0_prime;
|
|
51
|
+
};
|
|
52
|
+
|
|
53
|
+
static id<MTLDevice> g_device = nil;
|
|
54
|
+
static id<MTLCommandQueue> g_queue = nil;
|
|
55
|
+
|
|
56
|
+
static id<MTLComputePipelineState> g_pipe_gemm = nil;
|
|
57
|
+
static id<MTLComputePipelineState> g_pipe_scale = nil;
|
|
58
|
+
static id<MTLComputePipelineState> g_pipe_scale_adaptive = nil;
|
|
59
|
+
static id<MTLComputePipelineState> g_pipe_squaring_gate = nil;
|
|
60
|
+
static id<MTLComputePipelineState> g_pipe_copy = nil;
|
|
61
|
+
static id<MTLComputePipelineState> g_pipe_conj_trans = nil;
|
|
62
|
+
static id<MTLComputePipelineState> g_pipe_horner = nil;
|
|
63
|
+
static id<MTLComputePipelineState> g_pipe_horner_pair = nil;
|
|
64
|
+
static id<MTLComputePipelineState> g_pipe_sf = nil;
|
|
65
|
+
static id<MTLComputePipelineState> g_pipe_sm_scatter = nil;
|
|
66
|
+
static id<MTLComputePipelineState> g_pipe_sm_gather = nil;
|
|
67
|
+
|
|
68
|
+
static dispatch_once_t g_init_once;
|
|
69
|
+
|
|
70
|
+
static const double TAYLOR_COEFFS[19] = {
|
|
71
|
+
1.0,
|
|
72
|
+
1.0,
|
|
73
|
+
0.5,
|
|
74
|
+
1.0 / 6.0,
|
|
75
|
+
1.0 / 24.0,
|
|
76
|
+
1.0 / 120.0,
|
|
77
|
+
1.0 / 720.0,
|
|
78
|
+
1.0 / 5040.0,
|
|
79
|
+
1.0 / 40320.0,
|
|
80
|
+
1.0 / 362880.0,
|
|
81
|
+
1.0 / 3628800.0,
|
|
82
|
+
1.0 / 39916800.0,
|
|
83
|
+
1.0 / 479001600.0,
|
|
84
|
+
1.0 / 6227020800.0,
|
|
85
|
+
1.0 / 87178291200.0,
|
|
86
|
+
1.0 / 1307674368000.0,
|
|
87
|
+
1.0 / 20922789888000.0,
|
|
88
|
+
1.0 / 355687428096000.0,
|
|
89
|
+
1.0 / 6402373705728000.0
|
|
90
|
+
};
|
|
91
|
+
|
|
92
|
+
static const double THETA18 = 3.010066362817634;
|
|
93
|
+
|
|
94
|
+
static id<MTLComputePipelineState> make_pipeline(id<MTLLibrary> lib, NSString* name, NSError** err) {
|
|
95
|
+
id<MTLFunction> fn = [lib newFunctionWithName:name];
|
|
96
|
+
if (!fn) {
|
|
97
|
+
fprintf(stderr, "Metal error: kernel function %s not found\n", name.UTF8String);
|
|
98
|
+
return nil;
|
|
99
|
+
}
|
|
100
|
+
return [g_device newComputePipelineStateWithFunction:fn error:err];
|
|
101
|
+
}
|
|
102
|
+
|
|
103
|
+
static NSString* get_metallib_path(void) {
|
|
104
|
+
Dl_info info;
|
|
105
|
+
if (dladdr((const void*)get_metallib_path, &info) && info.dli_fname) {
|
|
106
|
+
NSString* dylibPath = [NSString stringWithUTF8String:info.dli_fname];
|
|
107
|
+
NSString* dir = [dylibPath stringByDeletingLastPathComponent];
|
|
108
|
+
NSString* metallibInDir = [dir stringByAppendingPathComponent:@"batch_cgemm.metallib"];
|
|
109
|
+
if ([[NSFileManager defaultManager] fileExistsAtPath:metallibInDir]) {
|
|
110
|
+
return metallibInDir;
|
|
111
|
+
}
|
|
112
|
+
}
|
|
113
|
+
NSArray* candidates = @[
|
|
114
|
+
@"bench/build/batch_cgemm.metallib",
|
|
115
|
+
@"src/batch_gemm/batch_cgemm.metallib",
|
|
116
|
+
@"../bench/build/batch_cgemm.metallib",
|
|
117
|
+
@"../../bench/build/batch_cgemm.metallib",
|
|
118
|
+
@"/Users/abhishekshivakumar/diffFlow/bench/build/batch_cgemm.metallib"
|
|
119
|
+
];
|
|
120
|
+
for (NSString* p in candidates) {
|
|
121
|
+
if ([[NSFileManager defaultManager] fileExistsAtPath:p]) {
|
|
122
|
+
return p;
|
|
123
|
+
}
|
|
124
|
+
}
|
|
125
|
+
return nil;
|
|
126
|
+
}
|
|
127
|
+
|
|
128
|
+
static void init_metal_state(void) {
|
|
129
|
+
dispatch_once(&g_init_once, ^{
|
|
130
|
+
@autoreleasepool {
|
|
131
|
+
g_device = MTLCreateSystemDefaultDevice();
|
|
132
|
+
if (!g_device) {
|
|
133
|
+
fprintf(stderr, "Metal error: failed to get default Metal device\n");
|
|
134
|
+
return;
|
|
135
|
+
}
|
|
136
|
+
g_queue = [g_device newCommandQueue];
|
|
137
|
+
|
|
138
|
+
NSError* error = nil;
|
|
139
|
+
NSString* libPath = get_metallib_path();
|
|
140
|
+
id<MTLLibrary> library = nil;
|
|
141
|
+
if (libPath) {
|
|
142
|
+
library = [g_device newLibraryWithURL:[NSURL fileURLWithPath:libPath] error:&error];
|
|
143
|
+
}
|
|
144
|
+
|
|
145
|
+
if (!library) {
|
|
146
|
+
fprintf(stderr, "Metal library error: %s\n", error ? error.localizedDescription.UTF8String : "library not found");
|
|
147
|
+
return;
|
|
148
|
+
}
|
|
149
|
+
|
|
150
|
+
g_pipe_gemm = make_pipeline(library, @"batch_cgemm_c64", &error);
|
|
151
|
+
g_pipe_scale = make_pipeline(library, @"cmat_scale_kernel", &error);
|
|
152
|
+
g_pipe_scale_adaptive = make_pipeline(library, @"cmat_scale_adaptive_kernel", &error);
|
|
153
|
+
g_pipe_squaring_gate = make_pipeline(library, @"cmat_squaring_gate_kernel", &error);
|
|
154
|
+
g_pipe_copy = make_pipeline(library, @"cmat_copy_kernel", &error);
|
|
155
|
+
g_pipe_conj_trans = make_pipeline(library, @"cmat_conjugate_transpose_kernel", &error);
|
|
156
|
+
g_pipe_horner = make_pipeline(library, @"cmat_taylor_horner_combine", &error);
|
|
157
|
+
g_pipe_horner_pair = make_pipeline(library, @"cmat_taylor_horner_combine_pair", &error);
|
|
158
|
+
g_pipe_sf = make_pipeline(library, @"structure_factors_c64_kernel", &error);
|
|
159
|
+
g_pipe_sm_scatter = make_pipeline(library, @"structure_matrix_scatter_kernel", &error);
|
|
160
|
+
g_pipe_sm_gather = make_pipeline(library, @"structure_matrix_gather_kernel", &error);
|
|
161
|
+
}
|
|
162
|
+
});
|
|
163
|
+
}
|
|
164
|
+
|
|
165
|
+
static inline int metal_trans_mode(BatchGemmTranspose t) {
|
|
166
|
+
switch (t) {
|
|
167
|
+
case BATCH_CGEMM_NO_TRANS: return 0;
|
|
168
|
+
case BATCH_CGEMM_TRANS: return 1;
|
|
169
|
+
case BATCH_CGEMM_CONJ_TRANS: return 2;
|
|
170
|
+
default: return 0;
|
|
171
|
+
}
|
|
172
|
+
}
|
|
173
|
+
|
|
174
|
+
static void encode_cgemm(
|
|
175
|
+
id<MTLComputeCommandEncoder> enc,
|
|
176
|
+
id<MTLBuffer> bufA, int64_t offA, int strideA, int transA,
|
|
177
|
+
id<MTLBuffer> bufB, int64_t offB, int strideB, int transB,
|
|
178
|
+
id<MTLBuffer> bufC, int64_t offC, int strideC,
|
|
179
|
+
int M, int N, int K,
|
|
180
|
+
float alpha_r, float alpha_i,
|
|
181
|
+
float beta_r, float beta_i,
|
|
182
|
+
int batch_count
|
|
183
|
+
) {
|
|
184
|
+
struct BatchGemmParams params;
|
|
185
|
+
params.M = M;
|
|
186
|
+
params.N = N;
|
|
187
|
+
params.K = K;
|
|
188
|
+
params.lda = (transA == 0) ? K : M;
|
|
189
|
+
params.ldb = (transB == 0) ? N : K;
|
|
190
|
+
params.ldc = N;
|
|
191
|
+
params.strideA = strideA;
|
|
192
|
+
params.strideB = strideB;
|
|
193
|
+
params.strideC = strideC;
|
|
194
|
+
params.transA = transA;
|
|
195
|
+
params.transB = transB;
|
|
196
|
+
params.pad = 0;
|
|
197
|
+
params.alpha.x = alpha_r;
|
|
198
|
+
params.alpha.y = alpha_i;
|
|
199
|
+
params.beta.x = beta_r;
|
|
200
|
+
params.beta.y = beta_i;
|
|
201
|
+
params.batch_count = batch_count;
|
|
202
|
+
params.pad2 = 0;
|
|
203
|
+
|
|
204
|
+
[enc setComputePipelineState:g_pipe_gemm];
|
|
205
|
+
[enc setBuffer:bufA offset:offA atIndex:0];
|
|
206
|
+
[enc setBuffer:bufB offset:offB atIndex:1];
|
|
207
|
+
[enc setBuffer:bufC offset:offC atIndex:2];
|
|
208
|
+
[enc setBytes:¶ms length:sizeof(params) atIndex:3];
|
|
209
|
+
|
|
210
|
+
MTLSize threadsPerTG = MTLSizeMake(16, 16, 1);
|
|
211
|
+
MTLSize tgPerGrid = MTLSizeMake((N + 15) / 16, (M + 15) / 16, batch_count);
|
|
212
|
+
[enc dispatchThreadgroups:tgPerGrid threadsPerThreadgroup:threadsPerTG];
|
|
213
|
+
[enc memoryBarrierWithScope:MTLBarrierScopeBuffers];
|
|
214
|
+
}
|
|
215
|
+
|
|
216
|
+
static double matrix_norm1_c64(const complex64_t* A, int n) {
|
|
217
|
+
double worst = 0.0;
|
|
218
|
+
for (int j = 0; j < n; j++) {
|
|
219
|
+
double col_sum = 0.0;
|
|
220
|
+
for (int i = 0; i < n; i++) {
|
|
221
|
+
col_sum += cabsf(A[i * n + j]);
|
|
222
|
+
}
|
|
223
|
+
if (col_sum > worst) worst = col_sum;
|
|
224
|
+
}
|
|
225
|
+
return worst;
|
|
226
|
+
}
|
|
227
|
+
|
|
228
|
+
static double pair_norm1_c64(const complex64_t* Y, const complex64_t* L, int n) {
|
|
229
|
+
double worst = 0.0;
|
|
230
|
+
for (int j = 0; j < n; j++) {
|
|
231
|
+
double col_y = 0.0;
|
|
232
|
+
double col_l = 0.0;
|
|
233
|
+
for (int i = 0; i < n; i++) {
|
|
234
|
+
col_y += cabsf(Y[i * n + j]);
|
|
235
|
+
col_l += cabsf(L[i * n + j]);
|
|
236
|
+
}
|
|
237
|
+
if (col_y > worst) worst = col_y;
|
|
238
|
+
if (col_y + col_l > worst) worst = col_y + col_l;
|
|
239
|
+
}
|
|
240
|
+
return worst;
|
|
241
|
+
}
|
|
242
|
+
|
|
243
|
+
// -----------------------------------------------------------------------------
|
|
244
|
+
// Strided Batched Complex GEMM (c64)
|
|
245
|
+
// -----------------------------------------------------------------------------
|
|
246
|
+
int32_t batch_cgemm_strided_metal_c64(
|
|
247
|
+
BatchGemmTranspose transA,
|
|
248
|
+
BatchGemmTranspose transB,
|
|
249
|
+
int32_t M,
|
|
250
|
+
int32_t N,
|
|
251
|
+
int32_t K,
|
|
252
|
+
const complex64_t* alpha,
|
|
253
|
+
const complex64_t* A,
|
|
254
|
+
int32_t lda,
|
|
255
|
+
int64_t strideA,
|
|
256
|
+
const complex64_t* B,
|
|
257
|
+
int32_t ldb,
|
|
258
|
+
int64_t strideB,
|
|
259
|
+
const complex64_t* beta,
|
|
260
|
+
complex64_t* C,
|
|
261
|
+
int32_t ldc,
|
|
262
|
+
int64_t strideC,
|
|
263
|
+
int32_t batch_count
|
|
264
|
+
) {
|
|
265
|
+
if (batch_count <= 0) return 0;
|
|
266
|
+
init_metal_state();
|
|
267
|
+
if (!g_pipe_gemm || !g_queue || !g_device) return -1;
|
|
268
|
+
|
|
269
|
+
@autoreleasepool {
|
|
270
|
+
size_t bytesA = (size_t)batch_count * strideA * sizeof(complex64_t);
|
|
271
|
+
size_t bytesB = (size_t)batch_count * strideB * sizeof(complex64_t);
|
|
272
|
+
size_t bytesC = (size_t)batch_count * strideC * sizeof(complex64_t);
|
|
273
|
+
|
|
274
|
+
id<MTLBuffer> bufA = [g_device newBufferWithBytes:A length:bytesA options:MTLResourceStorageModeShared];
|
|
275
|
+
id<MTLBuffer> bufB = [g_device newBufferWithBytes:B length:bytesB options:MTLResourceStorageModeShared];
|
|
276
|
+
id<MTLBuffer> bufC = [g_device newBufferWithBytes:C length:bytesC options:MTLResourceStorageModeShared];
|
|
277
|
+
|
|
278
|
+
id<MTLCommandBuffer> cmd = [g_queue commandBuffer];
|
|
279
|
+
id<MTLComputeCommandEncoder> enc = [cmd computeCommandEncoder];
|
|
280
|
+
|
|
281
|
+
encode_cgemm(
|
|
282
|
+
enc,
|
|
283
|
+
bufA, 0, (int)strideA, metal_trans_mode(transA),
|
|
284
|
+
bufB, 0, (int)strideB, metal_trans_mode(transB),
|
|
285
|
+
bufC, 0, (int)strideC,
|
|
286
|
+
M, N, K,
|
|
287
|
+
crealf(*alpha), cimagf(*alpha),
|
|
288
|
+
crealf(*beta), cimagf(*beta),
|
|
289
|
+
batch_count
|
|
290
|
+
);
|
|
291
|
+
|
|
292
|
+
[enc endEncoding];
|
|
293
|
+
[cmd commit];
|
|
294
|
+
[cmd waitUntilCompleted];
|
|
295
|
+
|
|
296
|
+
memcpy(C, bufC.contents, bytesC);
|
|
297
|
+
}
|
|
298
|
+
return 0;
|
|
299
|
+
}
|
|
300
|
+
|
|
301
|
+
// -----------------------------------------------------------------------------
|
|
302
|
+
// Metal Batched Matrix Exponential (c64)
|
|
303
|
+
// -----------------------------------------------------------------------------
|
|
304
|
+
int32_t metal_matrix_exp_c64(
|
|
305
|
+
const complex64_t* A,
|
|
306
|
+
int32_t n,
|
|
307
|
+
int32_t batch_count,
|
|
308
|
+
complex64_t* out
|
|
309
|
+
) {
|
|
310
|
+
if (batch_count <= 0 || n <= 0) return 0;
|
|
311
|
+
init_metal_state();
|
|
312
|
+
if (!g_pipe_gemm || !g_queue || !g_device) return -1;
|
|
313
|
+
|
|
314
|
+
@autoreleasepool {
|
|
315
|
+
int matrix_elements = n * n;
|
|
316
|
+
int total_elements = batch_count * matrix_elements;
|
|
317
|
+
size_t buf_bytes = (size_t)total_elements * sizeof(complex64_t);
|
|
318
|
+
|
|
319
|
+
int max_squarings = 0;
|
|
320
|
+
int* squarings = (int*)malloc((size_t)batch_count * sizeof(int));
|
|
321
|
+
float* scales = (float*)malloc((size_t)batch_count * sizeof(float));
|
|
322
|
+
|
|
323
|
+
for (int b = 0; b < batch_count; b++) {
|
|
324
|
+
double norm = matrix_norm1_c64(A + b * matrix_elements, n);
|
|
325
|
+
int sq = 0;
|
|
326
|
+
if (norm > THETA18) {
|
|
327
|
+
sq = (int)ceil(log2(norm / THETA18));
|
|
328
|
+
}
|
|
329
|
+
squarings[b] = sq;
|
|
330
|
+
scales[b] = (sq > 0) ? (float)exp(-sq * log(2.0)) : 1.0f;
|
|
331
|
+
if (sq > max_squarings) max_squarings = sq;
|
|
332
|
+
}
|
|
333
|
+
|
|
334
|
+
id<MTLBuffer> bufM = [g_device newBufferWithBytes:A length:buf_bytes options:MTLResourceStorageModeShared];
|
|
335
|
+
id<MTLBuffer> bufM2 = [g_device newBufferWithLength:buf_bytes options:MTLResourceStorageModeShared];
|
|
336
|
+
id<MTLBuffer> bufM3 = [g_device newBufferWithLength:buf_bytes options:MTLResourceStorageModeShared];
|
|
337
|
+
id<MTLBuffer> bufAcc = [g_device newBufferWithLength:buf_bytes options:MTLResourceStorageModeShared];
|
|
338
|
+
id<MTLBuffer> bufTmp = [g_device newBufferWithLength:buf_bytes options:MTLResourceStorageModeShared];
|
|
339
|
+
id<MTLBuffer> bufScales = [g_device newBufferWithBytes:scales length:(size_t)batch_count * sizeof(float) options:MTLResourceStorageModeShared];
|
|
340
|
+
id<MTLBuffer> bufSquarings = [g_device newBufferWithBytes:squarings length:(size_t)batch_count * sizeof(int) options:MTLResourceStorageModeShared];
|
|
341
|
+
|
|
342
|
+
free(squarings);
|
|
343
|
+
free(scales);
|
|
344
|
+
|
|
345
|
+
id<MTLCommandBuffer> cmd = [g_queue commandBuffer];
|
|
346
|
+
id<MTLComputeCommandEncoder> enc = [cmd computeCommandEncoder];
|
|
347
|
+
|
|
348
|
+
struct ElementwiseParams ep;
|
|
349
|
+
ep.total_elements = total_elements;
|
|
350
|
+
ep.n = n;
|
|
351
|
+
ep.batch_count = batch_count;
|
|
352
|
+
ep.scalar_a = 0.0f;
|
|
353
|
+
ep.scalar_b = 0.0f;
|
|
354
|
+
ep.scalar_c = 0.0f;
|
|
355
|
+
|
|
356
|
+
// 1. Adaptive per-matrix scaling
|
|
357
|
+
[enc setComputePipelineState:g_pipe_scale_adaptive];
|
|
358
|
+
[enc setBuffer:bufM offset:0 atIndex:0];
|
|
359
|
+
[enc setBuffer:bufM offset:0 atIndex:1];
|
|
360
|
+
[enc setBuffer:bufScales offset:0 atIndex:2];
|
|
361
|
+
[enc setBytes:&ep length:sizeof(ep) atIndex:3];
|
|
362
|
+
[enc dispatchThreads:MTLSizeMake(total_elements, 1, 1) threadsPerThreadgroup:MTLSizeMake(256, 1, 1)];
|
|
363
|
+
[enc memoryBarrierWithScope:MTLBarrierScopeBuffers];
|
|
364
|
+
|
|
365
|
+
// 2. M2 = M * M, M3 = M2 * M
|
|
366
|
+
encode_cgemm(enc, bufM, 0, matrix_elements, 0, bufM, 0, matrix_elements, 0, bufM2, 0, matrix_elements, n, n, n, 1.0f, 0.0f, 0.0f, 0.0f, batch_count);
|
|
367
|
+
encode_cgemm(enc, bufM2, 0, matrix_elements, 0, bufM, 0, matrix_elements, 0, bufM3, 0, matrix_elements, n, n, n, 1.0f, 0.0f, 0.0f, 0.0f, batch_count);
|
|
368
|
+
|
|
369
|
+
// 3. Horner evaluation in M3
|
|
370
|
+
for (int j = 5; j >= 0; j--) {
|
|
371
|
+
if (j < 5) {
|
|
372
|
+
encode_cgemm(enc, bufAcc, 0, matrix_elements, 0, bufM3, 0, matrix_elements, 0, bufTmp, 0, matrix_elements, n, n, n, 1.0f, 0.0f, 0.0f, 0.0f, batch_count);
|
|
373
|
+
[enc setComputePipelineState:g_pipe_copy];
|
|
374
|
+
[enc setBuffer:bufTmp offset:0 atIndex:0];
|
|
375
|
+
[enc setBuffer:bufAcc offset:0 atIndex:1];
|
|
376
|
+
[enc setBytes:&ep length:sizeof(ep) atIndex:2];
|
|
377
|
+
[enc dispatchThreads:MTLSizeMake(total_elements, 1, 1) threadsPerThreadgroup:MTLSizeMake(256, 1, 1)];
|
|
378
|
+
[enc memoryBarrierWithScope:MTLBarrierScopeBuffers];
|
|
379
|
+
} else {
|
|
380
|
+
ep.scalar_a = 0.0f;
|
|
381
|
+
[enc setComputePipelineState:g_pipe_scale];
|
|
382
|
+
[enc setBuffer:bufAcc offset:0 atIndex:0];
|
|
383
|
+
[enc setBuffer:bufAcc offset:0 atIndex:1];
|
|
384
|
+
[enc setBytes:&ep length:sizeof(ep) atIndex:2];
|
|
385
|
+
[enc dispatchThreads:MTLSizeMake(total_elements, 1, 1) threadsPerThreadgroup:MTLSizeMake(256, 1, 1)];
|
|
386
|
+
[enc memoryBarrierWithScope:MTLBarrierScopeBuffers];
|
|
387
|
+
}
|
|
388
|
+
|
|
389
|
+
ep.scalar_a = (float)TAYLOR_COEFFS[j * 3 + 0];
|
|
390
|
+
ep.scalar_b = (float)TAYLOR_COEFFS[j * 3 + 1];
|
|
391
|
+
ep.scalar_c = (float)TAYLOR_COEFFS[j * 3 + 2];
|
|
392
|
+
|
|
393
|
+
[enc setComputePipelineState:g_pipe_horner];
|
|
394
|
+
[enc setBuffer:bufAcc offset:0 atIndex:0];
|
|
395
|
+
[enc setBuffer:bufM offset:0 atIndex:1];
|
|
396
|
+
[enc setBuffer:bufM2 offset:0 atIndex:2];
|
|
397
|
+
[enc setBytes:&ep length:sizeof(ep) atIndex:3];
|
|
398
|
+
MTLSize grid = MTLSizeMake(n, n, batch_count);
|
|
399
|
+
MTLSize tg = MTLSizeMake(16, 16, 1);
|
|
400
|
+
[enc dispatchThreads:grid threadsPerThreadgroup:tg];
|
|
401
|
+
[enc memoryBarrierWithScope:MTLBarrierScopeBuffers];
|
|
402
|
+
}
|
|
403
|
+
|
|
404
|
+
// 4. Squaring loop with adaptive gating
|
|
405
|
+
for (int k = 0; k < max_squarings; k++) {
|
|
406
|
+
encode_cgemm(enc, bufAcc, 0, matrix_elements, 0, bufAcc, 0, matrix_elements, 0, bufTmp, 0, matrix_elements, n, n, n, 1.0f, 0.0f, 0.0f, 0.0f, batch_count);
|
|
407
|
+
int current_k = k;
|
|
408
|
+
[enc setComputePipelineState:g_pipe_squaring_gate];
|
|
409
|
+
[enc setBuffer:bufTmp offset:0 atIndex:0];
|
|
410
|
+
[enc setBuffer:bufAcc offset:0 atIndex:1];
|
|
411
|
+
[enc setBuffer:bufAcc offset:0 atIndex:2];
|
|
412
|
+
[enc setBuffer:bufSquarings offset:0 atIndex:3];
|
|
413
|
+
[enc setBytes:¤t_k length:sizeof(int) atIndex:4];
|
|
414
|
+
[enc setBytes:&ep length:sizeof(ep) atIndex:5];
|
|
415
|
+
[enc dispatchThreads:MTLSizeMake(total_elements, 1, 1) threadsPerThreadgroup:MTLSizeMake(256, 1, 1)];
|
|
416
|
+
[enc memoryBarrierWithScope:MTLBarrierScopeBuffers];
|
|
417
|
+
}
|
|
418
|
+
|
|
419
|
+
[enc endEncoding];
|
|
420
|
+
[cmd commit];
|
|
421
|
+
[cmd waitUntilCompleted];
|
|
422
|
+
|
|
423
|
+
memcpy(out, bufAcc.contents, buf_bytes);
|
|
424
|
+
}
|
|
425
|
+
return 0;
|
|
426
|
+
}
|
|
427
|
+
|
|
428
|
+
// -----------------------------------------------------------------------------
|
|
429
|
+
// Metal Batched Matrix Exponential Adjoint (c64)
|
|
430
|
+
// -----------------------------------------------------------------------------
|
|
431
|
+
int32_t metal_matrix_exp_backward_c64(
|
|
432
|
+
const complex64_t* M,
|
|
433
|
+
const complex64_t* Ebar,
|
|
434
|
+
int32_t n,
|
|
435
|
+
int32_t batch_count,
|
|
436
|
+
int32_t dense,
|
|
437
|
+
complex64_t* Mbar
|
|
438
|
+
) {
|
|
439
|
+
if (batch_count <= 0 || n <= 0) return 0;
|
|
440
|
+
init_metal_state();
|
|
441
|
+
if (!g_pipe_gemm || !g_queue || !g_device) return -1;
|
|
442
|
+
|
|
443
|
+
@autoreleasepool {
|
|
444
|
+
int matrix_elements = n * n;
|
|
445
|
+
int total_elements = batch_count * matrix_elements;
|
|
446
|
+
size_t buf_bytes = (size_t)total_elements * sizeof(complex64_t);
|
|
447
|
+
|
|
448
|
+
complex64_t* host_Y = (complex64_t*)malloc(buf_bytes);
|
|
449
|
+
for (int b = 0; b < batch_count; b++) {
|
|
450
|
+
const complex64_t* cur_m = M + b * matrix_elements;
|
|
451
|
+
complex64_t* cur_y = host_Y + b * matrix_elements;
|
|
452
|
+
for (int i = 0; i < n; i++) {
|
|
453
|
+
for (int j = 0; j < n; j++) {
|
|
454
|
+
complex64_t z = cur_m[j * n + i];
|
|
455
|
+
cur_y[i * n + j] = crealf(z) - I * cimagf(z);
|
|
456
|
+
}
|
|
457
|
+
}
|
|
458
|
+
}
|
|
459
|
+
|
|
460
|
+
int max_squarings = 0;
|
|
461
|
+
int* squarings = (int*)malloc((size_t)batch_count * sizeof(int));
|
|
462
|
+
float* scales = (float*)malloc((size_t)batch_count * sizeof(float));
|
|
463
|
+
|
|
464
|
+
for (int b = 0; b < batch_count; b++) {
|
|
465
|
+
double pnorm = pair_norm1_c64(host_Y + b * matrix_elements, Ebar + b * matrix_elements, n);
|
|
466
|
+
int sq = 0;
|
|
467
|
+
if (pnorm > THETA18) {
|
|
468
|
+
sq = (int)ceil(log2(pnorm / THETA18));
|
|
469
|
+
}
|
|
470
|
+
squarings[b] = sq;
|
|
471
|
+
scales[b] = (sq > 0) ? (float)exp(-sq * log(2.0)) : 1.0f;
|
|
472
|
+
if (sq > max_squarings) max_squarings = sq;
|
|
473
|
+
}
|
|
474
|
+
|
|
475
|
+
id<MTLBuffer> bufY = [g_device newBufferWithBytes:host_Y length:buf_bytes options:MTLResourceStorageModeShared];
|
|
476
|
+
id<MTLBuffer> bufL = [g_device newBufferWithBytes:Ebar length:buf_bytes options:MTLResourceStorageModeShared];
|
|
477
|
+
id<MTLBuffer> bufY2 = [g_device newBufferWithLength:buf_bytes options:MTLResourceStorageModeShared];
|
|
478
|
+
id<MTLBuffer> bufL2 = [g_device newBufferWithLength:buf_bytes options:MTLResourceStorageModeShared];
|
|
479
|
+
id<MTLBuffer> bufY3 = [g_device newBufferWithLength:buf_bytes options:MTLResourceStorageModeShared];
|
|
480
|
+
id<MTLBuffer> bufL3 = [g_device newBufferWithLength:buf_bytes options:MTLResourceStorageModeShared];
|
|
481
|
+
id<MTLBuffer> bufAccY = [g_device newBufferWithLength:buf_bytes options:MTLResourceStorageModeShared];
|
|
482
|
+
id<MTLBuffer> bufAccL = [g_device newBufferWithLength:buf_bytes options:MTLResourceStorageModeShared];
|
|
483
|
+
id<MTLBuffer> bufTmpY = [g_device newBufferWithLength:buf_bytes options:MTLResourceStorageModeShared];
|
|
484
|
+
id<MTLBuffer> bufTmpL = [g_device newBufferWithLength:buf_bytes options:MTLResourceStorageModeShared];
|
|
485
|
+
id<MTLBuffer> bufScales = [g_device newBufferWithBytes:scales length:(size_t)batch_count * sizeof(float) options:MTLResourceStorageModeShared];
|
|
486
|
+
id<MTLBuffer> bufSquarings = [g_device newBufferWithBytes:squarings length:(size_t)batch_count * sizeof(int) options:MTLResourceStorageModeShared];
|
|
487
|
+
|
|
488
|
+
free(host_Y);
|
|
489
|
+
free(squarings);
|
|
490
|
+
free(scales);
|
|
491
|
+
|
|
492
|
+
id<MTLCommandBuffer> cmd = [g_queue commandBuffer];
|
|
493
|
+
id<MTLComputeCommandEncoder> enc = [cmd computeCommandEncoder];
|
|
494
|
+
|
|
495
|
+
struct ElementwiseParams ep;
|
|
496
|
+
ep.total_elements = total_elements;
|
|
497
|
+
ep.n = n;
|
|
498
|
+
ep.batch_count = batch_count;
|
|
499
|
+
ep.scalar_a = 0.0f;
|
|
500
|
+
ep.scalar_b = 0.0f;
|
|
501
|
+
ep.scalar_c = 0.0f;
|
|
502
|
+
|
|
503
|
+
// 1. Adaptive scaling of Y and L
|
|
504
|
+
[enc setComputePipelineState:g_pipe_scale_adaptive];
|
|
505
|
+
[enc setBuffer:bufY offset:0 atIndex:0];
|
|
506
|
+
[enc setBuffer:bufY offset:0 atIndex:1];
|
|
507
|
+
[enc setBuffer:bufScales offset:0 atIndex:2];
|
|
508
|
+
[enc setBytes:&ep length:sizeof(ep) atIndex:3];
|
|
509
|
+
[enc dispatchThreads:MTLSizeMake(total_elements, 1, 1) threadsPerThreadgroup:MTLSizeMake(256, 1, 1)];
|
|
510
|
+
|
|
511
|
+
[enc setBuffer:bufL offset:0 atIndex:0];
|
|
512
|
+
[enc setBuffer:bufL offset:0 atIndex:1];
|
|
513
|
+
[enc setBuffer:bufScales offset:0 atIndex:2];
|
|
514
|
+
[enc setBytes:&ep length:sizeof(ep) atIndex:3];
|
|
515
|
+
[enc dispatchThreads:MTLSizeMake(total_elements, 1, 1) threadsPerThreadgroup:MTLSizeMake(256, 1, 1)];
|
|
516
|
+
[enc memoryBarrierWithScope:MTLBarrierScopeBuffers];
|
|
517
|
+
|
|
518
|
+
void (^encode_pair_mul)(id<MTLBuffer>, id<MTLBuffer>, id<MTLBuffer>, id<MTLBuffer>, id<MTLBuffer>, id<MTLBuffer>) =
|
|
519
|
+
^(id<MTLBuffer> y1, id<MTLBuffer> l1, id<MTLBuffer> y2, id<MTLBuffer> l2, id<MTLBuffer> yo, id<MTLBuffer> lo) {
|
|
520
|
+
encode_cgemm(enc, y1, 0, matrix_elements, 0, y2, 0, matrix_elements, 0, yo, 0, matrix_elements, n, n, n, 1.0f, 0.0f, 0.0f, 0.0f, batch_count);
|
|
521
|
+
encode_cgemm(enc, y1, 0, matrix_elements, 0, l2, 0, matrix_elements, 0, lo, 0, matrix_elements, n, n, n, 1.0f, 0.0f, 0.0f, 0.0f, batch_count);
|
|
522
|
+
encode_cgemm(enc, l1, 0, matrix_elements, 0, y2, 0, matrix_elements, 0, lo, 0, matrix_elements, n, n, n, 1.0f, 0.0f, 1.0f, 0.0f, batch_count);
|
|
523
|
+
};
|
|
524
|
+
|
|
525
|
+
encode_pair_mul(bufY, bufL, bufY, bufL, bufY2, bufL2);
|
|
526
|
+
encode_pair_mul(bufY2, bufL2, bufY, bufL, bufY3, bufL3);
|
|
527
|
+
|
|
528
|
+
for (int j = 5; j >= 0; j--) {
|
|
529
|
+
if (j < 5) {
|
|
530
|
+
encode_pair_mul(bufAccY, bufAccL, bufY3, bufL3, bufTmpY, bufTmpL);
|
|
531
|
+
[enc setComputePipelineState:g_pipe_copy];
|
|
532
|
+
[enc setBuffer:bufTmpY offset:0 atIndex:0];
|
|
533
|
+
[enc setBuffer:bufAccY offset:0 atIndex:1];
|
|
534
|
+
[enc setBytes:&ep length:sizeof(ep) atIndex:2];
|
|
535
|
+
[enc dispatchThreads:MTLSizeMake(total_elements, 1, 1) threadsPerThreadgroup:MTLSizeMake(256, 1, 1)];
|
|
536
|
+
|
|
537
|
+
[enc setBuffer:bufTmpL offset:0 atIndex:0];
|
|
538
|
+
[enc setBuffer:bufAccL offset:0 atIndex:1];
|
|
539
|
+
[enc setBytes:&ep length:sizeof(ep) atIndex:2];
|
|
540
|
+
[enc dispatchThreads:MTLSizeMake(total_elements, 1, 1) threadsPerThreadgroup:MTLSizeMake(256, 1, 1)];
|
|
541
|
+
[enc memoryBarrierWithScope:MTLBarrierScopeBuffers];
|
|
542
|
+
} else {
|
|
543
|
+
ep.scalar_a = 0.0f;
|
|
544
|
+
[enc setComputePipelineState:g_pipe_scale];
|
|
545
|
+
[enc setBuffer:bufAccY offset:0 atIndex:0];
|
|
546
|
+
[enc setBuffer:bufAccY offset:0 atIndex:1];
|
|
547
|
+
[enc setBytes:&ep length:sizeof(ep) atIndex:2];
|
|
548
|
+
[enc dispatchThreads:MTLSizeMake(total_elements, 1, 1) threadsPerThreadgroup:MTLSizeMake(256, 1, 1)];
|
|
549
|
+
|
|
550
|
+
[enc setBuffer:bufAccL offset:0 atIndex:0];
|
|
551
|
+
[enc setBuffer:bufAccL offset:0 atIndex:1];
|
|
552
|
+
[enc setBytes:&ep length:sizeof(ep) atIndex:2];
|
|
553
|
+
[enc dispatchThreads:MTLSizeMake(total_elements, 1, 1) threadsPerThreadgroup:MTLSizeMake(256, 1, 1)];
|
|
554
|
+
[enc memoryBarrierWithScope:MTLBarrierScopeBuffers];
|
|
555
|
+
}
|
|
556
|
+
|
|
557
|
+
ep.scalar_a = (float)TAYLOR_COEFFS[j * 3 + 0];
|
|
558
|
+
ep.scalar_b = (float)TAYLOR_COEFFS[j * 3 + 1];
|
|
559
|
+
ep.scalar_c = (float)TAYLOR_COEFFS[j * 3 + 2];
|
|
560
|
+
|
|
561
|
+
[enc setComputePipelineState:g_pipe_horner_pair];
|
|
562
|
+
[enc setBuffer:bufAccY offset:0 atIndex:0];
|
|
563
|
+
[enc setBuffer:bufAccL offset:0 atIndex:1];
|
|
564
|
+
[enc setBuffer:bufY offset:0 atIndex:2];
|
|
565
|
+
[enc setBuffer:bufL offset:0 atIndex:3];
|
|
566
|
+
[enc setBuffer:bufY2 offset:0 atIndex:4];
|
|
567
|
+
[enc setBuffer:bufL2 offset:0 atIndex:5];
|
|
568
|
+
[enc setBytes:&ep length:sizeof(ep) atIndex:6];
|
|
569
|
+
MTLSize grid = MTLSizeMake(n, n, batch_count);
|
|
570
|
+
MTLSize tg = MTLSizeMake(16, 16, 1);
|
|
571
|
+
[enc dispatchThreads:grid threadsPerThreadgroup:tg];
|
|
572
|
+
[enc memoryBarrierWithScope:MTLBarrierScopeBuffers];
|
|
573
|
+
}
|
|
574
|
+
|
|
575
|
+
for (int k = 0; k < max_squarings; k++) {
|
|
576
|
+
encode_pair_mul(bufAccY, bufAccL, bufAccY, bufAccL, bufTmpY, bufTmpL);
|
|
577
|
+
int current_k = k;
|
|
578
|
+
[enc setComputePipelineState:g_pipe_squaring_gate];
|
|
579
|
+
[enc setBuffer:bufTmpY offset:0 atIndex:0];
|
|
580
|
+
[enc setBuffer:bufAccY offset:0 atIndex:1];
|
|
581
|
+
[enc setBuffer:bufAccY offset:0 atIndex:2];
|
|
582
|
+
[enc setBuffer:bufSquarings offset:0 atIndex:3];
|
|
583
|
+
[enc setBytes:¤t_k length:sizeof(int) atIndex:4];
|
|
584
|
+
[enc setBytes:&ep length:sizeof(ep) atIndex:5];
|
|
585
|
+
[enc dispatchThreads:MTLSizeMake(total_elements, 1, 1) threadsPerThreadgroup:MTLSizeMake(256, 1, 1)];
|
|
586
|
+
|
|
587
|
+
[enc setBuffer:bufTmpL offset:0 atIndex:0];
|
|
588
|
+
[enc setBuffer:bufAccL offset:0 atIndex:1];
|
|
589
|
+
[enc setBuffer:bufAccL offset:0 atIndex:2];
|
|
590
|
+
[enc setBuffer:bufSquarings offset:0 atIndex:3];
|
|
591
|
+
[enc setBytes:¤t_k length:sizeof(int) atIndex:4];
|
|
592
|
+
[enc setBytes:&ep length:sizeof(ep) atIndex:5];
|
|
593
|
+
[enc dispatchThreads:MTLSizeMake(total_elements, 1, 1) threadsPerThreadgroup:MTLSizeMake(256, 1, 1)];
|
|
594
|
+
[enc memoryBarrierWithScope:MTLBarrierScopeBuffers];
|
|
595
|
+
}
|
|
596
|
+
|
|
597
|
+
[enc endEncoding];
|
|
598
|
+
[cmd commit];
|
|
599
|
+
[cmd waitUntilCompleted];
|
|
600
|
+
|
|
601
|
+
memcpy(Mbar, bufAccL.contents, buf_bytes);
|
|
602
|
+
}
|
|
603
|
+
return 0;
|
|
604
|
+
}
|
|
605
|
+
|
|
606
|
+
// -----------------------------------------------------------------------------
|
|
607
|
+
// Metal Structure Factors (c128 in host, float compute on GPU)
|
|
608
|
+
// -----------------------------------------------------------------------------
|
|
609
|
+
int32_t metal_structure_factors_c128(
|
|
610
|
+
complex128_t* fgb,
|
|
611
|
+
const double* g2,
|
|
612
|
+
const double* hkl,
|
|
613
|
+
const double* lobato_a,
|
|
614
|
+
const double* lobato_b,
|
|
615
|
+
const double* uij,
|
|
616
|
+
const double* pos,
|
|
617
|
+
const double* occ,
|
|
618
|
+
int32_t n_grid,
|
|
619
|
+
int32_t n_atoms,
|
|
620
|
+
double volume
|
|
621
|
+
) {
|
|
622
|
+
if (n_grid <= 0 || n_atoms <= 0) return 0;
|
|
623
|
+
init_metal_state();
|
|
624
|
+
if (!g_pipe_sf || !g_queue || !g_device) return -1;
|
|
625
|
+
|
|
626
|
+
@autoreleasepool {
|
|
627
|
+
float* f_g2 = (float*)malloc((size_t)n_grid * sizeof(float));
|
|
628
|
+
float* f_hkl = (float*)malloc((size_t)n_grid * 3 * sizeof(float));
|
|
629
|
+
for (int i = 0; i < n_grid; i++) f_g2[i] = (float)g2[i];
|
|
630
|
+
for (int i = 0; i < n_grid * 3; i++) f_hkl[i] = (float)hkl[i];
|
|
631
|
+
|
|
632
|
+
float* f_la = (float*)malloc((size_t)n_atoms * 5 * sizeof(float));
|
|
633
|
+
float* f_lb = (float*)malloc((size_t)n_atoms * 5 * sizeof(float));
|
|
634
|
+
float* f_uij = (float*)malloc((size_t)n_atoms * 9 * sizeof(float));
|
|
635
|
+
float* f_pos = (float*)malloc((size_t)n_atoms * 3 * sizeof(float));
|
|
636
|
+
float* f_occ = (float*)malloc((size_t)n_atoms * sizeof(float));
|
|
637
|
+
|
|
638
|
+
for (int i = 0; i < n_atoms * 5; i++) f_la[i] = (float)lobato_a[i];
|
|
639
|
+
for (int i = 0; i < n_atoms * 5; i++) f_lb[i] = (float)lobato_b[i];
|
|
640
|
+
for (int i = 0; i < n_atoms * 9; i++) f_uij[i] = (float)uij[i];
|
|
641
|
+
for (int i = 0; i < n_atoms * 3; i++) f_pos[i] = (float)pos[i];
|
|
642
|
+
for (int i = 0; i < n_atoms; i++) f_occ[i] = (float)occ[i];
|
|
643
|
+
|
|
644
|
+
id<MTLBuffer> bufFgb = [g_device newBufferWithLength:(size_t)n_grid * sizeof(complex64_t) options:MTLResourceStorageModeShared];
|
|
645
|
+
id<MTLBuffer> bufG2 = [g_device newBufferWithBytes:f_g2 length:(size_t)n_grid * sizeof(float) options:MTLResourceStorageModeShared];
|
|
646
|
+
id<MTLBuffer> bufHkl = [g_device newBufferWithBytes:f_hkl length:(size_t)n_grid * 3 * sizeof(float) options:MTLResourceStorageModeShared];
|
|
647
|
+
id<MTLBuffer> bufLa = [g_device newBufferWithBytes:f_la length:(size_t)n_atoms * 5 * sizeof(float) options:MTLResourceStorageModeShared];
|
|
648
|
+
id<MTLBuffer> bufLb = [g_device newBufferWithBytes:f_lb length:(size_t)n_atoms * 5 * sizeof(float) options:MTLResourceStorageModeShared];
|
|
649
|
+
id<MTLBuffer> bufUij = [g_device newBufferWithBytes:f_uij length:(size_t)n_atoms * 9 * sizeof(float) options:MTLResourceStorageModeShared];
|
|
650
|
+
id<MTLBuffer> bufPos = [g_device newBufferWithBytes:f_pos length:(size_t)n_atoms * 3 * sizeof(float) options:MTLResourceStorageModeShared];
|
|
651
|
+
id<MTLBuffer> bufOcc = [g_device newBufferWithBytes:f_occ length:(size_t)n_atoms * sizeof(float) options:MTLResourceStorageModeShared];
|
|
652
|
+
|
|
653
|
+
free(f_g2); free(f_hkl); free(f_la); free(f_lb); free(f_uij); free(f_pos); free(f_occ);
|
|
654
|
+
|
|
655
|
+
struct StructureFactorParams p;
|
|
656
|
+
p.n_grid = n_grid;
|
|
657
|
+
p.n_atoms = n_atoms;
|
|
658
|
+
p.volume = (float)volume;
|
|
659
|
+
p.two_pi = 6.28318530717958647692f;
|
|
660
|
+
p.minus_two_pi_sq = -19.73920880217871723766f;
|
|
661
|
+
p.absorption = 0;
|
|
662
|
+
|
|
663
|
+
id<MTLCommandBuffer> cmd = [g_queue commandBuffer];
|
|
664
|
+
id<MTLComputeCommandEncoder> enc = [cmd computeCommandEncoder];
|
|
665
|
+
|
|
666
|
+
[enc setComputePipelineState:g_pipe_sf];
|
|
667
|
+
[enc setBuffer:bufFgb offset:0 atIndex:0];
|
|
668
|
+
[enc setBuffer:bufG2 offset:0 atIndex:1];
|
|
669
|
+
[enc setBuffer:bufHkl offset:0 atIndex:2];
|
|
670
|
+
[enc setBuffer:bufLa offset:0 atIndex:3];
|
|
671
|
+
[enc setBuffer:bufLb offset:0 atIndex:4];
|
|
672
|
+
[enc setBuffer:bufUij offset:0 atIndex:5];
|
|
673
|
+
[enc setBuffer:bufPos offset:0 atIndex:6];
|
|
674
|
+
[enc setBuffer:bufOcc offset:0 atIndex:7];
|
|
675
|
+
[enc setBytes:&p length:sizeof(p) atIndex:8];
|
|
676
|
+
|
|
677
|
+
[enc dispatchThreads:MTLSizeMake(n_grid, 1, 1) threadsPerThreadgroup:MTLSizeMake(256, 1, 1)];
|
|
678
|
+
[enc endEncoding];
|
|
679
|
+
|
|
680
|
+
[cmd commit];
|
|
681
|
+
[cmd waitUntilCompleted];
|
|
682
|
+
|
|
683
|
+
complex64_t* res_c64 = (complex64_t*)bufFgb.contents;
|
|
684
|
+
for (int i = 0; i < n_grid; i++) {
|
|
685
|
+
fgb[i] = (double)crealf(res_c64[i]) + I * (double)cimagf(res_c64[i]);
|
|
686
|
+
}
|
|
687
|
+
}
|
|
688
|
+
return 0;
|
|
689
|
+
}
|
|
690
|
+
|
|
691
|
+
// -----------------------------------------------------------------------------
|
|
692
|
+
// Metal Structure Matrix Assembly (c128 in host, float compute on GPU)
|
|
693
|
+
// -----------------------------------------------------------------------------
|
|
694
|
+
int32_t metal_structure_matrix_c128(
|
|
695
|
+
const complex128_t* fgb,
|
|
696
|
+
int32_t n_grid,
|
|
697
|
+
const int32_t* source,
|
|
698
|
+
int32_t buffer_size,
|
|
699
|
+
const int32_t* destination,
|
|
700
|
+
int32_t n_beams,
|
|
701
|
+
const double* mii,
|
|
702
|
+
const double* diagonal,
|
|
703
|
+
int32_t n_batch,
|
|
704
|
+
double prefactor,
|
|
705
|
+
int32_t absorption,
|
|
706
|
+
complex128_t* out
|
|
707
|
+
) {
|
|
708
|
+
if (n_beams <= 0 || n_batch <= 0 || n_grid <= 0 || buffer_size <= 0) return 0;
|
|
709
|
+
init_metal_state();
|
|
710
|
+
if (!g_pipe_sm_scatter || !g_pipe_sm_gather || !g_queue || !g_device) return -1;
|
|
711
|
+
|
|
712
|
+
@autoreleasepool {
|
|
713
|
+
int pairs = n_beams * n_beams;
|
|
714
|
+
int total_pairs = n_batch * pairs;
|
|
715
|
+
|
|
716
|
+
complex64_t* f_fgb = (complex64_t*)malloc((size_t)n_grid * sizeof(complex64_t));
|
|
717
|
+
for (int i = 0; i < n_grid; i++) f_fgb[i] = (float)creal(fgb[i]) + I * (float)cimag(fgb[i]);
|
|
718
|
+
|
|
719
|
+
float* f_mii = (float*)malloc((size_t)n_batch * n_beams * sizeof(float));
|
|
720
|
+
float* f_diag = (float*)malloc((size_t)n_batch * n_beams * sizeof(float));
|
|
721
|
+
for (int i = 0; i < n_batch * n_beams; i++) f_mii[i] = (float)mii[i];
|
|
722
|
+
for (int i = 0; i < n_batch * n_beams; i++) f_diag[i] = (float)diagonal[i];
|
|
723
|
+
|
|
724
|
+
id<MTLBuffer> bufBuffer = [g_device newBufferWithLength:(size_t)buffer_size * sizeof(complex64_t) options:MTLResourceStorageModeShared];
|
|
725
|
+
id<MTLBuffer> bufFgb = [g_device newBufferWithBytes:f_fgb length:(size_t)n_grid * sizeof(complex64_t) options:MTLResourceStorageModeShared];
|
|
726
|
+
id<MTLBuffer> bufSource = [g_device newBufferWithBytes:source length:(size_t)n_grid * sizeof(int32_t) options:MTLResourceStorageModeShared];
|
|
727
|
+
id<MTLBuffer> bufDst = [g_device newBufferWithBytes:destination length:(size_t)pairs * sizeof(int32_t) options:MTLResourceStorageModeShared];
|
|
728
|
+
id<MTLBuffer> bufMii = [g_device newBufferWithBytes:f_mii length:(size_t)n_batch * n_beams * sizeof(float) options:MTLResourceStorageModeShared];
|
|
729
|
+
id<MTLBuffer> bufDiag = [g_device newBufferWithBytes:f_diag length:(size_t)n_batch * n_beams * sizeof(float) options:MTLResourceStorageModeShared];
|
|
730
|
+
id<MTLBuffer> bufOut = [g_device newBufferWithLength:(size_t)total_pairs * sizeof(complex64_t) options:MTLResourceStorageModeShared];
|
|
731
|
+
|
|
732
|
+
free(f_fgb); free(f_mii); free(f_diag);
|
|
733
|
+
|
|
734
|
+
struct StructureMatrixParams p;
|
|
735
|
+
p.n_beams = n_beams;
|
|
736
|
+
p.n_batch = n_batch;
|
|
737
|
+
p.buffer_size = buffer_size;
|
|
738
|
+
p.n_grid = n_grid;
|
|
739
|
+
p.prefactor = (float)prefactor;
|
|
740
|
+
p.u0_prime = 0.0f;
|
|
741
|
+
|
|
742
|
+
id<MTLCommandBuffer> cmd = [g_queue commandBuffer];
|
|
743
|
+
id<MTLComputeCommandEncoder> enc = [cmd computeCommandEncoder];
|
|
744
|
+
|
|
745
|
+
// 1. Scatter into buffer
|
|
746
|
+
[enc setComputePipelineState:g_pipe_sm_scatter];
|
|
747
|
+
[enc setBuffer:bufBuffer offset:0 atIndex:0];
|
|
748
|
+
[enc setBuffer:bufFgb offset:0 atIndex:1];
|
|
749
|
+
[enc setBuffer:bufSource offset:0 atIndex:2];
|
|
750
|
+
[enc setBytes:&p length:sizeof(p) atIndex:3];
|
|
751
|
+
[enc dispatchThreads:MTLSizeMake(n_grid, 1, 1) threadsPerThreadgroup:MTLSizeMake(256, 1, 1)];
|
|
752
|
+
[enc memoryBarrierWithScope:MTLBarrierScopeBuffers];
|
|
753
|
+
|
|
754
|
+
// 2. Gather into A
|
|
755
|
+
[enc setComputePipelineState:g_pipe_sm_gather];
|
|
756
|
+
[enc setBuffer:bufOut offset:0 atIndex:0];
|
|
757
|
+
[enc setBuffer:bufBuffer offset:0 atIndex:1];
|
|
758
|
+
[enc setBuffer:bufDst offset:0 atIndex:2];
|
|
759
|
+
[enc setBuffer:bufMii offset:0 atIndex:3];
|
|
760
|
+
[enc setBuffer:bufDiag offset:0 atIndex:4];
|
|
761
|
+
[enc setBytes:&p length:sizeof(p) atIndex:5];
|
|
762
|
+
MTLSize grid = MTLSizeMake(n_beams, n_beams, n_batch);
|
|
763
|
+
MTLSize tg = MTLSizeMake(16, 16, 1);
|
|
764
|
+
[enc dispatchThreads:grid threadsPerThreadgroup:tg];
|
|
765
|
+
|
|
766
|
+
[enc endEncoding];
|
|
767
|
+
[cmd commit];
|
|
768
|
+
[cmd waitUntilCompleted];
|
|
769
|
+
|
|
770
|
+
complex64_t* res_c64 = (complex64_t*)bufOut.contents;
|
|
771
|
+
for (int i = 0; i < total_pairs; i++) {
|
|
772
|
+
out[i] = (double)crealf(res_c64[i]) + I * (double)cimagf(res_c64[i]);
|
|
773
|
+
}
|
|
774
|
+
}
|
|
775
|
+
return 0;
|
|
776
|
+
}
|