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/common.py ADDED
@@ -0,0 +1,643 @@
1
+ import torch
2
+ import torch.nn as nn
3
+ import torch.nn.functional as F
4
+ import math
5
+ from scipy.constants import c
6
+ from scipy.constants import mu_0 as m0
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
+
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
+
40
+ def initialization(device, er,se,mr,source_amplitudes,source_location,receiver_location,dx,dt,pmlthick):
41
+ dtype=torch.float32
42
+ if er.min()<1 :
43
+ raise ValueError('The values of epsilon is incorrect.(should be greater than 1)')
44
+ if se.min()<0:
45
+ raise ValueError('The values of sigma is incorrect.(should be non-negative)')
46
+
47
+ if len(er.shape) == 2:
48
+ er = er.reshape(*er.shape, 1)
49
+ elif len(er.shape) != 3:
50
+ raise ValueError('The shape of epsilon should be 2-d or 3-d.')
51
+
52
+ if len(se.shape) == 2:
53
+ se = se.reshape(*se.shape, 1)
54
+ elif len(se.shape) != 3:
55
+ raise ValueError('The shape of epsilon should be 2-d or 3-d.')
56
+
57
+ if er.shape == se.shape:
58
+ nx=er.shape[0]
59
+ ny=er.shape[1]
60
+ nz=er.shape[2]
61
+ if nz==1:
62
+ mode=2
63
+ else:
64
+ mode=3
65
+ er=er.to(device)
66
+ se=se.to(device)
67
+ if mr is None:
68
+ mr=torch.ones_like(er, device=device)
69
+ else:
70
+ if mr.shape == er.shape:
71
+ mr=mr.to(device)
72
+ else:
73
+ raise ValueError('The shape of miu should be the same as epsilon and sigma.')
74
+ else:
75
+ raise ValueError('The shape of epsilon and sigma should be the same.')
76
+
77
+ if source_location.shape[0] == receiver_location.shape[0]:
78
+ source_location=source_location.to(torch.int)
79
+ receiver_location=receiver_location.to(torch.int)
80
+
81
+ source_check = (source_location >= 0).all()
82
+ receiver_check = (receiver_location >= 0).all()
83
+
84
+ source_check &= (source_location[..., 0] < nx).all()
85
+ source_check &= (source_location[..., 1] < ny).all()
86
+ source_check &= (source_location[..., 2] < nz).all()
87
+
88
+ if not (source_check):
89
+ raise ValueError(
90
+ "Error: Source coordinates out of range! "
91
+ f"Valid ranges are x∈[0,{nx}), y∈[0,{ny}), z∈[0,{nz})"
92
+ )
93
+
94
+ receiver_check &= (receiver_location[..., 0] < nx).all()
95
+ receiver_check &= (receiver_location[..., 1] < ny).all()
96
+ receiver_check &= (receiver_location[..., 2] < nz).all()
97
+
98
+ if not (receiver_check):
99
+ raise ValueError(
100
+ "Error: Receiver coordinates out of range! "
101
+ f"Valid ranges are x∈[0,{nx}), y∈[0,{ny}), z∈[0,{nz})"
102
+ )
103
+ nstep=source_location.shape[0]
104
+
105
+ nsr=source_location.shape[1]
106
+ nrx=receiver_location.shape[1]
107
+
108
+ source_location=source_location.to(device)
109
+ receiver_location=receiver_location.to(device)
110
+ else:
111
+ raise ValueError('The first dimension (nstep) of source_location and receiver_location should be the same.')
112
+
113
+ source_amplitudes=source_amplitudes.to(device).contiguous()
114
+
115
+ if (source_amplitudes.shape[0]>1 and source_amplitudes.shape[0]<nsr) or source_amplitudes.shape[0]>nsr :
116
+ raise ValueError('The number of source waveforms is incorrect.')
117
+ elif source_amplitudes.shape[0]==1 and nsr!=1:
118
+ source_amplitudes=source_amplitudes.repeat(nsr,1,1).contiguous()
119
+ print('Tips: The number of source waveforms is 1, but the number of sources is ',nsr,'. The source waveform is repeated for all sources.')
120
+
121
+ check_cfl(dx, dt,nx,ny,nz)
122
+
123
+ nt=source_amplitudes.shape[1]
124
+
125
+
126
+ pmlthick=pmlthick_revert(pmlthick,er)
127
+ ere=F.pad(er, (0, 1, 0, 1, 0, 1)).to(dtype)
128
+ see=F.pad(se, (0, 1, 0, 1, 0, 1)).to(dtype)
129
+ mr=F.pad(mr, (0, 1, 0, 1, 0, 1)).to(dtype)
130
+
131
+ return er,se,nx,ny,nz,nt,nstep,nsr,nrx,ere,see,mr,mode,dtype,pmlthick,source_amplitudes
132
+
133
+
134
+ def check_cfl(dx, dt, nx,ny,nz):
135
+
136
+ dy=dx
137
+ dz=dx
138
+
139
+ if nz==1:
140
+ dt_max = 1.0 / (c * math.sqrt(1/dx**2 + 1/dy**2))
141
+ elif nx==1:
142
+ dt_max = 1.0 / (c * math.sqrt(1/dy**2 + 1/dz**2))
143
+ elif ny==1:
144
+ dt_max = 1.0 / (c * math.sqrt(1/dx**2 + 1/dz**2))
145
+ else:
146
+ dt_max = 1.0 / (c * math.sqrt(1/dx**2 + 1/dy**2 + 1/dz**2))
147
+
148
+ if dt > dt_max:
149
+ raise ValueError(f"Does not meet CFL conditions: dt={dt:.3e} > dt_max={dt_max:.3e}")
150
+
151
+
152
+ def pmlthick_revert(p, er):
153
+ if isinstance(p, int): # 如果是 int
154
+ if er.shape[2] == 1:
155
+ return torch.tensor([p, p, p, p, 0, 0], dtype=torch.int32)
156
+ return torch.tensor([p]*6, dtype=torch.int32)
157
+
158
+ elif isinstance(p, list):
159
+ if len(p) == 6:
160
+ return torch.tensor(p, dtype=torch.int32)
161
+ elif len(p) == 4:
162
+ return torch.tensor(p + [0, 0], dtype=torch.int32)
163
+ else:
164
+ raise ValueError(f"Unsupported list length: {len(p)}. Must be 4 or 6.")
165
+
166
+ elif isinstance(p, torch.Tensor):
167
+ return p.to(dtype=torch.int32)
168
+
169
+ else:
170
+ raise TypeError(f"Unsupported type: {type(p)}")
171
+
172
+
173
+ class TVRegularization(nn.Module):
174
+ # tv_criterion = TVRegularization(weight_ep=1,weight_sigma=0.1).to(device)
175
+ def __init__(self, weight_ep=1, weight_sigma=0.001, method='anisotropic'):
176
+ super(TVRegularization, self).__init__()
177
+ self.weight_ep = weight_ep
178
+ self.weight_sigma = weight_sigma
179
+ self.method = method
180
+
181
+ def _compute_tv(self, data):
182
+ """
183
+ 自动适应 2D [H, W] 或 3D [D, H, W] / [H, W, C] 的 TV 计算
184
+ """
185
+ # 1. 维度清洗:如果是 [H, W, 1],先去掉那个 1
186
+ if data.dim() == 3 and data.shape[-1] == 1:
187
+ data = data.squeeze(-1)
188
+
189
+ # 2. 根据维度分情况计算梯度
190
+ if data.dim() == 2:
191
+ # --- 2D 情况 [H, W] ---
192
+ # X方向差分 (行与行之间)
193
+ d_x = data[1:, :] - data[:-1, :]
194
+ # Y方向差分 (列与列之间)
195
+ d_y = data[:, 1:] - data[:, :-1]
196
+
197
+ # 2D 肯定没有 Z 方向
198
+ loss_x = torch.sum(torch.abs(d_x))
199
+ loss_y = torch.sum(torch.abs(d_y))
200
+ loss_z = 0.0
201
+
202
+ elif data.dim() == 3:
203
+ # --- 3D 情况 [D, H, W] ---
204
+ # 假设第0维是深度/X,第1维是高度/Y,第2维是宽度/Z
205
+ d_x = data[1:, :, :] - data[:-1, :, :]
206
+ d_y = data[:, 1:, :] - data[:, :-1, :]
207
+ d_z = data[:, :, 1:] - data[:, :, :-1]
208
+
209
+ loss_x = torch.sum(torch.abs(d_x))
210
+ loss_y = torch.sum(torch.abs(d_y))
211
+ loss_z = torch.sum(torch.abs(d_z))
212
+
213
+ elif data.dim() == 4:
214
+ # --- 4D 情况 [Batch, Channel, H, W] ---
215
+ # 常见于图像处理习惯
216
+ d_x = data[..., 1:, :] - data[..., :-1, :]
217
+ d_y = data[..., :, 1:] - data[..., :, :-1]
218
+
219
+ loss_x = torch.sum(torch.abs(d_x))
220
+ loss_y = torch.sum(torch.abs(d_y))
221
+ loss_z = 0.0
222
+
223
+ else:
224
+ # 避免标量或1D数据报错
225
+ return torch.tensor(0.0, device=data.device)
226
+
227
+ # 3. 汇总 Loss
228
+ if self.method == 'anisotropic':
229
+ # 各向异性:直接相加 (|dx| + |dy|)
230
+ total_tv = loss_x + loss_y + loss_z
231
+ else:
232
+ # 各向同性:平方和开根 (sqrt(dx^2 + dy^2)) - 简化版近似
233
+ # 注意:严谨的各向同性需要对每个像素点求平方和再sum,这里简化处理以保持计算图简单
234
+ total_tv = (loss_x**2 + loss_y**2 + loss_z**2 + 1e-8).sqrt()
235
+
236
+ return total_tv
237
+
238
+ def forward(self, ep=None, sigma=None):
239
+ loss = torch.tensor(0.0, device=ep.device if ep is not None else sigma.device)
240
+
241
+ # 计算 ep (介电常数) 的 TV
242
+ if ep is not None and self.weight_ep > 0:
243
+ loss += self.weight_ep * self._compute_tv(ep)
244
+
245
+ # 计算 sigma (电导率) 的 TV
246
+ if sigma is not None and self.weight_sigma > 0:
247
+ loss += self.weight_sigma * self._compute_tv(sigma)
248
+
249
+ return loss
250
+
251
+
252
+ def create_or_separate(tuple:tuple, nx,ny,nz,nstep,device: torch.device,
253
+ dtype: torch.dtype):
254
+ if tuple == None:
255
+ return torch.zeros((nstep,nx+1,ny+1,nz+1), device=device, dtype=dtype).contiguous(),torch.zeros((nstep,nx+1,ny+1,nz+1), device=device, dtype=dtype).contiguous(),torch.zeros((nstep,nx+1,ny+1,nz+1), device=device, dtype=dtype).contiguous()
256
+ # else:
257
+ # if tensor[0].shape[1]==nx+1 and tensor[0].shape[2]==ny+1 and tensor[0].shape[3]==nz+1 and tensor[0].shape[0]==nstep:
258
+ # return tensor[0].contiguous(),tensor[1].contiguous(),tensor[2].contiguous()
259
+ # else:
260
+ # print(nstep,nx,ny,nz)
261
+ # raise ValueError('The shape of E and H should be (nstep,nx+1,ny+1,nz+1).')
262
+ else:
263
+ condition = (
264
+ tuple[0].shape[0] == nstep and
265
+ tuple[0].shape[1] == nx + 1 and
266
+ tuple[0].shape[2] == ny + 1
267
+ )
268
+
269
+ if tuple[0].ndim > 3:
270
+ condition = condition and (tuple[0].shape[3] == nz + 1)
271
+
272
+ if condition:
273
+ return (
274
+ tuple[0].contiguous(),
275
+ tuple[1].contiguous(),
276
+ tuple[2].contiguous()
277
+ )
278
+ else:
279
+ actual_shape = list(tuple[0].shape)
280
+ expected_min_shape = [nstep, nx + 1, ny + 1, nz + 1]
281
+ raise ValueError(
282
+ f"Shape not match! \n"
283
+ f"Actual shape: {actual_shape} \n"
284
+ f"Expected shape (at least): {expected_min_shape[:3]} and (if nz!=1) {nz+1}"
285
+ )
286
+
287
+
288
+
289
+ def check_tensors_for_nan_inf(d,**tensors):
290
+ """
291
+ check_tensors_for_nan_inf(
292
+ gEx=gEx, gEy=gEy, gEz=gEz,
293
+ gHx=gHx, gHy=gHy, gHz=gHz,
294
+ gx0EPhi1=gx0EPhi1, ...
295
+ )
296
+ """
297
+ found_issue = False
298
+
299
+ for name, tensor in tensors.items():
300
+ if tensor is None:
301
+ print(f"[WARNING]{d}: {name} is None.")
302
+ continue
303
+
304
+ if not isinstance(tensor, torch.Tensor):
305
+ print(f"[WARNING]{d}: {name} is not a tensor: {type(tensor)}")
306
+ continue
307
+
308
+ has_nan = torch.isnan(tensor).any().item()
309
+ has_inf = torch.isinf(tensor).any().item()
310
+
311
+ if has_nan or has_inf:
312
+ found_issue = True
313
+ print(f"❌ [ERROR]{d}: Tensor `{name}` contains:", end=" ")
314
+ if has_nan:
315
+ print("NaN ", end="")
316
+ if has_inf:
317
+ print("Inf ", end="")
318
+ print(f"| shape={tuple(tensor.shape)} | dtype={tensor.dtype}")
319
+
320
+
321
+ def buildpmlcoeffs(er,mr,dt,dx,nx,ny,nz,pmlthick,device,dtype):
322
+ averageer=torch.zeros(6, device=device, dtype=dtype)
323
+ averagemr=torch.zeros(6, device=device, dtype=dtype)
324
+ lencfs=1
325
+ x0 = torch.empty(0)
326
+ xm = torch.empty(0)
327
+ y0 = torch.empty(0)
328
+ ym = torch.empty(0)
329
+ z0 = torch.empty(0)
330
+ zm = torch.empty(0)
331
+ x01 = torch.empty(0)
332
+ x02 = torch.empty(0)
333
+ xm1 = torch.empty(0)
334
+ xm2 = torch.empty(0)
335
+ y01 = torch.empty(0)
336
+ y02 = torch.empty(0)
337
+ ym1 = torch.empty(0)
338
+ ym2 = torch.empty(0)
339
+ z01 = torch.empty(0)
340
+ z02 = torch.empty(0)
341
+ zm1 = torch.empty(0)
342
+ zm2 = torch.empty(0)
343
+ if pmlthick[0]>0:
344
+ x0=torch.tensor((pmlthick[0],0,pmlthick[0],0,ny,0,nz), device=device, dtype=torch.int)
345
+ averageer[0]=er[x0[1],:ny,:nz].mean()
346
+ averagemr[0]=mr[x0[1],:ny,:nz].mean()
347
+ CFS0=CFS(device=device)
348
+ x01=torch.zeros((4,lencfs,pmlthick[0]), device=device, dtype=dtype)
349
+ x02=torch.zeros((4,lencfs,pmlthick[0]), device=device, dtype=dtype)
350
+ calculate_pml_update_coeffs(CFS0,x01,x02, averageer[0], averagemr[0], dt,dx,pmlthick[0])
351
+
352
+ if pmlthick[1]>0:
353
+ xm=torch.tensor((pmlthick[1],nx-pmlthick[1],nx,0,ny,0,nz), device=device, dtype=torch.int)
354
+ averageer[1]=er[xm[1],:ny,:nz].mean()
355
+ averagemr[1]=mr[xm[1],:ny,:nz].mean()
356
+ CFS1=CFS(device=device)
357
+ xm1=torch.zeros((4,lencfs,pmlthick[1]), device=device, dtype=dtype)
358
+ xm2=torch.zeros((4,lencfs,pmlthick[1]), device=device, dtype=dtype)
359
+ calculate_pml_update_coeffs(CFS1,xm1,xm2, averageer[1], averagemr[1], dt,dx,pmlthick[1])
360
+
361
+ if pmlthick[2]>0:
362
+ y0=torch.tensor((pmlthick[2],0,nx,0,pmlthick[2],0,nz), device=device, dtype=torch.int)
363
+ averageer[2]=er[:nx,y0[3],:nz].mean()
364
+ averagemr[2]=mr[:nx,y0[3],:nz].mean()
365
+ CFS2=CFS(device=device)
366
+ y01=torch.zeros((4,lencfs,pmlthick[2]), device=device, dtype=dtype)
367
+ y02=torch.zeros((4,lencfs,pmlthick[2]), device=device, dtype=dtype)
368
+ calculate_pml_update_coeffs(CFS2,y01,y02, averageer[2], averagemr[2], dt,dx,pmlthick[2])
369
+
370
+ if pmlthick[3]>0:
371
+ ym=torch.tensor((pmlthick[3],0,nx,ny-pmlthick[3],ny,0,nz), device=device, dtype=torch.int)
372
+ averageer[3]=er[:nx,ym[3],:nz].mean()
373
+ averagemr[3]=mr[:nx,ym[3],:nz].mean()
374
+ CFS3=CFS(device=device)
375
+ ym1=torch.zeros((4,lencfs,pmlthick[3]), device=device, dtype=dtype)
376
+ ym2=torch.zeros((4,lencfs,pmlthick[3]), device=device, dtype=dtype)
377
+ calculate_pml_update_coeffs(CFS3,ym1,ym2, averageer[3], averagemr[3], dt,dx,pmlthick[3])
378
+
379
+ if pmlthick[4]>0:
380
+ z0=torch.tensor((pmlthick[4],0,nx,0,ny,0,pmlthick[4]), device=device, dtype=torch.int)
381
+ averageer[4]=er[:nx,:ny,z0[5]].mean()
382
+ averagemr[4]=mr[:nx,:ny,z0[5]].mean()
383
+ CFS4=CFS(device=device)
384
+ z01=torch.zeros((4,lencfs,pmlthick[4]), device=device, dtype=dtype)
385
+ z02=torch.zeros((4,lencfs,pmlthick[4]), device=device, dtype=dtype)
386
+ calculate_pml_update_coeffs(CFS4,z01,z02, averageer[4], averagemr[4], dt,dx,pmlthick[4])
387
+
388
+ if pmlthick[5]>0:
389
+ zm=torch.tensor((pmlthick[5],0,nx,0,ny,nz-pmlthick[5],nz), device=device, dtype=torch.int)
390
+ averageer[5]=er[:nx,:ny,zm[5]].mean()
391
+ averagemr[5]=mr[:nx,:ny,zm[5]].mean()
392
+ CFS5=CFS(device=device)
393
+ zm1=torch.zeros((4,lencfs,pmlthick[5]), device=device, dtype=dtype)
394
+ zm2=torch.zeros((4,lencfs,pmlthick[5]), device=device, dtype=dtype)
395
+ calculate_pml_update_coeffs(CFS5,zm1,zm2, averageer[5], averagemr[5], dt,dx,pmlthick[5])
396
+ return x0,xm,y0,ym,z0,zm,x01,x02,xm1,xm2,y01,y02,ym1,ym2,z01,z02,zm1,zm2
397
+
398
+
399
+
400
+ class CFSParameter(object):
401
+ scalingprofiles = {'constant': 0, 'linear': 1, 'quadratic': 2, 'cubic': 3, 'quartic': 4, 'quintic': 5, 'sextic': 6, 'septic': 7, 'octic': 8}
402
+
403
+ def __init__(self,ID =None, scaling='polynomial', scalingprofile=None, min=0, max=0):
404
+ self.ID = ID
405
+ self.scaling = scaling
406
+ self.scalingprofile = scalingprofile
407
+ self.min = min
408
+ self.max = max
409
+
410
+
411
+ class CFS(object):
412
+
413
+ def __init__(self, device):
414
+ self.alpha = CFSParameter(ID='alpha', scalingprofile='constant')
415
+ self.kappa = CFSParameter(ID='kappa', scalingprofile='constant', min=1, max=1)
416
+ self.sigma = CFSParameter(ID='sigma', scalingprofile='quartic', min=0, max=None)
417
+ self.device = device
418
+
419
+ def calculate_sigmamax(self, d, er, mr):
420
+ with torch.no_grad():
421
+ m = CFSParameter.scalingprofiles[self.sigma.scalingprofile]
422
+ self.sigma.max = (0.8 * (m + 1)) / (((m0 / e0) ** 0.5) * d * torch.sqrt(er * mr))
423
+
424
+
425
+ def scaling_polynomial(self, order, Evalues, Hvalues):
426
+ tmp = (torch.linspace(0, (len(Evalues) - 1) + 0.5, steps=2 * len(Evalues)) / (len(Evalues) - 1)) ** order
427
+ Evalues = tmp[0:-1:2].to(self.device)
428
+ Hvalues = tmp[1::2].to(self.device)
429
+ return Evalues, Hvalues
430
+
431
+ def calculate_values(self, thickness, parameter):
432
+
433
+ Evalues = torch.zeros(thickness + 1, device=self.device)
434
+ Hvalues = torch.zeros(thickness + 1, device=self.device)
435
+ if parameter.scalingprofile == 'constant':
436
+ Evalues += parameter.max
437
+ Hvalues += parameter.max
438
+ elif parameter.scaling == 'polynomial':
439
+ Evalues, Hvalues = self.scaling_polynomial(CFSParameter.scalingprofiles[parameter.scalingprofile], Evalues, Hvalues)
440
+ if parameter.ID == 'alpha':
441
+ Evalues = Evalues * (self.alpha.max - self.alpha.min) + self.alpha.min
442
+ Hvalues = Hvalues * (self.alpha.max - self.alpha.min) + self.alpha.min
443
+ elif parameter.ID == 'kappa':
444
+ Evalues = Evalues * (self.kappa.max - self.kappa.min) + self.kappa.min
445
+ Hvalues = Hvalues * (self.kappa.max - self.kappa.min) + self.kappa.min
446
+ elif parameter.ID == 'sigma':
447
+ Evalues = Evalues * (self.sigma.max - self.sigma.min) + self.sigma.min
448
+ Hvalues = Hvalues * (self.sigma.max - self.sigma.min) + self.sigma.min
449
+
450
+ Evalues = Evalues[:-1]
451
+ Hvalues = Hvalues[:-1]
452
+
453
+ return Evalues, Hvalues
454
+
455
+ def calculate_pml_update_coeffs(cfs,R1,R2, aver, avmr, dt,d,thickness):
456
+ if not cfs.sigma.max:
457
+ cfs.calculate_sigmamax(d, aver, avmr)
458
+
459
+ Ealpha, Halpha = cfs.calculate_values(thickness, cfs.alpha)
460
+ Ekappa, Hkappa = cfs.calculate_values(thickness, cfs.kappa)
461
+ Esigma, Hsigma = cfs.calculate_values(thickness, cfs.sigma)
462
+
463
+ R1=R1.contiguous()
464
+ R2=R2.contiguous()
465
+
466
+ tmp = (2 * e0 * Ekappa) + dt * (Ealpha * Ekappa + Esigma)
467
+ R1[0,0, :] = (2 * e0 + dt * Ealpha) / tmp
468
+ R1[1,0, :] = (2 * e0 * Ekappa) / tmp
469
+ R1[2,0, :] = ((2 * e0 * Ekappa) - dt * (Ealpha * Ekappa + Esigma)) / tmp
470
+ R1[3,0, :] = (2 * Esigma * dt) / (Ekappa * tmp)
471
+
472
+ tmp = (2 * e0 * Hkappa) + dt * (Halpha * Hkappa + Hsigma)
473
+ R2[0,0, :] = (2 * e0 + dt * Halpha) / tmp
474
+ R2[1,0, :] = (2 * e0 * Hkappa) / tmp
475
+ R2[2,0, :] = ((2 * e0 * Hkappa) - dt * (Halpha * Hkappa + Hsigma)) / tmp
476
+ R2[3,0, :] = (2 * Hsigma * dt) / (Hkappa * tmp)
477
+
478
+
479
+
480
+ def build_pml_phi(x0,xm,y0,ym,z0,zm,nstep,PML,device):
481
+
482
+ (x0EPhi1, x0EPhi2, x0HPhi1, x0HPhi2,
483
+ xmEPhi1, xmEPhi2, xmHPhi1, xmHPhi2,
484
+ y0EPhi1, y0EPhi2, y0HPhi1, y0HPhi2,
485
+ ymEPhi1, ymEPhi2, ymHPhi1, ymHPhi2,
486
+ z0EPhi1, z0EPhi2, z0HPhi1, z0HPhi2,
487
+ zmEPhi1, zmEPhi2, zmHPhi1, zmHPhi2) = [torch.empty(0) for _ in range(24)]
488
+
489
+ if PML==None:
490
+ PML=(None,None,None,None,None,None,None,None,None,None,None,None,None,None,None,None,None,None,None,None,None,None,None,None)
491
+
492
+ if x0.numel()!=0 and PML[0]==None:
493
+ x0EPhi1=torch.zeros((nstep, int(x0[2]-x0[1]+1), int(x0[4]-x0[3]), int(x0[6]-x0[5]+1)),dtype=torch.float, device=device)
494
+ x0EPhi2=torch.zeros((nstep, int(x0[2]-x0[1]+1), int(x0[4]-x0[3]+1), int(x0[6]-x0[5])),dtype=torch.float, device=device)
495
+ x0HPhi1=torch.zeros((nstep, int(x0[2]-x0[1]), int(x0[4]-x0[3]+1), int(x0[6]-x0[5])), dtype=torch.float, device=device)
496
+ x0HPhi2=torch.zeros((nstep, int(x0[2]-x0[1]), int(x0[4]-x0[3]), int(x0[6]-x0[5]+1)), dtype=torch.float, device=device)
497
+ elif PML[0]!=None:
498
+ x0EPhi1=PML[0].contiguous()
499
+ x0EPhi2=PML[1].contiguous()
500
+ x0HPhi1=PML[2].contiguous()
501
+ x0HPhi2=PML[3].contiguous()
502
+
503
+ if xm.numel()!=0 and PML[4]==None:
504
+ xmEPhi1=torch.zeros((nstep, xm[2]-xm[1]+1, xm[4]-xm[3], xm[6]-xm[5]+1), dtype=torch.float, device=device)
505
+ xmEPhi2=torch.zeros((nstep, xm[2]-xm[1]+1, xm[4]-xm[3]+1, xm[6]-xm[5]), dtype=torch.float, device=device)
506
+ xmHPhi1=torch.zeros((nstep, xm[2]-xm[1], xm[4]-xm[3]+1, xm[6]-xm[5]), dtype=torch.float, device=device)
507
+ xmHPhi2=torch.zeros((nstep, xm[2]-xm[1], xm[4]-xm[3], xm[6]-xm[5]+1), dtype=torch.float, device=device)
508
+ elif PML[4]!=None:
509
+ xmEPhi1=PML[4].contiguous()
510
+ xmEPhi2=PML[5].contiguous()
511
+ xmHPhi1=PML[6].contiguous()
512
+ xmHPhi2=PML[7].contiguous()
513
+
514
+ if y0.numel()!=0 and PML[8]==None:
515
+ y0EPhi1=torch.zeros((nstep, y0[2]-y0[1], y0[4]-y0[3]+1, y0[6]-y0[5]+1), dtype=torch.float, device=device)
516
+ y0EPhi2=torch.zeros((nstep, y0[2]-y0[1]+1, y0[4]-y0[3]+1, y0[6]-y0[5]), dtype=torch.float, device=device)
517
+ y0HPhi1=torch.zeros((nstep, y0[2]-y0[1]+1, y0[4]-y0[3], y0[6]-y0[5]), dtype=torch.float, device=device)
518
+ y0HPhi2=torch.zeros((nstep, y0[2]-y0[1], y0[4]-y0[3], y0[6]-y0[5]+1), dtype=torch.float, device=device)
519
+ elif PML[8]!=None:
520
+ y0EPhi1=PML[8].contiguous()
521
+ y0EPhi2=PML[9].contiguous()
522
+ y0HPhi1=PML[10].contiguous()
523
+ y0HPhi2=PML[11].contiguous()
524
+
525
+ if ym.numel()!=0 and PML[12]==None:
526
+ ymEPhi1=torch.zeros((nstep, ym[2]-ym[1], ym[4]-ym[3]+1, ym[6]-ym[5]+1),dtype=torch.float, device=device)
527
+ ymEPhi2=torch.zeros((nstep, ym[2]-ym[1]+1, ym[4]-ym[3]+1, ym[6]-ym[5]), dtype=torch.float, device=device)
528
+ ymHPhi1=torch.zeros((nstep, ym[2]-ym[1]+1, ym[4]-ym[3], ym[6]-ym[5]), dtype=torch.float, device=device)
529
+ ymHPhi2=torch.zeros((nstep, ym[2]-ym[1], ym[4]-ym[3], ym[6]-ym[5]+1), dtype=torch.float, device=device)
530
+ elif PML[12]!=None:
531
+ ymEPhi1=PML[12].contiguous()
532
+ ymEPhi2=PML[13].contiguous()
533
+ ymHPhi1=PML[14].contiguous()
534
+ ymHPhi2=PML[15].contiguous()
535
+
536
+ if z0.numel()!=0 and PML[16]==None:
537
+ z0EPhi1=torch.zeros((nstep, z0[2]-z0[1], z0[4]-z0[3]+1, z0[6]-z0[5]+1), dtype=torch.float, device=device)
538
+ z0EPhi2=torch.zeros((nstep, z0[2]-z0[1]+1, z0[4]-z0[3], z0[6]-z0[5]+1), dtype=torch.float, device=device)
539
+ z0HPhi1=torch.zeros((nstep, z0[2]-z0[1]+1, z0[4]-z0[3], z0[6]-z0[5]), dtype=torch.float, device=device)
540
+ z0HPhi2=torch.zeros((nstep, z0[2]-z0[1], z0[4]-z0[3]+1, z0[6]-z0[5]), dtype=torch.float, device=device)
541
+ elif PML[16]!=None:
542
+ z0EPhi1=PML[16].contiguous()
543
+ z0EPhi2=PML[17].contiguous()
544
+ z0HPhi1=PML[18].contiguous()
545
+ z0HPhi2=PML[19].contiguous()
546
+
547
+ if zm.numel()!=0 and PML[20]==None:
548
+ zmEPhi1=torch.zeros((nstep, zm[2]-zm[1], zm[4]-zm[3]+1, zm[6]-zm[5]+1), dtype=torch.float, device=device)
549
+ zmEPhi2=torch.zeros((nstep, zm[2]-zm[1]+1, zm[4]-zm[3], zm[6]-zm[5]+1), dtype=torch.float, device=device)
550
+ zmHPhi1=torch.zeros((nstep, zm[2]-zm[1]+1, zm[4]-zm[3], zm[6]-zm[5]), dtype=torch.float, device=device)
551
+ zmHPhi2=torch.zeros((nstep, zm[2]-zm[1], zm[4]-zm[3]+1, zm[6]-zm[5]), dtype=torch.float, device=device)
552
+ elif PML[20]!=None:
553
+ zmEPhi1=PML[20].contiguous()
554
+ zmEPhi2=PML[21].contiguous()
555
+ zmHPhi1=PML[22].contiguous()
556
+ zmHPhi2=PML[23].contiguous()
557
+
558
+ return x0EPhi1,x0EPhi2,x0HPhi1,x0HPhi2,xmEPhi1,xmEPhi2,xmHPhi1,xmHPhi2,y0EPhi1,y0EPhi2,y0HPhi1,y0HPhi2,ymEPhi1,ymEPhi2,ymHPhi1,ymHPhi2,z0EPhi1,z0EPhi2,z0HPhi1,z0HPhi2,zmEPhi1,zmEPhi2,zmHPhi1,zmHPhi2
559
+
560
+
561
+ def checkpoint_initial_field(device=None,per_nstep=None, dx=None, dt=None,
562
+ source_amplitudes=None,
563
+ source_location=None,
564
+ receiver_location=None,
565
+ er=None, se=None,mr=None,
566
+ pmlthick=10):
567
+ E=None
568
+ H=None
569
+ PML=None
570
+
571
+ er,se,nx,ny,nz,_,nstep,_,_,_,_,mr,_,dtype,pmlthick,source_amplitudes=initialization(device,er,se,mr,source_amplitudes,source_location,receiver_location,dx,dt,pmlthick)
572
+
573
+ Ex,Ey,Ez=create_or_separate(E,nx,ny,nz,nstep,device,dtype)
574
+ Hx,Hy,Hz=create_or_separate(H,nx,ny,nz,nstep,device,dtype)
575
+
576
+ 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)
577
+
578
+
579
+ 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)
580
+
581
+ del x01,x02,xm1,xm2,y01,y02,ym1,ym2,z01,z02,zm1,zm2
582
+
583
+ print("per_nstep:"+str(per_nstep))
584
+ print("total step:"+str(nstep))
585
+ # print_field_shapes((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))
586
+
587
+ if per_nstep==None:
588
+ 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)
589
+ elif er.shape[2]==1:
590
+ return (Ex[:per_nstep,:,:,:],Ey[:per_nstep,:,:,:],Ez[:per_nstep,:,:,:]),(Hx[:per_nstep,:,:,:],Hy[:per_nstep,:,:,:],Hz[:per_nstep,:,:,:]),(x0EPhi1[:per_nstep,:,:,:],x0EPhi2[:per_nstep,:,:,:],x0HPhi1[:per_nstep,:,:,:],x0HPhi2[:per_nstep,:,:,:],xmEPhi1[:per_nstep,:,:,:],xmEPhi2[:per_nstep,:,:,:],xmHPhi1[:per_nstep,:,:,:],xmHPhi2[:per_nstep,:,:,:],y0EPhi1[:per_nstep,:,:,:],y0EPhi2[:per_nstep,:,:,:],y0HPhi1[:per_nstep,:,:,:],y0HPhi2[:per_nstep,:,:,:],ymEPhi1[:per_nstep,:,:,:],ymEPhi2[:per_nstep,:,:,:],ymHPhi1[:per_nstep,:,:,:],ymHPhi2[:per_nstep,:,:,:],z0EPhi1,z0EPhi2,z0HPhi1,z0HPhi2,zmEPhi1,zmEPhi2,zmHPhi1,zmHPhi2)
591
+ else:
592
+ return (Ex[:per_nstep,:,:,:],Ey[:per_nstep,:,:,:],Ez[:per_nstep,:,:,:]),(Hx[:per_nstep,:,:,:],Hy[:per_nstep,:,:,:],Hz[:per_nstep,:,:,:]),(x0EPhi1[:per_nstep,:,:,:],x0EPhi2[:per_nstep,:,:,:],x0HPhi1[:per_nstep,:,:,:],x0HPhi2[:per_nstep,:,:,:],xmEPhi1[:per_nstep,:,:,:],xmEPhi2[:per_nstep,:,:,:],xmHPhi1[:per_nstep,:,:,:],xmHPhi2[:per_nstep,:,:,:],y0EPhi1[:per_nstep,:,:,:],y0EPhi2[:per_nstep,:,:,:],y0HPhi1[:per_nstep,:,:,:],y0HPhi2[:per_nstep,:,:,:],ymEPhi1[:per_nstep,:,:,:],ymEPhi2[:per_nstep,:,:,:],ymHPhi1[:per_nstep,:,:,:],ymHPhi2[:per_nstep,:,:,:],z0EPhi1[:per_nstep,:,:,:],z0EPhi2[:per_nstep,:,:,:],z0HPhi1[:per_nstep,:,:,:],z0HPhi2[:per_nstep,:,:,:],zmEPhi1[:per_nstep,:,:,:],zmEPhi2[:per_nstep,:,:,:],zmHPhi1[:per_nstep,:,:,:],zmHPhi2[:per_nstep,:,:,:])
593
+
594
+
595
+ def zero_field(*tensors):
596
+ 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)