Repository navigation
Expand file tree
/
Copy pathangular_spectrum_solver.py
More file actions
2313 lines (2002 loc) · 99.1 KB
/
Copy pathangular_spectrum_solver.py
File metadata and controls
2313 lines (2002 loc) · 99.1 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428
429
430
431
432
433
434
435
436
437
438
439
440
441
442
443
444
445
446
447
448
449
450
451
452
453
454
455
456
457
458
459
460
461
462
463
464
465
466
467
468
469
470
471
472
473
474
475
476
477
478
479
480
481
482
483
484
485
486
487
488
489
490
491
492
493
494
495
496
497
498
499
500
501
502
503
504
505
506
507
508
509
510
511
512
513
514
515
516
517
518
519
520
521
522
523
524
525
526
527
528
529
530
531
532
533
534
535
536
537
538
539
540
541
542
543
544
545
546
547
548
549
550
551
552
553
554
555
556
557
558
559
560
561
562
563
564
565
566
567
568
569
570
571
572
573
574
575
576
577
578
579
580
581
582
583
584
585
586
587
588
589
590
591
592
593
594
595
596
597
598
599
600
601
602
603
604
605
606
607
608
609
610
611
612
613
614
615
616
617
618
619
620
621
622
623
624
625
626
627
628
629
630
631
632
633
634
635
636
637
638
639
640
641
642
643
644
645
646
647
648
649
650
651
652
653
654
655
656
657
658
659
660
661
662
663
664
665
666
667
668
669
670
671
672
673
674
675
676
677
678
679
680
681
682
683
684
685
686
687
688
689
690
691
692
693
694
695
696
697
698
699
700
701
702
703
704
705
706
707
708
709
710
711
712
713
714
715
716
717
718
719
720
721
722
723
724
725
726
727
728
729
730
731
732
733
734
735
736
737
738
739
740
741
742
743
744
745
746
747
748
749
750
751
752
753
754
755
756
757
758
759
760
761
762
763
764
765
766
767
768
769
770
771
772
773
774
775
776
777
778
779
780
781
782
783
784
785
786
787
788
789
790
791
792
793
794
795
796
797
798
799
800
801
802
803
804
805
806
807
808
809
810
811
812
813
814
815
816
817
818
819
820
821
822
823
824
825
826
827
828
829
830
831
832
833
834
835
836
837
838
839
840
841
842
843
844
845
846
847
848
849
850
851
852
853
854
855
856
857
858
859
860
861
862
863
864
865
866
867
868
869
870
871
872
873
874
875
876
877
878
879
880
881
882
883
884
885
886
887
888
889
890
891
892
893
894
895
896
897
898
899
900
901
902
903
904
905
906
907
908
909
910
911
912
913
914
915
916
917
918
919
920
921
922
923
924
925
926
927
928
929
930
931
932
933
934
935
936
937
938
939
940
941
942
943
944
945
946
947
948
949
950
951
952
953
954
955
956
957
958
959
960
961
962
963
964
965
966
967
968
969
970
971
972
973
974
975
976
977
978
979
980
981
982
983
984
985
986
987
988
989
990
991
992
993
994
995
996
997
998
999
1000
"""
Angular Spectrum Solver — refactored Python/JAX port of angular_spectrum_solver.m
Implements:
- Strang split-step propagation (2nd-order splitting accuracy)
- TVD slope limiting (generalized minmod) for the nonlinear flux
- Adaptive transverse-wavenumber filtering in the angular spectrum operator
- Intensity-loss tracking due to attenuation
- Wendland C2 smooth boundary profiles
- Frequency-weighted boundary damping (quasi-PML)
- Super-absorbing boundary condition (Engquist-Majda directional decomposition)
Original MATLAB code by Gianmarco Pinton (2017-2023).
Python port and refactoring 2024-2026.
"""
import os
import numpy as np
import jax
import jax.numpy as jnp
from jax import jit
from functools import partial
from dataclasses import dataclass, field, replace as _dc_replace
from typing import Optional, Callable
import time as _time
from tof_extraction import (
extract_tof_envelope as _extract_tof,
extract_tof_matched_filter_parabolic as _extract_tof_mf,
)
# ---------------------------------------------------------------------------
# Parameter container
# ---------------------------------------------------------------------------
@dataclass
class SolverParams:
"""All solver parameters with defaults matching the MATLAB refactored code."""
dX: float = 0.0
dY: float = 0.0
dT: float = 0.0
c0: float = 1500.0
rho0: float = 1000.0
beta: float = 3.5
alpha0: float = 0.5 # dB/MHz^pow/cm (negative → water)
attenPow: float = 1.0
f0: float = 5e6
propDist: float = 0.08
boundaryFactor: float = 0.2
useSplitStep: bool = True
useAdaptiveFiltering: bool = True # k-space cosine taper
useTVD: bool = True # TVD slope limiter (independent of k-filtering)
freqFilterThreshold: float = 0.05
adaptiveFilterStrength: float = 0.7
stabilityThreshold: float = 0.2
stabilityRecoveryFactor: float = 0.15
dZmin: float = 1e-3
# --- boundary improvements ---
useBoundaryLayer: bool = True
boundaryProfile: str = 'quadratic' # 'quadratic' | 'wendland'
useFreqWeightedBoundary: bool = False
useSuperAbsorbing: bool = False
superAbsorbingStrength: float = 0.8 # 0–1, fraction of incoming wave removed
# --- flux scheme ---
fluxScheme: str = 'rusanov' # 'rusanov' | 'kt' (Kurganov-Tadmor)
# --- obliquity correction on attenuation/dispersion filter ---
useObliquityCorrection: bool = True # scale alpha, alphaStar by k/k_z per mode
# --- beam-averaged obliquity correction on the nonlinear operator ---
# Scales the Burgers coefficient N by the power-weighted <k/k_z> of the
# current plane-wave spectrum, so that the effective nonlinear path length
# matches the beam's mean obliquity rather than the axial step dz. Adds one
# FFT per march step when enabled.
useNonlinearityObliquity: bool = False
# --- phase screens for heterogeneous propagation ---
# List of per-plane screens, applied when the march crosses their z.
# Accepted tuple forms:
# (z, phase) — phase only, radians at f0
# (z, phase, amp) — plus amplitude transmission at f0 (0..1],
# frequency-scaled as amp^(f/f0) (α ∝ f law)
# (z, phase, amp, y) — plus attenuation power-law exponent: the
# transmission becomes amp^((f/f0)^y), i.e. a
# local α ∝ f^y law. y is a scalar or an
# (nX, nY) map (e.g. cortical bone vs marrow),
# valid on Szabo's power-law range 0 < y < 3.
phaseScreens: object = None
# When True, every amplitude screen additionally applies the
# Kramers-Kronig dispersive phase implied by its power law (Szabo),
# using the same formulas as precalculate_ad_pow2: with A0 = -ln(amp)
# the Np loss at f0 and fs = f/f0,
# φ(f) = tan(πy/2) · A0 · (fs^y - fs) for y ≠ 1
# φ(f) = -(2/π) · A0 · fs · ln(fs) for y = 1
# When False (default, legacy) screens are amplitude-only and add no
# dispersive phase beyond the sound-speed term already in `phase`.
screenKKDispersion: bool = False
# --- distributed source injection (bowl transducer) ---
sourcePlanes: object = None # list of (z_position, field_slice) from make_bowl_source_planes
# --- diagnostic / validation imagery ---
# ``diagnostic`` defaults to True: the solver writes runtime PNG
# dashboards every ``diagnosticInterval`` steps and a final summary
# report under ``diagnosticDir``. ``pdur`` (pulse duration, s) is used
# for the Isppa = pI * dT / (c0*rho0*pdur) conversion; if None it is
# estimated from the analytic-signal envelope of the central trace.
# Set ``diagnostic=False`` for production sweeps where you don't need
# the imagery (also mandatory when ``useGPUReductions=True``).
diagnostic: bool = True
diagnosticInterval: int = 10
diagnosticDir: str = './diagnostic_frames'
diagnosticSummary: bool = True
diagnosticInitialConditions: bool = True
# Lab-frame (conventional space-time) wavefront validation. When
# ``diagnostic`` and ``diagnosticLabframe`` are both on, the solver
# captures a lightweight xz time-history of the SIGNED mid-y field per
# march step (time axis subsampled to <= ``diagnosticLabframeMaxFrames``
# samples) and the end-of-run summary de-shears it from the retarded
# frame (tau = t - z/c0) into laboratory space-time, adding a
# ``labframe_wavefront.png`` / ``labframe_validation.pdf`` page under
# ``<diagnosticDir>/summary/``. Costs one small (nX, nT_sub) device->host
# copy per step; fully skipped when ``diagnostic`` is False.
diagnosticLabframe: bool = True
diagnosticLabframeMaxFrames: int = 300
pdur: Optional[float] = None
# Pre-flight sanity report. Defaults to True: a single-page LaTeX/PDF
# report is generated under ``preflightDir`` before the solve starts.
# Set ``preflightDir = diagnosticDir`` to co-locate the report with
# the runtime frames. Set ``preflight=False`` to skip (e.g. inside
# parameter-sweep loops).
preflight: bool = True
preflightDir: str = './preflight'
preflightScenario: str = ''
preflightCompilePdf: bool = True
# --- performance: keep tracking-array reductions on the GPU ---
# When False (default) the per-step pnp/ppp/pI/pIloss/pax reductions
# and the stability-check max(|field|) are computed on the host after
# transferring the full (nX, nY, nT) field across PCIe. On large grids
# this is the dominant cost (~5 s/step at 600^3 on an A6000) and
# leaves the GPU idle.
# When True the reductions run on the GPU and only the small reduced
# arrays (and a single scalar for stability) are transferred back.
# Output shapes and contents are unchanged. Forced off internally when
# TOF extraction is requested, since TOF needs the full host field.
useGPUReductions: bool = False
# --- physically-correct attenuation-only loss tracking ---
# When False (default, legacy), pIloss[:, :, cc] = max(0, I_before_step
# - I_after_step) per pixel, where I = ∫|p|² dt at z. This formula
# mixes true attenuation with diffractive lateral redistribution and
# then one-sidedly clips, so it is NOT a faithful local Q(r). It can
# be spiky on-axis (Fresnel-zone interference contributes apparent
# "loss") and discards energy gained at pixels where the wave focuses.
# When True, pIloss[:, :, cc] is computed from the attenuation step(s)
# ONLY: I_pre_atten - I_post_atten. This is the manuscript's local
# Q(r), per-pixel non-negative (modulo small obliquity cross-talk).
# Currently implemented for the split-step KT march (with/without
# nonlinearity-obliquity correction); other march variants fall back
# to the legacy formula with a printed warning.
useAttenLoss: bool = False
# --- cubic nonlinearity (shear-shock regime) ---
# When 2 (default), the solver uses the standard quadratic Burgers
# operator with coefficient N = β/(2·c₀³·ρ₀) — appropriate for
# longitudinal acoustic waves in fluids/soft tissue. When 3, the
# solver uses the cubic Burgers operator ∂p/∂z = -N₃·∂(p³)/∂t with
# coefficient N₃ = β₃/(3·c₀⁵·ρ₀²). Cubic nonlinearity dominates for
# transverse (shear) waves where β₃ ~ 100--200 in soft tissue
# (Catheline 2003, Gennisson 2007). Shock formation distance is
# 1/(β₃·M²·k₀) — much shorter than for longitudinal waves at the same
# Mach number, so a sub-wavelength shock is typical for shear.
nonlinearityOrder: int = 2
beta3: float = 0.0
# ---------------------------------------------------------------------------
# Absorbing boundary profiles
# ---------------------------------------------------------------------------
def ablvec(N: int, n: int) -> np.ndarray:
"""Original quadratic boundary profile: C0 continuous."""
if N == 1:
return np.ones(1) # singleton axis (2-D mode): no absorbing layer
vec = np.zeros(N)
for nn in range(n):
x = (n - nn - 1) / n
vec[nn] = x ** 2
for nn in range(N - n, N):
x = (nn - (N - n - 1)) / n
vec[nn] = x ** 2
return 1.0 - vec
def ablvec_wendland(N: int, n: int) -> np.ndarray:
"""Wendland C2 boundary profile: 1 - 10s^3 + 15s^4 - 6s^5.
The value, first derivative, and second derivative are all zero at
the domain edge (s=1) and smoothly transition to 1 in the interior
(s=0). This greatly reduces spurious reflections compared to the
quadratic profile whose second derivative is discontinuous at the
taper onset.
"""
if N == 1:
return np.ones(1) # singleton axis (2-D mode): no absorbing layer
vec = np.ones(N)
for nn in range(n):
s = (n - nn - 1) / n # 1 at edge, 0 at interior
vec[nn] = 1.0 - 10*s**3 + 15*s**4 - 6*s**5
for nn in range(N - n, N):
s = (nn - (N - n - 1)) / n
vec[nn] = 1.0 - 10*s**3 + 15*s**4 - 6*s**5
return np.clip(vec, 0.0, 1.0)
# ---------------------------------------------------------------------------
# Precalculate absorbing boundary layer (spatial × temporal)
# ---------------------------------------------------------------------------
def precalculate_abl(nX, nY, nT, boundary_factor=0.2, split_step=False,
profile='quadratic'):
print(f'Precalculating absorbing boundary layer ({profile})...')
_vec_fn = ablvec_wendland if profile == 'wendland' else ablvec
x_bw = max(round(nX * boundary_factor), 1)
y_bw = max(round(nY * boundary_factor), 1)
t_bw = max(round(nT * boundary_factor), 1)
abl_tmp = np.outer(_vec_fn(nX, x_bw), _vec_fn(nY, y_bw))
abl_vec = _vec_fn(nT, t_bw)
abl = (abl_tmp[:, :, np.newaxis] * abl_vec[np.newaxis, np.newaxis, :]).astype(np.float32)
abl_half = np.sqrt(abl) if split_step else None
print('done.')
return abl, abl_half
def precalculate_identity_abl(nX, nY, nT, split_step=False):
"""Return a neutral boundary mask for studies that should exclude boundary damping."""
abl = np.ones((nX, nY, nT), dtype=np.float32)
abl_half = abl.copy() if split_step else None
return abl, abl_half
# ---------------------------------------------------------------------------
# Frequency-weighted boundary damping (improvement 3 — quasi-PML)
# ---------------------------------------------------------------------------
def precalculate_abl_freq(nX, nY, nT, dT, c0, f0, boundary_factor=0.2,
profile='quadratic'):
"""Precompute a frequency-dependent spatial damping mask.
At low temporal frequencies the boundary is fewer wavelengths wide,
so the damping must be stronger to prevent reflection. The mask is
stored in the rfft frequency layout: shape (nX, nY, nT//2+1).
For each frequency bin f_m the 2D spatial taper is raised to the
power s_m = f_ref / max(f_m, f_min) so that:
* f_m = f_ref → normal damping (exponent 1)
* f_m < f_ref → stronger damping (exponent > 1)
* f_m > f_ref → weaker damping (exponent < 1)
"""
print('Precalculating frequency-weighted boundary mask...')
_vec_fn = ablvec_wendland if profile == 'wendland' else ablvec
x_bw = max(round(nX * boundary_factor), 1)
y_bw = max(round(nY * boundary_factor), 1)
abl_xy = np.outer(_vec_fn(nX, x_bw), _vec_fn(nY, y_bw)) # (nX, nY)
n_freq = nT // 2 + 1
freqs = np.fft.rfftfreq(nT, dT) # (n_freq,)
f_ref = f0
f_min = freqs[1] if len(freqs) > 1 else 1.0 # avoid div-by-zero
abl_freq = np.zeros((nX, nY, n_freq), dtype=np.float32)
for m in range(n_freq):
fm = max(freqs[m], f_min)
exponent = f_ref / fm
# Clamp exponent to avoid extreme values
exponent = np.clip(exponent, 0.1, 10.0)
abl_freq[:, :, m] = abl_xy ** exponent
print('done.')
return abl_freq
# ---------------------------------------------------------------------------
# Super-absorbing boundary (improvement 4 — Engquist-Majda directional)
# ---------------------------------------------------------------------------
@jit
def _super_absorbing_step(field, bdy_weight_x, bdy_weight_y, dX, dY, dT,
c0, strength):
"""Remove the incoming wave component at each spatial boundary.
Uses a first-order Engquist-Majda decomposition in the temporal
frequency domain. At each frequency omega the incoming pressure at
the left x-boundary is
P_in = 0.5 * (P + (c0 / (i omega)) * dP/dx )
and at the right boundary
P_in = 0.5 * (P - (c0 / (i omega)) * dP/dx )
The incoming component is then subtracted (scaled by *strength*) in
the boundary region defined by bdy_weight_x / bdy_weight_y.
Parameters
----------
field : (nX, nY, nT) real field
bdy_weight_x : (nX,) spatial weight: 0 in interior, >0 in boundary
bdy_weight_y : (nY,) spatial weight
dX, dY, dT : grid spacings
c0 : sound speed
strength : 0-1, fraction of incoming wave removed
"""
nX, nY, nT_full = field.shape
# Transform to temporal frequency domain (rfft along axis 2)
F = jnp.fft.rfft(field, axis=2) # (nX, nY, n_freq)
n_freq = F.shape[2]
# Frequency vector (angular)
freqs = jnp.fft.rfftfreq(nT_full, dT)
omega = 2.0 * jnp.pi * freqs # (n_freq,)
# Coefficient c0 / (i omega) = -i c0 / omega
# At DC (omega=0) and very low frequencies the decomposition is
# ill-conditioned, so we zero the coefficient there.
omega_min = 2.0 * jnp.pi * freqs[1] if n_freq > 1 else 1.0
valid = jnp.abs(omega) > 0.5 * omega_min
coeff = jnp.where(valid, -1j * c0 / jnp.where(valid, omega, 1.0), 0.0 + 0j)
coeff = coeff[jnp.newaxis, jnp.newaxis, :] # (1, 1, n_freq)
# --- x-boundaries ---
# Forward spatial gradient dF/dx (central differences, one-sided at edges)
dFdx = jnp.zeros_like(F)
dFdx = dFdx.at[1:-1, :, :].set((F[2:, :, :] - F[:-2, :, :]) / (2 * dX))
dFdx = dFdx.at[0, :, :].set((F[1, :, :] - F[0, :, :]) / dX)
dFdx = dFdx.at[-1, :, :].set((F[-1, :, :] - F[-2, :, :]) / dX)
# Incoming at left (positive x direction = into domain)
P_in_left = 0.5 * (F + coeff * dFdx)
# Incoming at right (negative x direction = into domain)
P_in_right = 0.5 * (F - coeff * dFdx)
# Build weight for left boundary (bdy_weight > 0 near edges)
# Left half gets left correction, right half gets right correction
wx = bdy_weight_x # (nX,)
wx_left = wx * (jnp.arange(nX) < nX // 2)
wx_right = wx * (jnp.arange(nX) >= nX // 2)
wx_left = wx_left[:, jnp.newaxis, jnp.newaxis] # (nX,1,1)
wx_right = wx_right[:, jnp.newaxis, jnp.newaxis]
F = F - strength * wx_left * P_in_left
F = F - strength * wx_right * P_in_right
# --- y-boundaries ---
dFdy = jnp.zeros_like(F)
dFdy = dFdy.at[:, 1:-1, :].set((F[:, 2:, :] - F[:, :-2, :]) / (2 * dY))
dFdy = dFdy.at[:, 0, :].set((F[:, 1, :] - F[:, 0, :]) / dY)
dFdy = dFdy.at[:, -1, :].set((F[:, -1, :] - F[:, -2, :]) / dY)
P_in_bottom = 0.5 * (F + coeff * dFdy)
P_in_top = 0.5 * (F - coeff * dFdy)
wy = bdy_weight_y # (nY,)
wy_bottom = wy * (jnp.arange(nY) < nY // 2)
wy_top = wy * (jnp.arange(nY) >= nY // 2)
wy_bottom = wy_bottom[jnp.newaxis, :, jnp.newaxis]
wy_top = wy_top[jnp.newaxis, :, jnp.newaxis]
F = F - strength * wy_bottom * P_in_bottom
F = F - strength * wy_top * P_in_top
# Transform back
return jnp.fft.irfft(F, n=nT_full, axis=2)
def _make_boundary_weights(N, n_bdy):
"""Smooth weight that is 1 at the edge and 0 in the interior."""
if N == 1:
return np.zeros(1, dtype=np.float32) # singleton axis: no edges
w = np.zeros(N, dtype=np.float32)
for nn in range(n_bdy):
s = (n_bdy - nn - 1) / n_bdy # 1 at edge → 0 at interior boundary
w[nn] = s
for nn in range(N - n_bdy, N):
s = (nn - (N - n_bdy - 1)) / n_bdy
w[nn] = s
return w
# ---------------------------------------------------------------------------
# Phase screen model for heterogeneous propagation
# ---------------------------------------------------------------------------
def generate_phase_screen(nX, nY, dX, dY, c0, f0,
correlation_length, speed_std,
thickness=None, seed=None):
"""Generate a random phase screen for modelling tissue heterogeneity.
The screen represents a thin layer of tissue with spatially-varying
sound speed. The speed fluctuations are drawn from a Gaussian random
field with a specified correlation length and standard deviation.
Parameters
----------
nX, nY : int — grid dimensions
dX, dY : float — grid spacing (m)
c0 : float — background sound speed (m/s)
f0 : float — center frequency (Hz) — sets k0
correlation_length : float — spatial correlation length (m)
speed_std : float — standard deviation of sound-speed
fluctuations (m/s)
thickness : float — effective layer thickness (m);
default = correlation_length
seed : int — random seed for reproducibility
Returns
-------
phase_shift : ndarray (nX, nY) — phase shift in radians at each
grid point. Applied to the field as
p *= exp(i * phase_shift).
c_map : ndarray (nX, nY) — the sound-speed map (m/s).
"""
if thickness is None:
thickness = correlation_length
rng = np.random.default_rng(seed)
# Spatial frequency grid
kx = np.fft.fftfreq(nX, dX) * 2 * np.pi
ky = np.fft.fftfreq(nY, dY) * 2 * np.pi
KX, KY = np.meshgrid(kx, ky, indexing='ij')
K2 = KX**2 + KY**2
# Gaussian power spectrum with correlation length L:
# S(k) ∝ exp(-k² L² / 4)
L = correlation_length
spectrum = np.exp(-K2 * L**2 / 4.0)
# Generate complex white noise and filter
noise = rng.standard_normal((nX, nY)) + 1j * rng.standard_normal((nX, nY))
filtered = np.fft.ifft2(np.fft.fft2(noise) * np.sqrt(spectrum))
delta_c = np.real(filtered)
# Normalize to desired standard deviation
delta_c = delta_c / (np.std(delta_c) + 1e-30) * speed_std
# Sound speed map
c_map = c0 + delta_c
# Phase shift: Δφ = ω * thickness * (1/c(x,y) - 1/c0)
omega = 2 * np.pi * f0
phase_shift = omega * thickness * (1.0 / c_map - 1.0 / c0)
return phase_shift.astype(np.float32), c_map.astype(np.float32)
@partial(jit, static_argnames=('kk_dispersion',))
def _apply_phase_screen(field, phase_screen, f0_bin, amplitude_screen=None,
atten_pow=None, kk_dispersion=False):
"""Apply a phase+amplitude screen in the temporal frequency domain.
The phase screen shifts each frequency component by phase * f/f0.
The amplitude screen attenuates each frequency component by
amp^((f/f0)^y) — a local α ∝ f^y power law whose transmission at f0
is ``amp``. With the default y = 1 (atten_pow None) this reduces to
the legacy amp^(f/f0): higher frequencies see more attenuation
through bone.
Parameters
----------
field : (nX, nY, nT) real-valued pressure field
phase_screen : (nX, nY) phase shift at f0 in radians
f0_bin : float — the rfft bin index corresponding to f0
amplitude_screen : (nX, nY) transmission factor at f0 (0 to 1), optional
atten_pow : power-law exponent y, scalar or (nX, nY) map, optional
(None → y = 1). Valid on Szabo's power-law range 0 < y < 3.
kk_dispersion : bool, static — when True, also apply the
Kramers-Kronig phase implied by the screen's power law, using the
same Szabo formulas as precalculate_ad_pow2 (the y = 1 log form
is selected per-pixel where |y - 1| < 1e-6, mirroring the
volumetric filter's odd-power branch; y = 2 gives tan(π) ≈ 0, no
dispersion, as it should). No-op without an amplitude screen.
"""
nT = field.shape[2]
F = jnp.fft.rfft(field, axis=2) # (nX, nY, n_freq)
n_freq = F.shape[2]
# Frequency bin indices; f_bin = k * df, f0 = f0_bin * df
# Scale: f / f0 = k / f0_bin
freq_idx = jnp.arange(n_freq, dtype=jnp.float32)
f_scale = freq_idx / jnp.maximum(f0_bin, 1.0)
# Phase modulation: exp(-i * phase_screen(x,y) * f/f0)
# Negative sign: phase = omega*dz*(1/c - 1/c0) is the extra spatial
# wavenumber*distance. In the temporal rfft domain the correction is
# exp(-j * delta_phi) because spatial propagation uses exp(-j*kz*z).
ps = phase_screen[:, :, jnp.newaxis] # (nX, nY, 1)
fs = f_scale[jnp.newaxis, jnp.newaxis, :] # (1, 1, n_freq)
F = F * jnp.exp(-1j * ps * fs)
# Amplitude modulation: amp^((f/f0)^y) — frequency-dependent attenuation
if amplitude_screen is not None:
amp = jnp.clip(amplitude_screen[:, :, jnp.newaxis], 1e-6, 1.0)
if atten_pow is None:
y = None
fs_pow = fs # legacy y = 1
else:
y = jnp.asarray(atten_pow, dtype=jnp.float32)
y = y[:, :, jnp.newaxis] if y.ndim == 2 else y
fs_pow = fs ** y # 0^y = 0 keeps DC untouched
# amp^((f/f0)^y): at f0 the nominal transmission, higher f more loss
F = F * jnp.power(amp, fs_pow)
if kk_dispersion:
# Szabo dispersion consistent with precalculate_ad_pow2,
# rewritten in screen variables: A0 = -ln(amp) is the Np loss
# at f0, so conv·dz = A0/f0^y and the dispersive phase
# α*(f)·dz = ω·dz·(1/c(ω) - 1/c0) collapses to pure functions
# of fs = f/f0. Same exp(-j·φ) sign convention as the phase
# screen and the volumetric filter.
y_kk = jnp.float32(1.0) if y is None else y
A0 = -jnp.log(amp)
fs_pos = fs > 0
fs_safe = jnp.where(fs_pos, fs, 1.0)
y_is_one = jnp.abs(y_kk - 1.0) < 1e-6
y_tan = jnp.where(y_is_one, 2.0, y_kk) # dodge tan(π/2) inf
phi_tan = jnp.tan(jnp.pi * y_tan / 2.0) * (fs ** y_tan - fs)
phi_log = -(2.0 / jnp.pi) * fs * jnp.log(fs_safe)
phi = A0 * jnp.where(y_is_one, phi_log, phi_tan)
phi = jnp.where(fs_pos, phi, 0.0)
F = F * jnp.exp(-1j * phi)
return jnp.fft.irfft(F, n=nT, axis=2)
# ---------------------------------------------------------------------------
# Precalculate modified angular spectrum (with optional adaptive filtering)
# ---------------------------------------------------------------------------
def _centered_k_axis(n, d):
"""Centered spatial-frequency axis in the solver's grid convention.
A singleton axis carries only the k=0 mode (this is how 2-D operation
collapses the y dimension), so it returns [0] instead of dividing by
n - 1.
"""
if n == 1:
return np.zeros(1)
k = np.linspace(0, n - 1, n) / (n - 1) / d * 2 * np.pi
return k - np.mean(k)
def precalculate_mas(nX, nY, nT, dX, dY, dZ, dT, c0,
split_step=False,
adaptive_filtering=False,
filter_threshold=0.05,
filter_strength=0.7):
print('Precalculating modified angular spectrum...')
kt = np.linspace(0, nT - 1, nT) / (nT - 1) / dT * 2 * np.pi / c0
kt -= np.mean(kt)
kx = _centered_k_axis(nX, dX)
ky = _centered_k_axis(nY, dY)
# Transverse wavenumber grid (precomputed once)
kk = kx[:, None] ** 2 + ky[None, :] ** 2 # (nX, nY)
# Adaptive filtering setup — kmax is the actual max spatial freq in the
# centered grid. Singleton axes (2-D operation) carry no bandwidth and
# must not constrain the cutoff.
_kmax_axes = [np.pi / d for n_ax, d in ((nX, dX), (nY, dY)) if n_ax > 1]
kmax = min(_kmax_axes) if _kmax_axes else np.pi / min(dX, dY)
lambda_char = c0 / np.mean(np.abs(kt[kt != 0])) if np.any(kt != 0) else c0 / (1.0 / dT)
norm_step = dZ / lambda_char
if adaptive_filtering:
# Empirically tuned: the factor 5 sets how fast the cutoff widens
# with step size (in characteristic wavelengths); the 0.95 cap
# always discards the outermost 5% of k-space (near-Nyquist modes
# whose propagator phase is unreliable); the cosine taper starts at
# 80% of the cutoff to avoid Gibbs ringing from a hard edge.
cutoff = filter_threshold * (1 + filter_strength * norm_step * 5)
k_trans_max = min(kmax * (1 - cutoff), kmax * 0.95)
k_trans_start = k_trans_max * 0.8
print(f' adaptive filtering: cutoff={cutoff:.4f}, norm_step={norm_step:.4f}')
else:
# Fixed filter: same 5% near-Nyquist guard, with a narrower (10%)
# cosine taper since the cutoff does not move with step size.
k_trans_max = kmax * 0.95
k_trans_start = k_trans_max * 0.9
k_trans = np.sqrt(kk)
filt_mask = np.ones_like(k_trans)
trans = (k_trans > k_trans_start) & (k_trans <= k_trans_max)
if np.any(trans):
norm_pos = (k_trans[trans] - k_trans_start) / (k_trans_max - k_trans_start)
filt_mask[trans] = 0.5 * (1 + np.cos(np.pi * norm_pos))
filt_mask[k_trans > k_trans_max] = 0.0
HH = np.zeros((nX, nY, nT), dtype=np.complex128)
HH_half = np.zeros((nX, nY, nT), dtype=np.complex128) if split_step else None
for m in range(nT):
k = kt[m]
H2 = np.exp(dZ * (1j * k - np.sqrt(kk - k ** 2 + 0j)))
H1 = np.exp(dZ * (1j * k - 1j * np.sqrt(k ** 2 - kk + 0j)))
H = np.where(kk < k ** 2, H1, H2)
H *= filt_mask
HH[:, :, m] = H
if split_step:
H2h = np.exp(dZ / 2 * (1j * k - np.sqrt(kk - k ** 2 + 0j)))
H1h = np.exp(dZ / 2 * (1j * k - 1j * np.sqrt(k ** 2 - kk + 0j)))
Hh = np.where(kk < k ** 2, H1h, H2h)
Hh *= filt_mask
HH_half[:, :, m] = Hh
# Zero negative frequencies & DC bin, then double positive frequencies.
# +1 so the DC bin (index nT/2 for even nT) is included in the zeroed half.
HH[:, :, :nT // 2 + 1] = 0
HH *= 2
if split_step:
HH_half[:, :, :nT // 2 + 1] = 0
HH_half *= 2
print('done.')
return HH, HH_half
# ---------------------------------------------------------------------------
# Attenuation / dispersion
# ---------------------------------------------------------------------------
def _build_oblique_atten_filter(alpha, alphaStar, nX, nY, nT, dX, dY, dZ, c0,
dT, scale=1.0, obliquity=True,
clip_negative=True):
"""Construct the obliquity-corrected attenuation/dispersion filter.
Per plane-wave mode (k_x, k_y, omega), the physical path length over one
axial step dZ is dZ/cos(theta) = dZ * k/k_z, where k = omega/c_0 and
k_z = sqrt(k^2 - k_perp^2). Applying alpha and alphaStar as pure
omega-only filters underestimates attenuation for oblique components;
multiplying the argument by k/k_z restores the correct path length.
Returned array has shape (nX, nY, nT), complex, in fftshifted-centered
(k_x, k_y) order and natural temporal-frequency order (matches the HH
propagator convention; _attenuation_step slices to the rfft support).
For evanescent modes (k_perp^2 > k^2), no obliquity correction is
applied (the diffraction propagator already decays these); for the
DC/zero-frequency bin the filter is unity.
"""
# Transverse wavenumber grid in the same centered convention as HH
kx = _centered_k_axis(nX, dX)
ky = _centered_k_axis(nY, dY)
kk = kx[:, None] ** 2 + ky[None, :] ** 2 # (nX, nY)
f = np.arange(nT) / (nT * dT)
omega = 2 * np.pi * f
k = omega / c0 # |k| per temporal-frequency bin
afilt3d = np.ones((nX, nY, nT), dtype=np.complex128)
eps = 1e-6
for m in range(nT):
km = k[m]
if km == 0:
continue # DC: no attenuation, filter stays unity
base = -(alpha[m] + 1j * alphaStar[m]) * dZ * scale
if obliquity:
kz_sq = km ** 2 - kk
prop = kz_sq > 0
kz_safe = np.where(prop, np.maximum(np.sqrt(np.maximum(kz_sq, 0.0)),
np.abs(km) * eps), 1.0)
obl = np.where(prop, np.abs(km) / kz_safe, 0.0)
factor = np.exp(base * obl)
# Evanescent modes: no extra attenuation beyond what HH does
factor = np.where(prop, factor, 1.0)
else:
factor = np.broadcast_to(np.exp(base), kk.shape)
afilt3d[:, :, m] = factor
if clip_negative:
afilt3d = np.where(np.real(afilt3d) < 0, 0, afilt3d)
return afilt3d
def precalculate_ad(alpha0, nX, nY, nT, dX, dY, dZ, dT, c0, f0,
split_step=False, obliquity=True):
print('Precalculating attenuation filter (water, f^2, '
f'{"obliquity-corrected" if obliquity else "omega-only"})...')
f = np.arange(nT) / (nT * dT)
conv = alpha0 / 1e12 * 1e2 / (20 * np.log10(np.e))
alpha = conv * f ** 2
alphaStar0 = (conv / (2 * np.pi) ** 2) * np.tan(np.pi) * \
((2 * np.pi * f) ** 1 - (2 * np.pi * f0) ** 1)
alphaStar = 2 * np.pi * alphaStar0 * f
alphaStar[0] = 0
afilt3d = _build_oblique_atten_filter(alpha, alphaStar, nX, nY, nT,
dX, dY, dZ, c0, dT,
scale=1.0, obliquity=obliquity)
afilt3d_half = None
if split_step:
afilt3d_half = _build_oblique_atten_filter(alpha, alphaStar,
nX, nY, nT,
dX, dY, dZ, c0, dT,
scale=0.5,
obliquity=obliquity)
print('done.')
return afilt3d, afilt3d_half
# ---------------------------------------------------------------------------
# Attenuation / dispersion — general power law
# ---------------------------------------------------------------------------
def precalculate_ad_pow2(alpha0, nX, nY, nT, dX, dY, dZ, dT, c0, f0, pw,
split_step=False, obliquity=True):
print(f'Precalculating attenuation filter (f^{pw} power law, '
f'{"obliquity-corrected" if obliquity else "omega-only"})...')
f = np.arange(nT) / (nT * dT)
f_safe = f.copy()
f_safe[0] = f_safe[1] / 2 # avoid log(0)
conv = alpha0 / (1e6 ** pw) * 1e2 / (20 * np.log10(np.e))
alpha = conv * f ** pw
if pw % 2 == 1:
alphaStar0 = (-2 * conv / ((2 * np.pi) ** pw) / np.pi) * \
(np.log(2 * np.pi * f_safe) - np.log(2 * np.pi * f0))
else:
alphaStar0 = (conv / (2 * np.pi) ** pw) * np.tan(np.pi * pw / 2) * \
((2 * np.pi * f) ** (pw - 1) - (2 * np.pi * f0) ** (pw - 1))
alphaStar = 2 * np.pi * alphaStar0 * f
alphaStar[0] = 0
afilt3d = _build_oblique_atten_filter(alpha, alphaStar, nX, nY, nT,
dX, dY, dZ, c0, dT,
scale=1.0, obliquity=obliquity,
clip_negative=False)
afilt3d_half = None
if split_step:
afilt3d_half = _build_oblique_atten_filter(alpha, alphaStar,
nX, nY, nT,
dX, dY, dZ, c0, dT,
scale=0.5,
obliquity=obliquity,
clip_negative=False)
dispersion = 1.0 / ((1.0 / c0) + (alphaStar / (2 * np.pi * f_safe)))
print('done.')
return afilt3d, afilt3d_half, f, alpha, dispersion
# ---------------------------------------------------------------------------
# Beam-averaged obliquity for the nonlinear operator
# ---------------------------------------------------------------------------
def precalculate_obliquity_map(nX, nY, nT, dX, dY, dT, c0):
"""Per-mode 1/cos(theta) = k/k_z on the rfft temporal grid.
Returns a float32 array of shape (nX, nY, nT//2+1) in fftshifted-centered
(k_x, k_y) order, zero for evanescent modes and for the DC bin.
"""
kx = _centered_k_axis(nX, dX)
ky = _centered_k_axis(nY, dY)
kk = kx[:, None] ** 2 + ky[None, :] ** 2 # (nX, nY)
n_freq = nT // 2 + 1
f = np.arange(n_freq) / (nT * dT)
k_tf = 2 * np.pi * f / c0 # (n_freq,)
obl_map = np.zeros((nX, nY, n_freq), dtype=np.float32)
eps = 1e-6
for m in range(n_freq):
km = k_tf[m]
if km == 0:
continue
kz_sq = km ** 2 - kk
prop = kz_sq > 0
kz_safe = np.where(prop,
np.maximum(np.sqrt(np.maximum(kz_sq, 0.0)),
np.abs(km) * eps),
1.0)
obl_map[:, :, m] = np.where(prop, np.abs(km) / kz_safe, 0.0).astype(np.float32)
return obl_map
@jit
def _beam_obliquity_scalar(field, obl_map):
"""Power-weighted <k/k_z> over the propagating plane-wave content.
Returns a scalar in [1, inf); equals 1 for an axially-propagating beam and
grows as rim rays take larger angles from z.
"""
F = jnp.fft.rfft(field, axis=2) # (nX, nY, nT//2+1)
F = jnp.fft.fftshift(jnp.fft.fft2(F, axes=(0, 1)), axes=(0, 1))
P = (F.real ** 2 + F.imag ** 2).astype(obl_map.dtype)
prop_mask = (obl_map > 0).astype(obl_map.dtype)
num = jnp.sum(P * obl_map)
denom = jnp.sum(P * prop_mask) + 1e-30
return num / denom
# ---------------------------------------------------------------------------
# JIT-compiled march steps
# ---------------------------------------------------------------------------
@jit
def _angular_spectrum_step(field, HH, abl):
"""Full-step angular spectrum propagation + spatial/temporal boundary layer."""
field = jnp.real(jnp.fft.ifftn(
jnp.fft.ifftshift(jnp.fft.fftshift(jnp.fft.fftn(field)) * HH)
))
return field * abl
@jit
def _freq_weighted_boundary_step(field, abl_freq):
"""Apply frequency-dependent spatial damping in the rfft domain."""
nT = field.shape[2]
F = jnp.fft.rfft(field, axis=2)
F = F * abl_freq
return jnp.fft.irfft(F, n=nT, axis=2)
@jit
def _attenuation_step(field, afilt3d):
"""Obliquity-corrected attenuation/dispersion.
afilt3d has shape (nX, nY, nT//2+1), complex, in fftshifted-centered
(k_x, k_y) order (same convention as the HH propagator). For each
propagating plane-wave mode the filter carries a k/k_z factor that
applies alpha and alphaStar over the correct per-mode path length
dZ/cos(theta); evanescent modes see an identity filter since the
diffraction step already handles their decay.
"""
nT = field.shape[2]
F = jnp.fft.rfft(field, axis=2) # (nX, nY, nT//2+1) — (x, y, omega)
F = jnp.fft.fft2(F, axes=(0, 1)) # (k_x, k_y, omega), natural order
F = jnp.fft.fftshift(F, axes=(0, 1)) # centered (k_x, k_y)
F = F * afilt3d
F = jnp.fft.ifftshift(F, axes=(0, 1))
F = jnp.fft.ifft2(F, axes=(0, 1))
return jnp.fft.irfft(F, n=nT, axis=2).real
@jit
def _rusanov_flux_standard(field, N, dZ, dT):
"""Standard Rusanov flux (no TVD limiting)."""
lambdahalf = jnp.maximum(jnp.abs(field[:, :, :-1]),
jnp.abs(field[:, :, 1:]))
fluxhalf = -(field[:, :, :-1] ** 2 + field[:, :, 1:] ** 2) / 2 - \
lambdahalf * (field[:, :, 1:] - field[:, :, :-1])
flux_diff = fluxhalf[:, :, 1:] - fluxhalf[:, :, :-1]
return field.at[:, :, 1:-1].add(-N * dZ / dT * flux_diff)
@jit
def _rusanov_flux_tvd(field, N, dZ, dT, beta_tvd):
"""Rusanov flux with TVD generalized minmod slope limiting."""
uL = field[:, :, :-1]
uR = field[:, :, 1:]
a = jnp.maximum(jnp.abs(uL), jnp.abs(uR))
fluxhalf = 0.5 * (-0.5 * (uL ** 2 + uR ** 2)) - 0.5 * a * (uR - uL)
# Flux differences at interior points
flux_diff = fluxhalf[:, :, 1:] - fluxhalf[:, :, :-1]
# TVD generalized minmod limiter
delta_minus = field[:, :, 1:-1] - field[:, :, :-2]
delta_plus = field[:, :, 2:] - field[:, :, 1:-1]
eps = 1e-30
r = delta_plus / (delta_minus + eps * jnp.sign(delta_minus + eps))
r_inv = delta_minus / (delta_plus + eps * jnp.sign(delta_plus + eps))
phi_plus = jnp.maximum(0.0, jnp.minimum(jnp.minimum(beta_tvd * r, 1.0),
jnp.minimum(r, beta_tvd)))
phi_minus = jnp.maximum(0.0, jnp.minimum(jnp.minimum(beta_tvd * r_inv, 1.0),
jnp.minimum(r_inv, beta_tvd)))
limited_flux = 0.5 * (phi_plus + phi_minus) * flux_diff
return field.at[:, :, 1:-1].add(-N * dZ / dT * limited_flux)
@jit
def _minmod(a, b):
"""Two-argument minmod limiter."""
same_sign = a * b > 0
return jnp.where(same_sign, jnp.sign(a) * jnp.minimum(jnp.abs(a), jnp.abs(b)), 0.0)
@jit
def _minmod3(a, b, c):
"""Three-argument minmod: returns the value with smallest magnitude
if all three have the same sign, else zero."""
same_sign = (a * b > 0) & (b * c > 0)
s = jnp.sign(a)
m = jnp.minimum(jnp.minimum(jnp.abs(a), jnp.abs(b)), jnp.abs(c))
return jnp.where(same_sign, s * m, 0.0)
@jit
def _mc_limiter(delta_minus, delta_plus):
"""Monotonized central (MC) limiter — less diffusive than minmod.
Returns minmod(2*delta_minus, (delta_minus+delta_plus)/2, 2*delta_plus).
Uses the centered slope when safe; falls back to one-sided slopes
near discontinuities. TVD for any convex combination of minmod
and centered differences.
"""
return _minmod3(2 * delta_minus,
0.5 * (delta_minus + delta_plus),
2 * delta_plus)
@jit
def _kt_rhs(field, N, dT):
"""Semidiscrete KT right-hand side with MUSCL-MC reconstruction.
Uses the MC (monotonized central) limiter for sharper shock
resolution while retaining the TVD property.
"""
# Limited slopes in each cell along the time axis.
delta_minus = field[:, :, 1:-1] - field[:, :, :-2]
delta_plus = field[:, :, 2:] - field[:, :, 1:-1]
sigma_mid = _mc_limiter(delta_minus, delta_plus)
sigma = jnp.pad(sigma_mid, ((0, 0), (0, 0), (1, 1)))
# MUSCL reconstruction at interfaces j+1/2.
u_minus = field[:, :, :-1] + 0.5 * sigma[:, :, :-1]
u_plus = field[:, :, 1:] - 0.5 * sigma[:, :, 1:]
# Local one-sided wave speeds for retarded-time Burgers: f'(u) = -u.
a_plus = jnp.maximum(jnp.maximum(-u_minus, -u_plus), 0.0)
a_minus = jnp.minimum(jnp.minimum(-u_minus, -u_plus), 0.0)
# Retarded-time Burgers flux f(u) = -u^2 / 2.
f_minus = -0.5 * u_minus ** 2
f_plus = -0.5 * u_plus ** 2
denom = a_plus - a_minus
safe_denom = jnp.where(jnp.abs(denom) < 1e-14, 1.0, denom)
kt_flux = (
a_plus * f_minus
- a_minus * f_plus
- a_plus * a_minus * (u_plus - u_minus)
) / safe_denom
# Fall back to Lax-Friedrichs when the local speeds collapse.
lf_flux = 0.5 * (f_minus + f_plus) - 0.5 * jnp.maximum(jnp.abs(u_minus), jnp.abs(u_plus)) * (u_plus - u_minus)
fluxhalf = jnp.where(jnp.abs(denom) < 1e-14, lf_flux, kt_flux)
flux_diff = fluxhalf[:, :, 1:] - fluxhalf[:, :, :-1]
rhs = jnp.zeros_like(field)
return rhs.at[:, :, 1:-1].set(-N / dT * flux_diff)
@jit
def _kt_rhs_minmod(field, N, dT):
"""KT right-hand side with the original minmod limiter."""
delta_minus = field[:, :, 1:-1] - field[:, :, :-2]
delta_plus = field[:, :, 2:] - field[:, :, 1:-1]
sigma_mid = _minmod(delta_minus, delta_plus)
sigma = jnp.pad(sigma_mid, ((0, 0), (0, 0), (1, 1)))
u_minus = field[:, :, :-1] + 0.5 * sigma[:, :, :-1]
u_plus = field[:, :, 1:] - 0.5 * sigma[:, :, 1:]
a_plus = jnp.maximum(jnp.maximum(-u_minus, -u_plus), 0.0)
a_minus = jnp.minimum(jnp.minimum(-u_minus, -u_plus), 0.0)
f_minus = -0.5 * u_minus ** 2
f_plus = -0.5 * u_plus ** 2
denom = a_plus - a_minus
safe_denom = jnp.where(jnp.abs(denom) < 1e-14, 1.0, denom)
kt_flux = (
a_plus * f_minus
- a_minus * f_plus
- a_plus * a_minus * (u_plus - u_minus)
) / safe_denom
lf_flux = 0.5 * (f_minus + f_plus) - 0.5 * jnp.maximum(jnp.abs(u_minus), jnp.abs(u_plus)) * (u_plus - u_minus)
fluxhalf = jnp.where(jnp.abs(denom) < 1e-14, lf_flux, kt_flux)
flux_diff = fluxhalf[:, :, 1:] - fluxhalf[:, :, :-1]
rhs = jnp.zeros_like(field)
return rhs.at[:, :, 1:-1].set(-N / dT * flux_diff)
@jit
def _kt_flux(field, N, dZ, dT):
"""Second-order KT nonlinear step with SSP-RK2 integration in z."""
rhs0 = _kt_rhs(field, N, dT)
stage1 = field + dZ * rhs0
rhs1 = _kt_rhs(stage1, N, dT)
return 0.5 * field + 0.5 * (stage1 + dZ * rhs1)
def _kt_flux_adaptive(field, N, dZ, dT, cfl_target=0.5):
"""KT nonlinear step with adaptive sub-cycling for CFL stability.
If the CFL number N*|u_max|*dZ/dT exceeds cfl_target, the step is
split into sub-steps that each satisfy the CFL condition. Each
sub-step uses SSP-RK2 for second-order accuracy.
"""