/
bari
/
PlasmaTorch
Обзор
Документация
Войти
/
bari
/
PlasmaTorch
Код
Запросы
0
Задачи
Вики
Пакеты
0
Релизы
0
CI/CD
Аналитика
Безопасность
master
src/plast/solver.py
381 строка
15 KB
bari
Fix parameter name for Magnus check threshold in Solver class
19 июн 2026, 12:27
19 июн 2026, 12:27
196f73a
Код
Авторство
О чём код?
from __future__ import annotations from dataclasses import dataclass, field import warnings import torch from plast.plasma import Plasma from plast.pulse import Pulse @dataclass(slots=True) class SolveResult: R: torch.Tensor T: torch.Tensor A_nodes: torch.Tensor | None = None B_nodes: torch.Tensor | None = None u_nodes: torch.Tensor | None = None p_nodes: torch.Tensor | None = None @dataclass(slots=True) class Solver: plasma: Plasma pulse: Pulse adiabatic_threshold: float = 1e3 magnus_check_threshold: float = torch.pi permittivity: torch.Tensor = field(init=False) wavenumber_longitudinal: torch.Tensor = field(init=False) adiabatic_mask: torch.Tensor = field(init=False) def __post_init__(self) -> None: self._check_input_errors() self.permittivity = self.plasma._get_permittivity(self.pulse.fftfreq) self.wavenumber_longitudinal = self._get_wavenumber_longitudinal(self.pulse, self.permittivity) self.adiabatic_mask = self._get_adiabatic_mask( self.wavenumber_longitudinal, self.permittivity, self.plasma.grid, self.pulse.polarization, torch.sin(self.pulse.angle_incidence) ** 2, self.plasma.offsets, self.adiabatic_threshold, ) self._check_errors() self._check_magnus() def _check_input_errors(self) -> None: if not isinstance(self.plasma, Plasma): raise TypeError("plasma must be a Plasma") if not isinstance(self.pulse, Pulse): raise TypeError("pulse must be a Pulse") if not isinstance(self.adiabatic_threshold, (int, float)): raise TypeError("adiabatic_threshold must be a number") @staticmethod def _get_wavenumber_longitudinal(pulse: Pulse, permittivity: torch.Tensor) -> torch.Tensor: sin2 = torch.sin(pulse.angle_incidence)[:, None] ** 2 k0 = 2 * torch.pi * pulse.fftfreq[:, None] if permittivity.ndim == 1: return k0**2 * (permittivity[None, :] - permittivity[0] * sin2) return k0**2 * (permittivity - permittivity[:, :1] * sin2) @staticmethod def _get_adiabatic_mask( wavenumber_longitudinal: torch.Tensor, permittivity: torch.Tensor, grid: torch.Tensor, polarization: str, sin2: torch.Tensor, offsets: torch.Tensor, adiabatic_threshold: float = 1.0, ) -> torch.Tensor: mask = torch.ones_like(wavenumber_longitudinal, dtype=torch.bool) for start, stop in zip(offsets[:-1].tolist(), offsets[1:].tolist()): for index in range(start, stop - 1): step = grid[index + 1] - grid[index] wavenumber = wavenumber_longitudinal[:, index + 1] wavenumber_grad = (wavenumber - wavenumber_longitudinal[:, index]) / step xi, _ = Solver._get_xi_eta(permittivity, polarization, sin2, index + 1) xi_previous, _ = Solver._get_xi_eta(permittivity, polarization, sin2, index) xi_grad = (xi - xi_previous) / step value = torch.abs(wavenumber_grad / (2 * wavenumber) - xi_grad / (2 * xi)) / torch.abs(wavenumber) mask[:, index + 1] = value < adiabatic_threshold return mask @staticmethod def _get_xi_eta( permittivity: torch.Tensor, polarization: str, sin2: torch.Tensor, index: int, ) -> tuple[torch.Tensor, torch.Tensor]: if permittivity.ndim == 1: epsilon = permittivity[index] * torch.ones_like(sin2) epsilon_left = permittivity[0] * torch.ones_like(sin2) else: epsilon = permittivity[:, index] epsilon_left = permittivity[:, 0] if polarization in {"s", "TE"}: xi = torch.ones_like(sin2) eta = epsilon - epsilon_left * sin2 else: xi = epsilon eta = 1 - epsilon_left * sin2 / epsilon return xi, eta @staticmethod def _ab_to_up(A: torch.Tensor, B: torch.Tensor, k: torch.Tensor, xi: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: return A + B, 1j * k / xi * (A - B) @staticmethod def _up_to_ab(u: torch.Tensor, p: torch.Tensor, k: torch.Tensor, xi: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: return 0.5 * (u - 1j * xi / k * p), 0.5 * (u + 1j * xi / k * p) @staticmethod def _get_adiabatic_step( A: torch.Tensor, B: torch.Tensor, k_left: torch.Tensor, k_right: torch.Tensor, xi_left: torch.Tensor, xi_right: torch.Tensor, dx: torch.Tensor, ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]: gamma = k_left * xi_right / (k_right * xi_left) phase = torch.exp(-1j * k_left * dx) phase_inv = torch.exp(1j * k_left * dx) A_next = 0.5 / gamma * ((1 + gamma) * phase * A - (1 - gamma) * phase_inv * B) B_next = 0.5 / gamma * (-(1 - gamma) * phase * A + (1 + gamma) * phase_inv * B) u_next, p_next = Solver._ab_to_up(A_next, B_next, k_left, xi_left) return A_next, B_next, u_next, p_next @staticmethod def _get_physical_step( u: torch.Tensor, p: torch.Tensor, xi: torch.Tensor, eta: torch.Tensor, k0_sq: torch.Tensor, dx: torch.Tensor, ) -> tuple[torch.Tensor, torch.Tensor]: X = xi * dx Z = eta * dx u_next = u - X * p p_next = k0_sq * Z * u + (1 - k0_sq * X * Z) * p return u_next, p_next @staticmethod def _get_adiabatic_input( A: torch.Tensor, B: torch.Tensor, u: torch.Tensor, p: torch.Tensor, k: torch.Tensor, xi: torch.Tensor, current_interval_is_adiabatic: torch.Tensor, from_ab_to_ab: torch.Tensor, from_up_to_ab: torch.Tensor, ) -> tuple[torch.Tensor, torch.Tensor]: A_for_adiabatic = torch.empty_like(A[current_interval_is_adiabatic]) B_for_adiabatic = torch.empty_like(B[current_interval_is_adiabatic]) from_ab_to_ab_local = from_ab_to_ab[current_interval_is_adiabatic] if torch.any(from_ab_to_ab): A_for_adiabatic[from_ab_to_ab_local] = A[from_ab_to_ab] B_for_adiabatic[from_ab_to_ab_local] = B[from_ab_to_ab] if torch.any(from_up_to_ab): from_up_to_ab_local = from_up_to_ab[current_interval_is_adiabatic] restored_A, restored_B = Solver._up_to_ab( u[from_up_to_ab], p[from_up_to_ab], k[from_up_to_ab], xi[from_up_to_ab], ) A_for_adiabatic[from_up_to_ab_local] = restored_A B_for_adiabatic[from_up_to_ab_local] = restored_B return A_for_adiabatic, B_for_adiabatic @staticmethod def _store_node( index: int, A_nodes: torch.Tensor, B_nodes: torch.Tensor, u_nodes: torch.Tensor, p_nodes: torch.Tensor, A: torch.Tensor, B: torch.Tensor, u: torch.Tensor, p: torch.Tensor, ) -> None: A_nodes[:, index] = A B_nodes[:, index] = B u_nodes[:, index] = u p_nodes[:, index] = p def solve(self, track_all_nodes: bool = False, renorm_threshold: float = 1e6) -> SolveResult: grid = self.plasma.grid permittivity = self.permittivity k = torch.sqrt(self.wavenumber_longitudinal) k0_sq = (2 * torch.pi * self.pulse.fftfreq) ** 2 sin2 = torch.sin(self.pulse.angle_incidence) ** 2 n_freq, n_nodes = k.shape right_index = n_nodes - 1 A = torch.ones(n_freq, device=k.device, dtype=k.dtype) B = torch.zeros_like(A) xi_right, _ = self._get_xi_eta(permittivity, self.pulse.polarization, sin2, right_index) u, p = self._ab_to_up(A, B, k[:, right_index], xi_right) normalization = torch.ones_like(A) if track_all_nodes: A_nodes = torch.empty((n_freq, n_nodes), device=k.device, dtype=k.dtype) B_nodes = torch.empty_like(A_nodes) u_nodes = torch.empty_like(A_nodes) p_nodes = torch.empty_like(A_nodes) self._store_node(right_index, A_nodes, B_nodes, u_nodes, p_nodes, A, B, u, p) else: A_nodes = None B_nodes = None u_nodes = None p_nodes = None for index in range(right_index - 1, -1, -1): dx = grid[index + 1] - grid[index] current_interval_is_adiabatic = self.adiabatic_mask[:, index + 1] xi_left, eta_left = self._get_xi_eta(permittivity, self.pulse.polarization, sin2, index) xi_right, _ = self._get_xi_eta(permittivity, self.pulse.polarization, sin2, index + 1) xi_left = xi_left * torch.ones_like(sin2) xi_right = xi_right * torch.ones_like(sin2) k_left = k[:, index] k_right = k[:, index + 1] previous_adiabatic = torch.isfinite(A) from_ab_to_ab = previous_adiabatic & current_interval_is_adiabatic from_up_to_ab = ~previous_adiabatic & current_interval_is_adiabatic to_up = ~current_interval_is_adiabatic A_next = torch.full_like(A, torch.nan) B_next = torch.full_like(B, torch.nan) u_next = torch.empty_like(u) p_next = torch.empty_like(p) if torch.any(current_interval_is_adiabatic): A_for_adiabatic, B_for_adiabatic = self._get_adiabatic_input( A, B, u, p, k_right, xi_right, current_interval_is_adiabatic, from_ab_to_ab, from_up_to_ab, ) A_from_ab, B_from_ab, u_from_ab, p_from_ab = self._get_adiabatic_step( A_for_adiabatic, B_for_adiabatic, k_left[current_interval_is_adiabatic], k_right[current_interval_is_adiabatic], xi_left[current_interval_is_adiabatic], xi_right[current_interval_is_adiabatic], dx, ) A_next[current_interval_is_adiabatic] = A_from_ab B_next[current_interval_is_adiabatic] = B_from_ab u_next[current_interval_is_adiabatic] = u_from_ab p_next[current_interval_is_adiabatic] = p_from_ab if torch.any(to_up): u_from_up, p_from_up = self._get_physical_step( u[to_up], p[to_up], xi_left[to_up], eta_left[to_up], k0_sq[to_up], dx, ) u_next[to_up] = u_from_up p_next[to_up] = p_from_up A = A_next B = B_next u = u_next p = p_next renorm_mask = current_interval_is_adiabatic & (torch.abs(A) > renorm_threshold) if torch.any(renorm_mask): scale = torch.where(renorm_mask, A, torch.ones_like(A)) A = A / scale B = B / scale u = u / scale p = p / scale normalization = normalization / scale if track_all_nodes: A_nodes = A_nodes / scale[:, None] B_nodes = B_nodes / scale[:, None] u_nodes = u_nodes / scale[:, None] p_nodes = p_nodes / scale[:, None] if track_all_nodes: self._store_node(index, A_nodes, B_nodes, u_nodes, p_nodes, A, B, u, p) final_scale = A A = A / final_scale B = B / final_scale u = u / final_scale p = p / final_scale T = normalization / final_scale if track_all_nodes: A_nodes = A_nodes / final_scale[:, None] B_nodes = B_nodes / final_scale[:, None] u_nodes = u_nodes / final_scale[:, None] p_nodes = p_nodes / final_scale[:, None] A_nodes[:, 0] = A B_nodes[:, 0] = B return SolveResult( R=B, T=T, A_nodes=A_nodes, B_nodes=B_nodes, u_nodes=u_nodes, p_nodes=p_nodes, ) def _check_magnus(self) -> None: grid = self.plasma.grid offsets = self.plasma.offsets for start, stop in zip(offsets[:-1].tolist(), offsets[1:].tolist()): for index in range(start, stop - 1): step = grid[index + 1] - grid[index] k0 = 2 * torch.pi * self.pulse.fftfreq k = self.wavenumber_longitudinal[:, index] k_grad = (self.wavenumber_longitudinal[:, index + 1] - k) / step sin2 = torch.sin(self.pulse.angle_incidence) ** 2 ksi, eta = self._get_xi_eta(self.permittivity, self.pulse.polarization, sin2, index) ksi_next, _ = self._get_xi_eta(self.permittivity, self.pulse.polarization, sin2, index + 1) ksi_grad = (ksi_next - ksi) / step adiabatic_magnus = (torch.abs(k) + torch.abs(k_grad / (2 * k) - ksi_grad / (2 * ksi))) * step nonadiabatic_magnus = (torch.abs(ksi) + torch.abs(k0**2 * eta)) * step magnus = torch.where(self.adiabatic_mask[:, index + 1], adiabatic_magnus, nonadiabatic_magnus) if torch.any(magnus > self.magnus_check_threshold): warnings.warn( f"convergence not guaranteed for grid interval [{index}, {index + 1}]. " "It is recommended to linearly interpolate the grid in this interval.", UserWarning, stacklevel=2, ) def _check_errors(self) -> None: self._check_input_errors() if not isinstance(self.permittivity, torch.Tensor): raise TypeError("permittivity must be a torch.Tensor") if self.permittivity.ndim not in {1, 2}: raise ValueError("permittivity must be 1D or 2D") if self.permittivity.ndim == 1 and self.permittivity.shape[0] != self.plasma.grid.shape[0]: raise ValueError("1D permittivity must have the same length as grid") if self.permittivity.ndim == 2 and self.permittivity.shape != self.wavenumber_longitudinal.shape: raise ValueError("2D permittivity must have shape [frequency, grid]") if not isinstance(self.wavenumber_longitudinal, torch.Tensor): raise TypeError("wavenumber_longitudinal must be a torch.Tensor") if not isinstance(self.adiabatic_mask, torch.Tensor): raise TypeError("adiabatic_mask must be a torch.Tensor")