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/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)
|