DeepGPR 0.0.1__tar.gz → 0.0.2__tar.gz

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,842 @@
1
+ #include <cuda_runtime.h>
2
+ #include <iostream>
3
+ #include <stdio.h>
4
+ #include <cfloat>
5
+
6
+ __constant__ float e0 = 8.8541878128e-12;
7
+ __constant__ float m0 = 1.25663706212e-06;
8
+
9
+ #define CEIL_DIV(x,y) (((x)+(y)-1)/(y))
10
+
11
+ #define CUDA_CHECK() {\
12
+ cudaError_t err = cudaGetLastError();\
13
+ if (err != cudaSuccess) {\
14
+ std::cerr << "CUDA Error: " << cudaGetErrorString(err) \
15
+ << " at " << __FILE__ << ":" << __LINE__ << std::endl;\
16
+ exit(EXIT_FAILURE);\
17
+ }\
18
+ }
19
+
20
+ // ---------------------------------------------------------
21
+ // 系数获取核函数
22
+ // ---------------------------------------------------------
23
+ __global__ void ucgetforward(const float* __restrict__ er, const float* __restrict__ se, const float* __restrict__ mr,
24
+ float* __restrict__ uE0, float* __restrict__ uE1, float* __restrict__ uE4,
25
+ float* __restrict__ uH0, float* __restrict__ uH1, float* __restrict__ uH4,
26
+ int NX_FIELDS, int NY_FIELDS, int NZ_FIELDS, float dt, float dx)
27
+ {
28
+ long long idx = blockIdx.x * blockDim.x + threadIdx.x;
29
+ long long ny_nz = (long long)NY_FIELDS * NZ_FIELDS;
30
+ if (idx >= (long long)NX_FIELDS * ny_nz) return;
31
+
32
+ long long i = idx / ny_nz;
33
+ long long rem = idx % ny_nz;
34
+ long long j = rem / NZ_FIELDS;
35
+ long long k = rem % NZ_FIELDS;
36
+
37
+ if (i < (NX_FIELDS-1) && j < (NY_FIELDS-1) && k < (NZ_FIELDS-1) ) {
38
+ float HA = m0 * mr[idx] / dt;
39
+ uH0[idx] = 1.0f;
40
+ uH1[idx] = (1.0f / dx) / HA;
41
+ uH4[idx] = 1.0f / HA;
42
+
43
+ if (se[idx] > 100.0f) {
44
+ uE0[idx] = 0.0f; uE1[idx] = 0.0f; uE4[idx] = 0.0f;
45
+ } else {
46
+ float e_term = e0 * er[idx] / dt;
47
+ float s_term = 0.5f * se[idx];
48
+ float EA = e_term + s_term;
49
+ float EB = e_term - s_term;
50
+ uE0[idx] = EB / EA;
51
+ uE1[idx] = (1.0f / dx) / EA;
52
+ uE4[idx] = 1.0f / EA;
53
+ }
54
+ }
55
+ }
56
+
57
+ __global__ void ucgetbackward(const float* __restrict__ er, const float* __restrict__ se, const float* __restrict__ mr,
58
+ float* __restrict__ uE0, float* __restrict__ uE1, float* __restrict__ uE4,
59
+ float* __restrict__ uH0, float* __restrict__ uH1, float* __restrict__ uH4,
60
+ int NX_FIELDS, int NY_FIELDS, int NZ_FIELDS, float dt, float dx)
61
+ {
62
+ long long idx = blockIdx.x * blockDim.x + threadIdx.x;
63
+ long long ny_nz = (long long)NY_FIELDS * NZ_FIELDS;
64
+ if (idx >= (long long)NX_FIELDS * ny_nz) return;
65
+
66
+ long long i = idx / ny_nz;
67
+ long long rem = idx % ny_nz;
68
+ long long j = rem / NZ_FIELDS;
69
+ long long k = rem % NZ_FIELDS;
70
+
71
+ if (i < (NX_FIELDS-1) && j < (NY_FIELDS-1) && k < (NZ_FIELDS-1) ) {
72
+ float HA = m0 * mr[idx] / dt;
73
+ uH0[idx] = 1.0f;
74
+ uH1[idx] = (1.0f / dx) / HA;
75
+ uH4[idx] = 1.0f / HA;
76
+
77
+ if (se[idx] > 100.0f) {
78
+ uE0[idx] = 0.0f; uE1[idx] = 0.0f; uE4[idx] = 0.0f;
79
+ } else {
80
+ float EA = (e0 * er[idx] / dt) + 0.5f * se[idx];
81
+ uE0[idx] = (2.0f * e0 * er[idx]) / (2.0f * e0 * er[idx] + se[idx] * dt);
82
+ uE1[idx] = (1.0f / dx) / EA;
83
+ uE4[idx] = 1.0f / EA;
84
+ }
85
+ }
86
+ }
87
+
88
+ // ---------------------------------------------------------
89
+ // 接收端保存
90
+ // ---------------------------------------------------------
91
+ __global__ void store_outputs(
92
+ int step, int NRX, int iteration,
93
+ const int* __restrict__ receiverlocation, float* __restrict__ rxs,
94
+ const float* __restrict__ Ex, const float* __restrict__ Ey, const float* __restrict__ Ez,
95
+ const float* __restrict__ Hx, const float* __restrict__ Hy, const float* __restrict__ Hz,
96
+ int NX, int NY, int NZ, int N_ITER)
97
+ {
98
+ long long rx = blockIdx.x * blockDim.x + threadIdx.x;
99
+ if (rx >= NRX) return;
100
+
101
+ long long field_stride = (long long)NX * NY * NZ;
102
+
103
+ for (int s = 0; s < step; ++s) {
104
+ long long i = receiverlocation[s * NRX * 3 + rx * 3 + 0];
105
+ long long j = receiverlocation[s * NRX * 3 + rx * 3 + 1];
106
+ long long k = receiverlocation[s * NRX * 3 + rx * 3 + 2];
107
+
108
+ long long id4 = s * field_stride + i * NY * NZ + j * NZ + k;
109
+
110
+ rxs[((s * 6 + 0) * N_ITER + iteration) * NRX + rx] = Ex[id4];
111
+ rxs[((s * 6 + 1) * N_ITER + iteration) * NRX + rx] = Ey[id4];
112
+ rxs[((s * 6 + 2) * N_ITER + iteration) * NRX + rx] = Ez[id4];
113
+ rxs[((s * 6 + 3) * N_ITER + iteration) * NRX + rx] = Hx[id4];
114
+ rxs[((s * 6 + 4) * N_ITER + iteration) * NRX + rx] = Hy[id4];
115
+ rxs[((s * 6 + 5) * N_ITER + iteration) * NRX + rx] = Hz[id4];
116
+ }
117
+ }
118
+
119
+ // ---------------------------------------------------------
120
+ // 震源更新
121
+ // ---------------------------------------------------------
122
+ __global__ void Update_hertzian_dipole(
123
+ int step, int iteration, float dx,
124
+ const int* __restrict__ sourcelocation, const float* __restrict__ srcwaveforms,
125
+ float* __restrict__ Ex, float* __restrict__ Ey, float* __restrict__ Ez, const float* __restrict__ uE4,
126
+ int NX, int NY, int NZ, int nsrc, int polarisation, int nt)
127
+ {
128
+ long long src = blockIdx.x * blockDim.x + threadIdx.x;
129
+ if (src >= nsrc) return;
130
+
131
+ float waveform_value = srcwaveforms[src * nt + iteration];
132
+ float scale = waveform_value * dx / (dx * dx * dx);
133
+ long long field_stride = (long long)NX * NY * NZ;
134
+
135
+ for (int s = 0; s < step; ++s) {
136
+ long long i = sourcelocation[s * nsrc * 3 + src * 3 + 0];
137
+ long long j = sourcelocation[s * nsrc * 3 + src * 3 + 1];
138
+ long long k = sourcelocation[s * nsrc * 3 + src * 3 + 2];
139
+
140
+ long long id3 = i * NY * NZ + j * NZ + k;
141
+ long long id4 = s * field_stride + id3;
142
+
143
+ if (polarisation == 0) Ex[id4] -= uE4[id3] * scale;
144
+ else if (polarisation == 1) Ey[id4] -= uE4[id3] * scale;
145
+ else if (polarisation == 2) Ez[id4] -= uE4[id3] * scale;
146
+ }
147
+ }
148
+
149
+ // ---------------------------------------------------------
150
+ // 融合:电场全局更新 + 全侧 PML 边界修正
151
+ // ---------------------------------------------------------
152
+ __global__ void fused_e_fields_updates_gpu(
153
+ const float* __restrict__ uE0, const float* __restrict__ uE1,
154
+ float* __restrict__ Ex, float* __restrict__ Ey, float* __restrict__ Ez,
155
+ const float* __restrict__ Hx, const float* __restrict__ Hy, const float* __restrict__ Hz,
156
+ float dx, float dy, float dz,
157
+ int step, int NX_FIELDS, int NY_FIELDS, int NZ_FIELDS,
158
+ int pml0, int pml1, int pml2, int pml3, int pml4, int pml5,
159
+ const float* __restrict__ x0ER, const float* __restrict__ xmER,
160
+ const float* __restrict__ y0ER, const float* __restrict__ ymER,
161
+ const float* __restrict__ z0ER, const float* __restrict__ zmER,
162
+ const float* __restrict__ updatecoeffsE,
163
+ float* __restrict__ x0EPhi1, float* __restrict__ x0EPhi2,
164
+ float* __restrict__ xmEPhi1, float* __restrict__ xmEPhi2,
165
+ float* __restrict__ y0EPhi1, float* __restrict__ y0EPhi2,
166
+ float* __restrict__ ymEPhi1, float* __restrict__ ymEPhi2,
167
+ float* __restrict__ z0EPhi1, float* __restrict__ z0EPhi2,
168
+ float* __restrict__ zmEPhi1, float* __restrict__ zmEPhi2)
169
+ {
170
+ long long idx = blockIdx.x * blockDim.x + threadIdx.x;
171
+ long long ny_nz = (long long)NY_FIELDS * NZ_FIELDS;
172
+ long long field_stride = (long long)NX_FIELDS * ny_nz;
173
+ if (idx >= field_stride) return;
174
+
175
+ long long i = idx / ny_nz;
176
+ long long rem = idx % ny_nz;
177
+ long long j = rem / NZ_FIELDS;
178
+ long long k = rem % NZ_FIELDS;
179
+
180
+ bool do_ex = (((NY_FIELDS-1) != 1 || (NZ_FIELDS-1) != 1) && i < (NX_FIELDS-1) && j > 0 && j < (NY_FIELDS-1) && k > 0 && k < (NZ_FIELDS-1));
181
+ bool do_ey = (((NX_FIELDS-1) != 1 || (NZ_FIELDS-1) != 1) && i > 0 && i < (NX_FIELDS-1) && j < (NY_FIELDS-1) && k > 0 && k < (NZ_FIELDS-1));
182
+ bool do_ez = (((NX_FIELDS-1) != 1 || (NY_FIELDS-1) != 1) && i > 0 && i < (NX_FIELDS-1) && j > 0 && j < (NY_FIELDS-1) && k < (NZ_FIELDS-1));
183
+
184
+ bool in_x0 = (pml0 > 0 && i <= pml0 && j < NY_FIELDS && k < NZ_FIELDS);
185
+ bool in_xm = (pml1 > 0 && i >= NX_FIELDS - 1 - pml1 && i < NX_FIELDS && j < NY_FIELDS && k < NZ_FIELDS);
186
+ bool in_y0 = (pml2 > 0 && i < NX_FIELDS && j <= pml2 && k < NZ_FIELDS);
187
+ bool in_ym = (pml3 > 0 && i < NX_FIELDS && j >= NY_FIELDS - 1 - pml3 && j < NY_FIELDS && k < NZ_FIELDS);
188
+ bool in_z0 = (pml4 > 0 && i < NX_FIELDS && j < NY_FIELDS && k <= pml4);
189
+ bool in_zm = (pml5 > 0 && i < NX_FIELDS && j < NY_FIELDS && k >= NZ_FIELDS - 1 - pml5 && k < NZ_FIELDS);
190
+
191
+ float ue0 = uE0[idx];
192
+ float ue1 = uE1[idx];
193
+ float upd = updatecoeffsE[idx];
194
+
195
+ long long id4 = idx;
196
+
197
+ for (int s = 0; s < step; ++s) {
198
+ if (do_ex) Ex[id4] = ue0 * Ex[id4] + ue1 * (Hz[id4] - Hz[id4 - NZ_FIELDS]) - ue1 * (Hy[id4] - Hy[id4 - 1]);
199
+ if (do_ey) Ey[id4] = ue0 * Ey[id4] + ue1 * (Hx[id4] - Hx[id4 - 1]) - ue1 * (Hz[id4] - Hz[id4 - ny_nz]);
200
+ if (do_ez) Ez[id4] = ue0 * Ez[id4] + ue1 * (Hy[id4] - Hy[id4 - ny_nz]) - ue1 * (Hx[id4] - Hx[id4 - NZ_FIELDS]);
201
+
202
+ if (in_x0) {
203
+ long long i1 = pml0 - i;
204
+ float RA01 = x0ER[i1] - 1.0f, RB0 = x0ER[pml0 + i1], RE0 = x0ER[2 * pml0 + i1], RF0 = x0ER[3 * pml0 + i1];
205
+ if (j < NY_FIELDS - 1 && i > 0) {
206
+ float dHz = (Hz[id4] - Hz[id4 - ny_nz]) / dx;
207
+ long long p_idx = ((long long)s * (pml0+1) * (NY_FIELDS-1) * NZ_FIELDS) + i1 * (NY_FIELDS-1) * NZ_FIELDS + j * NZ_FIELDS + k;
208
+ float phi = x0EPhi1[p_idx];
209
+ Ey[id4] -= upd * (RA01 * dHz + RB0 * phi);
210
+ x0EPhi1[p_idx] = RE0 * phi - RF0 * dHz;
211
+ }
212
+ if (k < NZ_FIELDS - 1 && i > 0) {
213
+ float dHy = (Hy[id4] - Hy[id4 - ny_nz]) / dx;
214
+ long long p_idx = ((long long)s * (pml0+1) * NY_FIELDS * (NZ_FIELDS-1)) + i1 * NY_FIELDS * (NZ_FIELDS-1) + j * (NZ_FIELDS-1) + k;
215
+ float phi = x0EPhi2[p_idx];
216
+ Ez[id4] += upd * (RA01 * dHy + RB0 * phi);
217
+ x0EPhi2[p_idx] = RE0 * phi - RF0 * dHy;
218
+ }
219
+ }
220
+
221
+ if (in_xm) {
222
+ long long i1 = i - (NX_FIELDS - 1 - pml1);
223
+ float RA01 = xmER[i1] - 1.0f, RB0 = xmER[pml1 + i1], RE0 = xmER[2 * pml1 + i1], RF0 = xmER[3 * pml1 + i1];
224
+ if (j < NY_FIELDS - 1 && i > 0) {
225
+ float dHz = (Hz[id4] - Hz[id4 - ny_nz]) / dx;
226
+ long long p_idx = ((long long)s * (pml1+1) * (NY_FIELDS-1) * NZ_FIELDS) + i1 * (NY_FIELDS-1) * NZ_FIELDS + j * NZ_FIELDS + k;
227
+ float phi = xmEPhi1[p_idx];
228
+ Ey[id4] -= upd * (RA01 * dHz + RB0 * phi);
229
+ xmEPhi1[p_idx] = RE0 * phi - RF0 * dHz;
230
+ }
231
+ if (k < NZ_FIELDS - 1 && i > 0) {
232
+ float dHy = (Hy[id4] - Hy[id4 - ny_nz]) / dx;
233
+ long long p_idx = ((long long)s * (pml1+1) * NY_FIELDS * (NZ_FIELDS-1)) + i1 * NY_FIELDS * (NZ_FIELDS-1) + j * (NZ_FIELDS-1) + k;
234
+ float phi = xmEPhi2[p_idx];
235
+ Ez[id4] += upd * (RA01 * dHy + RB0 * phi);
236
+ xmEPhi2[p_idx] = RE0 * phi - RF0 * dHy;
237
+ }
238
+ }
239
+
240
+ if (in_y0) {
241
+ long long j1 = pml2 - j;
242
+ float RA01 = y0ER[j1] - 1.0f, RB0 = y0ER[pml2 + j1], RE0 = y0ER[2 * pml2 + j1], RF0 = y0ER[3 * pml2 + j1];
243
+ if (i < NX_FIELDS - 1 && j > 0) {
244
+ float dHz = (Hz[id4] - Hz[id4 - NZ_FIELDS]) / dy;
245
+ long long p_idx = ((long long)s * (NX_FIELDS-1) * (pml2+1) * NZ_FIELDS) + i * (pml2+1) * NZ_FIELDS + j1 * NZ_FIELDS + k;
246
+ float phi = y0EPhi1[p_idx];
247
+ Ex[id4] += upd * (RA01 * dHz + RB0 * phi);
248
+ y0EPhi1[p_idx] = RE0 * phi - RF0 * dHz;
249
+ }
250
+ if (k < NZ_FIELDS - 1 && j > 0) {
251
+ float dHx = (Hx[id4] - Hx[id4 - NZ_FIELDS]) / dy;
252
+ long long p_idx = ((long long)s * NX_FIELDS * (pml2+1) * (NZ_FIELDS-1)) + i * (pml2+1) * (NZ_FIELDS-1) + j1 * (NZ_FIELDS-1) + k;
253
+ float phi = y0EPhi2[p_idx];
254
+ Ez[id4] -= upd * (RA01 * dHx + RB0 * phi);
255
+ y0EPhi2[p_idx] = RE0 * phi - RF0 * dHx;
256
+ }
257
+ }
258
+
259
+ if (in_ym) {
260
+ long long j1 = j - (NY_FIELDS - 1 - pml3);
261
+ float RA01 = ymER[j1] - 1.0f, RB0 = ymER[pml3 + j1], RE0 = ymER[2 * pml3 + j1], RF0 = ymER[3 * pml3 + j1];
262
+ if (i < NX_FIELDS - 1 && j > 0) {
263
+ float dHz = (Hz[id4] - Hz[id4 - NZ_FIELDS]) / dy;
264
+ long long p_idx = ((long long)s * (NX_FIELDS-1) * (pml3+1) * NZ_FIELDS) + i * (pml3+1) * NZ_FIELDS + j1 * NZ_FIELDS + k;
265
+ float phi = ymEPhi1[p_idx];
266
+ Ex[id4] += upd * (RA01 * dHz + RB0 * phi);
267
+ ymEPhi1[p_idx] = RE0 * phi - RF0 * dHz;
268
+ }
269
+ if (k < NZ_FIELDS - 1 && j > 0) {
270
+ float dHx = (Hx[id4] - Hx[id4 - NZ_FIELDS]) / dy;
271
+ long long p_idx = ((long long)s * NX_FIELDS * (pml3+1) * (NZ_FIELDS-1)) + i * (pml3+1) * (NZ_FIELDS-1) + j1 * (NZ_FIELDS-1) + k;
272
+ float phi = ymEPhi2[p_idx];
273
+ Ez[id4] -= upd * (RA01 * dHx + RB0 * phi);
274
+ ymEPhi2[p_idx] = RE0 * phi - RF0 * dHx;
275
+ }
276
+ }
277
+
278
+ if (in_z0) {
279
+ long long k1 = pml4 - k;
280
+ float RA01 = z0ER[k1] - 1.0f, RB0 = z0ER[pml4 + k1], RE0 = z0ER[2 * pml4 + k1], RF0 = z0ER[3 * pml4 + k1];
281
+ if (i < NX_FIELDS - 1 && k > 0) {
282
+ float dHy = (Hy[id4] - Hy[id4 - 1]) / dz;
283
+ long long p_idx = ((long long)s * (NX_FIELDS-1) * NY_FIELDS * (pml4+1)) + i * NY_FIELDS * (pml4+1) + j * (pml4+1) + k1;
284
+ float phi = z0EPhi1[p_idx];
285
+ Ex[id4] -= upd * (RA01 * dHy + RB0 * phi);
286
+ z0EPhi1[p_idx] = RE0 * phi - RF0 * dHy;
287
+ }
288
+ if (j < NY_FIELDS - 1 && k > 0) {
289
+ float dHx = (Hx[id4] - Hx[id4 - 1]) / dz;
290
+ long long p_idx = ((long long)s * NX_FIELDS * (NY_FIELDS-1) * (pml4+1)) + i * (NY_FIELDS-1) * (pml4+1) + j * (pml4+1) + k1;
291
+ float phi = z0EPhi2[p_idx];
292
+ Ey[id4] += upd * (RA01 * dHx + RB0 * phi);
293
+ z0EPhi2[p_idx] = RE0 * phi - RF0 * dHx;
294
+ }
295
+ }
296
+
297
+ if (in_zm) {
298
+ long long k1 = k - (NZ_FIELDS - 1 - pml5);
299
+ float RA01 = zmER[k1] - 1.0f, RB0 = zmER[pml5 + k1], RE0 = zmER[2 * pml5 + k1], RF0 = zmER[3 * pml5 + k1];
300
+ if (i < NX_FIELDS - 1 && k > 0) {
301
+ float dHy = (Hy[id4] - Hy[id4 - 1]) / dz;
302
+ long long p_idx = ((long long)s * (NX_FIELDS-1) * NY_FIELDS * (pml5+1)) + i * NY_FIELDS * (pml5+1) + j * (pml5+1) + k1;
303
+ float phi = zmEPhi1[p_idx];
304
+ Ex[id4] -= upd * (RA01 * dHy + RB0 * phi);
305
+ zmEPhi1[p_idx] = RE0 * phi - RF0 * dHy;
306
+ }
307
+ if (j < NY_FIELDS - 1 && k > 0) {
308
+ float dHx = (Hx[id4] - Hx[id4 - 1]) / dz;
309
+ long long p_idx = ((long long)s * NX_FIELDS * (NY_FIELDS-1) * (pml5+1)) + i * (NY_FIELDS-1) * (pml5+1) + j * (pml5+1) + k1;
310
+ float phi = zmEPhi2[p_idx];
311
+ Ey[id4] += upd * (RA01 * dHx + RB0 * phi);
312
+ zmEPhi2[p_idx] = RE0 * phi - RF0 * dHx;
313
+ }
314
+ }
315
+
316
+ id4 += field_stride;
317
+ }
318
+ }
319
+
320
+ // ---------------------------------------------------------
321
+ // 融合:磁场全局更新 + 全侧 PML 边界修正
322
+ // ---------------------------------------------------------
323
+ __global__ void fused_h_fields_updates_gpu(
324
+ const float* __restrict__ uH0, const float* __restrict__ uH1,
325
+ const float* __restrict__ Ex, const float* __restrict__ Ey, const float* __restrict__ Ez,
326
+ float* __restrict__ Hx, float* __restrict__ Hy, float* __restrict__ Hz,
327
+ float dx, float dy, float dz,
328
+ int step, int NX_FIELDS, int NY_FIELDS, int NZ_FIELDS,
329
+ int pml0, int pml1, int pml2, int pml3, int pml4, int pml5,
330
+ const float* __restrict__ x0HR, const float* __restrict__ xmHR,
331
+ const float* __restrict__ y0HR, const float* __restrict__ ymHR,
332
+ const float* __restrict__ z0HR, const float* __restrict__ zmHR,
333
+ const float* __restrict__ updatecoeffsH,
334
+ float* __restrict__ x0HPhi1, float* __restrict__ x0HPhi2,
335
+ float* __restrict__ xmHPhi1, float* __restrict__ xmHPhi2,
336
+ float* __restrict__ y0HPhi1, float* __restrict__ y0HPhi2,
337
+ float* __restrict__ ymHPhi1, float* __restrict__ ymHPhi2,
338
+ float* __restrict__ z0HPhi1, float* __restrict__ z0HPhi2,
339
+ float* __restrict__ zmHPhi1, float* __restrict__ zmHPhi2)
340
+ {
341
+ long long idx = blockIdx.x * blockDim.x + threadIdx.x;
342
+ long long ny_nz = (long long)NY_FIELDS * NZ_FIELDS;
343
+ long long field_stride = (long long)NX_FIELDS * ny_nz;
344
+ if (idx >= field_stride) return;
345
+
346
+ long long i = idx / ny_nz;
347
+ long long rem = idx % ny_nz;
348
+ long long j = rem / NZ_FIELDS;
349
+ long long k = rem % NZ_FIELDS;
350
+
351
+ bool do_hx = ((NX_FIELDS-1) != 1 && i > 0 && i < (NX_FIELDS-1) && j < (NY_FIELDS-1) && k < (NZ_FIELDS-1));
352
+ bool do_hy = ((NY_FIELDS-1) != 1 && i < (NX_FIELDS-1) && j > 0 && j < (NY_FIELDS-1) && k < (NZ_FIELDS-1));
353
+ bool do_hz = ((NZ_FIELDS-1) != 1 && i < (NX_FIELDS-1) && j < (NY_FIELDS-1) && k > 0 && k < (NZ_FIELDS-1));
354
+
355
+ bool in_x0 = (pml0 > 0 && i < pml0 && j < NY_FIELDS && k < NZ_FIELDS);
356
+ bool in_xm = (pml1 > 0 && i >= NX_FIELDS - 1 - pml1 && i < NX_FIELDS - 1 && j < NY_FIELDS && k < NZ_FIELDS);
357
+ bool in_y0 = (pml2 > 0 && i < NX_FIELDS && j < pml2 && k < NZ_FIELDS);
358
+ bool in_ym = (pml3 > 0 && i < NX_FIELDS && j >= NY_FIELDS - 1 - pml3 && j < NY_FIELDS - 1 && k < NZ_FIELDS);
359
+ bool in_z0 = (pml4 > 0 && i < NX_FIELDS && j < NY_FIELDS && k < pml4);
360
+ bool in_zm = (pml5 > 0 && i < NX_FIELDS && j < NY_FIELDS && k >= NZ_FIELDS - 1 - pml5 && k < NZ_FIELDS - 1);
361
+
362
+ float uh0 = uH0[idx];
363
+ float uh1 = uH1[idx];
364
+ float upd = updatecoeffsH[idx];
365
+
366
+ long long id4 = idx;
367
+
368
+ for (int s = 0; s < step; ++s) {
369
+ if (do_hx) Hx[id4] = uh0 * Hx[id4] - uh1 * (Ez[id4 + NZ_FIELDS] - Ez[id4]) + uh1 * (Ey[id4 + 1] - Ey[id4]);
370
+ if (do_hy) Hy[id4] = uh0 * Hy[id4] - uh1 * (Ex[id4 + 1] - Ex[id4]) + uh1 * (Ez[id4 + ny_nz] - Ez[id4]);
371
+ if (do_hz) Hz[id4] = uh0 * Hz[id4] - uh1 * (Ey[id4 + ny_nz] - Ey[id4]) + uh1 * (Ex[id4 + NZ_FIELDS] - Ex[id4]);
372
+
373
+ if (in_x0) {
374
+ long long i1 = pml0 - 1 - i;
375
+ float RA01 = x0HR[i1] - 1.0f, RB0 = x0HR[pml0 + i1], RE0 = x0HR[2 * pml0 + i1], RF0 = x0HR[3 * pml0 + i1];
376
+ if (k < NZ_FIELDS - 1) {
377
+ float dEz = (Ez[id4 + ny_nz] - Ez[id4]) / dx;
378
+ long long p_idx = ((long long)s * pml0 * NY_FIELDS * (NZ_FIELDS-1)) + i1 * NY_FIELDS * (NZ_FIELDS-1) + j * (NZ_FIELDS-1) + k;
379
+ float phi = x0HPhi1[p_idx];
380
+ Hy[id4] += upd * (RA01 * dEz + RB0 * phi);
381
+ x0HPhi1[p_idx] = RE0 * phi - RF0 * dEz;
382
+ }
383
+ if (j < NY_FIELDS - 1) {
384
+ float dEy = (Ey[id4 + ny_nz] - Ey[id4]) / dx;
385
+ long long p_idx = ((long long)s * pml0 * (NY_FIELDS-1) * NZ_FIELDS) + i1 * (NY_FIELDS-1) * NZ_FIELDS + j * NZ_FIELDS + k;
386
+ float phi = x0HPhi2[p_idx];
387
+ Hz[id4] -= upd * (RA01 * dEy + RB0 * phi);
388
+ x0HPhi2[p_idx] = RE0 * phi - RF0 * dEy;
389
+ }
390
+ }
391
+
392
+ if (in_xm) {
393
+ long long i1 = i - (NX_FIELDS - 1 - pml1);
394
+ float RA01 = xmHR[i1] - 1.0f, RB0 = xmHR[pml1 + i1], RE0 = xmHR[2 * pml1 + i1], RF0 = xmHR[3 * pml1 + i1];
395
+ if (k < NZ_FIELDS - 1) {
396
+ float dEz = (Ez[id4 + ny_nz] - Ez[id4]) / dx;
397
+ long long p_idx = ((long long)s * pml1 * NY_FIELDS * (NZ_FIELDS-1)) + i1 * NY_FIELDS * (NZ_FIELDS-1) + j * (NZ_FIELDS-1) + k;
398
+ float phi = xmHPhi1[p_idx];
399
+ Hy[id4] += upd * (RA01 * dEz + RB0 * phi);
400
+ xmHPhi1[p_idx] = RE0 * phi - RF0 * dEz;
401
+ }
402
+ if (j < NY_FIELDS - 1) {
403
+ float dEy = (Ey[id4 + ny_nz] - Ey[id4]) / dx;
404
+ long long p_idx = ((long long)s * pml1 * (NY_FIELDS-1) * NZ_FIELDS) + i1 * (NY_FIELDS-1) * NZ_FIELDS + j * NZ_FIELDS + k;
405
+ float phi = xmHPhi2[p_idx];
406
+ Hz[id4] -= upd * (RA01 * dEy + RB0 * phi);
407
+ xmHPhi2[p_idx] = RE0 * phi - RF0 * dEy;
408
+ }
409
+ }
410
+
411
+ if (in_y0) {
412
+ long long j1 = pml2 - 1 - j;
413
+ float RA01 = y0HR[j1] - 1.0f, RB0 = y0HR[pml2 + j1], RE0 = y0HR[2 * pml2 + j1], RF0 = y0HR[3 * pml2 + j1];
414
+ if (i < NX_FIELDS && k < NZ_FIELDS - 1) {
415
+ float dEz = (Ez[id4 + NZ_FIELDS] - Ez[id4]) / dy;
416
+ long long p_idx = ((long long)s * NX_FIELDS * pml2 * (NZ_FIELDS-1)) + i * pml2 * (NZ_FIELDS-1) + j1 * (NZ_FIELDS-1) + k;
417
+ float phi = y0HPhi1[p_idx];
418
+ Hx[id4] -= upd * (RA01 * dEz + RB0 * phi);
419
+ y0HPhi1[p_idx] = RE0 * phi - RF0 * dEz;
420
+ }
421
+ if (i < NX_FIELDS - 1 && k < NZ_FIELDS) {
422
+ float dEx = (Ex[id4 + NZ_FIELDS] - Ex[id4]) / dy;
423
+ long long p_idx = ((long long)s * (NX_FIELDS-1) * pml2 * NZ_FIELDS) + i * pml2 * NZ_FIELDS + j1 * NZ_FIELDS + k;
424
+ float phi = y0HPhi2[p_idx];
425
+ Hz[id4] += upd * (RA01 * dEx + RB0 * phi);
426
+ y0HPhi2[p_idx] = RE0 * phi - RF0 * dEx;
427
+ }
428
+ }
429
+
430
+ if (in_ym) {
431
+ long long j1 = j - (NY_FIELDS - 1 - pml3);
432
+ float RA01 = ymHR[j1] - 1.0f, RB0 = ymHR[pml3 + j1], RE0 = ymHR[2 * pml3 + j1], RF0 = ymHR[3 * pml3 + j1];
433
+ if (i < NX_FIELDS && k < NZ_FIELDS - 1) {
434
+ float dEz = (Ez[id4 + NZ_FIELDS] - Ez[id4]) / dy;
435
+ long long p_idx = ((long long)s * NX_FIELDS * pml3 * (NZ_FIELDS-1)) + i * pml3 * (NZ_FIELDS-1) + j1 * (NZ_FIELDS-1) + k;
436
+ float phi = ymHPhi1[p_idx];
437
+ Hx[id4] -= upd * (RA01 * dEz + RB0 * phi);
438
+ ymHPhi1[p_idx] = RE0 * phi - RF0 * dEz;
439
+ }
440
+ if (i < NX_FIELDS - 1 && k < NZ_FIELDS) {
441
+ float dEx = (Ex[id4 + NZ_FIELDS] - Ex[id4]) / dy;
442
+ long long p_idx = ((long long)s * (NX_FIELDS-1) * pml3 * NZ_FIELDS) + i * pml3 * NZ_FIELDS + j1 * NZ_FIELDS + k;
443
+ float phi = ymHPhi2[p_idx];
444
+ Hz[id4] += upd * (RA01 * dEx + RB0 * phi);
445
+ ymHPhi2[p_idx] = RE0 * phi - RF0 * dEx;
446
+ }
447
+ }
448
+
449
+ if (in_z0) {
450
+ long long k1 = pml4 - 1 - k;
451
+ float RA01 = z0HR[k1] - 1.0f, RB0 = z0HR[pml4 + k1], RE0 = z0HR[2 * pml4 + k1], RF0 = z0HR[3 * pml4 + k1];
452
+ if (i < NX_FIELDS && j < NY_FIELDS - 1) {
453
+ float dEy = (Ey[id4 + 1] - Ey[id4]) / dz;
454
+ long long p_idx = ((long long)s * NX_FIELDS * (NY_FIELDS-1) * pml4) + i * (NY_FIELDS-1) * pml4 + j * pml4 + k1;
455
+ float phi = z0HPhi1[p_idx];
456
+ Hx[id4] += upd * (RA01 * dEy + RB0 * phi);
457
+ z0HPhi1[p_idx] = RE0 * phi - RF0 * dEy;
458
+ }
459
+ if (i < NX_FIELDS - 1 && j < NY_FIELDS) {
460
+ float dEx = (Ex[id4 + 1] - Ex[id4]) / dz;
461
+ long long p_idx = ((long long)s * (NX_FIELDS-1) * NY_FIELDS * pml4) + i * NY_FIELDS * pml4 + j * pml4 + k1;
462
+ float phi = z0HPhi2[p_idx];
463
+ Hy[id4] -= upd * (RA01 * dEx + RB0 * phi);
464
+ z0HPhi2[p_idx] = RE0 * phi - RF0 * dEx;
465
+ }
466
+ }
467
+
468
+ if (in_zm) {
469
+ long long k1 = k - (NZ_FIELDS - 1 - pml5);
470
+ float RA01 = zmHR[k1] - 1.0f, RB0 = zmHR[pml5 + k1], RE0 = zmHR[2 * pml5 + k1], RF0 = zmHR[3 * pml5 + k1];
471
+ if (i < NX_FIELDS && j < NY_FIELDS - 1) {
472
+ float dEy = (Ey[id4 + 1] - Ey[id4]) / dz;
473
+ long long p_idx = ((long long)s * NX_FIELDS * (NY_FIELDS-1) * pml5) + i * (NY_FIELDS-1) * pml5 + j * pml5 + k1;
474
+ float phi = zmHPhi1[p_idx];
475
+ Hx[id4] += upd * (RA01 * dEy + RB0 * phi);
476
+ zmHPhi1[p_idx] = RE0 * phi - RF0 * dEy;
477
+ }
478
+ if (i < NX_FIELDS - 1 && j < NY_FIELDS) {
479
+ float dEx = (Ex[id4 + 1] - Ex[id4]) / dz;
480
+ long long p_idx = ((long long)s * (NX_FIELDS-1) * NY_FIELDS * pml5) + i * NY_FIELDS * pml5 + j * pml5 + k1;
481
+ float phi = zmHPhi2[p_idx];
482
+ Hy[id4] -= upd * (RA01 * dEx + RB0 * phi);
483
+ zmHPhi2[p_idx] = RE0 * phi - RF0 * dEx;
484
+ }
485
+ }
486
+
487
+ id4 += field_stride;
488
+ }
489
+ }
490
+
491
+ // ---------------------------------------------------------
492
+ // 反传:波场倒播
493
+ // ---------------------------------------------------------
494
+ __global__ void Back_source(
495
+ int step, int iteration, float dx,
496
+ const int* __restrict__ sourcelocation, const float* __restrict__ srcwaveforms,
497
+ float* Ex, float* Ey, float* Ez, float* uE4,
498
+ int NX, int NY, int NZ, int nsr, int polarisation, int iterations
499
+ ){
500
+ long long src = blockIdx.x * blockDim.x + threadIdx.x;
501
+ if (src >= nsr) return;
502
+ long long field_stride = (long long)NX * NY * NZ;
503
+ long long index_stride = (long long)iterations * nsr;
504
+ long long index = (long long)iteration * nsr + src;
505
+
506
+ for (int s = 0; s < step; ++s) {
507
+ long long i = sourcelocation[s * nsr * 3 + src * 3 + 0];
508
+ long long j = sourcelocation[s * nsr * 3 + src * 3 + 1];
509
+ long long k = sourcelocation[s * nsr * 3 + src * 3 + 2];
510
+
511
+ float waveform_value = srcwaveforms[index];
512
+ long long id4 = s * field_stride + i * NY * NZ + j * NZ + k;
513
+
514
+ if (polarisation == 0) Ex[id4] -= waveform_value;
515
+ else if (polarisation == 1) Ey[id4] -= waveform_value;
516
+ else if (polarisation == 2) Ez[id4] -= waveform_value;
517
+
518
+ index += index_stride;
519
+ }
520
+ }
521
+
522
+ // ---------------------------------------------------------
523
+ // 提取快照到缓冲区或全局内存
524
+ // ---------------------------------------------------------
525
+ __global__ void copy_to_Eall_single(
526
+ float* __restrict__ dst_ptr, int t_idx, const float* __restrict__ E,
527
+ int step, int NX, int NY, int NZ)
528
+ {
529
+ long long idx = blockIdx.x * blockDim.x + threadIdx.x;
530
+ long long nx1 = NX - 1, ny1 = NY - 1, nz1 = NZ - 1;
531
+ long long total = nx1 * ny1 * nz1;
532
+ if (idx >= total) return;
533
+
534
+ long long i = idx / (ny1 * nz1);
535
+ long long rem = idx % (ny1 * nz1);
536
+ long long j = rem / nz1;
537
+ long long k = rem % nz1;
538
+
539
+ long long src_idx = i * NY * NZ + j * NZ + k;
540
+ long long dst_idx = (long long)t_idx * step * total + idx;
541
+ long long field_stride = (long long)NX * NY * NZ;
542
+
543
+ for (int s = 0; s < step; ++s) {
544
+ dst_ptr[dst_idx] = E[src_idx];
545
+ src_idx += field_stride;
546
+ dst_idx += total;
547
+ }
548
+ }
549
+
550
+ // ---------------------------------------------------------
551
+ // 融合:伴随状态法梯度更新(支持降采样波场及异步滑窗/同步显存自适应)
552
+ // ---------------------------------------------------------
553
+ __global__ void accumulate_gradients(
554
+ const float* __restrict__ Ez, const float* __restrict__ Eall_ptr, const float* __restrict__ d_E_buf,
555
+ float* __restrict__ grader, float* __restrict__ gradse,
556
+ int i, int step, int NX, int NY, int NZ, float dt,int errequiregrad,int serequiregrad,
557
+ int S, int nt_saved, int use_async_offload
558
+ ) {
559
+ long long idx = blockIdx.x * blockDim.x + threadIdx.x;
560
+ long long sx = (NX - 1), sy = (NY - 1), sz = (NZ - 1);
561
+ long long total_cells = sx * sy * sz;
562
+
563
+ if (idx >= total_cells) return;
564
+
565
+ long long ix = idx / (sy * sz);
566
+ long long rem = idx % (sy * sz);
567
+ long long iy = rem / sz;
568
+ long long iz = rem % sz;
569
+
570
+ long long idx_Ez = ix * NY * NZ + iy * NZ + iz;
571
+
572
+ // 逻辑时刻索引
573
+ long long idx0_curr = i / S;
574
+ long long idx1_curr = min(idx0_curr + 1, (long long)nt_saved - 1);
575
+ float w1_curr = (float)(i % S) / S;
576
+ float w0_curr = 1.0f - w1_curr;
577
+
578
+ long long idx0_prev = (i - 1) / S;
579
+ long long idx1_prev = min(idx0_prev + 1, (long long)nt_saved - 1);
580
+ float w1_prev = (float)((i - 1) % S) / S;
581
+ float w0_prev = 1.0f - w1_prev;
582
+
583
+ long long ez_stride = (long long)NX * NY * NZ;
584
+ float local_grader = 0.0f;
585
+ float local_gradse = 0.0f;
586
+
587
+ for (int s = 0; s < step; ++s) {
588
+ long long base_idx = (long long)s * total_cells + idx;
589
+ float e0_c, e1_c, e0_p, e1_p;
590
+
591
+ if (use_async_offload) {
592
+ e0_c = d_E_buf[(idx0_curr % 3) * step * total_cells + base_idx];
593
+ e1_c = d_E_buf[(idx1_curr % 3) * step * total_cells + base_idx];
594
+ e0_p = d_E_buf[(idx0_prev % 3) * step * total_cells + base_idx];
595
+ e1_p = d_E_buf[(idx1_prev % 3) * step * total_cells + base_idx];
596
+ } else {
597
+ e0_c = Eall_ptr[idx0_curr * step * total_cells + base_idx];
598
+ e1_c = Eall_ptr[idx1_curr * step * total_cells + base_idx];
599
+ e0_p = Eall_ptr[idx0_prev * step * total_cells + base_idx];
600
+ e1_p = Eall_ptr[idx1_prev * step * total_cells + base_idx];
601
+ }
602
+
603
+ float e_curr = e0_c * w0_curr + e1_c * w1_curr;
604
+ float e_prev = e0_p * w0_prev + e1_p * w1_prev;
605
+
606
+ float ez_val = Ez[idx_Ez];
607
+
608
+ if (errequiregrad == 1) local_grader += (e_curr - e_prev) * ez_val / dt;
609
+ if (serequiregrad == 1) local_gradse += e_curr * ez_val * dt;
610
+
611
+ idx_Ez += ez_stride;
612
+ }
613
+
614
+ if (errequiregrad == 1) atomicAdd(&grader[idx], local_grader);
615
+ if (serequiregrad == 1) atomicAdd(&gradse[idx], local_gradse);
616
+ }
617
+
618
+ // ---------------------------------------------------------
619
+ // 主机 API
620
+ // ---------------------------------------------------------
621
+ extern "C" {
622
+
623
+ void forward(const float* __restrict__ er, const float* __restrict__ se, const float* __restrict__ mr,
624
+ float* __restrict__ Eall_ptr,
625
+ float* __restrict__ Ex, float* __restrict__ Ey, float* __restrict__ Ez,
626
+ float* __restrict__ Hx, float* __restrict__ Hy, float* __restrict__ Hz,
627
+ float* __restrict__ uE0, float* __restrict__ uE1, float* __restrict__ uE4,
628
+ float* __restrict__ uH0, float* __restrict__ uH1, float* __restrict__ uH4,
629
+
630
+ float* __restrict__ x0EPhi1,float* __restrict__ x0EPhi2, float* __restrict__ x0HPhi1,float* __restrict__ x0HPhi2,
631
+ float* __restrict__ xmEPhi1,float* __restrict__ xmEPhi2, float* __restrict__ xmHPhi1,float* __restrict__ xmHPhi2,
632
+ float* __restrict__ y0EPhi1,float* __restrict__ y0EPhi2, float* __restrict__ y0HPhi1,float* __restrict__ y0HPhi2,
633
+ float* __restrict__ ymEPhi1,float* __restrict__ ymEPhi2, float* __restrict__ ymHPhi1,float* __restrict__ ymHPhi2,
634
+ float* __restrict__ z0EPhi1,float* __restrict__ z0EPhi2, float* __restrict__ z0HPhi1,float* __restrict__ z0HPhi2,
635
+ float* __restrict__ zmEPhi1,float* __restrict__ zmEPhi2, float* __restrict__ zmHPhi1,float* __restrict__ zmHPhi2,
636
+
637
+ int pml0,int pml1,int pml2,int pml3,int pml4,int pml5,
638
+
639
+ const float* __restrict__ x0ER,const float* __restrict__ xmER, const float* __restrict__ y0ER,const float* __restrict__ ymER,
640
+ const float* __restrict__ z0ER,const float* __restrict__ zmER, const float* __restrict__ x0HR,const float* __restrict__ xmHR,
641
+ const float* __restrict__ y0HR,const float* __restrict__ ymHR, const float* __restrict__ z0HR,const float* __restrict__ zmHR,
642
+
643
+ float dt, int nt, int step, int nrx, float dx,
644
+ const int* __restrict__ receiverlocation, float* __restrict__ rxs,
645
+
646
+ int NX_FIELDS, int NY_FIELDS, int NZ_FIELDS, int nsrc,
647
+ const int* __restrict__ sourcelocation, const float* __restrict__ srcwaveforms, int polarisation,
648
+ int sampling_interval)
649
+ {
650
+ cudaPointerAttributes attr;
651
+ cudaError_t err = cudaPointerGetAttributes(&attr, Eall_ptr);
652
+ int use_async = 0;
653
+ if (err == cudaSuccess && attr.type == cudaMemoryTypeDevice) {
654
+ use_async = 0;
655
+ } else {
656
+ cudaGetLastError();
657
+ use_async = 1;
658
+ }
659
+
660
+ float* d_E_buf = nullptr;
661
+ long long snap_size = (long long)step * (NX_FIELDS - 1) * (NY_FIELDS - 1) * (NZ_FIELDS - 1);
662
+
663
+ cudaStream_t stream_comp = 0, stream_trans = 0;
664
+ cudaEvent_t event_comp;
665
+ if (use_async) {
666
+ cudaStreamCreate(&stream_comp);
667
+ cudaStreamCreate(&stream_trans);
668
+ cudaEventCreate(&event_comp);
669
+ cudaMalloc(&d_E_buf, 2 * snap_size * sizeof(float));
670
+ }
671
+
672
+ long long blockSize = 256;
673
+ long long total_fields = (long long)NX_FIELDS * NY_FIELDS * NZ_FIELDS;
674
+ dim3 grid_fields(CEIL_DIV(total_fields, blockSize));
675
+
676
+ ucgetforward<<<grid_fields, blockSize, 0, stream_comp>>>(er, se, mr, uE0, uE1, uE4, uH0, uH1, uH4, NX_FIELDS, NY_FIELDS, NZ_FIELDS, dt, dx);
677
+
678
+ dim3 grid_rx(CEIL_DIV(nrx, blockSize));
679
+ dim3 grid_src(CEIL_DIV(nsrc, blockSize));
680
+ long long total_copy = (long long)(NX_FIELDS - 1) * (NY_FIELDS - 1) * (NZ_FIELDS - 1);
681
+ dim3 grid_copy(CEIL_DIV(total_copy, blockSize));
682
+
683
+ for (int i = 0; i < nt; i++) {
684
+ long long rx_total = step * nrx;
685
+ long long gridSize_rx = (rx_total + blockSize - 1) / blockSize;
686
+ store_outputs<<<gridSize_rx, blockSize, 0, stream_comp>>>(step, nrx, i, receiverlocation, rxs, Ex, Ey, Ez, Hx, Hy, Hz, NX_FIELDS, NY_FIELDS, NZ_FIELDS, nt);
687
+
688
+ fused_h_fields_updates_gpu<<<grid_fields, blockSize, 0, stream_comp>>>(
689
+ uH0, uH1, Ex, Ey, Ez, Hx, Hy, Hz, dx, dx, dx, step, NX_FIELDS, NY_FIELDS, NZ_FIELDS,
690
+ pml0, pml1, pml2, pml3, pml4, pml5, x0HR, xmHR, y0HR, ymHR, z0HR, zmHR, uH4,
691
+ x0HPhi1, x0HPhi2, xmHPhi1, xmHPhi2, y0HPhi1, y0HPhi2, ymHPhi1, ymHPhi2, z0HPhi1, z0HPhi2, zmHPhi1, zmHPhi2);
692
+
693
+ fused_e_fields_updates_gpu<<<grid_fields, blockSize, 0, stream_comp>>>(
694
+ uE0, uE1, Ex, Ey, Ez, Hx, Hy, Hz, dx, dx, dx, step, NX_FIELDS, NY_FIELDS, NZ_FIELDS,
695
+ pml0, pml1, pml2, pml3, pml4, pml5, x0ER, xmER, y0ER, ymER, z0ER, zmER, uE4,
696
+ x0EPhi1, x0EPhi2, xmEPhi1, xmEPhi2, y0EPhi1, y0EPhi2, ymEPhi1, ymEPhi2, z0EPhi1, z0EPhi2, zmEPhi1, zmEPhi2);
697
+
698
+ Update_hertzian_dipole<<<grid_src, blockSize, 0, stream_comp>>>(step, i, dx, sourcelocation, srcwaveforms, Ex, Ey, Ez, uE4, NX_FIELDS, NY_FIELDS, NZ_FIELDS, nsrc, polarisation, nt);
699
+
700
+ if (i % sampling_interval == 0) {
701
+ int t_saved = i / sampling_interval;
702
+ if (use_async) {
703
+ int buf_idx = t_saved % 2;
704
+ cudaStreamSynchronize(stream_trans);
705
+ copy_to_Eall_single<<<grid_copy, blockSize, 0, stream_comp>>>(d_E_buf, buf_idx, Ez, step, NX_FIELDS, NY_FIELDS, NZ_FIELDS);
706
+ cudaEventRecord(event_comp, stream_comp);
707
+ cudaStreamWaitEvent(stream_trans, event_comp, 0);
708
+ cudaMemcpyAsync(Eall_ptr + t_saved * snap_size, d_E_buf + buf_idx * snap_size, snap_size * sizeof(float), cudaMemcpyDeviceToHost, stream_trans);
709
+ } else {
710
+ copy_to_Eall_single<<<grid_copy, blockSize, 0, stream_comp>>>(Eall_ptr, t_saved, Ez, step, NX_FIELDS, NY_FIELDS, NZ_FIELDS);
711
+ }
712
+ }
713
+ }
714
+
715
+ if (use_async) {
716
+ // 【核心修复】:必须先等流里的所有任务彻底做完,才能把底下的显存 free 掉!
717
+ cudaStreamSynchronize(stream_comp);
718
+ cudaStreamSynchronize(stream_trans);
719
+ cudaFree(d_E_buf);
720
+ cudaEventDestroy(event_comp);
721
+ cudaStreamDestroy(stream_comp);
722
+ cudaStreamDestroy(stream_trans);
723
+ }
724
+ }
725
+
726
+ void backward(const float* __restrict__ er, const float* __restrict__ se, const float* __restrict__ mr,
727
+ const float* __restrict__ Eall_ptr,
728
+ float* __restrict__ Ex, float* __restrict__ Ey, float* __restrict__ Ez,
729
+ float* __restrict__ Hx, float* __restrict__ Hy, float* __restrict__ Hz,
730
+ float* __restrict__ uE0, float* __restrict__ uE1, float* __restrict__ uE4,
731
+ float* __restrict__ uH0, float* __restrict__ uH1, float* __restrict__ uH4,
732
+
733
+ float* __restrict__ x0EPhi1,float* __restrict__ x0EPhi2, float* __restrict__ x0HPhi1,float* __restrict__ x0HPhi2,
734
+ float* __restrict__ xmEPhi1,float* __restrict__ xmEPhi2, float* __restrict__ xmHPhi1,float* __restrict__ xmHPhi2,
735
+ float* __restrict__ y0EPhi1,float* __restrict__ y0EPhi2, float* __restrict__ y0HPhi1,float* __restrict__ y0HPhi2,
736
+ float* __restrict__ ymEPhi1,float* __restrict__ ymEPhi2, float* __restrict__ ymHPhi1,float* __restrict__ ymHPhi2,
737
+ float* __restrict__ z0EPhi1,float* __restrict__ z0EPhi2, float* __restrict__ z0HPhi1,float* __restrict__ z0HPhi2,
738
+ float* __restrict__ zmEPhi1,float* __restrict__ zmEPhi2, float* __restrict__ zmHPhi1,float* __restrict__ zmHPhi2,
739
+
740
+ int pml0,int pml1,int pml2,int pml3,int pml4,int pml5,
741
+
742
+ float* __restrict__ x0ER,float* __restrict__ xmER, float* __restrict__ y0ER,float* __restrict__ ymER,
743
+ float* __restrict__ z0ER,float* __restrict__ zmER, float* __restrict__ x0HR,float* __restrict__ xmHR,
744
+ float* __restrict__ y0HR,float* __restrict__ ymHR, float* __restrict__ z0HR,float* __restrict__ zmHR,
745
+
746
+ float dt, int nt, int step, int nrx, float dx,
747
+ int NX_FIELDS, int NY_FIELDS, int NZ_FIELDS,
748
+ int nsrc, const int* __restrict__ sourcelocation, const float* __restrict__ srcwaveforms,
749
+ int polarisation,
750
+ float*__restrict__ grad_er,float*__restrict__ grad_se, int errequiregrad, int serequiregrad,
751
+ int sampling_interval)
752
+ {
753
+ cudaPointerAttributes attr;
754
+ cudaError_t err = cudaPointerGetAttributes(&attr, Eall_ptr);
755
+ int use_async = 0;
756
+ if (err == cudaSuccess && attr.type == cudaMemoryTypeDevice) {
757
+ use_async = 0;
758
+ } else {
759
+ cudaGetLastError();
760
+ use_async = 1;
761
+ }
762
+
763
+ float* d_E_buf = nullptr;
764
+ long long snap_size = (long long)step * (NX_FIELDS - 1) * (NY_FIELDS - 1) * (NZ_FIELDS - 1);
765
+
766
+ cudaStream_t stream_comp = 0, stream_trans = 0;
767
+ cudaEvent_t event_trans;
768
+ if (use_async) {
769
+ cudaStreamCreate(&stream_comp);
770
+ cudaStreamCreate(&stream_trans);
771
+ cudaEventCreate(&event_trans);
772
+ cudaMalloc(&d_E_buf, 3 * snap_size * sizeof(float));
773
+ }
774
+
775
+ long long blockSize = 256;
776
+ long long total_fields = (long long)NX_FIELDS * NY_FIELDS * NZ_FIELDS;
777
+ dim3 grid_fields(CEIL_DIV(total_fields, blockSize));
778
+
779
+ ucgetbackward<<<grid_fields, blockSize, 0, stream_comp>>>(er, se, mr, uE0, uE1, uE4, uH0, uH1, uH4, NX_FIELDS, NY_FIELDS, NZ_FIELDS, dt, dx);
780
+
781
+ long long total_src = step * nsrc;
782
+ long long src_blocks = (total_src + blockSize - 1) / blockSize;
783
+ dim3 grid_src(src_blocks);
784
+
785
+ long long total_grad = (long long)(NX_FIELDS-1) * (NY_FIELDS-1) * (NZ_FIELDS-1);
786
+ dim3 grid_grad(CEIL_DIV(total_grad, blockSize));
787
+
788
+ int nt_saved = (nt + sampling_interval - 1) / sampling_interval;
789
+
790
+ int max_t_needed = (nt - 1) / sampling_interval;
791
+ max_t_needed = min(max_t_needed + 1, nt_saved - 1);
792
+ int lowest_t_loaded = max_t_needed - 2;
793
+
794
+ if (use_async) {
795
+ for(int k = 0; k < 3; k++) {
796
+ int t_load = max_t_needed - k;
797
+ if(t_load >= 0) {
798
+ cudaMemcpyAsync(d_E_buf + (t_load % 3) * snap_size, Eall_ptr + t_load * snap_size, snap_size * sizeof(float), cudaMemcpyHostToDevice, stream_trans);
799
+ }
800
+ }
801
+ cudaStreamSynchronize(stream_trans);
802
+ }
803
+
804
+ for (int i = nt-1; i > 0; i--) {
805
+ if (use_async) {
806
+ int needed_t_min = (i - 1) / sampling_interval;
807
+ if (needed_t_min < lowest_t_loaded && needed_t_min >= 0) {
808
+ cudaStreamSynchronize(stream_trans);
809
+ cudaMemcpyAsync(d_E_buf + (needed_t_min % 3) * snap_size, Eall_ptr + needed_t_min * snap_size, snap_size * sizeof(float), cudaMemcpyHostToDevice, stream_trans);
810
+ lowest_t_loaded = needed_t_min;
811
+ }
812
+ cudaEventRecord(event_trans, stream_trans);
813
+ cudaStreamWaitEvent(stream_comp, event_trans, 0);
814
+ }
815
+
816
+ Back_source<<<grid_src, blockSize, 0, stream_comp>>>(step, i, dx, sourcelocation, srcwaveforms, Ex, Ey, Ez, uE4, NX_FIELDS, NY_FIELDS, NZ_FIELDS, nsrc, polarisation, nt);
817
+
818
+ fused_e_fields_updates_gpu<<<grid_fields, blockSize, 0, stream_comp>>>(
819
+ uE0, uE1, Ex, Ey, Ez, Hx, Hy, Hz, dx, dx, dx, step, NX_FIELDS, NY_FIELDS, NZ_FIELDS,
820
+ pml0, pml1, pml2, pml3, pml4, pml5, x0ER, xmER, y0ER, ymER, z0ER, zmER, uE4,
821
+ x0EPhi1, x0EPhi2, xmEPhi1, xmEPhi2, y0EPhi1, y0EPhi2, ymEPhi1, ymEPhi2, z0EPhi1, z0EPhi2, zmEPhi1, zmEPhi2);
822
+
823
+ fused_h_fields_updates_gpu<<<grid_fields, blockSize, 0, stream_comp>>>(
824
+ uH0, uH1, Ex, Ey, Ez, Hx, Hy, Hz, dx, dx, dx, step, NX_FIELDS, NY_FIELDS, NZ_FIELDS,
825
+ pml0, pml1, pml2, pml3, pml4, pml5, x0HR, xmHR, y0HR, ymHR, z0HR, zmHR, uH4,
826
+ x0HPhi1, x0HPhi2, xmHPhi1, xmHPhi2, y0HPhi1, y0HPhi2, ymHPhi1, ymHPhi2, z0HPhi1, z0HPhi2, zmHPhi1, zmHPhi2);
827
+
828
+ accumulate_gradients<<<grid_grad, blockSize, 0, stream_comp>>>(Ez, Eall_ptr, d_E_buf, grad_er, grad_se, i, step, NX_FIELDS, NY_FIELDS, NZ_FIELDS, dt, errequiregrad, serequiregrad, sampling_interval, nt_saved, use_async);
829
+ }
830
+
831
+ if (use_async) {
832
+ // 【核心修复】:必须先同步!
833
+ cudaStreamSynchronize(stream_comp);
834
+ cudaStreamSynchronize(stream_trans);
835
+ cudaFree(d_E_buf);
836
+ cudaEventDestroy(event_trans);
837
+ cudaStreamDestroy(stream_comp);
838
+ cudaStreamDestroy(stream_trans);
839
+ }
840
+ }
841
+
842
+ }
Binary file
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: DeepGPR
3
- Version: 0.0.1
3
+ Version: 0.0.2
4
4
  Summary: PyTorch and CUDA for GPR FWI
5
5
  Author-email: Lei Liu <liulei990222@gmail.com>
6
6
  Classifier: Programming Language :: Python :: 3
@@ -8,4 +8,6 @@ DeepGPR/visual.py
8
8
  DeepGPR.egg-info/PKG-INFO
9
9
  DeepGPR.egg-info/SOURCES.txt
10
10
  DeepGPR.egg-info/dependency_links.txt
11
- DeepGPR.egg-info/top_level.txt
11
+ DeepGPR.egg-info/top_level.txt
12
+ DeepGPR/lib/deepgpr.cu
13
+ DeepGPR/lib/deepgpr.so
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: DeepGPR
3
- Version: 0.0.1
3
+ Version: 0.0.2
4
4
  Summary: PyTorch and CUDA for GPR FWI
5
5
  Author-email: Lei Liu <liulei990222@gmail.com>
6
6
  Classifier: Programming Language :: Python :: 3
@@ -4,7 +4,7 @@ build-backend = "setuptools.build_meta"
4
4
 
5
5
  [project]
6
6
  name = "DeepGPR"
7
- version = "0.0.1"
7
+ version = "0.0.2"
8
8
  authors = [
9
9
  { name="Lei Liu", email="liulei990222@gmail.com" },
10
10
  ]
@@ -17,4 +17,7 @@ classifiers = [
17
17
  ]
18
18
 
19
19
  [tool.setuptools]
20
- packages = ["DeepGPR"]
20
+ packages = ["DeepGPR"]
21
+
22
+ [tool.setuptools.package-data]
23
+ "DeepGPR" = ["lib/*.cu", "lib/*.so", "lib/*.dll"]
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes