DeepGPR 0.0.1__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.
DeepGPR/compute2.py ADDED
@@ -0,0 +1,326 @@
1
+ import torch
2
+ import ctypes
3
+ from . import c_lib
4
+ from .common import initialization,build_pml_phi,create_or_separate,buildpmlcoeffs,check_tensors_for_nan_inf
5
+
6
+ def compute(device, dx=None, dt=None,
7
+ source_amplitudes=None,
8
+ source_location=None,
9
+ receiver_location=None,
10
+ er=None, se=None,mr=None,
11
+ E=None,H=None,
12
+ PML=None,
13
+ pmlthick=10, source_direction=2, reciever_direction=2,
14
+ model_gradient_sampling_interval=1,
15
+ use_async_offload=False):
16
+
17
+ er,se,nx,ny,nz,nt,nstep,nsr,nrx,ere,see,mr,mode,dtype,pmlthick,source_amplitudes=initialization(device,er,se,mr,source_amplitudes,source_location,receiver_location,dx,dt,pmlthick)
18
+
19
+ Ex,Ey,Ez=create_or_separate(E,nx,ny,nz,nstep,device,dtype)
20
+ Hx,Hy,Hz=create_or_separate(H,nx,ny,nz,nstep,device,dtype)
21
+
22
+ x0,xm,y0,ym,z0,zm,x01,x02,xm1,xm2,y01,y02,ym1,ym2,z01,z02,zm1,zm2=buildpmlcoeffs(er,mr,dt,dx,nx,ny,nz,pmlthick,device,dtype)
23
+
24
+ x0EPhi1,x0EPhi2,x0HPhi1,x0HPhi2,xmEPhi1,xmEPhi2,xmHPhi1,xmHPhi2,y0EPhi1,y0EPhi2,y0HPhi1,y0HPhi2,ymEPhi1,ymEPhi2,ymHPhi1,ymHPhi2,z0EPhi1,z0EPhi2,z0HPhi1,z0HPhi2,zmEPhi1,zmEPhi2,zmHPhi1,zmHPhi2=build_pml_phi(x0,xm,y0,ym,z0,zm,nstep,PML,device)
25
+
26
+ Ex,Ey,Ez,Hx,Hy,Hz,x0EPhi1,x0EPhi2,x0HPhi1,x0HPhi2,xmEPhi1,xmEPhi2,xmHPhi1,xmHPhi2,y0EPhi1,y0EPhi2,y0HPhi1,y0HPhi2,ymEPhi1,ymEPhi2,ymHPhi1,ymHPhi2,z0EPhi1,z0EPhi2,z0HPhi1,z0HPhi2,zmEPhi1,zmEPhi2,zmHPhi1,zmHPhi2,Eall,receiver_amplitudes = DeepGPR.apply(
27
+ er, se,Ex,Ey,Ez,Hx,Hy,Hz,x0EPhi1,x0EPhi2,x0HPhi1,x0HPhi2,xmEPhi1,xmEPhi2,xmHPhi1,xmHPhi2,y0EPhi1,y0EPhi2,y0HPhi1,y0HPhi2,ymEPhi1,ymEPhi2,ymHPhi1,ymHPhi2,z0EPhi1,z0EPhi2,z0HPhi1,z0HPhi2,zmEPhi1,zmEPhi2,zmHPhi1,zmHPhi2, mr,dx,nx,ny,nz,dt,nt,nstep,source_amplitudes,source_location,receiver_location,pmlthick,nsr,nrx,device,dtype,x0,xm,y0,ym,z0,zm,x01,x02,xm1,xm2,y01,y02,ym1,ym2,z01,z02,zm1,zm2,ere,see,source_direction, reciever_direction, model_gradient_sampling_interval, use_async_offload)
28
+
29
+ return Eall,(Ex,Ey,Ez),(Hx,Hy,Hz),(x0EPhi1,x0EPhi2,x0HPhi1,x0HPhi2,xmEPhi1,xmEPhi2,xmHPhi1,xmHPhi2,y0EPhi1,y0EPhi2,y0HPhi1,y0HPhi2,ymEPhi1,ymEPhi2,ymHPhi1,ymHPhi2,z0EPhi1,z0EPhi2,z0HPhi1,z0HPhi2,zmEPhi1,zmEPhi2,zmHPhi1,zmHPhi2),receiver_amplitudes
30
+
31
+
32
+ class DeepGPR(torch.autograd.Function):
33
+ @staticmethod
34
+ def forward(ctx, er, se,Ex,Ey,Ez,Hx,Hy,Hz,x0EPhi1,x0EPhi2,x0HPhi1,x0HPhi2,xmEPhi1,xmEPhi2,xmHPhi1,xmHPhi2,y0EPhi1,y0EPhi2,y0HPhi1,y0HPhi2,ymEPhi1,ymEPhi2,ymHPhi1,ymHPhi2,z0EPhi1,z0EPhi2,z0HPhi1,z0HPhi2,zmEPhi1,zmEPhi2,zmHPhi1,zmHPhi2, mr,dx,nx,ny,nz,dt,
35
+ nt,nstep,source_amplitudes,source_location,receiver_location,
36
+ pmlthick,nsr,nrx,device,dtype,x0,xm,
37
+ y0,ym,z0,zm,x01,x02,xm1,xm2,
38
+ y01,y02,ym1,ym2,z01,z02,zm1,zm2,
39
+ ere,see,source_direction, reciever_direction,
40
+ model_gradient_sampling_interval, use_async_offload):
41
+
42
+ source_amplitudes = source_amplitudes.contiguous()
43
+ source_location=source_location.to(torch.int32).contiguous()
44
+ receiver_location=receiver_location.to(torch.int32).contiguous()
45
+ ctx.save_for_backward(er, se, mr,receiver_location,x0,xm,y0,ym,z0,zm,x01,x02,xm1,xm2,y01,y02,ym1,ym2,z01,z02,zm1,zm2,ere,see)
46
+ ctx.dx=dx
47
+ ctx.nx=nx
48
+ ctx.ny=ny
49
+ ctx.nz=nz
50
+ ctx.dt=dt
51
+ ctx.nt=nt
52
+ ctx.nrx=nrx
53
+ ctx.nsr=nsr
54
+ ctx.nstep=nstep
55
+ ctx.pmlthick=pmlthick
56
+ ctx.device=device
57
+ ctx.dtype=dtype
58
+ ctx.model_gradient_sampling_interval = model_gradient_sampling_interval
59
+ ctx.use_async_offload = use_async_offload
60
+
61
+ nt_saved = (nt + model_gradient_sampling_interval - 1) // model_gradient_sampling_interval
62
+
63
+ # 仅仅控制显存分配策略,无需修改 C 端参数签名
64
+ if use_async_offload:
65
+ Eall = torch.zeros((nt_saved, nstep, nx, ny, nz), device='cpu', dtype=dtype).pin_memory()
66
+ else:
67
+ Eall = torch.zeros((nt_saved, nstep, nx, ny, nz), device=device, dtype=dtype).contiguous()
68
+
69
+ Eupdatecoffs0=torch.zeros((nx+1,ny+1,nz+1), device=device, dtype=dtype)
70
+ Eupdatecoffs1=torch.zeros((nx+1,ny+1,nz+1), device=device, dtype=dtype)
71
+ Eupdatecoffs4=torch.zeros((nx+1,ny+1,nz+1), device=device, dtype=dtype)
72
+ Hupdatecoffs0=torch.zeros((nx+1,ny+1,nz+1), device=device, dtype=dtype)
73
+ Hupdatecoffs1=torch.zeros((nx+1,ny+1,nz+1), device=device, dtype=dtype)
74
+ Hupdatecoffs4=torch.zeros((nx+1,ny+1,nz+1), device=device, dtype=dtype)
75
+
76
+ receiver_amplitudes = torch.zeros((nstep, 6, nt, nrx),device=device).contiguous()
77
+
78
+ c_lib.forward(
79
+ ctypes.cast(ere.data_ptr(), ctypes.POINTER(ctypes.c_float)),
80
+ ctypes.cast(see.data_ptr(), ctypes.POINTER(ctypes.c_float)),
81
+ ctypes.cast(mr.data_ptr(), ctypes.POINTER(ctypes.c_float)),
82
+ ctypes.cast(Eall.data_ptr(), ctypes.POINTER(ctypes.c_float)),
83
+ ctypes.cast(Ex.data_ptr(), ctypes.POINTER(ctypes.c_float)),
84
+ ctypes.cast(Ey.data_ptr(), ctypes.POINTER(ctypes.c_float)),
85
+ ctypes.cast(Ez.data_ptr(), ctypes.POINTER(ctypes.c_float)),
86
+ ctypes.cast(Hx.data_ptr(), ctypes.POINTER(ctypes.c_float)),
87
+ ctypes.cast(Hy.data_ptr(), ctypes.POINTER(ctypes.c_float)),
88
+ ctypes.cast(Hz.data_ptr(), ctypes.POINTER(ctypes.c_float)),
89
+
90
+ ctypes.cast(Eupdatecoffs0.data_ptr(), ctypes.POINTER(ctypes.c_float)), ctypes.cast(Eupdatecoffs1.data_ptr(), ctypes.POINTER(ctypes.c_float)),
91
+ ctypes.cast(Eupdatecoffs4.data_ptr(), ctypes.POINTER(ctypes.c_float)), ctypes.cast(Hupdatecoffs0.data_ptr(), ctypes.POINTER(ctypes.c_float)),
92
+ ctypes.cast(Hupdatecoffs1.data_ptr(), ctypes.POINTER(ctypes.c_float)), ctypes.cast(Hupdatecoffs4.data_ptr(), ctypes.POINTER(ctypes.c_float)),
93
+
94
+ ctypes.cast(x0EPhi1.data_ptr(), ctypes.POINTER(ctypes.c_float)), ctypes.cast(x0EPhi2.data_ptr(), ctypes.POINTER(ctypes.c_float)),
95
+ ctypes.cast(x0HPhi1.data_ptr(), ctypes.POINTER(ctypes.c_float)), ctypes.cast(x0HPhi2.data_ptr(), ctypes.POINTER(ctypes.c_float)),
96
+ ctypes.cast(xmEPhi1.data_ptr(), ctypes.POINTER(ctypes.c_float)), ctypes.cast(xmEPhi2.data_ptr(), ctypes.POINTER(ctypes.c_float)),
97
+ ctypes.cast(xmHPhi1.data_ptr(), ctypes.POINTER(ctypes.c_float)), ctypes.cast(xmHPhi2.data_ptr(), ctypes.POINTER(ctypes.c_float)),
98
+ ctypes.cast(y0EPhi1.data_ptr(), ctypes.POINTER(ctypes.c_float)), ctypes.cast(y0EPhi2.data_ptr(), ctypes.POINTER(ctypes.c_float)),
99
+ ctypes.cast(y0HPhi1.data_ptr(), ctypes.POINTER(ctypes.c_float)), ctypes.cast(y0HPhi2.data_ptr(), ctypes.POINTER(ctypes.c_float)),
100
+ ctypes.cast(ymEPhi1.data_ptr(), ctypes.POINTER(ctypes.c_float)), ctypes.cast(ymEPhi2.data_ptr(), ctypes.POINTER(ctypes.c_float)),
101
+ ctypes.cast(ymHPhi1.data_ptr(), ctypes.POINTER(ctypes.c_float)), ctypes.cast(ymHPhi2.data_ptr(), ctypes.POINTER(ctypes.c_float)),
102
+ ctypes.cast(z0EPhi1.data_ptr(), ctypes.POINTER(ctypes.c_float)), ctypes.cast(z0EPhi2.data_ptr(), ctypes.POINTER(ctypes.c_float)),
103
+ ctypes.cast(z0HPhi1.data_ptr(), ctypes.POINTER(ctypes.c_float)), ctypes.cast(z0HPhi2.data_ptr(), ctypes.POINTER(ctypes.c_float)),
104
+ ctypes.cast(zmEPhi1.data_ptr(), ctypes.POINTER(ctypes.c_float)), ctypes.cast(zmEPhi2.data_ptr(), ctypes.POINTER(ctypes.c_float)),
105
+ ctypes.cast(zmHPhi1.data_ptr(), ctypes.POINTER(ctypes.c_float)), ctypes.cast(zmHPhi2.data_ptr(), ctypes.POINTER(ctypes.c_float)),
106
+
107
+ pmlthick[0],pmlthick[1],pmlthick[2],
108
+ pmlthick[3],pmlthick[4],pmlthick[5],
109
+
110
+ ctypes.cast(x01.data_ptr(), ctypes.POINTER(ctypes.c_float)), ctypes.cast(xm1.data_ptr(), ctypes.POINTER(ctypes.c_float)),
111
+ ctypes.cast(y01.data_ptr(), ctypes.POINTER(ctypes.c_float)), ctypes.cast(ym1.data_ptr(), ctypes.POINTER(ctypes.c_float)),
112
+ ctypes.cast(z01.data_ptr(), ctypes.POINTER(ctypes.c_float)), ctypes.cast(zm1.data_ptr(), ctypes.POINTER(ctypes.c_float)),
113
+ ctypes.cast(x02.data_ptr(), ctypes.POINTER(ctypes.c_float)), ctypes.cast(xm2.data_ptr(), ctypes.POINTER(ctypes.c_float)),
114
+ ctypes.cast(y02.data_ptr(), ctypes.POINTER(ctypes.c_float)), ctypes.cast(ym2.data_ptr(), ctypes.POINTER(ctypes.c_float)),
115
+ ctypes.cast(z02.data_ptr(), ctypes.POINTER(ctypes.c_float)), ctypes.cast(zm2.data_ptr(), ctypes.POINTER(ctypes.c_float)),
116
+
117
+ dt, nt, nstep, nrx, dx,
118
+ ctypes.cast(receiver_location.data_ptr(), ctypes.POINTER(ctypes.c_int)), ctypes.cast(receiver_amplitudes.data_ptr(), ctypes.POINTER(ctypes.c_float)),
119
+ nx+1, ny+1, nz+1, nsr,
120
+ ctypes.cast(source_location.data_ptr(), ctypes.POINTER(ctypes.c_int)), ctypes.cast(source_amplitudes.data_ptr(), ctypes.POINTER(ctypes.c_float)),
121
+ source_direction,
122
+ model_gradient_sampling_interval)
123
+
124
+ check_tensors_for_nan_inf(d="forward",
125
+ Ex=Ex, Ey=Ey, Ez=Ez,
126
+ Hx=Hx, Hy=Hy, Hz=Hz,
127
+ x0EPhi1=x0EPhi1, x0EPhi2=x0EPhi2,
128
+ x0HPhi1=x0HPhi1, x0HPhi2=x0HPhi2,
129
+ xmEPhi1=xmEPhi1, xmEPhi2=xmEPhi2,
130
+ xmHPhi1=xmHPhi1, xmHPhi2=xmHPhi2,
131
+ y0EPhi1=y0EPhi1, y0EPhi2=y0EPhi2,
132
+ y0HPhi1=y0HPhi1, y0HPhi2=y0HPhi2,
133
+ ymEPhi1=ymEPhi1, ymEPhi2=ymEPhi2,
134
+ ymHPhi1=ymHPhi1, ymHPhi2=ymHPhi2,
135
+ z0EPhi1=z0EPhi1, z0EPhi2=z0EPhi2,
136
+ z0HPhi1=z0HPhi1, z0HPhi2=z0HPhi2,
137
+ zmEPhi1=zmEPhi1, zmEPhi2=zmEPhi2,
138
+ zmHPhi1=zmHPhi1, zmHPhi2=zmHPhi2
139
+ )
140
+
141
+ ctx.Eall = Eall
142
+ return (Ex,Ey,Ez,Hx,Hy,Hz,x0EPhi1,x0EPhi2,x0HPhi1,x0HPhi2,xmEPhi1,xmEPhi2,xmHPhi1,xmHPhi2,y0EPhi1,y0EPhi2,y0HPhi1,y0HPhi2,ymEPhi1,ymEPhi2,ymHPhi1,ymHPhi2,z0EPhi1,z0EPhi2,z0HPhi1,z0HPhi2,zmEPhi1,zmEPhi2,zmHPhi1,zmHPhi2,Eall,receiver_amplitudes[:,reciever_direction,:,:])
143
+
144
+ @staticmethod
145
+ def backward(ctx,gEx,gEy,gEz,gHx,gHy,gHz,gx0EPhi1,gx0EPhi2,gx0HPhi1,gx0HPhi2,gxmEPhi1,gxmEPhi2,gxmHPhi1,gxmHPhi2,gy0EPhi1,gy0EPhi2,gy0HPhi1,gy0HPhi2,gymEPhi1,gymEPhi2,gymHPhi1,gymHPhi2,gz0EPhi1,gz0EPhi2,gz0HPhi1,gz0HPhi2,gzmEPhi1,gzmEPhi2,gzmHPhi1,gzmHPhi2,gEall,gezreciver):
146
+
147
+ sourceamp=gezreciver
148
+ er, se, mr,receiver_location,x0,xm,y0,ym,z0,zm,x01,x02,xm1,xm2,y01,y02,ym1,ym2,z01,z02,zm1,zm2,ere,see=ctx.saved_tensors
149
+
150
+ ere=ere.contiguous()
151
+ see=see.contiguous()
152
+ er=er.contiguous()
153
+ se=se.contiguous()
154
+ mr=mr.contiguous()
155
+ receiver_location=receiver_location.contiguous()
156
+ x0=x0.contiguous()
157
+ xm=xm.contiguous()
158
+ y0=y0.contiguous()
159
+ ym=ym.contiguous()
160
+ z0=z0.contiguous()
161
+ zm=zm.contiguous()
162
+ x01=x01.contiguous()
163
+ x02=x02.contiguous()
164
+ xm1=xm1.contiguous()
165
+ xm2=xm2.contiguous()
166
+ y01=y01.contiguous()
167
+ y02=y02.contiguous()
168
+ ym1=ym1.contiguous()
169
+ ym2=ym2.contiguous()
170
+ z01=z01.contiguous()
171
+ z02=z02.contiguous()
172
+ zm1=zm1.contiguous()
173
+ zm2=zm2.contiguous()
174
+
175
+ dx=ctx.dx
176
+ nx=ctx.nx
177
+ ny=ctx.ny
178
+ nz=ctx.nz
179
+ dt=ctx.dt
180
+ nt=ctx.nt
181
+ nsr=ctx.nrx
182
+ nrx=ctx.nsr
183
+ dtype=ctx.dtype
184
+ nstep=ctx.nstep
185
+ pmlthick=ctx.pmlthick
186
+ device=ctx.device
187
+ Eall=ctx.Eall
188
+ model_gradient_sampling_interval = ctx.model_gradient_sampling_interval
189
+
190
+ Eall=Eall.contiguous()
191
+ gEx=gEx.contiguous()
192
+ gEy=gEy.contiguous()
193
+ gEz=gEz.contiguous()
194
+ gHx=gHx.contiguous()
195
+ gHy=gHy.contiguous()
196
+ gHz=gHz.contiguous()
197
+ gx0EPhi1=gx0EPhi1.contiguous()
198
+ gx0EPhi2=gx0EPhi2.contiguous()
199
+ gx0HPhi1=gx0HPhi1.contiguous()
200
+ gx0HPhi2=gx0HPhi2.contiguous()
201
+ gxmEPhi1=gxmEPhi1.contiguous()
202
+ gxmEPhi2=gxmEPhi2.contiguous()
203
+ gxmHPhi1=gxmHPhi1.contiguous()
204
+ gxmHPhi2=gxmHPhi2.contiguous()
205
+ gy0EPhi1=gy0EPhi1.contiguous()
206
+ gy0EPhi2=gy0EPhi2.contiguous()
207
+ gy0HPhi1=gy0HPhi1.contiguous()
208
+ gy0HPhi2=gy0HPhi2.contiguous()
209
+ gymEPhi1=gymEPhi1.contiguous()
210
+ gymEPhi2=gymEPhi2.contiguous()
211
+ gymHPhi1=gymHPhi1.contiguous()
212
+ gymHPhi2=gymHPhi2.contiguous()
213
+ gz0EPhi1=gz0EPhi1.contiguous()
214
+ gz0EPhi2=gz0EPhi2.contiguous()
215
+ gz0HPhi1=gz0HPhi1.contiguous()
216
+ gz0HPhi2=gz0HPhi2.contiguous()
217
+ gzmEPhi1=gzmEPhi1.contiguous()
218
+ gzmEPhi2=gzmEPhi2.contiguous()
219
+ gzmHPhi1=gzmHPhi1.contiguous()
220
+ gzmHPhi2=gzmHPhi2.contiguous()
221
+ sourceamp=sourceamp.contiguous()
222
+
223
+ Eupdatecoffs0=torch.zeros((nx+1,ny+1,nz+1), device=device, dtype=dtype)
224
+ Eupdatecoffs1=torch.zeros((nx+1,ny+1,nz+1), device=device, dtype=dtype)
225
+ Eupdatecoffs4=torch.zeros((nx+1,ny+1,nz+1), device=device, dtype=dtype)
226
+ Hupdatecoffs0=torch.zeros((nx+1,ny+1,nz+1), device=device, dtype=dtype)
227
+ Hupdatecoffs1=torch.zeros((nx+1,ny+1,nz+1), device=device, dtype=dtype)
228
+ Hupdatecoffs4=torch.zeros((nx+1,ny+1,nz+1), device=device, dtype=dtype)
229
+
230
+ if er.requires_grad:
231
+ grad_er=torch.zeros((nx,ny,nz),device=device,dtype=dtype).contiguous()
232
+ errequiregrad=1
233
+ else:
234
+ grad_er=torch.empty(0)
235
+ errequiregrad=0
236
+
237
+ if se.requires_grad:
238
+ grad_se=torch.zeros((nx,ny,nz),device=device,dtype=dtype).contiguous()
239
+ serequiregrad=1
240
+ else:
241
+ grad_se=torch.empty(0)
242
+ serequiregrad=0
243
+
244
+ c_lib.backward(
245
+ ctypes.cast(ere.data_ptr(), ctypes.POINTER(ctypes.c_float)),
246
+ ctypes.cast(see.data_ptr(), ctypes.POINTER(ctypes.c_float)),
247
+ ctypes.cast(mr.data_ptr(), ctypes.POINTER(ctypes.c_float)),
248
+ ctypes.cast(Eall.data_ptr(), ctypes.POINTER(ctypes.c_float)),
249
+ ctypes.cast(gEx.data_ptr(), ctypes.POINTER(ctypes.c_float)),
250
+ ctypes.cast(gEy.data_ptr(), ctypes.POINTER(ctypes.c_float)),
251
+ ctypes.cast(gEz.data_ptr(), ctypes.POINTER(ctypes.c_float)),
252
+ ctypes.cast(gHx.data_ptr(), ctypes.POINTER(ctypes.c_float)),
253
+ ctypes.cast(gHy.data_ptr(), ctypes.POINTER(ctypes.c_float)),
254
+ ctypes.cast(gHz.data_ptr(), ctypes.POINTER(ctypes.c_float)),
255
+
256
+ ctypes.cast(Eupdatecoffs0.data_ptr(), ctypes.POINTER(ctypes.c_float)), ctypes.cast(Eupdatecoffs1.data_ptr(), ctypes.POINTER(ctypes.c_float)),
257
+ ctypes.cast(Eupdatecoffs4.data_ptr(), ctypes.POINTER(ctypes.c_float)), ctypes.cast(Hupdatecoffs0.data_ptr(), ctypes.POINTER(ctypes.c_float)),
258
+ ctypes.cast(Hupdatecoffs1.data_ptr(), ctypes.POINTER(ctypes.c_float)), ctypes.cast(Hupdatecoffs4.data_ptr(), ctypes.POINTER(ctypes.c_float)),
259
+
260
+ ctypes.cast(gx0EPhi1.data_ptr(), ctypes.POINTER(ctypes.c_float)), ctypes.cast(gx0EPhi2.data_ptr(), ctypes.POINTER(ctypes.c_float)),
261
+ ctypes.cast(gx0HPhi1.data_ptr(), ctypes.POINTER(ctypes.c_float)), ctypes.cast(gx0HPhi2.data_ptr(), ctypes.POINTER(ctypes.c_float)),
262
+ ctypes.cast(gxmEPhi1.data_ptr(), ctypes.POINTER(ctypes.c_float)), ctypes.cast(gxmEPhi2.data_ptr(), ctypes.POINTER(ctypes.c_float)),
263
+ ctypes.cast(gxmHPhi1.data_ptr(), ctypes.POINTER(ctypes.c_float)), ctypes.cast(gxmHPhi2.data_ptr(), ctypes.POINTER(ctypes.c_float)),
264
+ ctypes.cast(gy0EPhi1.data_ptr(), ctypes.POINTER(ctypes.c_float)), ctypes.cast(gy0EPhi2.data_ptr(), ctypes.POINTER(ctypes.c_float)),
265
+ ctypes.cast(gy0HPhi1.data_ptr(), ctypes.POINTER(ctypes.c_float)), ctypes.cast(gy0HPhi2.data_ptr(), ctypes.POINTER(ctypes.c_float)),
266
+ ctypes.cast(gymEPhi1.data_ptr(), ctypes.POINTER(ctypes.c_float)), ctypes.cast(gymEPhi2.data_ptr(), ctypes.POINTER(ctypes.c_float)),
267
+ ctypes.cast(gymHPhi1.data_ptr(), ctypes.POINTER(ctypes.c_float)), ctypes.cast(gymHPhi2.data_ptr(), ctypes.POINTER(ctypes.c_float)),
268
+ ctypes.cast(gz0EPhi1.data_ptr(), ctypes.POINTER(ctypes.c_float)), ctypes.cast(gz0EPhi2.data_ptr(), ctypes.POINTER(ctypes.c_float)),
269
+ ctypes.cast(gz0HPhi1.data_ptr(), ctypes.POINTER(ctypes.c_float)), ctypes.cast(gz0HPhi2.data_ptr(), ctypes.POINTER(ctypes.c_float)),
270
+ ctypes.cast(gzmEPhi1.data_ptr(), ctypes.POINTER(ctypes.c_float)), ctypes.cast(gzmEPhi2.data_ptr(), ctypes.POINTER(ctypes.c_float)),
271
+ ctypes.cast(gzmHPhi1.data_ptr(), ctypes.POINTER(ctypes.c_float)), ctypes.cast(gzmHPhi2.data_ptr(), ctypes.POINTER(ctypes.c_float)),
272
+
273
+ pmlthick[0],pmlthick[1],pmlthick[2],
274
+ pmlthick[3],pmlthick[4],pmlthick[5],
275
+
276
+ ctypes.cast(x01.data_ptr(), ctypes.POINTER(ctypes.c_float)), ctypes.cast(xm1.data_ptr(), ctypes.POINTER(ctypes.c_float)),
277
+ ctypes.cast(y01.data_ptr(), ctypes.POINTER(ctypes.c_float)), ctypes.cast(ym1.data_ptr(), ctypes.POINTER(ctypes.c_float)),
278
+ ctypes.cast(z01.data_ptr(), ctypes.POINTER(ctypes.c_float)), ctypes.cast(zm1.data_ptr(), ctypes.POINTER(ctypes.c_float)),
279
+ ctypes.cast(x02.data_ptr(), ctypes.POINTER(ctypes.c_float)), ctypes.cast(xm2.data_ptr(), ctypes.POINTER(ctypes.c_float)),
280
+ ctypes.cast(y02.data_ptr(), ctypes.POINTER(ctypes.c_float)), ctypes.cast(ym2.data_ptr(), ctypes.POINTER(ctypes.c_float)),
281
+ ctypes.cast(z02.data_ptr(), ctypes.POINTER(ctypes.c_float)), ctypes.cast(zm2.data_ptr(), ctypes.POINTER(ctypes.c_float)),
282
+
283
+ dt, nt, nstep, nrx, dx,
284
+ nx+1, ny+1, nz+1, nsr,
285
+ ctypes.cast(receiver_location.data_ptr(), ctypes.POINTER(ctypes.c_int)), ctypes.cast(sourceamp.data_ptr(), ctypes.POINTER(ctypes.c_float)),
286
+ 2,
287
+ ctypes.cast(grad_er.data_ptr(), ctypes.POINTER(ctypes.c_float)), ctypes.cast(grad_se.data_ptr(), ctypes.POINTER(ctypes.c_float)),errequiregrad,serequiregrad,
288
+ model_gradient_sampling_interval)
289
+
290
+ check_tensors_for_nan_inf(d="backward",
291
+ gEx=gEx, gEy=gEy, gEz=gEz,
292
+ gHx=gHx, gHy=gHy, gHz=gHz,
293
+ gx0EPhi1=gx0EPhi1, gx0EPhi2=gx0EPhi2,
294
+ gx0HPhi1=gx0HPhi1, gx0HPhi2=gx0HPhi2,
295
+ gxmEPhi1=gxmEPhi1, gxmEPhi2=gxmEPhi2,
296
+ gxmHPhi1=gxmHPhi1, gxmHPhi2=gxmHPhi2,
297
+ gy0EPhi1=gy0EPhi1, gy0EPhi2=gy0EPhi2,
298
+ gy0HPhi1=gy0HPhi1, gy0HPhi2=gy0HPhi2,
299
+ gymEPhi1=gymEPhi1, gymEPhi2=gymEPhi2,
300
+ gymHPhi1=gymHPhi1, gymHPhi2=gymHPhi2,
301
+ gz0EPhi1=gz0EPhi1, gz0EPhi2=gz0EPhi2,
302
+ gz0HPhi1=gz0HPhi1, gz0HPhi2=gz0HPhi2,
303
+ gzmEPhi1=gzmEPhi1, gzmEPhi2=gzmEPhi2,
304
+ gzmHPhi1=gzmHPhi1, gzmHPhi2=gzmHPhi2
305
+ )
306
+
307
+ ctx.Eall = None
308
+ del Eall,er, se, mr,receiver_location,x0,xm,y0,ym,z0,zm,x01,x02,xm1,xm2,y01,y02,ym1,ym2,z01,z02,zm1,zm2,ere,see, Eupdatecoffs0, Eupdatecoffs1, Eupdatecoffs4, Hupdatecoffs0, Hupdatecoffs1, Hupdatecoffs4
309
+
310
+ # 返回与前向传播参数一一对应(最后补齐 use_async_offload 占位)
311
+ return (
312
+ grad_er, grad_se,
313
+ gEx,gEy,gEz, gHx,gHy,gHz,
314
+ gx0EPhi1,gx0EPhi2,gx0HPhi1,gx0HPhi2,
315
+ gxmEPhi1,gxmEPhi2,gxmHPhi1,gxmHPhi2,
316
+ gy0EPhi1,gy0EPhi2,gy0HPhi1,gy0HPhi2,
317
+ gymEPhi1,gymEPhi2,gymHPhi1,gymHPhi2,
318
+ gz0EPhi1,gz0EPhi2,gz0HPhi1,gz0HPhi2,
319
+ gzmEPhi1,gzmEPhi2,gzmHPhi1,gzmHPhi2,
320
+ None, None, None, None, None, None, None, None,
321
+ None, None, None, None, None, None, None,
322
+ None, None, None, None, None, None, None, None,
323
+ None, None, None, None, None, None, None, None,
324
+ None, None, None, None, None, None, None, None,
325
+ None
326
+ )
DeepGPR/multiscale.py ADDED
@@ -0,0 +1,63 @@
1
+
2
+ import torch
3
+
4
+ def design_fir_filter(cutoff: float, fs: float, numtaps: int) -> torch.Tensor:
5
+ """使用 Hamming 窗设计 FIR 滤波器"""
6
+ n = torch.arange(numtaps, dtype=torch.float32)
7
+ # 汉明窗
8
+ window = 0.54 - 0.46 * torch.cos(2 * torch.pi * n / (numtaps - 1))
9
+ # 正弦波
10
+ sinc = torch.sin(2 * torch.pi * (cutoff/fs) * (n - (numtaps-1)/2)) / (torch.pi * (n - (numtaps-1)/2))
11
+ # 处理中心点
12
+ center = (numtaps-1) // 2
13
+ sinc[center] = 2 * cutoff/fs
14
+ # 应用窗函数
15
+ h = window * sinc
16
+ # 归一化
17
+ return h / h.sum()
18
+
19
+ def apply_filter(data: torch.Tensor, fs: float, cutoff: float) -> torch.Tensor:
20
+ """应用 FIR 滤波器到数据"""
21
+ numtaps = int(1 * (fs / cutoff))
22
+ fir_coeff = design_fir_filter(cutoff, fs, numtaps)
23
+ fir_coeff = fir_coeff.to(data.device)
24
+
25
+ if data.ndim == 1:
26
+ # 1D 数据处理 - 添加维度以支持反射填充
27
+ data_2d = data.view(1, 1, -1)
28
+ padded_data = torch.nn.functional.pad(data_2d, (numtaps-1, 0), mode='reflect')
29
+ filtered = torch.nn.functional.conv1d(
30
+ padded_data,
31
+ fir_coeff.view(1, 1, -1),
32
+ padding=0
33
+ )
34
+ return filtered.view(-1)
35
+
36
+ elif data.ndim == 3:
37
+ # 3D 数据处理
38
+ step, iterations, nrx = data.shape
39
+ # 重塑数据以使用批量处理
40
+ reshaped_data = data.permute(0, 2, 1).reshape(-1, 1, iterations)
41
+ # 填充数据
42
+ padded_data = torch.nn.functional.pad(reshaped_data, (numtaps-1, 0), mode='reflect')
43
+ # 应用卷积
44
+ filtered = torch.nn.functional.conv1d(
45
+ padded_data,
46
+ fir_coeff.view(1, 1, -1),
47
+ padding=0
48
+ )
49
+ # 重塑回原始形状
50
+ return filtered.view(step, nrx, iterations).permute(0, 2, 1)
51
+
52
+ else:
53
+ raise ValueError(f"不支持的数据维度: {data.ndim}。期望 1D 或 3D tensor。")
54
+
55
+ def hilbert_transform(data_in, p=1):
56
+ ns, nt, nr = data_in.shape
57
+ transforms = torch.fft.fftn(data_in,dim=1)
58
+ #print(transforms.shape)
59
+ transforms[:,1:nt//2,:] *= 2.0
60
+ transforms[:,nt//2 + 1: nt,:] = 0+0j
61
+ transforms[:,0,:] = 0;
62
+ data_out = torch.abs(torch.fft.ifftn(transforms,dim=1))**p
63
+ return data_out
DeepGPR/visual.py ADDED
@@ -0,0 +1,105 @@
1
+ import torch
2
+ import numpy as np
3
+ import matplotlib.pyplot as plt
4
+
5
+
6
+ def plot_survey_geometry(model, sources, receivers, dx=0.02, title="Survey Geometry"):
7
+
8
+ # 1. 数据转换与展平
9
+ if isinstance(model, torch.Tensor):
10
+ model = model.detach().cpu().numpy()
11
+ if isinstance(sources, torch.Tensor):
12
+ sources = sources.detach().cpu().numpy()
13
+ if isinstance(receivers, torch.Tensor):
14
+ receivers = receivers.detach().cpu().numpy()
15
+
16
+ # 处理多炮维度 [Shots, N, 3] -> [Total, 3]
17
+ if sources.ndim == 3: sources = sources.reshape(-1, 3)
18
+ if receivers.ndim == 3: receivers = receivers.reshape(-1, 3)
19
+
20
+ nx, ny, nz = model.shape
21
+
22
+ # 2. 绘图
23
+ if nz == 1:
24
+ # ==========================================
25
+ # 2D 绘图模式 (垂直朝上,0在顶部)
26
+ # ==========================================
27
+ fig, ax = plt.subplots(figsize=(6, 4))
28
+
29
+ slice_data = model[:, :, 0] # Shape: (nx, ny)
30
+
31
+ # [修改 1] extent 乘以 dx,转换为真实物理距离
32
+ # extent=[左, 右, 下, 上]
33
+ # 垂直轴从 nx*dx (底部) 到 0 (顶部)
34
+ extent_real = [0, ny * dx, nx * dx, 0]
35
+
36
+ im = ax.imshow(slice_data, cmap='jet', origin='upper',
37
+ extent=extent_real, aspect='auto')
38
+
39
+ plt.colorbar(im, ax=ax, label='Relative permittivity')
40
+
41
+ # [修改 2] 散点坐标乘以 dx
42
+ # Scatter x (水平) = sources[:, 1] * dx
43
+ # Scatter y (垂直) = sources[:, 0] * dx
44
+ ax.scatter(sources[:, 1] * dx, sources[:, 0] * dx, c='red', marker='*', s=150,
45
+ label='Sources', edgecolors='k', zorder=10)
46
+
47
+ ax.scatter(receivers[:, 1] * dx, receivers[:, 0] * dx, c='white', marker='v', s=80,
48
+ label='Receivers', edgecolors='k', zorder=9, alpha=0.7)
49
+
50
+ # [修改 3] 标签改为物理单位
51
+ ax.set_xlabel('Distance (m)')
52
+ ax.set_ylabel('Depth (m)')
53
+
54
+ # 保持 Y 轴方向控制 (0 在上)
55
+ if ax.get_ylim()[0] < ax.get_ylim()[1]:
56
+ ax.invert_yaxis()
57
+
58
+ # ax.set_title(f"{title} ")
59
+ ax.legend(loc='lower right')
60
+ ax.grid(True, linestyle='--', alpha=0.3)
61
+
62
+ else:
63
+ # ==========================================
64
+ # 3D 绘图模式 (同步应用 dx 以保持一致)
65
+ # ==========================================
66
+ fig = plt.figure(figsize=(12, 10))
67
+ ax = fig.add_subplot(111, projection='3d')
68
+
69
+ # 散点坐标乘以 dx
70
+ ax.scatter(sources[:, 0] * dx, sources[:, 1] * dx, sources[:, 2] * dx,
71
+ c='red', marker='*', s=100, label='Sources')
72
+ ax.scatter(receivers[:, 0] * dx, receivers[:, 1] * dx, receivers[:, 2] * dx,
73
+ c='blue', marker='v', s=40, label='Receivers', alpha=0.6)
74
+
75
+ # 网格坐标乘以 dx
76
+ x_grid = np.arange(nx) * dx
77
+ y_grid = np.arange(ny) * dx
78
+ z_grid = np.arange(nz) * dx
79
+
80
+ cx, cy, cz = nx // 2, ny // 2, nz // 2
81
+
82
+ # 生成真实坐标的 Meshgrid
83
+ X, Y = np.meshgrid(x_grid, y_grid, indexing='ij')
84
+ # 注意 offset 也要乘以 dx
85
+ ax.contourf(X, Y, model[:, :, cz], zdir='z', offset=cz * dx, cmap='viridis', alpha=0.5)
86
+
87
+ X, Z = np.meshgrid(x_grid, z_grid, indexing='ij')
88
+ ax.contourf(X, model[:, cy, :], Z, zdir='y', offset=cy * dx, cmap='viridis', alpha=0.5)
89
+
90
+ Y, Z = np.meshgrid(y_grid, z_grid, indexing='ij')
91
+ ax.contourf(model[cx, :, :], Y, Z, zdir='x', offset=cx * dx, cmap='viridis', alpha=0.5)
92
+
93
+ # 设置真实坐标范围
94
+ ax.set_xlim(0, nx * dx)
95
+ ax.set_ylim(0, ny * dx)
96
+ ax.set_zlim(0, nz * dx)
97
+
98
+ ax.set_xlabel('X (m)')
99
+ ax.set_ylabel('Y (m)')
100
+ ax.set_zlabel('Z (m)')
101
+ ax.set_title(title)
102
+ ax.legend()
103
+
104
+ plt.tight_layout()
105
+ plt.show()
@@ -0,0 +1,86 @@
1
+ Metadata-Version: 2.4
2
+ Name: DeepGPR
3
+ Version: 0.0.1
4
+ Summary: PyTorch and CUDA for GPR FWI
5
+ Author-email: Lei Liu <liulei990222@gmail.com>
6
+ Classifier: Programming Language :: Python :: 3
7
+ Classifier: Operating System :: OS Independent
8
+ Requires-Python: >=3.7
9
+ Description-Content-Type: text/markdown
10
+
11
+ # DeepGPR
12
+
13
+ DeepGPR provides a wave propagation module for PyTorch, designed for applications such as Ground Penetrating Radar (GPR) imaging and inversion. Its core concepts are derived from Deepwave. You can use it to perform both forward modeling and backpropagation—thereby enabling the simulation of wave propagation to generate synthetic data—as well as for Full Waveform Inversion (FWI). Furthermore, you can integrate this wave propagation functionality into a larger operational pipeline—incorporating various wavelets, loss functions, and other components—to achieve end-to-end forward and reverse propagation, powered by automatic differentiation and our high-performance operators.
14
+
15
+
16
+ ## Features
17
+
18
+ Supports 2D and 3D forward modeling of Maxwell's equations—via the Finite-Difference Time-Domain (FDTD) method—for both single and multiple excitation scenarios.
19
+
20
+ Gradients of the output receiver data can be computed with respect to model parameters (relative permittivity, conductivity), the initial wavefield, and source amplitudes.
21
+
22
+ Utilizes CPML, allowing the width of the PML layer to be configured independently for each boundary.
23
+
24
+ All operations are executed on the GPU.
25
+
26
+ Supports techniques such as checkpointing, DDP, and the utilization of CPU memory to minimize GPU memory consumption, thereby enabling the execution of large-scale models.
27
+
28
+
29
+ ## Start
30
+
31
+ Before use, you must ensure that you have an NVIDIA graphics card and have installed a CUDA-enabled version of PyTorch.
32
+
33
+ DeepGPR can then be installed using
34
+
35
+ ```bash
36
+ pip install DeepGPR
37
+ ```
38
+
39
+ A Small Forward Modeling Test
40
+
41
+ ```python
42
+ import torch
43
+ import DeepGPR
44
+ import matplotlib.pyplot as plt
45
+
46
+ # Set up the parameters and models
47
+ device=torch.device("cuda")
48
+ dx=0.02
49
+ dt=3e-11
50
+ nt=2000
51
+ er = torch.ones(100, 100,1) * 2
52
+ er[50:,:]=5
53
+ se = torch.zeros_like(er)
54
+ er.requires_grad_()
55
+ source_location=torch.tensor([[[10,10,0]]],device=device,dtype=torch.int)
56
+ receiver_location=torch.tensor([[[10,90,0]]],device=device,dtype=torch.int)
57
+ freq=2e8
58
+ peak_time = 1 / freq
59
+ source_amplitudes = torch.zeros((1,nt,1),device=device)
60
+ source_amplitudes[0,:,0]=DeepGPR.ricker(freq, nt, dt, peak_time).to(device)
61
+
62
+ DeepGPR.plot_survey_geometry(er,source_location, receiver_location,dx)
63
+
64
+ #forward modeling
65
+ r = DeepGPR.compute(
66
+ device=device, dx=dx, dt=dt,
67
+ source_amplitudes=source_amplitudes,
68
+ source_location=source_location,
69
+ receiver_location=receiver_location,
70
+ er=er, se=se
71
+ )
72
+
73
+ (r[-1]**2).sum().backward()
74
+
75
+ _, ax = plt.subplots(1, 2, figsize=(10, 3))
76
+ ax[0].plot(r[-1].detach().flatten().cpu().numpy())
77
+ ax[0].set_title("Receiver data")
78
+ ax[1].imshow(er.grad.detach())
79
+ ax[1].set_title("Gradient")
80
+ plt.show()
81
+ ```
82
+ ![result](./Fig/example1.png)
83
+
84
+ ![result](./Fig/example2.png)
85
+
86
+ There are more examples in the ./examples.
@@ -0,0 +1,9 @@
1
+ DeepGPR/__init__.py,sha256=mhXJhG4SQ2C1o8GAMygkKxR_annxX1Txku8rdQbmVi0,6663
2
+ DeepGPR/common.py,sha256=If4uIzdLFSUdHTY3bHxB9Ws4LkB5cJiNlOAaY667-es,27915
3
+ DeepGPR/compute2.py,sha256=4b1YWlzPhmD-_rFJ6yRB2iaKHlcOhJccgzbaOR9UJPQ,20230
4
+ DeepGPR/multiscale.py,sha256=_tSFC3alHPDmMBPbRm1pJy5nZHQ-rtvgTjhupfhrPlc,2240
5
+ DeepGPR/visual.py,sha256=Iynfiv_2DKVB7MCbgLvvmWn_m8jhrHoDLlHWyC9iHeY,3944
6
+ deepgpr-0.0.1.dist-info/METADATA,sha256=wijleiqiX73S5EMxPNOWznbu1OnVC9JODpc7EOFsEvE,3089
7
+ deepgpr-0.0.1.dist-info/WHEEL,sha256=aeYiig01lYGDzBgS8HxWXOg3uV61G9ijOsup-k9o1sk,91
8
+ deepgpr-0.0.1.dist-info/top_level.txt,sha256=Clp4TLG1b60qzO8D6ZYieWFKK39K1uU_Em9yRGwelTE,8
9
+ deepgpr-0.0.1.dist-info/RECORD,,
@@ -0,0 +1,5 @@
1
+ Wheel-Version: 1.0
2
+ Generator: setuptools (82.0.1)
3
+ Root-Is-Purelib: true
4
+ Tag: py3-none-any
5
+
@@ -0,0 +1 @@
1
+ DeepGPR