DeepGPR 0.0.2__tar.gz → 0.0.5__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.
- {deepgpr-0.0.2/DeepGPR.egg-info → deepgpr-0.0.5}/PKG-INFO +4 -5
- {deepgpr-0.0.2 → deepgpr-0.0.5}/README.md +1 -3
- {deepgpr-0.0.2 → deepgpr-0.0.5}/pyproject.toml +6 -4
- {deepgpr-0.0.2 → deepgpr-0.0.5/src}/DeepGPR/__init__.py +5 -13
- {deepgpr-0.0.2 → deepgpr-0.0.5/src}/DeepGPR/common.py +2 -79
- {deepgpr-0.0.2 → deepgpr-0.0.5/src}/DeepGPR/compute2.py +0 -2
- {deepgpr-0.0.2 → deepgpr-0.0.5/src}/DeepGPR/lib/deepgpr.cu +8 -30
- {deepgpr-0.0.2 → deepgpr-0.0.5/src}/DeepGPR/multiscale.py +1 -13
- deepgpr-0.0.5/src/DeepGPR/wavelet.py +34 -0
- {deepgpr-0.0.2 → deepgpr-0.0.5/src/DeepGPR.egg-info}/PKG-INFO +4 -5
- deepgpr-0.0.5/src/DeepGPR.egg-info/SOURCES.txt +13 -0
- deepgpr-0.0.2/DeepGPR/visual.py +0 -105
- deepgpr-0.0.2/DeepGPR.egg-info/SOURCES.txt +0 -13
- {deepgpr-0.0.2 → deepgpr-0.0.5}/setup.cfg +0 -0
- {deepgpr-0.0.2 → deepgpr-0.0.5/src}/DeepGPR/lib/deepgpr.so +0 -0
- {deepgpr-0.0.2 → deepgpr-0.0.5/src}/DeepGPR.egg-info/dependency_links.txt +0 -0
- {deepgpr-0.0.2 → deepgpr-0.0.5/src}/DeepGPR.egg-info/top_level.txt +0 -0
|
@@ -1,10 +1,11 @@
|
|
|
1
1
|
Metadata-Version: 2.4
|
|
2
2
|
Name: DeepGPR
|
|
3
|
-
Version: 0.0.
|
|
3
|
+
Version: 0.0.5
|
|
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
|
|
7
|
-
Classifier: Operating System ::
|
|
7
|
+
Classifier: Operating System :: Microsoft :: Windows
|
|
8
|
+
Classifier: Operating System :: POSIX :: Linux
|
|
8
9
|
Requires-Python: >=3.7
|
|
9
10
|
Description-Content-Type: text/markdown
|
|
10
11
|
|
|
@@ -59,7 +60,6 @@ peak_time = 1 / freq
|
|
|
59
60
|
source_amplitudes = torch.zeros((1,nt,1),device=device)
|
|
60
61
|
source_amplitudes[0,:,0]=DeepGPR.ricker(freq, nt, dt, peak_time).to(device)
|
|
61
62
|
|
|
62
|
-
DeepGPR.plot_survey_geometry(er,source_location, receiver_location,dx)
|
|
63
63
|
|
|
64
64
|
#forward modeling
|
|
65
65
|
r = DeepGPR.compute(
|
|
@@ -79,8 +79,7 @@ ax[1].imshow(er.grad.detach())
|
|
|
79
79
|
ax[1].set_title("Gradient")
|
|
80
80
|
plt.show()
|
|
81
81
|
```
|
|
82
|
-

|
|
83
82
|
|
|
84
|
-

|
|
85
84
|
|
|
86
85
|
There are more examples in the ./examples.
|
|
@@ -49,7 +49,6 @@ peak_time = 1 / freq
|
|
|
49
49
|
source_amplitudes = torch.zeros((1,nt,1),device=device)
|
|
50
50
|
source_amplitudes[0,:,0]=DeepGPR.ricker(freq, nt, dt, peak_time).to(device)
|
|
51
51
|
|
|
52
|
-
DeepGPR.plot_survey_geometry(er,source_location, receiver_location,dx)
|
|
53
52
|
|
|
54
53
|
#forward modeling
|
|
55
54
|
r = DeepGPR.compute(
|
|
@@ -69,8 +68,7 @@ ax[1].imshow(er.grad.detach())
|
|
|
69
68
|
ax[1].set_title("Gradient")
|
|
70
69
|
plt.show()
|
|
71
70
|
```
|
|
72
|
-

|
|
73
71
|
|
|
74
|
-

|
|
75
73
|
|
|
76
74
|
There are more examples in the ./examples.
|
|
@@ -4,7 +4,7 @@ build-backend = "setuptools.build_meta"
|
|
|
4
4
|
|
|
5
5
|
[project]
|
|
6
6
|
name = "DeepGPR"
|
|
7
|
-
version = "0.0.
|
|
7
|
+
version = "0.0.5"
|
|
8
8
|
authors = [
|
|
9
9
|
{ name="Lei Liu", email="liulei990222@gmail.com" },
|
|
10
10
|
]
|
|
@@ -13,11 +13,13 @@ readme = "README.md"
|
|
|
13
13
|
requires-python = ">=3.7"
|
|
14
14
|
classifiers = [
|
|
15
15
|
"Programming Language :: Python :: 3",
|
|
16
|
-
"Operating System ::
|
|
16
|
+
"Operating System :: Microsoft :: Windows",
|
|
17
|
+
"Operating System :: POSIX :: Linux"
|
|
17
18
|
]
|
|
18
19
|
|
|
19
|
-
|
|
20
|
-
packages
|
|
20
|
+
|
|
21
|
+
[tool.setuptools.packages.find]
|
|
22
|
+
where = ["src"]
|
|
21
23
|
|
|
22
24
|
[tool.setuptools.package-data]
|
|
23
25
|
"DeepGPR" = ["lib/*.cu", "lib/*.so", "lib/*.dll"]
|
|
@@ -14,7 +14,6 @@ if system_name == "Windows":
|
|
|
14
14
|
lib_filename = f'deepgpr{lib_extension}'
|
|
15
15
|
lib_path_obj = lib_dir / lib_filename
|
|
16
16
|
|
|
17
|
-
# Windows 下直接调用 nvcc 的命令参数
|
|
18
17
|
nvcc_cmd = [
|
|
19
18
|
'nvcc', '-shared',
|
|
20
19
|
'-o', str(lib_path_obj),
|
|
@@ -25,7 +24,6 @@ else:
|
|
|
25
24
|
lib_filename = f'deepgpr{lib_extension}'
|
|
26
25
|
lib_path_obj = lib_dir / lib_filename
|
|
27
26
|
|
|
28
|
-
# Linux 下直接调用 nvcc 的命令参数 (-fPIC 和 ABI 设置)
|
|
29
27
|
nvcc_cmd = [
|
|
30
28
|
'nvcc', '-shared', '-Xcompiler', '-fPIC',
|
|
31
29
|
'-D_GLIBCXX_USE_CXX11_ABI=0',
|
|
@@ -33,11 +31,9 @@ else:
|
|
|
33
31
|
str(cu_file)
|
|
34
32
|
]
|
|
35
33
|
|
|
36
|
-
# 2. 检查动态库是否存在,不存在则直接使用 Python 调用 nvcc 进行编译
|
|
37
34
|
if not lib_path_obj.is_file():
|
|
38
35
|
print(f'Compiling CUDA extension for {system_name} directly via nvcc...')
|
|
39
36
|
try:
|
|
40
|
-
# 直接执行 nvcc 命令,完全摆脱 make 依赖
|
|
41
37
|
subprocess.run(nvcc_cmd, check=True)
|
|
42
38
|
except FileNotFoundError:
|
|
43
39
|
raise RuntimeError(
|
|
@@ -47,11 +43,10 @@ if not lib_path_obj.is_file():
|
|
|
47
43
|
except subprocess.CalledProcessError as e:
|
|
48
44
|
raise RuntimeError(f"Compilation failed with error code {e.returncode}.")
|
|
49
45
|
|
|
50
|
-
# 3. 加载编译好的动态链接库 (变量名改为 c_lib,避免与 lib 文件夹冲突)
|
|
51
46
|
lib_path = str(lib_path_obj)
|
|
52
47
|
c_lib = ctypes.cdll.LoadLibrary(lib_path)
|
|
53
48
|
|
|
54
|
-
|
|
49
|
+
|
|
55
50
|
c_lib.forward.argtypes = [
|
|
56
51
|
ctypes.POINTER(ctypes.c_float), ctypes.POINTER(ctypes.c_float),
|
|
57
52
|
ctypes.POINTER(ctypes.c_float), ctypes.POINTER(ctypes.c_float),
|
|
@@ -92,7 +87,7 @@ c_lib.forward.argtypes = [
|
|
|
92
87
|
ctypes.POINTER(ctypes.c_int), ctypes.POINTER(ctypes.c_float),
|
|
93
88
|
ctypes.c_int
|
|
94
89
|
]
|
|
95
|
-
c_lib.forward.restype = None
|
|
90
|
+
c_lib.forward.restype = None
|
|
96
91
|
|
|
97
92
|
c_lib.backward.argtypes = [
|
|
98
93
|
ctypes.POINTER(ctypes.c_float), ctypes.POINTER(ctypes.c_float),
|
|
@@ -134,14 +129,11 @@ c_lib.backward.argtypes = [
|
|
|
134
129
|
ctypes.c_int ,ctypes.POINTER(ctypes.c_float), ctypes.POINTER(ctypes.c_float),
|
|
135
130
|
ctypes.c_int, ctypes.c_int
|
|
136
131
|
]
|
|
137
|
-
c_lib.backward.restype = None
|
|
138
|
-
|
|
132
|
+
c_lib.backward.restype = None
|
|
139
133
|
__all__ = ['c_lib']
|
|
140
134
|
|
|
141
|
-
|
|
142
|
-
# 第二步:在动态库环境就绪后再导入其余模块
|
|
143
|
-
# ==============================================================================
|
|
135
|
+
|
|
144
136
|
from .common import *
|
|
145
137
|
from .compute2 import *
|
|
146
138
|
from .multiscale import *
|
|
147
|
-
from .
|
|
139
|
+
from .wavelet import *
|
|
@@ -5,37 +5,7 @@ import math
|
|
|
5
5
|
from scipy.constants import c
|
|
6
6
|
from scipy.constants import mu_0 as m0
|
|
7
7
|
from scipy.constants import epsilon_0 as e0
|
|
8
|
-
from typing import Optional
|
|
9
|
-
|
|
10
|
-
def ricker(
|
|
11
|
-
freq: float,
|
|
12
|
-
length: int,
|
|
13
|
-
dt: float,
|
|
14
|
-
peak_time: float,
|
|
15
|
-
dtype: Optional[torch.dtype] = None,
|
|
16
|
-
) -> torch.Tensor:
|
|
17
|
-
"""Return a Ricker wavelet with the specified central frequency.
|
|
18
|
-
|
|
19
|
-
Args:
|
|
20
|
-
freq: The central frequency.
|
|
21
|
-
length: The number of time samples.
|
|
22
|
-
dt: The time sample spacing.
|
|
23
|
-
peak_time: The time (in secs) of the peak amplitude.
|
|
24
|
-
dtype: The PyTorch datatype to use. Optional, defaults to PyTorch's
|
|
25
|
-
default (float32).
|
|
26
|
-
Returns:
|
|
27
|
-
A PyTorch tensor representing the Ricker wavelet.
|
|
28
|
-
"""
|
|
29
|
-
if dt == 0:
|
|
30
|
-
raise ValueError("dt cannot be zero.")
|
|
31
8
|
|
|
32
|
-
t: torch.Tensor = torch.arange(float(length), dtype=dtype) * dt - peak_time
|
|
33
|
-
y: torch.Tensor = (1 - 2 * math.pi**2 * freq**2 * t**2) * torch.exp(
|
|
34
|
-
-(math.pi**2) * freq**2 * t**2,
|
|
35
|
-
)
|
|
36
|
-
if dtype is not None:
|
|
37
|
-
return y.to(dtype)
|
|
38
|
-
return y
|
|
39
9
|
|
|
40
10
|
def initialization(device, er,se,mr,source_amplitudes,source_location,receiver_location,dx,dt,pmlthick):
|
|
41
11
|
dtype=torch.float32
|
|
@@ -150,7 +120,7 @@ def check_cfl(dx, dt, nx,ny,nz):
|
|
|
150
120
|
|
|
151
121
|
|
|
152
122
|
def pmlthick_revert(p, er):
|
|
153
|
-
if isinstance(p, int):
|
|
123
|
+
if isinstance(p, int):
|
|
154
124
|
if er.shape[2] == 1:
|
|
155
125
|
return torch.tensor([p, p, p, p, 0, 0], dtype=torch.int32)
|
|
156
126
|
return torch.tensor([p]*6, dtype=torch.int32)
|
|
@@ -310,7 +280,7 @@ def check_tensors_for_nan_inf(d,**tensors):
|
|
|
310
280
|
|
|
311
281
|
if has_nan or has_inf:
|
|
312
282
|
found_issue = True
|
|
313
|
-
print(f"
|
|
283
|
+
print(f"[ERROR]{d}: Tensor `{name}` contains:", end=" ")
|
|
314
284
|
if has_nan:
|
|
315
285
|
print("NaN ", end="")
|
|
316
286
|
if has_inf:
|
|
@@ -594,50 +564,3 @@ def checkpoint_initial_field(device=None,per_nstep=None, dx=None, dt=None,
|
|
|
594
564
|
|
|
595
565
|
def zero_field(*tensors):
|
|
596
566
|
return tuple(torch.zeros_like(t) if t is not None else None for t in tensors)
|
|
597
|
-
|
|
598
|
-
def print_field_shapes(E, H, PML):
|
|
599
|
-
"""
|
|
600
|
-
打印 FDTD 场变量的形状
|
|
601
|
-
E: (Ex, Ey, Ez)
|
|
602
|
-
H: (Hx, Hy, Hz)
|
|
603
|
-
PML: 包含 24 个 Phi 张量的元组
|
|
604
|
-
"""
|
|
605
|
-
print("="*30)
|
|
606
|
-
print(" FIELD SHAPES CHECK ")
|
|
607
|
-
print("="*30)
|
|
608
|
-
|
|
609
|
-
# 1. 检查电场 E
|
|
610
|
-
e_names = ['Ex', 'Ey', 'Ez']
|
|
611
|
-
for name, tensor in zip(e_names, E):
|
|
612
|
-
shape = tensor.shape if torch.is_tensor(tensor) else "Not a Tensor"
|
|
613
|
-
print(f"{name:10} : {shape}")
|
|
614
|
-
|
|
615
|
-
print("-" * 20)
|
|
616
|
-
|
|
617
|
-
# 2. 检查磁场 H
|
|
618
|
-
h_names = ['Hx', 'Hy', 'Hz']
|
|
619
|
-
for name, tensor in zip(h_names, H):
|
|
620
|
-
shape = tensor.shape if torch.is_tensor(tensor) else "Not a Tensor"
|
|
621
|
-
print(f"{name:10} : {shape}")
|
|
622
|
-
|
|
623
|
-
print("-" * 20)
|
|
624
|
-
|
|
625
|
-
# 3. 检查 PML 场
|
|
626
|
-
pml_names = [
|
|
627
|
-
"x0EPhi1", "x0EPhi2", "x0HPhi1", "x0HPhi2",
|
|
628
|
-
"xmEPhi1", "xmEPhi2", "xmHPhi1", "xmHPhi2",
|
|
629
|
-
"y0EPhi1", "y0EPhi2", "y0HPhi1", "y0HPhi2",
|
|
630
|
-
"ymEPhi1", "ymEPhi2", "ymHPhi1", "ymHPhi2",
|
|
631
|
-
"z0EPhi1", "z0EPhi2", "z0HPhi1", "z0HPhi2",
|
|
632
|
-
"zmEPhi1", "zmEPhi2", "zmHPhi1", "zmHPhi2"
|
|
633
|
-
]
|
|
634
|
-
|
|
635
|
-
for i, name in enumerate(pml_names):
|
|
636
|
-
if i < len(PML):
|
|
637
|
-
tensor = PML[i]
|
|
638
|
-
shape = tensor.shape if torch.is_tensor(tensor) else "Not a Tensor"
|
|
639
|
-
print(f"{name:10} : {shape}")
|
|
640
|
-
else:
|
|
641
|
-
print(f"{name:10} : Missing in PML tuple")
|
|
642
|
-
|
|
643
|
-
print("="*30)
|
|
@@ -60,7 +60,6 @@ class DeepGPR(torch.autograd.Function):
|
|
|
60
60
|
|
|
61
61
|
nt_saved = (nt + model_gradient_sampling_interval - 1) // model_gradient_sampling_interval
|
|
62
62
|
|
|
63
|
-
# 仅仅控制显存分配策略,无需修改 C 端参数签名
|
|
64
63
|
if use_async_offload:
|
|
65
64
|
Eall = torch.zeros((nt_saved, nstep, nx, ny, nz), device='cpu', dtype=dtype).pin_memory()
|
|
66
65
|
else:
|
|
@@ -307,7 +306,6 @@ class DeepGPR(torch.autograd.Function):
|
|
|
307
306
|
ctx.Eall = None
|
|
308
307
|
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
308
|
|
|
310
|
-
# 返回与前向传播参数一一对应(最后补齐 use_async_offload 占位)
|
|
311
309
|
return (
|
|
312
310
|
grad_er, grad_se,
|
|
313
311
|
gEx,gEy,gEz, gHx,gHy,gHz,
|
|
@@ -17,9 +17,6 @@ __constant__ float m0 = 1.25663706212e-06;
|
|
|
17
17
|
}\
|
|
18
18
|
}
|
|
19
19
|
|
|
20
|
-
// ---------------------------------------------------------
|
|
21
|
-
// 系数获取核函数
|
|
22
|
-
// ---------------------------------------------------------
|
|
23
20
|
__global__ void ucgetforward(const float* __restrict__ er, const float* __restrict__ se, const float* __restrict__ mr,
|
|
24
21
|
float* __restrict__ uE0, float* __restrict__ uE1, float* __restrict__ uE4,
|
|
25
22
|
float* __restrict__ uH0, float* __restrict__ uH1, float* __restrict__ uH4,
|
|
@@ -85,9 +82,7 @@ __global__ void ucgetbackward(const float* __restrict__ er, const float* __restr
|
|
|
85
82
|
}
|
|
86
83
|
}
|
|
87
84
|
|
|
88
|
-
|
|
89
|
-
// 接收端保存
|
|
90
|
-
// ---------------------------------------------------------
|
|
85
|
+
|
|
91
86
|
__global__ void store_outputs(
|
|
92
87
|
int step, int NRX, int iteration,
|
|
93
88
|
const int* __restrict__ receiverlocation, float* __restrict__ rxs,
|
|
@@ -116,9 +111,7 @@ __global__ void store_outputs(
|
|
|
116
111
|
}
|
|
117
112
|
}
|
|
118
113
|
|
|
119
|
-
|
|
120
|
-
// 震源更新
|
|
121
|
-
// ---------------------------------------------------------
|
|
114
|
+
|
|
122
115
|
__global__ void Update_hertzian_dipole(
|
|
123
116
|
int step, int iteration, float dx,
|
|
124
117
|
const int* __restrict__ sourcelocation, const float* __restrict__ srcwaveforms,
|
|
@@ -146,9 +139,7 @@ __global__ void Update_hertzian_dipole(
|
|
|
146
139
|
}
|
|
147
140
|
}
|
|
148
141
|
|
|
149
|
-
|
|
150
|
-
// 融合:电场全局更新 + 全侧 PML 边界修正
|
|
151
|
-
// ---------------------------------------------------------
|
|
142
|
+
|
|
152
143
|
__global__ void fused_e_fields_updates_gpu(
|
|
153
144
|
const float* __restrict__ uE0, const float* __restrict__ uE1,
|
|
154
145
|
float* __restrict__ Ex, float* __restrict__ Ey, float* __restrict__ Ez,
|
|
@@ -317,9 +308,7 @@ __global__ void fused_e_fields_updates_gpu(
|
|
|
317
308
|
}
|
|
318
309
|
}
|
|
319
310
|
|
|
320
|
-
|
|
321
|
-
// 融合:磁场全局更新 + 全侧 PML 边界修正
|
|
322
|
-
// ---------------------------------------------------------
|
|
311
|
+
|
|
323
312
|
__global__ void fused_h_fields_updates_gpu(
|
|
324
313
|
const float* __restrict__ uH0, const float* __restrict__ uH1,
|
|
325
314
|
const float* __restrict__ Ex, const float* __restrict__ Ey, const float* __restrict__ Ez,
|
|
@@ -488,9 +477,7 @@ __global__ void fused_h_fields_updates_gpu(
|
|
|
488
477
|
}
|
|
489
478
|
}
|
|
490
479
|
|
|
491
|
-
|
|
492
|
-
// 反传:波场倒播
|
|
493
|
-
// ---------------------------------------------------------
|
|
480
|
+
|
|
494
481
|
__global__ void Back_source(
|
|
495
482
|
int step, int iteration, float dx,
|
|
496
483
|
const int* __restrict__ sourcelocation, const float* __restrict__ srcwaveforms,
|
|
@@ -519,9 +506,7 @@ __global__ void Back_source(
|
|
|
519
506
|
}
|
|
520
507
|
}
|
|
521
508
|
|
|
522
|
-
|
|
523
|
-
// 提取快照到缓冲区或全局内存
|
|
524
|
-
// ---------------------------------------------------------
|
|
509
|
+
|
|
525
510
|
__global__ void copy_to_Eall_single(
|
|
526
511
|
float* __restrict__ dst_ptr, int t_idx, const float* __restrict__ E,
|
|
527
512
|
int step, int NX, int NY, int NZ)
|
|
@@ -547,9 +532,7 @@ __global__ void copy_to_Eall_single(
|
|
|
547
532
|
}
|
|
548
533
|
}
|
|
549
534
|
|
|
550
|
-
|
|
551
|
-
// 融合:伴随状态法梯度更新(支持降采样波场及异步滑窗/同步显存自适应)
|
|
552
|
-
// ---------------------------------------------------------
|
|
535
|
+
|
|
553
536
|
__global__ void accumulate_gradients(
|
|
554
537
|
const float* __restrict__ Ez, const float* __restrict__ Eall_ptr, const float* __restrict__ d_E_buf,
|
|
555
538
|
float* __restrict__ grader, float* __restrict__ gradse,
|
|
@@ -569,7 +552,6 @@ __global__ void accumulate_gradients(
|
|
|
569
552
|
|
|
570
553
|
long long idx_Ez = ix * NY * NZ + iy * NZ + iz;
|
|
571
554
|
|
|
572
|
-
// 逻辑时刻索引
|
|
573
555
|
long long idx0_curr = i / S;
|
|
574
556
|
long long idx1_curr = min(idx0_curr + 1, (long long)nt_saved - 1);
|
|
575
557
|
float w1_curr = (float)(i % S) / S;
|
|
@@ -615,9 +597,7 @@ __global__ void accumulate_gradients(
|
|
|
615
597
|
if (serequiregrad == 1) atomicAdd(&gradse[idx], local_gradse);
|
|
616
598
|
}
|
|
617
599
|
|
|
618
|
-
|
|
619
|
-
// 主机 API
|
|
620
|
-
// ---------------------------------------------------------
|
|
600
|
+
|
|
621
601
|
extern "C" {
|
|
622
602
|
|
|
623
603
|
void forward(const float* __restrict__ er, const float* __restrict__ se, const float* __restrict__ mr,
|
|
@@ -713,7 +693,6 @@ void forward(const float* __restrict__ er, const float* __restrict__ se, const f
|
|
|
713
693
|
}
|
|
714
694
|
|
|
715
695
|
if (use_async) {
|
|
716
|
-
// 【核心修复】:必须先等流里的所有任务彻底做完,才能把底下的显存 free 掉!
|
|
717
696
|
cudaStreamSynchronize(stream_comp);
|
|
718
697
|
cudaStreamSynchronize(stream_trans);
|
|
719
698
|
cudaFree(d_E_buf);
|
|
@@ -829,7 +808,6 @@ void backward(const float* __restrict__ er, const float* __restrict__ se, const
|
|
|
829
808
|
}
|
|
830
809
|
|
|
831
810
|
if (use_async) {
|
|
832
|
-
// 【核心修复】:必须先同步!
|
|
833
811
|
cudaStreamSynchronize(stream_comp);
|
|
834
812
|
cudaStreamSynchronize(stream_trans);
|
|
835
813
|
cudaFree(d_E_buf);
|
|
@@ -2,28 +2,21 @@
|
|
|
2
2
|
import torch
|
|
3
3
|
|
|
4
4
|
def design_fir_filter(cutoff: float, fs: float, numtaps: int) -> torch.Tensor:
|
|
5
|
-
"""使用 Hamming 窗设计 FIR 滤波器"""
|
|
6
5
|
n = torch.arange(numtaps, dtype=torch.float32)
|
|
7
|
-
# 汉明窗
|
|
8
6
|
window = 0.54 - 0.46 * torch.cos(2 * torch.pi * n / (numtaps - 1))
|
|
9
|
-
# 正弦波
|
|
10
7
|
sinc = torch.sin(2 * torch.pi * (cutoff/fs) * (n - (numtaps-1)/2)) / (torch.pi * (n - (numtaps-1)/2))
|
|
11
|
-
# 处理中心点
|
|
12
8
|
center = (numtaps-1) // 2
|
|
13
9
|
sinc[center] = 2 * cutoff/fs
|
|
14
|
-
# 应用窗函数
|
|
15
10
|
h = window * sinc
|
|
16
|
-
|
|
11
|
+
|
|
17
12
|
return h / h.sum()
|
|
18
13
|
|
|
19
14
|
def apply_filter(data: torch.Tensor, fs: float, cutoff: float) -> torch.Tensor:
|
|
20
|
-
"""应用 FIR 滤波器到数据"""
|
|
21
15
|
numtaps = int(1 * (fs / cutoff))
|
|
22
16
|
fir_coeff = design_fir_filter(cutoff, fs, numtaps)
|
|
23
17
|
fir_coeff = fir_coeff.to(data.device)
|
|
24
18
|
|
|
25
19
|
if data.ndim == 1:
|
|
26
|
-
# 1D 数据处理 - 添加维度以支持反射填充
|
|
27
20
|
data_2d = data.view(1, 1, -1)
|
|
28
21
|
padded_data = torch.nn.functional.pad(data_2d, (numtaps-1, 0), mode='reflect')
|
|
29
22
|
filtered = torch.nn.functional.conv1d(
|
|
@@ -34,19 +27,14 @@ def apply_filter(data: torch.Tensor, fs: float, cutoff: float) -> torch.Tensor:
|
|
|
34
27
|
return filtered.view(-1)
|
|
35
28
|
|
|
36
29
|
elif data.ndim == 3:
|
|
37
|
-
# 3D 数据处理
|
|
38
30
|
step, iterations, nrx = data.shape
|
|
39
|
-
# 重塑数据以使用批量处理
|
|
40
31
|
reshaped_data = data.permute(0, 2, 1).reshape(-1, 1, iterations)
|
|
41
|
-
# 填充数据
|
|
42
32
|
padded_data = torch.nn.functional.pad(reshaped_data, (numtaps-1, 0), mode='reflect')
|
|
43
|
-
# 应用卷积
|
|
44
33
|
filtered = torch.nn.functional.conv1d(
|
|
45
34
|
padded_data,
|
|
46
35
|
fir_coeff.view(1, 1, -1),
|
|
47
36
|
padding=0
|
|
48
37
|
)
|
|
49
|
-
# 重塑回原始形状
|
|
50
38
|
return filtered.view(step, nrx, iterations).permute(0, 2, 1)
|
|
51
39
|
|
|
52
40
|
else:
|
|
@@ -0,0 +1,34 @@
|
|
|
1
|
+
from typing import Optional
|
|
2
|
+
import torch
|
|
3
|
+
import math
|
|
4
|
+
|
|
5
|
+
|
|
6
|
+
def ricker(
|
|
7
|
+
freq: float,
|
|
8
|
+
length: int,
|
|
9
|
+
dt: float,
|
|
10
|
+
peak_time: float,
|
|
11
|
+
dtype: Optional[torch.dtype] = None,
|
|
12
|
+
) -> torch.Tensor:
|
|
13
|
+
"""Return a Ricker wavelet with the specified central frequency.
|
|
14
|
+
|
|
15
|
+
Args:
|
|
16
|
+
freq: The central frequency.
|
|
17
|
+
length: The number of time samples.
|
|
18
|
+
dt: The time sample spacing.
|
|
19
|
+
peak_time: The time (in secs) of the peak amplitude.
|
|
20
|
+
dtype: The PyTorch datatype to use. Optional, defaults to PyTorch's
|
|
21
|
+
default (float32).
|
|
22
|
+
Returns:
|
|
23
|
+
A PyTorch tensor representing the Ricker wavelet.
|
|
24
|
+
"""
|
|
25
|
+
if dt == 0:
|
|
26
|
+
raise ValueError("dt cannot be zero.")
|
|
27
|
+
|
|
28
|
+
t: torch.Tensor = torch.arange(float(length), dtype=dtype) * dt - peak_time
|
|
29
|
+
y: torch.Tensor = (1 - 2 * math.pi**2 * freq**2 * t**2) * torch.exp(
|
|
30
|
+
-(math.pi**2) * freq**2 * t**2,
|
|
31
|
+
)
|
|
32
|
+
if dtype is not None:
|
|
33
|
+
return y.to(dtype)
|
|
34
|
+
return y
|
|
@@ -1,10 +1,11 @@
|
|
|
1
1
|
Metadata-Version: 2.4
|
|
2
2
|
Name: DeepGPR
|
|
3
|
-
Version: 0.0.
|
|
3
|
+
Version: 0.0.5
|
|
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
|
|
7
|
-
Classifier: Operating System ::
|
|
7
|
+
Classifier: Operating System :: Microsoft :: Windows
|
|
8
|
+
Classifier: Operating System :: POSIX :: Linux
|
|
8
9
|
Requires-Python: >=3.7
|
|
9
10
|
Description-Content-Type: text/markdown
|
|
10
11
|
|
|
@@ -59,7 +60,6 @@ peak_time = 1 / freq
|
|
|
59
60
|
source_amplitudes = torch.zeros((1,nt,1),device=device)
|
|
60
61
|
source_amplitudes[0,:,0]=DeepGPR.ricker(freq, nt, dt, peak_time).to(device)
|
|
61
62
|
|
|
62
|
-
DeepGPR.plot_survey_geometry(er,source_location, receiver_location,dx)
|
|
63
63
|
|
|
64
64
|
#forward modeling
|
|
65
65
|
r = DeepGPR.compute(
|
|
@@ -79,8 +79,7 @@ ax[1].imshow(er.grad.detach())
|
|
|
79
79
|
ax[1].set_title("Gradient")
|
|
80
80
|
plt.show()
|
|
81
81
|
```
|
|
82
|
-

|
|
83
82
|
|
|
84
|
-

|
|
85
84
|
|
|
86
85
|
There are more examples in the ./examples.
|
|
@@ -0,0 +1,13 @@
|
|
|
1
|
+
README.md
|
|
2
|
+
pyproject.toml
|
|
3
|
+
src/DeepGPR/__init__.py
|
|
4
|
+
src/DeepGPR/common.py
|
|
5
|
+
src/DeepGPR/compute2.py
|
|
6
|
+
src/DeepGPR/multiscale.py
|
|
7
|
+
src/DeepGPR/wavelet.py
|
|
8
|
+
src/DeepGPR.egg-info/PKG-INFO
|
|
9
|
+
src/DeepGPR.egg-info/SOURCES.txt
|
|
10
|
+
src/DeepGPR.egg-info/dependency_links.txt
|
|
11
|
+
src/DeepGPR.egg-info/top_level.txt
|
|
12
|
+
src/DeepGPR/lib/deepgpr.cu
|
|
13
|
+
src/DeepGPR/lib/deepgpr.so
|
deepgpr-0.0.2/DeepGPR/visual.py
DELETED
|
@@ -1,105 +0,0 @@
|
|
|
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()
|
|
@@ -1,13 +0,0 @@
|
|
|
1
|
-
README.md
|
|
2
|
-
pyproject.toml
|
|
3
|
-
DeepGPR/__init__.py
|
|
4
|
-
DeepGPR/common.py
|
|
5
|
-
DeepGPR/compute2.py
|
|
6
|
-
DeepGPR/multiscale.py
|
|
7
|
-
DeepGPR/visual.py
|
|
8
|
-
DeepGPR.egg-info/PKG-INFO
|
|
9
|
-
DeepGPR.egg-info/SOURCES.txt
|
|
10
|
-
DeepGPR.egg-info/dependency_links.txt
|
|
11
|
-
DeepGPR.egg-info/top_level.txt
|
|
12
|
-
DeepGPR/lib/deepgpr.cu
|
|
13
|
-
DeepGPR/lib/deepgpr.so
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|