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/__init__.py +147 -0
- DeepGPR/common.py +643 -0
- DeepGPR/compute2.py +326 -0
- DeepGPR/multiscale.py +63 -0
- DeepGPR/visual.py +105 -0
- deepgpr-0.0.1.dist-info/METADATA +86 -0
- deepgpr-0.0.1.dist-info/RECORD +9 -0
- deepgpr-0.0.1.dist-info/WHEEL +5 -0
- deepgpr-0.0.1.dist-info/top_level.txt +1 -0
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
|
+

|
|
83
|
+
|
|
84
|
+

|
|
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 @@
|
|
|
1
|
+
DeepGPR
|