"""Spatial and velocity-ladder kernels for parallel streaming."""
from __future__ import annotations
import jax
import jax.numpy as jnp
import numpy as np
from gkx.core.velocity import hermite_ladder_coeffs
def _is_tracer(x) -> bool:
return isinstance(x, jax.core.Tracer)
def _check_positive(x, name: str) -> None:
arr = jnp.asarray(x)
if _is_tracer(x) or _is_tracer(arr):
return
if arr.ndim == 0:
if float(arr) <= 0.0:
raise ValueError(f"{name} must be > 0")
return
if np.any(np.asarray(arr) <= 0.0):
raise ValueError(f"{name} must be > 0")
def _fft_ik_multiplier(kz: jnp.ndarray, like: jnp.ndarray) -> jnp.ndarray:
"""Return ``1j * kz`` with the same complex dtype as an FFT array."""
real_dtype = jnp.real(like).dtype
kz_cast = jnp.asarray(kz, dtype=real_dtype).astype(like.dtype)
return jnp.asarray(1j, dtype=like.dtype) * kz_cast
def _fft_abs_multiplier(kz: jnp.ndarray, like: jnp.ndarray) -> jnp.ndarray:
"""Return ``abs(kz)`` as a complex multiplier matching an FFT array."""
real_dtype = jnp.real(like).dtype
return jnp.asarray(jnp.abs(kz), dtype=real_dtype).astype(like.dtype)
[docs]
def grad_z_periodic(
f: jnp.ndarray, dz: float | jnp.ndarray | None = None, kz: jnp.ndarray | None = None
) -> jnp.ndarray:
"""Spectral periodic derivative along the last axis."""
if kz is None:
if dz is None:
raise ValueError("Either dz or kz must be provided")
_check_positive(dz, "dz")
n = f.shape[-1]
if kz is None:
dz_val = jnp.asarray(dz, dtype=jnp.real(f).dtype)
kz = 2.0 * jnp.pi * jnp.fft.fftfreq(n, d=dz_val)
f_hat = jnp.fft.fft(f, axis=-1)
df_hat = _fft_ik_multiplier(kz, f_hat) * f_hat
return jnp.fft.ifft(df_hat, axis=-1)
def _shift_kx_linked(
f: jnp.ndarray,
kx_link: jnp.ndarray,
kx_mask: jnp.ndarray,
) -> jnp.ndarray:
"""Shift along kx for each ky using precomputed link indices."""
f_ky = jnp.moveaxis(f, -2, 0)
kx_mask = kx_mask.astype(f.dtype)
def _gather_ky(f_slice: jnp.ndarray, idx: jnp.ndarray, mask: jnp.ndarray) -> jnp.ndarray:
gathered = jnp.take(f_slice, idx, axis=-1)
return gathered * mask
shifted = jax.vmap(_gather_ky, in_axes=(0, 0, 0))(f_ky, kx_link, kx_mask)
return jnp.moveaxis(shifted, 0, -2)
def _grad_z_linked_fd(
f: jnp.ndarray,
dz: float | jnp.ndarray,
kx_link_plus: jnp.ndarray,
kx_link_minus: jnp.ndarray,
kx_mask_plus: jnp.ndarray,
kx_mask_minus: jnp.ndarray,
) -> jnp.ndarray:
"""Finite-difference z-derivative with twist-shift kx linking at the ends."""
_check_positive(dz, "dz")
dz_val = jnp.asarray(dz, dtype=jnp.real(f).dtype)
f_z0 = f[..., 0]
f_zm1 = f[..., -1]
f_z0_shift = _shift_kx_linked(f_z0, kx_link_plus, kx_mask_plus)
f_zm1_shift = _shift_kx_linked(f_zm1, kx_link_minus, kx_mask_minus)
f_roll_p1 = jnp.concatenate([f[..., 1:], f_z0_shift[..., None]], axis=-1)
f_roll_m1 = jnp.concatenate([f_zm1_shift[..., None], f[..., :-1]], axis=-1)
return (f_roll_p1 - f_roll_m1) / (2.0 * dz_val)
def _restore_linked_real_fft_conjugates(
out: jnp.ndarray,
*,
covered_rows: jnp.ndarray,
) -> jnp.ndarray:
"""Restore the conjugate ``-ky`` rows on a full real-FFT spectral grid.
The linked-FFT chains are built on the unique dealiased positive-``ky``
block. When the runtime carries the full real-FFT-expanded ``ky`` layout, the
untouched negative rows must be reconstructed by real-FFT conjugate
symmetry so the linked derivative acts on the physical Hermitian state.
"""
Ny = out.shape[-3]
if Ny <= 1:
return out
row_idx = jnp.arange(Ny, dtype=jnp.int32)
src_rows = jnp.mod(-row_idx, Ny)
covered = jnp.asarray(covered_rows, dtype=bool)
source_covered = jnp.take(covered, src_rows, axis=0)
fill_mask = (~covered) & source_covered & (row_idx != 0)
mirrored = jnp.take(out, src_rows, axis=-3)
Nx = out.shape[-2]
if Nx > 1:
kx_neg = jnp.concatenate(
(
jnp.asarray([0], dtype=jnp.int32),
jnp.arange(Nx - 1, 0, -1, dtype=jnp.int32),
)
)
mirrored = jnp.take(mirrored, kx_neg, axis=-2)
mirrored = jnp.conj(mirrored)
mask_shape = (1,) * (out.ndim - 3) + (Ny, 1, 1)
return jnp.where(fill_mask.reshape(mask_shape), mirrored, out)
def _validate_linked_fft_inputs(
linked_indices: tuple[jnp.ndarray, ...],
linked_kz: tuple[jnp.ndarray, ...],
*,
operator: str,
) -> None:
if operator not in {"grad", "abs"}:
raise ValueError(f"unsupported linked FFT operator {operator!r}")
if len(linked_indices) != len(linked_kz):
raise ValueError("linked_indices and linked_kz must have the same length")
if not linked_indices:
suffix = "derivative" if operator == "grad" else "operator"
raise ValueError(f"linked_indices cannot be empty for linked FFT {suffix}")
def _flatten_linked_fft_state(
f: jnp.ndarray,
) -> tuple[jnp.ndarray, tuple[int, ...], int, int, int]:
"""Move spectral ``kx`` next to ``ky`` and flatten linked chain indices."""
Ny = f.shape[-3]
Nx = f.shape[-2]
Nz = f.shape[-1]
lead_shape = f.shape[:-3]
f_perm = jnp.swapaxes(f, -3, -2)
return f_perm.reshape(*lead_shape, Nx * Ny, Nz), lead_shape, Ny, Nx, Nz
def _scatter_unique_linked_modes(
target: jnp.ndarray,
idx_flat: jnp.ndarray,
updates: jnp.ndarray,
) -> jnp.ndarray:
"""Scatter chain updates into unique linked ``(kx, ky)`` rows."""
idx = jnp.asarray(idx_flat, dtype=jnp.int32)
target_t = jnp.moveaxis(target, -2, 0)
updates_t = jnp.moveaxis(updates, -2, 0)
idx = idx[:, None]
dnums = jax.lax.ScatterDimensionNumbers(
update_window_dims=tuple(range(1, updates_t.ndim)),
inserted_window_dims=(0,),
scatter_dims_to_operand_dims=(0,),
)
out_t = jax.lax.scatter(
target_t,
idx,
updates_t,
dnums,
unique_indices=True,
)
return jnp.moveaxis(out_t, 0, -2)
def _linked_fft_chain_update(
f_flat: jnp.ndarray,
idx_map: jnp.ndarray,
kz_link: jnp.ndarray,
*,
lead_shape: tuple[int, ...],
Nz: int,
operator: str,
) -> tuple[jnp.ndarray, jnp.ndarray]:
"""Apply one linked-chain FFT derivative or absolute-``k_z`` operator."""
if idx_map.ndim != 2:
raise ValueError("linked index maps must have shape (nChains, nLinks)")
nChains, nLinks = idx_map.shape
idx_flat = idx_map.reshape(-1)
f_link = jnp.take(f_flat, idx_flat, axis=-2)
f_link = f_link.reshape(*lead_shape, nChains, nLinks * Nz)
f_hat = jnp.fft.fft(f_link, axis=-1)
multiplier = (
_fft_ik_multiplier(kz_link, f_hat)
if operator == "grad"
else _fft_abs_multiplier(kz_link, f_hat)
)
df_hat = multiplier * f_hat
df_link = jnp.fft.ifft(df_hat, axis=-1)
df_link = df_link.reshape(*lead_shape, nChains * nLinks, Nz)
return idx_flat, jnp.asarray(df_link, dtype=f_flat.dtype)
def _linked_fft_chain_outputs(
f_flat: jnp.ndarray,
linked_indices: tuple[jnp.ndarray, ...],
linked_kz: tuple[jnp.ndarray, ...],
*,
lead_shape: tuple[int, ...],
Nz: int,
operator: str,
) -> tuple[list[jnp.ndarray], list[jnp.ndarray]]:
chain_updates: list[jnp.ndarray] = []
chain_indices: list[jnp.ndarray] = []
for idx_map, kz_link in zip(linked_indices, linked_kz):
idx_flat, df_link = _linked_fft_chain_update(
f_flat,
idx_map,
kz_link,
lead_shape=lead_shape,
Nz=Nz,
operator=operator,
)
chain_indices.append(idx_flat)
chain_updates.append(df_link)
return chain_indices, chain_updates
def _linked_fft_covered_rows(
chain_indices: list[jnp.ndarray],
*,
Ny: int,
) -> jnp.ndarray:
idx_cat = (
chain_indices[0]
if len(chain_indices) == 1
else jnp.concatenate(chain_indices, axis=0)
)
return jnp.zeros((Ny,), dtype=bool).at[jnp.mod(idx_cat, Ny)].set(True)
def _linked_fft_gather_output(
chain_updates: list[jnp.ndarray],
*,
linked_gather_map: jnp.ndarray | None,
linked_gather_mask: jnp.ndarray | None,
lead_shape: tuple[int, ...],
Nx: int,
Ny: int,
Nz: int,
) -> jnp.ndarray:
updates_cat = jnp.concatenate(chain_updates, axis=-2)
gather_map = jnp.asarray(linked_gather_map, dtype=jnp.int32)
gather_mask = jnp.asarray(linked_gather_mask, dtype=updates_cat.dtype)
updates_full = jnp.take(updates_cat, gather_map, axis=-2)
mask_shape = (1,) * (updates_full.ndim - 2) + (gather_mask.shape[0], 1)
updates_full = updates_full * gather_mask.reshape(mask_shape)
updates_full = updates_full.reshape(*lead_shape, Nx, Ny, Nz)
return jnp.swapaxes(updates_full, -3, -2)
def _linked_fft_full_cover_output(
chain_updates: list[jnp.ndarray],
*,
linked_inverse_permutation: jnp.ndarray | None,
lead_shape: tuple[int, ...],
Nx: int,
Ny: int,
Nz: int,
) -> jnp.ndarray:
if linked_inverse_permutation is None:
raise ValueError("linked_inverse_permutation required when linked_full_cover is True")
updates_cat = jnp.concatenate(chain_updates, axis=-2)
inv = jnp.asarray(linked_inverse_permutation, dtype=jnp.int32)
df_flat = jnp.take(updates_cat, inv, axis=-2)
df_full = df_flat.reshape(*lead_shape, Nx, Ny, Nz)
return jnp.swapaxes(df_full, -3, -2)
def _linked_fft_scatter_output(
f_flat: jnp.ndarray,
chain_indices: list[jnp.ndarray],
chain_updates: list[jnp.ndarray],
*,
lead_shape: tuple[int, ...],
Nx: int,
Ny: int,
Nz: int,
) -> jnp.ndarray:
df_flat = jnp.zeros_like(f_flat)
for idx_flat, df_link in zip(chain_indices, chain_updates):
df_flat = _scatter_unique_linked_modes(df_flat, idx_flat, df_link)
df_full = df_flat.reshape(*lead_shape, Nx, Ny, Nz)
return jnp.swapaxes(df_full, -3, -2)
def _linked_fft_apply(
f: jnp.ndarray,
linked_indices: tuple[jnp.ndarray, ...],
linked_kz: tuple[jnp.ndarray, ...],
*,
operator: str,
linked_inverse_permutation: jnp.ndarray | None = None,
linked_full_cover: bool = False,
linked_gather_map: jnp.ndarray | None = None,
linked_gather_mask: jnp.ndarray | None = None,
linked_use_gather: bool = False,
) -> jnp.ndarray:
_validate_linked_fft_inputs(linked_indices, linked_kz, operator=operator)
f_flat, lead_shape, Ny, Nx, Nz = _flatten_linked_fft_state(f)
chain_indices, chain_updates = _linked_fft_chain_outputs(
f_flat,
linked_indices,
linked_kz,
lead_shape=lead_shape,
Nz=Nz,
operator=operator,
)
covered_rows = _linked_fft_covered_rows(chain_indices, Ny=Ny)
if linked_use_gather:
out = _linked_fft_gather_output(
chain_updates,
linked_gather_map=linked_gather_map,
linked_gather_mask=linked_gather_mask,
lead_shape=lead_shape,
Nx=Nx,
Ny=Ny,
Nz=Nz,
)
return _restore_linked_real_fft_conjugates(out, covered_rows=covered_rows)
if linked_full_cover:
out = _linked_fft_full_cover_output(
chain_updates,
linked_inverse_permutation=linked_inverse_permutation,
lead_shape=lead_shape,
Nx=Nx,
Ny=Ny,
Nz=Nz,
)
return _restore_linked_real_fft_conjugates(out, covered_rows=covered_rows)
out = _linked_fft_scatter_output(
f_flat,
chain_indices,
chain_updates,
lead_shape=lead_shape,
Nx=Nx,
Ny=Ny,
Nz=Nz,
)
return _restore_linked_real_fft_conjugates(out, covered_rows=covered_rows)
def grad_z_linked_fft(
f: jnp.ndarray,
dz: float | jnp.ndarray,
linked_indices: tuple[jnp.ndarray, ...],
linked_kz: tuple[jnp.ndarray, ...],
linked_inverse_permutation: jnp.ndarray | None = None,
linked_full_cover: bool = False,
linked_gather_map: jnp.ndarray | None = None,
linked_gather_mask: jnp.ndarray | None = None,
linked_use_gather: bool = False,
) -> jnp.ndarray:
"""Spectral z-derivative using linked-chain FFT modes."""
_check_positive(dz, "dz")
return _linked_fft_apply(
f,
linked_indices,
linked_kz,
operator="grad",
linked_inverse_permutation=linked_inverse_permutation,
linked_full_cover=linked_full_cover,
linked_gather_map=linked_gather_map,
linked_gather_mask=linked_gather_mask,
linked_use_gather=linked_use_gather,
)
def abs_z_linked_fft(
f: jnp.ndarray,
linked_indices: tuple[jnp.ndarray, ...],
linked_kz: tuple[jnp.ndarray, ...],
linked_inverse_permutation: jnp.ndarray | None = None,
linked_full_cover: bool = False,
linked_gather_map: jnp.ndarray | None = None,
linked_gather_mask: jnp.ndarray | None = None,
linked_use_gather: bool = False,
) -> jnp.ndarray:
"""Apply |kz| in linked-FFT space."""
return _linked_fft_apply(
f,
linked_indices,
linked_kz,
operator="abs",
linked_inverse_permutation=linked_inverse_permutation,
linked_full_cover=linked_full_cover,
linked_gather_map=linked_gather_map,
linked_gather_mask=linked_gather_mask,
linked_use_gather=linked_use_gather,
)
[docs]
def shift_axis(arr: jnp.ndarray, offset: int, axis: int) -> jnp.ndarray:
"""Shift an array along an axis with zero padding (non-periodic)."""
axis = axis % arr.ndim
if offset == 0:
return arr
axis_len = arr.shape[axis]
if abs(offset) >= axis_len:
return jnp.zeros_like(arr)
out = jnp.zeros_like(arr)
if offset > 0:
body = jax.lax.slice_in_dim(arr, offset, axis_len, axis=axis)
starts = [0] * arr.ndim
starts[axis] = 0
return jax.lax.dynamic_update_slice(out, body, starts)
body = jax.lax.slice_in_dim(arr, 0, axis_len + offset, axis=axis)
starts = [0] * arr.ndim
starts[axis] = -offset
return jax.lax.dynamic_update_slice(out, body, starts)
def _shift_with_zeros(arr: jnp.ndarray, axis: int, offset: int) -> jnp.ndarray:
"""Shift along ``axis`` and fill the exposed entries with zeros."""
return shift_axis(arr, offset, axis)
[docs]
def apply_hermite_v(G: jnp.ndarray) -> jnp.ndarray:
"""Multiply Hermite coefficients by v_parallel (ladder form)."""
axis_m = -4
Nm = G.shape[axis_m]
sqrt_p, sqrt_m = hermite_ladder_coeffs(Nm - 1)
sqrt_p = sqrt_p[:Nm]
sqrt_m = sqrt_m[:Nm]
G_plus = _shift_with_zeros(G, axis_m, 1)
G_minus = _shift_with_zeros(G, axis_m, -1)
shape = [1] * G.ndim
shape[axis_m] = Nm
sqrt_p = sqrt_p.reshape(shape)
sqrt_m = sqrt_m.reshape(shape)
return sqrt_p * G_plus + sqrt_m * G_minus
[docs]
def apply_hermite_v2(G: jnp.ndarray) -> jnp.ndarray:
"""Multiply Hermite coefficients by v_parallel^2."""
return apply_hermite_v(apply_hermite_v(G))
[docs]
def apply_laguerre_x(G: jnp.ndarray) -> jnp.ndarray:
"""Multiply Laguerre coefficients by the perpendicular energy variable."""
axis_l = -5
Nl = G.shape[axis_l]
ell = jnp.arange(Nl)
G_plus = _shift_with_zeros(G, axis_l, 1)
G_minus = _shift_with_zeros(G, axis_l, -1)
l_shape = [1] * G.ndim
l_shape[axis_l] = Nl
l_col = ell.reshape(l_shape)
return (
(2.0 * l_col + 1.0) * G
- (l_col + 1.0) * G_plus
- l_col * G_minus
)
def streaming_ladder_term(
H: jnp.ndarray,
kz: jnp.ndarray,
vth: float | jnp.ndarray,
sqrt_p: jnp.ndarray,
sqrt_m: jnp.ndarray,
*,
dz: float | jnp.ndarray | None = None,
kx_link_plus: jnp.ndarray | None = None,
kx_link_minus: jnp.ndarray | None = None,
kx_mask_plus: jnp.ndarray | None = None,
kx_mask_minus: jnp.ndarray | None = None,
linked_indices: tuple[jnp.ndarray, ...] | None = None,
linked_kz: tuple[jnp.ndarray, ...] | None = None,
linked_inverse_permutation: jnp.ndarray | None = None,
linked_full_cover: bool = False,
linked_gather_map: jnp.ndarray | None = None,
linked_gather_mask: jnp.ndarray | None = None,
linked_use_gather: bool = False,
use_twist_shift: bool = False,
) -> jnp.ndarray:
"""Apply streaming with precomputed Hermite-ladder coefficients."""
_check_positive(vth, "vth")
if use_twist_shift:
if dz is None:
raise ValueError("dz must be provided for twist-shift boundaries")
if linked_indices is not None and linked_kz is not None:
dH_dz = grad_z_linked_fft(
H,
dz=dz,
linked_indices=linked_indices,
linked_kz=linked_kz,
linked_inverse_permutation=linked_inverse_permutation,
linked_full_cover=linked_full_cover,
linked_gather_map=linked_gather_map,
linked_gather_mask=linked_gather_mask,
linked_use_gather=linked_use_gather,
)
else:
if kx_link_plus is None or kx_link_minus is None:
raise ValueError("kx_link arrays must be provided for twist-shift boundaries")
if kx_mask_plus is None or kx_mask_minus is None:
raise ValueError("kx_link masks must be provided for twist-shift boundaries")
dH_dz = _grad_z_linked_fd(
H,
dz=dz,
kx_link_plus=kx_link_plus,
kx_link_minus=kx_link_minus,
kx_mask_plus=kx_mask_plus,
kx_mask_minus=kx_mask_minus,
)
else:
dH_dz = grad_z_periodic(H, kz=kz)
axis_m = -4
pad = [(0, 0)] * H.ndim
pad[axis_m] = (1, 1)
H_pad = jnp.pad(dH_dz, pad)
slc_plus = [slice(None)] * H.ndim
slc_minus = [slice(None)] * H.ndim
slc_plus[axis_m] = slice(2, None)
slc_minus[axis_m] = slice(0, -2)
H_plus = H_pad[tuple(slc_plus)]
H_minus = H_pad[tuple(slc_minus)]
return vth * (sqrt_p * H_plus + sqrt_m * H_minus)