/
bari
/
PlasmaTorch
Обзор
Документация
Войти
/
bari
/
PlasmaTorch
Код
Запросы
0
Задачи
Вики
Пакеты
0
Релизы
0
CI/CD
Аналитика
Безопасность
master
test/test_solver.py
206 строк
8 KB
bari
Update plasma state equation API
18 июн 2026, 12:45
18 июн 2026, 12:45
8b10769
Код
Авторство
О чём код?
import pytest import torch from plast.plasma import Plasma from plast.pulse import Pulse from plast.solver import Solver def test_solver_stores_inputs(): # This test checks that the solver keeps the given objects. grid = torch.tensor([0.0, 0.0, 1.0, 1.0]) permittivity = torch.tensor([1.0, 1.0, 1.1, 1.1]) plasma = Plasma(grid=grid, permittivity=permittivity) pulse = Pulse( fft=torch.tensor([1.0]), fftfreq=torch.tensor([0.0]), angle_incidence=torch.tensor([0.2]), ) solver = Solver(plasma=plasma, pulse=pulse) assert solver.plasma is plasma assert solver.pulse is pulse def test_solver_rejects_bad_plasma(): # This test checks that plasma must be a Plasma object. pulse = Pulse( fft=torch.tensor([1.0]), fftfreq=torch.tensor([0.0]), angle_incidence=torch.tensor([0.2]), ) with pytest.raises(TypeError, match="plasma"): Solver(plasma="bad", pulse=pulse) def test_solver_computes_term(): # This test checks the solver builds the longitudinal wavenumber. grid = torch.tensor([0.0, 0.0, 1.0, 1.0]) permittivity = torch.tensor([4.0, 4.0, 4.0, 4.0]) plasma = Plasma(grid=grid, permittivity=permittivity) pulse = Pulse( fft=torch.tensor([1.0, 2.0]), fftfreq=torch.tensor([0.01, 0.02]), angle_incidence=torch.tensor([0.0, torch.pi / 2]), ) solver = Solver(plasma=plasma, pulse=pulse) expected = (2 * torch.pi * pulse.fftfreq[:, None]) ** 2 * ( solver.permittivity[None, :] - solver.permittivity[0] * torch.sin(pulse.angle_incidence)[:, None] ** 2 ) assert torch.equal(solver.wavenumber_longitudinal, expected) def test_solver_accepts_frequency_dependent_permittivity(): # This test checks that permittivity can be different for each pulse frequency. grid = torch.tensor([0.0, 0.0, 1.0, 1.0]) frequency = torch.tensor([0.0, 0.0, 0.001, 0.001]) fftfreq = torch.tensor([0.01, 0.02]) pulse = Pulse( fft=torch.ones(2, dtype=torch.complex64), fftfreq=fftfreq, angle_incidence=torch.zeros(2), polarization="TE", ) def state_equation(frequency, fftfreq): return 1.0 - frequency[None, :] ** 2 / fftfreq[:, None] ** 2 plasma = Plasma(grid=grid, frequency=frequency, state_equation=state_equation) solver = Solver(plasma=plasma, pulse=pulse) assert solver.permittivity.shape == (2, 4) assert solver.wavenumber_longitudinal.shape == (2, 4) def test_solver_with_explicit_pulse_and_plasma_values(): # This test checks the solver on the exact values from the request. grid = torch.tensor([0.0, 0.0, 1.0, 2.0, 3.0, 4.0, 4.0]) permittivity = torch.tensor([1.0, 2.0, 2.0, 2.0, 2.0, 2.0, 1.0]) plasma = Plasma(grid=grid, permittivity=permittivity) pulse = Pulse( fft=torch.tensor([1.0, 2.0, 3.0, 4.0]), fftfreq=torch.tensor([1.0, 2.0, 3.0, 4.0]), angle_incidence=torch.tensor([0.0, torch.pi / 6, torch.pi / 4, torch.pi / 3]), ) with pytest.warns(UserWarning, match="convergence not guaranteed"): solver = Solver(plasma=plasma, pulse=pulse) expected = (2 * torch.pi * pulse.fftfreq[:, None]) ** 2 * ( solver.permittivity[None, :] - solver.permittivity[0] * torch.sin(pulse.angle_incidence)[:, None] ** 2 ) assert torch.equal(solver.wavenumber_longitudinal, expected) def test_solver_warns_for_magnus_check(): # This test checks that the solver warns on the inner grid points. grid = torch.tensor([0.0, 0.0, 1.0, 2.0, 3.0, 4.0, 4.0]) permittivity = torch.tensor([10.0, 10.0, 10.0, 10.0, 10.0, 10.0, 10.0]) plasma = Plasma(grid=grid, permittivity=permittivity) pulse = Pulse( fft=torch.tensor([10.0]), fftfreq=torch.tensor([2.0]), angle_incidence=torch.tensor([0.0]), ) with pytest.warns(UserWarning, match="convergence not guaranteed"): Solver(plasma=plasma, pulse=pulse) def test_solver_builds_adiabatic_mask_for_s_polarization(): # This test checks the adiabatic mask for s polarization. grid = torch.tensor([0.0, 0.0, 1.0, 2.0, 3.0, 4.0, 4.0]) permittivity = torch.tensor([2.0, 2.0, 2.0, 2.0, 2.0, 2.0, 2.0]) plasma = Plasma(grid=grid, permittivity=permittivity) pulse = Pulse( fft=torch.tensor([1.0]), fftfreq=torch.tensor([0.02]), angle_incidence=torch.tensor([0.0]), polarization="s", ) solver = Solver(plasma=plasma, pulse=pulse, adiabatic_threshold=1.0) expected = torch.ones_like(solver.adiabatic_mask, dtype=torch.bool) start, stop = plasma.offsets[1].item(), plasma.offsets[2].item() for index in range(start, stop - 1): step = grid[index + 1] - grid[index] wavenumber = solver.wavenumber_longitudinal[:, index + 1] wavenumber_grad = (wavenumber - solver.wavenumber_longitudinal[:, index]) / step xi, _ = solver._get_xi_eta(solver.permittivity, pulse.polarization, torch.sin(pulse.angle_incidence) ** 2, index + 1) xi_previous, _ = solver._get_xi_eta(solver.permittivity, pulse.polarization, torch.sin(pulse.angle_incidence) ** 2, index) xi_grad = (xi - xi_previous) / step expected[:, index + 1] = torch.abs(wavenumber_grad / (2 * wavenumber) - xi_grad / (2 * xi)) / torch.abs(wavenumber) < 1.0 assert torch.equal(solver.adiabatic_mask, expected) def test_solver_builds_adiabatic_mask_for_p_polarization(): # This test checks the adiabatic mask for p polarization. grid = torch.tensor([0.0, 0.0, 1.0, 2.0, 3.0, 4.0, 4.0]) permittivity = torch.tensor([1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0]) plasma = Plasma(grid=grid, permittivity=permittivity) pulse = Pulse( fft=torch.tensor([1.0]), fftfreq=torch.tensor([0.02]), angle_incidence=torch.tensor([0.0]), polarization="p", ) solver = Solver(plasma=plasma, pulse=pulse, adiabatic_threshold=1.0) expected = torch.ones_like(solver.adiabatic_mask, dtype=torch.bool) start, stop = plasma.offsets[1].item(), plasma.offsets[2].item() for index in range(start, stop - 1): step = grid[index + 1] - grid[index] wavenumber = solver.wavenumber_longitudinal[:, index + 1] wavenumber_grad = (wavenumber - solver.wavenumber_longitudinal[:, index]) / step xi, _ = solver._get_xi_eta(solver.permittivity, pulse.polarization, torch.sin(pulse.angle_incidence) ** 2, index + 1) xi_previous, _ = solver._get_xi_eta(solver.permittivity, pulse.polarization, torch.sin(pulse.angle_incidence) ** 2, index) xi_grad = (xi - xi_previous) / step expected[:, index + 1] = torch.abs(wavenumber_grad / (2 * wavenumber) - xi_grad / (2 * xi)) / torch.abs(wavenumber) < 1.0 assert torch.equal(solver.adiabatic_mask, expected) def test_solver_builds_short_adiabatic_mask_example(): # This test checks one short mask example with three pulse points. grid = torch.tensor([0.0, 0.0, 1.0, 2.0, 3.0, 4.0, 4.0, 5.0, 6.0, 7.0, 7.0]) permittivity = torch.tensor([1.0, 1.0, 0.5, 0.5, 0.5, 0.5, 1.0, 1.0, 1.0, 1.0, 1.0]) plasma = Plasma(grid=grid, permittivity=permittivity) pulse = Pulse( fft=torch.tensor([10.0, 1.0, 0.1]), fftfreq=torch.tensor([0.01, 0.1, 0.2]), angle_incidence=torch.tensor([0.9, 0.9, 0.9]), polarization="s", ) solver = Solver(plasma=plasma, pulse=pulse) expected = torch.tensor( [ [True, True, False, True, True, True, True, True, True, True, True], [True, True, True, True, True, True, True, True, True, True, True], [True, True, True, True, True, True, True, True, True, True, True], ] ) assert torch.equal(solver.adiabatic_mask, expected) def test_solver_checks_two_point_piece(): # This test checks that a two point piece is also used. grid = torch.tensor([0.0, 0.0, 1.0, 2.0, 2.0, 3.0, 3.0]) permittivity = torch.tensor([0.01, 0.01, 0.01, 0.01, 1.0, 1.0, 1.0]) plasma = Plasma(grid=grid, permittivity=permittivity) pulse = Pulse( fft=torch.tensor([1.0]), fftfreq=torch.tensor([1.0]), angle_incidence=torch.tensor([0.0]), polarization="s", ) with pytest.warns(UserWarning, match=r"\[4, 5\]"): Solver(plasma=plasma, pulse=pulse)