From a40827be287595b890571f3673f78372cbe5c1d2 Mon Sep 17 00:00:00 2001 From: vsnever Date: Fri, 30 Aug 2024 19:44:50 +0200 Subject: [PATCH 01/91] Bump version to 1.6.0.dev1 --- cherab/core/VERSION | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/cherab/core/VERSION b/cherab/core/VERSION index bc80560f..4c39f7c4 100644 --- a/cherab/core/VERSION +++ b/cherab/core/VERSION @@ -1 +1 @@ -1.5.0 +1.6.0.dev1 From ae9c476913872c58fd13fc74c99a722f7e7b269d Mon Sep 17 00:00:00 2001 From: vsnever Date: Fri, 30 Aug 2024 19:48:43 +0200 Subject: [PATCH 02/91] Add e_field attribute to Plasma object for electric field vector. --- CHANGELOG.md | 7 +++++++ cherab/core/plasma/node.pxd | 3 +++ cherab/core/plasma/node.pyx | 21 +++++++++++++++++++++ 3 files changed, 31 insertions(+) diff --git a/CHANGELOG.md b/CHANGELOG.md index 5fe4cc06..36ef451f 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -1,6 +1,13 @@ Project Changelog ================= +Release 1.6.0 (TBD) +------------------- + +New: +* Add e_field attribute to Plasma object for electric field vector. (#465) + + Release 1.5.0 (27 Aug 2024) ------------------- 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 From 0eb6cf2211bb4f1fabb46d3db2d98fe74c3d7689 Mon Sep 17 00:00:00 2001 From: MatejTomes Date: Mon, 13 Jan 2025 15:52:19 +0100 Subject: [PATCH 03/91] Add Integrator2D base class --- CHANGELOG.md | 6 ++ .../core/math/integrators/integrators2d.pxd | 27 ++++++++ .../core/math/integrators/integrators2d.pyx | 64 +++++++++++++++++++ 3 files changed, 97 insertions(+) create mode 100644 cherab/core/math/integrators/integrators2d.pxd create mode 100644 cherab/core/math/integrators/integrators2d.pyx diff --git a/CHANGELOG.md b/CHANGELOG.md index 5fe4cc06..e3647dad 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -1,6 +1,12 @@ Project Changelog ================= +Release 1.6.0 (TBD) +------------------- + +New: +* Add Integrator2D base class for integration of two-dimensional functions. (#472) + Release 1.5.0 (27 Aug 2024) ------------------- diff --git a/cherab/core/math/integrators/integrators2d.pxd b/cherab/core/math/integrators/integrators2d.pxd new file mode 100644 index 00000000..d956e697 --- /dev/null +++ b/cherab/core/math/integrators/integrators2d.pxd @@ -0,0 +1,27 @@ +# 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 Function2D + + +cdef class Integrator2D: + + cdef: + Function2D function + + cdef double evaluate(self,double x_lower, double x_upper, double y_lower, double y_upper) except? -1e999 \ 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..9eb2ac1a --- /dev/null +++ b/cherab/core/math/integrators/integrators2d.pyx @@ -0,0 +1,64 @@ +# 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.float cimport autowrap_function2d + +from libc.math cimport INFINITY +cimport cython + + +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: int + """ + 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, double y_lower, double y_upper) except? -1e999: + + raise NotImplementedError("The evaluate() method has not been implemented.") + + def __call__(self, double x_lower, double x_upper, double y_lower, double 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 double y_lower: Lower limit of integration in the y dimension. + :param double 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) \ No newline at end of file From b03a12c1472efb8991abf2857eb648dd1bb18b34 Mon Sep 17 00:00:00 2001 From: MatejTomes Date: Fri, 17 Jan 2025 12:19:13 +0100 Subject: [PATCH 04/91] Change y integration limits to functions of x --- cherab/core/math/integrators/__init__.pxd | 1 + cherab/core/math/integrators/__init__.py | 1 + cherab/core/math/integrators/integrators2d.pxd | 4 ++-- cherab/core/math/integrators/integrators2d.pyx | 15 ++++++--------- 4 files changed, 10 insertions(+), 11 deletions(-) diff --git a/cherab/core/math/integrators/__init__.pxd b/cherab/core/math/integrators/__init__.pxd index db06ef43..c4eff464 100644 --- a/cherab/core/math/integrators/__init__.pxd +++ b/cherab/core/math/integrators/__init__.pxd @@ -17,4 +17,5 @@ # under the Licence. from cherab.core.math.integrators.integrators1d cimport Integrator1D, GaussianQuadrature +from cherab.core.math.integrators.integrators2d cimport Integrator2D diff --git a/cherab/core/math/integrators/__init__.py b/cherab/core/math/integrators/__init__.py index 86b7d58d..b1fa4516 100644 --- a/cherab/core/math/integrators/__init__.py +++ b/cherab/core/math/integrators/__init__.py @@ -17,3 +17,4 @@ # under the Licence. from .integrators1d import Integrator1D, GaussianQuadrature +from .integrators2d import Integrator2D diff --git a/cherab/core/math/integrators/integrators2d.pxd b/cherab/core/math/integrators/integrators2d.pxd index d956e697..62ba8667 100644 --- a/cherab/core/math/integrators/integrators2d.pxd +++ b/cherab/core/math/integrators/integrators2d.pxd @@ -16,7 +16,7 @@ # See the Licence for the specific language governing permissions and limitations # under the Licence. -from raysect.core.math.function.float cimport Function2D +from raysect.core.math.function.float cimport Function1D, Function2D cdef class Integrator2D: @@ -24,4 +24,4 @@ cdef class Integrator2D: cdef: Function2D function - cdef double evaluate(self,double x_lower, double x_upper, double y_lower, double y_upper) except? -1e999 \ No newline at end of file + cdef double evaluate(self,double x_lower, double x_upper, Function1D y_lower, Function1D y_upper) except? -1e999 diff --git a/cherab/core/math/integrators/integrators2d.pyx b/cherab/core/math/integrators/integrators2d.pyx index 9eb2ac1a..58627afe 100644 --- a/cherab/core/math/integrators/integrators2d.pyx +++ b/cherab/core/math/integrators/integrators2d.pyx @@ -18,10 +18,7 @@ # See the Licence for the specific language governing permissions and limitations # under the Licence. -from raysect.core.math.function.float cimport autowrap_function2d - -from libc.math cimport INFINITY -cimport cython +from raysect.core.math.function.float cimport Function1D, autowrap_function2d cdef class Integrator2D: @@ -45,20 +42,20 @@ cdef class Integrator2D: self.function = autowrap_function2d(func) - cdef double evaluate(self, double x_lower, double x_upper, double y_lower, double y_upper) except? -1e999: + 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, double y_lower, double y_upper): + 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 double y_lower: Lower limit of integration in the y dimension. - :param double y_upper: Upper limit of integration in the y 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) \ No newline at end of file + return self.evaluate(x_lower, x_upper, y_lower, y_upper) From 9bce0fafd3eb0b46df39718068596bfe1e76be23 Mon Sep 17 00:00:00 2001 From: MatejTomes Date: Mon, 3 Feb 2025 22:56:20 +0100 Subject: [PATCH 05/91] Add Function6D framework derived from Raysect's Function3D --- cherab/core/math/__init__.py | 1 + cherab/core/math/function/__init__.pxd | 2 + cherab/core/math/function/__init__.py | 2 + cherab/core/math/function/float/__init__.pxd | 32 + cherab/core/math/function/float/__init__.py | 32 + .../function/float/function6d/__init__.pxd | 37 + .../function/float/function6d/__init__.py | 36 + .../math/function/float/function6d/arg.pxd | 27 + .../math/function/float/function6d/arg.pyx | 82 ++ .../function/float/function6d/autowrap.pxd | 26 + .../function/float/function6d/autowrap.pyx | 98 +++ .../math/function/float/function6d/base.pxd | 165 ++++ .../math/function/float/function6d/base.pyx | 737 ++++++++++++++++++ .../math/function/float/function6d/blend.pxd | 25 + .../math/function/float/function6d/blend.pyx | 65 ++ .../math/function/float/function6d/cmath.pxd | 61 ++ .../math/function/float/function6d/cmath.pyx | 168 ++++ .../function/float/function6d/constant.pxd | 25 + .../function/float/function6d/constant.pyx | 49 ++ .../float/function6d/tests/__init__.py | 5 + .../float/function6d/tests/test_arg.py | 54 ++ .../float/function6d/tests/test_autowrap.py | 37 + .../float/function6d/tests/test_base.py | 649 +++++++++++++++ .../float/function6d/tests/test_cmath.py | 155 ++++ .../float/function6d/tests/test_constant.py | 35 + 25 files changed, 2605 insertions(+) create mode 100644 cherab/core/math/function/float/__init__.pxd create mode 100644 cherab/core/math/function/float/__init__.py create mode 100644 cherab/core/math/function/float/function6d/__init__.pxd create mode 100644 cherab/core/math/function/float/function6d/__init__.py create mode 100644 cherab/core/math/function/float/function6d/arg.pxd create mode 100644 cherab/core/math/function/float/function6d/arg.pyx create mode 100644 cherab/core/math/function/float/function6d/autowrap.pxd create mode 100644 cherab/core/math/function/float/function6d/autowrap.pyx create mode 100644 cherab/core/math/function/float/function6d/base.pxd create mode 100644 cherab/core/math/function/float/function6d/base.pyx create mode 100644 cherab/core/math/function/float/function6d/blend.pxd create mode 100644 cherab/core/math/function/float/function6d/blend.pyx create mode 100644 cherab/core/math/function/float/function6d/cmath.pxd create mode 100644 cherab/core/math/function/float/function6d/cmath.pyx create mode 100644 cherab/core/math/function/float/function6d/constant.pxd create mode 100644 cherab/core/math/function/float/function6d/constant.pyx create mode 100644 cherab/core/math/function/float/function6d/tests/__init__.py create mode 100644 cherab/core/math/function/float/function6d/tests/test_arg.py create mode 100644 cherab/core/math/function/float/function6d/tests/test_autowrap.py create mode 100644 cherab/core/math/function/float/function6d/tests/test_base.py create mode 100644 cherab/core/math/function/float/function6d/tests/test_cmath.py create mode 100644 cherab/core/math/function/float/function6d/tests/test_constant.py 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/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..151eea65 --- /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, W, V + +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..c4883e0e --- /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", "w", or "v". + + >>> 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 + >>> argw = Arg6D("w") + >>> argw(2, 3, 5, 7, 11, 13) + 11.0 + >>> argv = Arg6D("v") + >>> argv(2, 3, 5, 7, 11, 13) + 13.0 + + :param str argument: either "x", "y", "z", "u", "w", or "v", 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 == "w": + self._argument = W + elif argument == "v": + self._argument = V + else: + raise ValueError("The argument to Arg6D must be either 'x', 'y', 'z', 'u', 'w' or 'v'") + + cdef double evaluate(self, double x, double y, double z, double u, double w, double v) 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 == W: + return w + else: # V + return v 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..9026a0d8 --- /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 w, double v) except? -1e999: + return self.function(x, y, z, u, w, v) + + +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..4f482baa --- /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 w, double v) 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..fd653cd9 --- /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 w, double v) except? -1e999: + raise NotImplementedError("The evaluate() method has not been implemented.") + + def __call__(self, double x, double y, double z, double u, double w, double v): + """ Evaluate the function f(x, y, z, u, w, v) + + :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 w: function parameter w + :param float v: function parameter v + :rtype: float + """ + return self.evaluate(x, y, z, u, w, v) + 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 w, double v) except? -1e999: + return self._function1.evaluate(x, y, z, u, w, v) + self._function2.evaluate(x, y, z, u, w, v) + + +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 w, double v) except? -1e999: + return self._function1.evaluate(x, y, z, u, w, v) - self._function2.evaluate(x, y, z, u, w, v) + + +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 w, double v) except? -1e999: + return self._function1.evaluate(x, y, z, u, w, v) * self._function2.evaluate(x, y, z, u, w, v) + + +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 w, double v) except? -1e999: + cdef double denominator = self._function2.evaluate(x, y, z, u, w, v) + 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, w, v) / 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 w, double v) except? -1e999: + cdef double divisor = self._function2.evaluate(x, y, z, u, w, v) + 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, w, v) % 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 w, double v) except? -1e999: + cdef double base, exponent + base = self._function1.evaluate(x, y, z, u, w, v) + exponent = self._function2.evaluate(x, y, z, u, w, v) + 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 w, double v) except? -1e999: + return abs(self._function.evaluate(x, y, z, u, w, v)) + + +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 w, double v) except? -1e999: + return self._function1.evaluate(x, y, z, u, w, v) == self._function2.evaluate(x, y, z, u, w, v) + + +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 w, double v) except? -1e999: + return self._function1.evaluate(x, y, z, u, w, v) != self._function2.evaluate(x, y, z, u, w, v) + + +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 w, double v) except? -1e999: + return self._function1.evaluate(x, y, z, u, w, v) < self._function2.evaluate(x, y, z, u, w, v) + + +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 w, double v) except? -1e999: + return self._function1.evaluate(x, y, z, u, w, v) > self._function2.evaluate(x, y, z, u, w, v) + + +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 w, double v) except? -1e999: + return self._function1.evaluate(x, y, z, u, w, v) <= self._function2.evaluate(x, y, z, u, w, v) + + +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 w, double v) except? -1e999: + return self._function1.evaluate(x, y, z, u, w, v) >= self._function2.evaluate(x, y, z, u, w, v) + + +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 w, double v) except? -1e999: + return self._value + self._function.evaluate(x, y, z, u, w, v) + + +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 w, double v) except? -1e999: + return self._value - self._function.evaluate(x, y, z, u, w, v) + + +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 w, double v) except? -1e999: + return self._value * self._function.evaluate(x, y, z, u, w, v) + + +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 w, double v) except? -1e999: + cdef double denominator = self._function.evaluate(x, y, z, u, w, v) + 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 w, double v) except? -1e999: + cdef double divisor = self._function.evaluate(x, y, z, u, w, v) + 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 w, double v) except? -1e999: + return self._function.evaluate(x, y, z, u, w, v) % 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 w, double v) except? -1e999: + cdef double exponent = self._function.evaluate(x, y, z, u, w, v) + 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 w, double v) except? -1e999: + cdef double base = self._function.evaluate(x, y, z, u, w, v) + 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 w, double v) except? -1e999: + return self._value == self._function.evaluate(x, y, z, u, w, v) + + +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 w, double v) except? -1e999: + return self._value != self._function.evaluate(x, y, z, u, w, v) + + +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 w, double v) except? -1e999: + return self._value < self._function.evaluate(x, y, z, u, w, v) + + +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 w, double v) except? -1e999: + return self._value > self._function.evaluate(x, y, z, u, w, v) + + +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 w, double v) except? -1e999: + return self._value <= self._function.evaluate(x, y, z, u, w, v) + + +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 w, double v) except? -1e999: + return self._value >= self._function.evaluate(x, y, z, u, w, v) 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..6630bbb2 --- /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)) f_1(x) + f_m(x) f_2(x) + + 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 w, double v) except? -1e999: + + cdef double t = clamp(self._mask.evaluate(x, y, z, u, w, v), 0.0, 1.0) + + # sample endpoints directly + if t == 0: + return self._f1.evaluate(x, y, z, u, w, v) + + if t == 1: + return self._f2.evaluate(x, y, z, u, w, v) + + # lerp between function values + cdef double f1 = self._f1.evaluate(x, y, z, u, w, v) + cdef double f2 = self._f2.evaluate(x, y, z, u, w, v) + 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..ccc5d9b6 --- /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 w, double v) except? -1e999: + return cmath.exp(self._function.evaluate(x, y, z, u, w, v)) + + +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 w, double v) except? -1e999: + return cmath.sin(self._function.evaluate(x, y, z, u, w, v)) + + +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 w, double v) except? -1e999: + return cmath.cos(self._function.evaluate(x, y, z, u, w, v)) + + +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 w, double v) except? -1e999: + return cmath.tan(self._function.evaluate(x, y, z, u, w, v)) + + +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 w, double v) except? -1e999: + cdef double val = self._function.evaluate(x, y, z, u, w, v) + 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 w, double v) except? -1e999: + cdef double val = self._function.evaluate(x, y, z, u, w, v) + 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 w, double v) except? -1e999: + return cmath.atan(self._function.evaluate(x, y, z, u, w, v)) + + +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 w, double v) except? -1e999: + return cmath.atan2(self._numerator.evaluate(x, y, z, u, w, v), + self._denominator.evaluate(x, y, z, u, w, v)) + + +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 w, double v) except? -1e999: + cdef double f = self._function.evaluate(x, y, z, u, w, v) + 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 w, double v) except? -1e999: + return cmath.erf(self._function.evaluate(x, y, z, u, w, v)) \ No newline at end of file diff --git a/cherab/core/math/function/float/function6d/constant.pxd b/cherab/core/math/function/float/function6d/constant.pxd new file mode 100644 index 00000000..2ff17e61 --- /dev/null +++ b/cherab/core/math/function/float/function6d/constant.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 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..1464f01c --- /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 w, double v) 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..c5f42cae --- /dev/null +++ b/cherab/core/math/function/float/function6d/tests/test_arg.py @@ -0,0 +1,54 @@ +# 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 +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 in testvals: + for y in testvals: + for z in testvals: + for u in testvals: + for w in testvals: + for v in testvals: + argx = Arg6D("x") + argy = Arg6D("y") + argz = Arg6D("z") + argu = Arg6D("u") + argw = Arg6D("w") + argv = Arg6D("v") + self.assertEqual(argx(x, y, z, u, w, v), x, "Arg6D('x') call did not match reference value.") + self.assertEqual(argy(x, y, z, u, w, v), y, "Arg6D('y') call did not match reference value.") + self.assertEqual(argz(x, y, z, u, w, v), z, "Arg6D('z') call did not match reference value.") + self.assertEqual(argu(x, y, z, u, w, v), u, "Arg6D('u') call did not match reference value.") + self.assertEqual(argw(x, y, z, u, w, v), w, "Arg6D('w') call did not match reference value.") + self.assertEqual(argv(x, y, z, u, w, v), v, "Arg6D('v') 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..2ebe0463 --- /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, w, v: 10*x + 5*y + 2*z + u + 3*w + 4*v) + 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..66a5e2a1 --- /dev/null +++ b/cherab/core/math/function/float/function6d/tests/test_base.py @@ -0,0 +1,649 @@ +# 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 +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, w, v: 10 * x + 5 * y + 2 * z + u + 3 * w + 4 * v + self.ref2 = lambda x, y, z, u, w, v: abs(x + y + z + u + w + v) + + 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 in testvals: + for y in testvals: + for z in testvals: + for u in testvals: + for w in testvals: + for v in testvals: + self.assertEqual(self.f1(x, y, z, u, w, v), self.ref1(x, y, z, u, w, v), + "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 in testvals: + for y in testvals: + for z in testvals: + for u in testvals: + for w in testvals: + for v in testvals: + self.assertEqual(r(x, y, z, u, w, v), -self.ref1(x, y, z, u, w, v), + "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 in testvals: + for y in testvals: + for z in testvals: + for u in testvals: + for w in testvals: + for v in testvals: + self.assertEqual(r1(x, y, z, u, w, v), 8 + self.ref1(x, y, z, u, w, v), + "Function6D add scalar (K + f()) did not match reference function value.") + self.assertEqual(r2(x, y, z, u, w, v), self.ref1(x, y, z, u, w, v) + 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 in testvals: + for y in testvals: + for z in testvals: + for u in testvals: + for w in testvals: + for v in testvals: + self.assertEqual(r1(x, y, z, u, w, v), 8 - self.ref1(x, y, z, u, w, v), + "Function6D subtract scalar (K - f()) did not match reference function value.") + self.assertEqual(r2(x, y, z, u, w, v), self.ref1(x, y, z, u, w, v) - 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 in testvals: + for y in testvals: + for z in testvals: + for u in testvals: + for w in testvals: + for v in testvals: + self.assertEqual(r1(x, y, z, u, w, v), 5 * self.ref1(x, y, z, u, w, v), + "Function6D multiply scalar (K * f()) did not match reference function value.") + self.assertEqual(r2(x, y, z, u, w, v), self.ref1(x, y, z, u, w, v) * -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 in testvals: + for y in testvals: + for z in testvals: + for u in testvals: + for w in testvals: + for v in testvals: + self.assertEqual(r1(x, y, z, u, w, v), 5.451 / self.ref1(x, y, z, u, w, v), + "Function6D divide scalar (K / f()) did not match reference function value.") + self.assertAlmostEqual(r2(x, y, z, u, w, v), self.ref1(x, y, z, u, w, v) / -7.8, + delta=abs(r2(x, y, z, u, w, v)) * 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 in testvals: + for y in testvals: + for z in testvals: + for u in testvals: + for w in testvals: + for v in testvals: + if self.ref1(x, y, z, u, w, v) == 0: + with self.assertRaises(ZeroDivisionError, msg="ZeroDivisionError not raised when function returns 0."): + r1(x, y, z, u, w, v) + else: + self.assertAlmostEqual(r1(x, y, z, u, w, v), math.fmod(5, self.ref1(x, y, z, u, w, v)), 15, "Function6D modulo scalar (K % f()) did not match reference function value.") + self.assertAlmostEqual(r2(x, y, z, u, w, v), math.fmod(self.ref1(x, y, z, u, w, v), -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 in testvals: + for y in testvals: + for z in testvals: + for u in testvals: + for w in testvals: + for v in testvals: + self.assertAlmostEqual(r1(x, y, z, u, w, v), 5. ** self.ref1(x, y, z, u, w, v), 15, "Function6D power scalar (K ** f()) did not match reference function value.") + if self.ref1(x, y, z, u, w, v) < 0: + with self.assertRaises(ValueError, msg="ValueError not raised when base is negative and exponent non-integral."): + r2(x, y, z, u, w, v) + elif not float(self.ref1(x, y, z, u, w, v)).is_integer(): + with self.assertRaises(ValueError, msg="ValueError not raised when base is negative and exponent non-integral."): + r3(x, y, z, u, w, v) + else: + if self.ref1(x, y, z, u, w, v) == 0: + with self.assertRaises(ZeroDivisionError, msg="ZeroDivisionError not raised when base is 0 and exponent negative."): + r2(x, y, z, u, w, v) + else: + self.assertAlmostEqual(r2(x, y, z, u, w, v), self.ref1(x, y, z, u, w, v) ** -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 in testvals: + for y in testvals: + for z in testvals: + for u in testvals: + for w in testvals: + for v in testvals: + ref_value = self.ref1(x, y, z, u, w, v) + 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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 in testvals: + for y in testvals: + for z in testvals: + for u in testvals: + for w in testvals: + for v in testvals: + self.assertEqual(r1(x, y, z, u, w, v), self.ref1(x, y, z, u, w, v) + self.ref2(x, y, z, u, w, v), "Function6D add function (f1() + f2()) did not match reference function value.") + self.assertEqual(r2(x, y, z, u, w, v), self.ref1(x, y, z, u, w, v) + self.ref2(x, y, z, u, w, v), "Function6D add function (p1() + f2()) did not match reference function value.") + self.assertEqual(r3(x, y, z, u, w, v), self.ref1(x, y, z, u, w, v) + self.ref2(x, y, z, u, w, v), "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 in testvals: + for y in testvals: + for z in testvals: + for u in testvals: + for w in testvals: + for v in testvals: + self.assertEqual(r1(x, y, z, u, w, v), self.ref1(x, y, z, u, w, v) - self.ref2(x, y, z, u, w, v), "Function6D subtract function (f1() - f2()) did not match reference function value.") + self.assertEqual(r2(x, y, z, u, w, v), self.ref1(x, y, z, u, w, v) - self.ref2(x, y, z, u, w, v), "Function6D subtract function (p1() - f2()) did not match reference function value.") + self.assertEqual(r3(x, y, z, u, w, v), self.ref1(x, y, z, u, w, v) - self.ref2(x, y, z, u, w, v), "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 in testvals: + for y in testvals: + for z in testvals: + for u in testvals: + for w in testvals: + for v in testvals: + self.assertEqual(r1(x, y, z, u, w, v), self.ref1(x, y, z, u, w, v) * self.ref2(x, y, z, u, w, v), "Function6D multiply function (f1() * f2()) did not match reference function value.") + self.assertEqual(r2(x, y, z, u, w, v), self.ref1(x, y, z, u, w, v) * self.ref2(x, y, z, u, w, v), "Function6D multiply function (p1() * f2()) did not match reference function value.") + self.assertEqual(r3(x, y, z, u, w, v), self.ref1(x, y, z, u, w, v) * self.ref2(x, y, z, u, w, v), "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 in testvals: + for y in testvals: + for z in testvals: + for u in testvals: + for w in testvals: + for v in testvals: + self.assertAlmostEqual(r1(x, y, z, u, w, v), self.ref1(x, y, z, u, w, v) / self.ref2(x, y, z, u, w, v), delta=abs(r1(x, y, z, u, w, v)) * 1e-12, msg="Function6D divide function (f1() / f2()) did not match reference function value.") + self.assertAlmostEqual(r2(x, y, z, u, w, v), self.ref1(x, y, z, u, w, v) / self.ref2(x, y, z, u, w, v), delta=abs(r2(x, y, z, u, w, v)) * 1e-12, msg="Function6D divide function (p1() / f2()) did not match reference function value.") + self.assertAlmostEqual(r3(x, y, z, u, w, v), self.ref1(x, y, z, u, w, v) / self.ref2(x, y, z, u, w, v), delta=abs(r3(x, y, z, u, w, v)) * 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 in testvals: + for y in testvals: + for z in testvals: + for u in testvals: + for w in testvals: + for v in testvals: + self.assertAlmostEqual(r1(x, y, z, u, w, v), math.fmod(self.ref1(x, y, z, u, w, v), self.ref2(x, y, z, u, w, v)), delta=abs(r1(x, y, z, u, w, v)) * 1e-12, msg="Function6D modulo function (f1() % f2()) did not match reference function value.") + self.assertAlmostEqual(r2(x, y, z, u, w, v), math.fmod(self.ref1(x, y, z, u, w, v), self.ref2(x, y, z, u, w, v)), delta=abs(r2(x, y, z, u, w, v)) * 1e-12, msg="Function6D modulo function (p1() % f2()) did not match reference function value.") + self.assertAlmostEqual(r3(x, y, z, u, w, v), math.fmod(self.ref1(x, y, z, u, w, v), self.ref2(x, y, z, u, w, v)), delta=abs(r3(x, y, z, u, w, v)) * 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 in testvals: + for y in testvals: + for z in testvals: + for u in testvals: + for w in testvals: + for v in testvals: + if self.ref1(x, y, z, u, w, v) < 0 and not float(self.ref2(x, y, z, u, w, v)).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, w, v) + with self.assertRaises(ValueError, msg="ValueError not raised when base is negative and exponent non-integral (2/3)."): + r2(x, y, z, u, w, v) + with self.assertRaises(ValueError, msg="ValueError not raised when base is negative and exponent non-integral (3/3)."): + r3(x, y, z, u, w, v) + else: + self.assertAlmostEqual(r1(x, y, z, u, w, v), self.ref1(x, y, z, u, w, v) ** self.ref2(x, y, z, u, w, v), 15, "Function6D power function (f1() ** f2()) did not match reference function value.") + self.assertAlmostEqual(r2(x, y, z, u, w, v), self.ref1(x, y, z, u, w, v) ** self.ref2(x, y, z, u, w, v), 15, "Function6D power function (p1() ** f2()) did not match reference function value.") + self.assertAlmostEqual(r3(x, y, z, u, w, v), self.ref1(x, y, z, u, w, v) ** self.ref2(x, y, z, u, w, v), 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, w, v: 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 in testvals: + for y in testvals: + for z in testvals: + for u in testvals: + for w in testvals: + for v in testvals: + self.assertEqual(r1(x, y, z, u, w, v), math.fmod(self.ref1(x, y, z, u, w, v) ** 5, 3), "Function6D 3 argument pow(f1(), A, B) did not match reference value.") + self.assertEqual(r2(x, y, z, u, w, v), math.fmod(5 ** self.ref1(x, y, z, u, w, v), 3), "Function6D 3 argument pow(A, f1(), B) did not match reference value.") + self.assertEqual(r3(x, y, z, u, w, v), math.fmod(5 ** self.ref1(x, y, z, u, w, v), self.ref2(x, y, z, u, w, v)), "Function6D 3 argument pow(A, f1(), f2()) did not match reference value.") + self.assertEqual(r4(x, y, z, u, w, v), math.fmod(self.ref2(x, y, z, u, w, v) ** self.ref1(x, y, z, u, w, v), self.ref2(x, y, z, u, w, v)), "Function6D 3 argument pow(f2(), f1(), f2()) did not match reference value.") + self.assertEqual(r5(x, y, z, u, w, v), math.fmod(self.ref2(x, y, z, u, w, v) ** self.ref1(x, y, z, u, w, v), self.ref2(x, y, z, u, w, v)), "Function6D 3 argument pow(f2(), p1(), p2()) did not match reference value.") + self.assertEqual(r6(x, y, z, u, w, v), math.fmod(self.ref2(x, y, z, u, w, v) ** self.ref1(x, y, z, u, w, v), self.ref2(x, y, z, u, w, v)), "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 in testvals: + for y in testvals: + for z in testvals: + for u in testvals: + for w in testvals: + for v in testvals: + self.assertEqual(abs(self.f1)(x, y, z, u, w, v), abs(self.ref1(x, y, z, u, w, v)), + 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 in testvals: + for y in testvals: + for z in testvals: + for u in testvals: + for w in testvals: + for v in testvals: + ref_value = self.ref1 + higher_value = lambda x, y, z, u, w, v: self.ref1(x, y, z, u, w, v) + abs(self.ref1(x, y, z, u, w, v)) + 1 + lower_value = lambda x, y, z, u, w, v: self.ref1(x, y, z, u, w, v) - abs(self.ref1(x, y, z, u, w, v)) - 1 + self.assertEqual( + (self.f1 == ref_value)(x, y, z, u, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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 in testvals: + for y in testvals: + for z in testvals: + for u in testvals: + for w in testvals: + for v in testvals: + ref_value = self.ref1 + higher_value = lambda x, y, z, u, w, v: self.ref1(x, y, z, u, w, v) + abs(self.ref1(x, y, z, u, w, v)) + 1 + lower_value = lambda x, y, z, u, w, v: self.ref1(x, y, z, u, w, v) - abs(self.ref1(x, y, z, u, w, v)) - 1 + self.assertEqual( + (ref_value == self.f1)(x, y, z, u, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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 in testvals: + for y in testvals: + for z in testvals: + for u in testvals: + for w in testvals: + for v in testvals: + 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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..5c9f7a81 --- /dev/null +++ b/cherab/core/math/function/float/function6d/tests/test_cmath.py @@ -0,0 +1,155 @@ +# 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 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, w, v: x / 10 + y + z + u/2 + w/3 + v/4) + self.f2 = PythonFunction6D(lambda x, y, z, u, w, v: x * x + y * y - z * z + u * u + w * w - v * v) + + def test_exp(self): + testvals = [-10.0, -7, -0.001, 0.0, 0.00003, 10, 23.4] + for x in testvals: + for y in testvals: + for z in testvals: + for u in testvals: + for w in testvals: + for v in testvals: + function = cmath6d.Exp6D(self.f1) + expected = math.exp(self.f1(x, y, z, u, w, v)) + self.assertEqual(function(x, y, z, u, w, v), 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 in testvals: + for y in testvals: + for z in testvals: + for u in testvals: + for w in testvals: + for v in testvals: + function = cmath6d.Sin6D(self.f1) + expected = math.sin(self.f1(x, y, z, u, w, v)) + self.assertEqual(function(x, y, z, u, w, v), 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 in testvals: + for y in testvals: + for z in testvals: + for u in testvals: + for w in testvals: + for v in testvals: + function = cmath6d.Cos6D(self.f1) + expected = math.cos(self.f1(x, y, z, u, w, v)) + self.assertEqual(function(x, y, z, u, w, v), 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 in testvals: + for y in testvals: + for z in testvals: + for u in testvals: + for w in testvals: + for v in testvals: + function = cmath6d.Tan6D(self.f1) + expected = math.tan(self.f1(x, y, z, u, w, v)) + self.assertEqual(function(x, y, z, u, w, v), 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 in testvals: + for y in testvals: + for z in testvals: + for u in testvals: + for w in testvals: + for v in testvals: + function = cmath6d.Atan6D(self.f1) + expected = math.atan(self.f1(x, y, z, u, w, v)) + self.assertEqual(function(x, y, z, u, w, v), 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 in testvals: + for y in testvals: + for z in testvals: + for u in testvals: + for w in testvals: + for v in testvals: + function = cmath6d.Atan4Q6D(self.f1, self.f2) + expected = math.atan2(self.f1(x, y, z, u, w, v), self.f2(x, y, z, u, w, v)) + self.assertEqual(function(x, y, z, u, w, v), 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 in testvals: + for y in testvals: + for z in testvals: + for u in testvals: + for w in testvals: + for v in testvals: + expected = math.erf(self.f1(x, y, z, u, w, v)) + self.assertAlmostEqual(function(x, y, z, u, w, v), 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 in testvals: + for y in testvals: + for z in testvals: + for u in testvals: + for w in testvals: + for v in testvals: + expected = math.sqrt(self.f1(x, y, z, u, w, v)) + self.assertEqual(function(x, y, z, u, w, v), 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..1247e8f2 --- /dev/null +++ b/cherab/core/math/function/float/function6d/tests/test_constant.py @@ -0,0 +1,35 @@ +# 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) + self.assertEqual(constant(500, 1.5, -3.14, 2.7, 1.8, 3.6), x, "Constant6D call did not match reference value.") From e65fcc3f3c82a51e83bac0fdfd304adbddc9c041 Mon Sep 17 00:00:00 2001 From: MatejTomes Date: Tue, 4 Feb 2025 11:41:38 +0100 Subject: [PATCH 06/91] Corrected arguments in the docstring formula --- cherab/core/math/function/float/function6d/blend.pyx | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/cherab/core/math/function/float/function6d/blend.pyx b/cherab/core/math/function/float/function6d/blend.pyx index 6630bbb2..e05c0452 100644 --- a/cherab/core/math/function/float/function6d/blend.pyx +++ b/cherab/core/math/function/float/function6d/blend.pyx @@ -31,7 +31,7 @@ cdef class Blend6D(Function6D): this function is as follows: .. math:: - v = (1 - f_m(x)) f_1(x) + f_m(x) f_2(x) + v = (1 - f_m(x, y, z, u, w, v)) f_1(x, y, z, u, w, v) + f_m(x, y, z, u, w, v) f_2(x, y, z, u, w, v) The value of the mask function is clamped to the range [0, 1] if the sampled value exceeds the required range. From 68fa029a5d51f15cab23022febd261627844a68f Mon Sep 17 00:00:00 2001 From: MatejTomes Date: Tue, 4 Feb 2025 11:43:09 +0100 Subject: [PATCH 07/91] Add function6d documentation --- docs/source/math/function.rst | 50 +++++++++++++++++++++++++++++++++++ 1 file changed, 50 insertions(+) diff --git a/docs/source/math/function.rst b/docs/source/math/function.rst index 6f1680e1..1d670ec4 100644 --- a/docs/source/math/function.rst +++ b/docs/source/math/function.rst @@ -10,6 +10,56 @@ 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. +The relation of Function6D to distribution functions makes it domain specific, and so it was included in Cherab's math module. +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 ------------------- From 246f8aca798284ad28e86d29019217b1544555b8 Mon Sep 17 00:00:00 2001 From: MatejTomes Date: Tue, 4 Feb 2025 12:01:28 +0100 Subject: [PATCH 08/91] Add CHANGELOG record --- CHANGELOG.md | 6 ++++++ 1 file changed, 6 insertions(+) diff --git a/CHANGELOG.md b/CHANGELOG.md index 5fe4cc06..49f91145 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -1,6 +1,12 @@ Project Changelog ================= +Release 1.6.0 (TBD) +------------------- + +New: +* Add Function6D framework. (#478) + Release 1.5.0 (27 Aug 2024) ------------------- From 2bc58751c85d302abb9491f8716d4337b6172635 Mon Sep 17 00:00:00 2001 From: Jack Lovell Date: Mon, 24 Mar 2025 15:16:49 +0000 Subject: [PATCH 09/91] Run CI on ubuntu 22.04 runner ubuntu-latest switched to 24.04, but 22.04 is the last version which supports Python 3.7 that we continue to test against for now. --- .github/workflows/ci.yml | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index b32bde9d..da0d082c 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -7,7 +7,7 @@ on: jobs: tests: name: Run tests - runs-on: ubuntu-latest + runs-on: ubuntu-22.04 # Needed for Python 3.7 compatibility strategy: fail-fast: false matrix: From c673241d90d601cf147eeca16e6b6def467fb532 Mon Sep 17 00:00:00 2001 From: MatejTomes Date: Tue, 25 Mar 2025 12:20:57 +0100 Subject: [PATCH 10/91] Update docummentation Reason for having Function6D in Cherab was removed --- docs/source/math/function.rst | 1 - 1 file changed, 1 deletion(-) diff --git a/docs/source/math/function.rst b/docs/source/math/function.rst index 1d670ec4..8b29a66e 100644 --- a/docs/source/math/function.rst +++ b/docs/source/math/function.rst @@ -11,7 +11,6 @@ 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. -The relation of Function6D to distribution functions makes it domain specific, and so it was included in Cherab's math module. It follows closely Raysect's function framework. 6D Scalar Functions From 9f15c4fba46fac83bae074065cf1254c08f75050 Mon Sep 17 00:00:00 2001 From: MatejTomes Date: Tue, 25 Mar 2025 13:13:56 +0100 Subject: [PATCH 11/91] Replace nested forloops with itertools --- .../float/function6d/tests/test_arg.py | 32 +- .../float/function6d/tests/test_base.py | 873 ++++++++---------- .../float/function6d/tests/test_cmath.py | 102 +- .../float/function6d/tests/test_constant.py | 1 + 4 files changed, 433 insertions(+), 575 deletions(-) diff --git a/cherab/core/math/function/float/function6d/tests/test_arg.py b/cherab/core/math/function/float/function6d/tests/test_arg.py index c5f42cae..a594b472 100644 --- a/cherab/core/math/function/float/function6d/tests/test_arg.py +++ b/cherab/core/math/function/float/function6d/tests/test_arg.py @@ -23,6 +23,7 @@ """ import unittest +import itertools from cherab.core.math.function.float.function6d.arg import Arg6D # TODO: expand tests to cover the cython interface @@ -30,24 +31,19 @@ class TestArg6D(unittest.TestCase): def test_arg(self): testvals = [-1e10, -7, -0.001, 0.0, 0.00003, 10, 2.3e49] - for x in testvals: - for y in testvals: - for z in testvals: - for u in testvals: - for w in testvals: - for v in testvals: - argx = Arg6D("x") - argy = Arg6D("y") - argz = Arg6D("z") - argu = Arg6D("u") - argw = Arg6D("w") - argv = Arg6D("v") - self.assertEqual(argx(x, y, z, u, w, v), x, "Arg6D('x') call did not match reference value.") - self.assertEqual(argy(x, y, z, u, w, v), y, "Arg6D('y') call did not match reference value.") - self.assertEqual(argz(x, y, z, u, w, v), z, "Arg6D('z') call did not match reference value.") - self.assertEqual(argu(x, y, z, u, w, v), u, "Arg6D('u') call did not match reference value.") - self.assertEqual(argw(x, y, z, u, w, v), w, "Arg6D('w') call did not match reference value.") - self.assertEqual(argv(x, y, z, u, w, v), v, "Arg6D('v') call did not match reference value.") + for (x, y, z, u, w, v) in itertools.product(testvals, repeat=6): + argx = Arg6D("x") + argy = Arg6D("y") + argz = Arg6D("z") + argu = Arg6D("u") + argw = Arg6D("w") + argv = Arg6D("v") + self.assertEqual(argx(x, y, z, u, w, v), x, "Arg6D('x') call did not match reference value.") + self.assertEqual(argy(x, y, z, u, w, v), y, "Arg6D('y') call did not match reference value.") + self.assertEqual(argz(x, y, z, u, w, v), z, "Arg6D('z') call did not match reference value.") + self.assertEqual(argu(x, y, z, u, w, v), u, "Arg6D('u') call did not match reference value.") + self.assertEqual(argw(x, y, z, u, w, v), w, "Arg6D('w') call did not match reference value.") + self.assertEqual(argv(x, y, z, u, w, v), v, "Arg6D('v') call did not match reference value.") def test_invalid_inputs(self): with self.assertRaises(ValueError, msg="Arg6D did not raise ValueError with incorrect string."): diff --git a/cherab/core/math/function/float/function6d/tests/test_base.py b/cherab/core/math/function/float/function6d/tests/test_base.py index 66a5e2a1..d9ec4b13 100644 --- a/cherab/core/math/function/float/function6d/tests/test_base.py +++ b/cherab/core/math/function/float/function6d/tests/test_base.py @@ -24,6 +24,7 @@ import math import unittest +import itertools from cherab.core.math.function.float.function6d.autowrap import PythonFunction6D # TODO: expand tests to cover the cython interface @@ -38,87 +39,57 @@ def setUp(self): def test_call(self): testvals = [-1e10, -7, -0.001, 0.0, 0.00003, 10, 2.3e49] - for x in testvals: - for y in testvals: - for z in testvals: - for u in testvals: - for w in testvals: - for v in testvals: - self.assertEqual(self.f1(x, y, z, u, w, v), self.ref1(x, y, z, u, w, v), - "Function6D call did not match reference function value.") + for (x, y, z, u, w, v) in itertools.product(testvals, repeat=6): + self.assertEqual(self.f1(x, y, z, u, w, v), self.ref1(x, y, z, u, w, v), + "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 in testvals: - for y in testvals: - for z in testvals: - for u in testvals: - for w in testvals: - for v in testvals: - self.assertEqual(r(x, y, z, u, w, v), -self.ref1(x, y, z, u, w, v), - "Function6D negate did not match reference function value.") + for (x, y, z, u, w, v) in itertools.product(testvals, repeat=6): + self.assertEqual(r(x, y, z, u, w, v), -self.ref1(x, y, z, u, w, v), + "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 in testvals: - for y in testvals: - for z in testvals: - for u in testvals: - for w in testvals: - for v in testvals: - self.assertEqual(r1(x, y, z, u, w, v), 8 + self.ref1(x, y, z, u, w, v), - "Function6D add scalar (K + f()) did not match reference function value.") - self.assertEqual(r2(x, y, z, u, w, v), self.ref1(x, y, z, u, w, v) + 65, - "Function6D add scalar (f() + K) did not match reference function value.") + for (x, y, z, u, w, v) in itertools.product(testvals, repeat=6): + self.assertEqual(r1(x, y, z, u, w, v), 8 + self.ref1(x, y, z, u, w, v), + "Function6D add scalar (K + f()) did not match reference function value.") + self.assertEqual(r2(x, y, z, u, w, v), self.ref1(x, y, z, u, w, v) + 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 in testvals: - for y in testvals: - for z in testvals: - for u in testvals: - for w in testvals: - for v in testvals: - self.assertEqual(r1(x, y, z, u, w, v), 8 - self.ref1(x, y, z, u, w, v), - "Function6D subtract scalar (K - f()) did not match reference function value.") - self.assertEqual(r2(x, y, z, u, w, v), self.ref1(x, y, z, u, w, v) - 65, - "Function6D subtract scalar (f() - K) did not match reference function value.") + for (x, y, z, u, w, v) in itertools.product(testvals, repeat=6): + self.assertEqual(r1(x, y, z, u, w, v), 8 - self.ref1(x, y, z, u, w, v), + "Function6D subtract scalar (K - f()) did not match reference function value.") + self.assertEqual(r2(x, y, z, u, w, v), self.ref1(x, y, z, u, w, v) - 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 in testvals: - for y in testvals: - for z in testvals: - for u in testvals: - for w in testvals: - for v in testvals: - self.assertEqual(r1(x, y, z, u, w, v), 5 * self.ref1(x, y, z, u, w, v), - "Function6D multiply scalar (K * f()) did not match reference function value.") - self.assertEqual(r2(x, y, z, u, w, v), self.ref1(x, y, z, u, w, v) * -7.8, - "Function6D multiply scalar (f() * K) did not match reference function value.") + for (x, y, z, u, w, v) in itertools.product(testvals, repeat=6): + self.assertEqual(r1(x, y, z, u, w, v), 5 * self.ref1(x, y, z, u, w, v), + "Function6D multiply scalar (K * f()) did not match reference function value.") + self.assertEqual(r2(x, y, z, u, w, v), self.ref1(x, y, z, u, w, v) * -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 in testvals: - for y in testvals: - for z in testvals: - for u in testvals: - for w in testvals: - for v in testvals: - self.assertEqual(r1(x, y, z, u, w, v), 5.451 / self.ref1(x, y, z, u, w, v), - "Function6D divide scalar (K / f()) did not match reference function value.") - self.assertAlmostEqual(r2(x, y, z, u, w, v), self.ref1(x, y, z, u, w, v) / -7.8, - delta=abs(r2(x, y, z, u, w, v)) * 1e-12, - msg="Function6D divide scalar (f() / K) did not match reference function value.") + for (x, y, z, u, w, v) in itertools.product(testvals, repeat=6): + self.assertEqual(r1(x, y, z, u, w, v), 5.451 / self.ref1(x, y, z, u, w, v), + "Function6D divide scalar (K / f()) did not match reference function value.") + self.assertAlmostEqual(r2(x, y, z, u, w, v), self.ref1(x, y, z, u, w, v) / -7.8, + delta=abs(r2(x, y, z, u, w, v)) * 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."): @@ -131,18 +102,13 @@ def test_mod_function6d_scalar(self): testvals = [-10, -7, -0.001, 0.00003, 10, 12.3] r1 = 5 % self.f1 r2 = self.f1 % -7.8 - for x in testvals: - for y in testvals: - for z in testvals: - for u in testvals: - for w in testvals: - for v in testvals: - if self.ref1(x, y, z, u, w, v) == 0: - with self.assertRaises(ZeroDivisionError, msg="ZeroDivisionError not raised when function returns 0."): - r1(x, y, z, u, w, v) - else: - self.assertAlmostEqual(r1(x, y, z, u, w, v), math.fmod(5, self.ref1(x, y, z, u, w, v)), 15, "Function6D modulo scalar (K % f()) did not match reference function value.") - self.assertAlmostEqual(r2(x, y, z, u, w, v), math.fmod(self.ref1(x, y, z, u, w, v), -7.8), 15, "Function6D modulo scalar (f() % K) did not match reference function value.") + for (x, y, z, u, w, v) in itertools.product(testvals, repeat=6): + if self.ref1(x, y, z, u, w, v) == 0: + with self.assertRaises(ZeroDivisionError, msg="ZeroDivisionError not raised when function returns 0."): + r1(x, y, z, u, w, v) + else: + self.assertAlmostEqual(r1(x, y, z, u, w, v), math.fmod(5, self.ref1(x, y, z, u, w, v)), 15, "Function6D modulo scalar (K % f()) did not match reference function value.") + self.assertAlmostEqual(r2(x, y, z, u, w, v), math.fmod(self.ref1(x, y, z, u, w, v), -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."): @@ -153,25 +119,20 @@ def test_pow_function6d_scalar(self): r1 = 5. ** self.f1 r2 = self.f1 ** -7.8 r3 = (-5.) ** self.f1 - for x in testvals: - for y in testvals: - for z in testvals: - for u in testvals: - for w in testvals: - for v in testvals: - self.assertAlmostEqual(r1(x, y, z, u, w, v), 5. ** self.ref1(x, y, z, u, w, v), 15, "Function6D power scalar (K ** f()) did not match reference function value.") - if self.ref1(x, y, z, u, w, v) < 0: - with self.assertRaises(ValueError, msg="ValueError not raised when base is negative and exponent non-integral."): - r2(x, y, z, u, w, v) - elif not float(self.ref1(x, y, z, u, w, v)).is_integer(): - with self.assertRaises(ValueError, msg="ValueError not raised when base is negative and exponent non-integral."): - r3(x, y, z, u, w, v) - else: - if self.ref1(x, y, z, u, w, v) == 0: - with self.assertRaises(ZeroDivisionError, msg="ZeroDivisionError not raised when base is 0 and exponent negative."): - r2(x, y, z, u, w, v) - else: - self.assertAlmostEqual(r2(x, y, z, u, w, v), self.ref1(x, y, z, u, w, v) ** -7.8, 15, "Function6D power scalar (f() ** K) did not match reference function value.") + for (x, y, z, u, w, v) in itertools.product(testvals, repeat=6): + self.assertAlmostEqual(r1(x, y, z, u, w, v), 5. ** self.ref1(x, y, z, u, w, v), 15, "Function6D power scalar (K ** f()) did not match reference function value.") + if self.ref1(x, y, z, u, w, v) < 0: + with self.assertRaises(ValueError, msg="ValueError not raised when base is negative and exponent non-integral."): + r2(x, y, z, u, w, v) + elif not float(self.ref1(x, y, z, u, w, v)).is_integer(): + with self.assertRaises(ValueError, msg="ValueError not raised when base is negative and exponent non-integral."): + r3(x, y, z, u, w, v) + else: + if self.ref1(x, y, z, u, w, v) == 0: + with self.assertRaises(ZeroDivisionError, msg="ZeroDivisionError not raised when base is 0 and exponent negative."): + r2(x, y, z, u, w, v) + else: + self.assertAlmostEqual(r2(x, y, z, u, w, v), self.ref1(x, y, z, u, w, v) ** -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."): @@ -180,187 +141,162 @@ def test_pow_function6d_scalar(self): def test_richcmp_scalar(self): testvals = [-1e10, -7, -0.001, 0.0, 0.00003, 10, 2.3e49] - for x in testvals: - for y in testvals: - for z in testvals: - for u in testvals: - for w in testvals: - for v in testvals: - ref_value = self.ref1(x, y, z, u, w, v) - 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 0.0, - msg="Scalar greater equals Function6D (K >= f()) did not return false when it should." - ) + for (x, y, z, u, w, v) in itertools.product(testvals, repeat=6): + ref_value = self.ref1(x, y, z, u, w, v) + 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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 in testvals: - for y in testvals: - for z in testvals: - for u in testvals: - for w in testvals: - for v in testvals: - self.assertEqual(r1(x, y, z, u, w, v), self.ref1(x, y, z, u, w, v) + self.ref2(x, y, z, u, w, v), "Function6D add function (f1() + f2()) did not match reference function value.") - self.assertEqual(r2(x, y, z, u, w, v), self.ref1(x, y, z, u, w, v) + self.ref2(x, y, z, u, w, v), "Function6D add function (p1() + f2()) did not match reference function value.") - self.assertEqual(r3(x, y, z, u, w, v), self.ref1(x, y, z, u, w, v) + self.ref2(x, y, z, u, w, v), "Function6D add function (f1() + p2()) did not match reference function value.") + for (x, y, z, u, w, v) in itertools.product(testvals, repeat=6): + self.assertEqual(r1(x, y, z, u, w, v), self.ref1(x, y, z, u, w, v) + self.ref2(x, y, z, u, w, v), "Function6D add function (f1() + f2()) did not match reference function value.") + self.assertEqual(r2(x, y, z, u, w, v), self.ref1(x, y, z, u, w, v) + self.ref2(x, y, z, u, w, v), "Function6D add function (p1() + f2()) did not match reference function value.") + self.assertEqual(r3(x, y, z, u, w, v), self.ref1(x, y, z, u, w, v) + self.ref2(x, y, z, u, w, v), "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 in testvals: - for y in testvals: - for z in testvals: - for u in testvals: - for w in testvals: - for v in testvals: - self.assertEqual(r1(x, y, z, u, w, v), self.ref1(x, y, z, u, w, v) - self.ref2(x, y, z, u, w, v), "Function6D subtract function (f1() - f2()) did not match reference function value.") - self.assertEqual(r2(x, y, z, u, w, v), self.ref1(x, y, z, u, w, v) - self.ref2(x, y, z, u, w, v), "Function6D subtract function (p1() - f2()) did not match reference function value.") - self.assertEqual(r3(x, y, z, u, w, v), self.ref1(x, y, z, u, w, v) - self.ref2(x, y, z, u, w, v), "Function6D subtract function (f1() - p2()) did not match reference function value.") + for (x, y, z, u, w, v) in itertools.product(testvals, repeat=6): + self.assertEqual(r1(x, y, z, u, w, v), self.ref1(x, y, z, u, w, v) - self.ref2(x, y, z, u, w, v), "Function6D subtract function (f1() - f2()) did not match reference function value.") + self.assertEqual(r2(x, y, z, u, w, v), self.ref1(x, y, z, u, w, v) - self.ref2(x, y, z, u, w, v), "Function6D subtract function (p1() - f2()) did not match reference function value.") + self.assertEqual(r3(x, y, z, u, w, v), self.ref1(x, y, z, u, w, v) - self.ref2(x, y, z, u, w, v), "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 in testvals: - for y in testvals: - for z in testvals: - for u in testvals: - for w in testvals: - for v in testvals: - self.assertEqual(r1(x, y, z, u, w, v), self.ref1(x, y, z, u, w, v) * self.ref2(x, y, z, u, w, v), "Function6D multiply function (f1() * f2()) did not match reference function value.") - self.assertEqual(r2(x, y, z, u, w, v), self.ref1(x, y, z, u, w, v) * self.ref2(x, y, z, u, w, v), "Function6D multiply function (p1() * f2()) did not match reference function value.") - self.assertEqual(r3(x, y, z, u, w, v), self.ref1(x, y, z, u, w, v) * self.ref2(x, y, z, u, w, v), "Function6D multiply function (f1() * p2()) did not match reference function value.") + for (x, y, z, u, w, v) in itertools.product(testvals, repeat=6): + self.assertEqual(r1(x, y, z, u, w, v), self.ref1(x, y, z, u, w, v) * self.ref2(x, y, z, u, w, v), "Function6D multiply function (f1() * f2()) did not match reference function value.") + self.assertEqual(r2(x, y, z, u, w, v), self.ref1(x, y, z, u, w, v) * self.ref2(x, y, z, u, w, v), "Function6D multiply function (p1() * f2()) did not match reference function value.") + self.assertEqual(r3(x, y, z, u, w, v), self.ref1(x, y, z, u, w, v) * self.ref2(x, y, z, u, w, v), "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 in testvals: - for y in testvals: - for z in testvals: - for u in testvals: - for w in testvals: - for v in testvals: - self.assertAlmostEqual(r1(x, y, z, u, w, v), self.ref1(x, y, z, u, w, v) / self.ref2(x, y, z, u, w, v), delta=abs(r1(x, y, z, u, w, v)) * 1e-12, msg="Function6D divide function (f1() / f2()) did not match reference function value.") - self.assertAlmostEqual(r2(x, y, z, u, w, v), self.ref1(x, y, z, u, w, v) / self.ref2(x, y, z, u, w, v), delta=abs(r2(x, y, z, u, w, v)) * 1e-12, msg="Function6D divide function (p1() / f2()) did not match reference function value.") - self.assertAlmostEqual(r3(x, y, z, u, w, v), self.ref1(x, y, z, u, w, v) / self.ref2(x, y, z, u, w, v), delta=abs(r3(x, y, z, u, w, v)) * 1e-12, msg="Function6D divide function (f1() / p2()) did not match reference function value.") + for (x, y, z, u, w, v) in itertools.product(testvals, repeat=6): + self.assertAlmostEqual(r1(x, y, z, u, w, v), self.ref1(x, y, z, u, w, v) / self.ref2(x, y, z, u, w, v), delta=abs(r1(x, y, z, u, w, v)) * 1e-12, msg="Function6D divide function (f1() / f2()) did not match reference function value.") + self.assertAlmostEqual(r2(x, y, z, u, w, v), self.ref1(x, y, z, u, w, v) / self.ref2(x, y, z, u, w, v), delta=abs(r2(x, y, z, u, w, v)) * 1e-12, msg="Function6D divide function (p1() / f2()) did not match reference function value.") + self.assertAlmostEqual(r3(x, y, z, u, w, v), self.ref1(x, y, z, u, w, v) / self.ref2(x, y, z, u, w, v), delta=abs(r3(x, y, z, u, w, v)) * 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) @@ -370,15 +306,10 @@ def test_mod_function6d(self): r1 = self.f1 % self.f2 r2 = self.ref1 % self.f2 r3 = self.f1 % self.ref2 - for x in testvals: - for y in testvals: - for z in testvals: - for u in testvals: - for w in testvals: - for v in testvals: - self.assertAlmostEqual(r1(x, y, z, u, w, v), math.fmod(self.ref1(x, y, z, u, w, v), self.ref2(x, y, z, u, w, v)), delta=abs(r1(x, y, z, u, w, v)) * 1e-12, msg="Function6D modulo function (f1() % f2()) did not match reference function value.") - self.assertAlmostEqual(r2(x, y, z, u, w, v), math.fmod(self.ref1(x, y, z, u, w, v), self.ref2(x, y, z, u, w, v)), delta=abs(r2(x, y, z, u, w, v)) * 1e-12, msg="Function6D modulo function (p1() % f2()) did not match reference function value.") - self.assertAlmostEqual(r3(x, y, z, u, w, v), math.fmod(self.ref1(x, y, z, u, w, v), self.ref2(x, y, z, u, w, v)), delta=abs(r3(x, y, z, u, w, v)) * 1e-12, msg="Function6D modulo function (f1() % p2()) did not match reference function value.") + for (x, y, z, u, w, v) in itertools.product(testvals, repeat=6): + self.assertAlmostEqual(r1(x, y, z, u, w, v), math.fmod(self.ref1(x, y, z, u, w, v), self.ref2(x, y, z, u, w, v)), delta=abs(r1(x, y, z, u, w, v)) * 1e-12, msg="Function6D modulo function (f1() % f2()) did not match reference function value.") + self.assertAlmostEqual(r2(x, y, z, u, w, v), math.fmod(self.ref1(x, y, z, u, w, v), self.ref2(x, y, z, u, w, v)), delta=abs(r2(x, y, z, u, w, v)) * 1e-12, msg="Function6D modulo function (p1() % f2()) did not match reference function value.") + self.assertAlmostEqual(r3(x, y, z, u, w, v), math.fmod(self.ref1(x, y, z, u, w, v), self.ref2(x, y, z, u, w, v)), delta=abs(r3(x, y, z, u, w, v)) * 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) @@ -388,23 +319,18 @@ def test_pow_function6d_function6d(self): r1 = self.f1 ** self.f2 r2 = self.ref1 ** self.f2 r3 = self.f1 ** self.ref2 - for x in testvals: - for y in testvals: - for z in testvals: - for u in testvals: - for w in testvals: - for v in testvals: - if self.ref1(x, y, z, u, w, v) < 0 and not float(self.ref2(x, y, z, u, w, v)).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, w, v) - with self.assertRaises(ValueError, msg="ValueError not raised when base is negative and exponent non-integral (2/3)."): - r2(x, y, z, u, w, v) - with self.assertRaises(ValueError, msg="ValueError not raised when base is negative and exponent non-integral (3/3)."): - r3(x, y, z, u, w, v) - else: - self.assertAlmostEqual(r1(x, y, z, u, w, v), self.ref1(x, y, z, u, w, v) ** self.ref2(x, y, z, u, w, v), 15, "Function6D power function (f1() ** f2()) did not match reference function value.") - self.assertAlmostEqual(r2(x, y, z, u, w, v), self.ref1(x, y, z, u, w, v) ** self.ref2(x, y, z, u, w, v), 15, "Function6D power function (p1() ** f2()) did not match reference function value.") - self.assertAlmostEqual(r3(x, y, z, u, w, v), self.ref1(x, y, z, u, w, v) ** self.ref2(x, y, z, u, w, v), 15, "Function6D power function (f1() ** p2()) did not match reference function value.") + for (x, y, z, u, w, v) in itertools.product(testvals, repeat=6): + if self.ref1(x, y, z, u, w, v) < 0 and not float(self.ref2(x, y, z, u, w, v)).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, w, v) + with self.assertRaises(ValueError, msg="ValueError not raised when base is negative and exponent non-integral (2/3)."): + r2(x, y, z, u, w, v) + with self.assertRaises(ValueError, msg="ValueError not raised when base is negative and exponent non-integral (3/3)."): + r3(x, y, z, u, w, v) + else: + self.assertAlmostEqual(r1(x, y, z, u, w, v), self.ref1(x, y, z, u, w, v) ** self.ref2(x, y, z, u, w, v), 15, "Function6D power function (f1() ** f2()) did not match reference function value.") + self.assertAlmostEqual(r2(x, y, z, u, w, v), self.ref1(x, y, z, u, w, v) ** self.ref2(x, y, z, u, w, v), 15, "Function6D power function (p1() ** f2()) did not match reference function value.") + self.assertAlmostEqual(r3(x, y, z, u, w, v), self.ref1(x, y, z, u, w, v) ** self.ref2(x, y, z, u, w, v), 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, w, v: 0) ** self.f1 @@ -420,230 +346,205 @@ def test_pow_3_arguments(self): 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 in testvals: - for y in testvals: - for z in testvals: - for u in testvals: - for w in testvals: - for v in testvals: - self.assertEqual(r1(x, y, z, u, w, v), math.fmod(self.ref1(x, y, z, u, w, v) ** 5, 3), "Function6D 3 argument pow(f1(), A, B) did not match reference value.") - self.assertEqual(r2(x, y, z, u, w, v), math.fmod(5 ** self.ref1(x, y, z, u, w, v), 3), "Function6D 3 argument pow(A, f1(), B) did not match reference value.") - self.assertEqual(r3(x, y, z, u, w, v), math.fmod(5 ** self.ref1(x, y, z, u, w, v), self.ref2(x, y, z, u, w, v)), "Function6D 3 argument pow(A, f1(), f2()) did not match reference value.") - self.assertEqual(r4(x, y, z, u, w, v), math.fmod(self.ref2(x, y, z, u, w, v) ** self.ref1(x, y, z, u, w, v), self.ref2(x, y, z, u, w, v)), "Function6D 3 argument pow(f2(), f1(), f2()) did not match reference value.") - self.assertEqual(r5(x, y, z, u, w, v), math.fmod(self.ref2(x, y, z, u, w, v) ** self.ref1(x, y, z, u, w, v), self.ref2(x, y, z, u, w, v)), "Function6D 3 argument pow(f2(), p1(), p2()) did not match reference value.") - self.assertEqual(r6(x, y, z, u, w, v), math.fmod(self.ref2(x, y, z, u, w, v) ** self.ref1(x, y, z, u, w, v), self.ref2(x, y, z, u, w, v)), "Function6D 3 argument pow(p2(), f1(), f2()) did not match reference value.") + for (x, y, z, u, w, v) in itertools.product(testvals, repeat=6): + self.assertEqual(r1(x, y, z, u, w, v), math.fmod(self.ref1(x, y, z, u, w, v) ** 5, 3), "Function6D 3 argument pow(f1(), A, B) did not match reference value.") + self.assertEqual(r2(x, y, z, u, w, v), math.fmod(5 ** self.ref1(x, y, z, u, w, v), 3), "Function6D 3 argument pow(A, f1(), B) did not match reference value.") + self.assertEqual(r3(x, y, z, u, w, v), math.fmod(5 ** self.ref1(x, y, z, u, w, v), self.ref2(x, y, z, u, w, v)), "Function6D 3 argument pow(A, f1(), f2()) did not match reference value.") + self.assertEqual(r4(x, y, z, u, w, v), math.fmod(self.ref2(x, y, z, u, w, v) ** self.ref1(x, y, z, u, w, v), self.ref2(x, y, z, u, w, v)), "Function6D 3 argument pow(f2(), f1(), f2()) did not match reference value.") + self.assertEqual(r5(x, y, z, u, w, v), math.fmod(self.ref2(x, y, z, u, w, v) ** self.ref1(x, y, z, u, w, v), self.ref2(x, y, z, u, w, v)), "Function6D 3 argument pow(f2(), p1(), p2()) did not match reference value.") + self.assertEqual(r6(x, y, z, u, w, v), math.fmod(self.ref2(x, y, z, u, w, v) ** self.ref1(x, y, z, u, w, v), self.ref2(x, y, z, u, w, v)), "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 in testvals: - for y in testvals: - for z in testvals: - for u in testvals: - for w in testvals: - for v in testvals: - self.assertEqual(abs(self.f1)(x, y, z, u, w, v), abs(self.ref1(x, y, z, u, w, v)), - msg="abs(Function6D) did not match reference value") + for (x, y, z, u, w, v) in itertools.product(testvals, repeat=6): + self.assertEqual(abs(self.f1)(x, y, z, u, w, v), abs(self.ref1(x, y, z, u, w, v)), + 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 in testvals: - for y in testvals: - for z in testvals: - for u in testvals: - for w in testvals: - for v in testvals: - ref_value = self.ref1 - higher_value = lambda x, y, z, u, w, v: self.ref1(x, y, z, u, w, v) + abs(self.ref1(x, y, z, u, w, v)) + 1 - lower_value = lambda x, y, z, u, w, v: self.ref1(x, y, z, u, w, v) - abs(self.ref1(x, y, z, u, w, v)) - 1 - self.assertEqual( - (self.f1 == ref_value)(x, y, z, u, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 0.0, - msg="Function6D equals callable (f1() >= f2()) did not return false when it should." - ) + for (x, y, z, u, w, v) in itertools.product(testvals, repeat=6): + ref_value = self.ref1 + higher_value = lambda x, y, z, u, w, v: self.ref1(x, y, z, u, w, v) + abs(self.ref1(x, y, z, u, w, v)) + 1 + lower_value = lambda x, y, z, u, w, v: self.ref1(x, y, z, u, w, v) - abs(self.ref1(x, y, z, u, w, v)) - 1 + self.assertEqual( + (self.f1 == ref_value)(x, y, z, u, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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 in testvals: - for y in testvals: - for z in testvals: - for u in testvals: - for w in testvals: - for v in testvals: - ref_value = self.ref1 - higher_value = lambda x, y, z, u, w, v: self.ref1(x, y, z, u, w, v) + abs(self.ref1(x, y, z, u, w, v)) + 1 - lower_value = lambda x, y, z, u, w, v: self.ref1(x, y, z, u, w, v) - abs(self.ref1(x, y, z, u, w, v)) - 1 - self.assertEqual( - (ref_value == self.f1)(x, y, z, u, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 0.0, - msg="Callable equals Function6D (f1() >= f2()) did not return false when it should." - ) + for (x, y, z, u, w, v) in itertools.product(testvals, repeat=6): + ref_value = self.ref1 + higher_value = lambda x, y, z, u, w, v: self.ref1(x, y, z, u, w, v) + abs(self.ref1(x, y, z, u, w, v)) + 1 + lower_value = lambda x, y, z, u, w, v: self.ref1(x, y, z, u, w, v) - abs(self.ref1(x, y, z, u, w, v)) - 1 + self.assertEqual( + (ref_value == self.f1)(x, y, z, u, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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 in testvals: - for y in testvals: - for z in testvals: - for u in testvals: - for w in testvals: - for v in testvals: - 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 0.0, - msg="Function6D equals Function6D (f1() >= f2()) did not return false when it should." - ) + for (x, y, z, u, w, v) 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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, w, v), 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 index 5c9f7a81..257dd0b7 100644 --- a/cherab/core/math/function/float/function6d/tests/test_cmath.py +++ b/cherab/core/math/function/float/function6d/tests/test_cmath.py @@ -24,6 +24,7 @@ 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 @@ -36,51 +37,31 @@ def setUp(self): def test_exp(self): testvals = [-10.0, -7, -0.001, 0.0, 0.00003, 10, 23.4] - for x in testvals: - for y in testvals: - for z in testvals: - for u in testvals: - for w in testvals: - for v in testvals: - function = cmath6d.Exp6D(self.f1) - expected = math.exp(self.f1(x, y, z, u, w, v)) - self.assertEqual(function(x, y, z, u, w, v), expected, "Exp6D call did not match reference value.") + for (x, y, z, u, w, v) in itertools.product(testvals, repeat=6): + function = cmath6d.Exp6D(self.f1) + expected = math.exp(self.f1(x, y, z, u, w, v)) + self.assertEqual(function(x, y, z, u, w, v), 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 in testvals: - for y in testvals: - for z in testvals: - for u in testvals: - for w in testvals: - for v in testvals: - function = cmath6d.Sin6D(self.f1) - expected = math.sin(self.f1(x, y, z, u, w, v)) - self.assertEqual(function(x, y, z, u, w, v), expected, "Sin6D call did not match reference value.") + for (x, y, z, u, w, v) in itertools.product(testvals, repeat=6): + function = cmath6d.Sin6D(self.f1) + expected = math.sin(self.f1(x, y, z, u, w, v)) + self.assertEqual(function(x, y, z, u, w, v), 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 in testvals: - for y in testvals: - for z in testvals: - for u in testvals: - for w in testvals: - for v in testvals: - function = cmath6d.Cos6D(self.f1) - expected = math.cos(self.f1(x, y, z, u, w, v)) - self.assertEqual(function(x, y, z, u, w, v), expected, "Cos6D call did not match reference value.") + for (x, y, z, u, w, v) in itertools.product(testvals, repeat=6): + function = cmath6d.Cos6D(self.f1) + expected = math.cos(self.f1(x, y, z, u, w, v)) + self.assertEqual(function(x, y, z, u, w, v), 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 in testvals: - for y in testvals: - for z in testvals: - for u in testvals: - for w in testvals: - for v in testvals: - function = cmath6d.Tan6D(self.f1) - expected = math.tan(self.f1(x, y, z, u, w, v)) - self.assertEqual(function(x, y, z, u, w, v), expected, "Tan6D call did not match reference value.") + for (x, y, z, u, w, v) in itertools.product(testvals, repeat=6): + function = cmath6d.Tan6D(self.f1) + expected = math.tan(self.f1(x, y, z, u, w, v)) + self.assertEqual(function(x, y, z, u, w, v), expected, "Tan6D call did not match reference value.") def test_asin(self): v = [-10, -6, -2, -0.001, 0, 0.001, 2, 6, 10] @@ -102,54 +83,33 @@ def test_acos(self): 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 in testvals: - for y in testvals: - for z in testvals: - for u in testvals: - for w in testvals: - for v in testvals: - function = cmath6d.Atan6D(self.f1) - expected = math.atan(self.f1(x, y, z, u, w, v)) - self.assertEqual(function(x, y, z, u, w, v), expected, "Atan6D call did not match reference value.") + for (x, y, z, u, w, v) in itertools.product(testvals, repeat=6): + function = cmath6d.Atan6D(self.f1) + expected = math.atan(self.f1(x, y, z, u, w, v)) + self.assertEqual(function(x, y, z, u, w, v), 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 in testvals: - for y in testvals: - for z in testvals: - for u in testvals: - for w in testvals: - for v in testvals: - function = cmath6d.Atan4Q6D(self.f1, self.f2) - expected = math.atan2(self.f1(x, y, z, u, w, v), self.f2(x, y, z, u, w, v)) - self.assertEqual(function(x, y, z, u, w, v), expected, "Atan4Q6D call did not match reference value.") + for (x, y, z, u, w, v) in itertools.product(testvals, repeat=6): + function = cmath6d.Atan4Q6D(self.f1, self.f2) + expected = math.atan2(self.f1(x, y, z, u, w, v), self.f2(x, y, z, u, w, v)) + self.assertEqual(function(x, y, z, u, w, v), 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 in testvals: - for y in testvals: - for z in testvals: - for u in testvals: - for w in testvals: - for v in testvals: - expected = math.erf(self.f1(x, y, z, u, w, v)) - self.assertAlmostEqual(function(x, y, z, u, w, v), expected, 10, "Erf6D call did not match reference value.") + for (x, y, z, u, w, v) in itertools.product(testvals, repeat=6): + expected = math.erf(self.f1(x, y, z, u, w, v)) + self.assertAlmostEqual(function(x, y, z, u, w, v), 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 in testvals: - for y in testvals: - for z in testvals: - for u in testvals: - for w in testvals: - for v in testvals: - expected = math.sqrt(self.f1(x, y, z, u, w, v)) - self.assertEqual(function(x, y, z, u, w, v), expected, "Sqrt6D call did not match reference value.") + for (x, y, z, u, w, v) in itertools.product(testvals, repeat=6): + expected = math.sqrt(self.f1(x, y, z, u, w, v)) + self.assertEqual(function(x, y, z, u, w, v), 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 index 1247e8f2..bc7e8cba 100644 --- a/cherab/core/math/function/float/function6d/tests/test_constant.py +++ b/cherab/core/math/function/float/function6d/tests/test_constant.py @@ -32,4 +32,5 @@ 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.") From 677456b0215d420260fdd7ce48304a73f5f58ef0 Mon Sep 17 00:00:00 2001 From: Koyo MUNECHIKA <51052381+munechika-koyo@users.noreply.github.com> Date: Mon, 21 Apr 2025 05:41:50 +0900 Subject: [PATCH 12/91] Modify the description about cherab-iter --- docs/source/available_modules.rst | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/docs/source/available_modules.rst b/docs/source/available_modules.rst index 9be25f1e..c99ffe1f 100644 --- a/docs/source/available_modules.rst +++ b/docs/source/available_modules.rst @@ -34,9 +34,9 @@ Fusion Experiment Packages - The Cherab configuration package for AUG. * - `cherab-compass `_ - The Cherab configuration package for COMPASS. - * - cherab-iter + * - `cherab-iter `_ - Integrates Cherab with IMAS and provides diagnostic configuration - for ITER. This package is under development but not yet publicly available. + for ITER. * - `cherab-jet `_ - Experiment configuration package for JET. * - `cherab-mastu `_ From 82e3014d0e8897a8caa83470c45095a47d3da750 Mon Sep 17 00:00:00 2001 From: MatejTomes Date: Thu, 5 Jun 2025 16:16:03 +0200 Subject: [PATCH 13/91] Add read access to line and lineshape attributes This should add the possibility to fully investigate what emission models an observation was performed with. --- cherab/core/model/plasma/impact_excitation.pyx | 10 ++++++++++ cherab/core/model/plasma/recombination.pyx | 10 ++++++++++ cherab/core/model/plasma/thermal_cx.pyx | 10 ++++++++++ cherab/core/model/plasma/total_radiated_power.pyx | 11 +++++++++++ 4 files changed, 41 insertions(+) diff --git a/cherab/core/model/plasma/impact_excitation.pyx b/cherab/core/model/plasma/impact_excitation.pyx index b26336f1..e10b6280 100644 --- a/cherab/core/model/plasma/impact_excitation.pyx +++ b/cherab/core/model/plasma/impact_excitation.pyx @@ -47,6 +47,8 @@ cdef class ExcitationLine(PlasmaModel): :ivar Plasma plasma: The plasma to which this emission model is attached. :ivar AtomicData atomic_data: The atomic data provider for this model. + :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): + return self._line + + @property + def lineshape(self): + 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..db00ad35 100644 --- a/cherab/core/model/plasma/recombination.pyx +++ b/cherab/core/model/plasma/recombination.pyx @@ -47,6 +47,8 @@ cdef class RecombinationLine(PlasmaModel): :ivar Plasma plasma: The plasma to which this emission model is attached. :ivar AtomicData atomic_data: The atomic data provider for this model. + :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): + return self._line + + @property + def lineshape(self): + 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/thermal_cx.pyx b/cherab/core/model/plasma/thermal_cx.pyx index 88af9ae8..aa1e9ffa 100644 --- a/cherab/core/model/plasma/thermal_cx.pyx +++ b/cherab/core/model/plasma/thermal_cx.pyx @@ -49,6 +49,8 @@ cdef class ThermalCXLine(PlasmaModel): :ivar Plasma plasma: The plasma to which this emission model is attached. :ivar AtomicData atomic_data: The atomic data provider for this model. + :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): + return self._line + + @property + def lineshape(self): + 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..742cb3db 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: The atomic element/isotope. + :ivar int charge: The charge state of the element/isotope. """ 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): + return self._element + + @property + def charge(self): + return self._charge cpdef Spectrum emission(self, Point3D point, Vector3D direction, Spectrum spectrum): From 717378752adf59fc1c5f20e529cf4392fa9e1c11 Mon Sep 17 00:00:00 2001 From: MatejTomes Date: Thu, 5 Jun 2025 16:21:44 +0200 Subject: [PATCH 14/91] Add documentation to readonly attributes --- cherab/core/atomic/line.pyx | 8 ++++++++ 1 file changed, 8 insertions(+) diff --git a/cherab/core/atomic/line.pyx b/cherab/core/atomic/line.pyx index 262c7a2f..1cb34584 100644 --- a/cherab/core/atomic/line.pyx +++ b/cherab/core/atomic/line.pyx @@ -33,6 +33,14 @@ 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: The atomic element/isotope to which this emission line belongs. + :ivar int charge: The charge state of the element/isotope that emits this line. + :ivar tuple transition: A two element tuple that defines the upper and lower electron + configuration states of the transition. For hydrogen-like ions it may be enough to + 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. .. code-block:: pycon From 58e48ddf7a80ec7eafb88f2fef1af11fc766f847 Mon Sep 17 00:00:00 2001 From: MatejTomes Date: Thu, 5 Jun 2025 16:23:02 +0200 Subject: [PATCH 15/91] Modify changelog --- CHANGELOG.md | 3 +++ 1 file changed, 3 insertions(+) diff --git a/CHANGELOG.md b/CHANGELOG.md index 4f22a5d4..35a926c1 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -4,6 +4,9 @@ Project Changelog Release 1.6.0 (TBD) ------------------- +API changes: +* Add emission model attribute access to line and lineshape . (#294) + New: * Add Function6D framework. (#478) * Add e_field attribute to Plasma object for electric field vector. (#465) From 8b9c74142e1820452216cc6c6afd33a6343187c7 Mon Sep 17 00:00:00 2001 From: munechika-koyo Date: Mon, 13 Oct 2025 11:30:07 +0200 Subject: [PATCH 16/91] Update dependencies: upgrade Cython to 3.1 and numpy to >=2; adjust raysect version to 0.9.1.* --- pyproject.toml | 2 +- requirements.txt | 6 +++--- setup.py | 17 +++++++++-------- 3 files changed, 13 insertions(+), 12 deletions(-) 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..7fea14d1 100644 --- a/requirements.txt +++ b/requirements.txt @@ -1,5 +1,5 @@ -cython~=3.0 -numpy>=1.14,<2.0 +cython~=3.1 +numpy>=2 scipy matplotlib -raysect==0.8.1.* +raysect==0.9.1.* diff --git a/setup.py b/setup.py index f11dd08f..8699c4f5 100644 --- a/setup.py +++ b/setup.py @@ -1,14 +1,15 @@ -from collections import defaultdict -import sys +import multiprocessing import os import os.path as path +import sys +from collections import defaultdict from pathlib import Path -import multiprocessing + import numpy -from setuptools import setup, find_packages, Extension from Cython.Build import cythonize +from setuptools import Extension, find_packages, setup -multiprocessing.set_start_method('fork') +multiprocessing.set_start_method("fork") force = False profile = False @@ -117,14 +118,14 @@ long_description=long_description, long_description_content_type="text/markdown", install_requires=[ - "numpy>=1.14,<2.0", + "numpy>=2", "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*"]), package_data={"": [ From d059ed82c30408a9f2b1f6104fd07b2677fa6b68 Mon Sep 17 00:00:00 2001 From: munechika-koyo Date: Mon, 13 Oct 2025 11:37:43 +0200 Subject: [PATCH 17/91] Update CI configuration: switch to ubuntu-latest, adjust Python and numpy versions, and upgrade Raysect to 0.9.* --- .github/workflows/ci.yml | 9 ++++----- 1 file changed, 4 insertions(+), 5 deletions(-) diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index da0d082c..73fcdbd9 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -7,12 +7,11 @@ on: jobs: tests: name: Run tests - runs-on: ubuntu-22.04 # Needed for Python 3.7 compatibility + 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 +22,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 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 From 3ec32a1b34f4e1798c561b6e94cb5294a12f4c42 Mon Sep 17 00:00:00 2001 From: munechika-koyo Date: Mon, 13 Oct 2025 11:38:26 +0200 Subject: [PATCH 18/91] Refactor: Correct spelling of 'Targetted' to 'Targeted' in bolometry and targettedpixel modules --- cherab/tools/observers/bolometry.py | 16 +++++------ .../tools/observers/group/targettedpixel.py | 28 +++++++++---------- cherab/tools/tests/test_observer_groups.py | 16 +++++------ 3 files changed, 30 insertions(+), 30 deletions(-) diff --git a/cherab/tools/observers/bolometry.py b/cherab/tools/observers/bolometry.py index 20179796..5cd8c1a2 100644 --- a/cherab/tools/observers/bolometry.py +++ b/cherab/tools/observers/bolometry.py @@ -22,12 +22,12 @@ 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 @@ -351,7 +351,7 @@ def curvature_radius(self): return self._curvature_radius -class BolometerFoil(TargettedPixel): +class BolometerFoil(TargetedPixel): """ A rectangular foil bolometer detector. @@ -447,7 +447,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) @@ -662,7 +662,7 @@ def calculate_etendue(self, ray_count=10000, batches=10, max_distance=1e999): 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) + targetted_sampler = TargetedHemisphereSampler(spheres) # instance rectangle pixel sampler to sample origins point_sampler = RectangleSampler3D(width=self.x_width, height=self.y_width) @@ -701,7 +701,7 @@ def etendue_single_run(_): return etendue, etendue_error -class BolometerIRVB(TargettedCCDArray): +class BolometerIRVB(TargetedCCDArray): """ A rectangular infra red video bolometer (IRVB). @@ -784,7 +784,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 diff --git a/cherab/tools/observers/group/targettedpixel.py b/cherab/tools/observers/group/targettedpixel.py index 07d13e7f..e440f70e 100644 --- a/cherab/tools/observers/group/targettedpixel.py +++ b/cherab/tools/observers/group/targettedpixel.py @@ -17,14 +17,14 @@ # under the Licence. from numpy import ndarray -from raysect.optical.observer import TargettedPixel +from raysect.optical.observer import TargetedPixel from .base import Observer0DGroup class TargettedPixelGroup(Observer0DGroup): """ - 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' observers as a scene-graph parent. Allows combined observation and display @@ -33,9 +33,10 @@ class TargettedPixelGroup(Observer0DGroup): :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 targetted_path_prob: Probability of ray being casted at the target + :ivar list targeted_path_prob: Probability of ray being casted at the target """ - _OBSERVER_TYPE = TargettedPixel + + _OBSERVER_TYPE = TargetedPixel @property def x_width(self): @@ -76,7 +77,7 @@ 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 + :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 @@ -99,18 +100,17 @@ def targets(self, value): pixel.targets = value @property - def targetted_path_prob(self): - return [pixel.targetted_path_prob for pixel in self._observers] - - @targetted_path_prob.setter - def targetted_path_prob(self, value): + 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.targetted_path_prob = v + pixel.targeted_path_prob = v else: - raise ValueError("The length of 'value' ({}) " - "mismatches the number of pixels ({}).".format(len(value), len(self._observers))) + 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 + pixel.targeted_path_prob = value diff --git a/cherab/tools/tests/test_observer_groups.py b/cherab/tools/tests/test_observer_groups.py index 1f5b7cb0..a04f095c 100644 --- a/cherab/tools/tests/test_observer_groups.py +++ b/cherab/tools/tests/test_observer_groups.py @@ -1,7 +1,7 @@ import unittest 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.base import Observer0DGroup @@ -352,7 +352,7 @@ class TargettedPixelGroupTestCase(PixelGroupTestCase): _GROUP_CLASS = TargettedPixelGroup 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) @@ -376,13 +376,13 @@ def test_targets(self): # targetted 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) From 10d690feb7abc365acbe4383a17bd1699957f14c Mon Sep 17 00:00:00 2001 From: munechika-koyo Date: Mon, 13 Oct 2025 11:44:29 +0200 Subject: [PATCH 19/91] Fix remaining typo `targetted` to `targeted` --- cherab/core/model/laser/profile.pyx | 60 +++--- cherab/tools/observers/__init__.py | 14 +- cherab/tools/observers/bolometry.py | 201 +++++++++--------- cherab/tools/observers/group/__init__.py | 2 +- .../tools/observers/group/targettedpixel.py | 25 ++- cherab/tools/tests/test_observer_groups.py | 70 +++--- docs/source/tools/observers.rst | 14 +- 7 files changed, 200 insertions(+), 186 deletions(-) diff --git a/cherab/core/model/laser/profile.pyx b/cherab/core/model/laser/profile.pyx index 82376980..80b73d02 100644 --- a/cherab/core/model/laser/profile.pyx +++ b/cherab/core/model/laser/profile.pyx @@ -4,7 +4,7 @@ from raysect.primitive import Cylinder from raysect.optical cimport Spectrum, Vector3D, translate from cherab.core.laser cimport Laser, LaserProfile -from cherab.core.model.laser.math_functions cimport ConstantAxisymmetricGaussian3D, ConstantBivariateGaussian3D, TrivariateGaussian3D, GaussianBeamModel +from cherab.core.model.laser.math_functions cimport ConstantAxisymmetricGaussian3D, ConstantBivariateGaussian3D, TrivariateGaussian3D, GaussianBeamModel from cherab.core.utility.constants cimport SPEED_OF_LIGHT @@ -22,20 +22,20 @@ cdef class UniformEnergyDensity(LaserProfile): The methods get_pointing, get_polarization and get_energy_density are not limited to the inside of the laser cylinder. If called alone for position (x, y, z) outisde the laser cylinder, they will still return non-zero values. - + In the following example, a laser of length of 2 m (extending from z=0 to z=2 m) with a radius of 3 cm and volumetric energy density of 5 J*m^-3 and polarisation in the y direction is created: .. code-block:: pycon - + >>> from raysect.core import Vector3D >>> from cherab.core.model.laser import UniformEnergyDensity - + >>> energy = 5 # energy density in J >>> radius = 3e-2 # laser radius in m >>> length = 2 # laser length in m >>> polarisation = Vector3D(0, 1, 0) # polarisation direction - + # create the laser profile >>> laser_profile = UniformEnergyDensity(energy, radius, length, polarisation) @@ -108,7 +108,7 @@ cdef class UniformEnergyDensity(LaserProfile): cpdef list generate_geometry(self): return generate_segmented_cylinder(self.laser_radius, self.laser_length) - + cdef class ConstantBivariateGaussian(LaserProfile): """ @@ -120,8 +120,8 @@ cdef class ConstantBivariateGaussian(LaserProfile): The model imitates a laser beam with a uniform power output within a single pulse. This results in the distribution of the energy density along the propagation direction of the laser (z-axis) to be also uniform. The integral value of laser energy Exy in an x-y plane is given by - - .. math:: + + .. math:: E_{xy} = \\frac{E_p}{(c * \\tau)}, where Ep is the energy of the laser pulse, tau is the temporal pulse length and c is the speed of light in vacuum. @@ -133,23 +133,23 @@ cdef class ConstantBivariateGaussian(LaserProfile): The sigma_x and sigma_y are standard deviations in x and y directions, respectively. .. note:: - The height of the cylinder, forming the laser beam, is given by the laser_length and is independent from the + The height of the cylinder, forming the laser beam, is given by the laser_length and is independent from the temporal length of the laser pulse given by pulse_length. This gives the possibility to independently control the size of the laser primitive and the value of the volumetric energy density. - + The methods get_pointing, get_polarization and get_energy_density are not limited to the inside of the laser cylinder. If called for position (x, y, z) outisde the laser cylinder, they can still return non-zero values. - + The following example shows how to create a laser with sigma_x= 1 cm and sigma_y=2 cm, which makes the laser profile in x-y plane to be elliptical. The pulse energy is 5 J and the laser temporal pulse length is 10 ns: .. code-block:: pycon - + >>> from raysect.core import Vector3D >>> from cherab.core.model.laser import ConstantBivariateGaussian - + >>> radius = 3e-2 # laser radius in m >>> length = 2 # laser length in m >>> polarisation = Vector3D(0, 1, 0) # polarisation direction @@ -157,7 +157,7 @@ cdef class ConstantBivariateGaussian(LaserProfile): >>> pulse_length = 1e-8 # pulse length in s >>> width_x = 1e-2 # standard deviation in x direction in m >>> width_y = 2e-2 # standard deviation in y direction in m - + # create the laser profile >>> laser_profile = ConstantBivariateGaussian(pulse_energy, pulse_length, radius, length, width_x, width_y, polarisation) @@ -323,7 +323,7 @@ cdef class TrivariateGaussian(LaserProfile): The sigma_x and sigma_y are standard deviations in x and y directions, respectively, and E_p is the energy deliverd by laser in a single laser pulse. The mu_z is the mean of the distribution in the z direction and controls th position of the laser pulse along the z direction. - The standard deviation in z direction sigma_z is calculated from the pulse length tau_p, which is the + The standard deviation in z direction sigma_z is calculated from the pulse length tau_p, which is the standard deviation of the Gaussian distributed ouput power of the laser within a single pulse: .. math:: @@ -332,24 +332,24 @@ cdef class TrivariateGaussian(LaserProfile): The c stands for the speed of light in vacuum. .. note:: - The height of the cylinder, forming the laser beam, is given by the laser_length and is independent from the + The height of the cylinder, forming the laser beam, is given by the laser_length and is independent from the temporal length of the laser pulse given by pulse_length. This gives the possibility to independently control the size of the laser primitive and the value of the volumetric energy density. - + The methods get_pointing, get_polarization and get_energy_density are not limited to the inside of the laser cylinder. If called alone for position (x, y, z) outisde the laser cylinder, they can still return non-zero values. - + The following example shows how to create a laser with sigma_x = 1 cm and sigma_y = 2 cm, which makes the laser profile in an x-y plane to be elliptical. The pulse energy is 5 J and the laser temporal pulse length is 10 ns. The position of the laser pulse maximum mean_z is set to 0.5: .. code-block:: pycon - + >>> from raysect.core import Vector3D >>> from cherab.core.model.laser import ConstantBivariateGaussian - + >>> radius = 3e-2 # laser radius in m >>> length = 2 # laser length in m >>> polarisation = Vector3D(0, 1, 0) # polarisation direction @@ -358,7 +358,7 @@ cdef class TrivariateGaussian(LaserProfile): >>> pulse_z = 0.5 # position of the pulse mean >>> width_x = 1e-2 # standard deviation in x direction in m >>> width_y = 2e-2 # standard deviation in y direction in m - + # create the laser profile >>> laser_profile = ConstantBivariateGaussian(pulse_energy, pulse_length, pulse_z, radius, length, width_x, width_y, polarisation) @@ -512,7 +512,7 @@ cdef class TrivariateGaussian(LaserProfile): self._distribution = TrivariateGaussian3D(self._mean_z, self._stddev_x, self._stddev_y, self._stddev_z) - normalisation = self._pulse_energy + normalisation = self._pulse_energy function = normalisation * self._distribution self.set_energy_density_function(function) @@ -541,16 +541,16 @@ cdef class GaussianBeamAxisymmetric(LaserProfile): .. math:: z_R = \\frac{\\pi \\omega_0^2 n}{\\lambda_l} - + where the omega_0 is the standard deviation in the xy plane in the focal point (beam waist) and lambda_l is the central wavelength of the laser. The E_xy stand for the laser energy in an xy plane and is calculated as: - + .. math:: E_{xy} = \\frac{E_p}{(c * \\tau)}, where the E_p is the energy in a single laser pulse and tau is the temporal pulse length. - .. note:: + .. note:: For more information about the Gaussian beam model see https://en.wikipedia.org/wiki/Gaussian_beam The methods get_pointing, get_polarization and get_energy_density are not limited to the inside @@ -562,10 +562,10 @@ cdef class GaussianBeamAxisymmetric(LaserProfile): waist is z=50 cm. The laser wavelength is 1060 nm. .. code-block:: pycon - + >>> from raysect.core import Vector3D >>> from cherab.core.model.laser import GaussianBeamAxisymmetric - + >>> radius = 5e-2 # laser radius in m >>> length = 2 # laser length in m >>> polarisation = Vector3D(0, 1, 0) # polarisation direction @@ -576,7 +576,7 @@ cdef class GaussianBeamAxisymmetric(LaserProfile): >>> width_x = 1e-2 # standard deviation in x direction in m >>> width_y = 2e-2 # standard deviation in y direction in m >>> laser_wlen = 1060 # laser wavelength in nm - + # create the laser profile >>> laser_profile = GaussianBeamAxisymmetric(pulse_energy, pulse_length, length, radius, waist_z, waist_width, laser_wlen) @@ -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 @@ -761,5 +761,5 @@ def generate_segmented_cylinder(radius, length): geometry.append(segment) else: raise ValueError("Incorrect number of segments calculated.") - + return geometry \ No newline at end of file diff --git a/cherab/tools/observers/__init__.py b/cherab/tools/observers/__init__.py index d134ef63..670dfce5 100644 --- a/cherab/tools/observers/__init__.py +++ b/cherab/tools/observers/__init__.py @@ -1,4 +1,3 @@ - # Copyright 2016-2018 Euratom # Copyright 2016-2018 United Kingdom Atomic Energy Authority # Copyright 2016-2018 Centro de Investigaciones Energéticas, Medioambientales y Tecnológicas @@ -17,8 +16,15 @@ # See the Licence for the specific language governing permissions and limitations # under the Licence. -from .bolometry import BolometerCamera, BolometerFoil, BolometerSlit, BolometerIRVB +from .bolometry import BolometerCamera, BolometerFoil, BolometerIRVB, BolometerSlit from .calcam import load_calcam_calibration +from .group import ( + FibreOpticGroup, + PixelGroup, + SightLineGroup, + SpectroscopicFibreOpticGroup, + SpectroscopicSightLineGroup, + TargetedPixelGroup, +) from .intersections import find_wall_intersection -from .spectroscopy import SpectroscopicSightLine, SpectroscopicFibreOptic -from .group import PixelGroup, TargettedPixelGroup, SightLineGroup, FibreOpticGroup, SpectroscopicFibreOpticGroup, SpectroscopicSightLineGroup +from .spectroscopy import SpectroscopicFibreOptic, SpectroscopicSightLine diff --git a/cherab/tools/observers/bolometry.py b/cherab/tools/observers/bolometry.py index 5cd8c1a2..532d2826 100644 --- a/cherab/tools/observers/bolometry.py +++ b/cherab/tools/observers/bolometry.py @@ -1,4 +1,3 @@ - # Copyright 2016-2018 Euratom # Copyright 2016-2018 United Kingdom Atomic Energy Authority # Copyright 2016-2018 Centro de Investigaciones Energéticas, Medioambientales y Tecnológicas @@ -17,23 +16,32 @@ # See the Licence for the specific language governing permissions and limitations # under the Licence. -from enum import Enum import functools -import numpy as np +from enum import Enum -from raysect.core import Node, translate, rotate_basis, Point3D, Vector3D, Ray as CoreRay, Primitive, World -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, TargetedPixel -from raysect.optical.observer import PowerPipeline2D, RadiancePipeline2D, \ - SpectralPowerPipeline2D, SpectralRadiancePipeline2D, TargetedCCDArray -from raysect.optical.material.material import NullMaterial +import numpy as np +from raysect.core import Node, Point3D, Primitive, Vector3D, World, rotate_basis, translate +from raysect.core import Ray as CoreRay +from raysect.core.math.sampler import RectangleSampler3D, TargetedHemisphereSampler from raysect.optical.material import AbsorbingSurface +from raysect.optical.material.material import NullMaterial +from raysect.optical.observer import ( + PowerPipeline0D, + PowerPipeline2D, + RadiancePipeline0D, + RadiancePipeline2D, + SightLine, + SpectralPowerPipeline0D, + SpectralPowerPipeline2D, + SpectralRadiancePipeline0D, + SpectralRadiancePipeline2D, + TargetedCCDArray, + TargetedPixel, +) +from raysect.primitive import Box, Cylinder, Subtract, Union from cherab.tools.inversions.voxels import VoxelCollection - R_2_PI = 1 / (2 * np.pi) @@ -71,8 +79,7 @@ class BolometerCamera(Node): >>> camera = BolometerCamera(name="MyBolometer", parent=world) """ - def __init__(self, camera_geometry=None, parent=None, transform=None, name=''): - + def __init__(self, camera_geometry=None, parent=None, transform=None, name=""): super().__init__(parent=parent, transform=transform, name=name) self._foil_detectors = [] @@ -133,12 +140,8 @@ def foil_detectors(self): @foil_detectors.setter def foil_detectors(self, value): - if not isinstance(value, list): - raise TypeError( - "The foil_detectors attribute of BolometerCamera must be a list of " - "BolometerFoils or BolometerIRVBs." - ) + raise TypeError("The foil_detectors attribute of BolometerCamera must be a list of BolometerFoils or BolometerIRVBs.") # Prevent external changes being made to this list value = value.copy() @@ -148,8 +151,8 @@ def foil_detectors(self, value): "The foil_detectors attribute of BolometerCamera must be a list of " "BolometerFoil or BolometerIRVB objects. Value {} is not a BolometerFoil " "or BolometerIRVB.".format(foil_detector) - ) - if not foil_detector.slit in self._slits: + ) + if foil_detector.slit not in self._slits: self._slits.append(foil_detector.slit) foil_detector.parent = self @@ -167,11 +170,9 @@ def add_foil_detector(self, foil_detector): """ if not isinstance(foil_detector, (BolometerFoil, BolometerIRVB)): - raise TypeError( - "The foil_detector argument must be of type BolometerFoil or BolometerIRVB." - ) + raise TypeError("The foil_detector argument must be of type BolometerFoil or BolometerIRVB.") - if not foil_detector.slit in self._slits: + if foil_detector.slit not in self._slits: self._slits.append(foil_detector.slit) foil_detector.parent = self @@ -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. @@ -255,9 +256,7 @@ class BolometerSlit(Node): >>> slit = BolometerSlit("slit", centre_point, basis_x, dx, basis_y, dy, parent=camera) """ - def __init__(self, slit_id, centre_point, basis_x, dx, basis_y, dy, dz=0.001, - parent=None, csg_aperture=False, curvature_radius=0): - + def __init__(self, slit_id, centre_point, basis_x, dx, basis_y, dy, dz=0.001, parent=None, csg_aperture=False, curvature_radius=0): # perform validation of input parameters if not isinstance(dx, (float, int)): @@ -274,11 +273,9 @@ def __init__(self, slit_id, centre_point, basis_x, dx, basis_y, dy, dz=0.001, raise TypeError("centre_point argument for BolometerSlit must be of type Point3D.") if not isinstance(curvature_radius, (float, int)): - raise TypeError("curvature_radius argument for BolometerSlit " - "must be of type float/int.") + raise TypeError("curvature_radius argument for BolometerSlit must be of type float/int.") if curvature_radius < 0: - raise ValueError("curvature_radius argument for BolometerSlit " - "must not be negative.") + raise ValueError("curvature_radius argument for BolometerSlit must not be negative.") if not isinstance(basis_x, Vector3D): raise TypeError("The basis vectors of BolometerSlit must be of type Vector3D.") @@ -300,8 +297,14 @@ def __init__(self, slit_id, centre_point, basis_x, dx, basis_y, dy, dz=0.001, super().__init__(parent=parent, transform=transform, name=slit_id) - self.target = Box(lower=Point3D(-dx/2*1.01, -dy/2*1.01, -dz/2), upper=Point3D(dx/2*1.01, dy/2*1.01, dz/2), - transform=None, material=NullMaterial(), parent=self, name=slit_id+' - target') + self.target = Box( + lower=Point3D(-dx / 2 * 1.01, -dy / 2 * 1.01, -dz / 2), + upper=Point3D(dx / 2 * 1.01, dy / 2 * 1.01, dz / 2), + transform=None, + material=NullMaterial(), + parent=self, + name=slit_id + " - target", + ) self._csg_aperture = None self.csg_aperture = csg_aperture @@ -332,14 +335,14 @@ def csg_aperture(self): @csg_aperture.setter def csg_aperture(self, value): - if value is True: width = max(self.dx, self.dy) - face = Box(Point3D(-width, -width, -self.dz/2), Point3D(width, width, self.dz/2)) - slit = Box(lower=Point3D(-self.dx/2, -self.dy/2, -self.dz/2 - self.dz*0.1), - upper=Point3D(self.dx/2, self.dy/2, self.dz/2 + self.dz*0.1)) - self._csg_aperture = Subtract(face, slit, parent=self, - material=AbsorbingSurface(), name=self.name+' - CSG Aperture') + face = Box(Point3D(-width, -width, -self.dz / 2), Point3D(width, width, self.dz / 2)) + slit = Box( + lower=Point3D(-self.dx / 2, -self.dy / 2, -self.dz / 2 - self.dz * 0.1), + upper=Point3D(self.dx / 2, self.dy / 2, self.dz / 2 + self.dz * 0.1), + ) + self._csg_aperture = Subtract(face, slit, parent=self, material=AbsorbingSurface(), name=self.name + " - CSG Aperture") else: if isinstance(self._csg_aperture, Primitive): @@ -403,9 +406,9 @@ class BolometerFoil(TargetedPixel): >>> detector = BolometerFoil("ch#1", centre_point, basis_x, dx, basis_y, dy, slit, parent=camera) """ - def __init__(self, detector_id, centre_point, basis_x, dx, basis_y, dy, slit, - parent=None, units="Power", accumulate=False, curvature_radius=0): - + def __init__( + self, detector_id, centre_point, basis_x, dx, basis_y, dy, slit, parent=None, units="Power", accumulate=False, curvature_radius=0 + ): # perform validation of input parameters if not isinstance(dx, (float, int)): @@ -425,11 +428,9 @@ def __init__(self, detector_id, centre_point, basis_x, dx, basis_y, dy, slit, raise TypeError("centre_point argument for BolometerFoil must be of type Point3D.") if not isinstance(curvature_radius, (float, int)): - raise TypeError("curvature_radius argument for BolometerFoil " - "must be of type float/int.") + raise TypeError("curvature_radius argument for BolometerFoil must be of type float/int.") if curvature_radius < 0: - raise ValueError("curvature_radius argument for BolometerFoil " - "must not be negative.") + raise ValueError("curvature_radius argument for BolometerFoil must not be negative.") if not isinstance(basis_x, Vector3D): raise TypeError("The basis vectors of BolometerFoil must be of type Vector3D.") @@ -447,9 +448,18 @@ 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], 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) + 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, + ) # Update pipeline based on units self.units = units @@ -530,16 +540,13 @@ def as_sightline(self): else: raise ValueError("The units argument of BolometerFoil must be one of 'Power' or 'Radiance'.") - los_observer = SightLine(pipelines=[pipeline], pixel_samples=1, quiet=True, - parent=self, name=self.name) + los_observer = SightLine(pipelines=[pipeline], pixel_samples=1, quiet=True, parent=self, name=self.name) los_observer.render_engine = self.render_engine los_observer.spectral_bins = self.spectral_bins los_observer.min_wavelength = self.min_wavelength los_observer.max_wavelength = self.max_wavelength # The observer's Z axis should be aligned along the line of sight vector - los_observer.transform = rotate_basis( - self.sightline_vector.transform(self.to_local()), self.basis_y - ) + los_observer.transform = rotate_basis(self.sightline_vector.transform(self.to_local()), self.basis_y) return los_observer @@ -560,7 +567,6 @@ def trace_sightline(self): direction = self.sightline_vector while True: - # Find the next intersection point of the ray with the world intersection = self.root.hit(CoreRay(origin, direction)) @@ -661,8 +667,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 = TargetedHemisphereSampler(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 +677,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) @@ -758,14 +764,10 @@ class BolometerIRVB(TargetedCCDArray): >>> detector = BolometerIRVB("irvb", width, pixels, slit, transform, parent=camera) """ - _PIPELINES = {_Units.POWER: PowerPipeline2D, - _Units.RADIANCE: RadiancePipeline2D} - _SPECTRAL_PIPELINES = {_Units.POWER: SpectralPowerPipeline2D, - _Units.RADIANCE: SpectralRadiancePipeline2D} - - def __init__(self, name, width, pixels, slit, transform, parent=None, - units="power", accumulate=False, curvature_radius=0): + _PIPELINES = {_Units.POWER: PowerPipeline2D, _Units.RADIANCE: RadiancePipeline2D} + _SPECTRAL_PIPELINES = {_Units.POWER: SpectralPowerPipeline2D, _Units.RADIANCE: SpectralRadiancePipeline2D} + def __init__(self, name, width, pixels, slit, transform, parent=None, units="power", accumulate=False, curvature_radius=0): # perform validation of input parameters width = float(width) if width < 0: @@ -776,16 +778,15 @@ def __init__(self, name, width, pixels, slit, transform, parent=None, curvature_radius = float(curvature_radius) if curvature_radius < 0: - raise ValueError("curvature_radius argument for BolometerIRVB " - "must not be negative.") + raise ValueError("curvature_radius argument for BolometerIRVB must not be negative.") self._slit = slit self._curvature_radius = curvature_radius self._accumulate = None # Will be set after pipeline is created. - super().__init__([slit.target], pixels=pixels, width=width, - targeted_path_prob=0.99, parent=parent, pipelines=[], - transform=transform, name=name) + super().__init__( + [slit.target], pixels=pixels, width=width, targeted_path_prob=0.99, parent=parent, pipelines=[], transform=transform, name=name + ) self.pixel_samples = 1000 self.spectral_bins = 1 self.quiet = True @@ -815,18 +816,22 @@ def pixels_as_foils(self): for x in range(nx): pixel_column = [] for y in range(ny): - pixel_centre = (foil_bottom_left - + (x + 0.5) * XAXIS * pixel_width - + (y + 0.5) * YAXIS * pixel_height) + pixel_centre = foil_bottom_left + (x + 0.5) * XAXIS * pixel_width + (y + 0.5) * YAXIS * pixel_height pixel = BolometerFoil( detector_id="IRVB pixel ({},{})".format(x + 1, y + 1), - centre_point=pixel_centre, basis_x=XAXIS, dx=pixel_width, - basis_y=YAXIS, dy=pixel_height, slit=self._slit, - units=self._units.value.capitalize(), accumulate=False, parent=self + centre_point=pixel_centre, + basis_x=XAXIS, + dx=pixel_width, + basis_y=YAXIS, + dy=pixel_height, + slit=self._slit, + units=self._units.value.capitalize(), + accumulate=False, + parent=self, ) pixel_column.append(pixel) pixels.append(pixel_column) - return np.asarray(pixels, dtype='object') + return np.asarray(pixels, dtype="object") @property def height(self): @@ -851,9 +856,8 @@ def basis_y(self): @property def sightline_vectors(self): return np.asarray( - [[pixel.centre_point.vector_to(self._slit.centre_point) for pixel in pixel_column] - for pixel_column in self.pixels_as_foils], - dtype='object' + [[pixel.centre_point.vector_to(self._slit.centre_point) for pixel in pixel_column] for pixel_column in self.pixels_as_foils], + dtype="object", ) @property @@ -876,8 +880,7 @@ def units(self, units): self._units = _Units.RADIANCE else: raise ValueError( - "The units property of BolometerIRVB must be one of {}" - .format([member.value for member in _Units.__members__]) + "The units property of BolometerIRVB must be one of {}".format([member.value for member in _Units.__members__]) ) pipeline_class = self._PIPELINES[self._units] pipeline = pipeline_class(accumulate=self.accumulate) @@ -904,7 +907,7 @@ def as_sightlines(self): """ pixels = self.pixels_as_foils sightlines = [[pixel.as_sightline() for pixel in pixel_column] for pixel_column in pixels] - return np.asarray(sightlines, dtype='object') + return np.asarray(sightlines, dtype="object") def trace_sightlines(self): """ @@ -918,7 +921,7 @@ def trace_sightlines(self): """ pixels = self.pixels_as_foils traces = [[pixel.trace_sightline() for pixel in pixel_column] for pixel_column in pixels] - return np.asarray(traces, dtype='object') + return np.asarray(traces, dtype="object") def calculate_sensitivity(self, voxel_collection, ray_count=None): r""" @@ -1043,26 +1046,24 @@ def mask_corners(element): # Make the elements to cut out from the cover slightly thicker than the # cover, to guard against rounding errors - long_box = Box(lower=Point3D(-dx/2 + rc, -dy/2, -0.5 * dz), - upper=Point3D(dx/2 - rc, dy/2, 1.5 * dz)) - shot_box = Box(lower=Point3D(-dx/2, -dy/2 + rc, -0.5 * dz), - upper=Point3D(dx/2, dy/2 - rc, 1.5 * dz)) + long_box = Box(lower=Point3D(-dx / 2 + rc, -dy / 2, -0.5 * dz), upper=Point3D(dx / 2 - rc, dy / 2, 1.5 * dz)) + shot_box = Box(lower=Point3D(-dx / 2, -dy / 2 + rc, -0.5 * dz), upper=Point3D(dx / 2, dy / 2 - rc, 1.5 * dz)) cylinder_template = Cylinder(radius=rc, height=2 * dz) top_left_cylinder = cylinder_template.instance() - top_left_cylinder.transform = translate(-dx/2 + rc, dy/2 - rc, -dz/2) + top_left_cylinder.transform = translate(-dx / 2 + rc, dy / 2 - rc, -dz / 2) top_right_cylinder = cylinder_template.instance() - top_right_cylinder.transform = translate(dx/2 - rc, dy/2 - rc, -dz/2) + top_right_cylinder.transform = translate(dx / 2 - rc, dy / 2 - rc, -dz / 2) bottom_right_cylinder = cylinder_template.instance() - bottom_right_cylinder.transform = translate(dx/2 - rc, -dy/2 + rc, -dz/2) + bottom_right_cylinder.transform = translate(dx / 2 - rc, -dy / 2 + rc, -dz / 2) bottom_left_cylinder = cylinder_template.instance() - bottom_left_cylinder.transform = translate(-dx/2 + rc, -dy/2 + rc, -dz/2) - cutout = functools.reduce(Union, (long_box, shot_box, top_left_cylinder, - top_right_cylinder, bottom_right_cylinder, - bottom_left_cylinder)) - cover = Box(lower=Point3D(-dx/2, -dy/2, 0), upper=Point3D(dx/2, dy/2, dz)) + bottom_left_cylinder.transform = translate(-dx / 2 + rc, -dy / 2 + rc, -dz / 2) + cutout = functools.reduce( + Union, (long_box, shot_box, top_left_cylinder, top_right_cylinder, bottom_right_cylinder, bottom_left_cylinder) + ) + cover = Box(lower=Point3D(-dx / 2, -dy / 2, 0), upper=Point3D(dx / 2, dy / 2, dz)) mask = Subtract(cover, cutout) mask.material = AbsorbingSurface() mask.transform = translate(0, 0, dz) - mask.name = element.name + ' - rounded edges mask' + mask.name = element.name + " - rounded edges mask" mask.parent = element diff --git a/cherab/tools/observers/group/__init__.py b/cherab/tools/observers/group/__init__.py index eca93585..8a24f80f 100644 --- a/cherab/tools/observers/group/__init__.py +++ b/cherab/tools/observers/group/__init__.py @@ -18,6 +18,6 @@ from .fibreoptic import FibreOpticGroup from .sightline import SightLineGroup -from .targettedpixel import TargettedPixelGroup +from .targetedpixel import TargetedPixelGroup from .pixel import PixelGroup from .spectroscopic import SpectroscopicFibreOpticGroup, SpectroscopicSightLineGroup diff --git a/cherab/tools/observers/group/targettedpixel.py b/cherab/tools/observers/group/targettedpixel.py index e440f70e..a9f8775e 100644 --- a/cherab/tools/observers/group/targettedpixel.py +++ b/cherab/tools/observers/group/targettedpixel.py @@ -22,11 +22,11 @@ from .base import Observer0DGroup -class TargettedPixelGroup(Observer0DGroup): +class TargetedPixelGroup(Observer0DGroup): """ A group of targeted pixel under a single scene-graph node. - A scene-graph object regrouping a series of 'TargettedPixel' + A scene-graph object regrouping a series of `TargetedPixel` observers as a scene-graph parent. Allows combined observation and display control simultaneously. @@ -49,8 +49,9 @@ def x_width(self, value): 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))) + 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 @@ -66,8 +67,9 @@ def y_width(self, value): 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))) + 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 @@ -92,8 +94,11 @@ def targets(self, value): 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))) + 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: @@ -110,7 +115,9 @@ def targeted_path_prob(self, value): 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))) + 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/tests/test_observer_groups.py b/cherab/tools/tests/test_observer_groups.py index a04f095c..f92ab7d3 100644 --- a/cherab/tools/tests/test_observer_groups.py +++ b/cherab/tools/tests/test_observer_groups.py @@ -1,11 +1,11 @@ import unittest from raysect.core.workflow import RenderEngine -from raysect.optical.observer import Observer0D, SightLine, FibreOptic, Pixel, TargetedPixel, PowerPipeline0D, SpectralPowerPipeline0D +from raysect.optical.observer import FibreOptic, Observer0D, Pixel, PowerPipeline0D, SightLine, SpectralPowerPipeline0D, TargetedPixel from raysect.primitive import Sphere +from cherab.tools.observers.group import FibreOpticGroup, PixelGroup, SightLineGroup, TargetedPixelGroup from cherab.tools.observers.group.base import Observer0DGroup -from cherab.tools.observers.group import SightLineGroup, FibreOpticGroup, PixelGroup, TargettedPixelGroup from cherab.tools.raytransfer import pipelines @@ -20,13 +20,13 @@ def setUp(self): def test_get_item(self): """Tests all inputs for the __get_item__ method""" group = self._GROUP_CLASS(observers=self.observers) - names = ['zero', 'one', 'two'] + names = ["zero", "one", "two"] group.names = names 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]) @@ -37,11 +37,11 @@ def test_get_item(self): group[1.2] with self.assertRaises(ValueError): - group['fail'] + group["fail"] - group.names = ['fail'] * len(group) + group.names = ["fail"] * len(group) with self.assertRaises(ValueError): - group['fail'] + group["fail"] def test_assignments(self): """Test assignments of all supported attributes of Observer0DGroup""" @@ -49,7 +49,7 @@ def test_assignments(self): group.observers = self.observers for grouped_observer, input_observer in zip(group.observers, self.observers): - self.assertIs(grouped_observer, input_observer, msg='Observers do not match') + self.assertIs(grouped_observer, input_observer, msg="Observers do not match") with self.assertRaises(ValueError): group.observers = [Sphere()] @@ -58,32 +58,32 @@ def test_assignments(self): group.observers = Sphere() # names - names = ['zero', 'one', 'two'] + names = ["zero", "one", "two"] group.names = names for grouped_observer, input_name in zip(group.observers, names): - self.assertEqual(grouped_observer.name, input_name, msg='Observer name do not match') + self.assertEqual(grouped_observer.name, input_name, msg="Observer name do not match") with self.assertRaises(ValueError): - group.names = ['fail'] + group.names = ["fail"] with self.assertRaises(TypeError): - group.names = 'fail' + group.names = "fail" # pipelines - ppln_0 = PowerPipeline0D(name='pipeline zero, observer zero') - ppln_1 = PowerPipeline0D(name='pipeline one, observer one') - ppln_2 = PowerPipeline0D(name='pipeline two, observer two') - ppln_3 = PowerPipeline0D(name='pipeline three, observer two') + ppln_0 = PowerPipeline0D(name="pipeline zero, observer zero") + ppln_1 = PowerPipeline0D(name="pipeline one, observer one") + ppln_2 = PowerPipeline0D(name="pipeline two, observer two") + ppln_3 = PowerPipeline0D(name="pipeline three, observer two") pipelist = [[ppln_0], [ppln_1], [ppln_2, ppln_3]] group.pipelines = pipelist - self.assertIs(group[0].pipelines[0], ppln_0, 'non matching pipeline') - self.assertIs(group[1].pipelines[0], ppln_1, 'non matching pipeline') - self.assertIs(group[2].pipelines[0], ppln_2, 'non matching pipeline') - self.assertIs(group[2].pipelines[1], ppln_3, 'non matching pipeline') + self.assertIs(group[0].pipelines[0], ppln_0, "non matching pipeline") + self.assertIs(group[1].pipelines[0], ppln_1, "non matching pipeline") + self.assertIs(group[2].pipelines[0], ppln_2, "non matching pipeline") + self.assertIs(group[2].pipelines[1], ppln_3, "non matching pipeline") 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,15 +102,15 @@ 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 self.assertListEqual(group.min_wavelength, [wvl - 100] * len(group)) self.assertListEqual(group.max_wavelength, [wvl + 100] * len(group)) - min_wvls = [90 + 10*i for i in range(len(group))] - max_wvls = [100 + 10*i for i in range(len(group))] + min_wvls = [90 + 10 * i for i in range(len(group))] + max_wvls = [100 + 10 * i for i in range(len(group))] group.min_wavelength = min_wvls group.max_wavelength = max_wvls self.assertListEqual(group.min_wavelength, min_wvls) @@ -122,7 +122,7 @@ def test_assignments(self): group.min_wavelength = [90] * (len(group) - 1) # spectral - bins = [200 + i*100 for i in range(len(group))] + bins = [200 + i * 100 for i in range(len(group))] rays = [2] * len(group) group.spectral_bins = bins group.spectral_rays = rays @@ -139,7 +139,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,8 +152,8 @@ def test_assignments(self): with self.assertRaises(ValueError): group.quiet = [False] * (len(group) + 1) - # rays - probs = [0.2 + i*0.1 for i in range(len(group))] + # 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))] sampling = [False] * len(group) @@ -196,10 +196,10 @@ 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))] + pixel_samples = [2000 + i * 500 for i in range(len(group))] + per_task = [5000 + i * 100 for i in range(len(group))] group.pixel_samples = pixel_samples group.samples_per_task = per_task self.assertListEqual(group.pixel_samples, pixel_samples) @@ -228,7 +228,7 @@ def test_connect_pipelines(self): group = self._GROUP_CLASS(observers=self.observers) ppln_classes = [PowerPipeline0D, SpectralPowerPipeline0D] - names = ['power', 'spectral'] + names = ["power", "spectral"] keywords = [ dict(name=names[0]), dict(name=names[1], display_progress=True), @@ -348,8 +348,8 @@ 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 = [TargetedPixel(targets=[Sphere()], pipelines=[PowerPipeline0D()]) for _ in range(self._NUM)] @@ -374,7 +374,7 @@ def test_targets(self): with self.assertRaises(ValueError): group.targets = targets - # targetted path prob + # targeted path prob prob = [0.9, 0.95, 1] group.targeted_path_prob = prob self.assertListEqual(group.targeted_path_prob, prob) diff --git a/docs/source/tools/observers.rst b/docs/source/tools/observers.rst index e93e2057..84f94f48 100644 --- a/docs/source/tools/observers.rst +++ b/docs/source/tools/observers.rst @@ -109,10 +109,10 @@ combined into a group. Group observers --------------- -Group observer is a collection of observers of the same type. All Observer0D classes -defined in Raysect are supoorted. The parameters of individual observers in a group +Group observer is a collection of observers of the same type. All Observer0D classes +defined in Raysect are supoorted. The parameters of individual observers in a group may differ. Group observer allows combined observation, namely, calling the observe -function for a group leads to a sequential call of this function for each observer +function for a group leads to a sequential call of this function for each observer in the group. .. autoclass:: cherab.tools.observers.group.base.Observer0DGroup @@ -127,7 +127,7 @@ in the group. .. autoclass:: cherab.tools.observers.group.PixelGroup :members: -.. autoclass:: cherab.tools.observers.group.TargettedPixelGroup +.. autoclass:: cherab.tools.observers.group.TargetedPixelGroup :members: Spectroscopic Groups @@ -136,9 +136,9 @@ Spectroscopic Groups .. deprecated:: 1.4.0 Use groups based on Raysect's observer classes instead -These groups take control of spectroscopic lines of sight observers. They support -direction and origin positioning and contain methods for plotting the power and -spectrum. Originally, these were called group observers and did not include the +These groups take control of spectroscopic lines of sight observers. They support +direction and origin positioning and contain methods for plotting the power and +spectrum. Originally, these were called group observers and did not include the Spectroscopic prefix in class name. .. autoclass:: cherab.tools.observers.SpectroscopicSightLine From 662eb0447fda4fd80a925a1b804d2a19b77c28ab Mon Sep 17 00:00:00 2001 From: munechika-koyo Date: Mon, 13 Oct 2025 11:45:35 +0200 Subject: [PATCH 20/91] Rename filename --- .../tools/observers/group/{targettedpixel.py => targetedpixel.py} | 0 1 file changed, 0 insertions(+), 0 deletions(-) rename cherab/tools/observers/group/{targettedpixel.py => targetedpixel.py} (100%) diff --git a/cherab/tools/observers/group/targettedpixel.py b/cherab/tools/observers/group/targetedpixel.py similarity index 100% rename from cherab/tools/observers/group/targettedpixel.py rename to cherab/tools/observers/group/targetedpixel.py From af3d5c2bc660362eec30d193d9b2d8eaf2ecb40c Mon Sep 17 00:00:00 2001 From: munechika-koyo Date: Mon, 13 Oct 2025 11:54:45 +0200 Subject: [PATCH 21/91] Fix numpy v2 error related thing --- cherab/tools/tests/test_voxels.py | 3 +-- 1 file changed, 1 insertion(+), 2 deletions(-) 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) From 3d53694233ece459c60b258eb73827f2f3de4ef2 Mon Sep 17 00:00:00 2001 From: Koyo MUNECHIKA <51052381+munechika-koyo@users.noreply.github.com> Date: Mon, 13 Oct 2025 12:06:57 +0200 Subject: [PATCH 22/91] Update ci.yml --- .github/workflows/ci.yml | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index 73fcdbd9..25e2467d 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -22,7 +22,7 @@ jobs: with: python-version: ${{ matrix.python-version }} - name: Install Python dependencies - run: python -m pip install --prefer-binary cython~=3.1 numpy>=2 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.9.* - name: Build cherab From 58708e0337f6ac38a47c7ff10764408ce924d10c Mon Sep 17 00:00:00 2001 From: munechika-koyo Date: Mon, 13 Oct 2025 13:15:00 +0200 Subject: [PATCH 23/91] Remove Python 3.14 from CI matrix until pyopencl support is available --- .github/workflows/ci.yml | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index 25e2467d..77b05087 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -11,7 +11,7 @@ jobs: strategy: fail-fast: false matrix: - python-version: ["3.9", "3.10", "3.11", "3.12", "3.13", "3.14"] + python-version: ["3.9", "3.10", "3.11", "3.12", "3.13"] #, "3.14"] # TODO: re-enable 3.14 when pyopencl supports it steps: - name: Checkout code uses: actions/checkout@v2 From 1a63b6d35a21a5af5cda4a7a76efe6d37e6efe21 Mon Sep 17 00:00:00 2001 From: munechika-koyo Date: Tue, 14 Oct 2025 17:21:15 +0200 Subject: [PATCH 24/91] Update CHANGELOG.md for Release 1.6.0: add API changes and new features --- CHANGELOG.md | 8 ++++++-- 1 file changed, 6 insertions(+), 2 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index 4f22a5d4..6a4b7e88 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -3,11 +3,15 @@ Project Changelog Release 1.6.0 (TBD) ------------------- +API changes: +* Rename 'targetted' to 'targeted' following Raysect change. (#486) New: * 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) +* 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) Release 1.5.0 (27 Aug 2024) ------------------- @@ -134,7 +138,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. @@ -162,7 +166,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) From fde76c4daa79bf37c20c98411940c761f65e9c34 Mon Sep 17 00:00:00 2001 From: munechika-koyo Date: Thu, 16 Oct 2025 10:00:28 +0200 Subject: [PATCH 25/91] Revert "Fix remaining typo `targetted` to `targeted`" This reverts commit 10d690feb7abc365acbe4383a17bd1699957f14c. --- cherab/core/model/laser/profile.pyx | 60 +++--- cherab/tools/observers/__init__.py | 14 +- cherab/tools/observers/bolometry.py | 201 +++++++++--------- cherab/tools/observers/group/__init__.py | 2 +- cherab/tools/observers/group/targetedpixel.py | 25 +-- cherab/tools/tests/test_observer_groups.py | 70 +++--- docs/source/tools/observers.rst | 14 +- 7 files changed, 186 insertions(+), 200 deletions(-) diff --git a/cherab/core/model/laser/profile.pyx b/cherab/core/model/laser/profile.pyx index 80b73d02..82376980 100644 --- a/cherab/core/model/laser/profile.pyx +++ b/cherab/core/model/laser/profile.pyx @@ -4,7 +4,7 @@ from raysect.primitive import Cylinder from raysect.optical cimport Spectrum, Vector3D, translate from cherab.core.laser cimport Laser, LaserProfile -from cherab.core.model.laser.math_functions cimport ConstantAxisymmetricGaussian3D, ConstantBivariateGaussian3D, TrivariateGaussian3D, GaussianBeamModel +from cherab.core.model.laser.math_functions cimport ConstantAxisymmetricGaussian3D, ConstantBivariateGaussian3D, TrivariateGaussian3D, GaussianBeamModel from cherab.core.utility.constants cimport SPEED_OF_LIGHT @@ -22,20 +22,20 @@ cdef class UniformEnergyDensity(LaserProfile): The methods get_pointing, get_polarization and get_energy_density are not limited to the inside of the laser cylinder. If called alone for position (x, y, z) outisde the laser cylinder, they will still return non-zero values. - + In the following example, a laser of length of 2 m (extending from z=0 to z=2 m) with a radius of 3 cm and volumetric energy density of 5 J*m^-3 and polarisation in the y direction is created: .. code-block:: pycon - + >>> from raysect.core import Vector3D >>> from cherab.core.model.laser import UniformEnergyDensity - + >>> energy = 5 # energy density in J >>> radius = 3e-2 # laser radius in m >>> length = 2 # laser length in m >>> polarisation = Vector3D(0, 1, 0) # polarisation direction - + # create the laser profile >>> laser_profile = UniformEnergyDensity(energy, radius, length, polarisation) @@ -108,7 +108,7 @@ cdef class UniformEnergyDensity(LaserProfile): cpdef list generate_geometry(self): return generate_segmented_cylinder(self.laser_radius, self.laser_length) - + cdef class ConstantBivariateGaussian(LaserProfile): """ @@ -120,8 +120,8 @@ cdef class ConstantBivariateGaussian(LaserProfile): The model imitates a laser beam with a uniform power output within a single pulse. This results in the distribution of the energy density along the propagation direction of the laser (z-axis) to be also uniform. The integral value of laser energy Exy in an x-y plane is given by - - .. math:: + + .. math:: E_{xy} = \\frac{E_p}{(c * \\tau)}, where Ep is the energy of the laser pulse, tau is the temporal pulse length and c is the speed of light in vacuum. @@ -133,23 +133,23 @@ cdef class ConstantBivariateGaussian(LaserProfile): The sigma_x and sigma_y are standard deviations in x and y directions, respectively. .. note:: - The height of the cylinder, forming the laser beam, is given by the laser_length and is independent from the + The height of the cylinder, forming the laser beam, is given by the laser_length and is independent from the temporal length of the laser pulse given by pulse_length. This gives the possibility to independently control the size of the laser primitive and the value of the volumetric energy density. - + The methods get_pointing, get_polarization and get_energy_density are not limited to the inside of the laser cylinder. If called for position (x, y, z) outisde the laser cylinder, they can still return non-zero values. - + The following example shows how to create a laser with sigma_x= 1 cm and sigma_y=2 cm, which makes the laser profile in x-y plane to be elliptical. The pulse energy is 5 J and the laser temporal pulse length is 10 ns: .. code-block:: pycon - + >>> from raysect.core import Vector3D >>> from cherab.core.model.laser import ConstantBivariateGaussian - + >>> radius = 3e-2 # laser radius in m >>> length = 2 # laser length in m >>> polarisation = Vector3D(0, 1, 0) # polarisation direction @@ -157,7 +157,7 @@ cdef class ConstantBivariateGaussian(LaserProfile): >>> pulse_length = 1e-8 # pulse length in s >>> width_x = 1e-2 # standard deviation in x direction in m >>> width_y = 2e-2 # standard deviation in y direction in m - + # create the laser profile >>> laser_profile = ConstantBivariateGaussian(pulse_energy, pulse_length, radius, length, width_x, width_y, polarisation) @@ -323,7 +323,7 @@ cdef class TrivariateGaussian(LaserProfile): The sigma_x and sigma_y are standard deviations in x and y directions, respectively, and E_p is the energy deliverd by laser in a single laser pulse. The mu_z is the mean of the distribution in the z direction and controls th position of the laser pulse along the z direction. - The standard deviation in z direction sigma_z is calculated from the pulse length tau_p, which is the + The standard deviation in z direction sigma_z is calculated from the pulse length tau_p, which is the standard deviation of the Gaussian distributed ouput power of the laser within a single pulse: .. math:: @@ -332,24 +332,24 @@ cdef class TrivariateGaussian(LaserProfile): The c stands for the speed of light in vacuum. .. note:: - The height of the cylinder, forming the laser beam, is given by the laser_length and is independent from the + The height of the cylinder, forming the laser beam, is given by the laser_length and is independent from the temporal length of the laser pulse given by pulse_length. This gives the possibility to independently control the size of the laser primitive and the value of the volumetric energy density. - + The methods get_pointing, get_polarization and get_energy_density are not limited to the inside of the laser cylinder. If called alone for position (x, y, z) outisde the laser cylinder, they can still return non-zero values. - + The following example shows how to create a laser with sigma_x = 1 cm and sigma_y = 2 cm, which makes the laser profile in an x-y plane to be elliptical. The pulse energy is 5 J and the laser temporal pulse length is 10 ns. The position of the laser pulse maximum mean_z is set to 0.5: .. code-block:: pycon - + >>> from raysect.core import Vector3D >>> from cherab.core.model.laser import ConstantBivariateGaussian - + >>> radius = 3e-2 # laser radius in m >>> length = 2 # laser length in m >>> polarisation = Vector3D(0, 1, 0) # polarisation direction @@ -358,7 +358,7 @@ cdef class TrivariateGaussian(LaserProfile): >>> pulse_z = 0.5 # position of the pulse mean >>> width_x = 1e-2 # standard deviation in x direction in m >>> width_y = 2e-2 # standard deviation in y direction in m - + # create the laser profile >>> laser_profile = ConstantBivariateGaussian(pulse_energy, pulse_length, pulse_z, radius, length, width_x, width_y, polarisation) @@ -512,7 +512,7 @@ cdef class TrivariateGaussian(LaserProfile): self._distribution = TrivariateGaussian3D(self._mean_z, self._stddev_x, self._stddev_y, self._stddev_z) - normalisation = self._pulse_energy + normalisation = self._pulse_energy function = normalisation * self._distribution self.set_energy_density_function(function) @@ -541,16 +541,16 @@ cdef class GaussianBeamAxisymmetric(LaserProfile): .. math:: z_R = \\frac{\\pi \\omega_0^2 n}{\\lambda_l} - + where the omega_0 is the standard deviation in the xy plane in the focal point (beam waist) and lambda_l is the central wavelength of the laser. The E_xy stand for the laser energy in an xy plane and is calculated as: - + .. math:: E_{xy} = \\frac{E_p}{(c * \\tau)}, where the E_p is the energy in a single laser pulse and tau is the temporal pulse length. - .. note:: + .. note:: For more information about the Gaussian beam model see https://en.wikipedia.org/wiki/Gaussian_beam The methods get_pointing, get_polarization and get_energy_density are not limited to the inside @@ -562,10 +562,10 @@ cdef class GaussianBeamAxisymmetric(LaserProfile): waist is z=50 cm. The laser wavelength is 1060 nm. .. code-block:: pycon - + >>> from raysect.core import Vector3D >>> from cherab.core.model.laser import GaussianBeamAxisymmetric - + >>> radius = 5e-2 # laser radius in m >>> length = 2 # laser length in m >>> polarisation = Vector3D(0, 1, 0) # polarisation direction @@ -576,7 +576,7 @@ cdef class GaussianBeamAxisymmetric(LaserProfile): >>> width_x = 1e-2 # standard deviation in x direction in m >>> width_y = 2e-2 # standard deviation in y direction in m >>> laser_wlen = 1060 # laser wavelength in nm - + # create the laser profile >>> laser_profile = GaussianBeamAxisymmetric(pulse_energy, pulse_length, length, radius, waist_z, waist_width, laser_wlen) @@ -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 - targeted and importance sampling. The height of a cylinder segments is roughly + targetted and importance sampling. The height of a cylinder segments is roughly 2 * cylinder radius. :return: List of cylinders @@ -761,5 +761,5 @@ def generate_segmented_cylinder(radius, length): geometry.append(segment) else: raise ValueError("Incorrect number of segments calculated.") - + return geometry \ No newline at end of file diff --git a/cherab/tools/observers/__init__.py b/cherab/tools/observers/__init__.py index 670dfce5..d134ef63 100644 --- a/cherab/tools/observers/__init__.py +++ b/cherab/tools/observers/__init__.py @@ -1,3 +1,4 @@ + # Copyright 2016-2018 Euratom # Copyright 2016-2018 United Kingdom Atomic Energy Authority # Copyright 2016-2018 Centro de Investigaciones Energéticas, Medioambientales y Tecnológicas @@ -16,15 +17,8 @@ # See the Licence for the specific language governing permissions and limitations # under the Licence. -from .bolometry import BolometerCamera, BolometerFoil, BolometerIRVB, BolometerSlit +from .bolometry import BolometerCamera, BolometerFoil, BolometerSlit, BolometerIRVB from .calcam import load_calcam_calibration -from .group import ( - FibreOpticGroup, - PixelGroup, - SightLineGroup, - SpectroscopicFibreOpticGroup, - SpectroscopicSightLineGroup, - TargetedPixelGroup, -) from .intersections import find_wall_intersection -from .spectroscopy import SpectroscopicFibreOptic, SpectroscopicSightLine +from .spectroscopy import SpectroscopicSightLine, SpectroscopicFibreOptic +from .group import PixelGroup, TargettedPixelGroup, SightLineGroup, FibreOpticGroup, SpectroscopicFibreOpticGroup, SpectroscopicSightLineGroup diff --git a/cherab/tools/observers/bolometry.py b/cherab/tools/observers/bolometry.py index 532d2826..5cd8c1a2 100644 --- a/cherab/tools/observers/bolometry.py +++ b/cherab/tools/observers/bolometry.py @@ -1,3 +1,4 @@ + # Copyright 2016-2018 Euratom # Copyright 2016-2018 United Kingdom Atomic Energy Authority # Copyright 2016-2018 Centro de Investigaciones Energéticas, Medioambientales y Tecnológicas @@ -16,32 +17,23 @@ # See the Licence for the specific language governing permissions and limitations # under the Licence. -import functools from enum import Enum - +import functools import numpy as np -from raysect.core import Node, Point3D, Primitive, Vector3D, World, rotate_basis, translate -from raysect.core import Ray as CoreRay -from raysect.core.math.sampler import RectangleSampler3D, TargetedHemisphereSampler -from raysect.optical.material import AbsorbingSurface -from raysect.optical.material.material import NullMaterial -from raysect.optical.observer import ( - PowerPipeline0D, - PowerPipeline2D, - RadiancePipeline0D, - RadiancePipeline2D, - SightLine, - SpectralPowerPipeline0D, - SpectralPowerPipeline2D, - SpectralRadiancePipeline0D, - SpectralRadiancePipeline2D, - TargetedCCDArray, - TargetedPixel, -) + +from raysect.core import Node, translate, rotate_basis, Point3D, Vector3D, Ray as CoreRay, Primitive, World +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, TargetedPixel +from raysect.optical.observer import PowerPipeline2D, RadiancePipeline2D, \ + SpectralPowerPipeline2D, SpectralRadiancePipeline2D, TargetedCCDArray +from raysect.optical.material.material import NullMaterial +from raysect.optical.material import AbsorbingSurface from cherab.tools.inversions.voxels import VoxelCollection + R_2_PI = 1 / (2 * np.pi) @@ -79,7 +71,8 @@ class BolometerCamera(Node): >>> camera = BolometerCamera(name="MyBolometer", parent=world) """ - def __init__(self, camera_geometry=None, parent=None, transform=None, name=""): + def __init__(self, camera_geometry=None, parent=None, transform=None, name=''): + super().__init__(parent=parent, transform=transform, name=name) self._foil_detectors = [] @@ -140,8 +133,12 @@ def foil_detectors(self): @foil_detectors.setter def foil_detectors(self, value): + if not isinstance(value, list): - raise TypeError("The foil_detectors attribute of BolometerCamera must be a list of BolometerFoils or BolometerIRVBs.") + raise TypeError( + "The foil_detectors attribute of BolometerCamera must be a list of " + "BolometerFoils or BolometerIRVBs." + ) # Prevent external changes being made to this list value = value.copy() @@ -151,8 +148,8 @@ def foil_detectors(self, value): "The foil_detectors attribute of BolometerCamera must be a list of " "BolometerFoil or BolometerIRVB objects. Value {} is not a BolometerFoil " "or BolometerIRVB.".format(foil_detector) - ) - if foil_detector.slit not in self._slits: + ) + if not foil_detector.slit in self._slits: self._slits.append(foil_detector.slit) foil_detector.parent = self @@ -170,9 +167,11 @@ def add_foil_detector(self, foil_detector): """ if not isinstance(foil_detector, (BolometerFoil, BolometerIRVB)): - raise TypeError("The foil_detector argument must be of type BolometerFoil or BolometerIRVB.") + raise TypeError( + "The foil_detector argument must be of type BolometerFoil or BolometerIRVB." + ) - if foil_detector.slit not in self._slits: + if not foil_detector.slit in self._slits: self._slits.append(foil_detector.slit) foil_detector.parent = self @@ -214,7 +213,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 targeted_path_prob, + foil-slit distance and slit size, and also the foil's targetted_path_prob, this may not be guaranteed. Supplying a proper mesh geometry for the camera is recommended instead of using a CSG aperture. @@ -256,7 +255,9 @@ class BolometerSlit(Node): >>> slit = BolometerSlit("slit", centre_point, basis_x, dx, basis_y, dy, parent=camera) """ - def __init__(self, slit_id, centre_point, basis_x, dx, basis_y, dy, dz=0.001, parent=None, csg_aperture=False, curvature_radius=0): + def __init__(self, slit_id, centre_point, basis_x, dx, basis_y, dy, dz=0.001, + parent=None, csg_aperture=False, curvature_radius=0): + # perform validation of input parameters if not isinstance(dx, (float, int)): @@ -273,9 +274,11 @@ def __init__(self, slit_id, centre_point, basis_x, dx, basis_y, dy, dz=0.001, pa raise TypeError("centre_point argument for BolometerSlit must be of type Point3D.") if not isinstance(curvature_radius, (float, int)): - raise TypeError("curvature_radius argument for BolometerSlit must be of type float/int.") + raise TypeError("curvature_radius argument for BolometerSlit " + "must be of type float/int.") if curvature_radius < 0: - raise ValueError("curvature_radius argument for BolometerSlit must not be negative.") + raise ValueError("curvature_radius argument for BolometerSlit " + "must not be negative.") if not isinstance(basis_x, Vector3D): raise TypeError("The basis vectors of BolometerSlit must be of type Vector3D.") @@ -297,14 +300,8 @@ def __init__(self, slit_id, centre_point, basis_x, dx, basis_y, dy, dz=0.001, pa super().__init__(parent=parent, transform=transform, name=slit_id) - self.target = Box( - lower=Point3D(-dx / 2 * 1.01, -dy / 2 * 1.01, -dz / 2), - upper=Point3D(dx / 2 * 1.01, dy / 2 * 1.01, dz / 2), - transform=None, - material=NullMaterial(), - parent=self, - name=slit_id + " - target", - ) + self.target = Box(lower=Point3D(-dx/2*1.01, -dy/2*1.01, -dz/2), upper=Point3D(dx/2*1.01, dy/2*1.01, dz/2), + transform=None, material=NullMaterial(), parent=self, name=slit_id+' - target') self._csg_aperture = None self.csg_aperture = csg_aperture @@ -335,14 +332,14 @@ def csg_aperture(self): @csg_aperture.setter def csg_aperture(self, value): + if value is True: width = max(self.dx, self.dy) - face = Box(Point3D(-width, -width, -self.dz / 2), Point3D(width, width, self.dz / 2)) - slit = Box( - lower=Point3D(-self.dx / 2, -self.dy / 2, -self.dz / 2 - self.dz * 0.1), - upper=Point3D(self.dx / 2, self.dy / 2, self.dz / 2 + self.dz * 0.1), - ) - self._csg_aperture = Subtract(face, slit, parent=self, material=AbsorbingSurface(), name=self.name + " - CSG Aperture") + face = Box(Point3D(-width, -width, -self.dz/2), Point3D(width, width, self.dz/2)) + slit = Box(lower=Point3D(-self.dx/2, -self.dy/2, -self.dz/2 - self.dz*0.1), + upper=Point3D(self.dx/2, self.dy/2, self.dz/2 + self.dz*0.1)) + self._csg_aperture = Subtract(face, slit, parent=self, + material=AbsorbingSurface(), name=self.name+' - CSG Aperture') else: if isinstance(self._csg_aperture, Primitive): @@ -406,9 +403,9 @@ class BolometerFoil(TargetedPixel): >>> detector = BolometerFoil("ch#1", centre_point, basis_x, dx, basis_y, dy, slit, parent=camera) """ - def __init__( - self, detector_id, centre_point, basis_x, dx, basis_y, dy, slit, parent=None, units="Power", accumulate=False, curvature_radius=0 - ): + def __init__(self, detector_id, centre_point, basis_x, dx, basis_y, dy, slit, + parent=None, units="Power", accumulate=False, curvature_radius=0): + # perform validation of input parameters if not isinstance(dx, (float, int)): @@ -428,9 +425,11 @@ def __init__( raise TypeError("centre_point argument for BolometerFoil must be of type Point3D.") if not isinstance(curvature_radius, (float, int)): - raise TypeError("curvature_radius argument for BolometerFoil must be of type float/int.") + raise TypeError("curvature_radius argument for BolometerFoil " + "must be of type float/int.") if curvature_radius < 0: - raise ValueError("curvature_radius argument for BolometerFoil must not be negative.") + raise ValueError("curvature_radius argument for BolometerFoil " + "must not be negative.") if not isinstance(basis_x, Vector3D): raise TypeError("The basis vectors of BolometerFoil must be of type Vector3D.") @@ -448,18 +447,9 @@ def __init__( translation = translate(centre_point.x, centre_point.y, centre_point.z) rotation = rotate_basis(normal_vec, basis_y) - 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, - ) + 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) # Update pipeline based on units self.units = units @@ -540,13 +530,16 @@ def as_sightline(self): else: raise ValueError("The units argument of BolometerFoil must be one of 'Power' or 'Radiance'.") - los_observer = SightLine(pipelines=[pipeline], pixel_samples=1, quiet=True, parent=self, name=self.name) + los_observer = SightLine(pipelines=[pipeline], pixel_samples=1, quiet=True, + parent=self, name=self.name) los_observer.render_engine = self.render_engine los_observer.spectral_bins = self.spectral_bins los_observer.min_wavelength = self.min_wavelength los_observer.max_wavelength = self.max_wavelength # The observer's Z axis should be aligned along the line of sight vector - los_observer.transform = rotate_basis(self.sightline_vector.transform(self.to_local()), self.basis_y) + los_observer.transform = rotate_basis( + self.sightline_vector.transform(self.to_local()), self.basis_y + ) return los_observer @@ -567,6 +560,7 @@ def trace_sightline(self): direction = self.sightline_vector while True: + # Find the next intersection point of the ray with the world intersection = self.root.hit(CoreRay(origin, direction)) @@ -667,8 +661,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 targeted pixel sampler to sample directions - targeted_sampler = TargetedHemisphereSampler(spheres) + # instance targetted pixel sampler to sample directions + targetted_sampler = TargetedHemisphereSampler(spheres) # instance rectangle pixel sampler to sample origins point_sampler = RectangleSampler3D(width=self.x_width, height=self.y_width) @@ -677,8 +671,8 @@ def etendue_single_run(_): origins = point_sampler(samples=ray_count) passed = 0.0 for origin in origins: - # obtain targeted vector sample - direction, pdf = targeted_sampler(origin, pdf=True) + # obtain targetted vector sample + direction, pdf = targetted_sampler(origin, pdf=True) path_weight = R_2_PI * direction.z / pdf # Transform to world space origin = origin.transform(detector_transform) @@ -764,10 +758,14 @@ class BolometerIRVB(TargetedCCDArray): >>> detector = BolometerIRVB("irvb", width, pixels, slit, transform, parent=camera) """ - _PIPELINES = {_Units.POWER: PowerPipeline2D, _Units.RADIANCE: RadiancePipeline2D} - _SPECTRAL_PIPELINES = {_Units.POWER: SpectralPowerPipeline2D, _Units.RADIANCE: SpectralRadiancePipeline2D} + _PIPELINES = {_Units.POWER: PowerPipeline2D, + _Units.RADIANCE: RadiancePipeline2D} + _SPECTRAL_PIPELINES = {_Units.POWER: SpectralPowerPipeline2D, + _Units.RADIANCE: SpectralRadiancePipeline2D} + + def __init__(self, name, width, pixels, slit, transform, parent=None, + units="power", accumulate=False, curvature_radius=0): - def __init__(self, name, width, pixels, slit, transform, parent=None, units="power", accumulate=False, curvature_radius=0): # perform validation of input parameters width = float(width) if width < 0: @@ -778,15 +776,16 @@ def __init__(self, name, width, pixels, slit, transform, parent=None, units="pow curvature_radius = float(curvature_radius) if curvature_radius < 0: - raise ValueError("curvature_radius argument for BolometerIRVB must not be negative.") + raise ValueError("curvature_radius argument for BolometerIRVB " + "must not be negative.") self._slit = slit self._curvature_radius = curvature_radius self._accumulate = None # Will be set after pipeline is created. - super().__init__( - [slit.target], pixels=pixels, width=width, targeted_path_prob=0.99, parent=parent, pipelines=[], transform=transform, name=name - ) + super().__init__([slit.target], pixels=pixels, width=width, + targeted_path_prob=0.99, parent=parent, pipelines=[], + transform=transform, name=name) self.pixel_samples = 1000 self.spectral_bins = 1 self.quiet = True @@ -816,22 +815,18 @@ def pixels_as_foils(self): for x in range(nx): pixel_column = [] for y in range(ny): - pixel_centre = foil_bottom_left + (x + 0.5) * XAXIS * pixel_width + (y + 0.5) * YAXIS * pixel_height + pixel_centre = (foil_bottom_left + + (x + 0.5) * XAXIS * pixel_width + + (y + 0.5) * YAXIS * pixel_height) pixel = BolometerFoil( detector_id="IRVB pixel ({},{})".format(x + 1, y + 1), - centre_point=pixel_centre, - basis_x=XAXIS, - dx=pixel_width, - basis_y=YAXIS, - dy=pixel_height, - slit=self._slit, - units=self._units.value.capitalize(), - accumulate=False, - parent=self, + centre_point=pixel_centre, basis_x=XAXIS, dx=pixel_width, + basis_y=YAXIS, dy=pixel_height, slit=self._slit, + units=self._units.value.capitalize(), accumulate=False, parent=self ) pixel_column.append(pixel) pixels.append(pixel_column) - return np.asarray(pixels, dtype="object") + return np.asarray(pixels, dtype='object') @property def height(self): @@ -856,8 +851,9 @@ def basis_y(self): @property def sightline_vectors(self): return np.asarray( - [[pixel.centre_point.vector_to(self._slit.centre_point) for pixel in pixel_column] for pixel_column in self.pixels_as_foils], - dtype="object", + [[pixel.centre_point.vector_to(self._slit.centre_point) for pixel in pixel_column] + for pixel_column in self.pixels_as_foils], + dtype='object' ) @property @@ -880,7 +876,8 @@ def units(self, units): self._units = _Units.RADIANCE else: raise ValueError( - "The units property of BolometerIRVB must be one of {}".format([member.value for member in _Units.__members__]) + "The units property of BolometerIRVB must be one of {}" + .format([member.value for member in _Units.__members__]) ) pipeline_class = self._PIPELINES[self._units] pipeline = pipeline_class(accumulate=self.accumulate) @@ -907,7 +904,7 @@ def as_sightlines(self): """ pixels = self.pixels_as_foils sightlines = [[pixel.as_sightline() for pixel in pixel_column] for pixel_column in pixels] - return np.asarray(sightlines, dtype="object") + return np.asarray(sightlines, dtype='object') def trace_sightlines(self): """ @@ -921,7 +918,7 @@ def trace_sightlines(self): """ pixels = self.pixels_as_foils traces = [[pixel.trace_sightline() for pixel in pixel_column] for pixel_column in pixels] - return np.asarray(traces, dtype="object") + return np.asarray(traces, dtype='object') def calculate_sensitivity(self, voxel_collection, ray_count=None): r""" @@ -1046,24 +1043,26 @@ def mask_corners(element): # Make the elements to cut out from the cover slightly thicker than the # cover, to guard against rounding errors - long_box = Box(lower=Point3D(-dx / 2 + rc, -dy / 2, -0.5 * dz), upper=Point3D(dx / 2 - rc, dy / 2, 1.5 * dz)) - shot_box = Box(lower=Point3D(-dx / 2, -dy / 2 + rc, -0.5 * dz), upper=Point3D(dx / 2, dy / 2 - rc, 1.5 * dz)) + long_box = Box(lower=Point3D(-dx/2 + rc, -dy/2, -0.5 * dz), + upper=Point3D(dx/2 - rc, dy/2, 1.5 * dz)) + shot_box = Box(lower=Point3D(-dx/2, -dy/2 + rc, -0.5 * dz), + upper=Point3D(dx/2, dy/2 - rc, 1.5 * dz)) cylinder_template = Cylinder(radius=rc, height=2 * dz) top_left_cylinder = cylinder_template.instance() - top_left_cylinder.transform = translate(-dx / 2 + rc, dy / 2 - rc, -dz / 2) + top_left_cylinder.transform = translate(-dx/2 + rc, dy/2 - rc, -dz/2) top_right_cylinder = cylinder_template.instance() - top_right_cylinder.transform = translate(dx / 2 - rc, dy / 2 - rc, -dz / 2) + top_right_cylinder.transform = translate(dx/2 - rc, dy/2 - rc, -dz/2) bottom_right_cylinder = cylinder_template.instance() - bottom_right_cylinder.transform = translate(dx / 2 - rc, -dy / 2 + rc, -dz / 2) + bottom_right_cylinder.transform = translate(dx/2 - rc, -dy/2 + rc, -dz/2) bottom_left_cylinder = cylinder_template.instance() - bottom_left_cylinder.transform = translate(-dx / 2 + rc, -dy / 2 + rc, -dz / 2) - cutout = functools.reduce( - Union, (long_box, shot_box, top_left_cylinder, top_right_cylinder, bottom_right_cylinder, bottom_left_cylinder) - ) - cover = Box(lower=Point3D(-dx / 2, -dy / 2, 0), upper=Point3D(dx / 2, dy / 2, dz)) + bottom_left_cylinder.transform = translate(-dx/2 + rc, -dy/2 + rc, -dz/2) + cutout = functools.reduce(Union, (long_box, shot_box, top_left_cylinder, + top_right_cylinder, bottom_right_cylinder, + bottom_left_cylinder)) + cover = Box(lower=Point3D(-dx/2, -dy/2, 0), upper=Point3D(dx/2, dy/2, dz)) mask = Subtract(cover, cutout) mask.material = AbsorbingSurface() mask.transform = translate(0, 0, dz) - mask.name = element.name + " - rounded edges mask" + mask.name = element.name + ' - rounded edges mask' mask.parent = element diff --git a/cherab/tools/observers/group/__init__.py b/cherab/tools/observers/group/__init__.py index 8a24f80f..eca93585 100644 --- a/cherab/tools/observers/group/__init__.py +++ b/cherab/tools/observers/group/__init__.py @@ -18,6 +18,6 @@ from .fibreoptic import FibreOpticGroup from .sightline import SightLineGroup -from .targetedpixel import TargetedPixelGroup +from .targettedpixel import TargettedPixelGroup from .pixel import PixelGroup from .spectroscopic import SpectroscopicFibreOpticGroup, SpectroscopicSightLineGroup diff --git a/cherab/tools/observers/group/targetedpixel.py b/cherab/tools/observers/group/targetedpixel.py index a9f8775e..e440f70e 100644 --- a/cherab/tools/observers/group/targetedpixel.py +++ b/cherab/tools/observers/group/targetedpixel.py @@ -22,11 +22,11 @@ from .base import Observer0DGroup -class TargetedPixelGroup(Observer0DGroup): +class TargettedPixelGroup(Observer0DGroup): """ A group of targeted pixel under a single scene-graph node. - A scene-graph object regrouping a series of `TargetedPixel` + A scene-graph object regrouping a series of 'TargettedPixel' observers as a scene-graph parent. Allows combined observation and display control simultaneously. @@ -49,9 +49,8 @@ def x_width(self, value): 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)) - ) + 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 @@ -67,9 +66,8 @@ def y_width(self, value): 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)) - ) + 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 @@ -94,11 +92,8 @@ def targets(self, value): 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) - ) - ) + 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: @@ -115,9 +110,7 @@ def targeted_path_prob(self, value): 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)) - ) + 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/tests/test_observer_groups.py b/cherab/tools/tests/test_observer_groups.py index f92ab7d3..a04f095c 100644 --- a/cherab/tools/tests/test_observer_groups.py +++ b/cherab/tools/tests/test_observer_groups.py @@ -1,11 +1,11 @@ import unittest from raysect.core.workflow import RenderEngine -from raysect.optical.observer import FibreOptic, Observer0D, Pixel, PowerPipeline0D, SightLine, SpectralPowerPipeline0D, TargetedPixel +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 from cherab.tools.observers.group.base import Observer0DGroup +from cherab.tools.observers.group import SightLineGroup, FibreOpticGroup, PixelGroup, TargettedPixelGroup from cherab.tools.raytransfer import pipelines @@ -20,13 +20,13 @@ def setUp(self): def test_get_item(self): """Tests all inputs for the __get_item__ method""" group = self._GROUP_CLASS(observers=self.observers) - names = ["zero", "one", "two"] + names = ['zero', 'one', 'two'] group.names = names 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]) @@ -37,11 +37,11 @@ def test_get_item(self): group[1.2] with self.assertRaises(ValueError): - group["fail"] + group['fail'] - group.names = ["fail"] * len(group) + group.names = ['fail'] * len(group) with self.assertRaises(ValueError): - group["fail"] + group['fail'] def test_assignments(self): """Test assignments of all supported attributes of Observer0DGroup""" @@ -49,7 +49,7 @@ def test_assignments(self): group.observers = self.observers for grouped_observer, input_observer in zip(group.observers, self.observers): - self.assertIs(grouped_observer, input_observer, msg="Observers do not match") + self.assertIs(grouped_observer, input_observer, msg='Observers do not match') with self.assertRaises(ValueError): group.observers = [Sphere()] @@ -58,32 +58,32 @@ def test_assignments(self): group.observers = Sphere() # names - names = ["zero", "one", "two"] + names = ['zero', 'one', 'two'] group.names = names for grouped_observer, input_name in zip(group.observers, names): - self.assertEqual(grouped_observer.name, input_name, msg="Observer name do not match") + self.assertEqual(grouped_observer.name, input_name, msg='Observer name do not match') with self.assertRaises(ValueError): - group.names = ["fail"] + group.names = ['fail'] with self.assertRaises(TypeError): - group.names = "fail" + group.names = 'fail' # pipelines - ppln_0 = PowerPipeline0D(name="pipeline zero, observer zero") - ppln_1 = PowerPipeline0D(name="pipeline one, observer one") - ppln_2 = PowerPipeline0D(name="pipeline two, observer two") - ppln_3 = PowerPipeline0D(name="pipeline three, observer two") + ppln_0 = PowerPipeline0D(name='pipeline zero, observer zero') + ppln_1 = PowerPipeline0D(name='pipeline one, observer one') + ppln_2 = PowerPipeline0D(name='pipeline two, observer two') + ppln_3 = PowerPipeline0D(name='pipeline three, observer two') pipelist = [[ppln_0], [ppln_1], [ppln_2, ppln_3]] group.pipelines = pipelist - self.assertIs(group[0].pipelines[0], ppln_0, "non matching pipeline") - self.assertIs(group[1].pipelines[0], ppln_1, "non matching pipeline") - self.assertIs(group[2].pipelines[0], ppln_2, "non matching pipeline") - self.assertIs(group[2].pipelines[1], ppln_3, "non matching pipeline") + self.assertIs(group[0].pipelines[0], ppln_0, 'non matching pipeline') + self.assertIs(group[1].pipelines[0], ppln_1, 'non matching pipeline') + self.assertIs(group[2].pipelines[0], ppln_2, 'non matching pipeline') + self.assertIs(group[2].pipelines[1], ppln_3, 'non matching pipeline') 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,15 +102,15 @@ 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 self.assertListEqual(group.min_wavelength, [wvl - 100] * len(group)) self.assertListEqual(group.max_wavelength, [wvl + 100] * len(group)) - min_wvls = [90 + 10 * i for i in range(len(group))] - max_wvls = [100 + 10 * i for i in range(len(group))] + min_wvls = [90 + 10*i for i in range(len(group))] + max_wvls = [100 + 10*i for i in range(len(group))] group.min_wavelength = min_wvls group.max_wavelength = max_wvls self.assertListEqual(group.min_wavelength, min_wvls) @@ -122,7 +122,7 @@ def test_assignments(self): group.min_wavelength = [90] * (len(group) - 1) # spectral - bins = [200 + i * 100 for i in range(len(group))] + bins = [200 + i*100 for i in range(len(group))] rays = [2] * len(group) group.spectral_bins = bins group.spectral_rays = rays @@ -139,7 +139,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,8 +152,8 @@ def test_assignments(self): with self.assertRaises(ValueError): group.quiet = [False] * (len(group) + 1) - # rays - probs = [0.2 + i * 0.1 for i in range(len(group))] + # 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))] sampling = [False] * len(group) @@ -196,10 +196,10 @@ 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))] + pixel_samples = [2000 + i*500 for i in range(len(group))] + per_task = [5000 + i*100 for i in range(len(group))] group.pixel_samples = pixel_samples group.samples_per_task = per_task self.assertListEqual(group.pixel_samples, pixel_samples) @@ -228,7 +228,7 @@ def test_connect_pipelines(self): group = self._GROUP_CLASS(observers=self.observers) ppln_classes = [PowerPipeline0D, SpectralPowerPipeline0D] - names = ["power", "spectral"] + names = ['power', 'spectral'] keywords = [ dict(name=names[0]), dict(name=names[1], display_progress=True), @@ -348,8 +348,8 @@ def test_widths(self): group.y_width = [1e-1] * (len(group) + 1) -class TargetedPixelGroupTestCase(PixelGroupTestCase): - _GROUP_CLASS = TargetedPixelGroup +class TargettedPixelGroupTestCase(PixelGroupTestCase): + _GROUP_CLASS = TargettedPixelGroup def setUp(self): self.observers = [TargetedPixel(targets=[Sphere()], pipelines=[PowerPipeline0D()]) for _ in range(self._NUM)] @@ -374,7 +374,7 @@ def test_targets(self): with self.assertRaises(ValueError): group.targets = targets - # targeted path prob + # targetted path prob prob = [0.9, 0.95, 1] group.targeted_path_prob = prob self.assertListEqual(group.targeted_path_prob, prob) diff --git a/docs/source/tools/observers.rst b/docs/source/tools/observers.rst index 84f94f48..e93e2057 100644 --- a/docs/source/tools/observers.rst +++ b/docs/source/tools/observers.rst @@ -109,10 +109,10 @@ combined into a group. Group observers --------------- -Group observer is a collection of observers of the same type. All Observer0D classes -defined in Raysect are supoorted. The parameters of individual observers in a group +Group observer is a collection of observers of the same type. All Observer0D classes +defined in Raysect are supoorted. The parameters of individual observers in a group may differ. Group observer allows combined observation, namely, calling the observe -function for a group leads to a sequential call of this function for each observer +function for a group leads to a sequential call of this function for each observer in the group. .. autoclass:: cherab.tools.observers.group.base.Observer0DGroup @@ -127,7 +127,7 @@ in the group. .. autoclass:: cherab.tools.observers.group.PixelGroup :members: -.. autoclass:: cherab.tools.observers.group.TargetedPixelGroup +.. autoclass:: cherab.tools.observers.group.TargettedPixelGroup :members: Spectroscopic Groups @@ -136,9 +136,9 @@ Spectroscopic Groups .. deprecated:: 1.4.0 Use groups based on Raysect's observer classes instead -These groups take control of spectroscopic lines of sight observers. They support -direction and origin positioning and contain methods for plotting the power and -spectrum. Originally, these were called group observers and did not include the +These groups take control of spectroscopic lines of sight observers. They support +direction and origin positioning and contain methods for plotting the power and +spectrum. Originally, these were called group observers and did not include the Spectroscopic prefix in class name. .. autoclass:: cherab.tools.observers.SpectroscopicSightLine From 97a634cbfa97602f1dd7ab00449464a2f395a8ac Mon Sep 17 00:00:00 2001 From: munechika-koyo Date: Thu, 16 Oct 2025 10:00:57 +0200 Subject: [PATCH 26/91] Revert "Rename filename" This reverts commit 662eb0447fda4fd80a925a1b804d2a19b77c28ab. --- .../tools/observers/group/{targetedpixel.py => targettedpixel.py} | 0 1 file changed, 0 insertions(+), 0 deletions(-) rename cherab/tools/observers/group/{targetedpixel.py => targettedpixel.py} (100%) diff --git a/cherab/tools/observers/group/targetedpixel.py b/cherab/tools/observers/group/targettedpixel.py similarity index 100% rename from cherab/tools/observers/group/targetedpixel.py rename to cherab/tools/observers/group/targettedpixel.py From 60d76bff22f57f3ed6ca1c6db5cc3df624a163ed Mon Sep 17 00:00:00 2001 From: munechika-koyo Date: Thu, 16 Oct 2025 10:15:16 +0200 Subject: [PATCH 27/91] Fix typo: change `targetted` to `targeted` in the internal code of Bolometer classes --- cherab/tools/observers/bolometry.py | 10 +++++----- 1 file changed, 5 insertions(+), 5 deletions(-) diff --git a/cherab/tools/observers/bolometry.py b/cherab/tools/observers/bolometry.py index 5cd8c1a2..b69fef29 100644 --- a/cherab/tools/observers/bolometry.py +++ b/cherab/tools/observers/bolometry.py @@ -213,7 +213,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. @@ -661,8 +661,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 = TargetedHemisphereSampler(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 +671,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) From ab0b0fd51577a37b05e36f90798de517580be2e5 Mon Sep 17 00:00:00 2001 From: munechika-koyo Date: Thu, 16 Oct 2025 10:17:15 +0200 Subject: [PATCH 28/91] Bump version to 1.6.0.dev2 --- cherab/core/VERSION | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/cherab/core/VERSION b/cherab/core/VERSION index 4c39f7c4..8a3469e1 100644 --- a/cherab/core/VERSION +++ b/cherab/core/VERSION @@ -1 +1 @@ -1.6.0.dev1 +1.6.0.dev2 From 3dbad49df1a15d309e130eaa87358706e6fb0bba Mon Sep 17 00:00:00 2001 From: munechika-koyo Date: Thu, 16 Oct 2025 10:20:38 +0200 Subject: [PATCH 29/91] Fix typo: change `targetted` to `targeted` in generate_segmented_cylinder function --- cherab/core/model/laser/profile.pyx | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) 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 From 7d1e59abb6966d7bf9fc02b8a7b8199a69a322c0 Mon Sep 17 00:00:00 2001 From: munechika-koyo Date: Thu, 16 Oct 2025 10:57:52 +0200 Subject: [PATCH 30/91] Refactor TargettedPixelGroup to TargetedPixelGroup and add deprecation warnings --- cherab/tools/observers/__init__.py | 2 +- cherab/tools/observers/group/__init__.py | 5 +- cherab/tools/observers/group/targetedpixel.py | 123 ++++++++++++++++++ .../tools/observers/group/targettedpixel.py | 100 +++----------- cherab/tools/tests/test_observer_groups.py | 29 ++++- 5 files changed, 167 insertions(+), 92 deletions(-) create mode 100644 cherab/tools/observers/group/targetedpixel.py 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/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..165921b5 --- /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 pixel 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..3d612b2f 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 a future version. + 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,11 @@ 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 - @property - def targetted_path_prob(self): - return [pixel.targetted_path_prob for pixel in self._observers] - - @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 + def __init__(self, *args, **kwargs): + warnings.warn( + "TargettedPixelGroup is deprecated and will be removed in a future version. Use TargetedPixelGroup instead.", + DeprecationWarning, + stacklevel=2, + ) + super().__init__(*args, **kwargs) diff --git a/cherab/tools/tests/test_observer_groups.py b/cherab/tools/tests/test_observer_groups.py index 1f5b7cb0..649e5a96 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.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 @@ -348,8 +349,8 @@ 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)] @@ -374,7 +375,7 @@ 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) @@ -385,4 +386,22 @@ def test_targets(self): self.assertEqual(group_targetted_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)) From 24779d59290ec0d41afd85c0280fda885eb46e4e Mon Sep 17 00:00:00 2001 From: munechika-koyo Date: Thu, 16 Oct 2025 10:58:19 +0200 Subject: [PATCH 31/91] Fix formatting issues by removing unnecessary blank lines --- cherab/tools/tests/test_observer_groups.py | 12 ++++++------ 1 file changed, 6 insertions(+), 6 deletions(-) diff --git a/cherab/tools/tests/test_observer_groups.py b/cherab/tools/tests/test_observer_groups.py index 649e5a96..09527634 100644 --- a/cherab/tools/tests/test_observer_groups.py +++ b/cherab/tools/tests/test_observer_groups.py @@ -27,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]) @@ -84,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: @@ -103,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 @@ -140,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) @@ -153,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))] @@ -197,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))] From 1a1aaafcf8f79239137b8221bdd2a28ba5aa2a7d Mon Sep 17 00:00:00 2001 From: munechika-koyo Date: Thu, 16 Oct 2025 11:07:46 +0200 Subject: [PATCH 32/91] Fix typo in TargetedPixelGroup class docstring --- cherab/tools/observers/group/targetedpixel.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/cherab/tools/observers/group/targetedpixel.py b/cherab/tools/observers/group/targetedpixel.py index 165921b5..9f4fb814 100644 --- a/cherab/tools/observers/group/targetedpixel.py +++ b/cherab/tools/observers/group/targetedpixel.py @@ -24,7 +24,7 @@ class TargetedPixelGroup(Observer0DGroup): """ - A group of targeted pixel under a single scene-graph node. + 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 From 866740b9a51729ab71de1556f0c2a87939c71120 Mon Sep 17 00:00:00 2001 From: munechika-koyo Date: Thu, 16 Oct 2025 11:21:21 +0200 Subject: [PATCH 33/91] Correct imported class `TargetedPixel` not using `TargettedPixel` --- cherab/tools/tests/test_observer_groups.py | 14 +++++++------- 1 file changed, 7 insertions(+), 7 deletions(-) diff --git a/cherab/tools/tests/test_observer_groups.py b/cherab/tools/tests/test_observer_groups.py index 09527634..a602531e 100644 --- a/cherab/tools/tests/test_observer_groups.py +++ b/cherab/tools/tests/test_observer_groups.py @@ -2,7 +2,7 @@ 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 @@ -353,7 +353,7 @@ 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) @@ -377,13 +377,13 @@ def test_targets(self): # 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.targeted_path_prob = [0.7] * (len(group) + 1) From 1a640cf69d9193fc4366d9552e46bfba758dc833 Mon Sep 17 00:00:00 2001 From: munechika-koyo Date: Thu, 16 Oct 2025 11:24:10 +0200 Subject: [PATCH 34/91] Revert previous property to keep it in deprecated class --- cherab/tools/observers/group/targettedpixel.py | 8 ++++++++ 1 file changed, 8 insertions(+) diff --git a/cherab/tools/observers/group/targettedpixel.py b/cherab/tools/observers/group/targettedpixel.py index 3d612b2f..b08ea996 100644 --- a/cherab/tools/observers/group/targettedpixel.py +++ b/cherab/tools/observers/group/targettedpixel.py @@ -46,3 +46,11 @@ def __init__(self, *args, **kwargs): stacklevel=2, ) super().__init__(*args, **kwargs) + + @property + def targetted_path_prob(self): + return self.targeted_path_prob + + @targetted_path_prob.setter + def targetted_path_prob(self, value): + self.targeted_path_prob = value From 503aa22fb2d270948b7227505d698119b0804255 Mon Sep 17 00:00:00 2001 From: munechika-koyo Date: Thu, 16 Oct 2025 11:28:54 +0200 Subject: [PATCH 35/91] Restore `targetted_path_prob` property --- cherab/tools/observers/group/targettedpixel.py | 8 ++++---- 1 file changed, 4 insertions(+), 4 deletions(-) diff --git a/cherab/tools/observers/group/targettedpixel.py b/cherab/tools/observers/group/targettedpixel.py index e440f70e..f21694aa 100644 --- a/cherab/tools/observers/group/targettedpixel.py +++ b/cherab/tools/observers/group/targettedpixel.py @@ -33,7 +33,7 @@ class TargettedPixelGroup(Observer0DGroup): :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 + :ivar list targetted_path_prob: Probability of ray being casted at the target """ _OBSERVER_TYPE = TargetedPixel @@ -100,11 +100,11 @@ def targets(self, value): pixel.targets = value @property - def targeted_path_prob(self): + def targetted_path_prob(self): return [pixel.targeted_path_prob for pixel in self._observers] - @targeted_path_prob.setter - def targeted_path_prob(self, value): + @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): From 5182d8b7a56955ee8704e36fb0316bec93462bf8 Mon Sep 17 00:00:00 2001 From: munechika-koyo Date: Thu, 16 Oct 2025 11:31:17 +0200 Subject: [PATCH 36/91] Restore `targetted_path_prob` in test --- cherab/tools/tests/test_observer_groups.py | 20 ++++++++++---------- 1 file changed, 10 insertions(+), 10 deletions(-) diff --git a/cherab/tools/tests/test_observer_groups.py b/cherab/tools/tests/test_observer_groups.py index a04f095c..c75e8630 100644 --- a/cherab/tools/tests/test_observer_groups.py +++ b/cherab/tools/tests/test_observer_groups.py @@ -26,7 +26,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 +83,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 +102,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 +139,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 +152,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 +196,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))] @@ -376,13 +376,13 @@ def test_targets(self): # targetted path prob prob = [0.9, 0.95, 1] - group.targeted_path_prob = prob - self.assertListEqual(group.targeted_path_prob, prob) + group.targetted_path_prob = prob + self.assertListEqual(group.targetted_path_prob, prob) prob = 0.8 group.targeted_path_prob = prob - for group_targeted_path_prob in group.targeted_path_prob: + for group_targeted_path_prob in group.targetted_path_prob: self.assertEqual(group_targeted_path_prob, prob) with self.assertRaises(ValueError): - group.targeted_path_prob = [0.7] * (len(group) + 1) + group.targetted_path_prob = [0.7] * (len(group) + 1) From 362a6e3b8ba190a0e389aaaffbef7b5ed3674da9 Mon Sep 17 00:00:00 2001 From: munechika-koyo Date: Thu, 16 Oct 2025 11:34:43 +0200 Subject: [PATCH 37/91] Remove API changes section --- CHANGELOG.md | 2 -- 1 file changed, 2 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index 6a4b7e88..417bd613 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -3,8 +3,6 @@ Project Changelog Release 1.6.0 (TBD) ------------------- -API changes: -* Rename 'targetted' to 'targeted' following Raysect change. (#486) New: * Add Function6D framework. (#478) From fd044a64ecd7c8d84622e255a1548b2c2e82abb0 Mon Sep 17 00:00:00 2001 From: munechika-koyo Date: Thu, 16 Oct 2025 11:37:33 +0200 Subject: [PATCH 38/91] Update CHANGELOG.md --- CHANGELOG.md | 7 +++++-- 1 file changed, 5 insertions(+), 2 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index 4f22a5d4..1c132f58 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -4,6 +4,9 @@ Project Changelog Release 1.6.0 (TBD) ------------------- +API changes: +* Rename `TargettedPixelGroup` to `TargetedPixelGroup` for correct spelling. Still keep `TargettedPixelGroup` as an alias for backwards compatibility. (#487) + New: * Add Function6D framework. (#478) * Add e_field attribute to Plasma object for electric field vector. (#465) @@ -134,7 +137,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. @@ -162,7 +165,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) From 21d64497f1a46ae1d330e76bbcafffb6084ebcfb Mon Sep 17 00:00:00 2001 From: munechika-koyo Date: Thu, 16 Oct 2025 11:45:57 +0200 Subject: [PATCH 39/91] Fix typo --- cherab/tools/tests/test_observer_groups.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/cherab/tools/tests/test_observer_groups.py b/cherab/tools/tests/test_observer_groups.py index c75e8630..0418eab9 100644 --- a/cherab/tools/tests/test_observer_groups.py +++ b/cherab/tools/tests/test_observer_groups.py @@ -380,7 +380,7 @@ def test_targets(self): self.assertListEqual(group.targetted_path_prob, prob) prob = 0.8 - group.targeted_path_prob = prob + group.targetted_path_prob = prob for group_targeted_path_prob in group.targetted_path_prob: self.assertEqual(group_targeted_path_prob, prob) From 7d285d161fd72b9968b5b5fa1a0a8acf2496c8fd Mon Sep 17 00:00:00 2001 From: munechika-koyo Date: Thu, 16 Oct 2025 16:25:20 +0200 Subject: [PATCH 40/91] 1st commit for introducing `pixi` Add pixi manifest file and git-related files which are automatically generated by `pixi` --- .gitattributes | 2 + .gitignore | 9 ++- pixi.toml | 153 +++++++++++++++++++++++++++++++++++++++++++++++++ 3 files changed, 163 insertions(+), 1 deletion(-) create mode 100644 .gitattributes create mode 100644 pixi.toml 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/.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/pixi.toml b/pixi.toml new file mode 100644 index 00000000..f75a8e15 --- /dev/null +++ b/pixi.toml @@ -0,0 +1,153 @@ +[workspace] +channels = ["https://prefix.dev/conda-forge"] +platforms = ["linux-64", "osx-arm64", "osx-64"] +preview = ["pixi-build"] + +# ------------------------------- +# === Packaging Configuration === +# ------------------------------- +[package] +name = "cherab" +version = "dynamic" + +[package.build] +backend = { name = "pixi-build-python", version = "*" } + +[package.build.config] +noarch = false +compilers = ["c"] + +[workspace.build-variants] +python = ["3.9", "3.10.*", "3.11.*", "3.12.*", "3.13.*"] + +[package.host-dependencies] +uv = "*" +python = "*" +setuptools = "*" +cython = ">=3.1" +numpy = "*" +raysect = "0.9.*" + +[package.run-dependencies] +scipy = "*" +matplotlib-base = "*" +pyopencl = "*" +pocl = "*" + +# --------------------------------- +# === Development Configuration === +# --------------------------------- +[dependencies] +ipython = "*" + +[tasks] +clean = { cmd = [ + "find", + "cherab/", + "-type", + "f", + "\\(", + "-name", + "'*.c'", + "-o", + "-name", + "'*.so'", + "-o", + "-name", + "'*.dylib'", + "\\)", + "-delete", +], description = "🔥 Remove in-place build artifacts and temporary files (*.c, *.so, *.dylib)" } + +# 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", + "8000", + "--directory", + "build/html", +], cwd = "docs", 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" } + +# === Documentation feature === +[feature.docs.dependencies] +cherab = { path = "." } +sphinx = "*" +sphinx_rtd_theme = "<1" + +[feature.docs.pypi-dependencies] +sphinx-tabs = "*" # >=3.4.4 has not yet been released to conda-forge + +[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 = "*" +taplo = "*" + +[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" } +dprint = { cmd = "dprint fmt", description = "Format with dprint" } +typos = { cmd = "typos --write-changes --force-exclude", description = "Fix typos" } +taplo = { cmd = "taplo fmt", description = "Format toml files with taplo" } +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.py39.dependencies] +python = "3.9.*" +[feature.py313.dependencies] +python = "3.13.*" + +[environments] +default = { features = ["py313"], solve-group = "py313" } +test = { features = ["test"], solve-group = "py313" } +docs = { features = [ + "py39", + "docs", +], solve-group = "py39" } # TODO: change to py313 when bumping RTD theme to >=1.0 +test-py313 = { features = [ + "py313", + "test", +], solve-group = "py313" } # alias of tests +test-py39 = { features = [ + "py39", + "test", +], solve-group = "py39" } # alias of tests +lint = { features = ["lint"], no-default-feature = true } From 088a369a6a648234f1f924798cdfee49f83caa6a Mon Sep 17 00:00:00 2001 From: munechika-koyo Date: Fri, 17 Oct 2025 10:35:53 +0200 Subject: [PATCH 41/91] Update deprecation notice for TargettedPixelGroup to specify removal in version 2.0 --- CHANGELOG.md | 2 +- cherab/tools/observers/group/targettedpixel.py | 9 +++++---- 2 files changed, 6 insertions(+), 5 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index 1c132f58..0737af44 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -5,7 +5,7 @@ Release 1.6.0 (TBD) ------------------- API changes: -* Rename `TargettedPixelGroup` to `TargetedPixelGroup` for correct spelling. Still keep `TargettedPixelGroup` as an alias for backwards compatibility. (#487) +* Rename `TargettedPixelGroup` to `TargetedPixelGroup` for correct spelling. Still keep `TargettedPixelGroup` as an alias for backwards compatibility until the next major release. (#487) New: * Add Function6D framework. (#478) diff --git a/cherab/tools/observers/group/targettedpixel.py b/cherab/tools/observers/group/targettedpixel.py index b08ea996..5d1ac69f 100644 --- a/cherab/tools/observers/group/targettedpixel.py +++ b/cherab/tools/observers/group/targettedpixel.py @@ -26,10 +26,10 @@ class TargettedPixelGroup(_TargetedPixelGroup): A group of targeted pixel under a single scene-graph node. .. deprecated:: - TargettedPixelGroup is deprecated and will be removed in a future version. - Use TargetedPixelGroup instead. + `TargettedPixelGroup` is deprecated and will be removed in version 2.0. + Use `TargetedPixelGroup` instead. - A scene-graph object regrouping a series of 'TargetedPixel' + A scene-graph object regrouping a series of `TargetedPixel` observers as a scene-graph parent. Allows combined observation and display control simultaneously. @@ -41,7 +41,8 @@ class TargettedPixelGroup(_TargetedPixelGroup): def __init__(self, *args, **kwargs): warnings.warn( - "TargettedPixelGroup is deprecated and will be removed in a future version. Use TargetedPixelGroup instead.", + "TargettedPixelGroup is deprecated and will be removed in version 2.0. " + + "Use TargetedPixelGroup instead.", DeprecationWarning, stacklevel=2, ) From f270fd6d19dcccf79fb382b46e37e00b5b6966e9 Mon Sep 17 00:00:00 2001 From: munechika-koyo Date: Fri, 17 Oct 2025 11:26:15 +0200 Subject: [PATCH 42/91] Enhance CI workflow to include draft pull request checks --- .github/workflows/ci.yml | 6 ++++++ 1 file changed, 6 insertions(+) diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index da0d082c..e2f913f9 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -3,9 +3,15 @@ name: CI on: push: pull_request: + types: + - opened + - synchronize + - reopened + - ready_for_review jobs: tests: + if: ${{ !github.event.pull_request.draft }} name: Run tests runs-on: ubuntu-22.04 # Needed for Python 3.7 compatibility strategy: From 2271d14d2f71a0993d4e45083a03085ce0f38c0e Mon Sep 17 00:00:00 2001 From: munechika-koyo Date: Fri, 17 Oct 2025 11:32:28 +0200 Subject: [PATCH 43/91] Update CI workflow to run tests on push events in addition to non-draft pull requests --- .github/workflows/ci.yml | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index e2f913f9..f06c406e 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -11,7 +11,7 @@ on: jobs: tests: - if: ${{ !github.event.pull_request.draft }} + if: ${{ github.event_name == 'push' || !github.event.pull_request.draft }} name: Run tests runs-on: ubuntu-22.04 # Needed for Python 3.7 compatibility strategy: From 9b20d374dbf1e88aa8f87e8ade416b7b249f9cb8 Mon Sep 17 00:00:00 2001 From: munechika-koyo Date: Fri, 17 Oct 2025 18:01:15 +0200 Subject: [PATCH 44/91] Revert import lines in setup.py --- setup.py | 11 +++++------ 1 file changed, 5 insertions(+), 6 deletions(-) diff --git a/setup.py b/setup.py index 8699c4f5..15bc291a 100644 --- a/setup.py +++ b/setup.py @@ -1,15 +1,14 @@ -import multiprocessing +from collections import defaultdict +import sys import os import os.path as path -import sys -from collections import defaultdict from pathlib import Path - +import multiprocessing import numpy +from setuptools import setup, find_packages, Extension from Cython.Build import cythonize -from setuptools import Extension, find_packages, setup -multiprocessing.set_start_method("fork") +multiprocessing.set_start_method('fork') force = False profile = False From 780364f40251fd652dcf3e95c495a7d251871714 Mon Sep 17 00:00:00 2001 From: munechika-koyo Date: Fri, 17 Oct 2025 18:01:51 +0200 Subject: [PATCH 45/91] Update numpy version requirement to 2.0 in setup.py --- setup.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/setup.py b/setup.py index 15bc291a..c57bb0d3 100644 --- a/setup.py +++ b/setup.py @@ -117,7 +117,7 @@ long_description=long_description, long_description_content_type="text/markdown", install_requires=[ - "numpy>=2", + "numpy>=2.0", "scipy", "matplotlib", "raysect==0.9.1.*", From cae511476ff9450494bcc70f841ad6e3a8be53d9 Mon Sep 17 00:00:00 2001 From: munechika-koyo Date: Fri, 17 Oct 2025 18:02:29 +0200 Subject: [PATCH 46/91] Update numpy version requirement to 2.0 in requirements.txt --- requirements.txt | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/requirements.txt b/requirements.txt index 7fea14d1..c99e710f 100644 --- a/requirements.txt +++ b/requirements.txt @@ -1,5 +1,5 @@ cython~=3.1 -numpy>=2 +numpy>=2.0 scipy matplotlib raysect==0.9.1.* From 8d48d541c15bf51761e79796904cf5983a17a136 Mon Sep 17 00:00:00 2001 From: munechika-koyo Date: Mon, 27 Oct 2025 09:42:22 +0100 Subject: [PATCH 47/91] Add deprecation warnings for `targetted_path_prob` property in Bolometer classes --- cherab/tools/observers/bolometry.py | 37 +++++++++++++++++++++++++++++ 1 file changed, 37 insertions(+) diff --git a/cherab/tools/observers/bolometry.py b/cherab/tools/observers/bolometry.py index b69fef29..33870b0a 100644 --- a/cherab/tools/observers/bolometry.py +++ b/cherab/tools/observers/bolometry.py @@ -18,6 +18,7 @@ # under the Licence. from enum import Enum +from warnings import warn import functools import numpy as np @@ -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. @@ -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. From 63b9a8c5132d8a29121aae30b4a5af53fbd9bc71 Mon Sep 17 00:00:00 2001 From: MatejTomes Date: Wed, 29 Oct 2025 21:04:21 +0100 Subject: [PATCH 48/91] Chage docstrings - reference existing parameters Kindo of following DRY in docstrings... --- cherab/core/atomic/line.pyx | 10 +++------- cherab/core/model/plasma/impact_excitation.pyx | 4 ++-- cherab/core/model/plasma/recombination.pyx | 4 ++-- cherab/core/model/plasma/thermal_cx.pyx | 4 ++-- cherab/core/model/plasma/total_radiated_power.pyx | 4 ++-- 5 files changed, 11 insertions(+), 15 deletions(-) diff --git a/cherab/core/atomic/line.pyx b/cherab/core/atomic/line.pyx index 1cb34584..cbd07bba 100644 --- a/cherab/core/atomic/line.pyx +++ b/cherab/core/atomic/line.pyx @@ -34,13 +34,9 @@ cdef class Line: 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: The atomic element/isotope to which this emission line belongs. - :ivar int charge: The charge state of the element/isotope that emits this line. - :ivar tuple transition: A two element tuple that defines the upper and lower electron - configuration states of the transition. For hydrogen-like ions it may be enough to - 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/model/plasma/impact_excitation.pyx b/cherab/core/model/plasma/impact_excitation.pyx index e10b6280..1e40754a 100644 --- a/cherab/core/model/plasma/impact_excitation.pyx +++ b/cherab/core/model/plasma/impact_excitation.pyx @@ -45,8 +45,8 @@ 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. """ diff --git a/cherab/core/model/plasma/recombination.pyx b/cherab/core/model/plasma/recombination.pyx index db00ad35..f92c157c 100644 --- a/cherab/core/model/plasma/recombination.pyx +++ b/cherab/core/model/plasma/recombination.pyx @@ -45,8 +45,8 @@ 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. """ diff --git a/cherab/core/model/plasma/thermal_cx.pyx b/cherab/core/model/plasma/thermal_cx.pyx index aa1e9ffa..ef9819a8 100644 --- a/cherab/core/model/plasma/thermal_cx.pyx +++ b/cherab/core/model/plasma/thermal_cx.pyx @@ -47,8 +47,8 @@ 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. """ diff --git a/cherab/core/model/plasma/total_radiated_power.pyx b/cherab/core/model/plasma/total_radiated_power.pyx index 742cb3db..545a3727 100644 --- a/cherab/core/model/plasma/total_radiated_power.pyx +++ b/cherab/core/model/plasma/total_radiated_power.pyx @@ -54,8 +54,8 @@ cdef class TotalRadiatedPower(PlasmaModel): :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: The atomic element/isotope. - :ivar int charge: The charge state of the element/isotope. + :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): From 45cce623d4edc3970517a0584d6d6c6f91d9079f Mon Sep 17 00:00:00 2001 From: MatejTomes Date: Wed, 29 Oct 2025 21:09:45 +0100 Subject: [PATCH 49/91] Add Line class tests --- cherab/core/atomic/tests/test_line.py | 27 +++++++++++++++++++++++++++ 1 file changed, 27 insertions(+) create mode 100644 cherab/core/atomic/tests/test_line.py 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 From 8edb74767ed6fc7797ffaff73d02c57e5367a9ce Mon Sep 17 00:00:00 2001 From: MatejTomes Date: Wed, 29 Oct 2025 21:58:47 +0100 Subject: [PATCH 50/91] Add plasma model tests --- cherab/core/model/plasma/tests/test_models.py | 58 +++++++++++++++++++ 1 file changed, 58 insertions(+) create mode 100644 cherab/core/model/plasma/tests/test_models.py 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..56203db0 --- /dev/null +++ b/cherab/core/model/plasma/tests/test_models.py @@ -0,0 +1,58 @@ +import unittest + +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 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) + + 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) + + 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) From 9d9a1608fca4534bd938c70d4e79cfeeb006d00d Mon Sep 17 00:00:00 2001 From: MatejTomes Date: Wed, 29 Oct 2025 21:59:33 +0100 Subject: [PATCH 51/91] Add forgotten __init__.py files --- cherab/core/atomic/tests/__init__.py | 0 cherab/core/model/plasma/tests/__init__.py | 0 2 files changed, 0 insertions(+), 0 deletions(-) create mode 100644 cherab/core/atomic/tests/__init__.py create mode 100644 cherab/core/model/plasma/tests/__init__.py 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/model/plasma/tests/__init__.py b/cherab/core/model/plasma/tests/__init__.py new file mode 100644 index 00000000..e69de29b From a277d02c53933e6f80f07ea7dcec0ec5640f5776 Mon Sep 17 00:00:00 2001 From: MatejTomes Date: Wed, 29 Oct 2025 22:10:02 +0100 Subject: [PATCH 52/91] Add typing to the properties --- cherab/core/model/plasma/impact_excitation.pyx | 4 ++-- cherab/core/model/plasma/recombination.pyx | 4 ++-- cherab/core/model/plasma/thermal_cx.pyx | 4 ++-- cherab/core/model/plasma/total_radiated_power.pyx | 4 ++-- 4 files changed, 8 insertions(+), 8 deletions(-) diff --git a/cherab/core/model/plasma/impact_excitation.pyx b/cherab/core/model/plasma/impact_excitation.pyx index 1e40754a..5c12d94f 100644 --- a/cherab/core/model/plasma/impact_excitation.pyx +++ b/cherab/core/model/plasma/impact_excitation.pyx @@ -78,11 +78,11 @@ cdef class ExcitationLine(PlasmaModel): return ''.format(self._line.element.name, self._line.charge, self._line.transition) @property - def line(self): + def line(self) -> Line: return self._line @property - def lineshape(self): + def lineshape(self) -> LineShapeModel: return self._lineshape cpdef Spectrum emission(self, Point3D point, Vector3D direction, Spectrum spectrum): diff --git a/cherab/core/model/plasma/recombination.pyx b/cherab/core/model/plasma/recombination.pyx index f92c157c..7f85dad8 100644 --- a/cherab/core/model/plasma/recombination.pyx +++ b/cherab/core/model/plasma/recombination.pyx @@ -78,11 +78,11 @@ cdef class RecombinationLine(PlasmaModel): return ''.format(self._line.element.name, self._line.charge, self._line.transition) @property - def line(self): + def line(self) -> Line: return self._line @property - def lineshape(self): + def lineshape(self) -> LineShapeModel: return self._lineshape cpdef Spectrum emission(self, Point3D point, Vector3D direction, Spectrum spectrum): diff --git a/cherab/core/model/plasma/thermal_cx.pyx b/cherab/core/model/plasma/thermal_cx.pyx index ef9819a8..e0943a14 100644 --- a/cherab/core/model/plasma/thermal_cx.pyx +++ b/cherab/core/model/plasma/thermal_cx.pyx @@ -80,11 +80,11 @@ cdef class ThermalCXLine(PlasmaModel): return ''.format(self._line.element.name, self._line.charge, self._line.transition) @property - def line(self): + def line(self) -> Line: return self._line @property - def lineshape(self): + def lineshape(self) -> LineShapeModel: return self._lineshape cpdef Spectrum emission(self, Point3D point, Vector3D direction, Spectrum spectrum): diff --git a/cherab/core/model/plasma/total_radiated_power.pyx b/cherab/core/model/plasma/total_radiated_power.pyx index 545a3727..f0c94281 100644 --- a/cherab/core/model/plasma/total_radiated_power.pyx +++ b/cherab/core/model/plasma/total_radiated_power.pyx @@ -73,11 +73,11 @@ cdef class TotalRadiatedPower(PlasmaModel): self._change() @property - def element(self): + def element(self) -> Element: return self._element @property - def charge(self): + def charge(self) -> int: return self._charge cpdef Spectrum emission(self, Point3D point, Vector3D direction, Spectrum spectrum): From b24018f31dcc4fcadcb37e5ed02316742415e1d6 Mon Sep 17 00:00:00 2001 From: MatejTomes Date: Mon, 3 Nov 2025 00:21:02 +0100 Subject: [PATCH 53/91] Add hooks to avoid file reading --- cherab/core/model/plasma/tests/test_models.py | 67 ++++++++++++++++--- 1 file changed, 57 insertions(+), 10 deletions(-) diff --git a/cherab/core/model/plasma/tests/test_models.py b/cherab/core/model/plasma/tests/test_models.py index 56203db0..06dbd2d0 100644 --- a/cherab/core/model/plasma/tests/test_models.py +++ b/cherab/core/model/plasma/tests/test_models.py @@ -1,35 +1,78 @@ 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.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): +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 test_excitation(self): + def setUp(self): + # setup mock to avoid reading the data from the repository + self.patcher_pec = 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_pec = self.patcher_pec.start() + + self.patcher_rec = 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_rec = self.patcher_rec.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_pec.stop() + self.patcher_rec.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 + # check exc has the correct line self.assertEqual(exc.line, self.balmer_alpha) - #check exc has the correct lineshape + # check exc has the correct lineshape self.assertIsInstance(exc.lineshape, GaussianLine) - + + # check the mock was called + self.mock_get_pec.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] @@ -37,12 +80,16 @@ def test_recombination(self): # 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 + # check rec has the correct line self.assertEqual(rec.line, self.balmer_alpha) - #check rec has the correct lineshape + # check rec has the correct lineshape self.assertIsInstance(rec.lineshape, GaussianLine) + # check the mock was called + self.mock_get_rec.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] @@ -53,6 +100,6 @@ def test_total_radiated_power(self): with self.assertRaises(ValueError): TotalRadiatedPower(hydrogen, -1) - #check trp has the correct element and charge + # check trp has the correct element and charge self.assertEqual(trp.element, hydrogen) self.assertEqual(trp.charge, 0) From c3b936089bbc0c9cb9f8b032d5cfc92f3757f7ea Mon Sep 17 00:00:00 2001 From: MatejTomes Date: Mon, 3 Nov 2025 00:33:05 +0100 Subject: [PATCH 54/91] Improve naming --- cherab/core/model/plasma/tests/test_models.py | 16 ++++++++-------- 1 file changed, 8 insertions(+), 8 deletions(-) diff --git a/cherab/core/model/plasma/tests/test_models.py b/cherab/core/model/plasma/tests/test_models.py index 06dbd2d0..b0a15417 100644 --- a/cherab/core/model/plasma/tests/test_models.py +++ b/cherab/core/model/plasma/tests/test_models.py @@ -25,7 +25,7 @@ class TestPlasmaModels(unittest.TestCase): def setUp(self): # setup mock to avoid reading the data from the repository - self.patcher_pec = patch( + self.patcher_excitation = patch( "cherab.openadas.openadas.repository.get_pec_excitation_rate", return_value={ "ne": np.linspace(1e18, 1e20, 10), @@ -33,9 +33,9 @@ def setUp(self): "rate": np.ones((10, 12)), }, ) - self.mock_get_pec = self.patcher_pec.start() + self.mock_get_excitation = self.patcher_excitation.start() - self.patcher_rec = patch( + self.patcher_recombination = patch( "cherab.openadas.openadas.repository.get_pec_recombination_rate", return_value={ "ne": np.linspace(1e18, 1e20, 10), @@ -43,7 +43,7 @@ def setUp(self): "rate": np.ones((10, 12)), }, ) - self.mock_get_rec = self.patcher_rec.start() + self.mock_get_recombination = self.patcher_recombination.start() self.patcher_wl = patch( "cherab.openadas.openadas.repository.get_wavelength", return_value=656.28 @@ -52,8 +52,8 @@ def setUp(self): def tearDown(self): # stop the mocks after a test is run - self.patcher_pec.stop() - self.patcher_rec.stop() + self.patcher_excitation.stop() + self.patcher_recombination.stop() self.patcher_wl.stop() def test_excitation(self): @@ -70,7 +70,7 @@ def test_excitation(self): self.assertIsInstance(exc.lineshape, GaussianLine) # check the mock was called - self.mock_get_pec.assert_called_once() + self.mock_get_excitation.assert_called_once() self.assertEqual(self.mock_get_wavelength.call_count, 2) def test_recombination(self): @@ -87,7 +87,7 @@ def test_recombination(self): self.assertIsInstance(rec.lineshape, GaussianLine) # check the mock was called - self.mock_get_rec.assert_called_once() + self.mock_get_recombination.assert_called_once() self.assertEqual(self.mock_get_wavelength.call_count, 2) def test_total_radiated_power(self): From 5810951d0af0ece11a0789341783d038b9d5ab1a Mon Sep 17 00:00:00 2001 From: munechika-koyo Date: Mon, 3 Nov 2025 10:19:05 +0100 Subject: [PATCH 55/91] Re-enable Python 3.14 in CI workflow matrix --- .github/workflows/ci.yml | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index 77b05087..25e2467d 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -11,7 +11,7 @@ jobs: strategy: fail-fast: false matrix: - python-version: ["3.9", "3.10", "3.11", "3.12", "3.13"] #, "3.14"] # TODO: re-enable 3.14 when pyopencl supports it + python-version: ["3.9", "3.10", "3.11", "3.12", "3.13", "3.14"] steps: - name: Checkout code uses: actions/checkout@v2 From b548957df07c03400be87c867cd56da6e890010a Mon Sep 17 00:00:00 2001 From: MatejTomes Date: Thu, 13 Nov 2025 17:24:42 +0100 Subject: [PATCH 56/91] Update TMC and developer list --- README.md | 11 +++++++++-- docs/source/welcome.rst | 17 +++++++++++++---- 2 files changed, 22 insertions(+), 6 deletions(-) diff --git a/README.md b/README.md index 177feb9b..e9508595 100644 --- a/README.md +++ b/README.md @@ -85,12 +85,19 @@ of code / physics algorithm quality standards. TMC Members ----------- +- Jack Lovell (Oak Ridge, USA) +- Matej Tomes (IPP, Czechia) +- Koyo Munechika (ITER Organisation) +- Jakub Svoboda (IPP, Czechia) + + +Honorary 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/docs/source/welcome.rst b/docs/source/welcome.rst index 4e4a00e1..845956d1 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, USA) +* Vlad Neverov +* Matej Tomes (IPP, Czechia) +* Koyo Munechika (ITER Organisation) +* Jakub Svoboda (IPP, Czechia) Contributors @@ -37,6 +38,14 @@ Contributors * Andy Meigs (Physics) +Honorary Developers +----------------------------------- +* Matthew Carr (Core Developer) +* Alex Meakins (Architect/Core Developer) +* Alfonso Baciero (Model development) +* Carine Giroud (JET Project Management) + + Project History --------------- From cab1c4321ab07ffe9b585cd497780ac651be7dc2 Mon Sep 17 00:00:00 2001 From: MatejTomes Date: Thu, 13 Nov 2025 22:21:55 +0100 Subject: [PATCH 57/91] Use the full and correct affiliation Use Former instead of Honorary --- README.md | 8 ++++---- docs/source/welcome.rst | 8 ++++---- 2 files changed, 8 insertions(+), 8 deletions(-) diff --git a/README.md b/README.md index e9508595..69e076d8 100644 --- a/README.md +++ b/README.md @@ -85,13 +85,13 @@ of code / physics algorithm quality standards. TMC Members ----------- -- Jack Lovell (Oak Ridge, USA) -- Matej Tomes (IPP, Czechia) +- 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 (IPP, Czechia) +- Jakub Svoboda (Institute of Plasma Physics of the Czech Academy of Sciences, Czechia) -Honorary TMC Members +Former TMC Members ----------- - Alys Brett (chairwoman, master account holder, responsible for delegation, UKAEA, UK) diff --git a/docs/source/welcome.rst b/docs/source/welcome.rst index 845956d1..0317f954 100644 --- a/docs/source/welcome.rst +++ b/docs/source/welcome.rst @@ -20,11 +20,11 @@ The following authors have contributed to the project: Current Development Team ------------------------ -* Jack Lovell (Oak Ridge, USA) +* Jack Lovell (Oak Ridge National Laboratory, USA) * Vlad Neverov -* Matej Tomes (IPP, Czechia) +* Matej Tomes (Institute of Plasma Physics of the Czech Academy of Sciences, Czechia) * Koyo Munechika (ITER Organisation) -* Jakub Svoboda (IPP, Czechia) +* Jakub Svoboda (Institute of Plasma Physics of the Czech Academy of Sciences, Czechia) Contributors @@ -38,7 +38,7 @@ Contributors * Andy Meigs (Physics) -Honorary Developers +Former Developers ----------------------------------- * Matthew Carr (Core Developer) * Alex Meakins (Architect/Core Developer) From f0a38db645a442250d13b8ed23d12544a7134dc9 Mon Sep 17 00:00:00 2001 From: munechika-koyo Date: Thu, 27 Nov 2025 10:40:52 +0100 Subject: [PATCH 58/91] Add TargetedPixelGroup class to observers documentation --- docs/source/tools/observers.rst | 3 +++ 1 file changed, 3 insertions(+) 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: From f2ee0d3754eb87477255b3ffef01dc37bcc17556 Mon Sep 17 00:00:00 2001 From: munechika-koyo Date: Fri, 28 Nov 2025 14:40:33 +0100 Subject: [PATCH 59/91] Update raysect dependency version to 0.9.1.* --- pixi.toml | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/pixi.toml b/pixi.toml index f75a8e15..6ff8ae88 100644 --- a/pixi.toml +++ b/pixi.toml @@ -26,7 +26,7 @@ python = "*" setuptools = "*" cython = ">=3.1" numpy = "*" -raysect = "0.9.*" +raysect = "0.9.1.*" [package.run-dependencies] scipy = "*" From 6f9dddd37d3eb0391a356a245ec8a1abef524314 Mon Sep 17 00:00:00 2001 From: munechika-koyo Date: Fri, 28 Nov 2025 14:50:20 +0100 Subject: [PATCH 60/91] Revert "Update raysect dependency version to 0.9.1.*" This reverts commit f2ee0d3754eb87477255b3ffef01dc37bcc17556. --- pixi.toml | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/pixi.toml b/pixi.toml index 6ff8ae88..f75a8e15 100644 --- a/pixi.toml +++ b/pixi.toml @@ -26,7 +26,7 @@ python = "*" setuptools = "*" cython = ">=3.1" numpy = "*" -raysect = "0.9.1.*" +raysect = "0.9.*" [package.run-dependencies] scipy = "*" From 01dc361ce5bfc195bf6d49528cb118c952f4a7b2 Mon Sep 17 00:00:00 2001 From: Koyo MUNECHIKA <51052381+munechika-koyo@users.noreply.github.com> Date: Thu, 28 May 2026 17:19:37 +0200 Subject: [PATCH 61/91] Fix ADF15 parser and refine regex patterns (#497) * Fix ADF15 parser Update regex patterns to handle other raw files (e.g. pec40#w_ic#w0.dat) * Fix regressions for config regex pattern in _scrape_metadata_full * Move regex pattern definition out of for-loop * Refine regex pattern for configuration string in _scrape_metadata_full * Add unit tests for ADF15 parser with mock data for hydrogen, carbon, and tungsten formats * Refactor string literals to use double quotes in ADF15 parser and tests --- cherab/openadas/parse/adf15.py | 165 ++++++++-------- cherab/openadas/tests/test_adf15.py | 296 ++++++++++++++++++++++++++++ 2 files changed, 381 insertions(+), 80 deletions(-) create mode 100644 cherab/openadas/tests/test_adf15.py 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() From 4adb0aaf0bbef89d92288c8cf793ca964874e02f Mon Sep 17 00:00:00 2001 From: Koyo MUNECHIKA <51052381+munechika-koyo@users.noreply.github.com> Date: Fri, 5 Jun 2026 16:47:23 +0200 Subject: [PATCH 62/91] =?UTF-8?q?=F0=9F=97=91=EF=B8=8F=20Remove=20=5F=5Fin?= =?UTF-8?q?it=5F=5F.py=20to=20comply=20with=20PEP=20420=20(#500)?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit * Remove __init__.py file to follow pep420 * Update test script to use unittest discovery to handle the namespace packge * Remove namespace_packages from setup.py --- cherab/__init__.py | 19 ------------------- dev/test.sh | 2 +- setup.py | 1 - 3 files changed, 1 insertion(+), 21 deletions(-) delete mode 100644 cherab/__init__.py diff --git a/cherab/__init__.py b/cherab/__init__.py deleted file mode 100644 index 53af8cd3..00000000 --- a/cherab/__init__.py +++ /dev/null @@ -1,19 +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__('pkg_resources').declare_namespace(__name__) 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/setup.py b/setup.py index c57bb0d3..bee897ed 100644 --- a/setup.py +++ b/setup.py @@ -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", From e7f6415c85df39a197b2d16b1ed2792f34a1bf94 Mon Sep 17 00:00:00 2001 From: Koyo MUNECHIKA <51052381+munechika-koyo@users.noreply.github.com> Date: Mon, 15 Jun 2026 11:44:00 +0200 Subject: [PATCH 63/91] Update documentation for cherab-iter and cherab-imas packages (#501) --- docs/source/available_modules.rst | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/docs/source/available_modules.rst b/docs/source/available_modules.rst index c99ffe1f..88957276 100644 --- a/docs/source/available_modules.rst +++ b/docs/source/available_modules.rst @@ -35,8 +35,7 @@ Fusion Experiment Packages * - `cherab-compass `_ - The Cherab configuration package for COMPASS. * - `cherab-iter `_ - - Integrates Cherab with IMAS and provides diagnostic configuration - for 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). From cf03ee7c1d8beec5b9950bfb9bf0d7fef23ba8ec Mon Sep 17 00:00:00 2001 From: munechika-koyo Date: Mon, 15 Jun 2026 17:26:46 +0200 Subject: [PATCH 64/91] Refactor function signatures and improve docstring clarity across multiple modules --- cherab/core/atomic/gaunt.pxd | 2 +- cherab/core/atomic/gaunt.pyx | 10 +++++----- cherab/core/atomic/zeeman.pyx | 2 +- cherab/core/math/integrators/integrators1d.pyx | 2 +- cherab/core/math/integrators/integrators2d.pyx | 2 +- cherab/core/math/mask.pyx | 3 --- cherab/core/math/samplers.pyx | 12 ++++++------ cherab/core/math/transform/periodic.pyx | 4 ++-- cherab/tools/primitives/axisymmetric_mesh.pyx | 6 +++--- 9 files changed, 20 insertions(+), 23 deletions(-) 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/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/math/integrators/integrators1d.pyx b/cherab/core/math/integrators/integrators1d.pyx index 7ff9be74..6b52157c 100644 --- a/cherab/core/math/integrators/integrators1d.pyx +++ b/cherab/core/math/integrators/integrators1d.pyx @@ -39,7 +39,7 @@ cdef class Integrator1D: """ A 1D function to integrate. - :rtype: int + :rtype: Function1D """ return self.function diff --git a/cherab/core/math/integrators/integrators2d.pyx b/cherab/core/math/integrators/integrators2d.pyx index 58627afe..993decc2 100644 --- a/cherab/core/math/integrators/integrators2d.pyx +++ b/cherab/core/math/integrators/integrators2d.pyx @@ -33,7 +33,7 @@ cdef class Integrator2D: """ A 2D function to integrate. - :rtype: int + :rtype: Function2D """ return self.function 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/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/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 From c2cfe4c16528ce836ef490a604faba06d360c757 Mon Sep 17 00:00:00 2001 From: munechika-koyo Date: Tue, 7 Jul 2026 17:13:22 +0200 Subject: [PATCH 65/91] Add unit tests for ZeemanStructure initialization and method behavior --- cherab/core/atomic/tests/test_zeeman.py | 70 +++++++++++++++++++++++++ 1 file changed, 70 insertions(+) create mode 100644 cherab/core/atomic/tests/test_zeeman.py 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') From 15ae09b50ba036cf64ac9078991d62792520afdd Mon Sep 17 00:00:00 2001 From: Jack Lovell Date: Tue, 7 Jul 2026 16:30:49 +0100 Subject: [PATCH 66/91] Fix namespace package detection in setup.py (#504) Setuptools deprecated the `namespace_packages` option to `setup()`, so it was removed in #500. But `find_packages` then ignores any directories without an `__init__.py` which meant that most of the Cherab files (with the exception of the compiled .so libraries) were left out when producing a wheel. Fix this by replacing `find_packages` with `find_namespace_packages`, which correctly handles PEP420-style namespace packages. --- setup.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/setup.py b/setup.py index bee897ed..952cfbbe 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') @@ -125,7 +125,7 @@ # Running ./dev/build_docs.sh runs setup.py, which requires cython. "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 From c145a4d1776d87fd93e2545a57d69446ec041430 Mon Sep 17 00:00:00 2001 From: Koyo MUNECHIKA <51052381+munechika-koyo@users.noreply.github.com> Date: Wed, 15 Jul 2026 16:45:32 +0900 Subject: [PATCH 67/91] Update Python version range in pixi.toml --- pixi.toml | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/pixi.toml b/pixi.toml index f75a8e15..c6d78fb8 100644 --- a/pixi.toml +++ b/pixi.toml @@ -18,7 +18,7 @@ noarch = false compilers = ["c"] [workspace.build-variants] -python = ["3.9", "3.10.*", "3.11.*", "3.12.*", "3.13.*"] +python = ["3.9.*", "3.10.*", "3.11.*", "3.12.*", "3.13.*", "3.14.*"] [package.host-dependencies] uv = "*" From 8c8a189047bd9c51c72a58f9212c92cf3d028a98 Mon Sep 17 00:00:00 2001 From: Koyo MUNECHIKA <51052381+munechika-koyo@users.noreply.github.com> Date: Wed, 15 Jul 2026 16:47:17 +0900 Subject: [PATCH 68/91] Add HTML file removal to build artifacts cleanup Added support for removing HTML build artifacts. --- pixi.toml | 3 +++ 1 file changed, 3 insertions(+) diff --git a/pixi.toml b/pixi.toml index c6d78fb8..1a5d3f9e 100644 --- a/pixi.toml +++ b/pixi.toml @@ -55,6 +55,9 @@ clean = { cmd = [ "-o", "-name", "'*.dylib'", + "-o", + "-name", + "'*.html'", "\\)", "-delete", ], description = "🔥 Remove in-place build artifacts and temporary files (*.c, *.so, *.dylib)" } From 053cf8c37e98bb2de562c458d3aac7680339fc26 Mon Sep 17 00:00:00 2001 From: Koyo MUNECHIKA <51052381+munechika-koyo@users.noreply.github.com> Date: Wed, 15 Jul 2026 16:53:00 +0900 Subject: [PATCH 69/91] Update Python version features in pixi.toml Renaming python-featured environments to avoid specific Python versions --- pixi.toml | 26 +++++++++++++------------- 1 file changed, 13 insertions(+), 13 deletions(-) diff --git a/pixi.toml b/pixi.toml index 1a5d3f9e..d913d686 100644 --- a/pixi.toml +++ b/pixi.toml @@ -133,24 +133,24 @@ 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.py39.dependencies] +[feature.pyoldest.dependencies] python = "3.9.*" -[feature.py313.dependencies] -python = "3.13.*" +[feature.pylatest.dependencies] +python = "3.14.*" [environments] -default = { features = ["py313"], solve-group = "py313" } -test = { features = ["test"], solve-group = "py313" } +default = { features = ["pylatest"], solve-group = "pylatest" } +test = { features = ["test"], solve-group = "pylatest" } docs = { features = [ - "py39", + "pyoldest", "docs", -], solve-group = "py39" } # TODO: change to py313 when bumping RTD theme to >=1.0 -test-py313 = { features = [ - "py313", +], solve-group = "pyoldest" } # TODO: change to pylatest when bumping RTD theme to >=1.0 +test-pylatest = { features = [ + "pylatest", "test", -], solve-group = "py313" } # alias of tests -test-py39 = { features = [ - "py39", +], solve-group = "pylatest" } # alias of test +test-pyoldest = { features = [ + "pyoldest", "test", -], solve-group = "py39" } # alias of tests +], solve-group = "pyoldest" } lint = { features = ["lint"], no-default-feature = true } From f90feae427895d6433c5a7223be54ba99275eeb9 Mon Sep 17 00:00:00 2001 From: Koyo MUNECHIKA <51052381+munechika-koyo@users.noreply.github.com> Date: Wed, 15 Jul 2026 16:56:47 +0900 Subject: [PATCH 70/91] Update channels and remove sphinx-tabs dependency Updated channels to use 'conda-forge' instead of the prefix's URL. Removed sphinx-tabs from pypi-dependencies. --- pixi.toml | 7 ++----- 1 file changed, 2 insertions(+), 5 deletions(-) diff --git a/pixi.toml b/pixi.toml index d913d686..1b84805b 100644 --- a/pixi.toml +++ b/pixi.toml @@ -1,5 +1,5 @@ [workspace] -channels = ["https://prefix.dev/conda-forge"] +channels = ["conda-forge"] platforms = ["linux-64", "osx-arm64", "osx-64"] preview = ["pixi-build"] @@ -21,7 +21,6 @@ compilers = ["c"] python = ["3.9.*", "3.10.*", "3.11.*", "3.12.*", "3.13.*", "3.14.*"] [package.host-dependencies] -uv = "*" python = "*" setuptools = "*" cython = ">=3.1" @@ -89,9 +88,7 @@ test = { cmd = "python -m unittest discover cherab -v", description = "🧪 Run cherab = { path = "." } sphinx = "*" sphinx_rtd_theme = "<1" - -[feature.docs.pypi-dependencies] -sphinx-tabs = "*" # >=3.4.4 has not yet been released to conda-forge +sphinx-tabs = "*" [feature.docs.tasks] doc-build = { cmd = [ From 1db12a2a3b781874ec6289657a0d55890b3bcaa6 Mon Sep 17 00:00:00 2001 From: munechika-koyo Date: Tue, 28 Jul 2026 11:26:48 +0200 Subject: [PATCH 71/91] Fix import statement for netcdf_file in calcam.py --- cherab/tools/observers/calcam.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) 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 From 6a73f81725927a046d430847597f5839ff74d963 Mon Sep 17 00:00:00 2001 From: munechika-koyo Date: Wed, 29 Jul 2026 12:57:21 +0200 Subject: [PATCH 72/91] Update CHANGELOG.md to include bug fix for netcdf_file import statement in calcam.py --- CHANGELOG.md | 3 +++ 1 file changed, 3 insertions(+) diff --git a/CHANGELOG.md b/CHANGELOG.md index d55ef76e..e1071975 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -8,6 +8,9 @@ API changes: * 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) +Bug fixes: +* Fix the import statement for `netcdf_file` in `calcam.py` for compatibility with the upcoming `scipy` v2.0.0. (#510) + New: * Add Function6D framework. (#478) * Add e_field attribute to Plasma object for electric field vector. (#465) From a684e5a249da28271ad41309458103e18ef0493a Mon Sep 17 00:00:00 2001 From: Jack Lovell Date: Wed, 29 Jul 2026 12:09:19 -0400 Subject: [PATCH 73/91] Make constants in cherab.core.utility.constants accessible to Python This enables single-sourcing constant values in both Cython code and Python code (including tests). Fixes #509. --- CHANGELOG.md | 1 + cherab/core/utility/constants.pyx | 49 ++++++++++++ cherab/core/utility/tests/test_constants.py | 84 +++++++++++++++++++++ 3 files changed, 134 insertions(+) create mode 100644 cherab/core/utility/tests/test_constants.py diff --git a/CHANGELOG.md b/CHANGELOG.md index e1071975..04ef3fc3 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -17,6 +17,7 @@ New: * Add Integrator2D base class for integration of two-dimensional functions. (#472) * 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) +* Make values in `cherab.core.utility.constants` accessible to Python. (#509) Release 1.5.0 (27 Aug 2024) ------------------- 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 From 76a6d1219df13bf0e73f87097f4280b8f6533eeb Mon Sep 17 00:00:00 2001 From: Matej Tomes Date: Mon, 3 Aug 2026 12:37:09 +0200 Subject: [PATCH 74/91] Reorder parameters of Function6D to x,y,z,u,v,w (#513) * Reorder parameters of Function6D to x,y,z,u,v,w --- .../math/function/float/function6d/arg.pxd | 2 +- .../math/function/float/function6d/arg.pyx | 24 +- .../function/float/function6d/autowrap.pyx | 4 +- .../math/function/float/function6d/base.pxd | 2 +- .../math/function/float/function6d/base.pyx | 124 ++++---- .../math/function/float/function6d/blend.pyx | 14 +- .../math/function/float/function6d/cmath.pyx | 42 +-- .../function/float/function6d/constant.pyx | 2 +- .../float/function6d/tests/test_arg.py | 18 +- .../float/function6d/tests/test_autowrap.py | 2 +- .../float/function6d/tests/test_base.py | 300 +++++++++--------- .../float/function6d/tests/test_cmath.py | 52 +-- 12 files changed, 293 insertions(+), 293 deletions(-) diff --git a/cherab/core/math/function/float/function6d/arg.pxd b/cherab/core/math/function/float/function6d/arg.pxd index 151eea65..527a9c0b 100644 --- a/cherab/core/math/function/float/function6d/arg.pxd +++ b/cherab/core/math/function/float/function6d/arg.pxd @@ -21,7 +21,7 @@ from cherab.core.math.function.float.function6d.base cimport Function6D cdef enum ArgLabel: - X, Y, Z, U, W, V + 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 index c4883e0e..f49053c7 100644 --- a/cherab/core/math/function/float/function6d/arg.pyx +++ b/cherab/core/math/function/float/function6d/arg.pyx @@ -28,7 +28,7 @@ cdef class Arg6D(Function6D): 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", "w", or "v". + Valid options for argument are "x", "y", "z", "u", "v", or "w". >>> argx = Arg6D("x") >>> argx(2, 3, 5, 7, 11, 13) @@ -42,14 +42,14 @@ cdef class Arg6D(Function6D): >>> argu = Arg6D("u") >>> argu(2, 3, 5, 7, 11, 13) 7.0 - >>> argw = Arg6D("w") - >>> argw(2, 3, 5, 7, 11, 13) - 11.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", "w", or "v", the argument to return + :param str argument: either "x", "y", "z", "u", "v", or "w", the argument to return """ def __init__(self, object argument): if argument == "x": @@ -60,14 +60,14 @@ cdef class Arg6D(Function6D): self._argument = Z elif argument == "u": self._argument = U - elif argument == "w": - self._argument = W 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', 'w' or 'v'") + 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 w, double v) except? -1e999: + 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: @@ -76,7 +76,7 @@ cdef class Arg6D(Function6D): return z elif self._argument == U: return u - elif self._argument == W: - return w - else: # V + elif self._argument == V: return v + else: # W + return w diff --git a/cherab/core/math/function/float/function6d/autowrap.pyx b/cherab/core/math/function/float/function6d/autowrap.pyx index 9026a0d8..080f2a96 100644 --- a/cherab/core/math/function/float/function6d/autowrap.pyx +++ b/cherab/core/math/function/float/function6d/autowrap.pyx @@ -46,8 +46,8 @@ cdef class PythonFunction6D(Function6D): def __init__(self, object function): self.function = function - cdef double evaluate(self, double x, double y, double z, double u, double w, double v) except? -1e999: - return self.function(x, y, z, u, w, v) + 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): diff --git a/cherab/core/math/function/float/function6d/base.pxd b/cherab/core/math/function/float/function6d/base.pxd index 4f482baa..ef05f0f7 100644 --- a/cherab/core/math/function/float/function6d/base.pxd +++ b/cherab/core/math/function/float/function6d/base.pxd @@ -23,7 +23,7 @@ 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 w, double v) except? -1e999 + cdef double evaluate(self, double x, double y, double z, double u, double v, double w) except? -1e999 cdef class AddFunction6D(Function6D): diff --git a/cherab/core/math/function/float/function6d/base.pyx b/cherab/core/math/function/float/function6d/base.pyx index fd653cd9..a97d9583 100644 --- a/cherab/core/math/function/float/function6d/base.pyx +++ b/cherab/core/math/function/float/function6d/base.pyx @@ -40,21 +40,21 @@ cdef class Function6D(FloatFunction): that accepts a function object. """ - cdef double evaluate(self, double x, double y, double z, double u, double w, double v) except? -1e999: + 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 w, double v): - """ Evaluate the function f(x, y, z, u, w, v) + 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 w: function parameter w :param float v: function parameter v + :param float w: function parameter w :rtype: float """ - return self.evaluate(x, y, z, u, w, v) + return self.evaluate(x, y, z, u, v, w) def __add__(self, object b): if is_callable(b): # a() + b() @@ -224,8 +224,8 @@ cdef class AddFunction6D(Function6D): self._function1 = autowrap_function6d(function1) self._function2 = autowrap_function6d(function2) - cdef double evaluate(self, double x, double y, double z, double u, double w, double v) except? -1e999: - return self._function1.evaluate(x, y, z, u, w, v) + self._function2.evaluate(x, y, z, u, w, v) + 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): @@ -243,8 +243,8 @@ cdef class SubtractFunction6D(Function6D): self._function1 = autowrap_function6d(function1) self._function2 = autowrap_function6d(function2) - cdef double evaluate(self, double x, double y, double z, double u, double w, double v) except? -1e999: - return self._function1.evaluate(x, y, z, u, w, v) - self._function2.evaluate(x, y, z, u, w, v) + 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): @@ -262,8 +262,8 @@ cdef class MultiplyFunction6D(Function6D): self._function1 = autowrap_function6d(function1) self._function2 = autowrap_function6d(function2) - cdef double evaluate(self, double x, double y, double z, double u, double w, double v) except? -1e999: - return self._function1.evaluate(x, y, z, u, w, v) * self._function2.evaluate(x, y, z, u, w, v) + 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): @@ -282,11 +282,11 @@ cdef class DivideFunction6D(Function6D): self._function2 = autowrap_function6d(function2) @cython.cdivision(True) - cdef double evaluate(self, double x, double y, double z, double u, double w, double v) except? -1e999: - cdef double denominator = self._function2.evaluate(x, y, z, u, w, v) + 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, w, v) / denominator + return self._function1.evaluate(x, y, z, u, v, w) / denominator cdef class ModuloFunction6D(Function6D): @@ -304,11 +304,11 @@ cdef class ModuloFunction6D(Function6D): self._function2 = autowrap_function6d(function2) @cython.cdivision(True) - cdef double evaluate(self, double x, double y, double z, double u, double w, double v) except? -1e999: - cdef double divisor = self._function2.evaluate(x, y, z, u, w, v) + 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, w, v) % divisor + return self._function1.evaluate(x, y, z, u, v, w) % divisor cdef class PowFunction6D(Function6D): @@ -325,10 +325,10 @@ cdef class PowFunction6D(Function6D): self._function1 = autowrap_function6d(function1) self._function2 = autowrap_function6d(function2) - cdef double evaluate(self, double x, double y, double z, double u, double w, double v) except? -1e999: + 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, w, v) - exponent = self._function2.evaluate(x, y, z, u, w, v) + 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: @@ -348,8 +348,8 @@ cdef class AbsFunction6D(Function6D): def __init__(self, object function): self._function = autowrap_function6d(function) - cdef double evaluate(self, double x, double y, double z, double u, double w, double v) except? -1e999: - return abs(self._function.evaluate(x, y, z, u, w, v)) + 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): @@ -366,8 +366,8 @@ cdef class EqualsFunction6D(Function6D): self._function1 = autowrap_function6d(function1) self._function2 = autowrap_function6d(function2) - cdef double evaluate(self, double x, double y, double z, double u, double w, double v) except? -1e999: - return self._function1.evaluate(x, y, z, u, w, v) == self._function2.evaluate(x, y, z, u, w, v) + 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): @@ -384,8 +384,8 @@ cdef class NotEqualsFunction6D(Function6D): self._function1 = autowrap_function6d(function1) self._function2 = autowrap_function6d(function2) - cdef double evaluate(self, double x, double y, double z, double u, double w, double v) except? -1e999: - return self._function1.evaluate(x, y, z, u, w, v) != self._function2.evaluate(x, y, z, u, w, v) + 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): @@ -402,8 +402,8 @@ cdef class LessThanFunction6D(Function6D): self._function1 = autowrap_function6d(function1) self._function2 = autowrap_function6d(function2) - cdef double evaluate(self, double x, double y, double z, double u, double w, double v) except? -1e999: - return self._function1.evaluate(x, y, z, u, w, v) < self._function2.evaluate(x, y, z, u, w, v) + 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): @@ -420,8 +420,8 @@ cdef class GreaterThanFunction6D(Function6D): self._function1 = autowrap_function6d(function1) self._function2 = autowrap_function6d(function2) - cdef double evaluate(self, double x, double y, double z, double u, double w, double v) except? -1e999: - return self._function1.evaluate(x, y, z, u, w, v) > self._function2.evaluate(x, y, z, u, w, v) + 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): @@ -438,8 +438,8 @@ cdef class LessEqualsFunction6D(Function6D): self._function1 = autowrap_function6d(function1) self._function2 = autowrap_function6d(function2) - cdef double evaluate(self, double x, double y, double z, double u, double w, double v) except? -1e999: - return self._function1.evaluate(x, y, z, u, w, v) <= self._function2.evaluate(x, y, z, u, w, v) + 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): @@ -456,8 +456,8 @@ cdef class GreaterEqualsFunction6D(Function6D): self._function1 = autowrap_function6d(function1) self._function2 = autowrap_function6d(function2) - cdef double evaluate(self, double x, double y, double z, double u, double w, double v) except? -1e999: - return self._function1.evaluate(x, y, z, u, w, v) >= self._function2.evaluate(x, y, z, u, w, v) + 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): @@ -475,8 +475,8 @@ cdef class AddScalar6D(Function6D): self._value = value self._function = autowrap_function6d(function) - cdef double evaluate(self, double x, double y, double z, double u, double w, double v) except? -1e999: - return self._value + self._function.evaluate(x, y, z, u, w, v) + 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): @@ -494,8 +494,8 @@ cdef class SubtractScalar6D(Function6D): self._value = value self._function = autowrap_function6d(function) - cdef double evaluate(self, double x, double y, double z, double u, double w, double v) except? -1e999: - return self._value - self._function.evaluate(x, y, z, u, w, v) + 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): @@ -513,8 +513,8 @@ cdef class MultiplyScalar6D(Function6D): self._value = value self._function = autowrap_function6d(function) - cdef double evaluate(self, double x, double y, double z, double u, double w, double v) except? -1e999: - return self._value * self._function.evaluate(x, y, z, u, w, v) + 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): @@ -533,8 +533,8 @@ cdef class DivideScalar6D(Function6D): self._function = autowrap_function6d(function) @cython.cdivision(True) - cdef double evaluate(self, double x, double y, double z, double u, double w, double v) except? -1e999: - cdef double denominator = self._function.evaluate(x, y, z, u, w, v) + 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 @@ -555,8 +555,8 @@ cdef class ModuloScalarFunction6D(Function6D): self._function = autowrap_function6d(function) @cython.cdivision(True) - cdef double evaluate(self, double x, double y, double z, double u, double w, double v) except? -1e999: - cdef double divisor = self._function.evaluate(x, y, z, u, w, v) + 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 @@ -579,8 +579,8 @@ cdef class ModuloFunctionScalar6D(Function6D): self._function = autowrap_function6d(function) @cython.cdivision(True) - cdef double evaluate(self, double x, double y, double z, double u, double w, double v) except? -1e999: - return self._function.evaluate(x, y, z, u, w, v) % self._value + 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): @@ -597,8 +597,8 @@ cdef class PowScalarFunction6D(Function6D): self._value = value self._function = autowrap_function6d(function) - cdef double evaluate(self, double x, double y, double z, double u, double w, double v) except? -1e999: - cdef double exponent = self._function.evaluate(x, y, z, u, w, v) + 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: @@ -620,8 +620,8 @@ cdef class PowFunctionScalar6D(Function6D): self._value = value self._function = autowrap_function6d(function) - cdef double evaluate(self, double x, double y, double z, double u, double w, double v) except? -1e999: - cdef double base = self._function.evaluate(x, y, z, u, w, v) + 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: @@ -643,8 +643,8 @@ cdef class EqualsScalar6D(Function6D): self._value = value self._function = autowrap_function6d(function) - cdef double evaluate(self, double x, double y, double z, double u, double w, double v) except? -1e999: - return self._value == self._function.evaluate(x, y, z, u, w, v) + 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): @@ -661,8 +661,8 @@ cdef class NotEqualsScalar6D(Function6D): self._value = value self._function = autowrap_function6d(function) - cdef double evaluate(self, double x, double y, double z, double u, double w, double v) except? -1e999: - return self._value != self._function.evaluate(x, y, z, u, w, v) + 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): @@ -679,8 +679,8 @@ cdef class LessThanScalar6D(Function6D): self._value = value self._function = autowrap_function6d(function) - cdef double evaluate(self, double x, double y, double z, double u, double w, double v) except? -1e999: - return self._value < self._function.evaluate(x, y, z, u, w, v) + 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): @@ -697,8 +697,8 @@ cdef class GreaterThanScalar6D(Function6D): self._value = value self._function = autowrap_function6d(function) - cdef double evaluate(self, double x, double y, double z, double u, double w, double v) except? -1e999: - return self._value > self._function.evaluate(x, y, z, u, w, v) + 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): @@ -715,8 +715,8 @@ cdef class LessEqualsScalar6D(Function6D): self._value = value self._function = autowrap_function6d(function) - cdef double evaluate(self, double x, double y, double z, double u, double w, double v) except? -1e999: - return self._value <= self._function.evaluate(x, y, z, u, w, v) + 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): @@ -733,5 +733,5 @@ cdef class GreaterEqualsScalar6D(Function6D): self._value = value self._function = autowrap_function6d(function) - cdef double evaluate(self, double x, double y, double z, double u, double w, double v) except? -1e999: - return self._value >= self._function.evaluate(x, y, z, u, w, v) + 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.pyx b/cherab/core/math/function/float/function6d/blend.pyx index e05c0452..550feddd 100644 --- a/cherab/core/math/function/float/function6d/blend.pyx +++ b/cherab/core/math/function/float/function6d/blend.pyx @@ -31,7 +31,7 @@ cdef class Blend6D(Function6D): this function is as follows: .. math:: - v = (1 - f_m(x, y, z, u, w, v)) f_1(x, y, z, u, w, v) + f_m(x, y, z, u, w, v) f_2(x, y, z, u, w, v) + 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. @@ -48,18 +48,18 @@ cdef class Blend6D(Function6D): self._f2 = autowrap_function6d(f2) self._mask = autowrap_function6d(mask) - cdef double evaluate(self, double x, double y, double z, double u, double w, double v) except? -1e999: + 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, w, v), 0.0, 1.0) + 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, w, v) + return self._f1.evaluate(x, y, z, u, v, w) if t == 1: - return self._f2.evaluate(x, y, z, u, w, v) + return self._f2.evaluate(x, y, z, u, v, w) # lerp between function values - cdef double f1 = self._f1.evaluate(x, y, z, u, w, v) - cdef double f2 = self._f2.evaluate(x, y, z, u, w, v) + 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.pyx b/cherab/core/math/function/float/function6d/cmath.pyx index ccc5d9b6..b33e1680 100644 --- a/cherab/core/math/function/float/function6d/cmath.pyx +++ b/cherab/core/math/function/float/function6d/cmath.pyx @@ -32,8 +32,8 @@ cdef class Exp6D(Function6D): def __init__(self, object function): self._function = autowrap_function6d(function) - cdef double evaluate(self, double x, double y, double z, double u, double w, double v) except? -1e999: - return cmath.exp(self._function.evaluate(x, y, z, u, w, v)) + 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): @@ -45,8 +45,8 @@ cdef class Sin6D(Function6D): def __init__(self, object function): self._function = autowrap_function6d(function) - cdef double evaluate(self, double x, double y, double z, double u, double w, double v) except? -1e999: - return cmath.sin(self._function.evaluate(x, y, z, u, w, v)) + 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): @@ -58,8 +58,8 @@ cdef class Cos6D(Function6D): def __init__(self, object function): self._function = autowrap_function6d(function) - cdef double evaluate(self, double x, double y, double z, double u, double w, double v) except? -1e999: - return cmath.cos(self._function.evaluate(x, y, z, u, w, v)) + 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): @@ -71,8 +71,8 @@ cdef class Tan6D(Function6D): def __init__(self, object function): self._function = autowrap_function6d(function) - cdef double evaluate(self, double x, double y, double z, double u, double w, double v) except? -1e999: - return cmath.tan(self._function.evaluate(x, y, z, u, w, v)) + 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): @@ -84,8 +84,8 @@ cdef class Asin6D(Function6D): def __init__(self, object function): self._function = autowrap_function6d(function) - cdef double evaluate(self, double x, double y, double z, double u, double w, double v) except? -1e999: - cdef double val = self._function.evaluate(x, y, z, u, w, v) + 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].") @@ -100,8 +100,8 @@ cdef class Acos6D(Function6D): def __init__(self, object function): self._function = autowrap_function6d(function) - cdef double evaluate(self, double x, double y, double z, double u, double w, double v) except? -1e999: - cdef double val = self._function.evaluate(x, y, z, u, w, v) + 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].") @@ -116,8 +116,8 @@ cdef class Atan6D(Function6D): def __init__(self, object function): self._function = autowrap_function6d(function) - cdef double evaluate(self, double x, double y, double z, double u, double w, double v) except? -1e999: - return cmath.atan(self._function.evaluate(x, y, z, u, w, v)) + 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): @@ -134,9 +134,9 @@ cdef class Atan4Q6D(Function6D): self._numerator = autowrap_function6d(numerator) self._denominator = autowrap_function6d(denominator) - cdef double evaluate(self, double x, double y, double z, double u, double w, double v) except? -1e999: - return cmath.atan2(self._numerator.evaluate(x, y, z, u, w, v), - self._denominator.evaluate(x, y, z, u, w, v)) + 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): @@ -148,8 +148,8 @@ cdef class Sqrt6D(Function6D): def __init__(self, object function): self._function = autowrap_function6d(function) - cdef double evaluate(self, double x, double y, double z, double u, double w, double v) except? -1e999: - cdef double f = self._function.evaluate(x, y, z, u, w, v) + 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) @@ -164,5 +164,5 @@ cdef class Erf6D(Function6D): def __init__(self, object function): self._function = autowrap_function6d(function) - cdef double evaluate(self, double x, double y, double z, double u, double w, double v) except? -1e999: - return cmath.erf(self._function.evaluate(x, y, z, u, w, v)) \ No newline at end of file + 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/core/math/function/float/function6d/constant.pyx b/cherab/core/math/function/float/function6d/constant.pyx index 1464f01c..be25b44b 100644 --- a/cherab/core/math/function/float/function6d/constant.pyx +++ b/cherab/core/math/function/float/function6d/constant.pyx @@ -45,5 +45,5 @@ cdef class Constant6D(Function6D): def __init__(self, double value): self._value = value - cdef double evaluate(self, double x, double y, double z, double u, double w, double v) except? -1e999: + 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/test_arg.py b/cherab/core/math/function/float/function6d/tests/test_arg.py index a594b472..fb559064 100644 --- a/cherab/core/math/function/float/function6d/tests/test_arg.py +++ b/cherab/core/math/function/float/function6d/tests/test_arg.py @@ -31,19 +31,19 @@ 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, w, v) in itertools.product(testvals, repeat=6): + 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("w") - argv = Arg6D("v") - self.assertEqual(argx(x, y, z, u, w, v), x, "Arg6D('x') call did not match reference value.") - self.assertEqual(argy(x, y, z, u, w, v), y, "Arg6D('y') call did not match reference value.") - self.assertEqual(argz(x, y, z, u, w, v), z, "Arg6D('z') call did not match reference value.") - self.assertEqual(argu(x, y, z, u, w, v), u, "Arg6D('u') call did not match reference value.") - self.assertEqual(argw(x, y, z, u, w, v), w, "Arg6D('w') call did not match reference value.") - self.assertEqual(argv(x, y, z, u, w, v), v, "Arg6D('v') call did not match reference value.") + 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."): diff --git a/cherab/core/math/function/float/function6d/tests/test_autowrap.py b/cherab/core/math/function/float/function6d/tests/test_autowrap.py index 2ebe0463..5682f14e 100644 --- a/cherab/core/math/function/float/function6d/tests/test_autowrap.py +++ b/cherab/core/math/function/float/function6d/tests/test_autowrap.py @@ -33,5 +33,5 @@ def test_constant(self): self.assertIsInstance(function, Constant6D, "Autowrapped scalar float is not a Constant6D.") def test_python_function(self): - function = _autowrap_function6d(lambda x, y, z, u, w, v: 10*x + 5*y + 2*z + u + 3*w + 4*v) + 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 index d9ec4b13..0a572220 100644 --- a/cherab/core/math/function/float/function6d/tests/test_base.py +++ b/cherab/core/math/function/float/function6d/tests/test_base.py @@ -31,64 +31,64 @@ class TestFunction6D(unittest.TestCase): def setUp(self): - self.ref1 = lambda x, y, z, u, w, v: 10 * x + 5 * y + 2 * z + u + 3 * w + 4 * v - self.ref2 = lambda x, y, z, u, w, v: abs(x + y + z + u + w + v) + 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, w, v) in itertools.product(testvals, repeat=6): - self.assertEqual(self.f1(x, y, z, u, w, v), self.ref1(x, y, z, u, w, v), + 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, w, v) in itertools.product(testvals, repeat=6): - self.assertEqual(r(x, y, z, u, w, v), -self.ref1(x, y, z, u, w, v), + 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, w, v) in itertools.product(testvals, repeat=6): - self.assertEqual(r1(x, y, z, u, w, v), 8 + self.ref1(x, y, z, u, w, v), + 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, w, v), self.ref1(x, y, z, u, w, v) + 65, + 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, w, v) in itertools.product(testvals, repeat=6): - self.assertEqual(r1(x, y, z, u, w, v), 8 - self.ref1(x, y, z, u, w, v), + 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, w, v), self.ref1(x, y, z, u, w, v) - 65, + 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, w, v) in itertools.product(testvals, repeat=6): - self.assertEqual(r1(x, y, z, u, w, v), 5 * self.ref1(x, y, z, u, w, v), + 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, w, v), self.ref1(x, y, z, u, w, v) * -7.8, + 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, w, v) in itertools.product(testvals, repeat=6): - self.assertEqual(r1(x, y, z, u, w, v), 5.451 / self.ref1(x, y, z, u, w, v), + 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, w, v), self.ref1(x, y, z, u, w, v) / -7.8, - delta=abs(r2(x, y, z, u, w, v)) * 1e-12, + 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 @@ -102,13 +102,13 @@ def test_mod_function6d_scalar(self): testvals = [-10, -7, -0.001, 0.00003, 10, 12.3] r1 = 5 % self.f1 r2 = self.f1 % -7.8 - for (x, y, z, u, w, v) in itertools.product(testvals, repeat=6): - if self.ref1(x, y, z, u, w, v) == 0: + 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, w, v) + r1(x, y, z, u, v, w) else: - self.assertAlmostEqual(r1(x, y, z, u, w, v), math.fmod(5, self.ref1(x, y, z, u, w, v)), 15, "Function6D modulo scalar (K % f()) did not match reference function value.") - self.assertAlmostEqual(r2(x, y, z, u, w, v), math.fmod(self.ref1(x, y, z, u, w, v), -7.8), 15, "Function6D modulo scalar (f() % K) did not match reference function value.") + 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."): @@ -119,20 +119,20 @@ def test_pow_function6d_scalar(self): r1 = 5. ** self.f1 r2 = self.f1 ** -7.8 r3 = (-5.) ** self.f1 - for (x, y, z, u, w, v) in itertools.product(testvals, repeat=6): - self.assertAlmostEqual(r1(x, y, z, u, w, v), 5. ** self.ref1(x, y, z, u, w, v), 15, "Function6D power scalar (K ** f()) did not match reference function value.") - if self.ref1(x, y, z, u, w, v) < 0: + 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, w, v) - elif not float(self.ref1(x, y, z, u, w, v)).is_integer(): + 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, w, v) + r3(x, y, z, u, v, w) else: - if self.ref1(x, y, z, u, w, v) == 0: + 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, w, v) + r2(x, y, z, u, v, w) else: - self.assertAlmostEqual(r2(x, y, z, u, w, v), self.ref1(x, y, z, u, w, v) ** -7.8, 15, "Function6D power scalar (f() ** K) 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, 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."): @@ -141,120 +141,120 @@ def test_pow_function6d_scalar(self): def test_richcmp_scalar(self): testvals = [-1e10, -7, -0.001, 0.0, 0.00003, 10, 2.3e49] - for (x, y, z, u, w, v) in itertools.product(testvals, repeat=6): - ref_value = self.ref1(x, y, z, u, w, v) + 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, w, v), 1.0, + (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, w, v), 1.0, + (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, w, v), 0.0, + (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, w, v), 0.0, + (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, w, v), 1.0, + (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, w, v), 1.0, + (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, w, v), 0.0, + (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, w, v), 0.0, + (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, w, v), 1.0, + (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, w, v), 1.0, + (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, w, v), 0.0, + (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, w, v), 0.0, + (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, w, v), 1.0, + (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, w, v), 1.0, + (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, w, v), 0.0, + (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, w, v), 0.0, + (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, w, v), 1.0, + (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, w, v), 1.0, + (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, w, v), 1.0, + (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, w, v), 1.0, + (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, w, v), 0.0, + (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, w, v), 0.0, + (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, w, v), 1.0, + (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, w, v), 1.0, + (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, w, v), 1.0, + (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, w, v), 1.0, + (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, w, v), 0.0, + (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, w, v), 0.0, + (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." ) @@ -263,40 +263,40 @@ def test_add_function6d(self): r1 = self.f1 + self.f2 r2 = self.ref1 + self.f2 r3 = self.f1 + self.ref2 - for (x, y, z, u, w, v) in itertools.product(testvals, repeat=6): - self.assertEqual(r1(x, y, z, u, w, v), self.ref1(x, y, z, u, w, v) + self.ref2(x, y, z, u, w, v), "Function6D add function (f1() + f2()) did not match reference function value.") - self.assertEqual(r2(x, y, z, u, w, v), self.ref1(x, y, z, u, w, v) + self.ref2(x, y, z, u, w, v), "Function6D add function (p1() + f2()) did not match reference function value.") - self.assertEqual(r3(x, y, z, u, w, v), self.ref1(x, y, z, u, w, v) + self.ref2(x, y, z, u, w, v), "Function6D add function (f1() + p2()) did not match reference function value.") + 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, w, v) in itertools.product(testvals, repeat=6): - self.assertEqual(r1(x, y, z, u, w, v), self.ref1(x, y, z, u, w, v) - self.ref2(x, y, z, u, w, v), "Function6D subtract function (f1() - f2()) did not match reference function value.") - self.assertEqual(r2(x, y, z, u, w, v), self.ref1(x, y, z, u, w, v) - self.ref2(x, y, z, u, w, v), "Function6D subtract function (p1() - f2()) did not match reference function value.") - self.assertEqual(r3(x, y, z, u, w, v), self.ref1(x, y, z, u, w, v) - self.ref2(x, y, z, u, w, v), "Function6D subtract function (f1() - p2()) did not match reference function value.") + 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, w, v) in itertools.product(testvals, repeat=6): - self.assertEqual(r1(x, y, z, u, w, v), self.ref1(x, y, z, u, w, v) * self.ref2(x, y, z, u, w, v), "Function6D multiply function (f1() * f2()) did not match reference function value.") - self.assertEqual(r2(x, y, z, u, w, v), self.ref1(x, y, z, u, w, v) * self.ref2(x, y, z, u, w, v), "Function6D multiply function (p1() * f2()) did not match reference function value.") - self.assertEqual(r3(x, y, z, u, w, v), self.ref1(x, y, z, u, w, v) * self.ref2(x, y, z, u, w, v), "Function6D multiply function (f1() * p2()) did not match reference function value.") + 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, w, v) in itertools.product(testvals, repeat=6): - self.assertAlmostEqual(r1(x, y, z, u, w, v), self.ref1(x, y, z, u, w, v) / self.ref2(x, y, z, u, w, v), delta=abs(r1(x, y, z, u, w, v)) * 1e-12, msg="Function6D divide function (f1() / f2()) did not match reference function value.") - self.assertAlmostEqual(r2(x, y, z, u, w, v), self.ref1(x, y, z, u, w, v) / self.ref2(x, y, z, u, w, v), delta=abs(r2(x, y, z, u, w, v)) * 1e-12, msg="Function6D divide function (p1() / f2()) did not match reference function value.") - self.assertAlmostEqual(r3(x, y, z, u, w, v), self.ref1(x, y, z, u, w, v) / self.ref2(x, y, z, u, w, v), delta=abs(r3(x, y, z, u, w, v)) * 1e-12, msg="Function6D divide function (f1() / p2()) did not match reference function value.") + 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) @@ -306,10 +306,10 @@ def test_mod_function6d(self): r1 = self.f1 % self.f2 r2 = self.ref1 % self.f2 r3 = self.f1 % self.ref2 - for (x, y, z, u, w, v) in itertools.product(testvals, repeat=6): - self.assertAlmostEqual(r1(x, y, z, u, w, v), math.fmod(self.ref1(x, y, z, u, w, v), self.ref2(x, y, z, u, w, v)), delta=abs(r1(x, y, z, u, w, v)) * 1e-12, msg="Function6D modulo function (f1() % f2()) did not match reference function value.") - self.assertAlmostEqual(r2(x, y, z, u, w, v), math.fmod(self.ref1(x, y, z, u, w, v), self.ref2(x, y, z, u, w, v)), delta=abs(r2(x, y, z, u, w, v)) * 1e-12, msg="Function6D modulo function (p1() % f2()) did not match reference function value.") - self.assertAlmostEqual(r3(x, y, z, u, w, v), math.fmod(self.ref1(x, y, z, u, w, v), self.ref2(x, y, z, u, w, v)), delta=abs(r3(x, y, z, u, w, v)) * 1e-12, msg="Function6D modulo function (f1() % p2()) did not match reference function value.") + 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) @@ -319,21 +319,21 @@ def test_pow_function6d_function6d(self): r1 = self.f1 ** self.f2 r2 = self.ref1 ** self.f2 r3 = self.f1 ** self.ref2 - for (x, y, z, u, w, v) in itertools.product(testvals, repeat=6): - if self.ref1(x, y, z, u, w, v) < 0 and not float(self.ref2(x, y, z, u, w, v)).is_integer(): + 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, w, v) + 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, w, v) + 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, w, v) + r3(x, y, z, u, v, w) else: - self.assertAlmostEqual(r1(x, y, z, u, w, v), self.ref1(x, y, z, u, w, v) ** self.ref2(x, y, z, u, w, v), 15, "Function6D power function (f1() ** f2()) did not match reference function value.") - self.assertAlmostEqual(r2(x, y, z, u, w, v), self.ref1(x, y, z, u, w, v) ** self.ref2(x, y, z, u, w, v), 15, "Function6D power function (p1() ** f2()) did not match reference function value.") - self.assertAlmostEqual(r3(x, y, z, u, w, v), self.ref1(x, y, z, u, w, v) ** self.ref2(x, y, z, u, w, v), 15, "Function6D power function (f1() ** p2()) did not match reference function value.") + 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, w, v: 0) ** self.f1 + 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): @@ -346,205 +346,205 @@ def test_pow_3_arguments(self): 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, w, v) in itertools.product(testvals, repeat=6): - self.assertEqual(r1(x, y, z, u, w, v), math.fmod(self.ref1(x, y, z, u, w, v) ** 5, 3), "Function6D 3 argument pow(f1(), A, B) did not match reference value.") - self.assertEqual(r2(x, y, z, u, w, v), math.fmod(5 ** self.ref1(x, y, z, u, w, v), 3), "Function6D 3 argument pow(A, f1(), B) did not match reference value.") - self.assertEqual(r3(x, y, z, u, w, v), math.fmod(5 ** self.ref1(x, y, z, u, w, v), self.ref2(x, y, z, u, w, v)), "Function6D 3 argument pow(A, f1(), f2()) did not match reference value.") - self.assertEqual(r4(x, y, z, u, w, v), math.fmod(self.ref2(x, y, z, u, w, v) ** self.ref1(x, y, z, u, w, v), self.ref2(x, y, z, u, w, v)), "Function6D 3 argument pow(f2(), f1(), f2()) did not match reference value.") - self.assertEqual(r5(x, y, z, u, w, v), math.fmod(self.ref2(x, y, z, u, w, v) ** self.ref1(x, y, z, u, w, v), self.ref2(x, y, z, u, w, v)), "Function6D 3 argument pow(f2(), p1(), p2()) did not match reference value.") - self.assertEqual(r6(x, y, z, u, w, v), math.fmod(self.ref2(x, y, z, u, w, v) ** self.ref1(x, y, z, u, w, v), self.ref2(x, y, z, u, w, v)), "Function6D 3 argument pow(p2(), f1(), f2()) did not match reference value.") + 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, w, v) in itertools.product(testvals, repeat=6): - self.assertEqual(abs(self.f1)(x, y, z, u, w, v), abs(self.ref1(x, y, z, u, w, v)), + 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, w, v) in itertools.product(testvals, repeat=6): + for (x, y, z, u, v, w) in itertools.product(testvals, repeat=6): ref_value = self.ref1 - higher_value = lambda x, y, z, u, w, v: self.ref1(x, y, z, u, w, v) + abs(self.ref1(x, y, z, u, w, v)) + 1 - lower_value = lambda x, y, z, u, w, v: self.ref1(x, y, z, u, w, v) - abs(self.ref1(x, y, z, u, w, v)) - 1 + 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, w, v), 1.0, + (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, w, v), 0.0, + (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, w, v), 1.0, + (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, w, v), 0.0, + (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, w, v), 1.0, + (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, w, v), 0.0, + (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, w, v), 1.0, + (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, w, v), 0.0, + (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, w, v), 1.0, + (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, w, v), 1.0, + (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, w, v), 0.0, + (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, w, v), 1.0, + (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, w, v), 1.0, + (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, w, v), 0.0, + (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, w, v) in itertools.product(testvals, repeat=6): + for (x, y, z, u, v, w) in itertools.product(testvals, repeat=6): ref_value = self.ref1 - higher_value = lambda x, y, z, u, w, v: self.ref1(x, y, z, u, w, v) + abs(self.ref1(x, y, z, u, w, v)) + 1 - lower_value = lambda x, y, z, u, w, v: self.ref1(x, y, z, u, w, v) - abs(self.ref1(x, y, z, u, w, v)) - 1 + 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, w, v), 1.0, + (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, w, v), 0.0, + (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, w, v), 1.0, + (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, w, v), 0.0, + (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, w, v), 1.0, + (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, w, v), 0.0, + (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, w, v), 1.0, + (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, w, v), 0.0, + (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, w, v), 1.0, + (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, w, v), 1.0, + (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, w, v), 0.0, + (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, w, v), 1.0, + (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, w, v), 1.0, + (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, w, v), 0.0, + (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, w, v) in itertools.product(testvals, repeat=6): + 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, w, v), 1.0, + (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, w, v), 0.0, + (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, w, v), 1.0, + (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, w, v), 0.0, + (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, w, v), 1.0, + (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, w, v), 0.0, + (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, w, v), 1.0, + (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, w, v), 0.0, + (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, w, v), 1.0, + (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, w, v), 1.0, + (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, w, v), 0.0, + (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, w, v), 1.0, + (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, w, v), 1.0, + (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, w, v), 0.0, + (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 index 257dd0b7..4593587c 100644 --- a/cherab/core/math/function/float/function6d/tests/test_cmath.py +++ b/cherab/core/math/function/float/function6d/tests/test_cmath.py @@ -32,36 +32,36 @@ class TestCmath6D(unittest.TestCase): def setUp(self): - self.f1 = PythonFunction6D(lambda x, y, z, u, w, v: x / 10 + y + z + u/2 + w/3 + v/4) - self.f2 = PythonFunction6D(lambda x, y, z, u, w, v: x * x + y * y - z * z + u * u + w * w - v * v) + 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, w, v) in itertools.product(testvals, repeat=6): + 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, w, v)) - self.assertEqual(function(x, y, z, u, w, v), expected, "Exp6D call did not match reference value.") + 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, w, v) in itertools.product(testvals, repeat=6): + 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, w, v)) - self.assertEqual(function(x, y, z, u, w, v), expected, "Sin6D call did not match reference value.") + 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, w, v) in itertools.product(testvals, repeat=6): + 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, w, v)) - self.assertEqual(function(x, y, z, u, w, v), expected, "Cos6D call did not match reference value.") + 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, w, v) in itertools.product(testvals, repeat=6): + 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, w, v)) - self.assertEqual(function(x, y, z, u, w, v), expected, "Tan6D call did not match reference value.") + 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] @@ -85,31 +85,31 @@ def test_acos(self): def test_atan(self): testvals = [-10.0, -7, -0.001, 0.0, 0.00003, 10, 23.4] - for (x, y, z, u, w, v) in itertools.product(testvals, repeat=6): + 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, w, v)) - self.assertEqual(function(x, y, z, u, w, v), expected, "Atan6D call did not match reference value.") + 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, w, v) in itertools.product(testvals, repeat=6): + 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, w, v), self.f2(x, y, z, u, w, v)) - self.assertEqual(function(x, y, z, u, w, v), expected, "Atan4Q6D call did not match reference value.") + 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, w, v) in itertools.product(testvals, repeat=6): - expected = math.erf(self.f1(x, y, z, u, w, v)) - self.assertAlmostEqual(function(x, y, z, u, w, v), expected, 10, "Erf6D call did not match reference value.") + 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, w, v) in itertools.product(testvals, repeat=6): - expected = math.sqrt(self.f1(x, y, z, u, w, v)) - self.assertEqual(function(x, y, z, u, w, v), expected, "Sqrt6D call did not match reference value.") + 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 From 451606288cdf4f4f2e440d01a020481dd7e63b61 Mon Sep 17 00:00:00 2001 From: Matej Tomes Date: Mon, 3 Aug 2026 18:07:28 +0200 Subject: [PATCH 75/91] Feature/generic dist (#494) * Add generic distribution function * Add tests to GenericDistribution, ZeroDistribution; move Maxwellian tests from test_maxwellian.py to test_distribution.py * Add ZeroDistribution, GenericDistribution to docs * Add demo for the generic distribution --------- Co-authored-by: Jack Lovell --- CHANGELOG.md | 1 + cherab/core/distribution.pxd | 8 + cherab/core/distribution.pyx | 103 +++++ cherab/core/tests/test_distribution.py | 368 ++++++++++++++++++ cherab/core/tests/test_maxwellian.py | 109 ------ .../generic_distribution.py | 298 ++++++++++++++ docs/source/plasmas/core_plasma_classes.rst | 10 + 7 files changed, 788 insertions(+), 109 deletions(-) create mode 100644 cherab/core/tests/test_distribution.py delete mode 100644 cherab/core/tests/test_maxwellian.py create mode 100644 demos/particle_distribution/generic_distribution.py diff --git a/CHANGELOG.md b/CHANGELOG.md index 04ef3fc3..8026d323 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -12,6 +12,7 @@ Bug fixes: * Fix the import statement for `netcdf_file` in `calcam.py` for compatibility with the upcoming `scipy` v2.0.0. (#510) New: +* 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) 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/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_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/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/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: + From 150a8ec4f21d946ce2b2d611958dc24cb6659bda Mon Sep 17 00:00:00 2001 From: Matej Tomes Date: Tue, 18 Aug 2026 10:28:46 +0200 Subject: [PATCH 76/91] Add GaussianQuadrature2D (#476) * Add 2D Gauss-Legendre quadrature integrator * Add tests for GaussianQuadrature2D * Add Integrator2D documentation * Rename GaussianQuadrature to GaussianQuadrature1D To keep consistent naming with function frameword and quadrature2D integrator. GaussianQuadrature class marked as deprecated. * Add GaussianQuadrature rename to changelog * Replace deprecated GausianQuadrature usage with GausianQuadrature1D --- CHANGELOG.md | 4 + cherab/core/math/integrators/__init__.pxd | 4 +- cherab/core/math/integrators/__init__.py | 4 +- .../core/math/integrators/integrators1d.pxd | 6 +- .../core/math/integrators/integrators1d.pyx | 44 +- .../core/math/integrators/integrators2d.pxd | 15 + .../core/math/integrators/integrators2d.pyx | 391 +++++++++++++++++- cherab/core/math/tests/test_integrators.py | 145 ++++++- cherab/core/model/lineshape/stark.pyx | 6 +- cherab/core/model/plasma/bremsstrahlung.pyx | 6 +- cherab/core/tests/test_bremsstrahlung.py | 4 +- cherab/core/tests/test_lineshapes.py | 4 +- docs/source/math/integrators.rst | 11 + docs/source/math/math.rst | 1 + 14 files changed, 614 insertions(+), 31 deletions(-) create mode 100644 docs/source/math/integrators.rst diff --git a/CHANGELOG.md b/CHANGELOG.md index 8026d323..bd3cf741 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -5,6 +5,7 @@ Release 1.6.0 (TBD) ------------------- 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) @@ -12,6 +13,9 @@ Bug fixes: * Fix the import statement for `netcdf_file` in `calcam.py` for compatibility with the upcoming `scipy` v2.0.0. (#510) New: +* 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) diff --git a/cherab/core/math/integrators/__init__.pxd b/cherab/core/math/integrators/__init__.pxd index c4eff464..a0bf9d43 100644 --- a/cherab/core/math/integrators/__init__.pxd +++ b/cherab/core/math/integrators/__init__.pxd @@ -16,6 +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.integrators2d cimport Integrator2D +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 b1fa4516..200c3d82 100644 --- a/cherab/core/math/integrators/__init__.py +++ b/cherab/core/math/integrators/__init__.py @@ -16,5 +16,5 @@ # See the Licence for the specific language governing permissions and limitations # under the Licence. -from .integrators1d import Integrator1D, GaussianQuadrature -from .integrators2d import Integrator2D +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 6b52157c..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 @@ -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 index 62ba8667..34238397 100644 --- a/cherab/core/math/integrators/integrators2d.pxd +++ b/cherab/core/math/integrators/integrators2d.pxd @@ -25,3 +25,18 @@ cdef class Integrator2D: 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 index 993decc2..080c4c9c 100644 --- a/cherab/core/math/integrators/integrators2d.pyx +++ b/cherab/core/math/integrators/integrators2d.pyx @@ -18,7 +18,13 @@ # See the Licence for the specific language governing permissions and limitations # under the Licence. -from raysect.core.math.function.float cimport Function1D, autowrap_function2d +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: @@ -59,3 +65,386 @@ cdef class Integrator2D: """ 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/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/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/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_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/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 From fcdeaf72ccb126cbb10348dbddfdf2907e60e57d Mon Sep 17 00:00:00 2001 From: Jack Lovell Date: Thu, 20 Aug 2026 10:10:50 +0100 Subject: [PATCH 77/91] Improve regularisation documentation and add ADMT demo (#427) * Add a Regularisation section to the tomography documentation, describing the functionality inside the admt_utils module. * Fixup the docstrings in admt_utils. * Add calculation of "skewed" second derivative operators, which operate along the diagonals of the grid. * Add a bolometer diagonstic system to Generomak. * Add a new example using the Generomak bolometers and the regularisation operators to perform isotropic and ADMT inversions. * Improve generate_derivative_operators function: * More thorough commenting of the function to make it easier to follow what is going on and why. * Exploit the sparsity of the derivative opreators but using a dictionary-of-keys representation instead of dense numpy arrays. A dictionary is faster to insert single elements into inside the loop than a numpy array too. * Remove the `np.isnan` checks which are slow for single elements, replace with a cheap check for `None`. * Support returning the operators as Scipy sparse arrays for use downstream. Default to returning as Numpy arrays for backwards compatibility. * Support generating sparse ADMT operators. * Add a regularised NNLS inversion using sparse matrices asn alternative to `invert_regularised_nnls` using sparse weight and penalty matrices. The numerical algorithm used is slightly different to `scipy.optimize.nnls`, as it uses the TRF variant of a bounded `scipy.optimize.lsq_linear` with the lower bound set to 0 to enforce positivity. Since the results differ due to floating point precision, the sparse variant is implemented as a separate function to maintain backwards compatibility. * Enable generating derivative operators with only 1 mapping: The 1D-to-2D and 2D-to-1D maps are the inverse of one another, so one can be computed from the other. Depending on how the end user forms their inversion grid it may be simpler to calculate the 1D-to-2D or the 2D-to-1D, so allow the caller to pass either and compute any missing mapping. If the caller already has both mappings (machine packages like cherab-jet, cherab-mastu and cherab-aug already generate both) then accept both too. * Add tests for auto-computing missing mappings. * Tests uncovered a bug converting the Dyy sparse operator to dense when `sparse=False` was passed: fix that and complete test coverage for `admt_utils`. * Update changelog. --- CHANGELOG.md | 7 + cherab/generomak/diagnostics/__init__.py | 1 + cherab/generomak/diagnostics/bolometers.py | 247 +++++++++++ cherab/tools/inversions/__init__.py | 2 +- cherab/tools/inversions/admt_utils.py | 201 ++++++--- cherab/tools/inversions/nnls.py | 63 ++- cherab/tools/tests/test_admt.py | 53 +++ .../bolometry/admt_tomographic_inversion.py | 385 ++++++++++++++++++ docs/source/tools/tomography.rst | 35 ++ 9 files changed, 944 insertions(+), 50 deletions(-) create mode 100644 cherab/generomak/diagnostics/__init__.py create mode 100644 cherab/generomak/diagnostics/bolometers.py create mode 100644 demos/observers/bolometry/admt_tomographic_inversion.py diff --git a/CHANGELOG.md b/CHANGELOG.md index bd3cf741..513c1bc1 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -8,6 +8,9 @@ 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) @@ -23,6 +26,10 @@ New: * 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) * 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) ------------------- 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/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/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/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/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 From 31aca80267265aab853df4e5f215172043719d9b Mon Sep 17 00:00:00 2001 From: munechika-koyo Date: Mon, 24 Aug 2026 11:44:15 +0200 Subject: [PATCH 78/91] Refactor pixi.toml: reorganize build configuration and clean up commands --- pixi.toml | 101 +++++++++++++++++++++++++----------------------------- 1 file changed, 46 insertions(+), 55 deletions(-) diff --git a/pixi.toml b/pixi.toml index 1b84805b..187b4ba2 100644 --- a/pixi.toml +++ b/pixi.toml @@ -3,6 +3,9 @@ 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 === # ------------------------------- @@ -10,16 +13,13 @@ preview = ["pixi-build"] name = "cherab" version = "dynamic" -[package.build] -backend = { name = "pixi-build-python", version = "*" } +[package.build.backend] +name = "pixi-build-python" +version = "*" [package.build.config] -noarch = false compilers = ["c"] -[workspace.build-variants] -python = ["3.9.*", "3.10.*", "3.11.*", "3.12.*", "3.13.*", "3.14.*"] - [package.host-dependencies] python = "*" setuptools = "*" @@ -30,6 +30,8 @@ raysect = "0.9.*" [package.run-dependencies] scipy = "*" matplotlib-base = "*" + +[package.extra-dependencies.opencl] pyopencl = "*" pocl = "*" @@ -39,49 +41,40 @@ pocl = "*" [dependencies] ipython = "*" +[dev] +cherab = { path = "." } + [tasks] -clean = { cmd = [ - "find", - "cherab/", - "-type", - "f", - "\\(", - "-name", - "'*.c'", - "-o", - "-name", - "'*.so'", - "-o", - "-name", - "'*.dylib'", - "-o", - "-name", - "'*.html'", - "\\)", - "-delete", -], description = "🔥 Remove in-place build artifacts and temporary files (*.c, *.so, *.dylib)" } +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", - "8000", - "--directory", - "build/html", -], cwd = "docs", description = "🚀 Start a local server for the docs" } +doc-clean = { + cmd = "rm -rf build", + cwd = "docs", + description = "🔥 Clean the docs build directory", +} # === Testing feature === [feature.test.dependencies] cherab = { path = "." } [feature.test.tasks] -test = { cmd = "python -m unittest discover cherab -v", description = "🧪 Run the tests" } +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] @@ -91,15 +84,20 @@ sphinx_rtd_theme = "<1" sphinx-tabs = "*" [feature.docs.tasks] -doc-build = { cmd = [ +doc-build = { + cmd = [ "sphinx-build", "-b", "{{ target }}", "source", "build/{{ target }}", -], cwd = "docs", args = [ + ], + cwd = "docs", + args = [ { arg = "target", default = "html" }, -], description = "📝 Build the docs" } + ], + description = "📝 Build the docs" +} # === Linting feature === [feature.lint.dependencies] @@ -132,22 +130,15 @@ lint = { cmd = "lefthook run pre-commit --all-files --force", description = " # === 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" } +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 } From 36763a053e5eb6d9e230a285f3ab190a62b0982f Mon Sep 17 00:00:00 2001 From: munechika-koyo Date: Mon, 24 Aug 2026 13:11:47 +0200 Subject: [PATCH 79/91] Add doc-serve task to start a local server for documentation --- pixi.toml | 15 +++++++++++++++ 1 file changed, 15 insertions(+) diff --git a/pixi.toml b/pixi.toml index 187b4ba2..8ae3b640 100644 --- a/pixi.toml +++ b/pixi.toml @@ -56,6 +56,21 @@ doc-clean = { 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] From e68716147cfe4e5ee8d143a176ec89a35115640e Mon Sep 17 00:00:00 2001 From: munechika-koyo Date: Mon, 24 Aug 2026 13:28:22 +0200 Subject: [PATCH 80/91] Replace taplo with tombi for TOML formatting in lint tasks --- pixi.toml | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/pixi.toml b/pixi.toml index 8ae3b640..0c90842e 100644 --- a/pixi.toml +++ b/pixi.toml @@ -125,7 +125,7 @@ shellcheck = "*" validate-pyproject = "*" cython-lint = "*" blacken-docs = "*" -taplo = "*" +tombi = "*" [feature.lint.tasks] lefthook = { cmd = "lefthook", description = "🔗 Run lefthook" } @@ -133,9 +133,9 @@ 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" } -taplo = { cmd = "taplo fmt", description = "Format toml files with taplo" } 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" } From b8d5f284a101322289079382ea33d7de8843ca6f Mon Sep 17 00:00:00 2001 From: munechika-koyo Date: Mon, 24 Aug 2026 13:38:28 +0200 Subject: [PATCH 81/91] Add Pixi developer guide with environment and task documentation --- dev/pixi.md | 180 ++++++++++++++++++++++++++++++++++++++++++++++++++++ 1 file changed, 180 insertions(+) create mode 100644 dev/pixi.md diff --git a/dev/pixi.md b/dev/pixi.md new file mode 100644 index 00000000..0856e819 --- /dev/null +++ b/dev/pixi.md @@ -0,0 +1,180 @@ +# 🧰 Pixi developer guide + +This document describes the development environments and tasks configured in +[`pixi.toml`](../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`](../pixi.toml#L145-L159) 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 +``` From e69f53cfd50aecd7abed2f3c58c118116be556f0 Mon Sep 17 00:00:00 2001 From: munechika-koyo Date: Mon, 24 Aug 2026 14:45:35 +0200 Subject: [PATCH 82/91] Fix links in Pixi developer guide and remove top header's emoji --- dev/pixi.md | 13 +++++++------ 1 file changed, 7 insertions(+), 6 deletions(-) diff --git a/dev/pixi.md b/dev/pixi.md index 0856e819..2786e641 100644 --- a/dev/pixi.md +++ b/dev/pixi.md @@ -1,9 +1,10 @@ -# 🧰 Pixi developer guide +# Pixi developer guide This document describes the development environments and tasks configured in -[`pixi.toml`](../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.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 @@ -40,8 +41,8 @@ explicitly. | `lint` | Run formatting and static-analysis tools without installing Cherab | See the [`pyoldest` and `pylatest` features and environment definitions in -`pixi.toml`](../pixi.toml#L145-L159) for the Python versions used by each -environment. +`pixi.toml`](https://github.com/cherab/core/blob/development/pixi.toml#L146-L160) +for the Python versions used by each environment. ## 🛠️ Basic development tasks From 3890c0ed9c67627c591b6cf06fa53cb4940427c9 Mon Sep 17 00:00:00 2001 From: munechika-koyo Date: Mon, 24 Aug 2026 14:46:21 +0200 Subject: [PATCH 83/91] Add myst-parser dependency for documentation generation --- pixi.toml | 1 + 1 file changed, 1 insertion(+) diff --git a/pixi.toml b/pixi.toml index 0c90842e..250fc291 100644 --- a/pixi.toml +++ b/pixi.toml @@ -97,6 +97,7 @@ cherab = { path = "." } sphinx = "*" sphinx_rtd_theme = "<1" sphinx-tabs = "*" +myst-parser = "*" [feature.docs.tasks] doc-build = { From c596630a10496b2cb79d42aea6e83c441e285745 Mon Sep 17 00:00:00 2001 From: munechika-koyo Date: Mon, 24 Aug 2026 14:47:09 +0200 Subject: [PATCH 84/91] Add myst_parser extension and update documentation structure for development --- docs/source/conf.py | 1 + docs/source/development/pixi.rst | 2 ++ docs/source/index.rst | 9 ++++++++- 3 files changed, 11 insertions(+), 1 deletion(-) create mode 100644 docs/source/development/pixi.rst diff --git a/docs/source/conf.py b/docs/source/conf.py index b2f22c70..c268831e 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. 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` - From 05ca189198645fe85e5ce1395d1b7837e574a674 Mon Sep 17 00:00:00 2001 From: munechika-koyo Date: Mon, 24 Aug 2026 14:57:59 +0200 Subject: [PATCH 85/91] Update link to Python version features in Pixi developer guide --- dev/pixi.md | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/dev/pixi.md b/dev/pixi.md index 2786e641..6b19fa1a 100644 --- a/dev/pixi.md +++ b/dev/pixi.md @@ -41,7 +41,7 @@ explicitly. | `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#L146-L160) +`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 From 5f63cec0be3e4df7de9f218a15cfbdd1bd8574bd Mon Sep 17 00:00:00 2001 From: munechika-koyo Date: Mon, 24 Aug 2026 14:59:20 +0200 Subject: [PATCH 86/91] Add pixi feature into changelog --- CHANGELOG.md | 1 + 1 file changed, 1 insertion(+) diff --git a/CHANGELOG.md b/CHANGELOG.md index 513c1bc1..bc9d8b42 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -16,6 +16,7 @@ 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) From bff0e59b663b38f5b97ecf26e9dc54cd93f87ab7 Mon Sep 17 00:00:00 2001 From: Koyo MUNECHIKA <51052381+munechika-koyo@users.noreply.github.com> Date: Mon, 24 Aug 2026 16:27:55 +0200 Subject: [PATCH 87/91] Fix import order and improve code consistency in utility.py (#514) --- cherab/core/math/caching/utility.py | 12 +++++------- 1 file changed, 5 insertions(+), 7 deletions(-) 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() From 0c5fbed5d8988de07e437215c844a8f5677845a7 Mon Sep 17 00:00:00 2001 From: munechika-koyo Date: Tue, 25 Aug 2026 08:03:21 +0200 Subject: [PATCH 88/91] Add TODO comment about sphinx_rtd_theme version constraint --- pixi.toml | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/pixi.toml b/pixi.toml index 250fc291..7f9dd08d 100644 --- a/pixi.toml +++ b/pixi.toml @@ -95,7 +95,7 @@ test-opencl = { [feature.docs.dependencies] cherab = { path = "." } sphinx = "*" -sphinx_rtd_theme = "<1" +sphinx_rtd_theme = "<1" # TODO: change to >=1.0 when our docs layout is compatible with the new theme sphinx-tabs = "*" myst-parser = "*" From 9aee81b4a3328b51cbfa2624e6be1e83ec4061b7 Mon Sep 17 00:00:00 2001 From: munechika-koyo Date: Tue, 25 Aug 2026 08:05:41 +0200 Subject: [PATCH 89/91] Format python version constraints in pixi.toml for better readability --- pixi.toml | 9 ++++++++- 1 file changed, 8 insertions(+), 1 deletion(-) diff --git a/pixi.toml b/pixi.toml index 7f9dd08d..64d3c551 100644 --- a/pixi.toml +++ b/pixi.toml @@ -4,7 +4,14 @@ 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.*"] +python = [ + "3.9.*", + "3.10.*", + "3.11.*", + "3.12.*", + "3.13.*", + "3.14.*", +] # ------------------------------- # === Packaging Configuration === From 6694e48b5121fdeaa10d349b7fdec465f9b90bcc Mon Sep 17 00:00:00 2001 From: Jack Lovell Date: Tue, 8 Sep 2026 16:40:24 +0200 Subject: [PATCH 90/91] Add an action to build distributions on tagged releases (#518) * Add an action to build distributions on tagged releases Use the CIbuildwheel Github action to build wheels for all currently-supported Python versions, plus an SDist tarball. Manylinux wheels are produced using the manylinux2014 image where possible (to support JET/UKAEA computers which still run SL7), with manylinux_2_28 for Python 3.14 as that isn't available in manylinux2014. Only Linux wheels are built for now: MacOS (and Windows should Raysect ever support it) builds are deferred. But the strategy matrix for the wheels is written in such a way as to easily add other OS's or architectures later as necessary. To avoid wasting compute cycles, the workflow is configured to only run on published releases. This includes stable and pre-release publications. Automatic uploading to PyPI is not implemented: the distribution files should be downloaded from the `gather_artifacts` job's artifacts and checked locally before uploading to PyPI. This is to reduce the risk of uploading broken distributions: if the process proves robust and reliable then automatic uploading to PyPI can be revisited. --- .github/workflows/build.yml | 115 ++++++++++++++++++++++++++++++++++++ cherab/core/VERSION | 2 +- setup.py | 2 + 3 files changed, 118 insertions(+), 1 deletion(-) create mode 100644 .github/workflows/build.yml 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/cherab/core/VERSION b/cherab/core/VERSION index 8a3469e1..18a4f925 100644 --- a/cherab/core/VERSION +++ b/cherab/core/VERSION @@ -1 +1 @@ -1.6.0.dev2 +1.6.0.dev4 diff --git a/setup.py b/setup.py index 952cfbbe..c07eaff8 100644 --- a/setup.py +++ b/setup.py @@ -115,6 +115,8 @@ ), 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>=2.0", "scipy", From 5e2c1111350e7fbb94e3516577025395d5463234 Mon Sep 17 00:00:00 2001 From: Jack Lovell Date: Tue, 8 Sep 2026 10:48:00 -0400 Subject: [PATCH 91/91] Bump version to 1.6.0rc1 --- CHANGELOG.md | 3 +-- cherab/core/VERSION | 2 +- docs/source/conf.py | 6 +++--- 3 files changed, 5 insertions(+), 6 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index bc9d8b42..d1cd0320 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -1,7 +1,7 @@ Project Changelog ================= -Release 1.6.0 (TBD) +Release 1.6.0 (Sep 2026) ------------------- API changes: @@ -25,7 +25,6 @@ New: * 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) -* 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) * 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) diff --git a/cherab/core/VERSION b/cherab/core/VERSION index 18a4f925..40ab7ec2 100644 --- a/cherab/core/VERSION +++ b/cherab/core/VERSION @@ -1 +1 @@ -1.6.0.dev4 +1.6.0rc1 diff --git a/docs/source/conf.py b/docs/source/conf.py index c268831e..54bb264f 100644 --- a/docs/source/conf.py +++ b/docs/source/conf.py @@ -57,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.