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.
@@ -1,10 +1,11 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: DeepGPR
3
- Version: 0.0.2
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 :: OS Independent
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
- ![result](./Fig/example1.png)
83
82
 
84
- ![result](./Fig/example2.png)
83
+ ![result](./Fig/example.png)
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
- ![result](./Fig/example1.png)
73
71
 
74
- ![result](./Fig/example2.png)
72
+ ![result](./Fig/example.png)
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.2"
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 :: OS Independent"
16
+ "Operating System :: Microsoft :: Windows",
17
+ "Operating System :: POSIX :: Linux"
17
18
  ]
18
19
 
19
- [tool.setuptools]
20
- packages = ["DeepGPR"]
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
- # 4. 定义 C 函数的参数和返回值类型
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 .visual import *
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): # 如果是 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"❌ [ERROR]{d}: Tensor `{name}` contains:", end=" ")
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.2
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 :: OS Independent
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
- ![result](./Fig/example1.png)
83
82
 
84
- ![result](./Fig/example2.png)
83
+ ![result](./Fig/example.png)
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
@@ -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