diff --git a/.gitattributes b/.gitattributes new file mode 100644 index 00000000..887a2c18 --- /dev/null +++ b/.gitattributes @@ -0,0 +1,2 @@ +# SCM syntax highlighting & preventing 3-way merges +pixi.lock merge=binary linguist-language=YAML linguist-generated=true diff --git a/.github/workflows/build.yml b/.github/workflows/build.yml new file mode 100644 index 00000000..271e3b04 --- /dev/null +++ b/.github/workflows/build.yml @@ -0,0 +1,115 @@ +name: Build wheels + +# We don't want to spend CI compute cycles on every commit, so restrict building +# of wheels to published releases. This includes stable and beta/RC releases. +on: + release: + types: + - published + +jobs: + # Build a wheel for each Python version in a separate parallel job. + build_wheels: + name: Build wheels for ${{ matrix.python }} on ${{ matrix.target.name }} + runs-on: ${{ matrix.target.runner }} + strategy: + matrix: + python: [cp39, cp310, cp311, cp312, cp313, cp314] + target: + - name: Linux (x86-64) + runner: ubuntu-latest + platform: manylinux_x86_64 + arch: x86_64 + + # TODO: add other platforms once Raysect provides wheels on them. + # - name: Linux (ARM) + # runner: ubuntu-24.04-arm + # platform: manylinux_aarch64 + # arch: aarch64 + + # - name: macOS (Apple Silicon) + # runner: macos-latest + # platform: macosx_arm64 + # arch: arm64 + + # - name: macOS (Intel) + # runner: macos-15-intel + # platform: macosx_x86_64 + # arch: x86_64 + + # - name: Windows + # runner: windows-2022 + # platform: win_amd64 + # arch: AMD64 + + steps: + - uses: actions/checkout@v7 + with: + persist-credentials: false + + - name: Determine manylinux image + if: ${{ contains(matrix.target.name, 'Linux') }} + # Build on manylinux2014 where possible for greatest compatibility, + # but need newer manylinux_2_28 for Python >=3.14. + # TODO: use cp314|cp315 if we provide Python 3.15 wheels before dropping + # manylinux2014 (RHEL7) support. + # TODO: remove this step entirely once we drop manylinux2014 support. + run: | + case ${{ matrix.python }} in + cp314) + echo "CIBW_MANYLINUX_X86_64_IMAGE=manylinux_2_28" >> "$GITHUB_ENV" + ;; + *) + echo "CIBW_MANYLINUX_X86_64_IMAGE=manylinux2014" >> "$GITHUB_ENV" + ;; + esac + + - name: Build wheels + uses: pypa/cibuildwheel@v4.2.0 + env: + CIBW_BUILD: ${{ matrix.python }}-${{ matrix.target.platform }} + CIBW_ARCH: ${{ matrix.target.arch }} + # Avoid trying to compile e.g. Pillow (raysect->matplotlib dependency) from + # source on older manylinux with newer Python. An old version is fine. + CIBW_ENVIRONMENT: 'PIP_PREFER_BINARY=1' + + - uses: actions/upload-artifact@v4 + with: + name: dist-wheel-${{ matrix.target.platform }}-${{ matrix.python }} + path: ./wheelhouse/*.whl + + make_sdist: + name: Make SDist + runs-on: ubuntu-latest + steps: + - uses: actions/checkout@v6 + with: + persist-credentials: false + + - name: Build sdist + run: pipx run build --sdist + + - uses: actions/upload-artifact@v4 + with: + name: dist-sdist + path: dist/*.tar.gz + + # Wheels are built in parallel jobs: gather the outputs into a single artifact + # for more convenient retrieval. At some stage we may consider automatic + # upload to PyPI from this job (since its `needs` entry ensures the SDist and + # wheels all built successfully). + gather_artifacts: + name: Gather SDist and wheel outputs into a single artifact + runs-on: ubuntu-latest + needs: [build_wheels, make_sdist] + steps: + - uses: actions/download-artifact@v5 + with: + pattern: dist-* + path: dist + merge-multiple: true + + - uses: actions/upload-artifact@v4 + with: + name: artifacts + path: dist/* diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index b32bde9d..120c0701 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -3,16 +3,21 @@ name: CI on: push: pull_request: + types: + - opened + - synchronize + - reopened + - ready_for_review jobs: tests: + if: ${{ github.event_name == 'push' || !github.event.pull_request.draft }} name: Run tests runs-on: ubuntu-latest strategy: fail-fast: false matrix: - numpy-version: ["oldest-supported-numpy", "'numpy<2'"] - python-version: ["3.7", "3.8", "3.9", "3.10"] + python-version: ["3.9", "3.10", "3.11", "3.12", "3.13", "3.14"] steps: - name: Checkout code uses: actions/checkout@v2 @@ -23,9 +28,9 @@ jobs: with: python-version: ${{ matrix.python-version }} - name: Install Python dependencies - run: python -m pip install --prefer-binary cython~=3.0 ${{ matrix.numpy-version }} scipy matplotlib "pyopencl[pocl]>=2022.2.4" + run: python -m pip install --prefer-binary setuptools cython~=3.1 numpy>=2 scipy matplotlib "pyopencl[pocl]>=2022.2.4" - name: Install Raysect from pypi - run: pip install raysect==0.8.1.* + run: pip install raysect==0.9.* - name: Build cherab run: dev/build.sh - name: Run tests diff --git a/.gitignore b/.gitignore index 7045bc61..ea448d57 100644 --- a/.gitignore +++ b/.gitignore @@ -21,4 +21,11 @@ build/ .nfs* .coverage htmlcov* -cherab.egg-info/ \ No newline at end of file +cherab.egg-info/ + +# pixi environments +.pixi/* +!.pixi/config.toml + +# Ignore lock files until we start using them +pixi.lock diff --git a/CHANGELOG.md b/CHANGELOG.md index 5fe4cc06..d1cd0320 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -1,6 +1,36 @@ Project Changelog ================= +Release 1.6.0 (Sep 2026) +------------------- + +API changes: +* Rename `GaussianQuadrature` to `GaussianQuadrature1D` to conform to Cherab's naming convention. Original name kept as an alias for backwards compatibility until the next minor release. (#475) +* Rename `TargettedPixelGroup` to `TargetedPixelGroup` for correct spelling. Still keep `TargettedPixelGroup` as an alias for backwards compatibility until the next major release. (#487) +* Add emission model attribute access to line and lineshape . (#294) +* The `generate_derivative_operators` function in `admt_utils` can now return sparse matrices rather than dense if requested. (#427) +* Only 1 of the 1D-to-2D or 2D-to-1D voxel mappings is now required for `admt_utils.generate_derivative_operators`: if the other is missing it is computed automatically. (#427) +* The `calculate_admt` function in `admt_utils` will return a sparse matrix if the input derivative operators are themselves sparse. (#427) + +Bug fixes: +* Fix the import statement for `netcdf_file` in `calcam.py` for compatibility with the upcoming `scipy` v2.0.0. (#510) + +New: +* Add an optional Pixi workspace for package builds and isolated development environments, with tasks for testing, documentation, formatting, and static analysis, and include the corresponding developer guide in the Sphinx documentation. (#489) +* Add GaussianQuadrature2D integrator. (#475) +* Support Raysect 0.9. (#486) +* Test against Python 3.9, 3.10, 3.11, 3.12, 3.13 and latest released Numpy. Drop Python 3.7, 3.8 and older Numpy from tests. (#486) +* Add generic distribution function. (#481) +* Add Function6D framework. (#478) +* Add e_field attribute to Plasma object for electric field vector. (#465) +* Add Integrator2D base class for integration of two-dimensional functions. (#472) +* Support Raysect 0.9. (#486) +* Make values in `cherab.core.utility.constants` accessible to Python. (#509) +* Generomak now contains an example bolometer diagnostic. (#427) +* The regularisation utilities in `admt_utils` are now in the HTML documention. (#427) +* A new non-negative least squares inversion using sparse matrices, to complement the existing dense version. (#427) +* A demo performing bolometry inversions using both isotropic and anisotropic regularisation. (#427) + Release 1.5.0 (27 Aug 2024) ------------------- @@ -126,7 +156,7 @@ API changes: New: * Merged cherab-openadas package into the core cherab package to simplify installation. -* Beam object uses a cone primitive instead of a cylinder for the bounding volume of divergent beams. +* Beam object uses a cone primitive instead of a cylinder for the bounding volume of divergent beams. * Added Clamp functions. * Added ThermalCXRate. * Added optimised ray transfer grid calculation tools. @@ -154,7 +184,7 @@ New: Bug fixes: * Improved handling on non c-order arrays in various methods. -* Numerous minor bug fixes (see commit history) +* Numerous minor bug fixes (see commit history) Release 1.0.1 (1 Oct 2018) diff --git a/README.md b/README.md index 177feb9b..69e076d8 100644 --- a/README.md +++ b/README.md @@ -85,12 +85,19 @@ of code / physics algorithm quality standards. TMC Members ----------- +- Jack Lovell (Oak Ridge National Laboratory, USA) +- Matej Tomes (Institute of Plasma Physics of the Czech Academy of Sciences, Czechia) +- Koyo Munechika (ITER Organisation) +- Jakub Svoboda (Institute of Plasma Physics of the Czech Academy of Sciences, Czechia) + + +Former TMC Members +----------- + - Alys Brett (chairwoman, master account holder, responsible for delegation, UKAEA, UK) - Matt Carr (External consultant, diagnostic physics models) -- Jack Lovell (Oak Ridge, USA) - Alex Meakins (External consultant, Architecture, software integrity) - Vlad Neverov (NRC Kurchatov Institute, Moscow) -- Matej Tomes (Compass, IPP, Prague) Citing The Code diff --git a/cherab/core/VERSION b/cherab/core/VERSION index bc80560f..40ab7ec2 100644 --- a/cherab/core/VERSION +++ b/cherab/core/VERSION @@ -1 +1 @@ -1.5.0 +1.6.0rc1 diff --git a/cherab/core/atomic/gaunt.pxd b/cherab/core/atomic/gaunt.pxd index 827831b6..a18ca802 100644 --- a/cherab/core/atomic/gaunt.pxd +++ b/cherab/core/atomic/gaunt.pxd @@ -21,7 +21,7 @@ from cherab.core.math cimport Function2D cdef class FreeFreeGauntFactor(): - cpdef double evaluate(self, double z, double temperature, double wavelength) except? -1e999 + cdef double evaluate(self, double z, double temperature, double wavelength) except? -1e999 cdef class InterpolatedFreeFreeGauntFactor(FreeFreeGauntFactor): diff --git a/cherab/core/atomic/gaunt.pyx b/cherab/core/atomic/gaunt.pyx index eb1dbe59..001ff214 100644 --- a/cherab/core/atomic/gaunt.pyx +++ b/cherab/core/atomic/gaunt.pyx @@ -37,9 +37,9 @@ cdef class FreeFreeGauntFactor(): The base class for temperature-averaged free-free Gaunt factors. """ - cpdef double evaluate(self, double z, double temperature, double wavelength) except? -1e999: + cdef double evaluate(self, double z, double temperature, double wavelength) except? -1e999: """ - Returns the temperature-averaged free-free Gaunt factor for the supplied parameters. + Return the temperature-averaged free-free Gaunt factor for the supplied parameters. :param double z: Species charge or effective plasma charge. :param double temperature: Electron temperature in eV. @@ -51,7 +51,7 @@ cdef class FreeFreeGauntFactor(): def __call__(self, double z, double temperature, double wavelength): """ - Returns the temperature-averaged free-free Gaunt factor for the supplied parameters. + Return the temperature-averaged free-free Gaunt factor for the supplied parameters. :param double z: Species charge or effective plasma charge. :param double temperature: Electron temperature in eV. @@ -106,9 +106,9 @@ cdef class InterpolatedFreeFreeGauntFactor(FreeFreeGauntFactor): self._gaunt_factor = Interpolator2DArray(np.log10(u), np.log10(gamma2), gaunt_factor, 'cubic', 'none', 0, 0) @cython.cdivision(True) - cpdef double evaluate(self, double z, double temperature, double wavelength) except? -1e999: + cdef double evaluate(self, double z, double temperature, double wavelength) except? -1e999: """ - Returns the temperature-averaged free-free Gaunt factor for the supplied parameters. + Return the temperature-averaged free-free Gaunt factor for the supplied parameters. :param double z: Species charge or effective plasma charge. :param double temperature: Electron temperature in eV. diff --git a/cherab/core/atomic/line.pyx b/cherab/core/atomic/line.pyx index 262c7a2f..cbd07bba 100644 --- a/cherab/core/atomic/line.pyx +++ b/cherab/core/atomic/line.pyx @@ -33,6 +33,10 @@ cdef class Line: specify the n-levels with integers (e.g. (3,2)). For all other ions the full spectroscopic configuration string should be specified for both states. It is up to the atomic data provider package to define the exact notation. + + :ivar Element element: See parameter 'element'. + :ivar int charge: See parameter 'charge'. + :ivar tuple transition: See parameter 'transition'. .. code-block:: pycon diff --git a/cherab/core/atomic/tests/__init__.py b/cherab/core/atomic/tests/__init__.py new file mode 100644 index 00000000..e69de29b diff --git a/cherab/core/atomic/tests/test_line.py b/cherab/core/atomic/tests/test_line.py new file mode 100644 index 00000000..f268b82b --- /dev/null +++ b/cherab/core/atomic/tests/test_line.py @@ -0,0 +1,27 @@ +import unittest + +from cherab.core.atomic import Line, deuterium + + +class TestLine(unittest.TestCase): + + def test_initialisation(self): + line = Line(deuterium, 0, (3, 2)) + self.assertEqual(line.element, deuterium) + self.assertEqual(line.charge, 0) + self.assertEqual(line.transition, (3, 2)) + + # test invalid charge + with self.assertRaises(ValueError): + Line(deuterium, 2, (3, 2)) + with self.assertRaises(ValueError): + Line(deuterium, -1, (3, 2)) + + def test_properties(self): + element = deuterium + charge = 0 + transition = (3, 2) + line = Line(element, charge, transition) + self.assertEqual(line.element, element) + self.assertEqual(line.charge, charge) + self.assertEqual(line.transition, transition) \ No newline at end of file diff --git a/cherab/core/atomic/tests/test_zeeman.py b/cherab/core/atomic/tests/test_zeeman.py new file mode 100644 index 00000000..f2e6a62f --- /dev/null +++ b/cherab/core/atomic/tests/test_zeeman.py @@ -0,0 +1,70 @@ +import unittest + +import numpy as np +from raysect.core.math.function.float import Arg1D, Constant1D + +from cherab.core.atomic import ZeemanStructure + + +class TestZeemanStructure(unittest.TestCase): + def test_initialisation_rejects_invalid_component_shape(self): + with self.assertRaises(ValueError): + ZeemanStructure([(Constant1D(656.1),)], [], []) + + with self.assertRaises(ValueError): + ZeemanStructure([], [(Constant1D(656.1),)], []) + + with self.assertRaises(ValueError): + ZeemanStructure([], [], [(Constant1D(656.1),)]) + + def test_call_returns_expected_components_and_normalised_ratios(self): + pi_components = [ + (Constant1D(656.1), Constant1D(2.0)), + (Constant1D(656.2), Constant1D(6.0)), + ] + sigma_plus_components = [ + (656.0 + 0.01 * Arg1D(), Constant1D(3.0)), + (656.3 + 0.02 * Arg1D(), Constant1D(1.0)), + ] + sigma_minus_components = [ + (Constant1D(655.9), Constant1D(1.0)), + (Constant1D(656.4), Constant1D(1.0)), + ] + zeeman = ZeemanStructure(pi_components, sigma_plus_components, sigma_minus_components) + + b = 2.0 + + pi = zeeman(b, 'PI') + np.testing.assert_allclose(pi[0], np.array([656.1, 656.2])) + np.testing.assert_allclose(pi[1], np.array([0.25, 0.75])) + + sigma_plus = zeeman(b, 'SIGMA_PLUS') + np.testing.assert_allclose(sigma_plus[0], np.array([656.02, 656.34])) + np.testing.assert_allclose(sigma_plus[1], np.array([0.75, 0.25])) + + sigma_minus = zeeman(b, 'sigma_minus') + np.testing.assert_allclose(sigma_minus[0], np.array([655.9, 656.4])) + np.testing.assert_allclose(sigma_minus[1], np.array([0.5, 0.5])) + + def test_call_keeps_zero_ratios_when_sum_is_zero(self): + zeeman = ZeemanStructure( + pi_components=[ + (Constant1D(656.1), Constant1D(0.0)), + (Constant1D(656.2), Constant1D(0.0)), + ], + sigma_plus_components=[], + sigma_minus_components=[], + ) + + pi = zeeman(0.0, 'pi') + np.testing.assert_allclose(pi[0], np.array([656.1, 656.2])) + np.testing.assert_allclose(pi[1], np.array([0.0, 0.0])) + + def test_call_raises_for_invalid_arguments(self): + zeeman = ZeemanStructure([], [], []) + + with self.assertRaises(ValueError): + zeeman(-1.0, 'pi') + + with self.assertRaises(ValueError): + zeeman(1.0, 'sigma') diff --git a/cherab/core/atomic/zeeman.pyx b/cherab/core/atomic/zeeman.pyx index 68382d32..10c1928a 100644 --- a/cherab/core/atomic/zeeman.pyx +++ b/cherab/core/atomic/zeeman.pyx @@ -139,4 +139,4 @@ cdef class ZeemanStructure(): if polarisation.lower() == 'sigma_minus': return np.asarray(self.evaluate(b, SIGMA_MINUS_POLARISATION)) - raise ValueError('Argument "polarisation" must be "pi", "sigma_plus" or "sigma_minus", {} given.'.fotmat(polarisation)) + raise ValueError('Argument "polarisation" must be "pi", "sigma_plus" or "sigma_minus", {} given.'.format(polarisation)) diff --git a/cherab/core/distribution.pxd b/cherab/core/distribution.pxd index 062d7a45..af6f3844 100644 --- a/cherab/core/distribution.pxd +++ b/cherab/core/distribution.pxd @@ -19,6 +19,7 @@ from raysect.optical cimport Vector3D from cherab.core.math cimport Function3D, VectorFunction3D +from cherab.core.math.function.float cimport Function6D, autowrap_function6d cdef class DistributionFunction: @@ -45,3 +46,10 @@ cdef class Maxwellian(DistributionFunction): VectorFunction3D _velocity double _atomic_mass + +cdef class GenericDistribution(DistributionFunction): + + cdef readonly: + Function6D _phase_space_density + Function3D _density, _temperature + VectorFunction3D _velocity \ No newline at end of file diff --git a/cherab/core/distribution.pyx b/cherab/core/distribution.pyx index 869cb373..6f8bad61 100644 --- a/cherab/core/distribution.pyx +++ b/cherab/core/distribution.pyx @@ -25,6 +25,7 @@ from raysect.optical cimport Vector3D cimport cython from cherab.core.math cimport autowrap_function3d, autowrap_vectorfunction3d +from cherab.core.math.function.float cimport autowrap_function6d from cherab.core.utility.constants cimport ELEMENTARY_CHARGE @@ -301,3 +302,105 @@ cdef class Maxwellian(DistributionFunction): return self._density.evaluate(x, y, z) +cdef class GenericDistribution(DistributionFunction): + """ + A generic distribution function. + + This class implements a generic distribution function. The user supplies a 6D function + that provides the phase space density at a given point in 6D phase space, + a 3D function that provides the spatial density, a 3D function that provides the temperature, + and a 3D vector function that provides the bulk velocity. + + .. warning:: + The consistency of the provided functions is not checked and is the responsibilty of the user. + + :param Function6D phase_space_density: 6D function defining the phase space density in s^3/m^6. + :param Function3D density: 3D function defining the spatial density in m^-3. + :param Function3D temperature: 3D function defining the temperature in eV. + :param VectorFunction3D velocity: 3D vector function defining the bulk velocity in meters per second. + + .. code-block:: pycon + + >>> from cherab.core import GenericDistribution + >>> from cherab.core.math import Function6D, Function3D, VectorFunction3D + >>> + >>> # Setup distribution for a slab of plasma in thermodynamic equilibrium + >>> phase_space_density = Function6D(lambda x, y, z, vx, vy, vz: 1E17 * exp(-(vx**2 + vy**2 + vz**2) / (2 * 1))) + >>> density = Function3D(lambda x, y, z: 1E17) + >>> temperature = Function3D(lambda x, y, z: 1) + >>> velocity = VectorFunction3D(lambda x, y, z: Vector3D(0, 0, 0)) + >>> d0_distribution = GenericDistribution(phase_space_density, density, temperature, velocity) + """ + + def __init__(self, object phase_space_density, object density, object temperature, object velocity): + + super().__init__() + self._phase_space_density = autowrap_function6d(phase_space_density) + self._density = autowrap_function3d(density) + self._temperature = autowrap_function3d(temperature) + self._velocity = autowrap_vectorfunction3d(velocity) + + @cython.cdivision(True) + cdef double evaluate(self, double x, double y, double z, double vx, double vy, double vz) except? -1e999: + """ + Evaluates the phase space density at the specified point in 6D phase space. + + :param x: position in meters + :param y: position in meters + :param z: position in meters + :param vx: velocity in meters per second + :param vy: velocity in meters per second + :param vz: velocity in meters per second + :return: phase space density in s^3/m^6 + """ + + return self._phase_space_density.evaluate(x, y, z, vx, vy, vz) + + cpdef Vector3D bulk_velocity(self, double x, double y, double z): + """ + Evaluates the species' bulk velocity at the specified 3D coordinate. + + :param x: position in meters + :param y: position in meters + :param z: position in meters + :return: velocity vector in m/s + + .. code-block:: pycon + + >>> d0_distribution.bulk_velocity(1, 0, 0) + Vector3D(0.0, 0.0, 0.0) + """ + + return self._velocity.evaluate(x, y, z) + + cpdef double effective_temperature(self, double x, double y, double z) except? -1e999: + """ + + :param x: position in meters + :param y: position in meters + :param z: position in meters + :return: temperature in eV + + .. code-block:: pycon + + >>> d0_distribution.effective_temperature(1, 0, 0) + 1.0 + """ + + return self._temperature.evaluate(x, y, z) + + cpdef double density(self, double x, double y, double z) except? -1e999: + """ + + :param x: position in meters + :param y: position in meters + :param z: position in meters + :return: density in m^-3 + + .. code-block:: pycon + + >>> d0_distribution.density(1, 0, 0) + 1e+17 + """ + + return self._density.evaluate(x, y, z) \ No newline at end of file diff --git a/cherab/core/math/__init__.py b/cherab/core/math/__init__.py index 85336afa..567f540d 100644 --- a/cherab/core/math/__init__.py +++ b/cherab/core/math/__init__.py @@ -39,3 +39,4 @@ from .transform import CylindricalTransform, VectorCylindricalTransform from .transform import PeriodicTransform1D, PeriodicTransform2D, PeriodicTransform3D from .transform import VectorPeriodicTransform1D, VectorPeriodicTransform2D, VectorPeriodicTransform3D +from .function import * \ No newline at end of file diff --git a/cherab/core/math/caching/utility.py b/cherab/core/math/caching/utility.py index 84b910fc..d0fc342e 100644 --- a/cherab/core/math/caching/utility.py +++ b/cherab/core/math/caching/utility.py @@ -16,10 +16,10 @@ # See the Licence for the specific language governing permissions and limitations # under the Licence. -import numpy as np import matplotlib.pyplot as plt +import numpy as np -from core.math.caching import Caching2D +from . import Caching2D def auto_caching2d_optimiser(function2d, space_area, threshold): @@ -45,13 +45,12 @@ def auto_caching2d_optimiser(function2d, space_area, threshold): to_plot_resy = [] while current_error >= threshold: - resolutionx /= 2 resolutiony /= 2 to_plot_resx.append(resolutionx) to_plot_resy.append(resolutiony) cached_function = Caching2D(function2d, space_area, (resolutionx, resolutiony)) - current_error = 0. + current_error = 0.0 nb_zeros = 0 for x in np.linspace(minx, maxx, nb_samplesx): for y in np.linspace(miny, maxy, nb_samplesy): @@ -91,12 +90,11 @@ def mapping_caching2d_resolution(function2d, space_area): for i in range(20): for j in range(20): - print(i, j) resolutionx = resolutionsx[i] resolutiony = resolutionsy[j] cached_function = Caching2D(function2d, space_area, (resolutionx, resolutiony)) - error = 0. + error = 0.0 nb_zeros = 0 for x in np.linspace(minx, maxx, nb_samplesx): for y in np.linspace(miny, maxy, nb_samplesy): @@ -114,4 +112,4 @@ def mapping_caching2d_resolution(function2d, space_area): plt.xscale('log') plt.yscale('log') plt.colorbar() - plt.show() \ No newline at end of file + plt.show() diff --git a/cherab/core/math/function/__init__.pxd b/cherab/core/math/function/__init__.pxd index 9e5d8343..c4b2bbbe 100644 --- a/cherab/core/math/function/__init__.pxd +++ b/cherab/core/math/function/__init__.pxd @@ -16,6 +16,8 @@ # See the Licence for the specific language governing permissions and limitations # under the Licence. +from cherab.core.math.function cimport float + from raysect.core.math.function.float cimport Function1D, autowrap_function1d from raysect.core.math.function.float cimport Function2D, autowrap_function2d from raysect.core.math.function.float cimport Function3D, autowrap_function3d diff --git a/cherab/core/math/function/__init__.py b/cherab/core/math/function/__init__.py index 952e204c..98241d7a 100644 --- a/cherab/core/math/function/__init__.py +++ b/cherab/core/math/function/__init__.py @@ -16,6 +16,8 @@ # See the Licence for the specific language governing permissions and limitations # under the Licence. +from . import float + from raysect.core.math.function.float import Function1D, Function2D, Function3D from raysect.core.math.function.float import Constant1D, Constant2D, Constant3D from raysect.core.math.function.float import Discrete2DMesh, Interpolator2DMesh diff --git a/cherab/core/math/function/float/__init__.pxd b/cherab/core/math/function/float/__init__.pxd new file mode 100644 index 00000000..78e33ae5 --- /dev/null +++ b/cherab/core/math/function/float/__init__.pxd @@ -0,0 +1,32 @@ +# cython: language_level=3 + +# Copyright (c) 2014-2023, Dr Alex Meakins, Raysect Project +# All rights reserved. +# +# Redistribution and use in source and binary forms, with or without +# modification, are permitted provided that the following conditions are met: +# +# 1. Redistributions of source code must retain the above copyright notice, +# this list of conditions and the following disclaimer. +# +# 2. Redistributions in binary form must reproduce the above copyright +# notice, this list of conditions and the following disclaimer in the +# documentation and/or other materials provided with the distribution. +# +# 3. Neither the name of the Raysect Project nor the names of its +# contributors may be used to endorse or promote products derived from +# this software without specific prior written permission. +# +# THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS" +# AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE +# IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE +# ARE DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE +# LIABLE FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR +# CONSEQUENTIAL DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF +# SUBSTITUTE GOODS OR SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS +# INTERRUPTION) HOWEVER CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN +# CONTRACT, STRICT LIABILITY, OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) +# ARISING IN ANY WAY OUT OF THE USE OF THIS SOFTWARE, EVEN IF ADVISED OF THE +# POSSIBILITY OF SUCH DAMAGE. + +from cherab.core.math.function.float.function6d cimport * diff --git a/cherab/core/math/function/float/__init__.py b/cherab/core/math/function/float/__init__.py new file mode 100644 index 00000000..330f73c8 --- /dev/null +++ b/cherab/core/math/function/float/__init__.py @@ -0,0 +1,32 @@ +# cython: language_level=3 + +# Copyright (c) 2014-2023, Dr Alex Meakins, Raysect Project +# All rights reserved. +# +# Redistribution and use in source and binary forms, with or without +# modification, are permitted provided that the following conditions are met: +# +# 1. Redistributions of source code must retain the above copyright notice, +# this list of conditions and the following disclaimer. +# +# 2. Redistributions in binary form must reproduce the above copyright +# notice, this list of conditions and the following disclaimer in the +# documentation and/or other materials provided with the distribution. +# +# 3. Neither the name of the Raysect Project nor the names of its +# contributors may be used to endorse or promote products derived from +# this software without specific prior written permission. +# +# THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS" +# AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE +# IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE +# ARE DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE +# LIABLE FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR +# CONSEQUENTIAL DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF +# SUBSTITUTE GOODS OR SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS +# INTERRUPTION) HOWEVER CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN +# CONTRACT, STRICT LIABILITY, OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) +# ARISING IN ANY WAY OUT OF THE USE OF THIS SOFTWARE, EVEN IF ADVISED OF THE +# POSSIBILITY OF SUCH DAMAGE. + +from .function6d import * diff --git a/cherab/core/math/function/float/function6d/__init__.pxd b/cherab/core/math/function/float/function6d/__init__.pxd new file mode 100644 index 00000000..2e3336e9 --- /dev/null +++ b/cherab/core/math/function/float/function6d/__init__.pxd @@ -0,0 +1,37 @@ +# cython: language_level=3 + +# Copyright (c) 2014-2023, Dr Alex Meakins, Raysect Project +# All rights reserved. +# +# Redistribution and use in source and binary forms, with or without +# modification, are permitted provided that the following conditions are met: +# +# 1. Redistributions of source code must retain the above copyright notice, +# this list of conditions and the following disclaimer. +# +# 2. Redistributions in binary form must reproduce the above copyright +# notice, this list of conditions and the following disclaimer in the +# documentation and/or other materials provided with the distribution. +# +# 3. Neither the name of the Raysect Project nor the names of its +# contributors may be used to endorse or promote products derived from +# this software without specific prior written permission. +# +# THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS" +# AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE +# IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE +# ARE DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE +# LIABLE FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR +# CONSEQUENTIAL DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF +# SUBSTITUTE GOODS OR SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS +# INTERRUPTION) HOWEVER CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN +# CONTRACT, STRICT LIABILITY, OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) +# ARISING IN ANY WAY OUT OF THE USE OF THIS SOFTWARE, EVEN IF ADVISED OF THE +# POSSIBILITY OF SUCH DAMAGE. + +from cherab.core.math.function.float.function6d.base cimport Function6D +from cherab.core.math.function.float.function6d.constant cimport Constant6D +from cherab.core.math.function.float.function6d.blend cimport Blend6D +from cherab.core.math.function.float.function6d.autowrap cimport autowrap_function6d +from cherab.core.math.function.float.function6d.arg cimport Arg6D +from cherab.core.math.function.float.function6d.cmath cimport * diff --git a/cherab/core/math/function/float/function6d/__init__.py b/cherab/core/math/function/float/function6d/__init__.py new file mode 100644 index 00000000..0ababca8 --- /dev/null +++ b/cherab/core/math/function/float/function6d/__init__.py @@ -0,0 +1,36 @@ +# cython: language_level=3 + +# Copyright (c) 2014-2023, Dr Alex Meakins, Raysect Project +# All rights reserved. +# +# Redistribution and use in source and binary forms, with or without +# modification, are permitted provided that the following conditions are met: +# +# 1. Redistributions of source code must retain the above copyright notice, +# this list of conditions and the following disclaimer. +# +# 2. Redistributions in binary form must reproduce the above copyright +# notice, this list of conditions and the following disclaimer in the +# documentation and/or other materials provided with the distribution. +# +# 3. Neither the name of the Raysect Project nor the names of its +# contributors may be used to endorse or promote products derived from +# this software without specific prior written permission. +# +# THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS" +# AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE +# IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE +# ARE DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE +# LIABLE FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR +# CONSEQUENTIAL DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF +# SUBSTITUTE GOODS OR SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS +# INTERRUPTION) HOWEVER CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN +# CONTRACT, STRICT LIABILITY, OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) +# ARISING IN ANY WAY OUT OF THE USE OF THIS SOFTWARE, EVEN IF ADVISED OF THE +# POSSIBILITY OF SUCH DAMAGE. + +from .base import Function6D +from .constant import Constant6D +from .blend import Blend6D +from .arg import Arg6D +from .cmath import * diff --git a/cherab/core/math/function/float/function6d/arg.pxd b/cherab/core/math/function/float/function6d/arg.pxd new file mode 100644 index 00000000..527a9c0b --- /dev/null +++ b/cherab/core/math/function/float/function6d/arg.pxd @@ -0,0 +1,27 @@ +# cython: language_level=3 + +# Copyright 2016-2025 Euratom +# Copyright 2016-2025 United Kingdom Atomic Energy Authority +# Copyright 2016-2025 Centro de Investigaciones Energéticas, Medioambientales y Tecnológicas +# +# Licensed under the EUPL, Version 1.1 or – as soon they will be approved by the +# European Commission - subsequent versions of the EUPL (the "Licence"); +# You may not use this work except in compliance with the Licence. +# You may obtain a copy of the Licence at: +# +# https://joinup.ec.europa.eu/software/page/eupl5 +# +# Unless required by applicable law or agreed to in writing, software distributed +# under the Licence is distributed on an "AS IS" basis, WITHOUT WARRANTIES OR +# CONDITIONS OF ANY KIND, either express or implied. +# +# See the Licence for the specific language governing permissions and limitations +# under the Licence. + +from cherab.core.math.function.float.function6d.base cimport Function6D + +cdef enum ArgLabel: + X, Y, Z, U, V, W + +cdef class Arg6D(Function6D): + cdef ArgLabel _argument diff --git a/cherab/core/math/function/float/function6d/arg.pyx b/cherab/core/math/function/float/function6d/arg.pyx new file mode 100644 index 00000000..f49053c7 --- /dev/null +++ b/cherab/core/math/function/float/function6d/arg.pyx @@ -0,0 +1,82 @@ +# cython: language_level=3 + +# Copyright 2016-2025 Euratom +# Copyright 2016-2025 United Kingdom Atomic Energy Authority +# Copyright 2016-2025 Centro de Investigaciones Energéticas, Medioambientales y Tecnológicas +# +# Licensed under the EUPL, Version 1.1 or – as soon they will be approved by the +# European Commission - subsequent versions of the EUPL (the "Licence"); +# You may not use this work except in compliance with the Licence. +# You may obtain a copy of the Licence at: +# +# https://joinup.ec.europa.eu/software/page/eupl5 +# +# Unless required by applicable law or agreed to in writing, software distributed +# under the Licence is distributed on an "AS IS" basis, WITHOUT WARRANTIES OR +# CONDITIONS OF ANY KIND, either express or implied. +# +# See the Licence for the specific language governing permissions and limitations +# under the Licence. + +from cherab.core.math.function.float.function6d.base cimport Function6D + + +cdef class Arg6D(Function6D): + """ + Returns one of the arguments the function is passed, unmodified + + This is used to pass coordinates through to other functions in the + function framework which expect a Function6D object. + + Valid options for argument are "x", "y", "z", "u", "v", or "w". + + >>> argx = Arg6D("x") + >>> argx(2, 3, 5, 7, 11, 13) + 2.0 + >>> argy = Arg6D("y") + >>> argy(2, 3, 5, 7, 11, 13) + 3.0 + >>> argz = Arg6D("z") + >>> argz(2, 3, 5, 7, 11, 13) + 5.0 + >>> argu = Arg6D("u") + >>> argu(2, 3, 5, 7, 11, 13) + 7.0 + >>> argv = Arg6D("v") + >>> argv(2, 3, 5, 7, 11, 13) + 11.0 + >>> argw = Arg6D("w") + >>> argw(2, 3, 5, 7, 11, 13) + 13.0 + + :param str argument: either "x", "y", "z", "u", "v", or "w", the argument to return + """ + def __init__(self, object argument): + if argument == "x": + self._argument = X + elif argument == "y": + self._argument = Y + elif argument == "z": + self._argument = Z + elif argument == "u": + self._argument = U + elif argument == "v": + self._argument = V + elif argument == "w": + self._argument = W + else: + raise ValueError("The argument to Arg6D must be either 'x', 'y', 'z', 'u', 'v' or 'w'") + + cdef double evaluate(self, double x, double y, double z, double u, double v, double w) except? -1e999: + if self._argument == X: + return x + elif self._argument == Y: + return y + elif self._argument == Z: + return z + elif self._argument == U: + return u + elif self._argument == V: + return v + else: # W + return w diff --git a/cherab/core/math/function/float/function6d/autowrap.pxd b/cherab/core/math/function/float/function6d/autowrap.pxd new file mode 100644 index 00000000..a630393c --- /dev/null +++ b/cherab/core/math/function/float/function6d/autowrap.pxd @@ -0,0 +1,26 @@ +# cython: language_level=3 + +# Copyright 2016-2025 Euratom +# Copyright 2016-2025 United Kingdom Atomic Energy Authority +# Copyright 2016-2025 Centro de Investigaciones Energéticas, Medioambientales y Tecnológicas +# +# Licensed under the EUPL, Version 1.1 or – as soon they will be approved by the +# European Commission - subsequent versions of the EUPL (the "Licence"); +# You may not use this work except in compliance with the Licence. +# You may obtain a copy of the Licence at: +# +# https://joinup.ec.europa.eu/software/page/eupl5 +# +# Unless required by applicable law or agreed to in writing, software distributed +# under the Licence is distributed on an "AS IS" basis, WITHOUT WARRANTIES OR +# CONDITIONS OF ANY KIND, either express or implied. +# +# See the Licence for the specific language governing permissions and limitations +# under the Licence. + +from cherab.core.math.function.float.function6d.base cimport Function6D + +cdef class PythonFunction6D(Function6D): + cdef public object function + +cdef Function6D autowrap_function6d(object obj) diff --git a/cherab/core/math/function/float/function6d/autowrap.pyx b/cherab/core/math/function/float/function6d/autowrap.pyx new file mode 100644 index 00000000..080f2a96 --- /dev/null +++ b/cherab/core/math/function/float/function6d/autowrap.pyx @@ -0,0 +1,98 @@ +# cython: language_level=3 + +# Copyright 2016-2025 Euratom +# Copyright 2016-2025 United Kingdom Atomic Energy Authority +# Copyright 2016-2025 Centro de Investigaciones Energéticas, Medioambientales y Tecnológicas +# +# Licensed under the EUPL, Version 1.1 or – as soon they will be approved by the +# European Commission - subsequent versions of the EUPL (the "Licence"); +# You may not use this work except in compliance with the Licence. +# You may obtain a copy of the Licence at: +# +# https://joinup.ec.europa.eu/software/page/eupl5 +# +# Unless required by applicable law or agreed to in writing, software distributed +# under the Licence is distributed on an "AS IS" basis, WITHOUT WARRANTIES OR +# CONDITIONS OF ANY KIND, either express or implied. +# +# See the Licence for the specific language governing permissions and limitations +# under the Licence. + +import numbers +from cherab.core.math.function.float.function6d.base cimport Function6D +from cherab.core.math.function.float.function6d.constant cimport Constant6D +from raysect.core.math.function.base cimport Function + + +cdef class PythonFunction6D(Function6D): + """ + Wraps a python callable object with a Function6D object. + + This class allows a python object to interact with cython code that requires + a Function6D object. The python object must implement __call__() expecting + six arguments. + + This class is intended to be used to transparently wrap python objects that + are passed via constructors or methods into cython optimised code. It is not + intended that the users should need to directly interact with these wrapping + objects. Constructors and methods expecting a Function6D object should be + designed to accept a generic python object and then test that object to + determine if it is an instance of Function6D. If the object is not a + Function6D object it should be wrapped using this class for internal use. + + See also: autowrap_function6d() + """ + + def __init__(self, object function): + self.function = function + + cdef double evaluate(self, double x, double y, double z, double u, double v, double w) except? -1e999: + return self.function(x, y, z, u, v, w) + + +cdef Function6D autowrap_function6d(object obj): + """ + Automatically wraps the supplied python object in a PythonFunction6D or Constant6D object. + + If this function is passed a valid Function6D object, then the Function6D + object is simply returned without wrapping. + + If this function is passed a numerical scalar (int or float), a Constant6D + object is returned. + + This convenience function is provided to simplify the handling of Function6D + and python callable objects in constructors, functions and setters. + """ + + if isinstance(obj, Function6D): + return obj + elif isinstance(obj, Function): + raise TypeError('A Function6D object is required.') + elif isinstance(obj, numbers.Real): + return Constant6D(obj) + else: + return PythonFunction6D(obj) + + +def _autowrap_function6d(obj): + """Expose cython function for testing.""" + return autowrap_function6d(obj) + + +cdef inline bint is_callable(object f): + """ + Tests if an object is a python callable or a Function6D object. + """ + print(f"Checking if callable:", f) + print(f"isinstance(Function6D):", isinstance(f, Function6D)) + print(f"isinstance(Function):", isinstance(f, Function)) + print(f"callable():", callable(f)) + + if isinstance(f, Function6D): + return True + + # other function classes are incompatible + if isinstance(f, Function): + return False + + return callable(f) diff --git a/cherab/core/math/function/float/function6d/base.pxd b/cherab/core/math/function/float/function6d/base.pxd new file mode 100644 index 00000000..ef05f0f7 --- /dev/null +++ b/cherab/core/math/function/float/function6d/base.pxd @@ -0,0 +1,165 @@ +# cython: language_level=3 + +# Copyright 2016-2025 Euratom +# Copyright 2016-2025 United Kingdom Atomic Energy Authority +# Copyright 2016-2025 Centro de Investigaciones Energéticas, Medioambientales y Tecnológicas +# +# Licensed under the EUPL, Version 1.1 or – as soon they will be approved by the +# European Commission - subsequent versions of the EUPL (the "Licence"); +# You may not use this work except in compliance with the Licence. +# You may obtain a copy of the Licence at: +# +# https://joinup.ec.europa.eu/software/page/eupl5 +# +# Unless required by applicable law or agreed to in writing, software distributed +# under the Licence is distributed on an "AS IS" basis, WITHOUT WARRANTIES OR +# CONDITIONS OF ANY KIND, either express or implied. +# +# See the Licence for the specific language governing permissions and limitations +# under the Licence. + +from raysect.core.math.function.base cimport Function +from raysect.core.math.function.float.base cimport FloatFunction + + +cdef class Function6D(FloatFunction): + cdef double evaluate(self, double x, double y, double z, double u, double v, double w) except? -1e999 + + +cdef class AddFunction6D(Function6D): + cdef Function6D _function1, _function2 + + +cdef class SubtractFunction6D(Function6D): + cdef Function6D _function1, _function2 + + +cdef class MultiplyFunction6D(Function6D): + cdef Function6D _function1, _function2 + + +cdef class DivideFunction6D(Function6D): + cdef Function6D _function1, _function2 + + +cdef class ModuloFunction6D(Function6D): + cdef Function6D _function1, _function2 + + +cdef class PowFunction6D(Function6D): + cdef Function6D _function1, _function2 + + +cdef class AbsFunction6D(Function6D): + cdef Function6D _function + + +cdef class EqualsFunction6D(Function6D): + cdef Function6D _function1, _function2 + + +cdef class NotEqualsFunction6D(Function6D): + cdef Function6D _function1, _function2 + + +cdef class LessThanFunction6D(Function6D): + cdef Function6D _function1, _function2 + + +cdef class GreaterThanFunction6D(Function6D): + cdef Function6D _function1, _function2 + + +cdef class LessEqualsFunction6D(Function6D): + cdef Function6D _function1, _function2 + + +cdef class GreaterEqualsFunction6D(Function6D): + cdef Function6D _function1, _function2 + + +cdef class AddScalar6D(Function6D): + cdef double _value + cdef Function6D _function + + +cdef class SubtractScalar6D(Function6D): + cdef double _value + cdef Function6D _function + + +cdef class MultiplyScalar6D(Function6D): + cdef double _value + cdef Function6D _function + + +cdef class DivideScalar6D(Function6D): + cdef double _value + cdef Function6D _function + + +cdef class ModuloScalarFunction6D(Function6D): + cdef double _value + cdef Function6D _function + + +cdef class ModuloFunctionScalar6D(Function6D): + cdef double _value + cdef Function6D _function + + +cdef class PowScalarFunction6D(Function6D): + cdef double _value + cdef Function6D _function + + +cdef class PowFunctionScalar6D(Function6D): + cdef double _value + cdef Function6D _function + + +cdef class EqualsScalar6D(Function6D): + cdef double _value + cdef Function6D _function + + +cdef class NotEqualsScalar6D(Function6D): + cdef double _value + cdef Function6D _function + + +cdef class LessThanScalar6D(Function6D): + cdef double _value + cdef Function6D _function + + +cdef class GreaterThanScalar6D(Function6D): + cdef double _value + cdef Function6D _function + + +cdef class LessEqualsScalar6D(Function6D): + cdef double _value + cdef Function6D _function + + +cdef class GreaterEqualsScalar6D(Function6D): + cdef double _value + cdef Function6D _function + + +cdef inline bint is_callable(object f): + """ + Tests if an object is a python callable or a Function6D object. + + :param object f: Object to test. + :return: True if callable, False otherwise. + """ + if isinstance(f, Function6D): + return True + + # other function classes are incompatible + if isinstance(f, Function): + return False + + return callable(f) \ No newline at end of file diff --git a/cherab/core/math/function/float/function6d/base.pyx b/cherab/core/math/function/float/function6d/base.pyx new file mode 100644 index 00000000..a97d9583 --- /dev/null +++ b/cherab/core/math/function/float/function6d/base.pyx @@ -0,0 +1,737 @@ +# cython: language_level=3 + +# Copyright 2016-2025 Euratom +# Copyright 2016-2025 United Kingdom Atomic Energy Authority +# Copyright 2016-2025 Centro de Investigaciones Energéticas, Medioambientales y Tecnológicas +# +# Licensed under the EUPL, Version 1.1 or – as soon they will be approved by the +# European Commission - subsequent versions of the EUPL (the "Licence"); +# You may not use this work except in compliance with the Licence. +# You may obtain a copy of the Licence at: +# +# https://joinup.ec.europa.eu/software/page/eupl5 +# +# Unless required by applicable law or agreed to in writing, software distributed +# under the Licence is distributed on an "AS IS" basis, WITHOUT WARRANTIES OR +# CONDITIONS OF ANY KIND, either express or implied. +# +# See the Licence for the specific language governing permissions and limitations +# under the Licence. + +import numbers +from cpython.object cimport Py_LT, Py_EQ, Py_GT, Py_LE, Py_NE, Py_GE +cimport cython +from libc.math cimport floor +from .autowrap cimport autowrap_function6d + + +cdef class Function6D(FloatFunction): + """ + Cython optimised class for representing an arbitrary 6D function returning a float. + + Using __call__() in cython is slow. This class provides an overloadable + cython cdef evaluate() method which has much less overhead than a python + function call. + + For use in cython code only, this class cannot be extended via python. + + To create a new function object, inherit this class and implement the + evaluate() method. The new function object can then be used with any code + that accepts a function object. + """ + + cdef double evaluate(self, double x, double y, double z, double u, double v, double w) except? -1e999: + raise NotImplementedError("The evaluate() method has not been implemented.") + + def __call__(self, double x, double y, double z, double u, double v, double w): + """ Evaluate the function f(x, y, z, u, v, w) + + :param float x: function parameter x + :param float y: function parameter y + :param float z: function parameter z + :param float u: function parameter u + :param float v: function parameter v + :param float w: function parameter w + :rtype: float + """ + return self.evaluate(x, y, z, u, v, w) + def __add__(self, object b): + if is_callable(b): + # a() + b() + return AddFunction6D(self, b) + elif isinstance(b, numbers.Real): + # a() + B -> B + a() + return AddScalar6D( b, self) + return NotImplemented + + def __radd__(self, object a): + return self.__add__(a) + + def __sub__(self, object b): + if is_callable(b): + # a() - b() + return SubtractFunction6D(self, b) + elif isinstance(b, numbers.Real): + # a() - B -> -B + a() + return AddScalar6D(-( b), self) + return NotImplemented + + def __rsub__(self, object a): + if is_callable(a): + # a() - b() + return SubtractFunction6D(a, self) + elif isinstance(a, numbers.Real): + # A - b() + return SubtractScalar6D( a, self) + return NotImplemented + + def __mul__(self, object b): + if is_callable(b): + # a() * b() + return MultiplyFunction6D(self, b) + elif isinstance(b, numbers.Real): + # a() * B -> B * a() + return MultiplyScalar6D( b, self) + return NotImplemented + + def __rmul__(self, object a): + return self.__mul__(a) + + @cython.cdivision(True) + def __truediv__(self, object b): + cdef double v + if is_callable(b): + # a() / b() + return DivideFunction6D(self, b) + elif isinstance(b, numbers.Real): + # a() / B -> 1/B * a() + v = b + if v == 0.0: + raise ZeroDivisionError("Scalar used as the denominator of the division is zero valued.") + return MultiplyScalar6D(1/v, self) + return NotImplemented + + @cython.cdivision(True) + def __rtruediv__(self, object a): + if is_callable(a): + # a() / b() + return DivideFunction6D(a, self) + elif isinstance(a, numbers.Real): + # A / b() + return DivideScalar6D( a, self) + return NotImplemented + + def __mod__(self, object b): + cdef double v + if is_callable(b): + # a() % b() + return ModuloFunction6D(self, b) + elif isinstance(b, numbers.Real): + # a() % B + v = b + if v == 0.0: + raise ZeroDivisionError("Scalar used as the divisor of the division is zero valued.") + return ModuloFunctionScalar6D(self, v) + return NotImplemented + + def __rmod__(self, object a): + if is_callable(a): + # a() % b() + return ModuloFunction6D(a, self) + elif isinstance(a, numbers.Real): + # A % b() + return ModuloScalarFunction6D( a, self) + return NotImplemented + + def __neg__(self): + return MultiplyScalar6D(-1, self) + + def __pow__(self, object b, object c): + if c is not None: + # Optimised implementation of pow(a, b, c) not available: fall back + # to general implementation + return PowFunction6D(self, b) % c + if is_callable(b): + # a() ** b() + return PowFunction6D(self, b) + elif isinstance(b, numbers.Real): + # a() ** b + return PowFunctionScalar6D(self, b) + return NotImplemented + + def __rpow__(self, object a, object c): + if c is not None: + # Optimised implementation of pow(a, b, c) not available: fall back + # to general implementation + return PowFunction6D(a, self) % c + if is_callable(a): + # a() ** b() + return PowFunction6D(a, self) + elif isinstance(a, numbers.Real): + # A ** b() + return PowScalarFunction6D( a, self) + return NotImplemented + + def __abs__(self): + return AbsFunction6D(self) + + def __richcmp__(self, object other, int op): + if is_callable(other): + if op == Py_EQ: + return EqualsFunction6D(self, other) + if op == Py_NE: + return NotEqualsFunction6D(self, other) + if op == Py_LT: + return LessThanFunction6D(self, other) + if op == Py_GT: + return GreaterThanFunction6D(self, other) + if op == Py_LE: + return LessEqualsFunction6D(self, other) + if op == Py_GE: + return GreaterEqualsFunction6D(self, other) + if isinstance(other, numbers.Real): + if op == Py_EQ: + return EqualsScalar6D( other, self) + if op == Py_NE: + return NotEqualsScalar6D( other, self) + if op == Py_LT: + # f() < K -> K > f + return GreaterThanScalar6D( other, self) + if op == Py_GT: + # f() > K -> K < f + return LessThanScalar6D( other, self) + if op == Py_LE: + # f() <= K -> K >= f + return GreaterEqualsScalar6D( other, self) + if op == Py_GE: + # f() >= K -> K <= f + return LessEqualsScalar6D( other, self) + return NotImplemented + + +cdef class AddFunction6D(Function6D): + """ + A Function6D class that implements the addition of the results of two Function6D objects: f1() + f2() + + This class is not intended to be used directly, but rather returned as the result of an __add__() call on a + Function6D object. + + :param function1: A Function6D object. + :param function2: A Function6D object. + """ + + def __init__(self, object function1, object function2): + self._function1 = autowrap_function6d(function1) + self._function2 = autowrap_function6d(function2) + + cdef double evaluate(self, double x, double y, double z, double u, double v, double w) except? -1e999: + return self._function1.evaluate(x, y, z, u, v, w) + self._function2.evaluate(x, y, z, u, v, w) + + +cdef class SubtractFunction6D(Function6D): + """ + A Function6D class that implements the subtraction of the results of two Function6D objects: f1() - f2() + + This class is not intended to be used directly, but rather returned as the result of a __sub__() call on a + Function6D object. + + :param function1: A Function6D object. + :param function2: A Function6D object. + """ + + def __init__(self, object function1, object function2): + self._function1 = autowrap_function6d(function1) + self._function2 = autowrap_function6d(function2) + + cdef double evaluate(self, double x, double y, double z, double u, double v, double w) except? -1e999: + return self._function1.evaluate(x, y, z, u, v, w) - self._function2.evaluate(x, y, z, u, v, w) + + +cdef class MultiplyFunction6D(Function6D): + """ + A Function6D class that implements the multiplication of the results of two Function6D objects: f1() * f2() + + This class is not intended to be used directly, but rather returned as the result of a __mul__() call on a + Function6D object. + + :param function1: A Function6D object. + :param function2: A Function6D object. + """ + + def __init__(self, function1, function2): + self._function1 = autowrap_function6d(function1) + self._function2 = autowrap_function6d(function2) + + cdef double evaluate(self, double x, double y, double z, double u, double v, double w) except? -1e999: + return self._function1.evaluate(x, y, z, u, v, w) * self._function2.evaluate(x, y, z, u, v, w) + + +cdef class DivideFunction6D(Function6D): + """ + A Function6D class that implements the division of the results of two Function6D objects: f1() / f2() + + This class is not intended to be used directly, but rather returned as the result of a __truediv__() call on a + Function6D object. + + :param function1: A Function6D object. + :param function2: A Function6D object. + """ + + def __init__(self, function1, function2): + self._function1 = autowrap_function6d(function1) + self._function2 = autowrap_function6d(function2) + + @cython.cdivision(True) + cdef double evaluate(self, double x, double y, double z, double u, double v, double w) except? -1e999: + cdef double denominator = self._function2.evaluate(x, y, z, u, v, w) + if denominator == 0.0: + raise ZeroDivisionError("Function used as the denominator of the division returned a zero value.") + return self._function1.evaluate(x, y, z, u, v, w) / denominator + + +cdef class ModuloFunction6D(Function6D): + """ + A Function6D class that implements the modulo of the results of two Function6D objects: f1() % f2() + + This class is not intended to be used directly, but rather returned as the result of a __mod__() call on a + Function6D object. + + :param object function1: A Function6D object or Python callable. + :param object function2: A Function6D object or Python callable. + """ + def __init__(self, function1, function2): + self._function1 = autowrap_function6d(function1) + self._function2 = autowrap_function6d(function2) + + @cython.cdivision(True) + cdef double evaluate(self, double x, double y, double z, double u, double v, double w) except? -1e999: + cdef double divisor = self._function2.evaluate(x, y, z, u, v, w) + if divisor == 0.0: + raise ZeroDivisionError("Function used as the divisor of the modulo returned a zero value.") + return self._function1.evaluate(x, y, z, u, v, w) % divisor + + +cdef class PowFunction6D(Function6D): + """ + A Function6D class that implements the pow() operator on two Function6D objects. + + This class is not intended to be used directly, but rather returned as the result of a __pow__() call on a + Function6D object. + + :param object function1: A Function6D object or Python callable. + :param object function2: A Function6D object or Python callable. + """ + def __init__(self, function1, function2): + self._function1 = autowrap_function6d(function1) + self._function2 = autowrap_function6d(function2) + + cdef double evaluate(self, double x, double y, double z, double u, double v, double w) except? -1e999: + cdef double base, exponent + base = self._function1.evaluate(x, y, z, u, v, w) + exponent = self._function2.evaluate(x, y, z, u, v, w) + if base < 0 and floor(exponent) != exponent: # Would return a complex value rather than double + raise ValueError("Negative base and non-integral exponent is not supported") + if base == 0 and exponent < 0: + raise ZeroDivisionError("0.0 cannot be raised to a negative power") + return base ** exponent + + +cdef class AbsFunction6D(Function6D): + """ + A Function6D class that implements the absolute value of the result of a Function6D object: abs(f()). + + This class is not intended to be used directly, but rather returned as the + result of an __abs__() call on a Function6D object. + + :param object function: A Function6D object or Python callable. + """ + def __init__(self, object function): + self._function = autowrap_function6d(function) + + cdef double evaluate(self, double x, double y, double z, double u, double v, double w) except? -1e999: + return abs(self._function.evaluate(x, y, z, u, v, w)) + + +cdef class EqualsFunction6D(Function6D): + """ + A Function6D class that tests the equality of the results of two Function6D objects: f1() == f2() + + This class is not intended to be used directly, but rather returned as the result of an __eq__() call on a + Function6D object. + + :param object function1: A Function6D object or Python callable. + :param object function2: A Function6D object or Python callable. + """ + def __init__(self, object function1, object function2): + self._function1 = autowrap_function6d(function1) + self._function2 = autowrap_function6d(function2) + + cdef double evaluate(self, double x, double y, double z, double u, double v, double w) except? -1e999: + return self._function1.evaluate(x, y, z, u, v, w) == self._function2.evaluate(x, y, z, u, v, w) + + +cdef class NotEqualsFunction6D(Function6D): + """ + A Function6D class that tests the inequality of the results of two Function6D objects: f1() != f2() + + This class is not intended to be used directly, but rather returned as the result of an __ne__() call on a + Function6D object. + + :param object function1: A Function6D object or Python callable. + :param object function2: A Function6D object or Python callable. + """ + def __init__(self, object function1, object function2): + self._function1 = autowrap_function6d(function1) + self._function2 = autowrap_function6d(function2) + + cdef double evaluate(self, double x, double y, double z, double u, double v, double w) except? -1e999: + return self._function1.evaluate(x, y, z, u, v, w) != self._function2.evaluate(x, y, z, u, v, w) + + +cdef class LessThanFunction6D(Function6D): + """ + A Function6D class that implements < of the results of two Function6D objects: f1() < f2() + + This class is not intended to be used directly, but rather returned as the result of an __lt__() call on a + Function6D object. + + :param object function1: A Function6D object or Python callable. + :param object function2: A Function6D object or Python callable. + """ + def __init__(self, object function1, object function2): + self._function1 = autowrap_function6d(function1) + self._function2 = autowrap_function6d(function2) + + cdef double evaluate(self, double x, double y, double z, double u, double v, double w) except? -1e999: + return self._function1.evaluate(x, y, z, u, v, w) < self._function2.evaluate(x, y, z, u, v, w) + + +cdef class GreaterThanFunction6D(Function6D): + """ + A Function6D class that implements > of the results of two Function6D objects: f1() > f2() + + This class is not intended to be used directly, but rather returned as the result of a __gt__() call on a + Function6D object. + + :param object function1: A Function6D object or Python callable. + :param object function2: A Function6D object or Python callable. + """ + def __init__(self, object function1, object function2): + self._function1 = autowrap_function6d(function1) + self._function2 = autowrap_function6d(function2) + + cdef double evaluate(self, double x, double y, double z, double u, double v, double w) except? -1e999: + return self._function1.evaluate(x, y, z, u, v, w) > self._function2.evaluate(x, y, z, u, v, w) + + +cdef class LessEqualsFunction6D(Function6D): + """ + A Function6D class that implements <= of the results of two Function6D objects: f1() <= f2() + + This class is not intended to be used directly, but rather returned as the result of an __le__() call on a + Function6D object. + + :param object function1: A Function6D object or Python callable. + :param object function2: A Function6D object or Python callable. + """ + def __init__(self, object function1, object function2): + self._function1 = autowrap_function6d(function1) + self._function2 = autowrap_function6d(function2) + + cdef double evaluate(self, double x, double y, double z, double u, double v, double w) except? -1e999: + return self._function1.evaluate(x, y, z, u, v, w) <= self._function2.evaluate(x, y, z, u, v, w) + + +cdef class GreaterEqualsFunction6D(Function6D): + """ + A Function6D class that implements >= of the results of two Function6D objects: f1() >= f2() + + This class is not intended to be used directly, but rather returned as the result of an __ge__() call on a + Function6D object. + + :param object function1: A Function6D object or Python callable. + :param object function2: A Function6D object or Python callable. + """ + def __init__(self, object function1, object function2): + self._function1 = autowrap_function6d(function1) + self._function2 = autowrap_function6d(function2) + + cdef double evaluate(self, double x, double y, double z, double u, double v, double w) except? -1e999: + return self._function1.evaluate(x, y, z, u, v, w) >= self._function2.evaluate(x, y, z, u, v, w) + + +cdef class AddScalar6D(Function6D): + """ + A Function6D class that implements the addition of scalar and the result of a Function6D object: K + f() + + This class is not intended to be used directly, but rather returned as the result of an __add__() call on a + Function6D object. + + :param value: A double value. + :param function: A Function6D object or Python callable. + """ + + def __init__(self, double value, object function): + self._value = value + self._function = autowrap_function6d(function) + + cdef double evaluate(self, double x, double y, double z, double u, double v, double w) except? -1e999: + return self._value + self._function.evaluate(x, y, z, u, v, w) + + +cdef class SubtractScalar6D(Function6D): + """ + A Function6D class that implements the subtraction of scalar and the result of a Function6D object: K - f() + + This class is not intended to be used directly, but rather returned as the result of an __sub__() call on a + Function6D object. + + :param value: A double value. + :param function: A Function6D object or Python callable. + """ + + def __init__(self, double value, object function): + self._value = value + self._function = autowrap_function6d(function) + + cdef double evaluate(self, double x, double y, double z, double u, double v, double w) except? -1e999: + return self._value - self._function.evaluate(x, y, z, u, v, w) + + +cdef class MultiplyScalar6D(Function6D): + """ + A Function6D class that implements the multiplication of scalar and the result of a Function6D object: K * f() + + This class is not intended to be used directly, but rather returned as the result of an __mul__() call on a + Function6D object. + + :param value: A double value. + :param function: A Function6D object or Python callable. + """ + + def __init__(self, double value, object function): + self._value = value + self._function = autowrap_function6d(function) + + cdef double evaluate(self, double x, double y, double z, double u, double v, double w) except? -1e999: + return self._value * self._function.evaluate(x, y, z, u, v, w) + + +cdef class DivideScalar6D(Function6D): + """ + A Function6D class that implements the subtraction of scalar and the result of a Function6D object: K / f() + + This class is not intended to be used directly, but rather returned as the result of an __div__() call on a + Function6D object. + + :param value: A double value. + :param function: A Function6D object or Python callable. + """ + + def __init__(self, double value, object function): + self._value = value + self._function = autowrap_function6d(function) + + @cython.cdivision(True) + cdef double evaluate(self, double x, double y, double z, double u, double v, double w) except? -1e999: + cdef double denominator = self._function.evaluate(x, y, z, u, v, w) + if denominator == 0.0: + raise ZeroDivisionError("Function used as the denominator of the division returned a zero value.") + return self._value / denominator + + +cdef class ModuloScalarFunction6D(Function6D): + """ + A Function6D class that implements the modulo of scalar and the result of a Function6D object: K % f() + + This class is not intended to be used directly, but rather returned as the result of a __mod__() call on a + Function6D object. + + :param float value: A double value. + :param object function: A Function6D object or Python callable. + """ + def __init__(self, double value, object function): + self._value = value + self._function = autowrap_function6d(function) + + @cython.cdivision(True) + cdef double evaluate(self, double x, double y, double z, double u, double v, double w) except? -1e999: + cdef double divisor = self._function.evaluate(x, y, z, u, v, w) + if divisor == 0.0: + raise ZeroDivisionError("Function used as the divisor of the modulo returned a zero value.") + return self._value % divisor + + +cdef class ModuloFunctionScalar6D(Function6D): + """ + A Function6D class that implements the modulo of the result of a Function6D object and a scalar: f() % K + + This class is not intended to be used directly, but rather returned as the result of a __mod__() call on a + Function6D object. + + :param object function: A Function6D object or Python callable. + :param float value: A double value. + """ + def __init__(self, object function, double value): + if value == 0: + raise ValueError("Divisor cannot be zero") + self._value = value + self._function = autowrap_function6d(function) + + @cython.cdivision(True) + cdef double evaluate(self, double x, double y, double z, double u, double v, double w) except? -1e999: + return self._function.evaluate(x, y, z, u, v, w) % self._value + + +cdef class PowScalarFunction6D(Function6D): + """ + A Function6D class that implements the pow of scalar and the result of a Function6D object: K ** f() + + This class is not intended to be used directly, but rather returned as the result of an __pow__() call on a + Function6D object. + + :param float value: A double value. + :param object function: A Function6D object or Python callable. + """ + def __init__(self, double value, object function): + self._value = value + self._function = autowrap_function6d(function) + + cdef double evaluate(self, double x, double y, double z, double u, double v, double w) except? -1e999: + cdef double exponent = self._function.evaluate(x, y, z, u, v, w) + if self._value < 0 and floor(exponent) != exponent: + raise ValueError("Negative base and non-integral exponent is not supported") + if self._value == 0 and exponent < 0: + raise ZeroDivisionError("0.0 cannot be raised to a negative power") + return self._value ** exponent + + +cdef class PowFunctionScalar6D(Function6D): + """ + A Function6D class that implements the pow of the result of a Function6D object and a scalar: f() ** K + + This class is not intended to be used directly, but rather returned as the result of an __pow__() call on a + Function6D object. + + :param object function: A Function6D object or Python callable. + :param float value: A double value. + """ + def __init__(self, object function, double value): + self._value = value + self._function = autowrap_function6d(function) + + cdef double evaluate(self, double x, double y, double z, double u, double v, double w) except? -1e999: + cdef double base = self._function.evaluate(x, y, z, u, v, w) + if base < 0 and floor(self._value) != self._value: + raise ValueError("Negative base and non-integral exponent is not supported") + if base == 0 and self._value < 0: + raise ZeroDivisionError("0.0 cannot be raised to a negative power") + return base ** self._value + + +cdef class EqualsScalar6D(Function6D): + """ + A Function6D class that tests the equality of a scalar and the result of a Function6D object: K == f2() + + This class is not intended to be used directly, but rather returned as the result of an __eq__() call on a + Function6D object. + + :param value: A double value. + :param object function: A Function6D object or Python callable. + """ + def __init__(self, double value, object function): + self._value = value + self._function = autowrap_function6d(function) + + cdef double evaluate(self, double x, double y, double z, double u, double v, double w) except? -1e999: + return self._value == self._function.evaluate(x, y, z, u, v, w) + + +cdef class NotEqualsScalar6D(Function6D): + """ + A Function6D class that tests the inequality of a scalar and the result of a Function6D object: K != f2() + + This class is not intended to be used directly, but rather returned as the result of an __ne__() call on a + Function6D object. + + :param value: A double value. + :param object function: A Function6D object or Python callable. + """ + def __init__(self, double value, object function): + self._value = value + self._function = autowrap_function6d(function) + + cdef double evaluate(self, double x, double y, double z, double u, double v, double w) except? -1e999: + return self._value != self._function.evaluate(x, y, z, u, v, w) + + +cdef class LessThanScalar6D(Function6D): + """ + A Function6D class that implements < of a scalar and the result of a Function6D object: K < f2() + + This class is not intended to be used directly, but rather returned as the result of an __lt__() call on a + Function6D object. + + :param value: A double value. + :param object function: A Function6D object or Python callable. + """ + def __init__(self, double value, object function): + self._value = value + self._function = autowrap_function6d(function) + + cdef double evaluate(self, double x, double y, double z, double u, double v, double w) except? -1e999: + return self._value < self._function.evaluate(x, y, z, u, v, w) + + +cdef class GreaterThanScalar6D(Function6D): + """ + A Function6D class that implements > of a scalar and the result of a Function6D object: K > f2() + + This class is not intended to be used directly, but rather returned as the result of a __gt__() call on a + Function6D object. + + :param value: A double value. + :param object function: A Function6D object or Python callable. + """ + def __init__(self, double value, object function): + self._value = value + self._function = autowrap_function6d(function) + + cdef double evaluate(self, double x, double y, double z, double u, double v, double w) except? -1e999: + return self._value > self._function.evaluate(x, y, z, u, v, w) + + +cdef class LessEqualsScalar6D(Function6D): + """ + A Function6D class that implements <= of a scalar and the result of a Function6D object: K <= f2() + + This class is not intended to be used directly, but rather returned as the result of an __le__() call on a + Function6D object. + + :param value: A double value. + :param object function: A Function6D object or Python callable. + """ + def __init__(self, double value, object function): + self._value = value + self._function = autowrap_function6d(function) + + cdef double evaluate(self, double x, double y, double z, double u, double v, double w) except? -1e999: + return self._value <= self._function.evaluate(x, y, z, u, v, w) + + +cdef class GreaterEqualsScalar6D(Function6D): + """ + A Function6D class that implements >= of a scalar and the result of a Function6D object: K >= f2() + + This class is not intended to be used directly, but rather returned as the result of an __ge__() call on a + Function6D object. + + :param value: A double value. + :param object function: A Function6D object or Python callable. + """ + def __init__(self, double value, object function): + self._value = value + self._function = autowrap_function6d(function) + + cdef double evaluate(self, double x, double y, double z, double u, double v, double w) except? -1e999: + return self._value >= self._function.evaluate(x, y, z, u, v, w) diff --git a/cherab/core/math/function/float/function6d/blend.pxd b/cherab/core/math/function/float/function6d/blend.pxd new file mode 100644 index 00000000..3a634883 --- /dev/null +++ b/cherab/core/math/function/float/function6d/blend.pxd @@ -0,0 +1,25 @@ +# cython: language_level=3 + +# Copyright 2016-2025 Euratom +# Copyright 2016-2025 United Kingdom Atomic Energy Authority +# Copyright 2016-2025 Centro de Investigaciones Energéticas, Medioambientales y Tecnológicas +# +# Licensed under the EUPL, Version 1.1 or – as soon they will be approved by the +# European Commission - subsequent versions of the EUPL (the "Licence"); +# You may not use this work except in compliance with the Licence. +# You may obtain a copy of the Licence at: +# +# https://joinup.ec.europa.eu/software/page/eupl5 +# +# Unless required by applicable law or agreed to in writing, software distributed +# under the Licence is distributed on an "AS IS" basis, WITHOUT WARRANTIES OR +# CONDITIONS OF ANY KIND, either express or implied. +# +# See the Licence for the specific language governing permissions and limitations +# under the Licence. + +from cherab.core.math.function.float.function6d.base cimport Function6D + + +cdef class Blend6D(Function6D): + cdef Function6D _f1, _f2, _mask \ No newline at end of file diff --git a/cherab/core/math/function/float/function6d/blend.pyx b/cherab/core/math/function/float/function6d/blend.pyx new file mode 100644 index 00000000..550feddd --- /dev/null +++ b/cherab/core/math/function/float/function6d/blend.pyx @@ -0,0 +1,65 @@ +# cython: language_level=3 + +# Copyright 2016-2025 Euratom +# Copyright 2016-2025 United Kingdom Atomic Energy Authority +# Copyright 2016-2025 Centro de Investigaciones Energéticas, Medioambientales y Tecnológicas +# +# Licensed under the EUPL, Version 1.1 or – as soon they will be approved by the +# European Commission - subsequent versions of the EUPL (the "Licence"); +# You may not use this work except in compliance with the Licence. +# You may obtain a copy of the Licence at: +# +# https://joinup.ec.europa.eu/software/page/eupl5 +# +# Unless required by applicable law or agreed to in writing, software distributed +# under the Licence is distributed on an "AS IS" basis, WITHOUT WARRANTIES OR +# CONDITIONS OF ANY KIND, either express or implied. +# +# See the Licence for the specific language governing permissions and limitations +# under the Licence. + +from cherab.core.math.function.float.function6d.autowrap cimport autowrap_function6d +from raysect.core.math.cython cimport clamp + + +cdef class Blend6D(Function6D): + """ + Performs a linear interpolation between two scalar functions, modulated by a 3rd scalar function. + + The value of the scalar mask function is used to interpolated between the + values returned by the two functions. Mathematically the value returned by + this function is as follows: + + .. math:: + v = (1 - f_m(x, y, z, u, v, w)) f_1(x, y, z, u, v, w) + f_m(x, y, z, u, v, w) f_2(x, y, z, u, v, w) + + The value of the mask function is clamped to the range [0, 1] if the sampled + value exceeds the required range. + """ + + def __init__(self, object f1, object f2, object mask): + """ + :param float.Function6D f1: First scalar function. + :param float.Function6D f2: Second scalar function. + :param float.Function6D mask: Scalar function returning a value in the range [0, 1]. + """ + + self._f1 = autowrap_function6d(f1) + self._f2 = autowrap_function6d(f2) + self._mask = autowrap_function6d(mask) + + cdef double evaluate(self, double x, double y, double z, double u, double v, double w) except? -1e999: + + cdef double t = clamp(self._mask.evaluate(x, y, z, u, v, w), 0.0, 1.0) + + # sample endpoints directly + if t == 0: + return self._f1.evaluate(x, y, z, u, v, w) + + if t == 1: + return self._f2.evaluate(x, y, z, u, v, w) + + # lerp between function values + cdef double f1 = self._f1.evaluate(x, y, z, u, v, w) + cdef double f2 = self._f2.evaluate(x, y, z, u, v, w) + return (1 - t) * f1 + t * f2 diff --git a/cherab/core/math/function/float/function6d/cmath.pxd b/cherab/core/math/function/float/function6d/cmath.pxd new file mode 100644 index 00000000..bf1043a7 --- /dev/null +++ b/cherab/core/math/function/float/function6d/cmath.pxd @@ -0,0 +1,61 @@ +# cython: language_level=3 + +# Copyright 2016-2025 Euratom +# Copyright 2016-2025 United Kingdom Atomic Energy Authority +# Copyright 2016-2025 Centro de Investigaciones Energéticas, Medioambientales y Tecnológicas +# +# Licensed under the EUPL, Version 1.1 or – as soon they will be approved by the +# European Commission - subsequent versions of the EUPL (the "Licence"); +# You may not use this work except in compliance with the Licence. +# You may obtain a copy of the Licence at: +# +# https://joinup.ec.europa.eu/software/page/eupl5 +# +# Unless required by applicable law or agreed to in writing, software distributed +# under the Licence is distributed on an "AS IS" basis, WITHOUT WARRANTIES OR +# CONDITIONS OF ANY KIND, either express or implied. +# +# See the Licence for the specific language governing permissions and limitations +# under the Licence. + +from cherab.core.math.function.float.function6d.base cimport Function6D + + +cdef class Exp6D(Function6D): + cdef Function6D _function + + +cdef class Sin6D(Function6D): + cdef Function6D _function + + +cdef class Cos6D(Function6D): + cdef Function6D _function + + +cdef class Tan6D(Function6D): + cdef Function6D _function + + +cdef class Asin6D(Function6D): + cdef Function6D _function + + +cdef class Acos6D(Function6D): + cdef Function6D _function + + +cdef class Atan6D(Function6D): + cdef Function6D _function + + +cdef class Atan4Q6D(Function6D): + cdef Function6D _numerator, _denominator + + +cdef class Sqrt6D(Function6D): + cdef Function6D _function + + +cdef class Erf6D(Function6D): + cdef Function6D _function diff --git a/cherab/core/math/function/float/function6d/cmath.pyx b/cherab/core/math/function/float/function6d/cmath.pyx new file mode 100644 index 00000000..b33e1680 --- /dev/null +++ b/cherab/core/math/function/float/function6d/cmath.pyx @@ -0,0 +1,168 @@ +# cython: language_level=3 + +# Copyright 2016-2025 Euratom +# Copyright 2016-2025 United Kingdom Atomic Energy Authority +# Copyright 2016-2025 Centro de Investigaciones Energéticas, Medioambientales y Tecnológicas +# +# Licensed under the EUPL, Version 1.1 or – as soon they will be approved by the +# European Commission - subsequent versions of the EUPL (the "Licence"); +# You may not use this work except in compliance with the Licence. +# You may obtain a copy of the Licence at: +# +# https://joinup.ec.europa.eu/software/page/eupl5 +# +# Unless required by applicable law or agreed to in writing, software distributed +# under the Licence is distributed on an "AS IS" basis, WITHOUT WARRANTIES OR +# CONDITIONS OF ANY KIND, either express or implied. +# +# See the Licence for the specific language governing permissions and limitations +# under the Licence. + +cimport libc.math as cmath +from cherab.core.math.function.float.function6d.base cimport Function6D +from cherab.core.math.function.float.function6d.autowrap cimport autowrap_function6d + + +cdef class Exp6D(Function6D): + """ + A Function6D class that implements the exponential of the result of a Function6D object: exp(f()) + + :param Function6D function: A Function6D object. + """ + def __init__(self, object function): + self._function = autowrap_function6d(function) + + cdef double evaluate(self, double x, double y, double z, double u, double v, double w) except? -1e999: + return cmath.exp(self._function.evaluate(x, y, z, u, v, w)) + + +cdef class Sin6D(Function6D): + """ + A Function6D class that implements the sine of the result of a Function6D object: sin(f()) + + :param Function6D function: A Function6D object. + """ + def __init__(self, object function): + self._function = autowrap_function6d(function) + + cdef double evaluate(self, double x, double y, double z, double u, double v, double w) except? -1e999: + return cmath.sin(self._function.evaluate(x, y, z, u, v, w)) + + +cdef class Cos6D(Function6D): + """ + A Function6D class that implements the cosine of the result of a Function6D object: cos(f()) + + :param Function6D function: A Function6D object. + """ + def __init__(self, object function): + self._function = autowrap_function6d(function) + + cdef double evaluate(self, double x, double y, double z, double u, double v, double w) except? -1e999: + return cmath.cos(self._function.evaluate(x, y, z, u, v, w)) + + +cdef class Tan6D(Function6D): + """ + A Function6D class that implements the tangent of the result of a Function6D object: tan(f()) + + :param Function6D function: A Function6D object. + """ + def __init__(self, object function): + self._function = autowrap_function6d(function) + + cdef double evaluate(self, double x, double y, double z, double u, double v, double w) except? -1e999: + return cmath.tan(self._function.evaluate(x, y, z, u, v, w)) + + +cdef class Asin6D(Function6D): + """ + A Function6D class that implements the arcsine of the result of a Function6D object: asin(f()) + + :param Function6D function: A Function6D object. + """ + def __init__(self, object function): + self._function = autowrap_function6d(function) + + cdef double evaluate(self, double x, double y, double z, double u, double v, double w) except? -1e999: + cdef double val = self._function.evaluate(x, y, z, u, v, w) + if -1.0 <= val <= 1.0: + return cmath.asin(val) + raise ValueError("The function returned a value outside of the arcsine domain of [-1, 1].") + + +cdef class Acos6D(Function6D): + """ + A Function6D class that implements the arccosine of the result of a Function6D object: acos(f()) + + :param Function6D function: A Function6D object. + """ + def __init__(self, object function): + self._function = autowrap_function6d(function) + + cdef double evaluate(self, double x, double y, double z, double u, double v, double w) except? -1e999: + cdef double val = self._function.evaluate(x, y, z, u, v, w) + if -1.0 <= val <= 1.0: + return cmath.acos(val) + raise ValueError("The function returned a value outside of the arccosine domain of [-1, 1].") + + +cdef class Atan6D(Function6D): + """ + A Function6D class that implements the arctangent of the result of a Function6D object: atan(f()) + + :param Function6D function: A Function6D object. + """ + def __init__(self, object function): + self._function = autowrap_function6d(function) + + cdef double evaluate(self, double x, double y, double z, double u, double v, double w) except? -1e999: + return cmath.atan(self._function.evaluate(x, y, z, u, v, w)) + + +cdef class Atan4Q6D(Function6D): + """ + A Function6D class that implements the arctangent of the result of 2 Function6D objects: atan2(f1(), f2()) + + This differs from Atan6D in that it takes separate functions for the + numerator and denominator, in order to get the quadrant correct. + + :param Function6D numerator: A Function6D object representing the numerator + :param Function6D denominator: A Function6D object representing the denominator + """ + def __init__(self, object numerator, object denominator): + self._numerator = autowrap_function6d(numerator) + self._denominator = autowrap_function6d(denominator) + + cdef double evaluate(self, double x, double y, double z, double u, double v, double w) except? -1e999: + return cmath.atan2(self._numerator.evaluate(x, y, z, u, v, w), + self._denominator.evaluate(x, y, z, u, v, w)) + + +cdef class Sqrt6D(Function6D): + """ + A Function6D class that implements the square root of the result of a Function6D object: sqrt(f()) + + :param Function6D function: A Function6D object. + """ + def __init__(self, object function): + self._function = autowrap_function6d(function) + + cdef double evaluate(self, double x, double y, double z, double u, double v, double w) except? -1e999: + cdef double f = self._function.evaluate(x, y, z, u, v, w) + if f < 0: # complex values are not supported + raise ValueError("Math domain error in sqrt({0}). Sqrt of a negative value is not supported.".format(f)) + return cmath.sqrt(f) + + +cdef class Erf6D(Function6D): + """ + A Function6D class that implements the error function of the result of a Function6D object: erf(f()) + + :param Function6D function: A Function6D object. + """ + def __init__(self, object function): + self._function = autowrap_function6d(function) + + cdef double evaluate(self, double x, double y, double z, double u, double v, double w) except? -1e999: + return cmath.erf(self._function.evaluate(x, y, z, u, v, w)) \ No newline at end of file diff --git a/cherab/__init__.py b/cherab/core/math/function/float/function6d/constant.pxd similarity index 68% rename from cherab/__init__.py rename to cherab/core/math/function/float/function6d/constant.pxd index 53af8cd3..2ff17e61 100644 --- a/cherab/__init__.py +++ b/cherab/core/math/function/float/function6d/constant.pxd @@ -1,6 +1,8 @@ -# Copyright 2016-2018 Euratom -# Copyright 2016-2018 United Kingdom Atomic Energy Authority -# Copyright 2016-2018 Centro de Investigaciones Energéticas, Medioambientales y Tecnológicas +# cython: language_level=3 + +# Copyright 2016-2025 Euratom +# Copyright 2016-2025 United Kingdom Atomic Energy Authority +# Copyright 2016-2025 Centro de Investigaciones Energéticas, Medioambientales y Tecnológicas # # Licensed under the EUPL, Version 1.1 or – as soon they will be approved by the # European Commission - subsequent versions of the EUPL (the "Licence"); @@ -16,4 +18,8 @@ # See the Licence for the specific language governing permissions and limitations # under the Licence. -__import__('pkg_resources').declare_namespace(__name__) +from cherab.core.math.function.float.function6d.base cimport Function6D + + +cdef class Constant6D(Function6D): + cdef double _value diff --git a/cherab/core/math/function/float/function6d/constant.pyx b/cherab/core/math/function/float/function6d/constant.pyx new file mode 100644 index 00000000..be25b44b --- /dev/null +++ b/cherab/core/math/function/float/function6d/constant.pyx @@ -0,0 +1,49 @@ +# cython: language_level=3 + +# Copyright 2016-2025 Euratom +# Copyright 2016-2025 United Kingdom Atomic Energy Authority +# Copyright 2016-2025 Centro de Investigaciones Energéticas, Medioambientales y Tecnológicas +# +# Licensed under the EUPL, Version 1.1 or – as soon they will be approved by the +# European Commission - subsequent versions of the EUPL (the "Licence"); +# You may not use this work except in compliance with the Licence. +# You may obtain a copy of the Licence at: +# +# https://joinup.ec.europa.eu/software/page/eupl5 +# +# Unless required by applicable law or agreed to in writing, software distributed +# under the Licence is distributed on an "AS IS" basis, WITHOUT WARRANTIES OR +# CONDITIONS OF ANY KIND, either express or implied. +# +# See the Licence for the specific language governing permissions and limitations +# under the Licence. + +from cherab.core.math.function.float.function6d.base cimport Function6D + + +cdef class Constant6D(Function6D): + """ + Wraps a scalar constant with a Function6D object. + + This class allows a numeric Python scalar, such as a float or an integer, to + interact with cython code that requires a Function6D object. The scalar must + be convertible to double. The value of the scalar constant will be returned + independent of the arguments the function is called with. + + This class is intended to be used to transparently wrap python objects that + are passed via constructors or methods into cython optimised code. It is not + intended that the users should need to directly interact with these wrapping + objects. Constructors and methods expecting a Function6D object should be + designed to accept a generic python object and then test that object to + determine if it is an instance of Function6D. If the object is not a + Function6D object it should be wrapped using this class for internal use. + + See also: autowrap_function6d() + + :param float value: the constant value to return when called + """ + def __init__(self, double value): + self._value = value + + cdef double evaluate(self, double x, double y, double z, double u, double v, double w) except? -1e999: + return self._value diff --git a/cherab/core/math/function/float/function6d/tests/__init__.py b/cherab/core/math/function/float/function6d/tests/__init__.py new file mode 100644 index 00000000..dcc4669c --- /dev/null +++ b/cherab/core/math/function/float/function6d/tests/__init__.py @@ -0,0 +1,5 @@ +from .test_base import * +from .test_autowrap import * +from .test_constant import * +from .test_arg import * +from .test_cmath import * diff --git a/cherab/core/math/function/float/function6d/tests/test_arg.py b/cherab/core/math/function/float/function6d/tests/test_arg.py new file mode 100644 index 00000000..fb559064 --- /dev/null +++ b/cherab/core/math/function/float/function6d/tests/test_arg.py @@ -0,0 +1,50 @@ +# cython: language_level=3 + +# Copyright 2016-2025 Euratom +# Copyright 2016-2025 United Kingdom Atomic Energy Authority +# Copyright 2016-2025 Centro de Investigaciones Energéticas, Medioambientales y Tecnológicas +# +# Licensed under the EUPL, Version 1.1 or – as soon they will be approved by the +# European Commission - subsequent versions of the EUPL (the "Licence"); +# You may not use this work except in compliance with the Licence. +# You may obtain a copy of the Licence at: +# +# https://joinup.ec.europa.eu/software/page/eupl5 +# +# Unless required by applicable law or agreed to in writing, software distributed +# under the Licence is distributed on an "AS IS" basis, WITHOUT WARRANTIES OR +# CONDITIONS OF ANY KIND, either express or implied. +# +# See the Licence for the specific language governing permissions and limitations +# under the Licence. + +""" +Unit tests for the Arg6D class. +""" + +import unittest +import itertools +from cherab.core.math.function.float.function6d.arg import Arg6D + +# TODO: expand tests to cover the cython interface +class TestArg6D(unittest.TestCase): + + def test_arg(self): + testvals = [-1e10, -7, -0.001, 0.0, 0.00003, 10, 2.3e49] + for (x, y, z, u, v, w) in itertools.product(testvals, repeat=6): + argx = Arg6D("x") + argy = Arg6D("y") + argz = Arg6D("z") + argu = Arg6D("u") + argw = Arg6D("v") + argv = Arg6D("w") + self.assertEqual(argx(x, y, z, u, v, w), x, "Arg6D('x') call did not match reference value.") + self.assertEqual(argy(x, y, z, u, v, w), y, "Arg6D('y') call did not match reference value.") + self.assertEqual(argz(x, y, z, u, v, w), z, "Arg6D('z') call did not match reference value.") + self.assertEqual(argu(x, y, z, u, v, w), u, "Arg6D('u') call did not match reference value.") + self.assertEqual(argw(x, y, z, u, v, w), v, "Arg6D('v') call did not match reference value.") + self.assertEqual(argv(x, y, z, u, v, w), w, "Arg6D('w') call did not match reference value.") + + def test_invalid_inputs(self): + with self.assertRaises(ValueError, msg="Arg6D did not raise ValueError with incorrect string."): + Arg6D("q") diff --git a/cherab/core/math/function/float/function6d/tests/test_autowrap.py b/cherab/core/math/function/float/function6d/tests/test_autowrap.py new file mode 100644 index 00000000..5682f14e --- /dev/null +++ b/cherab/core/math/function/float/function6d/tests/test_autowrap.py @@ -0,0 +1,37 @@ +# cython: language_level=3 + +# Copyright 2016-2025 Euratom +# Copyright 2016-2025 United Kingdom Atomic Energy Authority +# Copyright 2016-2025 Centro de Investigaciones Energéticas, Medioambientales y Tecnológicas +# +# Licensed under the EUPL, Version 1.1 or – as soon they will be approved by the +# European Commission - subsequent versions of the EUPL (the "Licence"); +# You may not use this work except in compliance with the Licence. +# You may obtain a copy of the Licence at: +# +# https://joinup.ec.europa.eu/software/page/eupl5 +# +# Unless required by applicable law or agreed to in writing, software distributed +# under the Licence is distributed on an "AS IS" basis, WITHOUT WARRANTIES OR +# CONDITIONS OF ANY KIND, either express or implied. +# +# See the Licence for the specific language governing permissions and limitations +# under the Licence. + +""" +Unit tests for the autowrap_6d function +""" + +import unittest +from cherab.core.math.function.float.function6d.autowrap import _autowrap_function6d, PythonFunction6D +from cherab.core.math.function.float.function6d.constant import Constant6D + +class TestAutowrap6D(unittest.TestCase): + + def test_constant(self): + function = _autowrap_function6d(5.0) + self.assertIsInstance(function, Constant6D, "Autowrapped scalar float is not a Constant6D.") + + def test_python_function(self): + function = _autowrap_function6d(lambda x, y, z, u, v, w: 10*x + 5*y + 2*z + u + 3*v + 4*w) + self.assertIsInstance(function, PythonFunction6D, "Autowrapped function is not a PythonFunction6D.") diff --git a/cherab/core/math/function/float/function6d/tests/test_base.py b/cherab/core/math/function/float/function6d/tests/test_base.py new file mode 100644 index 00000000..0a572220 --- /dev/null +++ b/cherab/core/math/function/float/function6d/tests/test_base.py @@ -0,0 +1,550 @@ +# cython: language_level=3 + +# Copyright 2016-2025 Euratom +# Copyright 2016-2025 United Kingdom Atomic Energy Authority +# Copyright 2016-2025 Centro de Investigaciones Energéticas, Medioambientales y Tecnológicas +# +# Licensed under the EUPL, Version 1.1 or – as soon they will be approved by the +# European Commission - subsequent versions of the EUPL (the "Licence"); +# You may not use this work except in compliance with the Licence. +# You may obtain a copy of the Licence at: +# +# https://joinup.ec.europa.eu/software/page/eupl5 +# +# Unless required by applicable law or agreed to in writing, software distributed +# under the Licence is distributed on an "AS IS" basis, WITHOUT WARRANTIES OR +# CONDITIONS OF ANY KIND, either express or implied. +# +# See the Licence for the specific language governing permissions and limitations +# under the Licence. + +""" +Unit tests for the Function6D class. +""" + +import math +import unittest +import itertools +from cherab.core.math.function.float.function6d.autowrap import PythonFunction6D + +# TODO: expand tests to cover the cython interface +class TestFunction6D(unittest.TestCase): + + def setUp(self): + self.ref1 = lambda x, y, z, u, v, w: 10 * x + 5 * y + 2 * z + u + 3 * v + 4 * w + self.ref2 = lambda x, y, z, u, v, w: abs(x + y + z + u + v + w) + + self.f1 = PythonFunction6D(self.ref1) + self.f2 = PythonFunction6D(self.ref2) + + def test_call(self): + testvals = [-1e10, -7, -0.001, 0.0, 0.00003, 10, 2.3e49] + for (x, y, z, u, v, w) in itertools.product(testvals, repeat=6): + self.assertEqual(self.f1(x, y, z, u, v, w), self.ref1(x, y, z, u, v, w), + "Function6D call did not match reference function value.") + + def test_negate(self): + testvals = [-1e10, -7, -0.001, 0.0, 0.00003, 10, 2.3e49] + r = -self.f1 + for (x, y, z, u, v, w) in itertools.product(testvals, repeat=6): + self.assertEqual(r(x, y, z, u, v, w), -self.ref1(x, y, z, u, v, w), + "Function6D negate did not match reference function value.") + + def test_add_scalar(self): + testvals = [-1e10, -7, -0.001, 0.0, 0.00003, 10, 2.3e49] + r1 = 8 + self.f1 + r2 = self.f1 + 65 + for (x, y, z, u, v, w) in itertools.product(testvals, repeat=6): + self.assertEqual(r1(x, y, z, u, v, w), 8 + self.ref1(x, y, z, u, v, w), + "Function6D add scalar (K + f()) did not match reference function value.") + self.assertEqual(r2(x, y, z, u, v, w), self.ref1(x, y, z, u, v, w) + 65, + "Function6D add scalar (f() + K) did not match reference function value.") + + def test_sub_scalar(self): + testvals = [-1e10, -7, -0.001, 0.0, 0.00003, 10, 2.3e49] + r1 = 8 - self.f1 + r2 = self.f1 - 65 + for (x, y, z, u, v, w) in itertools.product(testvals, repeat=6): + self.assertEqual(r1(x, y, z, u, v, w), 8 - self.ref1(x, y, z, u, v, w), + "Function6D subtract scalar (K - f()) did not match reference function value.") + self.assertEqual(r2(x, y, z, u, v, w), self.ref1(x, y, z, u, v, w) - 65, + "Function6D subtract scalar (f() - K) did not match reference function value.") + + def test_mul_scalar(self): + testvals = [-1e10, -7, -0.001, 0.0, 0.00003, 10, 2.3e49] + r1 = 5 * self.f1 + r2 = self.f1 * -7.8 + for (x, y, z, u, v, w) in itertools.product(testvals, repeat=6): + self.assertEqual(r1(x, y, z, u, v, w), 5 * self.ref1(x, y, z, u, v, w), + "Function6D multiply scalar (K * f()) did not match reference function value.") + self.assertEqual(r2(x, y, z, u, v, w), self.ref1(x, y, z, u, v, w) * -7.8, + "Function6D multiply scalar (f() * K) did not match reference function value.") + + def test_div_scalar(self): + testvals = [-1e10, -7, -0.001, 0.000031, 10.3, 2.3e49] + r1 = 5.451 / self.f1 + r2 = self.f1 / -7.8 + for (x, y, z, u, v, w) in itertools.product(testvals, repeat=6): + self.assertEqual(r1(x, y, z, u, v, w), 5.451 / self.ref1(x, y, z, u, v, w), + "Function6D divide scalar (K / f()) did not match reference function value.") + self.assertAlmostEqual(r2(x, y, z, u, v, w), self.ref1(x, y, z, u, v, w) / -7.8, + delta=abs(r2(x, y, z, u, v, w)) * 1e-12, + msg="Function6D divide scalar (f() / K) did not match reference function value.") + + r = 5 / self.f1 + with self.assertRaises(ZeroDivisionError, msg="ZeroDivisionError not raised when function returns zero."): + r(0, 0, 0, 0, 0, 0) + + def test_mod_function6d_scalar(self): + # Note that Function6D objects work with doubles, so the floating modulo + # operator is used rather than the integer one. For accurate testing we + # therefore need to use the math.fmod operator rather than % in Python. + testvals = [-10, -7, -0.001, 0.00003, 10, 12.3] + r1 = 5 % self.f1 + r2 = self.f1 % -7.8 + for (x, y, z, u, v, w) in itertools.product(testvals, repeat=6): + if self.ref1(x, y, z, u, v, w) == 0: + with self.assertRaises(ZeroDivisionError, msg="ZeroDivisionError not raised when function returns 0."): + r1(x, y, z, u, v, w) + else: + self.assertAlmostEqual(r1(x, y, z, u, v, w), math.fmod(5, self.ref1(x, y, z, u, v, w)), 15, "Function6D modulo scalar (K % f()) did not match reference function value.") + self.assertAlmostEqual(r2(x, y, z, u, v, w), math.fmod(self.ref1(x, y, z, u, v, w), -7.8), 15, "Function6D modulo scalar (f() % K) did not match reference function value.") + with self.assertRaises(ZeroDivisionError, msg="ZeroDivisionError not raised when function returns 0."): + r1(0, 0, 0, 0, 0, 0) + with self.assertRaises(ZeroDivisionError, msg="ZeroDivisionError not raised when modulo scalar is 0."): + self.f1 % 0 + + def test_pow_function6d_scalar(self): + testvals = [-10, -7, -0.001, 0.00003, 10, 12.3] + r1 = 5. ** self.f1 + r2 = self.f1 ** -7.8 + r3 = (-5.) ** self.f1 + for (x, y, z, u, v, w) in itertools.product(testvals, repeat=6): + self.assertAlmostEqual(r1(x, y, z, u, v, w), 5. ** self.ref1(x, y, z, u, v, w), 15, "Function6D power scalar (K ** f()) did not match reference function value.") + if self.ref1(x, y, z, u, v, w) < 0: + with self.assertRaises(ValueError, msg="ValueError not raised when base is negative and exponent non-integral."): + r2(x, y, z, u, v, w) + elif not float(self.ref1(x, y, z, u, v, w)).is_integer(): + with self.assertRaises(ValueError, msg="ValueError not raised when base is negative and exponent non-integral."): + r3(x, y, z, u, v, w) + else: + if self.ref1(x, y, z, u, v, w) == 0: + with self.assertRaises(ZeroDivisionError, msg="ZeroDivisionError not raised when base is 0 and exponent negative."): + r2(x, y, z, u, v, w) + else: + self.assertAlmostEqual(r2(x, y, z, u, v, w), self.ref1(x, y, z, u, v, w) ** -7.8, 15, "Function6D power scalar (f() ** K) did not match reference function value.") + with self.assertRaises(ZeroDivisionError, msg="ZeroDivisionError not raised when base is 0 and exponent negative."): + r2(0, 0, 0, 0, 0, 0) + with self.assertRaises(ZeroDivisionError, msg="ZeroDivisionError not raised when base is zero and exponent negative."): + r4 = 0 ** self.f1 + r4(-1, 0, 0, 0, 0, 0) + + def test_richcmp_scalar(self): + testvals = [-1e10, -7, -0.001, 0.0, 0.00003, 10, 2.3e49] + for (x, y, z, u, v, w) in itertools.product(testvals, repeat=6): + ref_value = self.ref1(x, y, z, u, v, w) + higher_value = ref_value + abs(ref_value) + 1 + lower_value = ref_value - abs(ref_value) - 1 + self.assertEqual( + (self.f1 == ref_value)(x, y, z, u, v, w), 1.0, + msg="Function6D equals scalar (f() == K) did not return true when it should." + ) + self.assertEqual( + (ref_value == self.f1)(x, y, z, u, v, w), 1.0, + msg="Scalar equals Function6D (K == f()) did not return true when it should." + ) + self.assertEqual( + (self.f1 == higher_value)(x, y, z, u, v, w), 0.0, + msg="Function6D equals scalar (f() == K) did not return false when it should." + ) + self.assertEqual( + (higher_value == self.f1)(x, y, z, u, v, w), 0.0, + msg="Scalar equals Function6D (K == f()) did not return false when it should." + ) + self.assertEqual( + (self.f1 != higher_value)(x, y, z, u, v, w), 1.0, + msg="Function6D not equals scalar (f() != K) did not return true when it should." + ) + self.assertEqual( + (higher_value != self.f1)(x, y, z, u, v, w), 1.0, + msg="Scalar not equals Function6D (K != f()) did not return true when it should." + ) + self.assertEqual( + (self.f1 != ref_value)(x, y, z, u, v, w), 0.0, + msg="Function6D not equals scalar (f() != K) did not return false when it should." + ) + self.assertEqual( + (ref_value != self.f1)(x, y, z, u, v, w), 0.0, + msg="Scalar not equals Function6D (K != f()) did not return false when it should." + ) + self.assertEqual( + (self.f1 < higher_value)(x, y, z, u, v, w), 1.0, + msg="Function6D less than scalar (f() < K) did not return true when it should." + ) + self.assertEqual( + (lower_value < self.f1)(x, y, z, u, v, w), 1.0, + msg="Scalar less than Function6D (K < f()) did not return true when it should." + ) + self.assertEqual( + (self.f1 < lower_value)(x, y, z, u, v, w), 0.0, + msg="Function6D less than scalar (f() < K) did not return false when it should." + ) + self.assertEqual( + (higher_value < self.f1)(x, y, z, u, v, w), 0.0, + msg="Scalar less than Function6D (K < f()) did not return false when it should." + ) + self.assertEqual( + (self.f1 > lower_value)(x, y, z, u, v, w), 1.0, + msg="Function6D greater than scalar (f() > K) did not return true when it should." + ) + self.assertEqual( + (higher_value > self.f1)(x, y, z, u, v, w), 1.0, + msg="Scalar greater than Function6D (K > f()) did not return true when it should." + ) + self.assertEqual( + (self.f1 > higher_value)(x, y, z, u, v, w), 0.0, + msg="Function6D greater than scalar (f() > K) did not return false when it should." + ) + self.assertEqual( + (lower_value > self.f1)(x, y, z, u, v, w), 0.0, + msg="Scalar greater than Function6D (K > f()) did not return false when it should." + ) + self.assertEqual( + (self.f1 <= higher_value)(x, y, z, u, v, w), 1.0, + msg="Function6D less equals scalar (f() <= K) did not return true when it should." + ) + self.assertEqual( + (lower_value <= self.f1)(x, y, z, u, v, w), 1.0, + msg="Scalar less equals Function6D (K <= f()) did not return true when it should." + ) + self.assertEqual( + (self.f1 <= ref_value)(x, y, z, u, v, w), 1.0, + msg="Function6D less equals scalar (f() <= K) did not return true when it should." + ) + self.assertEqual( + (ref_value <= self.f1)(x, y, z, u, v, w), 1.0, + msg="Scalar less equals Function6D (K <= f()) did not return true when it should." + ) + self.assertEqual( + (self.f1 <= lower_value)(x, y, z, u, v, w), 0.0, + msg="Function6D less equals scalar (f() <= K) did not return false when it should." + ) + self.assertEqual( + (higher_value <= self.f1)(x, y, z, u, v, w), 0.0, + msg="Scalar less equals Function6D (K <= f()) did not return false when it should." + ) + self.assertEqual( + (self.f1 >= lower_value)(x, y, z, u, v, w), 1.0, + msg="Function6D greater equals scalar (f() >= K) did not return true when it should." + ) + self.assertEqual( + (higher_value >= self.f1)(x, y, z, u, v, w), 1.0, + msg="Scalar greater equals Function6D (K >= f()) did not return true when it should." + ) + self.assertEqual( + (self.f1 >= ref_value)(x, y, z, u, v, w), 1.0, + msg="Function6D greater equals scalar (f() >= K) did not return true when it should." + ) + self.assertEqual( + (ref_value >= self.f1)(x, y, z, u, v, w), 1.0, + msg="Scalar greater equals Function6D (K >= f()) did not return true when it should." + ) + self.assertEqual( + (self.f1 >= higher_value)(x, y, z, u, v, w), 0.0, + msg="Function6D greater equals scalar (f() >= K) did not return false when it should." + ) + self.assertEqual( + (lower_value >= self.f1)(x, y, z, u, v, w), 0.0, + msg="Scalar greater equals Function6D (K >= f()) did not return false when it should." + ) + + def test_add_function6d(self): + testvals = [-1e10, -7, -0.001, 0.0, 0.00003, 10, 2.3e49] + r1 = self.f1 + self.f2 + r2 = self.ref1 + self.f2 + r3 = self.f1 + self.ref2 + for (x, y, z, u, v, w) in itertools.product(testvals, repeat=6): + self.assertEqual(r1(x, y, z, u, v, w), self.ref1(x, y, z, u, v, w) + self.ref2(x, y, z, u, v, w), "Function6D add function (f1() + f2()) did not match reference function value.") + self.assertEqual(r2(x, y, z, u, v, w), self.ref1(x, y, z, u, v, w) + self.ref2(x, y, z, u, v, w), "Function6D add function (p1() + f2()) did not match reference function value.") + self.assertEqual(r3(x, y, z, u, v, w), self.ref1(x, y, z, u, v, w) + self.ref2(x, y, z, u, v, w), "Function6D add function (f1() + p2()) did not match reference function value.") + + def test_sub_function6d(self): + testvals = [-1e10, -7, -0.001, 0.0, 0.00003, 10, 2.3e49] + r1 = self.f1 - self.f2 + r2 = self.ref1 - self.f2 + r3 = self.f1 - self.ref2 + for (x, y, z, u, v, w) in itertools.product(testvals, repeat=6): + self.assertEqual(r1(x, y, z, u, v, w), self.ref1(x, y, z, u, v, w) - self.ref2(x, y, z, u, v, w), "Function6D subtract function (f1() - f2()) did not match reference function value.") + self.assertEqual(r2(x, y, z, u, v, w), self.ref1(x, y, z, u, v, w) - self.ref2(x, y, z, u, v, w), "Function6D subtract function (p1() - f2()) did not match reference function value.") + self.assertEqual(r3(x, y, z, u, v, w), self.ref1(x, y, z, u, v, w) - self.ref2(x, y, z, u, v, w), "Function6D subtract function (f1() - p2()) did not match reference function value.") + + def test_mul_function6d(self): + testvals = [-1e10, -7, -0.001, 0.0, 0.00003, 10, 2.3e49] + r1 = self.f1 * self.f2 + r2 = self.ref1 * self.f2 + r3 = self.f1 * self.ref2 + for (x, y, z, u, v, w) in itertools.product(testvals, repeat=6): + self.assertEqual(r1(x, y, z, u, v, w), self.ref1(x, y, z, u, v, w) * self.ref2(x, y, z, u, v, w), "Function6D multiply function (f1() * f2()) did not match reference function value.") + self.assertEqual(r2(x, y, z, u, v, w), self.ref1(x, y, z, u, v, w) * self.ref2(x, y, z, u, v, w), "Function6D multiply function (p1() * f2()) did not match reference function value.") + self.assertEqual(r3(x, y, z, u, v, w), self.ref1(x, y, z, u, v, w) * self.ref2(x, y, z, u, v, w), "Function6D multiply function (f1() * p2()) did not match reference function value.") + + def test_div_function6d(self): + testvals = [-1e10, -7, -0.001, 0.00003, 10, 2.3e49] + r1 = self.f1 / self.f2 + r2 = self.ref1 / self.f2 + r3 = self.f1 / self.ref2 + for (x, y, z, u, v, w) in itertools.product(testvals, repeat=6): + self.assertAlmostEqual(r1(x, y, z, u, v, w), self.ref1(x, y, z, u, v, w) / self.ref2(x, y, z, u, v, w), delta=abs(r1(x, y, z, u, v, w)) * 1e-12, msg="Function6D divide function (f1() / f2()) did not match reference function value.") + self.assertAlmostEqual(r2(x, y, z, u, v, w), self.ref1(x, y, z, u, v, w) / self.ref2(x, y, z, u, v, w), delta=abs(r2(x, y, z, u, v, w)) * 1e-12, msg="Function6D divide function (p1() / f2()) did not match reference function value.") + self.assertAlmostEqual(r3(x, y, z, u, v, w), self.ref1(x, y, z, u, v, w) / self.ref2(x, y, z, u, v, w), delta=abs(r3(x, y, z, u, v, w)) * 1e-12, msg="Function6D divide function (f1() / p2()) did not match reference function value.") + + with self.assertRaises(ZeroDivisionError, msg="ZeroDivisionError not raised when function returns zero."): + r1(0, 0, 0, 0, 0, 0) + + def test_mod_function6d(self): + testvals = [-1e10, -7, -0.001, 0.00003, 10, 2.3e49] + r1 = self.f1 % self.f2 + r2 = self.ref1 % self.f2 + r3 = self.f1 % self.ref2 + for (x, y, z, u, v, w) in itertools.product(testvals, repeat=6): + self.assertAlmostEqual(r1(x, y, z, u, v, w), math.fmod(self.ref1(x, y, z, u, v, w), self.ref2(x, y, z, u, v, w)), delta=abs(r1(x, y, z, u, v, w)) * 1e-12, msg="Function6D modulo function (f1() % f2()) did not match reference function value.") + self.assertAlmostEqual(r2(x, y, z, u, v, w), math.fmod(self.ref1(x, y, z, u, v, w), self.ref2(x, y, z, u, v, w)), delta=abs(r2(x, y, z, u, v, w)) * 1e-12, msg="Function6D modulo function (p1() % f2()) did not match reference function value.") + self.assertAlmostEqual(r3(x, y, z, u, v, w), math.fmod(self.ref1(x, y, z, u, v, w), self.ref2(x, y, z, u, v, w)), delta=abs(r3(x, y, z, u, v, w)) * 1e-12, msg="Function6D modulo function (f1() % p2()) did not match reference function value.") + + with self.assertRaises(ZeroDivisionError, msg="ZeroDivisionError not raised when function returns zero."): + r1(0, 0, 0, 0, 0, 0) + + def test_pow_function6d_function6d(self): + testvals = [-3.0, -0.7, -0.001, 0.00003, 2] + r1 = self.f1 ** self.f2 + r2 = self.ref1 ** self.f2 + r3 = self.f1 ** self.ref2 + for (x, y, z, u, v, w) in itertools.product(testvals, repeat=6): + if self.ref1(x, y, z, u, v, w) < 0 and not float(self.ref2(x, y, z, u, v, w)).is_integer(): + with self.assertRaises(ValueError, msg="ValueError not raised when base is negative and exponent non-integral (1/3)."): + r1(x, y, z, u, v, w) + with self.assertRaises(ValueError, msg="ValueError not raised when base is negative and exponent non-integral (2/3)."): + r2(x, y, z, u, v, w) + with self.assertRaises(ValueError, msg="ValueError not raised when base is negative and exponent non-integral (3/3)."): + r3(x, y, z, u, v, w) + else: + self.assertAlmostEqual(r1(x, y, z, u, v, w), self.ref1(x, y, z, u, v, w) ** self.ref2(x, y, z, u, v, w), 15, "Function6D power function (f1() ** f2()) did not match reference function value.") + self.assertAlmostEqual(r2(x, y, z, u, v, w), self.ref1(x, y, z, u, v, w) ** self.ref2(x, y, z, u, v, w), 15, "Function6D power function (p1() ** f2()) did not match reference function value.") + self.assertAlmostEqual(r3(x, y, z, u, v, w), self.ref1(x, y, z, u, v, w) ** self.ref2(x, y, z, u, v, w), 15, "Function6D power function (f1() ** p2()) did not match reference function value.") + + with self.assertRaises(ZeroDivisionError, msg="ZeroDivisionError not raised when f1() == 0 and f2() is negative."): + r4 = PythonFunction6D(lambda x, y, z, u, v, w: 0) ** self.f1 + r4(-1, 0, 0, 0, 0, 0) + + def test_pow_3_arguments(self): + testvals = [-10, -7, -0.001, 0.00003, 0.8] + r1 = pow(self.f1, 5, 3) + r2 = pow(5, self.f1, 3) + r3 = pow(5, self.f1, self.f2) + r4 = pow(self.f2, self.f1, self.f2) + r5 = pow(self.f2, self.ref1, self.ref2) + r6 = pow(self.ref2, self.f1, self.f2) + # Can't use 3 argument pow() if all arguments aren't integers, so + # use fmod(a, b) % c instead + for (x, y, z, u, v, w) in itertools.product(testvals, repeat=6): + self.assertEqual(r1(x, y, z, u, v, w), math.fmod(self.ref1(x, y, z, u, v, w) ** 5, 3), "Function6D 3 argument pow(f1(), A, B) did not match reference value.") + self.assertEqual(r2(x, y, z, u, v, w), math.fmod(5 ** self.ref1(x, y, z, u, v, w), 3), "Function6D 3 argument pow(A, f1(), B) did not match reference value.") + self.assertEqual(r3(x, y, z, u, v, w), math.fmod(5 ** self.ref1(x, y, z, u, v, w), self.ref2(x, y, z, u, v, w)), "Function6D 3 argument pow(A, f1(), f2()) did not match reference value.") + self.assertEqual(r4(x, y, z, u, v, w), math.fmod(self.ref2(x, y, z, u, v, w) ** self.ref1(x, y, z, u, v, w), self.ref2(x, y, z, u, v, w)), "Function6D 3 argument pow(f2(), f1(), f2()) did not match reference value.") + self.assertEqual(r5(x, y, z, u, v, w), math.fmod(self.ref2(x, y, z, u, v, w) ** self.ref1(x, y, z, u, v, w), self.ref2(x, y, z, u, v, w)), "Function6D 3 argument pow(f2(), p1(), p2()) did not match reference value.") + self.assertEqual(r6(x, y, z, u, v, w), math.fmod(self.ref2(x, y, z, u, v, w) ** self.ref1(x, y, z, u, v, w), self.ref2(x, y, z, u, v, w)), "Function6D 3 argument pow(p2(), f1(), f2()) did not match reference value.") + + def test_abs(self): + testvals = [-1e10, -7, -0.001, 0.0, 0.0003, 10, 2.3e49] + for (x, y, z, u, v, w) in itertools.product(testvals, repeat=6): + self.assertEqual(abs(self.f1)(x, y, z, u, v, w), abs(self.ref1(x, y, z, u, v, w)), + msg="abs(Function6D) did not match reference value") + + def test_richcmp_function_callable(self): + testvals = [-1e10, -7, -0.001, 0.0, 0.00003, 10, 2.3e49] + for (x, y, z, u, v, w) in itertools.product(testvals, repeat=6): + ref_value = self.ref1 + higher_value = lambda x, y, z, u, v, w: self.ref1(x, y, z, u, v, w) + abs(self.ref1(x, y, z, u, v, w)) + 1 + lower_value = lambda x, y, z, u, v, w: self.ref1(x, y, z, u, v, w) - abs(self.ref1(x, y, z, u, v, w)) - 1 + self.assertEqual( + (self.f1 == ref_value)(x, y, z, u, v, w), 1.0, + msg="Function6D equals callable (f1() == f2()) did not return true when it should." + ) + self.assertEqual( + (self.f1 == higher_value)(x, y, z, u, v, w), 0.0, + msg="Function6D equals callable (f1() == f2()) did not return false when it should." + ) + self.assertEqual( + (self.f1 != higher_value)(x, y, z, u, v, w), 1.0, + msg="Function6D not equals callable (f1() != f2()) did not return true when it should." + ) + self.assertEqual( + (self.f1 != ref_value)(x, y, z, u, v, w), 0.0, + msg="Function6D not equals callable (f1() != f2()) did not return false when it should." + ) + self.assertEqual( + (self.f1 < higher_value)(x, y, z, u, v, w), 1.0, + msg="Function6D less than callable (f1() < f2()) did not return true when it should." + ) + self.assertEqual( + (self.f1 < lower_value)(x, y, z, u, v, w), 0.0, + msg="Function6D less than callable (f1() < f2()) did not return false when it should." + ) + self.assertEqual( + (self.f1 > lower_value)(x, y, z, u, v, w), 1.0, + msg="Function6D greater than callable (f1() > f2()) did not return true when it should." + ) + self.assertEqual( + (self.f1 > higher_value)(x, y, z, u, v, w), 0.0, + msg="Function6D greater than callable (f1() > f2()) did not return false when it should." + ) + self.assertEqual( + (self.f1 <= higher_value)(x, y, z, u, v, w), 1.0, + msg="Function6D less equals callable (f1() <= f2()) did not return true when it should." + ) + self.assertEqual( + (self.f1 <= ref_value)(x, y, z, u, v, w), 1.0, + msg="Function6D less equals callable (f1() <= f2()) did not return true when it should." + ) + self.assertEqual( + (self.f1 <= lower_value)(x, y, z, u, v, w), 0.0, + msg="Function6D less equals callable (f1() <= f2()) did not return false when it should." + ) + self.assertEqual( + (self.f1 >= lower_value)(x, y, z, u, v, w), 1.0, + msg="Function6D equals callable (f1() >= f2()) did not return true when it should." + ) + self.assertEqual( + (self.f1 >= ref_value)(x, y, z, u, v, w), 1.0, + msg="Function6D greater equals callable (f1() >= f2()) did not return true when it should." + ) + self.assertEqual( + (self.f1 >= higher_value)(x, y, z, u, v, w), 0.0, + msg="Function6D equals callable (f1() >= f2()) did not return false when it should." + ) + + def test_richcmp_callable_function(self): + testvals = [-1e10, -7, -0.001, 0.0, 0.00003, 10, 2.3e49] + for (x, y, z, u, v, w) in itertools.product(testvals, repeat=6): + ref_value = self.ref1 + higher_value = lambda x, y, z, u, v, w: self.ref1(x, y, z, u, v, w) + abs(self.ref1(x, y, z, u, v, w)) + 1 + lower_value = lambda x, y, z, u, v, w: self.ref1(x, y, z, u, v, w) - abs(self.ref1(x, y, z, u, v, w)) - 1 + self.assertEqual( + (ref_value == self.f1)(x, y, z, u, v, w), 1.0, + msg="Callable equals Function6D (f1() == f2()) did not return true when it should." + ) + self.assertEqual( + (higher_value == self.f1)(x, y, z, u, v, w), 0.0, + msg="Callable equals Function6D (f1() == f2()) did not return false when it should." + ) + self.assertEqual( + (higher_value != self.f1)(x, y, z, u, v, w), 1.0, + msg="Callable not equals Function6D (f1() != f2()) did not return true when it should." + ) + self.assertEqual( + (ref_value != self.f1)(x, y, z, u, v, w), 0.0, + msg="Callable not equals Function6D (f1() != f2()) did not return false when it should." + ) + self.assertEqual( + (lower_value < self.f1)(x, y, z, u, v, w), 1.0, + msg="Callable less than Function6D (f1() < f2()) did not return true when it should." + ) + self.assertEqual( + (higher_value < self.f1)(x, y, z, u, v, w), 0.0, + msg="Callable less than Function6D (f1() < f2()) did not return false when it should." + ) + self.assertEqual( + (higher_value > self.f1)(x, y, z, u, v, w), 1.0, + msg="Callable greater than Function6D (f1() > f2()) did not return true when it should." + ) + self.assertEqual( + (lower_value > self.f1)(x, y, z, u, v, w), 0.0, + msg="Callable greater than Function6D (f1() > f2()) did not return false when it should." + ) + self.assertEqual( + (lower_value <= self.f1)(x, y, z, u, v, w), 1.0, + msg="Callable less equals Function6D (f1() <= f2()) did not return true when it should." + ) + self.assertEqual( + (ref_value <= self.f1)(x, y, z, u, v, w), 1.0, + msg="Callable less equals Function6D (f1() <= f2()) did not return true when it should." + ) + self.assertEqual( + (higher_value <= self.f1)(x, y, z, u, v, w), 0.0, + msg="Callable less equals Function6D (f1() <= f2()) did not return false when it should." + ) + self.assertEqual( + (higher_value >= self.f1)(x, y, z, u, v, w), 1.0, + msg="Callable equals Function6D (f1() >= f2()) did not return true when it should." + ) + self.assertEqual( + (ref_value >= self.f1)(x, y, z, u, v, w), 1.0, + msg="Callable greater equals Function6D (f1() >= f2()) did not return true when it should." + ) + self.assertEqual( + (lower_value >= self.f1)(x, y, z, u, v, w), 0.0, + msg="Callable equals Function6D (f1() >= f2()) did not return false when it should." + ) + + def test_richcmp_function_function(self): + testvals = [-1e10, -7, -0.001, 0.0, 0.00003, 10, 2.3e49] + for (x, y, z, u, v, w) in itertools.product(testvals, repeat=6): + ref_value = self.f1 + higher_value = self.f1 + abs(self.f1) + 1 + lower_value = self.f1 - abs(self.f1) - 1 + self.assertEqual( + (self.f1 == ref_value)(x, y, z, u, v, w), 1.0, + msg="Function6D equals Function6D (f1() == f2()) did not return true when it should." + ) + self.assertEqual( + (self.f1 == higher_value)(x, y, z, u, v, w), 0.0, + msg="Function6D equals Function6D (f1() == f2()) did not return false when it should." + ) + self.assertEqual( + (self.f1 != higher_value)(x, y, z, u, v, w), 1.0, + msg="Function6D not equals Function6D (f1() != f2()) did not return true when it should." + ) + self.assertEqual( + (self.f1 != ref_value)(x, y, z, u, v, w), 0.0, + msg="Function6D not equals Function6D (f1() != f2()) did not return false when it should." + ) + self.assertEqual( + (self.f1 < higher_value)(x, y, z, u, v, w), 1.0, + msg="Function6D less than Function6D (f1() < f2()) did not return true when it should." + ) + self.assertEqual( + (self.f1 < lower_value)(x, y, z, u, v, w), 0.0, + msg="Function6D less than Function6D (f1() < f2()) did not return false when it should." + ) + self.assertEqual( + (self.f1 > lower_value)(x, y, z, u, v, w), 1.0, + msg="Function6D greater than Function6D (f1() > f2()) did not return true when it should." + ) + self.assertEqual( + (self.f1 > higher_value)(x, y, z, u, v, w), 0.0, + msg="Function6D greater than Function6D (f1() > f2()) did not return false when it should." + ) + self.assertEqual( + (self.f1 <= higher_value)(x, y, z, u, v, w), 1.0, + msg="Function6D less equals Function6D (f1() <= f2()) did not return true when it should." + ) + self.assertEqual( + (self.f1 <= ref_value)(x, y, z, u, v, w), 1.0, + msg="Function6D less equals Function6D (f1() <= f2()) did not return true when it should." + ) + self.assertEqual( + (self.f1 <= lower_value)(x, y, z, u, v, w), 0.0, + msg="Function6D less equals Function6D (f1() <= f2()) did not return false when it should." + ) + self.assertEqual( + (self.f1 >= lower_value)(x, y, z, u, v, w), 1.0, + msg="Function6D equals Function6D (f1() >= f2()) did not return true when it should." + ) + self.assertEqual( + (self.f1 >= ref_value)(x, y, z, u, v, w), 1.0, + msg="Function6D greater equals Function6D (f1() >= f2()) did not return true when it should." + ) + self.assertEqual( + (self.f1 >= higher_value)(x, y, z, u, v, w), 0.0, + msg="Function6D equals Function6D (f1() >= f2()) did not return false when it should." + ) diff --git a/cherab/core/math/function/float/function6d/tests/test_cmath.py b/cherab/core/math/function/float/function6d/tests/test_cmath.py new file mode 100644 index 00000000..4593587c --- /dev/null +++ b/cherab/core/math/function/float/function6d/tests/test_cmath.py @@ -0,0 +1,115 @@ +# cython: language_level=3 + +# Copyright 2016-2025 Euratom +# Copyright 2016-2025 United Kingdom Atomic Energy Authority +# Copyright 2016-2025 Centro de Investigaciones Energéticas, Medioambientales y Tecnológicas +# +# Licensed under the EUPL, Version 1.1 or – as soon they will be approved by the +# European Commission - subsequent versions of the EUPL (the "Licence"); +# You may not use this work except in compliance with the Licence. +# You may obtain a copy of the Licence at: +# +# https://joinup.ec.europa.eu/software/page/eupl5 +# +# Unless required by applicable law or agreed to in writing, software distributed +# under the Licence is distributed on an "AS IS" basis, WITHOUT WARRANTIES OR +# CONDITIONS OF ANY KIND, either express or implied. +# +# See the Licence for the specific language governing permissions and limitations +# under the Licence. + +""" +Unit tests for the cmath wrapper classes. +""" + +import math +import unittest +import itertools +import cherab.core.math.function.float.function6d.cmath as cmath6d +from cherab.core.math.function.float.function6d.autowrap import PythonFunction6D + +# TODO: expand tests to cover the cython interface +class TestCmath6D(unittest.TestCase): + + def setUp(self): + self.f1 = PythonFunction6D(lambda x, y, z, u, v, w: x / 10 + y + z + u/2 + v/3 + w/4) + self.f2 = PythonFunction6D(lambda x, y, z, u, v, w: x * x + y * y - z * z + u * u + v * v - w * w) + + def test_exp(self): + testvals = [-10.0, -7, -0.001, 0.0, 0.00003, 10, 23.4] + for (x, y, z, u, v, w) in itertools.product(testvals, repeat=6): + function = cmath6d.Exp6D(self.f1) + expected = math.exp(self.f1(x, y, z, u, v, w)) + self.assertEqual(function(x, y, z, u, v, w), expected, "Exp6D call did not match reference value.") + + def test_sin(self): + testvals = [-10.0, -7, -0.001, 0.0, 0.00003, 10, 23.4] + for (x, y, z, u, v, w) in itertools.product(testvals, repeat=6): + function = cmath6d.Sin6D(self.f1) + expected = math.sin(self.f1(x, y, z, u, v, w)) + self.assertEqual(function(x, y, z, u, v, w), expected, "Sin6D call did not match reference value.") + + def test_cos(self): + testvals = [-10.0, -7, -0.001, 0.0, 0.00003, 10, 23.4] + for (x, y, z, u, v, w) in itertools.product(testvals, repeat=6): + function = cmath6d.Cos6D(self.f1) + expected = math.cos(self.f1(x, y, z, u, v, w)) + self.assertEqual(function(x, y, z, u, v, w), expected, "Cos6D call did not match reference value.") + + def test_tan(self): + testvals = [-10.0, -7, -0.001, 0.0, 0.00003, 10, 23.4] + for (x, y, z, u, v, w) in itertools.product(testvals, repeat=6): + function = cmath6d.Tan6D(self.f1) + expected = math.tan(self.f1(x, y, z, u, v, w)) + self.assertEqual(function(x, y, z, u, v, w), expected, "Tan6D call did not match reference value.") + + def test_asin(self): + v = [-10, -6, -2, -0.001, 0, 0.001, 2, 6, 10] + function = cmath6d.Asin6D(self.f1) + for x in v: + expected = math.asin(self.f1(x, 0, 0, 0, 0, 0)) + self.assertEqual(function(x, 0, 0, 0, 0, 0), expected, "Asin3D call did not match reference value.") + + with self.assertRaises(ValueError, msg="Asin3D did not raise a ValueError with value outside domain."): + function(100, 0, 0, 0, 0, 0) + + def test_acos(self): + v = [-10, -6, -2, -0.001, 0, 0.001, 2, 6, 10] + function = cmath6d.Acos6D(self.f1) + for x in v: + expected = math.acos(self.f1(x, 0, 0, 0, 0, 0)) + self.assertEqual(function(x, 0, 0, 0, 0, 0), expected, "Acos6D call did not match reference value.") + + with self.assertRaises(ValueError, msg="Acos3D did not raise a ValueError with value outside domain."): + function(100, 0, 0, 0, 0, 0) + + def test_atan(self): + testvals = [-10.0, -7, -0.001, 0.0, 0.00003, 10, 23.4] + for (x, y, z, u, v, w) in itertools.product(testvals, repeat=6): + function = cmath6d.Atan6D(self.f1) + expected = math.atan(self.f1(x, y, z, u, v, w)) + self.assertEqual(function(x, y, z, u, v, w), expected, "Atan6D call did not match reference value.") + + def test_atan2(self): + testvals = [-10.0, -7, -0.001, 0.0, 0.00003, 10, 23.4] + for (x, y, z, u, v, w) in itertools.product(testvals, repeat=6): + function = cmath6d.Atan4Q6D(self.f1, self.f2) + expected = math.atan2(self.f1(x, y, z, u, v, w), self.f2(x, y, z, u, v, w)) + self.assertEqual(function(x, y, z, u, v, w), expected, "Atan4Q6D call did not match reference value.") + + def test_erf(self): + testvals = [-1e5, -7, -0.001, 0.0, 0.00003, 10, 23.4, 1e5] + function = cmath6d.Erf6D(self.f1) + for (x, y, z, u, v, w) in itertools.product(testvals, repeat=6): + expected = math.erf(self.f1(x, y, z, u, v, w)) + self.assertAlmostEqual(function(x, y, z, u, v, w), expected, 10, "Erf6D call did not match reference value.") + + def test_sqrt(self): + testvals = [0.0, 0.00003, 10, 23.4, 1e5] + function = cmath6d.Sqrt6D(self.f1) + for (x, y, z, u, v, w) in itertools.product(testvals, repeat=6): + expected = math.sqrt(self.f1(x, y, z, u, v, w)) + self.assertEqual(function(x, y, z, u, v, w), expected, "Sqrt6D call did not match reference value.") + + with self.assertRaises(ValueError, msg="Sqrt6D did not raise a ValueError with value outside domain."): + function(-0.1, -0.1, -0.1, -0.1, -0.1, -0.1) \ No newline at end of file diff --git a/cherab/core/math/function/float/function6d/tests/test_constant.py b/cherab/core/math/function/float/function6d/tests/test_constant.py new file mode 100644 index 00000000..bc7e8cba --- /dev/null +++ b/cherab/core/math/function/float/function6d/tests/test_constant.py @@ -0,0 +1,36 @@ +# cython: language_level=3 + +# Copyright 2016-2025 Euratom +# Copyright 2016-2025 United Kingdom Atomic Energy Authority +# Copyright 2016-2025 Centro de Investigaciones Energéticas, Medioambientales y Tecnológicas +# +# Licensed under the EUPL, Version 1.1 or – as soon they will be approved by the +# European Commission - subsequent versions of the EUPL (the "Licence"); +# You may not use this work except in compliance with the Licence. +# You may obtain a copy of the Licence at: +# +# https://joinup.ec.europa.eu/software/page/eupl5 +# +# Unless required by applicable law or agreed to in writing, software distributed +# under the Licence is distributed on an "AS IS" basis, WITHOUT WARRANTIES OR +# CONDITIONS OF ANY KIND, either express or implied. +# +# See the Licence for the specific language governing permissions and limitations +# under the Licence. + +""" +Unit tests for the Constant6D class. +""" + +import unittest +from cherab.core.math.function.float.function6d.constant import Constant6D + +# TODO: expand tests to cover the cython interface +class TestConstant6D(unittest.TestCase): + + def test_constant(self): + testvals = [-1e10, -7, -0.001, 0.0, 0.00003, 10, 2.3e49] + for x in testvals: + constant = Constant6D(x) + # Test with a single set of values since it's a constant function + self.assertEqual(constant(500, 1.5, -3.14, 2.7, 1.8, 3.6), x, "Constant6D call did not match reference value.") diff --git a/cherab/core/math/integrators/__init__.pxd b/cherab/core/math/integrators/__init__.pxd index db06ef43..a0bf9d43 100644 --- a/cherab/core/math/integrators/__init__.pxd +++ b/cherab/core/math/integrators/__init__.pxd @@ -16,5 +16,6 @@ # See the Licence for the specific language governing permissions and limitations # under the Licence. -from cherab.core.math.integrators.integrators1d cimport Integrator1D, GaussianQuadrature +from cherab.core.math.integrators.integrators1d cimport Integrator1D, GaussianQuadrature1D, GaussianQuadrature +from cherab.core.math.integrators.integrators2d cimport Integrator2D, GaussianQuadrature2D diff --git a/cherab/core/math/integrators/__init__.py b/cherab/core/math/integrators/__init__.py index 86b7d58d..200c3d82 100644 --- a/cherab/core/math/integrators/__init__.py +++ b/cherab/core/math/integrators/__init__.py @@ -16,4 +16,5 @@ # See the Licence for the specific language governing permissions and limitations # under the Licence. -from .integrators1d import Integrator1D, GaussianQuadrature +from .integrators1d import Integrator1D, GaussianQuadrature1D, GaussianQuadrature +from .integrators2d import Integrator2D, GaussianQuadrature2D diff --git a/cherab/core/math/integrators/integrators1d.pxd b/cherab/core/math/integrators/integrators1d.pxd index c451285f..71a0ea6d 100644 --- a/cherab/core/math/integrators/integrators1d.pxd +++ b/cherab/core/math/integrators/integrators1d.pxd @@ -28,7 +28,7 @@ cdef class Integrator1D: cdef double evaluate(self, double a, double b) except? -1e999 -cdef class GaussianQuadrature(Integrator1D): +cdef class GaussianQuadrature1D(Integrator1D): cdef: int _min_order, _max_order @@ -37,3 +37,7 @@ cdef class GaussianQuadrature(Integrator1D): double[:] _roots_mv, _weights_mv cdef _build_cache(self) + + +cdef class GaussianQuadrature(GaussianQuadrature1D): + pass \ No newline at end of file diff --git a/cherab/core/math/integrators/integrators1d.pyx b/cherab/core/math/integrators/integrators1d.pyx index 7ff9be74..060773f9 100644 --- a/cherab/core/math/integrators/integrators1d.pyx +++ b/cherab/core/math/integrators/integrators1d.pyx @@ -18,6 +18,8 @@ # See the Licence for the specific language governing permissions and limitations # under the Licence. +import warnings + import numpy as np from scipy.special import roots_legendre @@ -39,7 +41,7 @@ cdef class Integrator1D: """ A 1D function to integrate. - :rtype: int + :rtype: Function1D """ return self.function @@ -65,23 +67,25 @@ cdef class Integrator1D: return self.evaluate(a, b) -cdef class GaussianQuadrature(Integrator1D): +cdef class GaussianQuadrature1D(Integrator1D): """ Compute an integral of a one-dimensional function over a finite interval using fixed-tolerance Gaussian quadrature. (see Scipy `quadrature `). + The integration is performed by iteratively increasing the order of the Gaussian quadrature until the relative tolerance is met or the maximum order is reached. + :param object integrand: A 1D function to integrate. Default is Constant1D(0). :param double relative_tolerance: Iteration stops when relative error between last two iterates is less than this value. Default is 1.e-5. - :param int max_order: Maximum order on Gaussian quadrature. Default is 50. - :param int min_order: Minimum order on Gaussian quadrature. Default is 1. + :param int max_order: Maximum order on Gaussian quadrature the integration stops at. Default is 50. + :param int min_order: Minimum order on Gaussian quadrature the integration starts from. Default is 1. :ivar Function1D integrand: A 1D function to integrate. :ivar double relative_tolerance: Iteration stops when relative error between last two iterates is less than this value. - :ivar int max_order: Maximum order on Gaussian quadrature. - :ivar int min_order: Minimum order on Gaussian quadrature. + :ivar int max_order: Maximum order on Gaussian quadrature the integration stops at. + :ivar int min_order: Minimum order on Gaussian quadrature the integration starts from. """ def __init__(self, object integrand=Constant1D(0), double relative_tolerance=1.e-5, int max_order=50, int min_order=1): @@ -169,6 +173,8 @@ cdef class GaussianQuadrature(Integrator1D): cdef: int order, n, i + # Store the variable-length roots and weights for each quadrature order + # consecutively in packed 1D arrays to avoid rectangular-array padding. n = (self._max_order + self._min_order) * (self._max_order - self._min_order + 1) // 2 self._roots = np.zeros(n, dtype=np.float64) @@ -222,3 +228,31 @@ cdef class GaussianQuadrature(Integrator1D): break return newval + + +cdef class GaussianQuadrature(GaussianQuadrature1D): + """ + Compute an integral of a one-dimensional function over a finite interval + using fixed-tolerance Gaussian quadrature. + (see Scipy `quadrature `). + + .. warning:: + This class is deprecated and will be removed in cherab 1.7. Use :class:`GaussianQuadrature1D` instead. + + :param object integrand: A 1D function to integrate. Default is Constant1D(0). + :param double relative_tolerance: Iteration stops when relative error between + last two iterates is less than this value. Default is 1.e-5. + :param int max_order: Maximum order on Gaussian quadrature. Default is 50. + :param int min_order: Minimum order on Gaussian quadrature. Default is 1. + + :ivar Function1D integrand: A 1D function to integrate. + :ivar double relative_tolerance: Iteration stops when relative error between + last two iterates is less than this value. + :ivar int max_order: Maximum order on Gaussian quadrature. + :ivar int min_order: Minimum order on Gaussian quadrature. + """ + + def __init__(self, object integrand=Constant1D(0), double relative_tolerance=1.e-5, int max_order=50, int min_order=1): + + warnings.warn("The GaussianQuadrature class is deprecated and will be removed in cherab 1.7. Use GaussianQuadrature1D instead.", DeprecationWarning, stacklevel=2) + super().__init__(integrand, relative_tolerance, max_order, min_order) \ No newline at end of file diff --git a/cherab/core/math/integrators/integrators2d.pxd b/cherab/core/math/integrators/integrators2d.pxd new file mode 100644 index 00000000..34238397 --- /dev/null +++ b/cherab/core/math/integrators/integrators2d.pxd @@ -0,0 +1,42 @@ +# Copyright 2016-2025 Euratom +# Copyright 2016-2025 United Kingdom Atomic Energy Authority +# Copyright 2016-2025 Centro de Investigaciones Energéticas, Medioambientales y Tecnológicas +# +# Licensed under the EUPL, Version 1.1 or – as soon they will be approved by the +# European Commission - subsequent versions of the EUPL (the "Licence"); +# You may not use this work except in compliance with the Licence. +# You may obtain a copy of the Licence at: +# +# https://joinup.ec.europa.eu/software/page/eupl5 +# +# Unless required by applicable law or agreed to in writing, software distributed +# under the Licence is distributed on an "AS IS" basis, WITHOUT WARRANTIES OR +# CONDITIONS OF ANY KIND, either express or implied. +# +# See the Licence for the specific language governing permissions and limitations +# under the Licence. + +from raysect.core.math.function.float cimport Function1D, Function2D + + +cdef class Integrator2D: + + cdef: + Function2D function + + cdef double evaluate(self,double x_lower, double x_upper, Function1D y_lower, Function1D y_upper) except? -1e999 + + +cdef class GaussianQuadrature2D(Integrator2D): + + cdef: + int _x_min_order, _x_max_order, _y_min_order, _y_max_order + double _rtol + object _x_roots, _x_weights, _y_roots, _y_weights + double[:] _x_roots_mv, _x_weights_mv, _y_roots_mv, _y_weights_mv + + cdef _build_cache(self) + + cdef double _evaluate_orders(self, double x_lower, double x_upper, Function1D y_lower, Function1D y_upper, int x_order, int y_order) except? -1e999 + + cdef inline Py_ssize_t _packed_offset(self, int order, int min_order) noexcept \ No newline at end of file diff --git a/cherab/core/math/integrators/integrators2d.pyx b/cherab/core/math/integrators/integrators2d.pyx new file mode 100644 index 00000000..080c4c9c --- /dev/null +++ b/cherab/core/math/integrators/integrators2d.pyx @@ -0,0 +1,450 @@ +# cython: language_level=3 + +# Copyright 2016-2025 Euratom +# Copyright 2016-2025 United Kingdom Atomic Energy Authority +# Copyright 2016-2025 Centro de Investigaciones Energéticas, Medioambientales y Tecnológicas +# +# Licensed under the EUPL, Version 1.1 or – as soon they will be approved by the +# European Commission - subsequent versions of the EUPL (the "Licence"); +# You may not use this work except in compliance with the Licence. +# You may obtain a copy of the Licence at: +# +# https://joinup.ec.europa.eu/software/page/eupl5 +# +# Unless required by applicable law or agreed to in writing, software distributed +# under the Licence is distributed on an "AS IS" basis, WITHOUT WARRANTIES OR +# CONDITIONS OF ANY KIND, either express or implied. +# +# See the Licence for the specific language governing permissions and limitations +# under the Licence. + +import numpy as np +from scipy.special import roots_legendre +cimport cython + +from raysect.core.math.function.float cimport Function1D, autowrap_function2d, Constant2D + +from libc.math cimport INFINITY + + +cdef class Integrator2D: + """ + Compute a definite integral of a two-dimensional function. + + :ivar Function2D integrand: A 2D function to integrate. + """ + + @property + def integrand(self): + """ + A 2D function to integrate. + + :rtype: Function2D + """ + return self.function + + @integrand.setter + def integrand(self, object func not None): + + self.function = autowrap_function2d(func) + + cdef double evaluate(self, double x_lower, double x_upper, Function1D y_lower, Function1D y_upper) except? -1e999: + + raise NotImplementedError("The evaluate() method has not been implemented.") + + def __call__(self, double x_lower, double x_upper, Function1D y_lower, Function1D y_upper): + """ + Integrates a two-dimensional function over a finite interval. + + :param double x_lower: Lower limit of integration in the x dimension. + :param double x_upper: Upper limit of integration in the x dimension. + :param Function1D y_lower: Lower limit of integration in the y dimension. + :param Function1D y_upper: Upper limit of integration in the y dimension. + + :returns: Definite integral of a two-dimensional function. + """ + + return self.evaluate(x_lower, x_upper, y_lower, y_upper) + + +cdef class GaussianQuadrature2D(Integrator2D): + r""" + Approximates an integral of a two-dimensional function over a finite interval. + + The integral is approximated with fixed-tolerance Gauss-Legendre quadrature. + The quadrature approximation is calculated as follows: + + .. math:: + \int_{x_{\mathrm{lower}}}^{x_{\mathrm{upper}}} \int_{y_{\mathrm{lower}}(x)}^{y_{\mathrm{upper}}(x)} f(x, y) \, dy \, dx + \approx \sum_{i=1}^{n_x} \sum_{j=1}^{n_y} w_{i} w_{j} \, + f\left( B \xi_i + A, \, D(x) \eta_j + C(x) \right) \, + B \cdot D(x) + + where: + - :math:`x_{\mathrm{lower}}`: Lower limit of integration for the x-dimension. + - :math:`x_{\mathrm{upper}}`: Upper limit of integration for the x-dimension. + - :math:`y_{\mathrm{lower}}(x)`: Lower limit of integration for the y-dimension, a function of :math:`x`. + - :math:`y_{\mathrm{upper}}(x)`: Upper limit of integration for the y-dimension, a function of :math:`x`. + - :math:`f(x, y)`: The function to be integrated over the specified region. + - :math:`\xi`: The transformed variable for the x-dimension, ranging from -1 to 1. + - :math:`\eta`: The transformed variable for the y-dimension, ranging from -1 to 1. + - :math:`w_i`: The weight corresponding to the :math:`i`-th root of the Legendre polynomial in the x-dimension. + - :math:`w_j`: The weight corresponding to the :math:`j`-th root of the Legendre polynomial in the y-dimension. + - :math:`\xi_i`: The :math:`i`-th root of the Legendre polynomial in the x-dimension. + - :math:`\eta_j`: The :math:`j`-th root of the Legendre polynomial in the y-dimension. + - :math:`n_x`: The number of roots (or nodes) in the x-dimension. + - :math:`n_y`: The number of roots (or nodes) in the y-dimension. + - :math:`A = \frac{x_{\mathrm{upper}} + x_{\mathrm{lower}}}{2}`: Midpoint of the x-interval. + - :math:`B = \frac{x_{\mathrm{upper}} - x_{\mathrm{lower}}}{2}`: Half-width of the x-interval. + - :math:`D(x) = \frac{y_{\mathrm{upper}}(x) - y_{\mathrm{lower}}(x)}{2}`: Half-width of the y-interval. + - :math:`C(x) = \frac{y_{\mathrm{upper}}(x) + y_{\mathrm{lower}}(x)}{2}`: Midpoint of the y-interval. + + The integration is performed by iteratively increasing the order of the Gaussian quadrature in both the x and y dimensions until the relative tolerance is met or the maximum orders are reached. + + :param Function2D integrand: A 2D function to integrate. Default is `Constant2D(0)`. + :param double relative_tolerance: Iteration stops when relative error between + last two iterates is less than this value. Default is `1.e-5`. + :param int x_max_order: Maximum order on Gaussian quadrature in the x dimension the integration stops at. Default is `50`. + :param int x_min_order: Minimum order on Gaussian quadrature in the x dimension the integration starts from. Default is `1`. + :param int y_max_order: Maximum order on Gaussian quadrature in the y dimension the integration stops at. Default is `50`. + :param int y_min_order: Minimum order on Gaussian quadrature in the y dimension the integration starts from. Default is `1`. + + :ivar Function1D integrand: A 1D function to integrate. + :ivar double relative_tolerance: Iteration stops when relative error between + last two iterates is less than this value. + :ivar int x_max_order: Maximum order on Gaussian quadrature in the x dimension the integration stops at. + :ivar int x_min_order: Minimum order on Gaussian quadrature in the x dimension the integration starts from. + :ivar int y_max_order: Maximum order on Gaussian quadrature in the y dimension the integration stops at. + :ivar int y_min_order: Minimum order on Gaussian quadrature in the y dimension the integration starts from. + """ + def __init__(self, object integrand=Constant2D(0), double relative_tolerance=1.e-5, int x_max_order=50, int x_min_order=1, + int y_max_order=50, int y_min_order=1): + + self._check_order(x_min_order, x_max_order, "x") + self._check_order(y_min_order, y_max_order, "y") + + self._x_min_order = x_min_order + self._x_max_order = x_max_order + + self._y_min_order = y_min_order + self._y_max_order = y_max_order + + self._build_cache() + + self.integrand = integrand + + self.relative_tolerance = relative_tolerance + + def _check_order(self, int min_order, int max_order, str dimension): + """ + Check the order of Gaussian quadrature. + + :param int min_order: Minimum order on Gaussian quadrature. + :param int max_order: Maximum order on Gaussian quadrature. + :param str dimension: Dimension of Gaussian quadrature. + + :raises ValueError: If the order of Gaussian quadrature is invalid. + """ + + if min_order < 1 or max_order < 1: + raise ValueError("Order of Gaussian quadrature in the {} dimension must be >= 1.".format(dimension)) + + if min_order > max_order: + raise ValueError("Minimum order of Gaussian quadrature in the {} dimension must be less than or equal to the maximum order.".format(dimension)) + + @property + def x_min_order(self): + """ + Minimum order on Gaussian quadrature in the x dimension. + + :rtype: int + """ + return self._x_min_order + + @x_min_order.setter + def x_min_order(self, int value): + + self._check_order(value, self._x_max_order, "x") + + self._x_min_order = value + + self._build_cache() + + @property + def x_max_order(self): + """ + Maximum order on Gaussian quadrature in the x dimension. + + :rtype: int + """ + return self._x_max_order + + @x_max_order.setter + def x_max_order(self, int value): + + self._check_order(self._x_min_order, value, "x") + + self._x_max_order = value + + self._build_cache() + + @property + def y_min_order(self): + """ + Minimum order on Gaussian quadrature in the y dimension. + + :rtype: int + """ + return self._y_min_order + + @y_min_order.setter + def y_min_order(self, int value): + + self._check_order(value, self._y_max_order, "y") + + self._y_min_order = value + + self._build_cache() + + @property + def y_max_order(self): + """ + Maximum order on Gaussian quadrature in the y dimension. + + :rtype: int + """ + return self._y_max_order + + @y_max_order.setter + def y_max_order(self, int value): + + self._check_order(self._y_min_order, value, "y") + + self._y_max_order = value + + self._build_cache() + + @property + def relative_tolerance(self): + """ + Iteration stops when relative error between last two iterates is less than this value. + + :rtype: double + """ + return self._rtol + + @relative_tolerance.setter + def relative_tolerance(self, double value): + + if value <= 0: + raise ValueError("Relative tolerance must be positive.") + + self._rtol = value + + cdef _build_cache(self): + """ + Caches the roots and weights of the Gauss-Legendre quadrature. + """ + + cdef: + int order, n, i + + # Pack the variable-length quadrature rules for each coordinate direction into + # contiguous 1D arrays, avoiding the unused padding of rectangular caches. + # x-direction + n = (self._x_max_order + self._x_min_order) * (self._x_max_order - self._x_min_order + 1) // 2 + + self._x_roots = np.zeros(n, dtype=np.float64) + self._x_weights = np.zeros(n, dtype=np.float64) + + i = 0 + for order in range(self._x_min_order, self._x_max_order + 1): + self._x_roots[i:i + order], self._x_weights[i:i + order] = roots_legendre(order) + i += order + + self._x_roots_mv = self._x_roots + self._x_weights_mv = self._x_weights + + # y-direction + n = (self._y_max_order + self._y_min_order) * (self._y_max_order - self._y_min_order + 1) // 2 + + self._y_roots = np.zeros(n, dtype=np.float64) + self._y_weights = np.zeros(n, dtype=np.float64) + + i = 0 + for order in range(self._y_min_order, self._y_max_order + 1): + self._y_roots[i:i + order], self._y_weights[i:i + order] = roots_legendre(order) + i += order + + self._y_roots_mv = self._y_roots + self._y_weights_mv = self._y_weights + + @cython.boundscheck(False) + @cython.wraparound(False) + @cython.cdivision(True) + @cython.initializedcheck(False) + cdef double evaluate(self, double x_lower, double x_upper, Function1D y_lower, Function1D y_upper) except? -1e999: + """ + Integrates a two-dimensional function over a finite interval. + + :param double x_lower: Lower limit of integration in the x dimension. + :param double x_upper: Upper limit of integration in the x dimension. + :param Function1D y_lower: Lower limit of integration in the y dimension as a function of x. + :param Function1D y_upper: Upper limit of integration in the y dimension as a function of x. + + :returns: Gaussian quadrature approximation to integral. + """ + + cdef: + int x_order = self._x_min_order + int y_order = self._y_min_order + + double previous_integral + double current_integral + double rtol = self._rtol + + current_integral = self._evaluate_orders( + x_lower, + x_upper, + y_lower, + y_upper, + x_order, + y_order, + ) + + while ( + x_order < self._x_max_order + or y_order < self._y_max_order + ): + previous_integral = current_integral + + if x_order < self._x_max_order: + x_order += 1 + + if y_order < self._y_max_order: + y_order += 1 + + current_integral = self._evaluate_orders( + x_lower, + x_upper, + y_lower, + y_upper, + x_order, + y_order, + ) + + if ( + abs(current_integral - previous_integral) + <= rtol * abs(current_integral) + ): + return current_integral + + return current_integral + + @cython.boundscheck(False) + @cython.wraparound(False) + @cython.cdivision(True) + @cython.initializedcheck(False) + cdef double _evaluate_orders( + self, + double x_lower, + double x_upper, + Function1D y_lower, + Function1D y_upper, + int x_order, + int y_order, + ) except? -1e999: + """ + Evaluate the quadrature using fixed x and y orders. + + :param double x_lower: Lower limit of integration in the x dimension. + :param double x_upper: Upper limit of integration in the x dimension. + :param Function1D y_lower: Lower limit of integration in the y dimension as a function of x. + :param Function1D y_upper: Upper limit of integration in the y dimension as a function of x. + :param int x_order: Order of Gaussian quadrature in the x dimension. + :param int y_order: Order of Gaussian quadrature in the y dimension. + + :returns: Gaussian quadrature approximation to integral. + """ + + cdef: + Py_ssize_t x_ibegin, y_ibegin + Py_ssize_t i, j + + double integral + double y_contribution + + double x, y + + double x_offset, x_slope + double y_offset, y_slope + double y_lower_val, y_upper_val + + x_ibegin = self._packed_offset( + x_order, + self._x_min_order, + ) + + y_ibegin = self._packed_offset( + y_order, + self._y_min_order, + ) + + # Transform the x interval from [-1, 1] to [x_lower, x_upper]. + x_offset = 0.5 * (x_lower + x_upper) + x_slope = 0.5 * (x_upper - x_lower) + + integral = 0. + + for i in range(x_ibegin, x_ibegin + x_order): + + x = x_offset + x_slope * self._x_roots_mv[i] + + y_lower_val = y_lower.evaluate(x) + y_upper_val = y_upper.evaluate(x) + + # Transform the y interval from [-1, 1] to + # [y_lower(x), y_upper(x)]. + y_offset = 0.5 * (y_lower_val + y_upper_val) + y_slope = 0.5 * (y_upper_val - y_lower_val) + + y_contribution = 0. + + for j in range(y_ibegin, y_ibegin + y_order): + + y = y_offset + y_slope * self._y_roots_mv[j] + + y_contribution += ( + self._y_weights_mv[j] + * self.function.evaluate(x, y) + ) + + # y_slope depends on x, so it must be applied separately + # for every x quadrature node. + integral += ( + self._x_weights_mv[i] + * y_slope + * y_contribution + ) + + return x_slope * integral + + cdef inline Py_ssize_t _packed_offset( + self, + int order, + int min_order, + ) noexcept: + """ + Return the start index of a quadrature rule in a packed cache of roots and weights. + + :param int order: Order of Gaussian quadrature. + :param int min_order: Minimum order of Gaussian quadrature. + + :returns: Start index of a quadrature rule in a packed cache of roots and weights. + """ + + return ( + (order - min_order) + * (order + min_order - 1) + // 2 + ) \ No newline at end of file diff --git a/cherab/core/math/mask.pyx b/cherab/core/math/mask.pyx index 14b8a3a1..91894908 100644 --- a/cherab/core/math/mask.pyx +++ b/cherab/core/math/mask.pyx @@ -65,6 +65,3 @@ cdef class PolygonMask2D(Function2D): cdef double evaluate(self, double x, double y) except? -1e999: return self._mesh.evaluate(x, y) - - - diff --git a/cherab/core/math/samplers.pyx b/cherab/core/math/samplers.pyx index 482dc813..caa027f7 100644 --- a/cherab/core/math/samplers.pyx +++ b/cherab/core/math/samplers.pyx @@ -42,7 +42,7 @@ cpdef tuple sample1d(object function1d, tuple x_range): :param function1d: a Python function or Function1D object :param x_range: a tuple defining the sample range: (min, max, samples) :return: a tuple containing the sampled values: (x_points, function_samples) - + .. code-block:: pycon >>> from cherab.core.math import sample1d @@ -221,7 +221,7 @@ cpdef np.ndarray sample2d_points(object function2d, object points): .. code-block:: pycon - >>> from cherab.core.math import sample2d + >>> from cherab.core.math import sample2d_points >>> >>> def f1(x, y): >>> return x**2 + y @@ -316,7 +316,7 @@ cpdef tuple sample3d(object function3d, tuple x_range, tuple y_range, tuple z_ra """ Samples a 3D function over the specified range. - :param function3d: a Python function or Function2D object + :param function3d: a Python function or Function3D object :param x_range: a tuple defining the x sample range: (x_min, x_max, x_samples) :param y_range: a tuple defining the y sample range: (y_min, y_max, y_samples) :param z_range: a tuple defining the z sample range: (z_min, z_max, z_samples) @@ -335,7 +335,7 @@ cpdef tuple sample3d(object function3d, tuple x_range, tuple y_range, tuple z_ra >>> f_vals array([[[ 3., 4., 5.], [ 6., 7., 8.], - [11., 12., 13.]], + [11., 12., 13.]], [[10., 11., 12.], [13., 14., 15.], [18., 19., 20.]], @@ -415,7 +415,7 @@ cpdef np.ndarray sample3d_points(object function3d, object points): :param function3d: a Python function or Function3D object :param points: an Nx3 array of points at which to sample the function :return: a 1D array containing the sampled values at each point - + .. code-block:: pycon >>> from cherab.core.math import sample3d_points @@ -744,7 +744,7 @@ cpdef tuple samplevector3d(object function3d, tuple x_range, tuple y_range, tupl The function samples returns are an NxMxKx3 array where the last axis are the x, y, and z components of the vector respectively. - :param function3d: a Python function or Function2D object + :param function3d: a Python function or Function3D object :param x_range: a tuple defining the x sample range: (x_min, x_max, x_samples) :param y_range: a tuple defining the y sample range: (y_min, y_max, y_samples) :param z_range: a tuple defining the z sample range: (z_min, z_max, z_samples) diff --git a/cherab/core/math/tests/test_integrators.py b/cherab/core/math/tests/test_integrators.py index fa967165..0c2f9ef1 100644 --- a/cherab/core/math/tests/test_integrators.py +++ b/cherab/core/math/tests/test_integrators.py @@ -16,22 +16,28 @@ # See the Licence for the specific language governing permissions and limitations # under the Licence. -from raysect.core.math.function.float import Exp1D, Arg1D -from cherab.core.math.integrators import GaussianQuadrature +from raysect.core.math.function.float import Exp1D, Arg1D, Exp2D, Arg2D, Constant1D +from cherab.core.math.integrators import GaussianQuadrature1D, GaussianQuadrature2D from math import sqrt, pi from scipy.special import erf import unittest +import itertools -class TestGaussianQuadrature(unittest.TestCase): +class TestGaussianQuadrature1D(unittest.TestCase): """Gaussian quadrature integrator tests.""" def test_properties(self): """Test property assignment.""" min_order = 3 max_order = 30 - reltol = 1.e-6 - quadrature = GaussianQuadrature(integrand=Arg1D, relative_tolerance=reltol, max_order=max_order, min_order=min_order) + reltol = 1.0e-6 + quadrature = GaussianQuadrature1D( + integrand=Arg1D, + relative_tolerance=reltol, + max_order=max_order, + min_order=min_order, + ) self.assertEqual(quadrature.relative_tolerance, reltol) self.assertEqual(quadrature.max_order, max_order) @@ -53,7 +59,7 @@ def test_properties(self): min_order = 1 max_order = 20 - reltol = 1.e-5 + reltol = 1.0e-5 quadrature.relative_tolerance = reltol quadrature.min_order = min_order @@ -67,14 +73,133 @@ def test_properties(self): def test_integrate(self): """Test integration.""" - quadrature = GaussianQuadrature(relative_tolerance=1.e-8) + quadrature = GaussianQuadrature1D(relative_tolerance=1.0e-8) a = -0.5 - b = 3. - quadrature.integrand = (2 / sqrt(pi)) * Exp1D(- Arg1D() * Arg1D()) + b = 3.0 + quadrature.integrand = (2 / sqrt(pi)) * Exp1D(-Arg1D() * Arg1D()) exact_integral = erf(b) - erf(a) self.assertAlmostEqual(quadrature(a, b), exact_integral, places=8) -if __name__ == '__main__': +class TestGaussianQuadrature2D(unittest.TestCase): + """Gaussian quadrature 2D integrator tests.""" + + def test_properties(self): + """Test property assignment.""" + x_min_order = 3 + y_min_order = 4 + x_max_order = 30 + y_max_order = 40 + reltol = 1.0e-6 + integrand = Exp2D(Arg2D("x") + Arg2D("y")) + + quadrature = GaussianQuadrature2D( + integrand=integrand, + relative_tolerance=reltol, + x_max_order=x_max_order, + x_min_order=x_min_order, + y_max_order=y_max_order, + y_min_order=y_min_order, + ) + + self.assertEqual(quadrature.relative_tolerance, reltol) + self.assertEqual(quadrature.x_max_order, x_max_order) + self.assertEqual(quadrature.y_max_order, y_max_order) + self.assertEqual(quadrature.x_min_order, x_min_order) + self.assertEqual(quadrature.y_min_order, y_min_order) + self.assertEqual(quadrature.integrand, integrand) + + x_min_order = 50 # > x_max_order + x_max_order = 2 # < x_min_order + y_min_order = 50 # > y_max_order + y_max_order = 1 # < y_min_order + reltol = -1 + + with self.assertRaises(ValueError): + quadrature.x_max_order = x_max_order + + with self.assertRaises(ValueError): + quadrature.x_min_order = x_min_order + + with self.assertRaises(ValueError): + quadrature.y_max_order = y_max_order + + with self.assertRaises(ValueError): + quadrature.y_min_order = y_min_order + + with self.assertRaises(ValueError): + quadrature.relative_tolerance = reltol + + x_min_order = 0 + y_min_order = 0 + + with self.assertRaises(ValueError): + quadrature.x_min_order = x_min_order + + with self.assertRaises(ValueError): + quadrature.y_min_order = y_min_order + + x_min_order = 1 + x_max_order = 20 + y_min_order = 2 + y_max_order = 30 + reltol = 1.0e-5 + + quadrature.relative_tolerance = reltol + quadrature.x_min_order = x_min_order + quadrature.x_max_order = x_max_order + quadrature.y_min_order = y_min_order + quadrature.y_max_order = y_max_order + quadrature.integrand = Arg2D("x") + + self.assertEqual(quadrature.relative_tolerance, reltol) + self.assertEqual(quadrature.x_min_order, x_min_order) + self.assertEqual(quadrature.x_max_order, x_max_order) + self.assertEqual(quadrature.y_min_order, y_min_order) + self.assertEqual(quadrature.y_max_order, y_max_order) + self.assertEqual(quadrature.integrand, Arg2D("x")) + + def test_integrate(self): + """Test 2D integration.""" + + max_orders = [40, 50, 60] + min_orders = [10, 20, 30] + + for x_max_order, y_max_order in itertools.product(max_orders, repeat=2): + for x_min_order, y_min_order in itertools.product(min_orders, repeat=2): + quadrature = GaussianQuadrature2D( + relative_tolerance=1.0e-8, + x_max_order=x_max_order, + y_max_order=y_max_order, + x_min_order=x_min_order, + y_min_order=y_min_order, + ) + + # Integration limits + a_x, b_x = -2.0, 2.0 + a_y, b_y = -3.0, 3.0 + + # Bivariate Normal distribution with std_dev=1, mean=0 and no correlation + quadrature.integrand = ( + 1 / (2 * pi) * Exp2D(-0.5 * (Arg2D("x") ** 2 + Arg2D("y") ** 2)) + ) + + # Exact integral of the bivariate normal distribution + exact_integral = ( + 1 + / 4.0 + * (erf(b_x / sqrt(2)) - erf(a_x / sqrt(2))) + * (erf(b_y / sqrt(2)) - erf(a_y / sqrt(2))) + ) + + self.assertAlmostEqual( + quadrature(a_x, b_x, Constant1D(a_y), Constant1D(b_y)), + exact_integral, + places=8, + msg=f"x_max_order={x_max_order}, y_max_order={y_max_order}, x_min_order={x_min_order}, y_min_order={y_min_order}", + ) + + +if __name__ == "__main__": unittest.main() diff --git a/cherab/core/math/transform/periodic.pyx b/cherab/core/math/transform/periodic.pyx index 13869458..617abfce 100644 --- a/cherab/core/math/transform/periodic.pyx +++ b/cherab/core/math/transform/periodic.pyx @@ -306,7 +306,7 @@ cdef class VectorPeriodicTransform3D(VectorFunction3D): .. code-block:: pycon - >>> from cherab.core.math import PeriodicTransform3D + >>> from cherab.core.math import VectorPeriodicTransform3D >>> >>> def f1(x, y, z): >>> return Vector3D(x, y, z) @@ -327,7 +327,7 @@ cdef class VectorPeriodicTransform3D(VectorFunction3D): def __init__(self, object function3d, double period_x, double period_y, double period_z): if not callable(function3d): - raise TypeError("function2d is not callable.") + raise TypeError("function3d is not callable.") self.function3d = autowrap_vectorfunction3d(function3d) diff --git a/cherab/core/model/laser/profile.pyx b/cherab/core/model/laser/profile.pyx index 82376980..ed5f6024 100644 --- a/cherab/core/model/laser/profile.pyx +++ b/cherab/core/model/laser/profile.pyx @@ -738,7 +738,7 @@ def generate_segmented_cylinder(radius, length): Generates a segmented cylindrical laser geometry Approximates a long cylinder with a cylindrical segments to optimize - targetted and importance sampling. The height of a cylinder segments is roughly + targeted and importance sampling. The height of a cylinder segments is roughly 2 * cylinder radius. :return: List of cylinders diff --git a/cherab/core/model/lineshape/stark.pyx b/cherab/core/model/lineshape/stark.pyx index 9e333a5f..8b16f174 100644 --- a/cherab/core/model/lineshape/stark.pyx +++ b/cherab/core/model/lineshape/stark.pyx @@ -29,7 +29,7 @@ from cherab.core.species cimport Species from cherab.core.plasma cimport Plasma from cherab.core.atomic.elements import hydrogen, deuterium, tritium from cherab.core.math.function cimport autowrap_function1d, autowrap_function2d -from cherab.core.math.integrators cimport GaussianQuadrature +from cherab.core.math.integrators cimport GaussianQuadrature1D from cherab.core.utility.constants cimport BOHR_MAGNETON, HC_EV_NM from cherab.core.model.lineshape.doppler cimport doppler_shift, thermal_broadening from cherab.core.model.lineshape.gaussian cimport add_gaussian_line @@ -211,7 +211,7 @@ cdef class StarkBroadenedLine(ZeemanLineShapeModel): Default is None (will use `atomic_data.stark_model_coefficients`). :param Integrator1D integrator: Integrator1D instance to integrate the line shape - over the spectral bin. Default is `GaussianQuadrature()`. + over the spectral bin. Default is `GaussianQuadrature1D()`. :param str polarisation: Leaves only :math:`\pi`-/:math:`\sigma`-polarised components: "pi" - leave only :math:`\pi`-polarised components, "sigma" - leave only :math:`\sigma`-polarised components, @@ -219,7 +219,7 @@ cdef class StarkBroadenedLine(ZeemanLineShapeModel): """ def __init__(self, Line line, double wavelength, Species target_species, Plasma plasma, AtomicData atomic_data, - tuple stark_model_coefficients=None, Integrator1D integrator=GaussianQuadrature(), polarisation='no'): + tuple stark_model_coefficients=None, Integrator1D integrator=GaussianQuadrature1D(), polarisation='no'): super().__init__(line, wavelength, target_species, plasma, atomic_data, polarisation, integrator) diff --git a/cherab/core/model/plasma/bremsstrahlung.pyx b/cherab/core/model/plasma/bremsstrahlung.pyx index 404750a2..453489d5 100644 --- a/cherab/core/model/plasma/bremsstrahlung.pyx +++ b/cherab/core/model/plasma/bremsstrahlung.pyx @@ -21,7 +21,7 @@ import numpy as np from raysect.optical cimport Spectrum, Point3D, Vector3D from cherab.core cimport Plasma, AtomicData -from cherab.core.math.integrators cimport GaussianQuadrature +from cherab.core.math.integrators cimport GaussianQuadrature1D from cherab.core.species cimport Species from cherab.core.utility.constants cimport RECIP_4_PI, ELEMENTARY_CHARGE, SPEED_OF_LIGHT, PLANCK_CONSTANT, ELECTRON_REST_MASS, VACUUM_PERMITTIVITY from libc.math cimport sqrt, log, exp, M_PI @@ -123,7 +123,7 @@ cdef class Bremsstrahlung(PlasmaModel): wavelength. If not provided, the `atomic_data` is used. :ivar Integrator1D integrator: Integrator1D instance to integrate Bremsstrahlung radiation - over the spectral bin. Default is `GaussianQuadrature`. + over the spectral bin. Default is `GaussianQuadrature1D`. """ def __init__(self, Plasma plasma=None, AtomicData atomic_data=None, FreeFreeGauntFactor gaunt_factor=None, Integrator1D integrator=None): @@ -132,7 +132,7 @@ cdef class Bremsstrahlung(PlasmaModel): self._brems_func = BremsFunction.__new__(BremsFunction) self.gaunt_factor = gaunt_factor - self.integrator = integrator or GaussianQuadrature() + self.integrator = integrator or GaussianQuadrature1D() # ensure that cache is initialised self._change() diff --git a/cherab/core/model/plasma/impact_excitation.pyx b/cherab/core/model/plasma/impact_excitation.pyx index b26336f1..5c12d94f 100644 --- a/cherab/core/model/plasma/impact_excitation.pyx +++ b/cherab/core/model/plasma/impact_excitation.pyx @@ -45,8 +45,10 @@ cdef class ExcitationLine(PlasmaModel): :param object lineshape_args: A list of line shape model arguments. Default is None. :param object lineshape_kwargs: A dictionary of line shape model keyword arguments. Default is None. - :ivar Plasma plasma: The plasma to which this emission model is attached. - :ivar AtomicData atomic_data: The atomic data provider for this model. + :ivar Plasma plasma: See parameter 'plasma'. + :ivar AtomicData atomic_data: See parameter 'atomic_data'. + :ivar Line line: The emission line object. + :ivar LineShapeModel lineshape: The line shape model. """ def __init__(self, Line line, Plasma plasma=None, AtomicData atomic_data=None, object lineshape=None, @@ -75,6 +77,14 @@ cdef class ExcitationLine(PlasmaModel): def __repr__(self): return ''.format(self._line.element.name, self._line.charge, self._line.transition) + @property + def line(self) -> Line: + return self._line + + @property + def lineshape(self) -> LineShapeModel: + return self._lineshape + cpdef Spectrum emission(self, Point3D point, Vector3D direction, Spectrum spectrum): cdef double ne, ni, te, radiance diff --git a/cherab/core/model/plasma/recombination.pyx b/cherab/core/model/plasma/recombination.pyx index a33009d0..7f85dad8 100644 --- a/cherab/core/model/plasma/recombination.pyx +++ b/cherab/core/model/plasma/recombination.pyx @@ -45,8 +45,10 @@ cdef class RecombinationLine(PlasmaModel): :param object lineshape_args: A list of line shape model arguments. Default is None. :param object lineshape_kwargs: A dictionary of line shape model keyword arguments. Default is None. - :ivar Plasma plasma: The plasma to which this emission model is attached. - :ivar AtomicData atomic_data: The atomic data provider for this model. + :ivar Plasma plasma: See parameter 'plasma'. + :ivar AtomicData atomic_data: See parameter 'atomic_data'. + :ivar Line line: The emission line object. + :ivar LineShapeModel lineshape: The line shape model. """ def __init__(self, Line line, Plasma plasma=None, AtomicData atomic_data=None, object lineshape=None, @@ -75,6 +77,14 @@ cdef class RecombinationLine(PlasmaModel): def __repr__(self): return ''.format(self._line.element.name, self._line.charge, self._line.transition) + @property + def line(self) -> Line: + return self._line + + @property + def lineshape(self) -> LineShapeModel: + return self._lineshape + cpdef Spectrum emission(self, Point3D point, Vector3D direction, Spectrum spectrum): cdef double ne, ni, te, radiance diff --git a/cherab/core/model/plasma/tests/__init__.py b/cherab/core/model/plasma/tests/__init__.py new file mode 100644 index 00000000..e69de29b diff --git a/cherab/core/model/plasma/tests/test_models.py b/cherab/core/model/plasma/tests/test_models.py new file mode 100644 index 00000000..b0a15417 --- /dev/null +++ b/cherab/core/model/plasma/tests/test_models.py @@ -0,0 +1,105 @@ +import unittest +from unittest.mock import patch + +import numpy as np + +from raysect.optical import Point3D, Vector3D, Spectrum + +from cherab.core.model.plasma import ( + ExcitationLine, + RecombinationLine, + TotalRadiatedPower, +) +from cherab.core.atomic import Line, hydrogen +from cherab.core.model import GaussianLine +from cherab.tools.plasmas.slab import build_slab_plasma +from cherab.openadas import OpenADAS + + +class TestPlasmaModels(unittest.TestCase): + # make a slab plasma + + plasma = build_slab_plasma(peak_density=5e19) + plasma.atomic_data = OpenADAS(permit_extrapolation=True) + balmer_alpha = Line(hydrogen, 0, (3, 2)) + + def setUp(self): + # setup mock to avoid reading the data from the repository + self.patcher_excitation = patch( + "cherab.openadas.openadas.repository.get_pec_excitation_rate", + return_value={ + "ne": np.linspace(1e18, 1e20, 10), + "te": np.linspace(1, 1e3, 12), + "rate": np.ones((10, 12)), + }, + ) + self.mock_get_excitation = self.patcher_excitation.start() + + self.patcher_recombination = patch( + "cherab.openadas.openadas.repository.get_pec_recombination_rate", + return_value={ + "ne": np.linspace(1e18, 1e20, 10), + "te": np.linspace(1, 1e3, 12), + "rate": np.ones((10, 12)), + }, + ) + self.mock_get_recombination = self.patcher_recombination.start() + + self.patcher_wl = patch( + "cherab.openadas.openadas.repository.get_wavelength", return_value=656.28 + ) + self.mock_get_wavelength = self.patcher_wl.start() + + def tearDown(self): + # stop the mocks after a test is run + self.patcher_excitation.stop() + self.patcher_recombination.stop() + self.patcher_wl.stop() + + def test_excitation(self): + exc = ExcitationLine(self.balmer_alpha) + self.plasma.models = [exc] + + # sample emission to trigger the caching mechanism + exc.emission(Point3D(0, 0, 0), Vector3D(0, 0, 1), Spectrum(300, 1000, 1000)) + + # check exc has the correct line + self.assertEqual(exc.line, self.balmer_alpha) + + # check exc has the correct lineshape + self.assertIsInstance(exc.lineshape, GaussianLine) + + # check the mock was called + self.mock_get_excitation.assert_called_once() + self.assertEqual(self.mock_get_wavelength.call_count, 2) + + def test_recombination(self): + rec = RecombinationLine(self.balmer_alpha) + self.plasma.models = [rec] + + # sample emission to trigger the caching mechanism + rec.emission(Point3D(0, 0, 0), Vector3D(0, 0, 1), Spectrum(300, 1000, 1000)) + + # check rec has the correct line + self.assertEqual(rec.line, self.balmer_alpha) + + # check rec has the correct lineshape + self.assertIsInstance(rec.lineshape, GaussianLine) + + # check the mock was called + self.mock_get_recombination.assert_called_once() + self.assertEqual(self.mock_get_wavelength.call_count, 2) + + def test_total_radiated_power(self): + trp = TotalRadiatedPower(hydrogen, 0) + self.plasma.models = [trp] + + # check initialisation + with self.assertRaises(ValueError): + TotalRadiatedPower(hydrogen, 2) + with self.assertRaises(ValueError): + TotalRadiatedPower(hydrogen, -1) + + # check trp has the correct element and charge + self.assertEqual(trp.element, hydrogen) + self.assertEqual(trp.charge, 0) diff --git a/cherab/core/model/plasma/thermal_cx.pyx b/cherab/core/model/plasma/thermal_cx.pyx index 88af9ae8..e0943a14 100644 --- a/cherab/core/model/plasma/thermal_cx.pyx +++ b/cherab/core/model/plasma/thermal_cx.pyx @@ -47,8 +47,10 @@ cdef class ThermalCXLine(PlasmaModel): :param object lineshape_args: A list of line shape model arguments. Default is None. :param object lineshape_kwargs: A dictionary of line shape model keyword arguments. Default is None. - :ivar Plasma plasma: The plasma to which this emission model is attached. - :ivar AtomicData atomic_data: The atomic data provider for this model. + :ivar Plasma plasma: See parameter 'plasma'. + :ivar AtomicData atomic_data: See parameter 'atomic_data'. + :ivar Line line: The emission line object. + :ivar LineShapeModel lineshape: The line shape model. """ def __init__(self, Line line, Plasma plasma=None, AtomicData atomic_data=None, object lineshape=None, @@ -77,6 +79,14 @@ cdef class ThermalCXLine(PlasmaModel): def __repr__(self): return ''.format(self._line.element.name, self._line.charge, self._line.transition) + @property + def line(self) -> Line: + return self._line + + @property + def lineshape(self) -> LineShapeModel: + return self._lineshape + cpdef Spectrum emission(self, Point3D point, Vector3D direction, Spectrum spectrum): cdef: diff --git a/cherab/core/model/plasma/total_radiated_power.pyx b/cherab/core/model/plasma/total_radiated_power.pyx index 84b42424..f0c94281 100644 --- a/cherab/core/model/plasma/total_radiated_power.pyx +++ b/cherab/core/model/plasma/total_radiated_power.pyx @@ -53,6 +53,9 @@ cdef class TotalRadiatedPower(PlasmaModel): :param int charge: The charge state of the element/isotope. :param Plasma plasma: The plasma to which this emission model is attached. Default is None. :param AtomicData atomic_data: The atomic data provider for this model. Default is None. + + :ivar Element element: See parameter 'element'. + :ivar int charge: See parameter 'charge'. """ def __init__(self, Element element, int charge, Plasma plasma=None, AtomicData atomic_data=None): @@ -68,6 +71,14 @@ cdef class TotalRadiatedPower(PlasmaModel): # ensure that cache is initialised self._change() + + @property + def element(self) -> Element: + return self._element + + @property + def charge(self) -> int: + return self._charge cpdef Spectrum emission(self, Point3D point, Vector3D direction, Spectrum spectrum): diff --git a/cherab/core/plasma/node.pxd b/cherab/core/plasma/node.pxd index f7a9bb89..9d268070 100644 --- a/cherab/core/plasma/node.pxd +++ b/cherab/core/plasma/node.pxd @@ -60,6 +60,7 @@ cdef class Plasma(Node): readonly object notifier VectorFunction3D _b_field + VectorFunction3D _e_field DistributionFunction _electron_distribution Composition _composition AtomicData _atomic_data @@ -72,6 +73,8 @@ cdef class Plasma(Node): cdef VectorFunction3D get_b_field(self) + cdef VectorFunction3D get_e_field(self) + cdef DistributionFunction get_electron_distribution(self) cdef Composition get_composition(self) diff --git a/cherab/core/plasma/node.pyx b/cherab/core/plasma/node.pyx index 08d31069..96f687ca 100644 --- a/cherab/core/plasma/node.pyx +++ b/cherab/core/plasma/node.pyx @@ -258,6 +258,8 @@ cdef class Plasma(Node): All plasma emission from this plasma will be calculated with the same provider. :ivar VectorFunction3D b_field: A vector function in 3D space that returns the magnetic field vector at any requested point. + :ivar VectorFunction3D e_field: A vector function in 3D space that returns the + electric field vector at any requested point. :ivar Composition composition: The composition object manages all the atomic plasma species and provides access to their distribution functions. :ivar DistributionFunction electron_distribution: A distribution function object @@ -324,6 +326,7 @@ cdef class Plasma(Node): # plasma properties self.b_field = None + self.e_field = None self.electron_distribution = None # setup plasma composition handler and pass through notifications @@ -362,6 +365,24 @@ cdef class Plasma(Node): cdef VectorFunction3D get_b_field(self): return self._b_field + @property + def e_field(self): + return self._e_field + + @e_field.setter + def e_field(self, object value): + # assign Vector3D(0, 0, 0) if None is passed + if value is None: + self._e_field = autowrap_vectorfunction3d(Vector3D(0, 0, 0)) + else: + self._e_field = autowrap_vectorfunction3d(value) + + self._modified() + + # cython fast access + cdef VectorFunction3D get_e_field(self): + return self._e_field + @property def electron_distribution(self): return self._electron_distribution diff --git a/cherab/core/tests/test_bremsstrahlung.py b/cherab/core/tests/test_bremsstrahlung.py index a77a1da5..8fd9faf5 100644 --- a/cherab/core/tests/test_bremsstrahlung.py +++ b/cherab/core/tests/test_bremsstrahlung.py @@ -24,7 +24,7 @@ from raysect.optical import World, Ray from cherab.core.atomic import AtomicData, MaxwellianFreeFreeGauntFactor -from cherab.core.math.integrators import GaussianQuadrature +from cherab.core.math.integrators import GaussianQuadrature1D from cherab.core.atomic import deuterium, nitrogen from cherab.tools.plasmas.slab import build_constant_slab_plasma from cherab.core.model import Bremsstrahlung @@ -79,7 +79,7 @@ def brems_func(wvl): return brems_const * ni_gff_z2 * ne / (np.sqrt(te) * wvl * wvl) * np.exp(- exp_factor / (te * wvl)) - integrator = GaussianQuadrature(brems_func) + integrator = GaussianQuadrature1D(brems_func) test_samples = np.zeros(brems_spectrum.bins) delta_wavelength = (brems_spectrum.max_wavelength - brems_spectrum.min_wavelength) / brems_spectrum.bins diff --git a/cherab/core/tests/test_distribution.py b/cherab/core/tests/test_distribution.py new file mode 100644 index 00000000..849821f6 --- /dev/null +++ b/cherab/core/tests/test_distribution.py @@ -0,0 +1,368 @@ +# Copyright 2016-2018 Euratom +# Copyright 2016-2018 United Kingdom Atomic Energy Authority +# Copyright 2016-2018 Centro de Investigaciones Energéticas, Medioambientales y Tecnológicas +# +# Licensed under the EUPL, Version 1.1 or – as soon they will be approved by the +# European Commission - subsequent versions of the EUPL (the "Licence"); +# You may not use this work except in compliance with the Licence. +# You may obtain a copy of the Licence at: +# +# https://joinup.ec.europa.eu/software/page/eupl5 +# +# Unless required by applicable law or agreed to in writing, software distributed +# under the Licence is distributed on an "AS IS" basis, WITHOUT WARRANTIES OR +# CONDITIONS OF ANY KIND, either express or implied. +# +# See the Licence for the specific language governing permissions and limitations +# under the Licence. + +import unittest +from itertools import product + +import numpy as np + +from raysect.core import Vector3D + +from cherab.core.distribution import ZeroDistribution, GenericDistribution, Maxwellian +from cherab.core.utility.constants import ATOMIC_MASS, ELEMENTARY_CHARGE + + +# Note: DistributionFunction is a cdef class (abstract base class) that cannot be +# directly instantiated or subclassed from Python. The abstract methods raise +# NotImplementedError, which is tested implicitly through the concrete implementations +# (ZeroDistribution, Maxwellian, GenericDistribution) that inherit from it. + + +class TestZeroDistribution(unittest.TestCase): + """ + Test cases for the ZeroDistribution class. + + ZeroDistribution should return zero for all distribution properties. + """ + + def setUp(self): + self.distribution = ZeroDistribution() + self.x = np.linspace(-10, 10, 5) # m + self.y = np.linspace(-10, 10, 5) # m + self.z = np.linspace(-10, 10, 5) # m + self.vx = np.linspace(-10e5, 10e5, 5) # m/s + self.vy = np.linspace(-10e5, 10e5, 5) # m/s + self.vz = np.linspace(-10e5, 10e5, 5) # m/s + + def tearDown(self): + pass + + def test_call_returns_zero(self): + """Test that __call__() returns zero for all inputs.""" + # iterate over a subset of inputs to avoid long execution time + for x, y, z, vx, vy, vz in product( + self.x[::2], self.y[::2], self.z[::2], self.vx[::2], self.vy[::2], self.vz[::2] + ): + result = self.distribution(x, y, z, vx, vy, vz) + self.assertEqual( + result, 0.0, msg="__call__() should return 0.0 at ({}, {}, {}, {}, {}, {}).".format(x, y, z, vx, vy, vz) + ) + + def test_bulk_velocity_returns_zero_vector(self): + """Test that bulk_velocity() returns zero vector for all positions.""" + for x, y, z in product(self.x, self.y, self.z): # iterate over all positions + velocity = self.distribution.bulk_velocity(x, y, z) + self.assertAlmostEqual( + velocity.x, 0.0, delta=1e-10, msg="bulk_velocity().x should be 0.0 at ({}, {}, {}).".format(x, y, z) + ) + self.assertAlmostEqual( + velocity.y, 0.0, delta=1e-10, msg="bulk_velocity().y should be 0.0 at ({}, {}, {}).".format(x, y, z) + ) + self.assertAlmostEqual( + velocity.z, 0.0, delta=1e-10, msg="bulk_velocity().z should be 0.0 at ({}, {}, {}).".format(x, y, z) + ) + + def test_effective_temperature_returns_zero(self): + """Test that effective_temperature() returns zero for all positions.""" + for x, y, z in product(self.x, self.y, self.z): # iterate over all positions + temperature = self.distribution.effective_temperature(x, y, z) + self.assertEqual( + temperature, 0.0, msg="effective_temperature() should return 0.0 at ({}, {}, {}).".format(x, y, z) + ) + + def test_density_returns_zero(self): + """Test that density() returns zero for all positions.""" + for x, y, z in product(self.x, self.y, self.z): # iterate over all positions + density = self.distribution.density(x, y, z) + self.assertEqual(density, 0.0, msg="density() should return 0.0 at ({}, {}, {}).".format(x, y, z)) + + +class TestGenericDistribution(unittest.TestCase): + """ + Test cases for the GenericDistribution class. + + GenericDistribution allows users to provide custom 6D phase space density, + density, temperature, and velocity functions. + """ + + def setUp(self): + self.x = np.linspace(-5, 5, 3) # m + self.y = np.linspace(-5, 5, 3) # m + self.z = np.linspace(-5, 5, 3) # m + self.vx = np.linspace(-5e5, 5e5, 3) # m/s + self.vy = np.linspace(-5e5, 5e5, 3) # m/s + self.vz = np.linspace(-5e5, 5e5, 3) # m/s + + # Define atomic mass for Gaussian distribution (using deuterium mass) + self.atomic_mass = 2 * ATOMIC_MASS # kg + + # Define shared density and temperature functions for 3D Gaussian distribution + self.density = lambda x, y, z: 1e20 * (1 + 0.1 * np.sin(x) * np.sin(y) * np.sin(z)) # m^-3 + self.temperature = lambda x, y, z: 1e3 * (1 + 0.1 * np.sin(x + 1) * np.sin(y + 1) * np.sin(z + 1)) # eV + self.velocity = lambda x, y, z: Vector3D(1e5 * x, 2e5 * y, 3e5 * z) # m/s + + # Define 3D Gaussian phase space density function using density and temperature + # This implements a Maxwellian distribution: f = n * (m/(2*pi*e*T))^(3/2) * exp(-m*v^2/(2*e*T)) + def phase_space_density_gaussian(x, y, z, vx, vy, vz): + n = self.density(x, y, z) + T = self.temperature(x, y, z) + m = self.atomic_mass + + # Thermal velocity spread squared + sigma_sq = T * ELEMENTARY_CHARGE / m # (m/s)^2 + + # Velocity magnitude squared (assuming zero bulk velocity for simplicity) + v_sq = vx**2 + vy**2 + vz**2 + + # Normalization factor + norm = (m / (2 * np.pi * ELEMENTARY_CHARGE * T)) ** 1.5 + + # Gaussian distribution + return n * norm * np.exp(-v_sq / (2 * sigma_sq)) + + self.phase_space_density = phase_space_density_gaussian + + def test_bulk_velocity(self): + """Test that bulk_velocity() returns the correct velocity vector.""" + # Define velocity function + + distribution = GenericDistribution(self.phase_space_density, self.density, self.temperature, self.velocity) + + for x, y, z in product(self.x, self.y, self.z): + result = distribution.bulk_velocity(x, y, z) + expected = self.velocity(x, y, z) + self.assertAlmostEqual( + result.x, expected.x, delta=1e-10, msg="bulk_velocity().x is wrong at ({}, {}, {}).".format(x, y, z) + ) + self.assertAlmostEqual( + result.y, expected.y, delta=1e-10, msg="bulk_velocity().y is wrong at ({}, {}, {}).".format(x, y, z) + ) + self.assertAlmostEqual( + result.z, expected.z, delta=1e-10, msg="bulk_velocity().z is wrong at ({}, {}, {}).".format(x, y, z) + ) + + def test_effective_temperature(self): + """Test that effective_temperature() returns the correct temperature.""" + + distribution = GenericDistribution(self.phase_space_density, self.density, self.temperature, self.velocity) + + for x, y, z in product(self.x, self.y, self.z): # iterate over all positions + result = distribution.effective_temperature(x, y, z) + expected = self.temperature(x, y, z) + self.assertAlmostEqual( + result, expected, delta=1e-10, msg="effective_temperature() is wrong at ({}, {}, {}).".format(x, y, z) + ) + + def test_density(self): + """Test that density() returns the correct density.""" + + distribution = GenericDistribution(self.phase_space_density, self.density, self.temperature, self.velocity) + + for x, y, z in product(self.x, self.y, self.z): # iterate over all positions + result = distribution.density(x, y, z) + expected = self.density(x, y, z) + self.assertAlmostEqual( + result, expected, delta=1e-10, msg="density() is wrong at ({}, {}, {}).".format(x, y, z) + ) + + def test_call(self): + """Test that __call__() returns the correct phase space density using 3D Gaussian.""" + + distribution = GenericDistribution(self.phase_space_density, self.density, self.temperature, self.velocity) + + # Test subset to avoid long execution time + for x, y, z, vx, vy, vz in product( + self.x[::2], self.y[::2], self.z[::2], self.vx[::2], self.vy[::2], self.vz[::2] + ): + result = distribution(x, y, z, vx, vy, vz) + expected = self.phase_space_density(x, y, z, vx, vy, vz) + self.assertAlmostEqual( + result, + expected, + delta=1e-10, + msg="__call__() is wrong at ({}, {}, {}, {}, {}, {}).".format(x, y, z, vx, vy, vz), + ) + + def test_float_inputs(self): + """Test GenericDistribution with float inputs instead of lambda functions.""" + # Pass floats directly - they should be converted to constant functions + density_float = 1e20 # m^-3 + temperature_float = 1e3 # eV + velocity_constant = Vector3D(1e5, 2e5, 3e5) # m/s (constant vector function) + phase_space_density_float = 1e17 # s^3/m^6 + + distribution = GenericDistribution( + phase_space_density_float, density_float, temperature_float, velocity_constant + ) + + # Test that all methods work with constant float inputs + for x, y, z in product(self.x, self.y, self.z): # iterate over all positions + # Density should be constant + self.assertAlmostEqual( + distribution.density(x, y, z), + density_float, + delta=1e-10, + msg="density() should return constant value at ({}, {}, {}).".format(x, y, z), + ) + # Temperature should be constant + self.assertAlmostEqual( + distribution.effective_temperature(x, y, z), + temperature_float, + delta=1e-10, + msg="effective_temperature() should return constant value at ({}, {}, {}).".format(x, y, z), + ) + # Velocity should be constant + vel = distribution.bulk_velocity(x, y, z) + self.assertAlmostEqual( + vel.x, + velocity_constant.x, + delta=1e-10, + msg="bulk_velocity().x should return constant value at ({}, {}, {}).".format(x, y, z), + ) + self.assertAlmostEqual( + vel.y, + velocity_constant.y, + delta=1e-10, + msg="bulk_velocity().y should return constant value at ({}, {}, {}).".format(x, y, z), + ) + self.assertAlmostEqual( + vel.z, + velocity_constant.z, + delta=1e-10, + msg="bulk_velocity().z should return constant value at ({}, {}, {}).".format(x, y, z), + ) + + # Test phase space density is constant + for x, y, z, vx, vy, vz in product( + self.x[::2], self.y[::2], self.z[::2], self.vx[::2], self.vy[::2], self.vz[::2] + ): + result = distribution(x, y, z, vx, vy, vz) + self.assertAlmostEqual( + result, + phase_space_density_float, + delta=1e-10, + msg="__call__() should return constant value at ({}, {}, {}, {}, {}, {}).".format(x, y, z, vx, vy, vz), + ) + + +class TestMaxwellian(unittest.TestCase): + """ + Test cases for the Maxwellian class. + + Maxwellian implements a Maxwell-Boltzmann distribution function. + """ + + def setUp(self): + self.x = np.linspace(-10, 10, 5) # m + self.y = np.linspace(-10, 10, 5) # m + self.z = np.linspace(-10, 10, 5) # m + self.vx = np.linspace(-10e5, 10e5, 5) # m/s + self.vy = np.linspace(-10e5, 10e5, 5) # m/s + self.vz = np.linspace(-10e5, 10e5, 5) # m/s + + # Define shared density, temperature, velocity, and mass for all tests + self.density = lambda x, y, z: 6e19 * (1 + 0.1 * np.sin(x) * np.sin(y) * np.sin(z)) # m^-3 + self.temperature = lambda x, y, z: 3e3 * (1 + 0.1 * np.sin(x + 1) * np.sin(y + 1) * np.sin(z + 1)) # eV + self.velocity = lambda x, y, z: ( + 1.6e5 * (1 + 0.1 * np.sin(x + 2) * np.sin(y + 2) * np.sin(z + 2)) * Vector3D(1, 2, 3).normalise() + ) # m/s + self.mass = 4 * ATOMIC_MASS # kg + + # Define sigma and phase_space_density for test_value + self.sigma = lambda x, y, z: np.sqrt(self.temperature(x, y, z) * ELEMENTARY_CHARGE / self.mass) # m/s + self.phase_space_density = lambda x, y, z, vx, vy, vz: ( + self.density(x, y, z) + / (np.sqrt(2 * np.pi) * self.sigma(x, y, z)) ** 3 + * np.exp(-((Vector3D(vx, vy, vz) - self.velocity(x, y, z)).length ** 2) / (2 * self.sigma(x, y, z) ** 2)) + ) # s^3/m^6 + + def tearDown(self): + pass + + def test_bulk_velocity(self): + """Test that bulk_velocity() returns the correct velocity vector.""" + maxwellian = Maxwellian(self.density, self.temperature, self.velocity, self.mass) + + for x, y, z in product(self.x, self.y, self.z): # iterate over all positions + velocity = maxwellian.bulk_velocity(x, y, z) + expected = self.velocity(x, y, z) + self.assertAlmostEqual( + velocity.x, + expected.x, + delta=1e-10, + msg="bulk_velocity method gives a wrong value at ({}, {}, {}).".format(x, y, z), + ) + self.assertAlmostEqual( + velocity.y, + expected.y, + delta=1e-10, + msg="bulk_velocity method gives a wrong value at ({}, {}, {}).".format(x, y, z), + ) + self.assertAlmostEqual( + velocity.z, + expected.z, + delta=1e-10, + msg="bulk_velocity method gives a wrong value at ({}, {}, {}).".format(x, y, z), + ) + + def test_effective_temperature(self): + """Test that effective_temperature() returns the correct temperature.""" + maxwellian = Maxwellian(self.density, self.temperature, self.velocity, self.mass) + + for x, y, z in product(self.x, self.y, self.z): # iterate over all positions + temperature = maxwellian.effective_temperature(x, y, z) + expected = self.temperature(x, y, z) + self.assertAlmostEqual( + temperature, + expected, + delta=1e-10, + msg="effective_temperature method gives a wrong value at ({}, {}, {}).".format(x, y, z), + ) + + def test_density(self): + """Test that density() returns the correct density.""" + maxwellian = Maxwellian(self.density, self.temperature, self.velocity, self.mass) + + for x, y, z in product(self.x, self.y, self.z): # iterate over all positions + density = maxwellian.density(x, y, z) + expected = self.density(x, y, z) + self.assertAlmostEqual( + density, + expected, + delta=1e-10, + msg="density method gives a wrong value at ({}, {}, {}).".format(x, y, z), + ) + + def test_value(self): + """Test that __call__() returns the correct phase space density.""" + maxwellian = Maxwellian(self.density, self.temperature, self.velocity, self.mass) + + # testing only half the values to avoid huge execution time + for x, y, z, vx, vy, vz in product( + self.x[::2], self.y[::2], self.z[::2], self.vx[::2], self.vy[::2], self.vz[::2] + ): + result = maxwellian(x, y, z, vx, vy, vz) + expected = self.phase_space_density(x, y, z, vx, vy, vz) + self.assertAlmostEqual( + result, + expected, + delta=1e-10, + msg="call method gives a wrong phase space density at ({}, {}, {}, {}, {}, {}).".format( + x, y, z, vx, vy, vz + ), + ) diff --git a/cherab/core/tests/test_lineshapes.py b/cherab/core/tests/test_lineshapes.py index e2122d41..5d0edbb6 100644 --- a/cherab/core/tests/test_lineshapes.py +++ b/cherab/core/tests/test_lineshapes.py @@ -27,7 +27,7 @@ from raysect.optical import Spectrum from cherab.core import Beam, Line, AtomicData -from cherab.core.math.integrators import GaussianQuadrature +from cherab.core.math.integrators import GaussianQuadrature1D from cherab.core.atomic import deuterium, nitrogen, ZeemanStructure from cherab.tools.plasmas.slab import build_constant_slab_plasma from cherab.core.model import GaussianLine, MultipletLineShape, StarkBroadenedLine, ZeemanTriplet, ParametrisedZeemanTriplet, ZeemanMultiplet @@ -297,7 +297,7 @@ def test_stark_broadened_line(self): target_species = self.plasma.composition.get(line.element, line.charge) wavelength = 656.104 relative_tolerance = 1.e-8 - integrator = GaussianQuadrature(relative_tolerance=relative_tolerance) + integrator = GaussianQuadrature1D(relative_tolerance=relative_tolerance) stark_line = StarkBroadenedLine(line, wavelength, target_species, self.plasma, self.atomic_data, integrator=integrator) # spectrum parameters diff --git a/cherab/core/tests/test_maxwellian.py b/cherab/core/tests/test_maxwellian.py deleted file mode 100644 index 67339b4b..00000000 --- a/cherab/core/tests/test_maxwellian.py +++ /dev/null @@ -1,109 +0,0 @@ -# Copyright 2016-2018 Euratom -# Copyright 2016-2018 United Kingdom Atomic Energy Authority -# Copyright 2016-2018 Centro de Investigaciones Energéticas, Medioambientales y Tecnológicas -# -# Licensed under the EUPL, Version 1.1 or – as soon they will be approved by the -# European Commission - subsequent versions of the EUPL (the "Licence"); -# You may not use this work except in compliance with the Licence. -# You may obtain a copy of the Licence at: -# -# https://joinup.ec.europa.eu/software/page/eupl5 -# -# Unless required by applicable law or agreed to in writing, software distributed -# under the Licence is distributed on an "AS IS" basis, WITHOUT WARRANTIES OR -# CONDITIONS OF ANY KIND, either express or implied. -# -# See the Licence for the specific language governing permissions and limitations -# under the Licence. - -import unittest - -import numpy as np - -from cherab.core.distribution import Maxwellian -from raysect.core import Vector3D - -ATOMIC_MASS = 1.66053906660e-27 -ELEMENTARY_CHARGE = 1.602176634e-19 - - -class TestMaxwellian(unittest.TestCase): - - def setUp(self): - self.x = np.linspace(-10, 10, 5) # m - self.y = np.linspace(-10, 10, 5) # m - self.z = np.linspace(-10, 10, 5) # m - self.vx = np.linspace(-10e5, 10e5, 5) # m/s - self.vy = np.linspace(-10e5, 10e5, 5) # m/s - self.vz = np.linspace(-10e5, 10e5, 5) # m/s - - def tearDown(self): - pass - - def test_bulk_velocity(self): - density = lambda x, y, z: 6e19 * (1 + 0.1 * np.sin(x) * np.sin(y) * np.sin(z)) # m^-3 - temperature = lambda x, y, z: 3e3 * (1 + 0.1 * np.sin(x + 1) * np.sin(y + 1) * np.sin(z + 1)) # eV - velocity = lambda x, y, z: 1.6e5 * (1 + 0.1 * np.sin(x + 2) * np.sin(y + 2) * np.sin(z + 2)) * Vector3D(1, 2, 3).normalise() # m/s - mass = 4 * ATOMIC_MASS # kg - maxwellian = Maxwellian(density, temperature, velocity, mass) - - for x in self.x: - for y in self.y: - for z in self.z: - self.assertAlmostEqual(maxwellian.bulk_velocity(x, y, z).x, velocity(x, y, z).x, delta=1e-10, - msg='bulk_velocity method gives a wrong value at ({}, {}, {}).'.format(x, y, z)) - self.assertAlmostEqual(maxwellian.bulk_velocity(x, y, z).y, velocity(x, y, z).y, delta=1e-10, - msg='bulk_velocity method gives a wrong value at ({}, {}, {}).'.format(x, y, z)) - self.assertAlmostEqual(maxwellian.bulk_velocity(x, y, z).z, velocity(x, y, z).z, delta=1e-10, - msg='bulk_velocity method gives a wrong value at ({}, {}, {}).'.format(x, y, z)) - - def test_effective_temperature(self): - density = lambda x, y, z: 6e19 * (1 + 0.1 * np.sin(x) * np.sin(y) * np.sin(z)) # m^-3 - temperature = lambda x, y, z: 3e3 * (1 + 0.1 * np.sin(x + 1) * np.sin(y + 1) * np.sin(z + 1)) # eV - velocity = lambda x, y, z: 1.6e5 * (1 + 0.1 * np.sin(x + 2) * np.sin(y + 2) * np.sin(z + 2)) * Vector3D(1, 2, 3).normalise() # m/s - mass = 4 * ATOMIC_MASS # kg - maxwellian = Maxwellian(density, temperature, velocity, mass) - - for x in self.x: - for y in self.y: - for z in self.z: - self.assertAlmostEqual(maxwellian.effective_temperature(x, y, z), temperature(x, y, z), delta=1e-10, - msg='effective_temperature method gives a wrong value at ({}, {}, {}).'.format(x, y, z)) - - def test_density(self): - density = lambda x, y, z: 6e19 * (1 + 0.1 * np.sin(x) * np.sin(y) * np.sin(z)) # m^-3 - temperature = lambda x, y, z: 3e3 * (1 + 0.1 * np.sin(x + 1) * np.sin(y + 1) * np.sin(z + 1)) # eV - velocity = lambda x, y, z: 1.6e5 * (1 + 0.1 * np.sin(x + 2) * np.sin(y + 2) * np.sin(z + 2)) * Vector3D(1, 2, 3).normalise() # m/s - mass = 4 * ATOMIC_MASS # kg - maxwellian = Maxwellian(density, temperature, velocity, mass) - - for x in self.x: - for y in self.y: - for z in self.z: - self.assertAlmostEqual(maxwellian.density(x, y, z), density(x, y, z), delta=1e-10, - msg='density method gives a wrong value at ({}, {}, {}).'.format(x, y, z)) - - def test_value(self): - density = lambda x, y, z: 6e19 * (1 + 0.1 * np.sin(x) * np.sin(y) * np.sin(z)) # m^-3 - temperature = lambda x, y, z: 3e3 * (1 + 0.1 * np.sin(x + 1) * np.sin(y + 1) * np.sin(z + 1)) # eV - velocity = lambda x, y, z: 1.6e5 * (1 + 0.1 * np.sin(x + 2) * np.sin(y + 2) * np.sin(z + 2)) * Vector3D(1, 2, 3).normalise() # m/s - mass = 4 * ATOMIC_MASS # kg - maxwellian = Maxwellian(density, temperature, velocity, mass) - - sigma = lambda x, y, z: np.sqrt(temperature(x, y, z) * ELEMENTARY_CHARGE / mass) # m/s - phase_space_density = lambda x, y, z, vx, vy, vz: density(x, y, z) / (np.sqrt(2 * np.pi) * sigma(x, y, z)) ** 3 \ - * np.exp(-(Vector3D(vx, vy, vz) - velocity(x, y, z)).length ** 2 / (2 * sigma(x, y, z) ** 2)) # s^3/m^6 - - # testing only half the values to avoid huge execution time - for x in self.x[::2]: - for y in self.y[::2]: - for z in self.z[::2]: - for vx in self.vx[::2]: - for vy in self.vy[::2]: - for vz in self.vz[::2]: - self.assertAlmostEqual(maxwellian(x, y, z, vx, vy, vz), phase_space_density(x, y, z, vx, vy, vz), delta=1e-10, - msg='call method gives a wrong phase space density at ({}, {}, {}, {}, {}, {}).'.format(x, y, z, vx, vy, vz)) - - -if __name__ == '__main__': - unittest.main() \ No newline at end of file diff --git a/cherab/core/utility/constants.pyx b/cherab/core/utility/constants.pyx index 423d15e8..18f94a7d 100644 --- a/cherab/core/utility/constants.pyx +++ b/cherab/core/utility/constants.pyx @@ -15,6 +15,11 @@ # # See the Licence for the specific language governing permissions and limitations # under the Licence. +import sys +from types import ModuleType + +from libc.math cimport M_PI + cdef: @@ -35,3 +40,47 @@ cdef: double RYDBERG_CONSTANT_EV = 13.605693122994 double VACUUM_PERMITTIVITY = 8.8541878128e-12 double BOHR_MAGNETON = 5.78838180123e-5 # in eV/T + + +# Make the constants available to Python too. +# To ensure the Python and Cython constants do not got out of sync the exported +# Python attributes of the module are made read only using module getattr. +cdef dict _CONSTANTS = { + # c stdlib + "RECIP_2_PI": RECIP_2_PI, + "RECIP_4_PI": RECIP_4_PI, + "DEGREES_TO_RADIANS": DEGREES_TO_RADIANS, + "RADIANS_TO_DEGREES": RADIANS_TO_DEGREES, + # NIST 2018 + "ATOMIC_MASS": ATOMIC_MASS, + "ELEMENTARY_CHARGE": ELEMENTARY_CHARGE, + "SPEED_OF_LIGHT": SPEED_OF_LIGHT, + "PLANCK_CONSTANT": PLANCK_CONSTANT, + "HC_EV_NM": HC_EV_NM, + "ELECTRON_CLASSICAL_RADIUS": ELECTRON_CLASSICAL_RADIUS, + "ELECTRON_REST_MASS": ELECTRON_REST_MASS, + "RYDBERG_CONSTANT_EV": RYDBERG_CONSTANT_EV, + "VACUUM_PERMITTIVITY": VACUUM_PERMITTIVITY, + "BOHR_MAGNETON": BOHR_MAGNETON, +} + + +def __getattr__(name): + if name not in _CONSTANTS: + raise AttributeError() + return _CONSTANTS[name] + + +def __dir__(): + return list(_CONSTANTS.keys()) + + +class ReadOnlyModule(ModuleType): + def __setattr__(self, attr, value): + raise AttributeError("Constants are read-only") + + def __delattr__(self, attr): + raise AttributeError("Constants are read-only") + + +sys.modules[__name__].__class__ = ReadOnlyModule diff --git a/cherab/core/utility/tests/test_constants.py b/cherab/core/utility/tests/test_constants.py new file mode 100644 index 00000000..9d46c0d2 --- /dev/null +++ b/cherab/core/utility/tests/test_constants.py @@ -0,0 +1,84 @@ +# Copyright 2026 Oak Ridge National Laboratory +# +# Licensed under the EUPL, Version 1.1 or – as soon they will be approved by the +# European Commission - subsequent versions of the EUPL (the "Licence"); +# You may not use this work except in compliance with the Licence. +# You may obtain a copy of the Licence at: +# +# https://joinup.ec.europa.eu/software/page/eupl5 +# +# Unless required by applicable law or agreed to in writing, software distributed +# under the Licence is distributed on an "AS IS" basis, WITHOUT WARRANTIES OR +# CONDITIONS OF ANY KIND, either express or implied. +# +# See the Licence for the specific language governing permissions and limitations +# under the Licence. + +import math +import unittest +from cherab.core.utility import constants + + +class TestConstants(unittest.TestCase): + def setUp(self): + self._expected_constants = dict( + # sourced c standard maths library + # CPython wraps libc's math so uses the same constants as cython's + # cimport of libc.math. + RECIP_2_PI=1 / (2 * math.pi), + RECIP_4_PI=1 / (4 * math.pi), + DEGREES_TO_RADIANS=math.pi / 180, + RADIANS_TO_DEGREES=180 / math.pi, + + # sourced from NIST, CODATA 2018=https://physics.nist.gov/cuu/Constants/Table/allascii.txt + ATOMIC_MASS=1.66053906660e-27, + ELEMENTARY_CHARGE=1.602176634e-19, + SPEED_OF_LIGHT=299792458.0, + PLANCK_CONSTANT=6.62607015e-34, + HC_EV_NM=1239.8419738620933, # (Planck constant in eV s) x (speed of light in nm/s) + ELECTRON_CLASSICAL_RADIUS=2.8179403262e-15, + ELECTRON_REST_MASS=9.1093837015e-31, + RYDBERG_CONSTANT_EV=13.605693122994, + VACUUM_PERMITTIVITY=8.8541878128e-12, + BOHR_MAGNETON=5.78838180123e-5, # in eV/T + ) + + def test_all_exported(self): + """ + Test Cython constants exported as Python floats. + """ + for name, value in self._expected_constants.items(): + self.assertEqual(value, getattr(constants, name)) + + def test_exported_literal(self): + """ + Test a constant accessed by literal name. + """ + self.assertEqual(self._expected_constants['ATOMIC_MASS'], constants.ATOMIC_MASS) + + def test_exported_names(self): + """ + Test all exported names are as expected by this test class. + """ + self.assertEqual(sorted(self._expected_constants.keys()), sorted(dir(constants))) + + def test_readonly(self): + """ + Test that constants can't be modified or removed and new constants can't be added. + """ + with self.assertRaises(AttributeError): + constants.ATOMIC_MASS = 1.66e-27 + + def test_nonew(self): + """ + Test that attempting to assign a new constant from Python errors. + """ + with self.assertRaises(AttributeError): + constants.TAU = math.tau + + def test_nodel(self): + """ + Test that attempting to delete and constant from Python errors. + """ + with self.assertRaises(AttributeError): + del constants.RYDBERG_CONSTANT_EV diff --git a/cherab/generomak/diagnostics/__init__.py b/cherab/generomak/diagnostics/__init__.py new file mode 100644 index 00000000..ee4a0fc8 --- /dev/null +++ b/cherab/generomak/diagnostics/__init__.py @@ -0,0 +1 @@ +from .bolometers import load_bolometers diff --git a/cherab/generomak/diagnostics/bolometers.py b/cherab/generomak/diagnostics/bolometers.py new file mode 100644 index 00000000..d90060fc --- /dev/null +++ b/cherab/generomak/diagnostics/bolometers.py @@ -0,0 +1,247 @@ +""" +Some foil bolometers for measuring total radiated power. + +Each individual channel consists of a BolometerFoil which receives +radiation. 4 such channels are packaged into a single bolometer "head", +similar to the bolometer hardware used in many tokamaks worldwide. +Individual bolometer cameras consist of a box with an aperture and +several bolometer heads. The overall diagnostic is made up of multiple +cameras spaced around the vessel. + +A description of the camera positions and orientations can be found in +the CAMERA_GEOMETRY dictionary within this module, which has a +separate key for each camera. This is not the only way to define the +geometry, but is convenient for computing relative transforms between +the components of the bolometer system. + +The coordinate system conventions in CAMERA_GEOMETRY are as follows. +All angles are in degrees and increase clockwise when viewing along +the relevant axes: y axis for poloidal rotation, z axis for toroidal +rotation and x axis for radial rotation. + +- rotation_poloidal: viewing angle of the slit in the poloidal plane, + with 0 being horizontally inwards. +- rotation_toroidal: viewing angle of the slit in the toroidal plane, + with 0 being purely radial. +- rotation_radial: rotation about the radial axis, 0 being vertically upwards. +- origin: position of the slit relative to the (x, z) poloidal plane i.e. y=0. +- slit_head_separation: distance between slit and each 4-channel head. +- head_angles: angle between slit normal and bolometer head normal. +- head_rotations: rotation angle about the slit-head vector, enables + reversing the order of lines of sight spatially within + each bolometer head. +- toroidal_angle: the angle of the poloidal plane in which the origin is + definied, with 0 being the (x, z) plane. + + +All of the bolometer heads and foils are identical, and are defined by +other module-level constants. +""" +from raysect.core import (Node, Point3D, Vector3D, rotate_basis, + rotate_x, rotate_y, rotate_z, translate) +from raysect.optical.material import AbsorbingSurface +from raysect.primitive import Box, Subtract + +from cherab.tools.observers import BolometerCamera, BolometerSlit, BolometerFoil + + +# Convenient constants +XAXIS = Vector3D(1, 0, 0) +YAXIS = Vector3D(0, 1, 0) +ZAXIS = Vector3D(0, 0, 1) +ORIGIN = Point3D(0, 0, 0) +# Bolometer geometry, independent of camera. The foil shapes and separation are +# inspired by the 4-channel bolometer head currently used by many tokamaks. +BOX_WIDTH = 0.1 +BOX_HEIGHT = 0.07 +BOX_DEPTH = 0.2 +THICKNESS = 1e-3 +SLIT_WIDTH = 0.004 +SLIT_HEIGHT = 0.005 +FOIL_WIDTH = 0.0013 +FOIL_HEIGHT = 0.0038 +FOIL_CORNER_CURVATURE = 0.0005 +FOIL_SEPARATION = 0.00508 # 0.2 inch between foils + +CAMERA_GEOMETRY = { + 'HozPol1': {}, # Horizontal poloidal + 'HozPol2': {}, # Horizontal poloidal, + 'VertPol': {}, # Vertical poloidal + 'TanMid1': {}, # Tangential + 'TanPol1': {} # Combined poloidal/tangential +} + +# The camera geometry definitions are grouped by property here, to illustrate +# the relationship between the different cameras. The geometry can be viewed +# grouped by camera instead as follows: +# >>> from cherab.generomak.diagnostics.bolometers import CAMERA_GEOMETRY +# >>> from pprint import pprint +# >>> pprint(CAMERA_GEOMETRY) + +# poloidal rotations +CAMERA_GEOMETRY['HozPol1']['rotation_poloidal'] = 30 +CAMERA_GEOMETRY['HozPol2']['rotation_poloidal'] = -30 +CAMERA_GEOMETRY['VertPol']['rotation_poloidal'] = -90 +CAMERA_GEOMETRY['TanMid1']['rotation_poloidal'] = 0 +CAMERA_GEOMETRY['TanPol1']['rotation_poloidal'] = -25 +# toroidal rotation +CAMERA_GEOMETRY['HozPol1']['rotation_toroidal'] = 0 +CAMERA_GEOMETRY['HozPol2']['rotation_toroidal'] = 0 +CAMERA_GEOMETRY['VertPol']['rotation_toroidal'] = 0 +CAMERA_GEOMETRY['TanMid1']['rotation_toroidal'] = -40 +CAMERA_GEOMETRY['TanPol1']['rotation_toroidal'] = 40 +# radial rotation +CAMERA_GEOMETRY['HozPol1']['rotation_radial'] = -90 +CAMERA_GEOMETRY['HozPol2']['rotation_radial'] = -90 +CAMERA_GEOMETRY['VertPol']['rotation_radial'] = -90 +CAMERA_GEOMETRY['TanMid1']['rotation_radial'] = 0 +CAMERA_GEOMETRY['TanPol1']['rotation_radial'] = 0 +# origins relative to the poloidal (x, z) plane +CAMERA_GEOMETRY['HozPol1']['origin'] = Point3D(2.45, 0.05, 0) +CAMERA_GEOMETRY['HozPol2']['origin'] = Point3D(2.45, -0.05, 0) +CAMERA_GEOMETRY['VertPol']['origin'] = Point3D(1.3, 0, 1.42) +CAMERA_GEOMETRY['TanMid1']['origin'] = Point3D(2.5, 0, 0) +CAMERA_GEOMETRY['TanPol1']['origin'] = Point3D(2.2, 0, -0.8) +# slit-head separations +CAMERA_GEOMETRY['HozPol1']['slit_head_separation'] = 0.08 +CAMERA_GEOMETRY['HozPol2']['slit_head_separation'] = 0.08 +CAMERA_GEOMETRY['VertPol']['slit_head_separation'] = 0.05 +CAMERA_GEOMETRY['TanMid1']['slit_head_separation'] = 0.1 +CAMERA_GEOMETRY['TanPol1']['slit_head_separation'] = 0.15 +# bolometer head angles relative to the slit +CAMERA_GEOMETRY['HozPol1']['head_angles'] = [22.5, 7.5, -7.5, -22.5] +CAMERA_GEOMETRY['HozPol2']['head_angles'] = [22.5, 7.5, -7.5, -22.5] +CAMERA_GEOMETRY['VertPol']['head_angles'] = [36, 12, -12, -36] +CAMERA_GEOMETRY['TanMid1']['head_angles'] = [18, 6, -6, -18] +CAMERA_GEOMETRY['TanPol1']['head_angles'] = [-12, -4, 4, 12] +# bolometer head rotation relative to the slit +CAMERA_GEOMETRY['HozPol1']['head_rotations'] = [0, 0, 0, 0] +CAMERA_GEOMETRY['HozPol2']['head_rotations'] = [0, 0, 0, 0] +CAMERA_GEOMETRY['VertPol']['head_rotations'] = [0, 0, 0, 0] +CAMERA_GEOMETRY['TanMid1']['head_rotations'] = [0, 0, 0, 0] +CAMERA_GEOMETRY['TanPol1']['head_rotations'] = [180, 180, 180, 180] +# toroidal angles about which to rotate the poloidal plane +CAMERA_GEOMETRY['HozPol1']['toroidal_angle'] = 10 # need to avoid LFS limiters +CAMERA_GEOMETRY['HozPol2']['toroidal_angle'] = 10 # need to avoid LFS limiters +CAMERA_GEOMETRY['VertPol']['toroidal_angle'] = 0 # happy to hit LFS limiters +CAMERA_GEOMETRY['TanMid1']['toroidal_angle'] = -15 # avoid LFS limiters +CAMERA_GEOMETRY['TanPol1']['toroidal_angle'] = 15 # avoid LFS limiters + + +def _make_bolometer_camera(slit_head_separation, head_angles, head_rotations): + """ + Build a single bolometer camera. + + The camera consists of a box with a rectangular slit and 4 + bolometer heads, each of which has 4 foils. + + In its local coordinate system, the camera's slit is located at + the origin with its width along the X axis and its height along + the y axis, and the bolometer heads are below the z=0 plane + looking up towards the slit. + + The bolometer heads are rotated by head_angles about the y axis to + form a fan, and by head_rotations about the axis defined by the + line between the slit and the head. A rotation of 180 degrees + flips the head upside down and therefore reverses the spatial + ordering of lines of sight relative to a rotation of 0 degrees. + """ + camera_box = Box(lower=Point3D(-BOX_WIDTH / 2, -BOX_HEIGHT / 2, -BOX_DEPTH), + upper=Point3D(BOX_WIDTH / 2, BOX_HEIGHT / 2, 0)) + # Hollow out the box: it has 1 mm thick walls. + inside_box = Box(lower=camera_box.lower + Vector3D(THICKNESS, THICKNESS, THICKNESS), + upper=camera_box.upper - Vector3D(THICKNESS, THICKNESS, THICKNESS)) + camera_box = Subtract(camera_box, inside_box) + # The slit is a hole in the box. Make it thicker than the wall. + aperture = Box(lower=Point3D(-SLIT_WIDTH / 2, -SLIT_HEIGHT / 2, -1.1 * THICKNESS), + upper=Point3D(SLIT_WIDTH / 2, SLIT_HEIGHT / 2, 0.1 * THICKNESS)) + camera_box = Subtract(camera_box, aperture) + camera_box.material = AbsorbingSurface() + bolometer_camera = BolometerCamera(camera_geometry=camera_box) + # The bolometer slit in this instance just contains targeting information + # for the ray tracing, since we have already given our camera a geometry + # The slit is defined in the local coordinate system of the camera + slit = BolometerSlit(slit_id="Example slit", centre_point=ORIGIN, + basis_x=XAXIS, dx=SLIT_WIDTH, basis_y=YAXIS, dy=SLIT_HEIGHT, + parent=bolometer_camera) + for j, (angle, rotation) in enumerate(zip(head_angles, head_rotations)): + # 4 bolometer foils, spaced at equal intervals along the local X axis + head = Node(name="Bolometer head", parent=bolometer_camera) + head.transform = ( + rotate_y(angle) + * rotate_z(rotation) + * translate(0, 0, -slit_head_separation) + ) + for i, shift in enumerate([-1.5, -0.5, 0.5, 1.5]): + # Note that the foils will be parented to the camera rather than the bolometer + # head, so we need to define their transform relative to the camera. + foil_transform = head.transform * translate(shift * FOIL_SEPARATION, 0, 0) + foil = BolometerFoil(detector_id="Foil {} head {}".format(i + 1, j + 1), + centre_point=ORIGIN.transform(foil_transform), + basis_x=XAXIS.transform(foil_transform), dx=FOIL_WIDTH, + basis_y=YAXIS.transform(foil_transform), dy=FOIL_HEIGHT, + slit=slit, parent=bolometer_camera, units="Power", + accumulate=False, curvature_radius=FOIL_CORNER_CURVATURE) + bolometer_camera.add_foil_detector(foil) + return bolometer_camera + + +def load_bolometers(parent=None): + """ + Load the Generomak bolometers. + + The Generomak bolometer diagnostic consists of multiple 16-channel + cameras. Each camera has 4 4-channel bolometer heads inside. + + * 2 cameras are located at the midplane with purely-poloidal, + horizontal views. + * 1 camera is located at the top of the machine with purely-poloidal, + vertical views. + * 2 cameras have purely tangential views at the midplane. + * 1 camera has combined poloidal+tangential views, which look like + curved lines of sight in the poloidal plane. It looks at the lower + divertor. + + Channel ordering is as follows: + * Poloidal channels are ordered anti-clockwise by line-of-sight: + channel 1 of HozPol1 views the top of the machine and channel 16 + HozPol2 views the bottom of the machine. Similarly, channel 1 of + VertPol views the high field side and channel 16 views the low + field side. + * Tangential channels are ordered by increasing tangency radius: + channel 1 of TanMid1 has its tangency radius on the high field + side and channel 16 has its tangency radius on the low field side. + * The combined tangential/poloidal channels follow both conventions: + channel 1 views the high field side and channel 16 views the low + field side. + + :param parent: the scenegraph node the bolometers will belong to. + :return: a list of BolometerCamera instances, one for each of the + cameras described above. + """ + cameras = [] + for name, prop in CAMERA_GEOMETRY.items(): + camera = _make_bolometer_camera( + prop['slit_head_separation'], + prop['head_angles'], + prop['head_rotations'], + ) + # The transform is applied as follows: + # 1. Point the camera along the inward radial direction in the (x, z) plane. + # 2. Make the radial, poloidal and toroidal rotations while the camera is at + # the origin. + # 3. Move the camera to its position relative to the (x, z) plane. + # 4. Rotate the (x, z) plane to the correct toroidal angle. + # Transforms are applied right-to-left (or bottom-to-top with one per line): + camera.transform = ( + rotate_z(prop['toroidal_angle']) + * translate(prop['origin'].x, prop['origin'].y, prop['origin'].z) + * rotate_z(prop['rotation_toroidal']) + * rotate_y(prop['rotation_poloidal']) + * rotate_x(prop['rotation_radial']) + * rotate_basis(-XAXIS, ZAXIS) + ) + camera.parent = parent + camera.name = name + cameras.append(camera) + return cameras diff --git a/cherab/openadas/parse/adf15.py b/cherab/openadas/parse/adf15.py index 12aa01a9..0b05f42f 100644 --- a/cherab/openadas/parse/adf15.py +++ b/cherab/openadas/parse/adf15.py @@ -17,27 +17,52 @@ # under the Licence. import re + import numpy as np -from cherab.core.atomic import hydrogen, Element + +from cherab.core.atomic import Element, hydrogen from cherab.core.utility import RecursiveDict from cherab.core.utility.conversion import Cm3ToM3, PerCm3ToPerM3 +# Compiled regex patterns for ADF15 file parsing +_ADF_HEADER_MATCH = re.compile(r"^\s*(\d*) {4}/(.*)/?\s*$") +_PEC_INDEX_HEADER_MATCH_STANDARD = re.compile(r"^C\s*ISEL\s*(?:WAVELENGTH|WVLEN\(A\))\s*TRANSITION\s*TYPE", re.IGNORECASE) +_PEC_HYDROGEN_TRANSITION_MATCH = re.compile(r"^C\s*([0-9]*)\.\s*([0-9]*\.[0-9]*)\s*N=\s*([0-9]*) - N=\s*([0-9]*)\s*([A-Z]*)", re.IGNORECASE) +_PEC_FULL_TRANSITION_MATCH = re.compile(r"^[cC]\s*([0-9]*)\.?\s*([0-9]*\.[0-9]*)\s*([0-9]*)[\(\)\.0-9\s]*-\s*([0-9]*)[\(\)\.0-9\s]*([A-Z]*)", re.IGNORECASE) +_CONFIGURATION_HEADER_MATCH = re.compile(r"^C\s*(?:lv\s+)?Configuration\s*\(2S\+1\)L\(w-1/2\)\s*Energy\s*\(cm(?:\*\*|\^)-1\)\s*$", re.IGNORECASE) +_CONFIGURATION_STRING_MATCH = re.compile( + r"^[cC]\s*([0-9]+)\s*" + r"((?:[0-9][SPDFG][0-9](?:\s+[0-9][SPDFG][0-9])*)|(?:[0-9A-Z]+))\s*" + r"\(([0-9]*\.?[0-9]+)\)" + r"\s*([0-9]+)" + r"\(\s*([0-9]*\.?[0-9]+)\)", + re.IGNORECASE, +) +_WAVELENGTH_MATCH = re.compile(r"^\s*[0-9]*\.[0-9]* ?a?\s+[0-9]+\s+[0-9]+.*?/isel *= *[0-9]+$", re.IGNORECASE) +_BLOCK_ID_MATCH = re.compile(r"^\s*[0-9]*\.[0-9]* ?a?\s*([0-9]*)\s*([0-9]*).*/type *= *([a-zA-Z]*).*/isel *= * ([0-9]*)$", re.IGNORECASE) _L_LOOKUP = { - 0: 'S', - 1: 'P', - 2: 'D', - 3: 'F', - 4: 'G', - 5: 'H', - 6: 'I', - 7: 'K', - 8: 'L', - 9: 'M', - 10: 'N', - 11: 'O', - 12: 'Q', - 13: 'R', + 0: "S", + 1: "P", + 2: "D", + 3: "F", + 4: "G", + 5: "H", + 6: "I", + 7: "K", + 8: "L", + 9: "M", + 10: "N", + 11: "O", + 12: "Q", + 13: "R", + 14: "T", + 15: "U", + 16: "V", + 17: "W", + 18: "X", + 19: "Y", + 20: "Z", } @@ -52,26 +77,25 @@ def parse_adf15(element, charge, adf_file_path, header_format=None): """ if not isinstance(element, Element): - raise TypeError('The element must be an Element object.') + raise TypeError("The element must be an Element object.") charge = int(charge) with open(adf_file_path, "r") as file: - # for check header line header = file.readline() - if not re.match(r'^\s*(\d*) {4}/(.*)/?\s*$', header): - raise ValueError('The specified path does not point to a valid ADF15 file.') + if not _ADF_HEADER_MATCH.match(header): + raise ValueError("The specified path does not point to a valid ADF15 file.") # scrape transition information and wavelength # use simple electron configuration structure for hydrogen-like ions - if header_format == 'hydrogen' or element == hydrogen: + if header_format == "hydrogen" or element == hydrogen: config = _scrape_metadata_hydrogen(file, element, charge) - elif header_format == 'hydrogen-like': + elif header_format == "hydrogen-like": config = _scrape_metadata_hydrogen_like(file, element, charge) elif element.atomic_number - charge == 1: config = _scrape_metadata_hydrogen_like(file, element, charge) - if not config and 'bnd#' in adf_file_path: + if not config and "bnd#" in adf_file_path: # ADF15 files with the "bnd" suffix may have metadata in the "hydrogen" format config = _scrape_metadata_hydrogen(file, element, charge) else: @@ -82,14 +106,14 @@ def parse_adf15(element, charge, adf_file_path, header_format=None): # process rate data rates = RecursiveDict() - for cls in ('excitation', 'recombination', 'thermalcx'): + for cls in ("excitation", "recombination", "thermalcx"): for element, charge_states in config[cls].items(): for charge, transitions in charge_states.items(): for transition in transitions.keys(): block_num = config[cls][element][charge][transition] rates[cls][element][charge][transition] = _extract_rate(file, block_num) - wavelengths = config['wavelength'] + wavelengths = config["wavelength"] return rates, wavelengths @@ -104,15 +128,11 @@ def _scrape_metadata_hydrogen(file, element, charge): file.seek(0) lines = file.readlines() - pec_index_header_match = r'^C\s*ISEL\s*WAVELENGTH\s*TRANSITION\s*TYPE' - while not re.match(pec_index_header_match, lines[0], re.IGNORECASE): + while not _PEC_INDEX_HEADER_MATCH_STANDARD.match(lines[0]): lines.pop(0) index_lines = lines - for i in range(len(index_lines)): - - pec_hydrogen_transition_match = r'^C\s*([0-9]*)\.\s*([0-9]*\.[0-9]*)\s*N=\s*([0-9]*) - N=\s*([0-9]*)\s*([A-Z]*)' - match = re.match(pec_hydrogen_transition_match, index_lines[i], re.IGNORECASE) + match = _PEC_HYDROGEN_TRANSITION_MATCH.match(index_lines[i]) if not match: continue @@ -120,13 +140,13 @@ def _scrape_metadata_hydrogen(file, element, charge): wavelength = float(match.groups()[1]) / 10 # convert Angstroms to nm upper_level = int(match.groups()[2]) lower_level = int(match.groups()[3]) - rate_type_adas = match.groups()[4] - if rate_type_adas == 'EXCIT': - rate_type = 'excitation' - elif rate_type_adas == 'RECOM': - rate_type = 'recombination' - elif rate_type_adas == 'CHEXC': - rate_type = 'thermalcx' + rate_type_adas = match.groups()[4].upper() + if rate_type_adas == "EXCIT": + rate_type = "excitation" + elif rate_type_adas == "RECOM": + rate_type = "recombination" + elif rate_type_adas == "CHEXC": + rate_type = "thermalcx" else: raise ValueError("Unrecognised rate type - {}".format(rate_type_adas)) @@ -147,15 +167,11 @@ def _scrape_metadata_hydrogen_like(file, element, charge): file.seek(0) lines = file.readlines() - pec_index_header_match = r'^C\s*ISEL\s*WAVELENGTH\s*TRANSITION\s*TYPE' - while not re.match(pec_index_header_match, lines[0], re.IGNORECASE): + while not _PEC_INDEX_HEADER_MATCH_STANDARD.match(lines[0]): lines.pop(0) index_lines = lines - for i in range(len(index_lines)): - - pec_full_transition_match = r'^C\s*([0-9]*)\.\s*([0-9]*\.[0-9]*)\s*([0-9]*)[\(\)\.0-9\s]*-\s*([0-9]*)[\(\)\.0-9\s]*([A-Z]*)' - match = re.match(pec_full_transition_match, index_lines[i], re.IGNORECASE) + match = _PEC_FULL_TRANSITION_MATCH.match(index_lines[i]) if not match: continue @@ -163,13 +179,13 @@ def _scrape_metadata_hydrogen_like(file, element, charge): wavelength = float(match.groups()[1]) / 10 # convert Angstroms to nm upper_level = int(match.groups()[2]) lower_level = int(match.groups()[3]) - rate_type_adas = match.groups()[4] - if rate_type_adas == 'EXCIT': - rate_type = 'excitation' - elif rate_type_adas == 'RECOM': - rate_type = 'recombination' - elif rate_type_adas == 'CHEXC': - rate_type = 'thermalcx' + rate_type_adas = match.groups()[4].upper() + if rate_type_adas == "EXCIT": + rate_type = "excitation" + elif rate_type_adas == "RECOM": + rate_type = "recombination" + elif rate_type_adas == "CHEXC": + rate_type = "thermalcx" else: raise ValueError("Unrecognised rate type - {}".format(rate_type_adas)) @@ -193,19 +209,14 @@ def _scrape_metadata_full(file, element, charge): configuration_lines = [] configuration_dict = {} - configuration_header_match = r'^C\s*Configuration\s*\(2S\+1\)L\(w-1/2\)\s*Energy \(cm\*\*-1\)$' - while not re.match(configuration_header_match, lines[0], re.IGNORECASE): + while not _CONFIGURATION_HEADER_MATCH.match(lines[0]): lines.pop(0) - pec_index_header_match = r'^C\s*ISEL\s*WAVELENGTH\s*TRANSITION\s*TYPE' - while not re.match(pec_index_header_match, lines[0], re.IGNORECASE): + while not _PEC_INDEX_HEADER_MATCH_STANDARD.match(lines[0]): configuration_lines.append(lines[0]) lines.pop(0) index_lines = lines - for i in range(len(configuration_lines)): - - configuration_string_match = r"^C\s*([0-9]*)\s*((?:[0-9][SPDFG][0-9]\s)*)\s*\(([0-9]*\.?[0-9]*)\)([0-9]*)\(\s*([0-9]*\.?[0-9]*)\)" - match = re.match(configuration_string_match, configuration_lines[i], re.IGNORECASE) + match = _CONFIGURATION_STRING_MATCH.match(configuration_lines[i]) if not match: continue @@ -215,13 +226,10 @@ def _scrape_metadata_full(file, element, charge): total_orbital_quantum_number = _L_LOOKUP[int(match.groups()[3])] # L total_angular_momentum_quantum_number = match.groups()[4] # J - configuration_dict[config_id] = (electron_configuration + " " + spin_multiplicity + - total_orbital_quantum_number + total_angular_momentum_quantum_number) + configuration_dict[config_id] = electron_configuration + " " + spin_multiplicity + total_orbital_quantum_number + total_angular_momentum_quantum_number for i in range(len(index_lines)): - - pec_full_transition_match = r'^C\s*([0-9]*)\.?\s*([0-9]*\.[0-9]*)\s*([0-9]*)[\(\)\.0-9\s]*-\s*([0-9]*)[\(\)\.0-9\s]*([A-Z]*)' - match = re.match(pec_full_transition_match, index_lines[i], re.IGNORECASE) + match = _PEC_FULL_TRANSITION_MATCH.match(index_lines[i]) if not match: continue @@ -231,13 +239,13 @@ def _scrape_metadata_full(file, element, charge): upper_level = configuration_dict[upper_level_id] lower_level_id = int(match.groups()[3]) lower_level = configuration_dict[lower_level_id] - rate_type_adas = match.groups()[4] - if rate_type_adas == 'EXCIT': - rate_type = 'excitation' - elif rate_type_adas == 'RECOM': - rate_type = 'recombination' - elif rate_type_adas == 'CHEXC': - rate_type = 'thermalcx' + rate_type_adas = match.groups()[4].upper() + if rate_type_adas == "EXCIT": + rate_type = "excitation" + elif rate_type_adas == "RECOM": + rate_type = "recombination" + elif rate_type_adas == "CHEXC": + rate_type = "thermalcx" else: raise ValueError("Unrecognised rate type - {}".format(rate_type_adas)) @@ -255,11 +263,8 @@ def _extract_rate(file, block_num): # search from start of file file.seek(0) - wavelength_match = r"^\s*[0-9]*\.[0-9]* ?a? +.*$" - block_id_match = r"^\s*[0-9]*\.[0-9]* ?a?\s*([0-9]*)\s*([0-9]*).*/type *= *([a-zA-Z]*).*/isel *= * ([0-9]*)$" - - for block in _group_by_block(file, wavelength_match): - match = re.match(block_id_match, block[0], re.IGNORECASE) + for block in _group_by_block(file, _WAVELENGTH_MATCH): + match = _BLOCK_ID_MATCH.match(block[0]) if not match: continue @@ -311,24 +316,24 @@ def _extract_rate(file, block_num): density = PerCm3ToPerM3.to(density) rates = Cm3ToM3.to(rates) - return {'ne': density, 'te': temperature, 'rate': rates} + return {"ne": density, "te": temperature, "rate": rates} # If code gets to here, block wasn't found. - raise RuntimeError('Block number {} was not found in the ADF15 file.'.format(block_num)) + raise RuntimeError("Block number {} was not found in the ADF15 file.".format(block_num)) -def _group_by_block(source_file, match_string): +def _group_by_block(source_file, match_pattern): """ Generator the splits the ADF15 file into blocks. - Groups lines of file into blocks based on precursor ' 6561.9A 24...' + Groups lines of file into blocks based on wavelength pattern match. Note: comment section not filtered out of last block, don't over-read! """ buffer = [] for line in source_file: - if re.match(match_string, line, re.IGNORECASE): + if match_pattern.match(line): if buffer: yield buffer buffer = [line] diff --git a/cherab/openadas/tests/test_adf15.py b/cherab/openadas/tests/test_adf15.py new file mode 100644 index 00000000..ad333d63 --- /dev/null +++ b/cherab/openadas/tests/test_adf15.py @@ -0,0 +1,296 @@ +# Copyright 2016-2023 Euratom +# Copyright 2016-2023 United Kingdom Atomic Energy Authority +# Copyright 2016-2023 Centro de Investigaciones Energéticas, Medioambientales y Tecnológicas +# +# Licensed under the EUPL, Version 1.1 or – as soon they will be approved by the +# European Commission - subsequent versions of the EUPL (the "Licence"); +# You may not use this work except in compliance with the Licence. +# You may obtain a copy of the Licence at: +# +# https://joinup.ec.europa.eu/software/page/eupl5 +# +# Unless required by applicable law or agreed to in writing, software distributed +# under the Licence is distributed on an "AS IS" basis, WITHOUT WARRANTIES OR +# CONDITIONS OF ANY KIND, either express or implied. +# +# See the Licence for the specific language governing permissions and limitations +# under the Licence. + +import os +import tempfile +import unittest + +import numpy as np + +from cherab.core.atomic import carbon, hydrogen +from cherab.openadas.parse.adf15 import parse_adf15 + + +class MockADF15Files: + """Helper class to create mock ADF15 files for testing.""" + + @staticmethod + def create_hydrogen_adf15(): + """Create a mock ADF15 file in hydrogen format.""" + content = """ 0 /test hydrogen format/ +C +C TEST FILE FOR HYDROGEN +C +C ISEL WAVELENGTH TRANSITION TYPE +C +C 1. 656.3 N= 2 - N= 1 EXCIT +C 2. 486.1 N= 3 - N= 2 RECOM +C 3. 434.0 N= 4 - N= 2 CHEXC +C +C PHOTON EMISSIVITY COEFFICIENTS +C +656.3A 2 2 /type = excit /isel = 1 +1.0E+08 1.0E+09 +1000.0 5000.0 +1.0E-12 2.0E-12 3.0E-12 4.0E-12 +486.1A 2 2 /type = recom /isel = 2 +1.1E+08 1.1E+09 +2000.0 6000.0 +1.1E-12 2.1E-12 3.1E-12 4.1E-12 +434.0A 2 2 /type = chexc /isel = 3 +1.2E+08 1.2E+09 +3000.0 7000.0 +1.2E-12 2.2E-12 3.2E-12 4.2E-12 +""" + return content + + @staticmethod + def create_carbon_adf15(): + """Create a mock ADF15 file in carbon (full configuration) format.""" + content = """ 0 /test carbon format/ +C +C TEST FILE FOR CARBON +C +C lv Configuration (2S+1)L(w-1/2) Energy (cm^-1) +C --- ------------------- -------------- -------------- +c 1 60964A52B (1) 0( 0.0) 0.0 +c 2 60963A52B51C (3) 2( 3.0) 16720.4 +c 3 60964A51C51D (5) 4( 4.5) 32456.8 +c 4 60963A52B51D (2) 1( 1.5) 48932.1 +C +C ISEL WVLEN(A) TRANSITION TYPE ISPB NSPB +C ISPP NSPP SZ TG PR WR +C ----- ---------- ----------------------------------- ----- ---- ---- -- -- -- -- +C 1 1560.70 1(3)1( 1.0)- 2(2)2( 2.0) excit 1 1 12 569 12 1 +C 2 1657.80 1(3)1( 1.0)- 3(5)4( 4.5) excit 1 1 12 1068 7 2 +C 3 1329.50 1(3)1( 1.0)- 4(2)1( 1.5) excit 1 1 12 778 31 3 +C +C PHOTON EMISSIVITY COEFFICIENTS +C +1560.70A 2 2 /type = excit /isel = 1 +1.0E+08 1.0E+09 +1000.0 5000.0 +1.0E-12 2.0E-12 3.0E-12 4.0E-12 +1657.80A 2 2 /type = excit /isel = 2 +1.1E+08 1.1E+09 +2000.0 6000.0 +1.1E-12 2.1E-12 3.1E-12 4.1E-12 +1329.50A 2 2 /type = excit /isel = 3 +1.2E+08 1.2E+09 +3000.0 7000.0 +1.2E-12 2.2E-12 3.2E-12 4.2E-12 +""" + return content + + @staticmethod + def create_tungsten_adf15(): + """Create a mock ADF15 file in tungsten (extended) format.""" + content = """ 0 /test tungsten format/ +C +C TEST FILE FOR TUNGSTEN +C +C lv Configuration (2S+1)L(w-1/2) Energy (cm^-1) +C --- ------------------- -------------- -------------- +c 1074 60964A52B (1) 0( 0.0) 0.0 +c 1075 60963A52B51C (3) 2( 3.0) 16720.4 +c 414 60964A51C51D (5) 4( 4.5) 32456.8 +c 365 60963A52B51D (2) 1( 1.5) 48932.1 +C +C ISEL WVLEN(A) TRANSITION TYPE ISPB NSPB +C ISPP NSPP SZ TG PR WR +C ----- ---------- ----------------------------------- ----- ---- ---- -- -- -- -- +C 1 56.5300 1074(1)0( 0.0)- 1075(3)2( 3.0) excit 1 1 12 569 12 1 +C 2 56.5520 1075(3)2( 3.0)- 414(5)4( 4.5) excit 1 1 12 1068 7 2 +C 3 70.1877 414(5)4( 4.5)- 365(2)1( 1.5) excit 1 1 12 778 31 3 +C 4 70.5640 365(2)1( 1.5)- 1074(1)0( 0.0) excit 1 1 12 275 36 4 +C +C PHOTON EMISSIVITY COEFFICIENTS +C +56.5300A 2 2 /type = excit /isel = 1 +1.0E+08 1.0E+09 +1000.0 5000.0 +1.0E-12 2.0E-12 3.0E-12 4.0E-12 +56.5520A 2 2 /type = excit /isel = 2 +1.1E+08 1.1E+09 +2000.0 6000.0 +1.1E-12 2.1E-12 3.1E-12 4.1E-12 +70.1877A 2 2 /type = excit /isel = 3 +1.2E+08 1.2E+09 +3000.0 7000.0 +1.2E-12 2.2E-12 3.2E-12 4.2E-12 +70.5640A 2 2 /type = excit /isel = 4 +1.3E+08 1.3E+09 +4000.0 8000.0 +1.3E-12 2.3E-12 3.3E-12 4.3E-12 +""" + return content + + +class TestADF15Parser(unittest.TestCase): + """Unit tests for ADF15 parser.""" + + def setUp(self): + """Set up test fixtures.""" + self.temp_dir = tempfile.mkdtemp() + + def tearDown(self): + """Clean up temporary files.""" + for filename in os.listdir(self.temp_dir): + filepath = os.path.join(self.temp_dir, filename) + if os.path.isfile(filepath): + os.unlink(filepath) + os.rmdir(self.temp_dir) + + def _create_test_file(self, filename, content): + """Helper to create a test file.""" + filepath = os.path.join(self.temp_dir, filename) + with open(filepath, "w") as f: + f.write(content) + return filepath + + def test_parse_hydrogen_adf15(self): + """Test parsing of hydrogen format ADF15 file.""" + content = MockADF15Files.create_hydrogen_adf15() + filepath = self._create_test_file("test_h.adf15", content) + + rates, wavelengths = parse_adf15(hydrogen, 0, filepath, header_format="hydrogen") + + # Check that rates were extracted + self.assertIn("excitation", rates) + self.assertIn(hydrogen, rates["excitation"]) + self.assertIn(0, rates["excitation"][hydrogen]) + + # Check that wavelengths were extracted + self.assertIn(hydrogen, wavelengths) + self.assertIn(0, wavelengths[hydrogen]) + + # Check specific transitions + self.assertIn((2, 1), rates["excitation"][hydrogen][0]) + self.assertIn((3, 2), rates["recombination"][hydrogen][0]) + self.assertIn((4, 2), rates["thermalcx"][hydrogen][0]) + + def test_parse_carbon_adf15_full_config(self): + """Test parsing of carbon format ADF15 file with full configuration.""" + content = MockADF15Files.create_carbon_adf15() + filepath = self._create_test_file("test_c.adf15", content) + + rates, wavelengths = parse_adf15(carbon, 0, filepath) + + # Check that rates were extracted + self.assertIn("excitation", rates) + self.assertIn(carbon, rates["excitation"]) + self.assertIn(0, rates["excitation"][carbon]) + + # Check that wavelengths were extracted + self.assertIn(carbon, wavelengths) + + # Verify rate data has correct shape + transitions = list(rates["excitation"][carbon][0].keys()) + self.assertGreater(len(transitions), 0) + + for transition, rate_data in rates["excitation"][carbon][0].items(): + self.assertIn("ne", rate_data) + self.assertIn("te", rate_data) + self.assertIn("rate", rate_data) + self.assertTrue(isinstance(rate_data["ne"], np.ndarray)) + self.assertTrue(isinstance(rate_data["te"], np.ndarray)) + self.assertTrue(isinstance(rate_data["rate"], np.ndarray)) + + def test_parse_tungsten_adf15(self): + """Test parsing of tungsten format ADF15 file.""" + # Tungsten (W) has atomic number 74 + # Create a mock tungsten element for testing + from cherab.core.atomic import tungsten + + content = MockADF15Files.create_tungsten_adf15() + filepath = self._create_test_file("test_w.adf15", content) + + rates, wavelengths = parse_adf15(tungsten, 0, filepath) + + # Check that rates were extracted + self.assertIn("excitation", rates) + self.assertIn(tungsten, rates["excitation"]) + self.assertIn(0, rates["excitation"][tungsten]) + + # Check that wavelengths were extracted + self.assertIn(tungsten, wavelengths) + + def test_rate_data_structure(self): + """Test that rate data has correct structure and units.""" + content = MockADF15Files.create_hydrogen_adf15() + filepath = self._create_test_file("test_structure.adf15", content) + + rates, wavelengths = parse_adf15(hydrogen, 0, filepath, header_format="hydrogen") + + # Extract first rate + first_transition = list(rates["excitation"][hydrogen][0].keys())[0] + rate_data = rates["excitation"][hydrogen][0][first_transition] + + # Check structure + self.assertEqual(set(rate_data.keys()), {"ne", "te", "rate"}) + + # Check that units were converted (values should be large after conversion from cm^-3 to m^-3) + self.assertTrue(np.all(rate_data["ne"] >= 1e14)) # Should be in m^-3 + self.assertTrue(np.all(rate_data["te"] > 0)) # Temperature should be positive + self.assertTrue(np.all(rate_data["rate"] > 0)) # Rate should be positive + + # Check array dimensions match + ne_count = len(rate_data["ne"]) + te_count = len(rate_data["te"]) + rate_shape = rate_data["rate"].shape + self.assertEqual(rate_shape, (ne_count, te_count)) + + def test_invalid_adf15_file(self): + """Test that invalid ADF15 file raises appropriate error.""" + invalid_content = "This is not a valid ADF15 file\n" + filepath = self._create_test_file("invalid.adf15", invalid_content) + + with self.assertRaises(ValueError) as context: + parse_adf15(hydrogen, 0, filepath) + + self.assertIn("valid ADF15 file", str(context.exception)) + + def test_wavelength_conversion(self): + """Test that wavelengths are correctly converted from Angstroms to nm.""" + content = MockADF15Files.create_hydrogen_adf15() + filepath = self._create_test_file("test_wavelength.adf15", content) + + rates, wavelengths = parse_adf15(hydrogen, 0, filepath, header_format="hydrogen") + + # Check specific wavelengths (656.3 Angstrom = 65.63 nm) + transition_21 = (2, 1) + if transition_21 in wavelengths[hydrogen][0]: + wl = wavelengths[hydrogen][0][transition_21] + # Should be around 65.63 nm (converted from 656.3 Angstrom) + self.assertAlmostEqual(wl, 65.63, places=1) + + def test_multiple_rate_types(self): + """Test parsing file with multiple rate types.""" + content = MockADF15Files.create_hydrogen_adf15() + filepath = self._create_test_file("test_multitypes.adf15", content) + + rates, wavelengths = parse_adf15(hydrogen, 0, filepath, header_format="hydrogen") + + # Check both rate types exist + self.assertIn("excitation", rates) + self.assertIn("recombination", rates) + self.assertIn("thermalcx", rates) + + +if __name__ == "__main__": + unittest.main() diff --git a/cherab/tools/inversions/__init__.py b/cherab/tools/inversions/__init__.py index 00b67a9b..128a0a0d 100644 --- a/cherab/tools/inversions/__init__.py +++ b/cherab/tools/inversions/__init__.py @@ -19,7 +19,7 @@ from .sart import invert_sart, invert_constrained_sart from .opencl import SartOpencl -from .nnls import invert_regularised_nnls +from .nnls import invert_regularised_nnls, invert_sparse_regularised_nnls from .lstsq import invert_regularised_lstsq from .svd import invert_svd from .voxels import Voxel, AxisymmetricVoxel, VoxelCollection, ToroidalVoxelGrid, UnityVoxelEmitter diff --git a/cherab/tools/inversions/admt_utils.py b/cherab/tools/inversions/admt_utils.py index 549c4d42..d65ee694 100644 --- a/cherab/tools/inversions/admt_utils.py +++ b/cherab/tools/inversions/admt_utils.py @@ -28,23 +28,31 @@ from collections.abc import Mapping import numpy as np +from scipy.sparse import issparse +try: + from scipy.sparse import coo_array as coo, diags_array as diags +except ImportError: # Scipy < 1.8, deprecated from 1.18 + from scipy.sparse import coo_matrix as coo, diags -def generate_derivative_operators(voxel_vertices, grid_index_1d_to_2d_map, - grid_index_2d_to_1d_map): +def generate_derivative_operators(voxel_vertices, grid_index_1d_to_2d_map=None, + grid_index_2d_to_1d_map=None, sparse=False): r""" Generate the first and second derivative operators for a regular grid. :param ndarray voxel_vertices: an Nx4x2 array of coordinates of the - vertices of each voxel, (R, Z) + vertices of each voxel, (R, Z) :param dict grid_1d_to_2d_map: a mapping from the 1D array of - voxels in the grid to a 2D array of voxels if they were arranged - spatially. + voxels in the grid to a 2D array of voxels if they were arranged + spatially. Computed from grid_2d_to_1d_map if not given. :param dict grid_2d_to_1d_map: the inverse mapping from a 2D - spatially-arranged array of voxels to the 1D array. + spatially-arranged array of voxels to the 1D array. Computed from + grid_1d_to_2d_map if not given. + :param sparse: return the operators as sparse matrices if True, or + as dense matrices if False. - :return dict operators: a dictionary containing the derivative - operators: Dij for i, y ∊ (x, y) and Di for i ∊ (x, y). + :return: a dictionary containing the derivative operators: Dij for + i, j ∊ (x, y) and Di for i ∊ (x, y), Dsp and Dsm. This function assumes that all voxels are rectilinear, with their axes aligned to the coordinate axes. Additionally, all voxels are @@ -62,31 +70,48 @@ def generate_derivative_operators(voxel_vertices, grid_index_1d_to_2d_map, D_{xx} \equiv \frac{\partial^2}{\partial x^2}\\ D_{xy} \equiv \frac{\partial^2}{\partial x \partial y} - etc. + etc. It also produces two additional operators, Dsp and Dsm, for + second derivatives on the dy/dx = 1 and dy/dx = -1 diagonals + respectively. Note that the standard 2D laplacian (for isotropic regularisation) - can be trivially calculated as L = Dxx * dx + Dyy * dy, where dx and - dy are the voxel width and height respectively. This expression does - not however produce the 2D laplacian derived from the N-dimensional - case. + can be trivially calculated as follows: + + .. math:: + L = (1 - \alpha) (D_{xx} + D_{yy}) + (\alpha / 2) (D_{sp} + D_{sm}) + + α = 2/3 produces the operator used in Carr et. al. RSI 89, 083506 (2018). + α = 1/3 produces the operator with optimal isotropy. """ # Input argument validation: assume rectilinear voxels voxel_vertices = np.asarray(voxel_vertices) if voxel_vertices.ndim != 3 or voxel_vertices.shape[-2] != 4 or voxel_vertices.shape[-1] != 2: raise TypeError("voxel_vertices must be an NxMx2 array of vertices") - if not isinstance(grid_index_1d_to_2d_map, Mapping): + if not (isinstance(grid_index_1d_to_2d_map, Mapping) or grid_index_1d_to_2d_map is None): raise TypeError("grid_index_1d_to_2d_map should be dict-like") - if not isinstance(grid_index_2d_to_1d_map, Mapping): + if not (isinstance(grid_index_2d_to_1d_map, Mapping) or grid_index_2d_to_1d_map is None): raise TypeError("grid_index_2d_to_1d_map should be dict-like") + if grid_index_1d_to_2d_map is None and grid_index_2d_to_1d_map is None: + raise ValueError("At least one of grid_index_2d_to_1d_map or grid_index_1d_to_2d_map" + " must be given") + + # If only one of the mappings is given, compute the other one. + if grid_index_1d_to_2d_map is None and grid_index_2d_to_1d_map is not None: + grid_index_1d_to_2d_map = {k: rz for (rz, k) in grid_index_2d_to_1d_map.items()} + if grid_index_2d_to_1d_map is None and grid_index_1d_to_2d_map is not None: + grid_index_2d_to_1d_map = {rz: k for (k, rz) in grid_index_1d_to_2d_map.items()} num_cells = voxel_vertices.shape[0] cell_centres = np.mean(voxel_vertices, axis=1) # Individual derivative operators - Dx = np.zeros((num_cells, num_cells)) - Dy = np.zeros((num_cells, num_cells)) - Dxx = np.zeros((num_cells, num_cells)) - Dxy = np.zeros((num_cells, num_cells)) - Dyy = np.zeros((num_cells, num_cells)) + # Store derivative operators in dictionary-of-keys sparse array format. + Dx = {} + Dy = {} + Dxx = {} + Dxy = {} + Dyy = {} + Dsp = {} + Dsm = {} # TODO: for now, we assume all voxels have rectangular cross sections # which are approximately identical. As per Ingesson's notation, we # assume voxels are ordered from top left to bottom right, in column-major @@ -99,19 +124,29 @@ def generate_derivative_operators(voxel_vertices, grid_index_1d_to_2d_map, dx = np.min(abs(dx[dx != 0])).item() dy = np.min(abs(dy[dy != 0])).item() + # Work out how the voxels are ordered: increasing/decreasing in x/y. + xinc, yinc = np.sign(cell_centres[-1] - cell_centres[0]) + # Note that iy increases as y decreases (cells go from top to bottom), # which is the same as Ingesson's notation in equations 37-41 # Use the second version of the second derivative boundary formulae, so # that we only need to consider nearest neighbours for ith_cell in range(num_cells): at_top, at_bottom, at_left, at_right = False, False, False, False - n_left, n_right, n_below, n_above = np.nan, np.nan, np.nan, np.nan - n_above_left, n_above_right, n_below_left, n_below_right = np.nan, np.nan, np.nan, np.nan + n_left, n_right, n_below, n_above = None, None, None, None + n_above_left, n_above_right, n_below_left, n_below_right = None, None, None, None + # get the 2D mesh coordinates of this cell ix, iy = grid_index_1d_to_2d_map[ith_cell] + iright = ix + xinc + ileft = ix - xinc + iabove = iy + yinc + ibelow = iy - yinc + + # Handle voxels not at the edges/corners of the grid. try: - n_left = grid_index_2d_to_1d_map[ix - 1, iy] # left of n0 + n_left = grid_index_2d_to_1d_map[ileft, iy] # left of n0 except KeyError: at_left = True else: @@ -119,7 +154,7 @@ def generate_derivative_operators(voxel_vertices, grid_index_1d_to_2d_map, Dxx[ith_cell, n_left] = 1 try: - n_below_left = grid_index_2d_to_1d_map[ix - 1, iy + 1] # below left of n0 + n_below_left = grid_index_2d_to_1d_map[ileft, ibelow] # below left of n0 except KeyError: # KeyError does not necessarily mean bottom AND left pass @@ -127,7 +162,7 @@ def generate_derivative_operators(voxel_vertices, grid_index_1d_to_2d_map, Dxy[ith_cell, n_below_left] = 1 / 4 try: - n_below = grid_index_2d_to_1d_map[ix, iy + 1] + n_below = grid_index_2d_to_1d_map[ix, ibelow] except KeyError: at_bottom = True else: @@ -135,14 +170,14 @@ def generate_derivative_operators(voxel_vertices, grid_index_1d_to_2d_map, Dyy[ith_cell, n_below] = 1 try: - n_below_right = grid_index_2d_to_1d_map[ix + 1, iy + 1] + n_below_right = grid_index_2d_to_1d_map[iright, ibelow] except KeyError: pass else: Dxy[ith_cell, n_below_right] = -1 / 4 try: - n_right = grid_index_2d_to_1d_map[ix + 1, iy] + n_right = grid_index_2d_to_1d_map[iright, iy] except KeyError: at_right = True else: @@ -150,14 +185,14 @@ def generate_derivative_operators(voxel_vertices, grid_index_1d_to_2d_map, Dxx[ith_cell, n_right] = 1 try: - n_above_right = grid_index_2d_to_1d_map[ix + 1, iy - 1] + n_above_right = grid_index_2d_to_1d_map[iright, iabove] except KeyError: pass else: Dxy[ith_cell, n_above_right] = 1 / 4 try: - n_above = grid_index_2d_to_1d_map[ix, iy - 1] + n_above = grid_index_2d_to_1d_map[ix, iabove] except KeyError: at_top = True else: @@ -165,20 +200,24 @@ def generate_derivative_operators(voxel_vertices, grid_index_1d_to_2d_map, Dyy[ith_cell, n_above] = 1 try: - n_above_left = grid_index_2d_to_1d_map[ix - 1, iy - 1] + n_above_left = grid_index_2d_to_1d_map[ileft, iabove] except KeyError: pass else: Dxy[ith_cell, n_above_left] = -1 / 4 + + # Cases which are the same throughout the matrix. + Dxx[ith_cell, ith_cell] = -2 + Dyy[ith_cell, ith_cell] = -2 + + + # Handle cases at the edges/corners top_left = at_top and at_left top_right = at_top and at_right bottom_left = at_bottom and at_left bottom_right = at_bottom and at_right - Dxx[ith_cell, ith_cell] = -2 - Dyy[ith_cell, ith_cell] = -2 - if at_left: Dx[ith_cell, ith_cell] = -1 Dx[ith_cell, n_right] = 1 @@ -247,14 +286,71 @@ def generate_derivative_operators(voxel_vertices, grid_index_1d_to_2d_map, Dxy[ith_cell, ith_cell] = -1 Dxy[ith_cell, n_above_left] = -1 + + # Handle the "skewed" operators. + if n_above_left is None and n_below_right is not None: + Dsm[ith_cell, ith_cell] = -1 + Dsm[ith_cell, n_below_right] = 1 + elif n_below_right is None and n_above_left is not None: + Dsm[ith_cell, ith_cell] = -1 + Dsm[ith_cell, n_above_left] = 1 + elif n_above_left is None and n_below_right is None: + Dsm[ith_cell, ith_cell] = 0 + else: + Dsm[ith_cell, ith_cell] = -2 + Dsm[ith_cell, n_above_left] = 1 + Dsm[ith_cell, n_below_right] = 1 + + if n_above_right is None and n_below_left is not None: + Dsp[ith_cell, ith_cell] = -1 + Dsp[ith_cell, n_below_left] = 1 + elif n_below_left is None and n_above_right is not None: + Dsp[ith_cell, ith_cell] = -1 + Dsp[ith_cell, n_above_right] = 1 + elif n_below_left is None and n_above_right is None: + Dsp[ith_cell, ith_cell] = 0 + else: + Dsp[ith_cell, ith_cell] = -2 + Dsp[ith_cell, n_above_right] = 1 + Dsp[ith_cell, n_below_left] = 1 + + + # Although we've stored the operators as dictionaries of keys, it turns out to be + # more convenient to construct a COOrdinate sparse matrix rather than a DOK one + # in Scipy. We then convert that to CSR representation for efficient numerical + # operations later. + def dok_to_sparse(D): + row, col = zip(*D.keys()) + vals = list(D.values()) + return coo((vals, (row, col)), shape=(num_cells, num_cells)).tocsr() + + Dx = dok_to_sparse(Dx) + Dy = dok_to_sparse(Dy) + Dxx = dok_to_sparse(Dxx) + Dyy = dok_to_sparse(Dyy) + Dxy = dok_to_sparse(Dxy) + Dsp = dok_to_sparse(Dsp) + Dsm = dok_to_sparse(Dsm) Dx = Dx / dx Dy = Dy / dy Dxx = Dxx / dx**2 Dyy = Dyy / dy**2 Dxy = Dxy / (dx * dy) + Dsp = Dsp / (dx**2 + dy**2) + Dsm = Dsm / (dx**2 + dy**2) + + # If the user requests dense matrices, convert them after performing all the scaling. + if not sparse: + Dx = Dx.toarray() + Dy = Dy.toarray() + Dxx = Dxx.toarray() + Dyy = Dyy.toarray() + Dxy = Dxy.toarray() + Dsp = Dsp.toarray() + Dsm = Dsm.toarray() # Package all operators up into a dictionary - operators = dict(Dx=Dx, Dy=Dy, Dxx=Dxx, Dyy=Dyy, Dxy=Dxy) + operators = dict(Dx=Dx, Dy=Dy, Dxx=Dxx, Dyy=Dyy, Dxy=Dxy, Dsp=Dsp, Dsm=Dsm) return operators @@ -263,22 +359,21 @@ def calculate_admt(voxel_radii, derivative_operators, psi_at_voxels, dx, dy, ani Calculate the ADMT regularisation operator. :param ndarray voxel_radii: a 1D array of the radius at the centre - of each voxel in the grid - :param tuple derivative_operators: a named tuple with the derivative - operators for the grid, as returned by :func:generate_derivative_operators + of each voxel in the grid + :param dict derivative_operators: a dictionary with the derivative + operators for the grid, as returned by :func:generate_derivative_operators :param ndarray psi_at_voxels: the magnetic flux at the centre of - each voxel in the grid + each voxel in the grid :param float dx: the width of each voxel. :param float dy: the height of each voxel :param float anisotropy: the ratio of the smoothing in the parallel - and perpendicular directions. - - :return ndarray admt: the ADMT regularisation operator. + and perpendicular directions. + :return: the ADMT regularisation operator. The degree of anisotropy dictates the relative suppression of gradients in the directions parallel and perpendicular to the - magnetic field. For example, `anisotropy=10` implies parallel - gradients in solution are 10 times smaller than perpendicular + magnetic field. For example, ``anisotropy=10`` implies parallel + gradients in the solution are 10 times smaller than perpendicular gradients. This function assumes that all voxels are rectilinear, with their @@ -294,6 +389,10 @@ def calculate_admt(voxel_radii, derivative_operators, psi_at_voxels, dx, dy, ani This means it is suitable for use in Cherab's inversion methods, such as NNLS and SART. + + If the derivative operators are sparse matrices, the returned admt + operator is also a sparse matrix. Otherwise a dense matrix is + returned. """ Dpar = np.full(psi_at_voxels.shape, 1) Dperp = Dpar / anisotropy @@ -345,11 +444,17 @@ def calculate_admt(voxel_radii, derivative_operators, psi_at_voxels, dx, dy, ani + (Dperp - Dpar) * (dpsidxdy * dpsidx + dpsidxx * dpsidy) + ddiff_term_cy + dnorm_term_cy + toroidal_term_cy ) / normalisation - cx = np.diag(cx) - cy = np.diag(cy) - cxx = np.diag(cxx) - cyy = np.diag(cyy) - cxy = np.diag(cxy) + if all(issparse(d) for d in derivative_operators.values()): + # Make sparse versions of the diagonal matrices. + diag = diags + else: + # Dense versions using Numpy. + diag = np.diag + cx = diag(cx) + cy = diag(cy) + cxx = diag(cxx) + cyy = diag(cyy) + cxy = diag(cxy) admt_operator = cx @ Dx + cy @ Dy + cxx @ Dxx + 2 * cxy @ Dxy + cyy @ Dyy admt_operator *= np.sqrt(dx * dy) return admt_operator diff --git a/cherab/tools/inversions/nnls.py b/cherab/tools/inversions/nnls.py index 34779f71..e4dff19d 100644 --- a/cherab/tools/inversions/nnls.py +++ b/cherab/tools/inversions/nnls.py @@ -19,6 +19,10 @@ import numpy as np import scipy +try: + from scipy.sparse import lil_array as lil, eye_array as eye +except ImportError: # Scipy < 1.8, deprecated from 1.18 + from scipy.sparse import lil_matrix as lil, eye def invert_regularised_nnls(w_matrix, b_vector, alpha=0.01, tikhonov_matrix=None, **kwargs): @@ -29,7 +33,7 @@ def invert_regularised_nnls(w_matrix, b_vector, alpha=0.01, tikhonov_matrix=None This is a thin wrapper around scipy.optimize.nnls, which modifies the arguments to include the supplied Tikhonov regularisation matrix. - The values of w_matrix, b_vector and alpha * tikhonov_matrix are notmalised + The values of w_matrix, b_vector and alpha * tikhonov_matrix are normalised by max(b_vector) before passing them to scipy.optimize.nnls(). :param np.ndarray w_matrix: The sensitivity matrix describing the coupling between the @@ -70,3 +74,60 @@ def invert_regularised_nnls(w_matrix, b_vector, alpha=0.01, tikhonov_matrix=None x_vector, rnorm = scipy.optimize.nnls(c_matrix / vmax, d_vector / vmax, **kwargs) return x_vector, rnorm * vmax + + +def invert_sparse_regularised_nnls(w_matrix, b_vector, alpha=0.01, tikhonov_matrix=None, **kwargs): + r""" + Solves :math:`\mathbf{b} = \mathbf{W} \mathbf{x}` for the vector :math:`\mathbf{x}`, + using Tikhonov regulariastion. + + This is a thin wrapper around scipy.optimize.lsq_linear which modifies + the arguments to include the supplied Tikhonov regularisation matrix and + enforces bounds to avoid negativity. + + The values of w_matrix, b_vector and alpha * tikhonov_matrix are normalised + by max(b_vector) before passing them to scipy.optimize.lsq_linear(). + + :param w_matrix: The sensitivity matrix describing the coupling between the + detectors and the voxels. Must be an array with shape :math:`(N_d, N_s)`. May be either + a dense array or a sparse matrix or array. + :param np.ndarray b_vector: The measured power/radiance vector with shape :math:`(N_d)`. + :param float alpha: The regularisation hyperparameter :math:`\alpha` which determines + the regularisation strength of the tikhonov matrix. + :param np.ndarray tikhonov_matrix: The tikhonov regularisation matrix operator, an array + with shape :math:`(N_s, N_s)`. If None, the identity matrix is used. + :param \**kwargs: Keyword arguments passed to scipy.optimize.lsq_linear. + :return: (x, norm), the solution vector and the residual norm. + + .. code-block:: pycon + + >>> from cherab.tools.inversions import invert_sparse_regularised_nnls + >>> x, norm = invert_sparse_regularised_nnls(w_matrix, b_vector, tikhonov_matrix=tikhonov_matrix) + """ + + m, n = w_matrix.shape + + if tikhonov_matrix is None: + tikhonov_matrix = eye(n) + + tikhonov_matrix = alpha * tikhonov_matrix + + # Extend W to have form ... + c_matrix = lil((m+n, n)) + c_matrix[0:m, :] = w_matrix[:, :] + c_matrix[m:, :] = tikhonov_matrix[:, :] + c_matrix = c_matrix.tocsr() + + # Extend b to have form ... + d_vector = np.zeros(m+n) + d_vector[0:m] = b_vector[:] + + # Normalise c_matrix and d_vector to avoid possible issues with the inversion termination criteria. + vmax = d_vector.max() + + res = scipy.optimize.lsq_linear(c_matrix / vmax, d_vector / vmax, bounds=(0, np.inf), **kwargs) + + x_vector = res.x + rnorm = np.linalg.norm(res.fun) + + return x_vector, rnorm * vmax diff --git a/cherab/tools/observers/__init__.py b/cherab/tools/observers/__init__.py index d134ef63..99cf5ab7 100644 --- a/cherab/tools/observers/__init__.py +++ b/cherab/tools/observers/__init__.py @@ -21,4 +21,4 @@ from .calcam import load_calcam_calibration from .intersections import find_wall_intersection from .spectroscopy import SpectroscopicSightLine, SpectroscopicFibreOptic -from .group import PixelGroup, TargettedPixelGroup, SightLineGroup, FibreOpticGroup, SpectroscopicFibreOpticGroup, SpectroscopicSightLineGroup +from .group import PixelGroup, TargetedPixelGroup, TargettedPixelGroup, SightLineGroup, FibreOpticGroup, SpectroscopicFibreOpticGroup, SpectroscopicSightLineGroup diff --git a/cherab/tools/observers/bolometry.py b/cherab/tools/observers/bolometry.py index 20179796..33870b0a 100644 --- a/cherab/tools/observers/bolometry.py +++ b/cherab/tools/observers/bolometry.py @@ -18,16 +18,17 @@ # under the Licence. from enum import Enum +from warnings import warn import functools import numpy as np from raysect.core import Node, translate, rotate_basis, Point3D, Vector3D, Ray as CoreRay, Primitive, World -from raysect.core.math.sampler import TargettedHemisphereSampler, RectangleSampler3D +from raysect.core.math.sampler import TargetedHemisphereSampler, RectangleSampler3D from raysect.primitive import Box, Cylinder, Subtract, Union from raysect.optical.observer import PowerPipeline0D, RadiancePipeline0D, \ - SpectralPowerPipeline0D, SpectralRadiancePipeline0D, SightLine, TargettedPixel + SpectralPowerPipeline0D, SpectralRadiancePipeline0D, SightLine, TargetedPixel from raysect.optical.observer import PowerPipeline2D, RadiancePipeline2D, \ - SpectralPowerPipeline2D, SpectralRadiancePipeline2D, TargettedCCDArray + SpectralPowerPipeline2D, SpectralRadiancePipeline2D, TargetedCCDArray from raysect.optical.material.material import NullMaterial from raysect.optical.material import AbsorbingSurface @@ -213,7 +214,7 @@ class BolometerSlit(Node): larger than the slit dx and dy, which can cause partial occlusion of nearby primitives. It also relies on no rays being launched with directions outside the solid angle of the aperture's bounding sphere: depending on the - foil-slit distance and slit size, and also the foil's targetted_path_prob, + foil-slit distance and slit size, and also the foil's targeted_path_prob, this may not be guaranteed. Supplying a proper mesh geometry for the camera is recommended instead of using a CSG aperture. @@ -351,7 +352,7 @@ def curvature_radius(self): return self._curvature_radius -class BolometerFoil(TargettedPixel): +class BolometerFoil(TargetedPixel): """ A rectangular foil bolometer detector. @@ -447,7 +448,7 @@ def __init__(self, detector_id, centre_point, basis_x, dx, basis_y, dy, slit, translation = translate(centre_point.x, centre_point.y, centre_point.z) rotation = rotate_basis(normal_vec, basis_y) - super().__init__([slit.target], targetted_path_prob=1.0, + super().__init__([slit.target], targeted_path_prob=1.0, pixel_samples=1000, x_width=dx, y_width=dy, spectral_bins=1, quiet=True, parent=parent, transform=translation * rotation, name=detector_id) @@ -516,6 +517,24 @@ def accumulate(self, value): # Discard any samples from previous accumulate behaviour pipeline.value.clear() + @property + def targetted_path_prob(self): + warn( + "The 'targetted_path_prob' property is deprecated, use 'targeted_path_prob' instead.", + DeprecationWarning, + stacklevel=2 + ) + return self._targeted_path_prob + + @targetted_path_prob.setter + def targetted_path_prob(self, value): + warn( + "The 'targetted_path_prob' property is deprecated, use 'targeted_path_prob' instead.", + DeprecationWarning, + stacklevel=2 + ) + self.targeted_path_prob = value + def as_sightline(self): """ Constructs a SightLine observer for this bolometer. @@ -661,8 +680,8 @@ def calculate_etendue(self, ray_count=10000, batches=10, max_distance=1e999): # generate bounding sphere and convert to local coordinate system sphere = target.bounding_sphere() spheres = [(sphere.centre.transform(self.to_local()), sphere.radius, 1.0)] - # instance targetted pixel sampler to sample directions - targetted_sampler = TargettedHemisphereSampler(spheres) + # instance targeted pixel sampler to sample directions + targeted_sampler = TargetedHemisphereSampler(spheres) # instance rectangle pixel sampler to sample origins point_sampler = RectangleSampler3D(width=self.x_width, height=self.y_width) @@ -671,8 +690,8 @@ def etendue_single_run(_): origins = point_sampler(samples=ray_count) passed = 0.0 for origin in origins: - # obtain targetted vector sample - direction, pdf = targetted_sampler(origin, pdf=True) + # obtain targeted vector sample + direction, pdf = targeted_sampler(origin, pdf=True) path_weight = R_2_PI * direction.z / pdf # Transform to world space origin = origin.transform(detector_transform) @@ -701,7 +720,7 @@ def etendue_single_run(_): return etendue, etendue_error -class BolometerIRVB(TargettedCCDArray): +class BolometerIRVB(TargetedCCDArray): """ A rectangular infra red video bolometer (IRVB). @@ -784,7 +803,7 @@ def __init__(self, name, width, pixels, slit, transform, parent=None, self._accumulate = None # Will be set after pipeline is created. super().__init__([slit.target], pixels=pixels, width=width, - targetted_path_prob=0.99, parent=parent, pipelines=[], + targeted_path_prob=0.99, parent=parent, pipelines=[], transform=transform, name=name) self.pixel_samples = 1000 self.spectral_bins = 1 @@ -896,6 +915,24 @@ def accumulate(self, value): if pipeline.frame is not None: pipeline.frame.clear() + @property + def targetted_path_prob(self): + warn( + "The 'targetted_path_prob' property is deprecated, use 'targeted_path_prob' instead.", + DeprecationWarning, + stacklevel=2, + ) + return self._targeted_path_prob + + @targetted_path_prob.setter + def targetted_path_prob(self, value): + warn( + "The 'targetted_path_prob' property is deprecated, use 'targeted_path_prob' instead.", + DeprecationWarning, + stacklevel=2, + ) + self.targeted_path_prob = value + def as_sightlines(self): """ Constructs a SightLine observer for each pixel in this bolometer. diff --git a/cherab/tools/observers/calcam.py b/cherab/tools/observers/calcam.py index 3cfc22d4..1c8460eb 100644 --- a/cherab/tools/observers/calcam.py +++ b/cherab/tools/observers/calcam.py @@ -18,7 +18,7 @@ # under the Licence. import numpy as np -from scipy.io.netcdf import netcdf_file +from scipy.io import netcdf_file from raysect.core import Point3D, Vector3D diff --git a/cherab/tools/observers/group/__init__.py b/cherab/tools/observers/group/__init__.py index eca93585..bbcf8633 100644 --- a/cherab/tools/observers/group/__init__.py +++ b/cherab/tools/observers/group/__init__.py @@ -17,7 +17,8 @@ # under the Licence. from .fibreoptic import FibreOpticGroup -from .sightline import SightLineGroup -from .targettedpixel import TargettedPixelGroup from .pixel import PixelGroup +from .sightline import SightLineGroup from .spectroscopic import SpectroscopicFibreOpticGroup, SpectroscopicSightLineGroup +from .targetedpixel import TargetedPixelGroup +from .targettedpixel import TargettedPixelGroup diff --git a/cherab/tools/observers/group/targetedpixel.py b/cherab/tools/observers/group/targetedpixel.py new file mode 100644 index 00000000..9f4fb814 --- /dev/null +++ b/cherab/tools/observers/group/targetedpixel.py @@ -0,0 +1,123 @@ +# Copyright 2016-2021 Euratom +# Copyright 2016-2021 United Kingdom Atomic Energy Authority +# Copyright 2016-2021 Centro de Investigaciones Energéticas, Medioambientales y Tecnológicas +# +# Licensed under the EUPL, Version 1.1 or – as soon they will be approved by the +# European Commission - subsequent versions of the EUPL (the "Licence"); +# You may not use this work except in compliance with the Licence. +# You may obtain a copy of the Licence at: +# +# https://joinup.ec.europa.eu/software/page/eupl5 +# +# Unless required by applicable law or agreed to in writing, software distributed +# under the Licence is distributed on an "AS IS" basis, WITHOUT WARRANTIES OR +# CONDITIONS OF ANY KIND, either express or implied. +# +# See the Licence for the specific language governing permissions and limitations +# under the Licence. + +from numpy import ndarray +from raysect.optical.observer import TargetedPixel + +from .base import Observer0DGroup + + +class TargetedPixelGroup(Observer0DGroup): + """ + A group of targeted pixels under a single scene-graph node. + + A scene-graph object regrouping a series of 'TargetedPixel' + observers as a scene-graph parent. Allows combined observation and display + control simultaneously. + + :ivar list x_width: Width of pixel along local x axis + :ivar list y_width: Width of pixel along local y axis + :ivar list targets: Targets for preferential sampling + :ivar list targeted_path_prob: Probability of ray being casted at the target + """ + + _OBSERVER_TYPE = TargetedPixel + + @property + def x_width(self): + return [pixel.x_width for pixel in self._observers] + + @x_width.setter + def x_width(self, value): + if isinstance(value, (list, tuple, ndarray)): + if len(value) == len(self._observers): + for pixel, v in zip(self._observers, value): + pixel.x_width = v + else: + raise ValueError( + "The length of 'x_width' ({}) mismatches the number of pixels ({}).".format(len(value), len(self._observers)) + ) + else: + for pixel in self._observers: + pixel.x_width = value + + @property + def y_width(self): + return [pixel.y_width for pixel in self._observers] + + @y_width.setter + def y_width(self, value): + if isinstance(value, (list, tuple, ndarray)): + if len(value) == len(self._observers): + for pixel, v in zip(self._observers, value): + pixel.y_width = v + else: + raise ValueError( + "The length of 'y_width' ({}) mismatches the number of pixels ({}).".format(len(value), len(self._observers)) + ) + else: + for pixel in self._observers: + pixel.y_width = value + + @property + def targets(self): + """ + List of target lists used by pixels for preferential sampling + + :param list value: List of primitives to be set to each pixel or + list of lists containing targets specific for each pixel + in this case the number of lists must match number of pixels + + :rtype: list + """ + return [pixel.targets for pixel in self._observers] + + @targets.setter + def targets(self, value): + if all(isinstance(v, (list, tuple)) for v in value): + if len(value) == len(self._observers): + for pixel, v in zip(self._observers, value): + pixel.targets = v + else: + raise ValueError( + "The number of provided target lists' ({}) mismatches the number of pixels ({}).".format( + len(value), len(self._observers) + ) + ) + else: + # assuming a list of primitives, the pixel's setter will throw an error if not + for pixel in self._observers: + pixel.targets = value + + @property + def targeted_path_prob(self): + return [pixel.targeted_path_prob for pixel in self._observers] + + @targeted_path_prob.setter + def targeted_path_prob(self, value): + if isinstance(value, (list, tuple)): + if len(value) == len(self._observers): + for pixel, v in zip(self._observers, value): + pixel.targeted_path_prob = v + else: + raise ValueError( + "The length of 'value' ({}) mismatches the number of pixels ({}).".format(len(value), len(self._observers)) + ) + else: + for pixel in self._observers: + pixel.targeted_path_prob = value diff --git a/cherab/tools/observers/group/targettedpixel.py b/cherab/tools/observers/group/targettedpixel.py index 07d13e7f..5d1ac69f 100644 --- a/cherab/tools/observers/group/targettedpixel.py +++ b/cherab/tools/observers/group/targettedpixel.py @@ -16,17 +16,20 @@ # See the Licence for the specific language governing permissions and limitations # under the Licence. -from numpy import ndarray -from raysect.optical.observer import TargettedPixel +import warnings -from .base import Observer0DGroup +from .targetedpixel import TargetedPixelGroup as _TargetedPixelGroup -class TargettedPixelGroup(Observer0DGroup): +class TargettedPixelGroup(_TargetedPixelGroup): """ - A group of targetted pixel under a single scene-graph node. + A group of targeted pixel under a single scene-graph node. - A scene-graph object regrouping a series of 'TargettedPixel' + .. deprecated:: + `TargettedPixelGroup` is deprecated and will be removed in version 2.0. + Use `TargetedPixelGroup` instead. + + A scene-graph object regrouping a series of `TargetedPixel` observers as a scene-graph parent. Allows combined observation and display control simultaneously. @@ -35,82 +38,20 @@ class TargettedPixelGroup(Observer0DGroup): :ivar list targets: Targets for preferential sampling :ivar list targetted_path_prob: Probability of ray being casted at the target """ - _OBSERVER_TYPE = TargettedPixel - - @property - def x_width(self): - return [pixel.x_width for pixel in self._observers] - - @x_width.setter - def x_width(self, value): - if isinstance(value, (list, tuple, ndarray)): - if len(value) == len(self._observers): - for pixel, v in zip(self._observers, value): - pixel.x_width = v - else: - raise ValueError("The length of 'x_width' ({}) " - "mismatches the number of pixels ({}).".format(len(value), len(self._observers))) - else: - for pixel in self._observers: - pixel.x_width = value - - @property - def y_width(self): - return [pixel.y_width for pixel in self._observers] - @y_width.setter - def y_width(self, value): - if isinstance(value, (list, tuple, ndarray)): - if len(value) == len(self._observers): - for pixel, v in zip(self._observers, value): - pixel.y_width = v - else: - raise ValueError("The length of 'y_width' ({}) " - "mismatches the number of pixels ({}).".format(len(value), len(self._observers))) - else: - for pixel in self._observers: - pixel.y_width = value - - @property - def targets(self): - """ - List of target lists used by pixels for preferential sampling - - :param list value: List of primitives to be set to each pixel or - list of lists containing targets specific for each pixel - in this case the number of lists must match number of pixels - - :rtype: list - """ - return [pixel.targets for pixel in self._observers] - - @targets.setter - def targets(self, value): - if all(isinstance(v, (list, tuple)) for v in value): - if len(value) == len(self._observers): - for pixel, v in zip(self._observers, value): - pixel.targets = v - else: - raise ValueError("The number of provided target lists' ({}) " - "mismatches the number of pixels ({}).".format(len(value), len(self._observers))) - else: - # assuming a list of primitives, the pixel's setter will throw an error if not - for pixel in self._observers: - pixel.targets = value + def __init__(self, *args, **kwargs): + warnings.warn( + "TargettedPixelGroup is deprecated and will be removed in version 2.0. " + + "Use TargetedPixelGroup instead.", + DeprecationWarning, + stacklevel=2, + ) + super().__init__(*args, **kwargs) @property def targetted_path_prob(self): - return [pixel.targetted_path_prob for pixel in self._observers] - + return self.targeted_path_prob + @targetted_path_prob.setter def targetted_path_prob(self, value): - if isinstance(value, (list, tuple)): - if len(value) == len(self._observers): - for pixel, v in zip(self._observers, value): - pixel.targetted_path_prob = v - else: - raise ValueError("The length of 'value' ({}) " - "mismatches the number of pixels ({}).".format(len(value), len(self._observers))) - else: - for pixel in self._observers: - pixel.targetted_path_prob = value + self.targeted_path_prob = value diff --git a/cherab/tools/primitives/axisymmetric_mesh.pyx b/cherab/tools/primitives/axisymmetric_mesh.pyx index 81533cb0..d3c79c37 100644 --- a/cherab/tools/primitives/axisymmetric_mesh.pyx +++ b/cherab/tools/primitives/axisymmetric_mesh.pyx @@ -23,10 +23,10 @@ from .toroidal_mesh import toroidal_mesh_from_polygon cpdef Mesh axisymmetric_mesh_from_polygon(object polygon, int num_toroidal_segments=500): """ - Generates an Raysect Mesh primitive from the specified 2D polygon. + Generate a Raysect Mesh primitive from the specified 2D polygon. - :param object polygon: An object which can be converted to a numpy array with shape [N,2] - specifying the wall outline polygon in the R-Z plane. The polygon + :param object polygon: An object which can be converted to a numpy array with shape [N,2] + specifying the wall outline polygon in the R-Z plane. The polygon should not be closed, i.e. vertex i = 0 and i = N should not be the same vertex, but neighbours. :param int num_toroidal_segments: The number of repeating toroidal segments that will be used diff --git a/cherab/tools/tests/test_admt.py b/cherab/tools/tests/test_admt.py index b7a3a7dd..ac979698 100644 --- a/cherab/tools/tests/test_admt.py +++ b/cherab/tools/tests/test_admt.py @@ -107,6 +107,10 @@ class TestADMT(unittest.TestCase): VOXEL_VERTICES, GRID_1D_TO_2D_MAP, GRID_2D_TO_1D_MAP ) + SPARSE_DERIVATIVE_OPERATORS = generate_derivative_operators( + VOXEL_VERTICES, GRID_1D_TO_2D_MAP, GRID_2D_TO_1D_MAP, sparse=True, + ) + def test_dx(self): """D/Dx (Equations 37)""" DtestDx = self.DERIVATIVE_OPERATORS["Dx"] @ self.VOXEL_TEST_DATA @@ -234,6 +238,35 @@ def test_invalid_2d_1d_mapping(self): generate_derivative_operators(self.VOXEL_VERTICES, self.GRID_2D_TO_1D_MAP, self.TEST_DATA_2D) + def test_only_1d_2d_mapping_provided(self): + """Test auto-computing 2D-to-1D mapping""" + derivs = generate_derivative_operators( + voxel_vertices=self.VOXEL_VERTICES, + grid_index_1d_to_2d_map=self.GRID_1D_TO_2D_MAP, + ) + for key in derivs.keys(): + np.testing.assert_equal(derivs[key], self.DERIVATIVE_OPERATORS[key]) + + def test_only_2d_1d_mapping_provided(self): + """Test auto-computing 1D-to-2D mapping""" + derivs = generate_derivative_operators( + voxel_vertices=self.VOXEL_VERTICES, + grid_index_2d_to_1d_map=self.GRID_2D_TO_1D_MAP, + ) + for key in derivs.keys(): + np.testing.assert_equal(derivs[key], self.DERIVATIVE_OPERATORS[key]) + + def test_missing_mappings(self): + """Test for raising if neither mapping is provided.""" + with self.assertRaises(ValueError): + generate_derivative_operators(self.VOXEL_VERTICES) + + def test_sparse_derivatives(self): + """Test returning sparse arrays.""" + for key in self.DERIVATIVE_OPERATORS.keys(): + np.testing.assert_equal(self.SPARSE_DERIVATIVE_OPERATORS[key].toarray(), + self.DERIVATIVE_OPERATORS[key]) + def test_objective(self, debug=False): """Test that the objective function looks sensible.""" # Make a test equilibrium and an emission vector which corresponds @@ -284,6 +317,26 @@ def test_objective(self, debug=False): print(kernel.sum()) # Should be zero for large grids plot_kernel(kernel, self.VOXEL_VERTICES) + def test_sparse_objective(self): + theta = np.pi / 2 # Vertical field + points = self.VOXELS_2D.reshape((-1, 2)) + test_field = sample2d_points( + lambda x, y: x * np.sin(theta) + y * np.cos(theta), + points + ) + test_field_2d = test_field.reshape(self.VOXELS_2D[:, :, 0].shape) + voxel_radii = np.asarray(self.VOXEL_COORDS)[:, 0] + dense_admt_operator = calculate_admt( + voxel_radii, self.DERIVATIVE_OPERATORS, test_field, + self.DX, self.DY, anisotropy=10, + ) + sparse_admt_operator = calculate_admt( + voxel_radii, self.SPARSE_DERIVATIVE_OPERATORS, test_field, + self.DX, self.DY, anisotropy=10, + ) + # Sparse matrix math may differ from dense due to floating point precision. + np.testing.assert_allclose(dense_admt_operator, sparse_admt_operator.toarray(), rtol=1e-14) + def plot_kernel(kernel, voxel_vertices): """Plot a 1D grid function as a 2D image""" diff --git a/cherab/tools/tests/test_observer_groups.py b/cherab/tools/tests/test_observer_groups.py index 1f5b7cb0..a602531e 100644 --- a/cherab/tools/tests/test_observer_groups.py +++ b/cherab/tools/tests/test_observer_groups.py @@ -1,11 +1,12 @@ import unittest +import warnings from raysect.core.workflow import RenderEngine -from raysect.optical.observer import Observer0D, SightLine, FibreOptic, Pixel, TargettedPixel, PowerPipeline0D, SpectralPowerPipeline0D +from raysect.optical.observer import Observer0D, SightLine, FibreOptic, Pixel, TargetedPixel, PowerPipeline0D, SpectralPowerPipeline0D from raysect.primitive import Sphere +from cherab.tools.observers.group import FibreOpticGroup, PixelGroup, SightLineGroup, TargetedPixelGroup, TargettedPixelGroup from cherab.tools.observers.group.base import Observer0DGroup -from cherab.tools.observers.group import SightLineGroup, FibreOpticGroup, PixelGroup, TargettedPixelGroup from cherab.tools.raytransfer import pipelines @@ -26,7 +27,7 @@ def test_get_item(self): idx = slice(1, 3, 1) for observer, input_observer in zip(group[idx], self.observers[idx]): self.assertIs(observer, input_observer) - + for i, name in enumerate(names): self.assertIs(group[name], self.observers[i]) @@ -83,7 +84,7 @@ def test_assignments(self): with self.assertRaises(ValueError): group.pipelines = [ppln_0] - # render_engine + # render_engine engine = RenderEngine() group.render_engine = engine for group_engine in group.render_engine: @@ -102,7 +103,7 @@ def test_assignments(self): with self.assertRaises(ValueError): group.render_engine = [RenderEngine() for _ in range(len(group) - 1)] - # wavelengths + # wavelengths wvl = 500 group.min_wavelength = wvl - 100 group.max_wavelength = wvl + 100 @@ -139,7 +140,7 @@ def test_assignments(self): with self.assertRaises(ValueError): group.spectral_bins = [1000] * (len(group) + 1) - # quiet + # quiet quiet = [True] * len(group) group.quiet = quiet self.assertListEqual(group.quiet, quiet) @@ -152,7 +153,7 @@ def test_assignments(self): with self.assertRaises(ValueError): group.quiet = [False] * (len(group) + 1) - # rays + # rays probs = [0.2 + i*0.1 for i in range(len(group))] max_depths = [5 + i for i in range(len(group))] min_depths = [2 + i for i in range(len(group))] @@ -196,7 +197,7 @@ def test_assignments(self): group.ray_importance_sampling = [False] * (len(group) + 1) with self.assertRaises(ValueError): group.ray_important_path_weight = [0.7] * (len(group) + 1) - + # samples pixel_samples = [2000 + i*500 for i in range(len(group))] per_task = [5000 + i*100 for i in range(len(group))] @@ -348,11 +349,11 @@ def test_widths(self): group.y_width = [1e-1] * (len(group) + 1) -class TargettedPixelGroupTestCase(PixelGroupTestCase): - _GROUP_CLASS = TargettedPixelGroup +class TargetedPixelGroupTestCase(PixelGroupTestCase): + _GROUP_CLASS = TargetedPixelGroup def setUp(self): - self.observers = [TargettedPixel(targets=[Sphere()], pipelines=[PowerPipeline0D()]) for _ in range(self._NUM)] + self.observers = [TargetedPixel(targets=[Sphere()], pipelines=[PowerPipeline0D()]) for _ in range(self._NUM)] def test_targets(self): group = self._GROUP_CLASS(observers=self.observers) @@ -374,15 +375,33 @@ def test_targets(self): with self.assertRaises(ValueError): group.targets = targets - # targetted path prob + # targeted path prob prob = [0.9, 0.95, 1] - group.targetted_path_prob = prob - self.assertListEqual(group.targetted_path_prob, prob) + group.targeted_path_prob = prob + self.assertListEqual(group.targeted_path_prob, prob) prob = 0.8 - group.targetted_path_prob = prob - for group_targetted_path_prob in group.targetted_path_prob: - self.assertEqual(group_targetted_path_prob, prob) + group.targeted_path_prob = prob + for group_targeted_path_prob in group.targeted_path_prob: + self.assertEqual(group_targeted_path_prob, prob) with self.assertRaises(ValueError): - group.targetted_path_prob = [0.7] * (len(group) + 1) + group.targeted_path_prob = [0.7] * (len(group) + 1) + + +class TargettedPixelGroupTestCase(TargetedPixelGroupTestCase): + """Test case for deprecated TargettedPixelGroup class.""" + + _GROUP_CLASS = TargettedPixelGroup + + def test_deprecation_warning(self): + """Test that using TargettedPixelGroup raises a deprecation warning.""" + with warnings.catch_warnings(record=True) as w: + warnings.simplefilter("always") + group = TargettedPixelGroup(observers=self.observers) + + # Check that a warning was issued + self.assertEqual(len(w), 1) + self.assertTrue(issubclass(w[0].category, DeprecationWarning)) + self.assertIn("TargettedPixelGroup is deprecated", str(w[0].message)) + self.assertIn("Use TargetedPixelGroup instead", str(w[0].message)) diff --git a/cherab/tools/tests/test_voxels.py b/cherab/tools/tests/test_voxels.py index 046fffbe..62b4fcc6 100644 --- a/cherab/tools/tests/test_voxels.py +++ b/cherab/tools/tests/test_voxels.py @@ -280,8 +280,7 @@ def test_rectangle_area(self): for rectangle in RECTANGULAR_VOXEL_COORDS: coords = np.asarray(rectangle) voxel = AxisymmetricVoxel(coords) - dx = coords[:, 0].ptp() - dy = coords[:, 1].ptp() + dx, dy = np.ptp(coords, axis=0) expected_area = dx * dy self.assertEqual(voxel.cross_sectional_area, expected_area) diff --git a/demos/observers/bolometry/admt_tomographic_inversion.py b/demos/observers/bolometry/admt_tomographic_inversion.py new file mode 100644 index 00000000..c38cebfc --- /dev/null +++ b/demos/observers/bolometry/admt_tomographic_inversion.py @@ -0,0 +1,385 @@ +""" +This example demonstrates performing a tomographic reconstruction of a +radiation profile using Cherab's anisotropic diffusion (ADMT) regularisation +utilities. We use the machine geometry, sample bolometers and equilibrium +from Generomak. +""" +import matplotlib.pyplot as plt +import numpy as np + +from raysect.core.math.function.float import Exp2D, Arg2D, Atan4Q2D +from raysect.core.math import translate +from raysect.optical import World +from raysect.optical.material import AbsorbingSurface, VolumeTransform +from raysect.primitive import Cylinder, Subtract + +from cherab.generomak.machine import load_first_wall +from cherab.generomak.equilibrium import load_equilibrium +from cherab.generomak.diagnostics import load_bolometers +from cherab.core.math import sample2d, sample2d_grid, sample2d_points, AxisymmetricMapper +from cherab.tools.emitters import RadiationFunction +from cherab.tools.raytransfer import RayTransferCylinder, RayTransferPipeline0D +from cherab.tools.inversions import admt_utils as admt +from cherab.tools.inversions import invert_sparse_regularised_nnls + + +plt.ion() + +################################################################################ +# Define the emissivity profile. +################################################################################ +# The emissivity profile consists of a blob, a ring and part of a ring on the LFS. +# The blob and the ring are Gaussian flux functions. +# The ring is Gaussian in flux and poloidal angle. +# All have equal maximum emissivities, but not necessarily equal total power. +# We use Raysect's function framework to specify an analytic form for the +# emissivity profile, as this is very quick to sample and ray trace. +eq = load_equilibrium() +psin = eq.psi_normalised +axis = eq.magnetic_axis +blob_centre_psin = 0 +blob_width_psin = 0.1 +blob = Exp2D(-0.5 * (psin - blob_centre_psin)**2 / (blob_width_psin**2)) +ring_centre_psin = 0.5 +ring_width_psin = 0.05 +ring = Exp2D(-0.5 * (psin - ring_centre_psin)**2 / (ring_width_psin**2)) +theta = Atan4Q2D(Arg2D('y') - axis.y, Arg2D('x') - axis.x) +lfs_centre_psin = 0.85 +lfs_width_psin = 0.1 +lfs_centre_theta = 0 +lfs_width_theta = 0.5 +lfs = Exp2D(-0.5 * (((psin - lfs_centre_psin) / lfs_width_psin)**2 + + ((theta - lfs_centre_theta) / lfs_width_theta)**2)) +emissivity = blob + ring + lfs +# Assume no emission from these contributors outside the separatrix. +emissivity = emissivity * eq.inside_lcfs + +# Visualise the emissivity profile with the equilibrium overlayed. +plt.figure() +rsamp, zsamp, psisamp = sample2d(psin, (*eq.r_range, 500), (*eq.z_range, 1000)) +plt.contour(rsamp, zsamp, psisamp.T, linewidths=0.5, alpha=0.3, + levels=np.linspace(0, 1, 10), colors=['k']*9 + ['red']) +rsamp, zsamp, emsamp = sample2d(emissivity, (*eq.r_range, 500), (*eq.z_range, 1000)) +im = plt.imshow(emsamp.T, extent=(rsamp[0], rsamp[-1], zsamp[0], zsamp[-1]), cmap='Purples') +plt.xlabel("R[m]") +plt.ylabel("Z[m]") +plt.colorbar(im, label="Model emissivity [W/m3]") +plt.xlim([rsamp[0], rsamp[-1]]) +plt.ylim([zsamp[0], zsamp[-1]]) +plt.gca().set_aspect('equal') +plt.pause(0.5) + + +################################################################################ +# Load the machine wall and diagnostic. +################################################################################ +print("Loading the geometry...") +world = World() +load_first_wall(world, material=AbsorbingSurface()) +bolos = load_bolometers(world) +# Only consider the purely-poloidal cameras for now... +poloidal_bolos = bolos[:3] +tangential_bolos = bolos[3:] # Includes midplane and divertor tangential + +######################################################################## +# Produce a voxel grid +######################################################################## +print("Producing the voxel grid...") +# Define the centres of each voxel, as an (nr, nz, 2) array. +nr = 40 +nz = 85 +cell_r, cell_dr = np.linspace(0.7, 2.5, nr, retstep=True) +cell_z, cell_dz = np.linspace(-1.8, 1.6, nz, retstep=True) +cell_r_grid, cell_z_grid = np.broadcast_arrays(cell_r[:, None], cell_z[None, :]) +cell_centres = np.stack((cell_r_grid, cell_z_grid), axis=-1) # (nr, nz, 2) array + +# Define the positions of the vertices of the voxels. +cell_vertices_r = np.linspace(cell_r[0] - 0.5 * cell_dr, cell_r[-1] + 0.5 * cell_dr, nr + 1) +cell_vertices_z = np.linspace(cell_z[0] - 0.5 * cell_dz, cell_z[-1] + 0.5 * cell_dz, nz + 1) + +# Build a mask, only including cells within the wall. +mask_2d = sample2d_grid(eq.inside_limiter, cell_r, cell_z) +mask_3d = mask_2d[:, np.newaxis, :] +ncells = int(mask_3d.sum()) + +# We'll use the Ray Transfer frameworks as these voxels are rectangular +# and it's much faster than the Voxel framework for simple cases like this. +ray_transfer_grid = RayTransferCylinder( + radius_outer=cell_vertices_r[-1], + radius_inner=cell_vertices_r[0], + height=cell_vertices_z[-1] - cell_vertices_z[0], + n_radius=nr, n_height=nz, mask=mask_3d, n_polar=1, + transform=translate(0, 0, cell_vertices_z[0]), +) + +######################################################################## +# Calculate the geometry matrix for the grid +######################################################################## +print("Calculating the geometry matrix...") +# The ray transfer object must be in the same world as the bolometers +ray_transfer_grid.parent = world + +sensitivity_matrix = [] +for camera in poloidal_bolos: + for foil in camera: + # Temporarily override foil pipelines for the sensitivity calculation. + orig_pipelines = foil.pipelines + foil.pipelines = [RayTransferPipeline0D(kind=foil.units)] + # All objects in world have wavelength-independent material properties, + # so it doesn't matter which wavelength range we use (as long as + # max_wavelength - min_wavelength = 1) + foil.min_wavelength = 1 + foil.max_wavelength = 2 + foil.spectral_bins = ray_transfer_grid.bins + foil.observe() + sensitivity_matrix.append(foil.pipelines[0].matrix) + # Restore original pipelines for subsequent observe calls. + foil.pipelines = orig_pipelines +sensitivity_matrix = np.asarray(sensitivity_matrix) + +# Remove the ray transfer object from the world so it doesn't interfere with +# later observations. +ray_transfer_grid.parent = None + + +################################################################################ +# Generate the regularisation operators. +################################################################################ +print("Generating regularisation operators...") +# Generating the derivative operators requires two mappings, one from a flat +# list of voxels to the original 2D grid, and one for the 2D grid coordinates to +# the flat list of voxels. Since these are the inverse of one another then one +# can be computed from the other, and therefore we only need to provide one of +# the mappings. We could build these by hand - and in the general case they must +# be built by hand - but the RayTransferCylinder object we're using helpfully +# provides the data already so we just need to convert from arrays to +# dictionaries. The easist of these to convert is the inverse voxel map as it +# already excludes masked elements from the original regular grid to leave only +# the voxels actually used in the inversion. +grid_index_1d_to_2d_map = {} +for k, (ir, iphi, iz) in enumerate(ray_transfer_grid.invert_voxel_map()): + # We want the r and z elements, as the Ray Transfer grid is 3D and this + # inversion is going to be in 2D. + grid_index_1d_to_2d_map[k] = (ir.item(), iz.item()) + +# We now need an (Nx4x2) array of voxel vertices, which can be easily calculated. +voxel_centres = np.array([cell_centres[grid_index_1d_to_2d_map[i]] + for i in range(ray_transfer_grid.bins)]) +vertex_displacements = np.array([[-cell_dr/2, -cell_dz/2], + [-cell_dr/2, cell_dz/2], + [cell_dr/2, cell_dz/2], + [cell_dr/2, -cell_dz/2]]) +# Combine the (N,2) and (4,2) arrays to get an (N,4,2) array. +voxel_vertices = voxel_centres[:, None, :] + vertex_displacements[None, :, :] +# The derivative operators are (ncells x ncells) matrices which are sparse. We +# have quite a lot of cells (around 2100), so it's more efficient to generate +# and use sparse matrices here, though dense ones will be returned by default +# for backwards compatibility. +sparse = True +derivative_operators = admt.generate_derivative_operators( + voxel_vertices, grid_index_1d_to_2d_map, sparse=True, +) + +# As described in the docstring for generate_derivative_operators, we can +# calculate a 2D laplacian operator for "isotropic" smoothing easily: +alpha = 1/3 # Optimal isotropy +aligned = derivative_operators['Dxx'] * cell_dr**2 + derivative_operators['Dyy'] * cell_dz**2 +skewed = (derivative_operators['Dsp'] + derivative_operators['Dsm']) * (cell_dr**2 + cell_dz**2) +laplacian = (1 - alpha) * aligned + (alpha / 2) * skewed +# We could also use alpha = 2/3, which would produce an operator akin to the one +# used in Carr et. al. RSI 89, 083506 (2018). + +# We can also derive an anistoropic regularisation operator, which calculates the +# amount of un-smoothness parallel and perpendicular to the magnetic field lines. +# For this we need the radii of the voxels and the magnetic flux at each voxel, +# along with a few other inputs. +voxel_radii = voxel_centres[:, 0] +psi_at_voxels = sample2d_points(eq.psi_normalised, voxel_centres) +# We also need to decide on the degree of anisotropy we expect, i.e. how much more +# smooth the radiation is along the field lines vs perpendicular to them. +# The optimal value will depend on the problem at hand. +anisotropy = 50 +admt_operator = admt.calculate_admt( + voxel_radii, derivative_operators, psi_at_voxels, cell_dr, cell_dz, anisotropy +) + +################################################################################ +# Forward model the measurements. +################################################################################ +print("Modelling the measurement values...") +# Create an emitting object whose emission is defined by the analytic form we +# produced earlier. As the emission depends on the equilibrium, this object +# should have an extent no larger than the equilibrium reconstruction extent. +# We actually make the emitter slightly smaller than the equilibrium region to +# avoid numerical precision issues creating attempts to calculate the emissivity +# outside of the equlibrium domain. +CYLINDER_RADIUS = eq.r_range[-1] - 1e-6 +CYLINDER_HEIGHT = eq.z_range[-1] - eq.z_range[0] - 2e-6 +CYLINDER_SHIFT = eq.z_range[0] + 1e-6 +emitter = Cylinder(radius=CYLINDER_RADIUS, height=CYLINDER_HEIGHT, + transform=translate(0, 0, CYLINDER_SHIFT)) +# Cut out middle of cylinder as well: equilibrium not defined here. +emitter = Subtract(emitter, Cylinder(radius=eq.r_range[0] + 1e-6, height=10, + transform=translate(0, 0, -5))) +emission_function_3d = AxisymmetricMapper(emissivity) +emitting_material = VolumeTransform(RadiationFunction(emission_function_3d), + transform=emitter.transform.inverse()) +emitter.material = emitting_material +emitter.parent = world + +# Calculate the line-integral bolometer measurements by observing the emitter +# with all bolometers. The measurements should have the same channel order as +# the sensitivity matrix. +all_measurements = [] +for camera in poloidal_bolos: + all_measurements.extend(camera.observe()) + + +################################################################################ +# Perform the inversions. +################################################################################ +print("Performing inversions...") +# We'll use NNLS with regularisation. Since the number of voxels is reasonably +# large (around 2100), we'll use the sparse variant of the NNLS inversion for +# memory and computational efficiency. The hyperparameters have been chosen by +# hand but techniques such as the discrepancy principle or L curve optimisation +# could also be used to determine them. That is out of the scope of this demo. +isotropic_alpha = 1e-10 +isotropic_inversion, _ = invert_sparse_regularised_nnls( + sensitivity_matrix, all_measurements, alpha=isotropic_alpha, + tikhonov_matrix=laplacian, +) + +admt_alpha = 1e-10 +admt_inversion, _ = invert_sparse_regularised_nnls( + sensitivity_matrix, all_measurements, alpha=admt_alpha, + tikhonov_matrix=admt_operator, +) + + +################################################################################ +# Plot the inversion results. +################################################################################ +emiss2d = np.zeros((nr, nz)) + +# Isotropic +for index1d, indices2d in grid_index_1d_to_2d_map.items(): + emiss2d[indices2d] = isotropic_inversion[index1d] +emiss2d *= 4 * np.pi +plt.figure() +im = plt.imshow(emiss2d.T, extent=(cell_r[0], cell_r[-1], cell_z[0], cell_z[-1]), cmap='Purples') +plt.contour(rsamp, zsamp, psisamp.T, linewidths=0.5, alpha=0.3, + levels=np.linspace(0, 1, 10), colors=['k']*9 + ['red']) +plt.xlabel("R[m]") +plt.ylabel("Z[m]") +plt.colorbar(im, label="Inverted\nEmissivity [W/m3]") +plt.xlim([rsamp[0], rsamp[-1]]) +plt.ylim([zsamp[0], zsamp[-1]]) +plt.gca().set_aspect('equal') +plt.title("Isotropic regularisation,\npoloidal channels") + +# Anisotropic. +for index1d, indices2d in grid_index_1d_to_2d_map.items(): + emiss2d[indices2d] = admt_inversion[index1d] +emiss2d *= 4 * np.pi +plt.figure() +im = plt.imshow(emiss2d.T, extent=(cell_r[0], cell_r[-1], cell_z[0], cell_z[-1]), cmap='Purples') +plt.contour(rsamp, zsamp, psisamp.T, linewidths=0.5, alpha=0.3, + levels=np.linspace(0, 1, 10), colors=['k']*9 + ['red']) +plt.xlabel("R[m]") +plt.ylabel("Z[m]") +plt.colorbar(im, label="Inverted\nEmissivity [W/m3]") +plt.xlim([rsamp[0], rsamp[-1]]) +plt.ylim([zsamp[0], zsamp[-1]]) +plt.gca().set_aspect('equal') +plt.title("Anisotropic regularisation\npoloidal channels") + +plt.pause(0.5) + + +######################################################################## +# Can we get a better inversion by including tangential information? +######################################################################## +print("Augmenting the geometry matrix with tangential bolos...") +# The ray transfer object must be in the same world as the bolometers, +# and the plasma emitter must be absent. +ray_transfer_grid.parent = world +emitter.parent = None + + +# sensitivity_matrix = [] +sensitivity_matrix = sensitivity_matrix.tolist() +for camera in tangential_bolos: + for foil in camera: + # Temporarily override foil pipelines for the sensitivity calculation. + orig_pipelines = foil.pipelines + foil.pipelines = [RayTransferPipeline0D(kind=foil.units)] + # All objects in world have wavelength-independent material properties, + # so it doesn't matter which wavelength range we use (as long as + # max_wavelength - min_wavelength = 1) + foil.min_wavelength = 1 + foil.max_wavelength = 2 + foil.spectral_bins = ray_transfer_grid.bins + foil.observe() + sensitivity_matrix.append(foil.pipelines[0].matrix) + # Restore original pipelines for subsequent observe calls. + foil.pipelines = orig_pipelines +sensitivity_matrix = np.asarray(sensitivity_matrix) + +ray_transfer_grid.parent = None + + +print("Adding tangential bolometer measurements...") +emitter.parent = world +for camera in tangential_bolos: + all_measurements.extend(camera.observe()) + + +print("Performing new inversions...") +isotropic_inversion, _ = invert_sparse_regularised_nnls( + sensitivity_matrix, all_measurements, alpha=isotropic_alpha, + tikhonov_matrix=laplacian, +) + +admt_inversion, _ = invert_sparse_regularised_nnls( + sensitivity_matrix, all_measurements, alpha=admt_alpha, + tikhonov_matrix=admt_operator, +) + +print("Plotting results...") +emiss2d = np.zeros((nr, nz)) + +# Isotropic +for index1d, indices2d in grid_index_1d_to_2d_map.items(): + emiss2d[indices2d] = isotropic_inversion[index1d] +emiss2d *= 4 * np.pi +plt.figure() +im = plt.imshow(emiss2d.T, extent=(cell_r[0], cell_r[-1], cell_z[0], cell_z[-1]), cmap='Purples') +plt.contour(rsamp, zsamp, psisamp.T, linewidths=0.5, alpha=0.3, + levels=np.linspace(0, 1, 10), colors=['k']*9 + ['red']) +plt.xlabel("R[m]") +plt.ylabel("Z[m]") +plt.colorbar(im, label="Inverted\nEmissivity [W/m3]") +plt.xlim([rsamp[0], rsamp[-1]]) +plt.ylim([zsamp[0], zsamp[-1]]) +plt.gca().set_aspect('equal') +plt.title("Isotropic regularisation,\nall channels") + +# Anisotropic. +for index1d, indices2d in grid_index_1d_to_2d_map.items(): + emiss2d[indices2d] = admt_inversion[index1d] +emiss2d *= 4 * np.pi +plt.figure() +im = plt.imshow(emiss2d.T, extent=(cell_r[0], cell_r[-1], cell_z[0], cell_z[-1]), cmap='Purples') +plt.contour(rsamp, zsamp, psisamp.T, linewidths=0.5, alpha=0.3, + levels=np.linspace(0, 1, 10), colors=['k']*9 + ['red']) +plt.xlabel("R[m]") +plt.ylabel("Z[m]") +plt.colorbar(im, label="Inverted\nEmissivity [W/m3]") +plt.xlim([rsamp[0], rsamp[-1]]) +plt.ylim([zsamp[0], zsamp[-1]]) +plt.gca().set_aspect('equal') +plt.title("Anisotropic regularisation\nall channels") + +plt.ioff() +plt.show() diff --git a/demos/particle_distribution/generic_distribution.py b/demos/particle_distribution/generic_distribution.py new file mode 100644 index 00000000..de8f55ce --- /dev/null +++ b/demos/particle_distribution/generic_distribution.py @@ -0,0 +1,298 @@ +from scipy.constants import atomic_mass, elementary_charge, pi + +import numpy as np + +import matplotlib.pyplot as plt + +from raysect.core.math.function.float import Arg3D, Exp3D, Sqrt3D +from raysect.core.math.function.vector3d import FloatToVector3DFunction3D + +from cherab.core.atomic import deuterium +from cherab.core.distribution import GenericDistribution +from cherab.core.math.function.float import Arg6D, Exp6D, Sqrt6D + + +# To set up a generic distribution, we need to define the following: +# - define 3D scalar function defining the spatial distribution of the effective temperature +# - define 3D vector function defining the spatial distribution of the bulk velocity +# - define 3D scalar function defining the spatial distribution of the density +# - define 6D scalar function defining the phase space density + +# This example creates a toroidally symmetric distribution in R-Z coordinates +# where R = sqrt(X^2 + Y^2) and Z is the vertical coordinate. +# The distribution peaks at R=2, Z=0 with Gaussian-like profiles. + +# initialise the spatial arguments for the 3D functions +x3d, y3d, z3d = Arg3D("x"), Arg3D("y"), Arg3D("z") + +# Calculate R = sqrt(X^2 + Y^2) for 3D functions +r_3d = Sqrt3D(x3d**2 + y3d**2) + +# Peak location in R-Z space +r_peak = 2.0 # meters +z_peak = 0.0 # meters + +# set the properties of the temperature profile +maximum_temperature = 1000 # eV +temperature_peak_width_R = 0.5 # meters +temperature_peak_width_Z = 0.5 # meters + +# set up a 3D gaussian-like temperature profile in R-Z +# The temperature is defined only as a 3D function, +# and is not used within the phase space density function. +# Consistency with the phase space density function is not checked +# and is the responsibility of the user, if required. +temperature_3d = maximum_temperature * Exp3D( + -0.5 + * ( + ((r_3d - r_peak) ** 2 / temperature_peak_width_R**2) + + ((z3d - z_peak) ** 2 / temperature_peak_width_Z**2) + ) +) + +# set the properties of the density profile +maximum_density = 5e19 # m^-3 +density_peak_width_r = 0.5 # meters +density_peak_width_z = 0.5 # meters + +# set up the 3D function defining the spatial density in R-Z +# The density is defined only as a 3D function, +# and is not used within the phase space density function. +# Consistency with the phase space density function is not checked +# and is the responsibility of the user, if required. +density_3d = maximum_density * Exp3D( + -0.5 + * ( + ((r_3d - r_peak) ** 2 / density_peak_width_r**2) + + ((z3d - z_peak) ** 2 / density_peak_width_z**2) + ) +) + +# set the properties of the toroidal rotation velocity profile +# The bulk velocity is defined only as a 3D vector function, +# and is not used within the phase space density function. +# Consistency with the phase space density function is not checked +# and is the responsibility of the user, if required. +maximum_toroidal_velocity = 1e5 # m/s +toroidal_velocity_peak_width_R = 0.5 # meters +toroidal_velocity_peak_width_Z = 0.5 # meters + +# Toroidal velocity profile in R-Z (Gaussian-like) +# The toroidal direction is perpendicular to R and Z +# In Cartesian: v_toroidal * (-y/R, x/R, 0) +v_toroidal_profile_3d = maximum_toroidal_velocity * Exp3D( + -0.5 + * ( + ((r_3d - r_peak) ** 2 / toroidal_velocity_peak_width_R**2) + + ((z3d - z_peak) ** 2 / toroidal_velocity_peak_width_Z**2) + ) +) + +# Convert toroidal velocity to Cartesian components +# vx = -v_toroidal * y/R, vy = v_toroidal * x/R, vz = 0 +# Note: We need to handle the case where R=0, but for this example we assume R>0 +vx_profile = -v_toroidal_profile_3d * y3d / (r_3d + 1e-10) # small epsilon to avoid division by zero +vy_profile = v_toroidal_profile_3d * x3d / (r_3d + 1e-10) +vz_profile = 0.0 +bulk_velocity_profile = FloatToVector3DFunction3D(vx_profile, vy_profile, vz_profile) + +# initialise the arguments for the 6D function +x6d, y6d, z6d, vx6d, vy6d, vz6d = ( + Arg6D("x"), + Arg6D("y"), + Arg6D("z"), + Arg6D("u"), + Arg6D("v"), + Arg6D("w"), +) + +# Calculate R = sqrt(X^2 + Y^2) for 6D functions +r_6d = Sqrt6D(x6d**2 + y6d**2) + +# set the missing parameters of the distribution function +deuterium_mass = deuterium.atomic_weight * atomic_mass + +# re-define the spatial temperature profile with the 6D function parameters +te_6d = maximum_temperature * Exp6D( + -0.5 + * ( + ((r_6d - r_peak) ** 2 / temperature_peak_width_R**2) + + ((z6d - z_peak) ** 2 / temperature_peak_width_Z**2) + ) +) + +# re-define the spatial density profile with the 6D function parameters +density_6d = maximum_density * Exp6D( + -0.5 + * ( + ((r_6d - r_peak) ** 2 / density_peak_width_r**2) + + ((z6d - z_peak) ** 2 / density_peak_width_z**2) + ) +) + +# Toroidal velocity profile redefined with the 6D function parameters +v_toroidal_mean_6d = maximum_toroidal_velocity * Exp6D( + -0.5 + * ( + ((r_6d - r_peak) ** 2 / toroidal_velocity_peak_width_R**2) + + ((z6d - z_peak) ** 2 / toroidal_velocity_peak_width_Z**2) + ) +) * -1.0 + +# Convert toroidal velocity to Cartesian velocity components for the 6D function +# vx_mean = -v_toroidal * y/R, vy_mean = v_toroidal * x/R, vz_mean = 0 +vx_mean_6d = -v_toroidal_mean_6d * y6d / (r_6d + 1e-10) +vy_mean_6d = v_toroidal_mean_6d * x6d / (r_6d + 1e-10) +vz_mean_6d = 0.0 + +# define the 6D distribution function for the bulk population +factor_6d = (deuterium_mass / (2 * pi * elementary_charge * te_6d)) ** 1.5 + +thermal_exponential_6d = Exp6D( + -0.5 + * deuterium_mass + * ((vx6d - vx_mean_6d) ** 2 + (vy6d - vy_mean_6d) ** 2 + (vz6d - vz_mean_6d) ** 2) + / (elementary_charge * te_6d) +) +bulk_pdf_6d = ( + density_6d * factor_6d * thermal_exponential_6d +) # bulk particle distribution function + +# add a population of supra-thermal particles with higher toroidal rotation +# the 3D bulk velocity and temperature functions ignore this population, +# and is the responsibility of the user, if required. +supra_thermal_population_ratio = 0.005 # 10% of the particles are supra-thermal +suprathermal_toroidal_velocity_factor = 20.0 # Supra-thermal particles have 2x toroidal velocity +suprathermal_temperature = 100 # eV (higher temperature for supra-thermal particles) + +# Supra-thermal toroidal velocity profile (higher than bulk) +v_toroidal_suprathermal_6d = ( + maximum_toroidal_velocity + * suprathermal_toroidal_velocity_factor + * Exp6D( + -0.5 + * ( + ((r_6d - r_peak) ** 2 / toroidal_velocity_peak_width_R**2) + + ((z6d - z_peak) ** 2 / toroidal_velocity_peak_width_Z**2) + ) + ) +) + +# Convert supra-thermal toroidal velocity to Cartesian components +vx_suprathermal_6d = -v_toroidal_suprathermal_6d * y6d / (r_6d + 1e-10) +vy_suprathermal_6d = v_toroidal_suprathermal_6d * x6d / (r_6d + 1e-10) +vz_suprathermal_6d = 0.0 + +factor_st = (deuterium_mass / (2 * pi * elementary_charge * suprathermal_temperature)) ** 1.5 + +suprathermal_exponential_6d = Exp6D( + -0.5 + * deuterium_mass + * ( + (vx6d - vx_suprathermal_6d) ** 2 + + (vy6d - vy_suprathermal_6d) ** 2 + + (vz6d - vz_suprathermal_6d) ** 2 + ) + / (elementary_charge * suprathermal_temperature) +) +supra_thermal_pdf_6d = ( + supra_thermal_population_ratio * density_6d * factor_st * suprathermal_exponential_6d +) +phase_space_density_6d = bulk_pdf_6d + supra_thermal_pdf_6d + + +generic_distribution = GenericDistribution( + phase_space_density_6d, density_3d, temperature_3d, bulk_velocity_profile +) + +# Example evaluation: sample the distribution at a point in R-Z space +# At R=2, Z=0 (the peak), sample velocity distribution +# Convert R=2, Z=0 to Cartesian: x=2, y=0, z=0 +sample_x, sample_y, sample_z = 2.0, 0.0, 0.0 + +# Sample velocity distribution along the toroidal and vertical directions (vy and vz components) +v_vals = np.linspace(-1e6, 3e6, 1000) +n_particles_vy = np.zeros_like(v_vals) +n_particles_vz = np.zeros_like(v_vals) +for i in range(len(v_vals)): + n_particles_vy[i] = generic_distribution(sample_x, sample_y, sample_z, 0.0, v_vals[i], 0.0) + n_particles_vz[i] = generic_distribution(sample_x, sample_y, sample_z, 0.0, 0.0, v_vals[i]) + + +_, ax = plt.subplots() +ax.plot(v_vals, n_particles_vy, label="$\\mathrm{v}_\\mathrm{y}$") +ax.plot(v_vals, n_particles_vz, label="$\\mathrm{v}_\\mathrm{z}$") +ax.legend() +ax.set_xlabel("Velocity (m/s)") +ax.set_ylabel("Phase Space Density (s^3/m^6)") +ax.set_title("Velocity Distribution at R=2, Z=0 (Peak Location)") +ax.grid(True) + +# sample the temperature distribution in the R-Z plane +r_vals = np.linspace(1, 3, 100) +z_vals = np.linspace(-2, 2, 210) + +temperature_vals = np.zeros((r_vals.size, z_vals.size)) +density_vals = np.zeros((r_vals.size, z_vals.size)) +bulk_velocity_vals = np.zeros((r_vals.size, z_vals.size)) +for i in range(r_vals.size): + for j in range(z_vals.size): + temperature_vals[i, j] = generic_distribution.effective_temperature(r_vals[i], 0.0, z_vals[j]) + density_vals[i, j] = generic_distribution.density(r_vals[i], 0.0, z_vals[j]) + vector_velocity = generic_distribution.bulk_velocity(r_vals[i], 0.0, z_vals[j]) + bulk_velocity_vals[i, j] = np.sqrt(vector_velocity.x**2 + vector_velocity.y**2 + vector_velocity.z**2) + +_, ax = plt.subplots() +pcm = ax.pcolormesh(r_vals, z_vals, temperature_vals.transpose(), shading='gouraud') +plt.colorbar(pcm, ax=ax, label="Temperature (eV)") +ax.set_xlabel("R (m)") +ax.set_ylabel("Z (m)") +ax.set_title("Temperature Distribution") +ax.set_aspect('equal') +ax.grid(True) + +_, ax = plt.subplots() +pcm = ax.pcolormesh(r_vals, z_vals, density_vals.transpose(), shading='gouraud') +plt.colorbar(pcm, ax=ax, label="Density (m^-3)") +ax.set_xlabel("R (m)") +ax.set_ylabel("Z (m)") +ax.set_title("Density Distribution") +ax.set_aspect('equal') +ax.grid(True) + +_, ax = plt.subplots() +pcm = ax.pcolormesh(r_vals, z_vals, bulk_velocity_vals.transpose(), shading='gouraud') +plt.colorbar(pcm, ax=ax, label="velocity (m/s)") +ax.set_xlabel("R (m)") +ax.set_ylabel("Z (m)") +ax.set_title("Toroidal Bulk Velocity Distribution") +ax.set_aspect('equal') +ax.grid(True) + +# sample the x, y velocity components in the X-Y plane using arrow vectors +x_vals = np.linspace(-3, 3, 21) # Reduced resolution for clearer quiver plot +y_vals = np.linspace(-3, 3, 21) + +X, Y = np.meshgrid(x_vals, y_vals) +x_velocity_vals = np.zeros_like(X) +y_velocity_vals = np.zeros_like(Y) + +for i in range(x_vals.size): + for j in range(y_vals.size): + vector_velocity = generic_distribution.bulk_velocity(x_vals[i], y_vals[j], 0.0) + x_velocity_vals[j, i] = vector_velocity.x + y_velocity_vals[j, i] = vector_velocity.y + +# Calculate velocity magnitude for colormap +velocity_magnitude = np.sqrt(x_velocity_vals**2 + y_velocity_vals**2) + +_, ax = plt.subplots() +quiver = ax.quiver(X, Y, x_velocity_vals, y_velocity_vals, velocity_magnitude, + cmap='viridis', scale=1e6, width=0.003) +plt.colorbar(quiver, ax=ax, label="Velocity Magnitude (m/s)") +ax.set_xlabel("X (m)") +ax.set_ylabel("Y (m)") +ax.set_title("Bulk x, y Velocity Cmponents Field in X-Y Plane (Z=0)") +ax.set_aspect('equal') +ax.grid(True) +plt.show() \ No newline at end of file diff --git a/dev/pixi.md b/dev/pixi.md new file mode 100644 index 00000000..6b19fa1a --- /dev/null +++ b/dev/pixi.md @@ -0,0 +1,181 @@ +# Pixi developer guide + +This document describes the development environments and tasks configured in +[`pixi.toml`](https://github.com/cherab/core/blob/development/pixi.toml). Pixi +manages the development dependencies in isolated environments and provides a +common interface for running tests, building the documentation, and checking +the source tree. + +Pixi installs or updates the selected environment automatically when a command +is run. The workspace currently supports Linux x86-64, macOS x86-64, and macOS +Arm64. + +## 🔎 Discovering tasks + +List every task in the workspace with: + +```console +pixi task list +``` + +To show only the tasks available in a particular environment, pass its name: + +```console +pixi task list -e test +``` + +Use `pixi run --help` for general command help. When a task is available in +more than one environment, use `-e ` to select the environment +explicitly. + +## 🧩 Environments + +| Environment | Purpose | +| --- | --- | +| `default` | Basic development tools; Cherab is not installed | +| `test` | Run the complete test suite using the latest supported Python; currently equivalent to `test-pylatest` | +| `test-pylatest` | Run the complete test suite using the latest supported Python | +| `test-pyoldest` | Run the complete test suite using the oldest supported Python | +| `test-opencl` | Run the OpenCL SART tests with Cherab's `opencl` extra | +| `docs` | Build the documentation | +| `lint` | Run formatting and static-analysis tools without installing Cherab | + +See the [`pyoldest` and `pylatest` features and environment definitions in +`pixi.toml`](https://github.com/cherab/core/blob/development/pixi.toml#:~:text=Python%20Version%20Features) +for the Python versions used by each environment. + +## 🛠️ Basic development tasks + +Start an IPython session in the default environment: + +```console +pixi run ipython +``` + +Remove generated C/Cython libraries and HTML files from the `cherab/` source +tree: + +```console +pixi run clean +``` + +The `clean` task deletes files matching `*.c`, `*.so`, `*.pyd`, `*.dll`, and +`*.html` below `cherab/`. + +## 🧪 Testing + +Run the complete test suite with the latest supported Python: + +```console +pixi run -e test test +``` + +The `test` environment currently uses the same solve group as +`test-pylatest`. The explicit alias can also be used: + +```console +pixi run -e test-pylatest test +``` + +Run the suite with the oldest supported Python: + +```console +pixi run -e test-pyoldest test +``` + +Run the OpenCL SART tests: + +```console +pixi run -e test-opencl test-opencl +``` + +The regular test task runs `python -m unittest discover cherab -v`. The OpenCL +task runs `cherab.tools.tests.test_sart_opencl` only. + +## 📚 Documentation + +Build the HTML documentation: + +```console +pixi run -e docs doc-build +``` + +`html` is the default Sphinx builder. A different builder can be supplied as +the final argument; for example, check external and internal links with: + +```console +pixi run -e docs doc-build linkcheck +``` + +Build output is written below `docs/build/`. Remove all documentation +build output with: + +```console +pixi run doc-clean +``` + +After building the HTML documentation, serve it locally on port 8000 with: + +```console +pixi run doc-serve +``` + +Then open in a browser. To use a different port, pass +it as the final argument: + +```console +pixi run doc-serve 8080 +``` + +## 🧹 Formatting and static analysis + +The `lint` environment keeps code-quality tools separate from the environments +that build and install Cherab. Because the task names below are unique to this +environment, Pixi selects it automatically; `-e lint` is not required. + +| Task | Action | +| --- | --- | +| `lefthook` | Run Lefthook | +| `hooks` | Install the Git hooks managed by Lefthook | +| `pre-commit` | Run the Lefthook `pre-commit` group | +| `ruff-check` | Run `ruff check` | +| `ruff-format` | Run `ruff format` | +| `toml-format` | Run `tombi format` | +| `dprint` | Run `dprint fmt` | +| `typos` | Find and fix spelling errors | +| `actionlint` | Run Actionlint | +| `blacken-docs` | Format Python examples in documentation | +| `validate-pyproject` | Validate `pyproject.toml` | +| `cython-lint` | Run Cython-Lint | +| `lint` | Run the Lefthook `pre-commit` group on all files | + +For example: + +```console +pixi run ruff-check +pixi run toml-format +pixi run cython-lint +pixi run validate-pyproject +``` + +The `hooks`, `lefthook`, `pre-commit`, and aggregate `lint` tasks invoke +Lefthook. + +> [!WARNING] +> A Lefthook configuration file has not been added to the repository yet, so +> these tasks are not currently available. + +Install the Git hooks and run all configured checks with: + +```console +pixi run hooks +pixi run lint +``` + +Running `pixi run hooks` installs the Git hooks once. After installation, the +configured pre-commit checks are triggered automatically for every commit. To +remove the installed hooks and stop the automatic checks, run: + +```console +pixi run lefthook uninstall +``` diff --git a/dev/test.sh b/dev/test.sh index 6132427b..edfbc452 100755 --- a/dev/test.sh +++ b/dev/test.sh @@ -1,3 +1,3 @@ #!/bin/bash -python -m unittest $1 $2 $3 $4 $5 +python -m unittest discover cherab $1 $2 $3 $4 $5 diff --git a/docs/source/available_modules.rst b/docs/source/available_modules.rst index 9be25f1e..88957276 100644 --- a/docs/source/available_modules.rst +++ b/docs/source/available_modules.rst @@ -34,9 +34,8 @@ Fusion Experiment Packages - The Cherab configuration package for AUG. * - `cherab-compass `_ - The Cherab configuration package for COMPASS. - * - cherab-iter - - Integrates Cherab with IMAS and provides diagnostic configuration - for ITER. This package is under development but not yet publicly available. + * - `cherab-iter `_ + - The Cherab configuration package for ITER. * - `cherab-jet `_ - Experiment configuration package for JET. * - `cherab-mastu `_ @@ -59,4 +58,5 @@ and workflow management tools. - Module for providing OMFIT integration and example workflow scripts. * - `cherab-solps `_ - Allows loading of Cherab plasma objects from saved SOLPS simulations. - + * - `cherab-imas `_ + - Provides Cherab integration with IMAS (Integrated Modelling & Analysis Suite). diff --git a/docs/source/conf.py b/docs/source/conf.py index b2f22c70..54bb264f 100644 --- a/docs/source/conf.py +++ b/docs/source/conf.py @@ -40,6 +40,7 @@ 'sphinx.ext.autodoc', 'sphinx.ext.mathjax', 'sphinx_tabs.tabs', + 'myst_parser', ] # Add any paths that contain templates here, relative to this directory. @@ -56,16 +57,16 @@ # General information about the project. project = 'Cherab' -copyright = '2024, Cherab Team' +copyright = '2026, Cherab Team' # The version info for the project you're documenting, acts as replacement for # |version| and |release|, also used in various other places throughout the # built documents. # # The short X.Y version. -version = '1.5' +version = '1.6' # The full version, including alpha/beta/rc tags. -release = '1.5.0' +release = '1.6.0rc1' # The language for content autogenerated by Sphinx. Refer to documentation # for a list of supported languages. diff --git a/docs/source/development/pixi.rst b/docs/source/development/pixi.rst new file mode 100644 index 00000000..0be6bd02 --- /dev/null +++ b/docs/source/development/pixi.rst @@ -0,0 +1,2 @@ +.. include:: ../../../dev/pixi.md + :parser: myst_parser.sphinx_ diff --git a/docs/source/index.rst b/docs/source/index.rst index 3e3e90ce..72f6cda6 100644 --- a/docs/source/index.rst +++ b/docs/source/index.rst @@ -28,6 +28,14 @@ become stable until we have finished moving the source code to github. tools/tools +.. toctree:: + :maxdepth: 2 + :caption: Development + :name: development + + development/pixi + + .. toctree:: :maxdepth: 2 :caption: Demonstrations @@ -41,4 +49,3 @@ Indices and tables * :ref:`genindex` * :ref:`modindex` - diff --git a/docs/source/math/function.rst b/docs/source/math/function.rst index 6f1680e1..8b29a66e 100644 --- a/docs/source/math/function.rst +++ b/docs/source/math/function.rst @@ -10,6 +10,55 @@ documentation and the Cherab function tutorials. Cherab previously provided vector functions which were not present in Raysect. New codes should prefer the Raysect vector functions, but the old aliases are preserved for backwards compatibility. +The Function6D framework in Cherab aims to provide a framework for building six-dimensional distribution functions. +It follows closely Raysect's function framework. + +6D Scalar Functions +------------------- + +.. autoclass:: cherab.core.math.function.float.function6d.base.Function6D + :members: + :special-members: __call__ + +.. autoclass:: cherab.core.math.function.float.function6d.constant.Constant6D + :show-inheritance: + +.. autoclass:: cherab.core.math.function.float.function6d.arg.Arg6D + :show-inheritance: + +.. autoclass:: cherab.core.math.function.float.function6d.cmath.Exp6D + :show-inheritance: + +.. autoclass:: cherab.core.math.function.float.function6d.cmath.Sin6D + :show-inheritance: + +.. autoclass:: cherab.core.math.function.float.function6d.cmath.Cos6D + :show-inheritance: + +.. autoclass:: cherab.core.math.function.float.function6d.cmath.Tan6D + :show-inheritance: + +.. autoclass:: cherab.core.math.function.float.function6d.cmath.Asin6D + :show-inheritance: + +.. autoclass:: cherab.core.math.function.float.function6d.cmath.Acos6D + :show-inheritance: + +.. autoclass:: cherab.core.math.function.float.function6d.cmath.Atan6D + :show-inheritance: + +.. autoclass:: cherab.core.math.function.float.function6d.cmath.Atan4Q6D + :show-inheritance: + +.. autoclass:: cherab.core.math.function.float.function6d.cmath.Sqrt6D + :show-inheritance: + +.. autoclass:: cherab.core.math.function.float.function6d.cmath.Erf6D + :show-inheritance: + +.. autoclass:: cherab.core.math.function.float.function6d.blend.Blend6D + :show-inheritance: + 2D Vector Functions ------------------- diff --git a/docs/source/math/integrators.rst b/docs/source/math/integrators.rst new file mode 100644 index 00000000..f5753d1d --- /dev/null +++ b/docs/source/math/integrators.rst @@ -0,0 +1,11 @@ + +Integrators +------------- + +.. autoclass:: cherab.core.math.integrators.integrators1d.GaussianQuadrature1D + :members: + :show-inheritance: + +.. autoclass:: cherab.core.math.integrators.integrators2d.GaussianQuadrature2D + :members: + :show-inheritance: \ No newline at end of file diff --git a/docs/source/math/math.rst b/docs/source/math/math.rst index 836cbe84..e0a85775 100644 --- a/docs/source/math/math.rst +++ b/docs/source/math/math.rst @@ -22,3 +22,4 @@ utilities that Cherab provides for slicing, dicing and projecting these function mask samplers slice + integrators diff --git a/docs/source/plasmas/core_plasma_classes.rst b/docs/source/plasmas/core_plasma_classes.rst index a45a141d..46cd7e64 100644 --- a/docs/source/plasmas/core_plasma_classes.rst +++ b/docs/source/plasmas/core_plasma_classes.rst @@ -30,3 +30,13 @@ Distribution functions :special-members: __call__ :show-inheritance: +.. autoclass:: cherab.core.distribution.GenericDistribution + :members: + :special-members: __call__ + :show-inheritance: + +.. autoclass:: cherab.core.distribution.ZeroDistribution + :members: + :special-members: __call__ + :show-inheritance: + diff --git a/docs/source/tools/observers.rst b/docs/source/tools/observers.rst index e93e2057..b3d5e612 100644 --- a/docs/source/tools/observers.rst +++ b/docs/source/tools/observers.rst @@ -127,6 +127,9 @@ in the group. .. autoclass:: cherab.tools.observers.group.PixelGroup :members: +.. autoclass:: cherab.tools.observers.group.TargetedPixelGroup + :members: + .. autoclass:: cherab.tools.observers.group.TargettedPixelGroup :members: diff --git a/docs/source/tools/tomography.rst b/docs/source/tools/tomography.rst index 3f3984f6..f82ccccc 100644 --- a/docs/source/tools/tomography.rst +++ b/docs/source/tools/tomography.rst @@ -43,6 +43,8 @@ Inversion Methods .. autofunction:: cherab.tools.inversions.nnls.invert_regularised_nnls +.. autofunction:: cherab.tools.inversions.nnls.invert_sparse_regularised_nnls + .. autofunction:: cherab.tools.inversions.svd.invert_svd @@ -119,3 +121,36 @@ Use spectral pipelines from Raysect if you need these features. .. autoclass:: cherab.tools.raytransfer.pipelines.RayTransferPipeline1D .. autoclass:: cherab.tools.raytransfer.pipelines.RayTransferPipeline2D + + +Regularisation +-------------- + +Some of the inversion methods take a regularisation operator, which provides +additional constraints to help achieve unique solutions to ill-posed +tomography problems. Many regularisation schemes impose constraints on the smoothness +of the resulting solution, with this smoothness quantified by the second derivative +of the solution. Two such regularisation schemes are common in fusion applications: + +#. Isotropic smoothing, where the solution has the same smoothness in all directions. +#. Anisotropic smoothing, so-called "anisotropic diffusion model tomography" (ADMT), + where the solution is smoother parallel to the magnetic field and less smooth + perpendicular to the magnetic field. + +Cherab provides some utility functions to assist in calculating appropriate +operators using these (and other) derivative-based regularisation schemes. These can be used +on inversion grids defined using both the Voxel and Ray Transfer frameworks, and passed +directly to the inversion methods in Cherab which take regularisation operators, such as +cherab.tools.inversions.invert_constrained_sart and cherab.tools.inversions.invert_regularised_nnls. + +The routines to calculate derivative operators for inversion grids, and further to calculate +the ADMT operator for a given set of derivative operators and magnetic field, are taken from +work published by L. C. Ingesson in `JET-R(99)08`_. + + +.. autofunction:: cherab.tools.inversions.admt_utils.generate_derivative_operators + +.. autofunction:: cherab.tools.inversions.admt_utils.calculate_admt + + +.. _JET-R(99)08: http://www.euro-fusionscipub.org/wp-content/uploads/2014/11/JETR99008.pdf diff --git a/docs/source/welcome.rst b/docs/source/welcome.rst index 4e4a00e1..0317f954 100644 --- a/docs/source/welcome.rst +++ b/docs/source/welcome.rst @@ -20,10 +20,11 @@ The following authors have contributed to the project: Current Development Team ------------------------ -* Matthew Carr (Core Developer) -* Alex Meakins (Architect/Core Developer) -* Alfonso Baciero (Model development) -* Carine Giroud (JET Project Management) +* Jack Lovell (Oak Ridge National Laboratory, USA) +* Vlad Neverov +* Matej Tomes (Institute of Plasma Physics of the Czech Academy of Sciences, Czechia) +* Koyo Munechika (ITER Organisation) +* Jakub Svoboda (Institute of Plasma Physics of the Czech Academy of Sciences, Czechia) Contributors @@ -37,6 +38,14 @@ Contributors * Andy Meigs (Physics) +Former Developers +----------------------------------- +* Matthew Carr (Core Developer) +* Alex Meakins (Architect/Core Developer) +* Alfonso Baciero (Model development) +* Carine Giroud (JET Project Management) + + Project History --------------- diff --git a/pixi.toml b/pixi.toml new file mode 100644 index 00000000..64d3c551 --- /dev/null +++ b/pixi.toml @@ -0,0 +1,167 @@ +[workspace] +channels = ["conda-forge"] +platforms = ["linux-64", "osx-arm64", "osx-64"] +preview = ["pixi-build"] + +[workspace.build-variants] +python = [ + "3.9.*", + "3.10.*", + "3.11.*", + "3.12.*", + "3.13.*", + "3.14.*", +] + +# ------------------------------- +# === Packaging Configuration === +# ------------------------------- +[package] +name = "cherab" +version = "dynamic" + +[package.build.backend] +name = "pixi-build-python" +version = "*" + +[package.build.config] +compilers = ["c"] + +[package.host-dependencies] +python = "*" +setuptools = "*" +cython = ">=3.1" +numpy = "*" +raysect = "0.9.*" + +[package.run-dependencies] +scipy = "*" +matplotlib-base = "*" + +[package.extra-dependencies.opencl] +pyopencl = "*" +pocl = "*" + +# --------------------------------- +# === Development Configuration === +# --------------------------------- +[dependencies] +ipython = "*" + +[dev] +cherab = { path = "." } + +[tasks] +clean = { + cmd = "find cherab -type f \\( -name '*.c' -o -name '*.so' -o -name '*.pyd' -o -name '*.dll' -o -name '*.html' \\) -delete", + description = "🔥 Remove in-place build artifacts and temporary files (*.c, *.so, *.pyd, *.dll, *.html)", +} + +# The documentation-related tasks below do not require the source package. +doc-clean = { + cmd = "rm -rf build", + cwd = "docs", + description = "🔥 Clean the docs build directory", +} +doc-serve = { + cmd = [ + "python", + "-m", + "http.server", + "{{ port }}", + "--directory", + "build/html", + ], + cwd = "docs", + args = [ + { arg = "port", default = "8000" }, + ], + description = "🚀 Start a local server for the docs", +} + +# === Testing feature === +[feature.test.dependencies] +cherab = { path = "." } + +[feature.test.tasks] +test = { + cmd = "python -m unittest discover cherab -v", + description = "🧪 Run the tests", +} + +[feature.test-opencl.dependencies] +cherab = { path = ".", extras = ["opencl"] } + +[feature.test-opencl.tasks] +test-opencl = { + cmd = "python -m unittest cherab.tools.tests.test_sart_opencl -v", + description = "🧪 Run the OpenCL tests", +} + +# === Documentation feature === +[feature.docs.dependencies] +cherab = { path = "." } +sphinx = "*" +sphinx_rtd_theme = "<1" # TODO: change to >=1.0 when our docs layout is compatible with the new theme +sphinx-tabs = "*" +myst-parser = "*" + +[feature.docs.tasks] +doc-build = { + cmd = [ + "sphinx-build", + "-b", + "{{ target }}", + "source", + "build/{{ target }}", + ], + cwd = "docs", + args = [ + { arg = "target", default = "html" }, + ], + description = "📝 Build the docs" +} + +# === Linting feature === +[feature.lint.dependencies] +dprint = "*" +lefthook = "*" +ruff = "*" +typos = "*" +actionlint = "*" +shellcheck = "*" +validate-pyproject = "*" +cython-lint = "*" +blacken-docs = "*" +tombi = "*" + +[feature.lint.tasks] +lefthook = { cmd = "lefthook", description = "🔗 Run lefthook" } +hooks = { cmd = "lefthook install", description = "🔗 Install pre-commit hooks" } +pre-commit = { cmd = "lefthook run pre-commit", description = "🔗 Run pre-commit checks" } +ruff-check = { cmd = "ruff check", description = "Lint with ruff" } +ruff-format = { cmd = "ruff format", description = "Format with ruff" } +toml-format = { cmd = "tombi format", description = "Format TOML files" } +dprint = { cmd = "dprint fmt", description = "Format with dprint" } +typos = { cmd = "typos --write-changes --force-exclude", description = "Fix typos" } +actionlint = { cmd = "actionlint", description = "Lint actions with actionlint" } +blacken-docs = { cmd = "blacken-docs", description = "Format Python markdown blocks with Black" } +validate-pyproject = { cmd = "validate-pyproject pyproject.toml", description = "Validate pyproject.toml" } +cython-lint = { cmd = "cython-lint", description = "Lint Cython files" } +lint = { cmd = "lefthook run pre-commit --all-files --force", description = "🧹 Run all linters" } + +# === Python Version Features === +[feature.pyoldest.dependencies] +python = "3.9.*" + +[feature.pylatest.dependencies] +python = "3.14.*" + +[environments] +default = { features = ["pylatest"], solve-group = "pylatest" } +test = { features = ["test"], solve-group = "pylatest" } +docs = { features = ["pyoldest", "docs"], solve-group = "pyoldest" } # TODO: change to pylatest when bumping RTD theme to >=1.0 +test-pylatest = { features = ["pylatest", "test"], solve-group = "pylatest" } # alias of test +test-pyoldest = { features = ["pyoldest", "test"], solve-group = "pyoldest" } +test-opencl = { features = ["test-opencl"], solve-group = "pyoldest" } +lint = { features = ["lint"], no-default-feature = true } diff --git a/pyproject.toml b/pyproject.toml index 4849f0b5..e198bb5c 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -1,3 +1,3 @@ [build-system] -requires = ["setuptools>=62.3", "oldest-supported-numpy", "cython~=3.0", "raysect==0.8.1.*"] +requires = ["setuptools>=62.3", "numpy", "cython~=3.1", "raysect==0.9.1.*"] build-backend="setuptools.build_meta" diff --git a/requirements.txt b/requirements.txt index 9a13464d..c99e710f 100644 --- a/requirements.txt +++ b/requirements.txt @@ -1,5 +1,5 @@ -cython~=3.0 -numpy>=1.14,<2.0 +cython~=3.1 +numpy>=2.0 scipy matplotlib -raysect==0.8.1.* +raysect==0.9.1.* diff --git a/setup.py b/setup.py index f11dd08f..c07eaff8 100644 --- a/setup.py +++ b/setup.py @@ -5,7 +5,7 @@ from pathlib import Path import multiprocessing import numpy -from setuptools import setup, find_packages, Extension +from setuptools import setup, find_namespace_packages, Extension from Cython.Build import cythonize multiprocessing.set_start_method('fork') @@ -96,7 +96,6 @@ name="cherab", version=version, license="EUPL 1.1", - namespace_packages=["cherab"], description="Cherab spectroscopy framework", classifiers=[ "Development Status :: 5 - Production/Stable", @@ -116,17 +115,19 @@ ), long_description=long_description, long_description_content_type="text/markdown", + # Support Python versions where Raysect wheels are available. + requires_python=">=3.9", install_requires=[ - "numpy>=1.14,<2.0", + "numpy>=2.0", "scipy", "matplotlib", - "raysect==0.8.1.*", + "raysect==0.9.1.*", ], extras_require={ # Running ./dev/build_docs.sh runs setup.py, which requires cython. - "docs": ["cython~=3.0", "sphinx", "sphinx-rtd-theme", "sphinx-tabs"], + "docs": ["cython~=3.1", "sphinx", "sphinx-rtd-theme", "sphinx-tabs"], }, - packages=find_packages(include=["cherab*"]), + packages=find_namespace_packages(include=["cherab*"]), package_data={"": [ "**/*.pyx", "**/*.pxd", # Needed to build Cython extensions. "**/*.json", "**/*.cl", "**/*.npy", "**/*.obj", # Supplementary data