s2conv 0.1.0__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
s2conv/__init__.py ADDED
@@ -0,0 +1,35 @@
1
+ r"""
2
+ s2conv: differentiable and accelerated spherical convolutions with JAX.
3
+
4
+ s2conv implements the generalised lifted spherical convolution, in which a spin-:math:`s`
5
+ field is lifted to the rotation group, convolved there with a kernel :math:`\Psi` and
6
+ projected back to the sphere through the adjoint of the lift for a possibly different
7
+ spin :math:`s'`, and the spherical convolutions of the literature that it unifies. The
8
+ harmonic transforms are those of s2fft, which warns on import if JAX is not using 64-bit
9
+ precision.
10
+ """
11
+
12
+ from . import convolutions, kernels, lifting, transforms, utils
13
+ from .convolutions import (
14
+ classical,
15
+ diagonal,
16
+ group_convolution,
17
+ group_to_group,
18
+ group_to_spin,
19
+ lifted,
20
+ lifted_convolution,
21
+ physics,
22
+ sampled,
23
+ section,
24
+ spin_multiplier,
25
+ spin_to_group,
26
+ spin_to_spin,
27
+ )
28
+ from .convolutions.sampled import convolve, convolve_multiplier
29
+ from .lifting import lift, project
30
+ from .transforms import forward, inverse, wigner_forward, wigner_inverse
31
+
32
+ try:
33
+ from ._version import __version__
34
+ except ImportError: # pragma: no cover - only without an installed version file
35
+ __version__ = "unknown"
s2conv/_version.py ADDED
@@ -0,0 +1,24 @@
1
+ # file generated by vcs-versioning
2
+ # don't change, don't track in version control
3
+ from __future__ import annotations
4
+
5
+ __all__ = [
6
+ "__version__",
7
+ "__version_tuple__",
8
+ "version",
9
+ "version_tuple",
10
+ "__commit_id__",
11
+ "commit_id",
12
+ ]
13
+
14
+ version: str
15
+ __version__: str
16
+ __version_tuple__: tuple[int | str, ...]
17
+ version_tuple: tuple[int | str, ...]
18
+ commit_id: str | None
19
+ __commit_id__: str | None
20
+
21
+ __version__ = version = '0.1.0'
22
+ __version_tuple__ = version_tuple = (0, 1, 0)
23
+
24
+ __commit_id__ = commit_id = 'g03ddef413'
@@ -0,0 +1,40 @@
1
+ """The generalised lifted convolution and the spherical convolutions it unifies."""
2
+
3
+ from . import classical, diagonal, lifted, physics, sampled, section
4
+ from .classical import (
5
+ axisymmetric_convolution,
6
+ axisymmetric_synthesis,
7
+ directional_convolution,
8
+ directional_synthesis,
9
+ left_convolution,
10
+ spin_changing_convolution,
11
+ )
12
+ from .diagonal import harmonic_multiplication, sifting_convolution
13
+ from .lifted import (
14
+ group_convolution,
15
+ group_to_group,
16
+ group_to_group_weights,
17
+ group_to_spin,
18
+ group_to_spin_weights,
19
+ lifted_convolution,
20
+ spin_multiplier,
21
+ spin_to_group,
22
+ spin_to_group_weights,
23
+ spin_to_spin,
24
+ )
25
+ from .physics import (
26
+ eb_decomposition,
27
+ eb_to_polarisation,
28
+ eth,
29
+ ethbar,
30
+ kaiser_squires,
31
+ kaiser_squires_inverse,
32
+ )
33
+ from .sampled import convolve, convolve_multiplier
34
+ from .section import (
35
+ commutative_anisotropic_convolution,
36
+ section_components,
37
+ section_evaluation,
38
+ standard_correlation,
39
+ standard_correlation_fourier,
40
+ )
@@ -0,0 +1,175 @@
1
+ r"""
2
+ The rotation-equivariant spherical convolutions of the literature (Secs. III and VI).
3
+
4
+ Each is a generalised lifted convolution with a particular kernel (Table II), and is
5
+ written here in its closed harmonic form, with the constants of the article:
6
+
7
+ .. list-table::
8
+ :header-rows: 1
9
+
10
+ * - Convolution
11
+ - Lifted form
12
+ - Harmonic form
13
+ * - axisymmetric (Defs III.1, III.6)
14
+ - :math:`(2\pi)^{-1/2}\mathcal{C}_{s\to0}[\Lambda_s\psi]`
15
+ - :math:`\sqrt{4\pi/(2\ell+1)}\, {}_sf_{\ell m}\, {}_s\psi^*_{\ell0}`
16
+ * - left (Def. III.2)
17
+ - :math:`\mathcal{C}_{0\to0}[\Psi_h]`
18
+ - :math:`2\pi\sqrt{4\pi/(2\ell+1)}\, f_{\ell m} h_{\ell0}`
19
+ * - directional (Defs III.3, III.5)
20
+ - :math:`\mathcal{C}_{s\to G}[\Lambda_s\psi]`
21
+ - :math:`\frac{8\pi^2}{2\ell+1}\, {}_sf_{\ell m}\, {}_s\psi^*_{\ell n}`
22
+ * - synthesis (Def. III.10)
23
+ - :math:`\mathcal{C}_{G\to s}[\widetilde{\Lambda_s\psi}]`
24
+ - :math:`\sum_n W^{\ell}_{mn}\, {}_s\psi_{\ell n}`
25
+ * - spin-changing (Prop. VI.5)
26
+ - :math:`\mathcal{C}_{s\to s'}[\Psi_T]`
27
+ - :math:`c_\ell\, {}_sf_{\ell m}`
28
+
29
+ The group convolution (Definition III.4) is :func:`s2conv.convolutions.lifted.group_convolution`.
30
+ All functions act on harmonic coefficients in the s2fft layout, broadcast over leading
31
+ axes, and are differentiable and JIT compilable.
32
+ """
33
+
34
+ import jax.numpy as jnp
35
+ import numpy as np
36
+
37
+ from s2conv.convolutions.lifted import (
38
+ group_to_spin_weights,
39
+ spin_multiplier,
40
+ spin_to_group_weights,
41
+ )
42
+ from s2conv.utils.indexing import azimuthal_bandlimit, lm_mask
43
+
44
+
45
+ def _axisymmetric_weight(L: int) -> np.ndarray:
46
+ return np.sqrt(4 * np.pi / (2 * np.arange(L) + 1))
47
+
48
+
49
+ def axisymmetric_convolution(flm, psi_lm, spin: int = 0) -> jnp.ndarray:
50
+ r"""
51
+ Axisymmetric convolution :math:`(f \odot \psi)(\omega) = \langle f, \mathcal{R}_\omega \psi \rangle`.
52
+
53
+ For a scalar field (Definition III.1) or a spin-:math:`s` field (Definition III.6) with an
54
+ axisymmetric kernel, the output is a scalar field with
55
+ :math:`(f \odot \psi)_{\ell m} = \sqrt{4\pi/(2\ell+1)}\, {}_sf_{\ell m}\, {}_s\psi^*_{\ell 0}`
56
+ (Eqs. 14, 25). It is :math:`(2\pi)^{-1/2}\mathcal{C}_{s\to0}[\Lambda_s\psi]`, the factor being the
57
+ difference between evaluation on the section and the adjoint projection
58
+ (Proposition VI.3). Only the axisymmetric part :math:`\psi_{\ell0}` of the kernel enters.
59
+
60
+ Args:
61
+ flm: Spin-:math:`s` harmonic coefficients, ``[..., L, 2L - 1]``.
62
+ psi_lm: Spin-:math:`s` harmonic coefficients of the kernel, ``[L, 2L - 1]``.
63
+ spin (int, optional): Spin :math:`s`. Defaults to 0.
64
+
65
+ Returns:
66
+ jnp.ndarray: Scalar harmonic coefficients, ``[..., L, 2L - 1]``.
67
+
68
+ """
69
+ L = flm.shape[-2]
70
+ c = _axisymmetric_weight(L) * jnp.conj(psi_lm[..., :, L - 1])
71
+ return spin_multiplier(flm, c, spin, 0)
72
+
73
+
74
+ def left_convolution(flm, h_lm) -> jnp.ndarray:
75
+ r"""
76
+ Left convolution of Driscoll & Healy, :math:`(f \star_{\mathrm{L}} h)(\omega) = \int f(\rho\eta)\, h(R_\rho^{-1}\omega)\, \mathrm{d}\mu(\rho)`.
77
+
78
+ Its harmonic form is :math:`2\pi\sqrt{4\pi/(2\ell+1)}\, f_{\ell m} h_{\ell 0}` (Definition III.2,
79
+ Eq. 16): only the zonal part of the kernel enters, and the :math:`2\pi` is the volume of
80
+ the fibre. It is :math:`\mathcal{C}_{0\to0}[\Psi_h]` with :math:`\Psi_h(g) = h^*(g^{-1}\eta)`,
81
+ exactly (Proposition VI.4).
82
+
83
+ Args:
84
+ flm: Scalar harmonic coefficients, ``[..., L, 2L - 1]``.
85
+ h_lm: Scalar harmonic coefficients of the kernel, ``[L, 2L - 1]``.
86
+
87
+ Returns:
88
+ jnp.ndarray: Scalar harmonic coefficients, ``[..., L, 2L - 1]``.
89
+
90
+ """
91
+ L = flm.shape[-2]
92
+ c = 2 * np.pi * _axisymmetric_weight(L) * h_lm[..., :, L - 1]
93
+ return spin_multiplier(flm, c, 0, 0)
94
+
95
+
96
+ def directional_convolution(flm, psi_lm, spin: int = 0, N: int = None) -> jnp.ndarray:
97
+ r"""
98
+ Directional convolution :math:`(f \circledast \psi)(\rho) = \langle f, \mathcal{R}_\rho\psi\rangle`, a function on the rotation group.
99
+
100
+ For scalar (Definition III.3) or spin-:math:`s` fields (Definition III.5) its Wigner
101
+ coefficients are :math:`\frac{8\pi^2}{2\ell+1}\, {}_sf_{\ell m}\, {}_s\psi^*_{\ell n}`
102
+ (Eqs. 18, 23). It is :math:`\mathcal{C}_{s\to G}[\Lambda_s\psi]`, exactly, with no residual
103
+ constant (Proposition VI.1). For a steerable kernel, azimuthally band-limited at
104
+ :math:`N`, the output is the tuple of :math:`2N - 1` spin-changing convolutions of
105
+ Proposition VI.6 (see :func:`s2conv.lifting.spin_components`).
106
+
107
+ Args:
108
+ flm: Spin-:math:`s` harmonic coefficients, ``[..., L, 2L - 1]``.
109
+ psi_lm: Spin-:math:`s` harmonic coefficients of the kernel, ``[L, 2L - 1]``.
110
+ spin (int, optional): Spin :math:`s`. Defaults to 0.
111
+ N (int, optional): Azimuthal band-limit of the output, retaining the kernel's
112
+ orders :math:`|n| < N`. Defaults to :math:`L`, which is exact for any kernel.
113
+
114
+ Returns:
115
+ jnp.ndarray: Wigner coefficients, ``[..., 2N - 1, L, 2L - 1]``.
116
+
117
+ """
118
+ L = flm.shape[-2]
119
+ N = L if N is None else N
120
+ weight = 8 * np.pi**2 / (2 * np.arange(L) + 1)
121
+ u = (
122
+ jnp.conj(
123
+ psi_lm[..., :, L - N : L + N - 1] * lm_mask(L, spin)[:, L - N : L + N - 1]
124
+ )
125
+ * weight[:, None]
126
+ )
127
+ return spin_to_group_weights(flm, jnp.swapaxes(u, -1, -2), spin)
128
+
129
+
130
+ def directional_synthesis(flmn, psi_lm, spin: int = 0) -> jnp.ndarray:
131
+ r"""
132
+ Haar synthesis :math:`(W \circledast^\dagger \psi)(\omega) = \int W(\rho)\, (\mathcal{R}_\rho\psi)(\omega)\, \mathrm{d}\mu(\rho)`.
133
+
134
+ Its harmonic form is :math:`\sum_n W^{\ell}_{mn}\, {}_s\psi_{\ell n}` (Definition III.10,
135
+ Eq. 31). It is :math:`\mathcal{C}_{G\to s}[\widetilde{\Lambda_s\psi}]`, the adjoint of the
136
+ directional convolution with the same kernel (Proposition VI.2).
137
+
138
+ Args:
139
+ flmn: Wigner coefficients, ``[..., 2N - 1, L, 2L - 1]``.
140
+ psi_lm: Spin-:math:`s` harmonic coefficients of the kernel, ``[L, 2L - 1]``.
141
+ spin (int, optional): Spin :math:`s` of the output. Defaults to 0.
142
+
143
+ Returns:
144
+ jnp.ndarray: Spin-:math:`s` harmonic coefficients, ``[..., L, 2L - 1]``.
145
+
146
+ """
147
+ N, L = azimuthal_bandlimit(flmn), flmn.shape[-2]
148
+ v = (psi_lm * lm_mask(L, spin))[..., :, L - N : L + N - 1]
149
+ return group_to_spin_weights(flmn, jnp.swapaxes(v, -1, -2), spin)
150
+
151
+
152
+ def axisymmetric_synthesis(flm, psi_lm, spin: int = 0) -> jnp.ndarray:
153
+ r"""
154
+ Axisymmetric synthesis :math:`\int W(\omega')\, (\mathcal{R}_{\omega'}\psi)(\omega)\, \mathrm{d}\mu(\omega')`.
155
+
156
+ The adjoint of :func:`axisymmetric_convolution`, mapping a scalar field to a spin-:math:`s`
157
+ field, :math:`\sqrt{4\pi/(2\ell+1)}\, W_{\ell m}\, {}_s\psi_{\ell 0}`: the first term of the
158
+ spin wavelet synthesis (Eq. 26).
159
+
160
+ Args:
161
+ flm: Scalar harmonic coefficients, ``[..., L, 2L - 1]``.
162
+ psi_lm: Spin-:math:`s` harmonic coefficients of the kernel, ``[L, 2L - 1]``.
163
+ spin (int, optional): Spin :math:`s` of the output. Defaults to 0.
164
+
165
+ Returns:
166
+ jnp.ndarray: Spin-:math:`s` harmonic coefficients, ``[..., L, 2L - 1]``.
167
+
168
+ """
169
+ L = flm.shape[-2]
170
+ c = _axisymmetric_weight(L) * psi_lm[..., :, L - 1]
171
+ return spin_multiplier(flm, c, 0, spin)
172
+
173
+
174
+ spin_changing_convolution = spin_multiplier
175
+ r"""Spin-changing convolution with multiplier :math:`c_\ell` (Proposition VI.5); see :func:`s2conv.convolutions.lifted.spin_multiplier`."""
@@ -0,0 +1,41 @@
1
+ r"""
2
+ Operators diagonal in the spherical harmonic basis (Sec. VII B).
3
+
4
+ Harmonic multiplication of Kennedy et al. and the sifting convolution of Roddy & McEwen
5
+ multiply each coefficient by :math:`g_{\ell m}` or :math:`g^*_{\ell m}` (Definition III.9).
6
+ They commute with every polar rotation, and are equivariant under all rotations exactly
7
+ when :math:`g_{\ell m}` does not depend on :math:`m`, in which case they are the lifted
8
+ convolution :math:`\mathcal{C}_{0\to0}[\Psi_T]` with :math:`c_\ell = g_{\ell0}` (Proposition VII.2).
9
+ """
10
+
11
+ import jax.numpy as jnp
12
+
13
+
14
+ def harmonic_multiplication(flm, glm) -> jnp.ndarray:
15
+ r"""
16
+ Harmonic multiplication :math:`(g \odot_{\mathrm{K}} f)_{\ell m} = g_{\ell m} f_{\ell m}` (Definition III.9).
17
+
18
+ Args:
19
+ flm: Scalar harmonic coefficients, ``[..., L, 2L - 1]``.
20
+ glm: Scalar harmonic coefficients of the filter, ``[..., L, 2L - 1]``.
21
+
22
+ Returns:
23
+ jnp.ndarray: Scalar harmonic coefficients, ``[..., L, 2L - 1]``.
24
+
25
+ """
26
+ return glm * flm
27
+
28
+
29
+ def sifting_convolution(flm, glm) -> jnp.ndarray:
30
+ r"""
31
+ Sifting convolution :math:`(f \circledcirc g)_{\ell m} = f_{\ell m} g^*_{\ell m}` (Definition III.9).
32
+
33
+ Args:
34
+ flm: Scalar harmonic coefficients, ``[..., L, 2L - 1]``.
35
+ glm: Scalar harmonic coefficients of the filter, ``[..., L, 2L - 1]``.
36
+
37
+ Returns:
38
+ jnp.ndarray: Scalar harmonic coefficients, ``[..., L, 2L - 1]``.
39
+
40
+ """
41
+ return flm * jnp.conj(glm)
@@ -0,0 +1,221 @@
1
+ r"""
2
+ The generalised lifted spherical convolution (Sec. V).
3
+
4
+ For a kernel :math:`\Psi` on the rotation group and spins :math:`s, s'`, the four operators
5
+ of Definition V.1 are
6
+
7
+ .. math::
8
+
9
+ \mathcal{C}_{s\to s'}[\Psi] = \Lambda_{s'}^\dagger (\Psi \star \Lambda_s\,\cdot\,), \quad
10
+ \mathcal{C}_{s\to G}[\Psi] = \Psi \star \Lambda_s\,\cdot\,, \quad
11
+ \mathcal{C}_{G\to s'}[\Psi] = \Lambda_{s'}^\dagger (\Psi \star\,\cdot\,), \quad
12
+ \mathcal{C}_{G\to G}[\Psi] = \Psi \star\,\cdot\,,
13
+
14
+ with :math:`\star` the group convolution in correlation form,
15
+ :math:`(\Psi \star F)(\rho) = \int \Psi^*(\rho^{-1}\rho') F(\rho')\, \mathrm{d}\mu(\rho')`
16
+ (Definition III.4). In harmonic space they read a single entry, column, row or the
17
+ whole block of the kernel (Theorem V.2, Eqs. 42–45):
18
+
19
+ .. math::
20
+
21
+ (\mathcal{C}_{s\to G}[\Psi]\, {}_sf)^{\ell}_{mn} &= (-1)^s \sqrt{\tfrac{8\pi^2}{2\ell+1}}\, {}_sf_{\ell m}\, \Psi^{\ell *}_{n,-s}, \\
22
+ (\mathcal{C}_{s\to s'}[\Psi]\, {}_sf)_{\ell m} &= (-1)^{s+s'}\, {}_sf_{\ell m}\, \Psi^{\ell *}_{-s',-s}, \\
23
+ (\mathcal{C}_{G\to s'}[\Psi]\, F)_{\ell m} &= (-1)^{s'} \sqrt{\tfrac{2\ell+1}{8\pi^2}} \sum_{m'} F^{\ell}_{mm'}\, \Psi^{\ell *}_{-s',m'}, \\
24
+ (\mathcal{C}_{G\to G}[\Psi]\, F)^{\ell}_{mn} &= \sum_{m'} F^{\ell}_{mm'}\, \Psi^{\ell *}_{nm'} .
25
+
26
+ Every bounded rotation-equivariant map of the four types is of this form, and acts
27
+ through compact weights (Theorem V.5): a multiplier :math:`c_\ell`, a column
28
+ :math:`u^{\ell}_n`, a row :math:`v^{\ell}_k` or a block :math:`B^{\ell}`. This module applies
29
+ the operators in both parameterisations; :mod:`s2conv.kernels` converts between them.
30
+ All functions act on harmonic coefficients in the s2fft layout, broadcast over leading
31
+ axes, and are differentiable and JIT compilable (with spins and band-limits static).
32
+ """
33
+
34
+ import jax.numpy as jnp
35
+
36
+ from s2conv import kernels
37
+ from s2conv.utils.indexing import azimuthal_bandlimit, lm_mask, wigner_mask
38
+
39
+ # ------------------------------------------------------------------------------ weight form
40
+
41
+
42
+ def spin_multiplier(flm, c, s_in: int, s_out: int) -> jnp.ndarray:
43
+ r"""
44
+ A spin-changing multiplier, :math:`({}_{s'}g)_{\ell m} = c_\ell\, {}_sf_{\ell m}` (Theorem V.5).
45
+
46
+ This is the general bounded rotation-equivariant map from spin :math:`s` to spin
47
+ :math:`s'`; degrees :math:`\ell < \max(|s|, |s'|)` are absent from the target and are
48
+ annihilated.
49
+
50
+ Args:
51
+ flm: Spin-:math:`s` harmonic coefficients, ``[..., L, 2L - 1]``.
52
+ c: Multiplier :math:`c_\ell`, ``[..., L]``.
53
+ s_in (int): Input spin :math:`s`.
54
+ s_out (int): Output spin :math:`s'`.
55
+
56
+ Returns:
57
+ jnp.ndarray: Spin-:math:`s'` harmonic coefficients, ``[..., L, 2L - 1]``.
58
+
59
+ """
60
+ L = flm.shape[-2]
61
+ return jnp.asarray(c)[..., :, None] * flm * lm_mask(L, max(abs(s_in), abs(s_out)))
62
+
63
+
64
+ def spin_to_group_weights(flm, u, s_in: int) -> jnp.ndarray:
65
+ r"""
66
+ An equivariant map from spin :math:`s` to the rotation group, :math:`F^{\ell}_{mn} = u^{\ell}_n\, {}_sf_{\ell m}` (Theorem V.5(c)).
67
+
68
+ Args:
69
+ flm: Spin-:math:`s` harmonic coefficients, ``[..., L, 2L - 1]``.
70
+ u: Column weights ``[..., 2N - 1, L]``.
71
+ s_in (int): Input spin :math:`s`.
72
+
73
+ Returns:
74
+ jnp.ndarray: Wigner coefficients, ``[..., 2N - 1, L, 2L - 1]``.
75
+
76
+ """
77
+ L, N = flm.shape[-2], (u.shape[-2] + 1) // 2
78
+ return (
79
+ u[..., :, :, None]
80
+ * (flm * lm_mask(L, s_in))[..., None, :, :]
81
+ * wigner_mask(L, N)
82
+ )
83
+
84
+
85
+ def group_to_spin_weights(flmn, v, s_out: int) -> jnp.ndarray:
86
+ r"""
87
+ An equivariant map from the rotation group to spin :math:`s'`, :math:`g_{\ell m} = \sum_k v^{\ell}_k F^{\ell}_{mk}` (Theorem V.5(c)).
88
+
89
+ Args:
90
+ flmn: Wigner coefficients, ``[..., 2N - 1, L, 2L - 1]``.
91
+ v: Row weights ``[..., 2N - 1, L]``.
92
+ s_out (int): Output spin :math:`s'`.
93
+
94
+ Returns:
95
+ jnp.ndarray: Spin-:math:`s'` harmonic coefficients, ``[..., L, 2L - 1]``.
96
+
97
+ """
98
+ L = flmn.shape[-2]
99
+ return jnp.einsum("...kl,...klm->...lm", v, flmn) * lm_mask(L, s_out)
100
+
101
+
102
+ def group_to_group_weights(flmn, B) -> jnp.ndarray:
103
+ r"""
104
+ An equivariant map on the rotation group, :math:`G^{\ell} = F^{\ell} B^{\ell}` (Theorem V.5(c)).
105
+
106
+ Args:
107
+ flmn: Wigner coefficients, ``[..., 2N_in - 1, L, 2L - 1]``.
108
+ B: Block weights ``[..., L, 2N_in - 1, 2N_out - 1]``.
109
+
110
+ Returns:
111
+ jnp.ndarray: Wigner coefficients, ``[..., 2N_out - 1, L, 2L - 1]``.
112
+
113
+ """
114
+ L, N_out = flmn.shape[-2], (B.shape[-1] + 1) // 2
115
+ return jnp.einsum("...klm,...lkn->...nlm", flmn, B) * wigner_mask(L, N_out)
116
+
117
+
118
+ # ------------------------------------------------------------------------------ kernel form
119
+
120
+
121
+ def spin_to_spin(flm, kernel, s_in: int, s_out: int) -> jnp.ndarray:
122
+ r"""
123
+ :math:`\mathcal{C}_{s\to s'}[\Psi]`, which reads the single entry :math:`\Psi^{\ell}_{-s',-s}` (Eq. 43).
124
+
125
+ Args:
126
+ flm: Spin-:math:`s` harmonic coefficients, ``[..., L, 2L - 1]``.
127
+ kernel: The kernel :math:`\Psi`, ``[2N - 1, L, 2L - 1]``.
128
+ s_in (int): Input spin :math:`s`.
129
+ s_out (int): Output spin :math:`s'`.
130
+
131
+ Returns:
132
+ jnp.ndarray: Spin-:math:`s'` harmonic coefficients, ``[..., L, 2L - 1]``.
133
+
134
+ """
135
+ return spin_multiplier(flm, kernels.to_multiplier(kernel, s_in, s_out), s_in, s_out)
136
+
137
+
138
+ def spin_to_group(flm, kernel, s_in: int, N_out: int = None) -> jnp.ndarray:
139
+ r"""
140
+ :math:`\mathcal{C}_{s\to G}[\Psi]`, which reads the column :math:`k = -s` (Eq. 42).
141
+
142
+ Args:
143
+ flm: Spin-:math:`s` harmonic coefficients, ``[..., L, 2L - 1]``.
144
+ kernel: The kernel :math:`\Psi`, ``[2N - 1, L, 2L - 1]``.
145
+ s_in (int): Input spin :math:`s`.
146
+ N_out (int, optional): Azimuthal band-limit of the output. Defaults to :math:`L`.
147
+
148
+ Returns:
149
+ jnp.ndarray: Wigner coefficients, ``[..., 2N_out - 1, L, 2L - 1]``.
150
+
151
+ """
152
+ N_out = flm.shape[-2] if N_out is None else N_out
153
+ return spin_to_group_weights(flm, kernels.to_column(kernel, s_in, N_out), s_in)
154
+
155
+
156
+ def group_to_spin(flmn, kernel, s_out: int) -> jnp.ndarray:
157
+ r"""
158
+ :math:`\mathcal{C}_{G\to s'}[\Psi]`, which reads the row :math:`n = -s'` (Eq. 44).
159
+
160
+ Args:
161
+ flmn: Wigner coefficients, ``[..., 2N - 1, L, 2L - 1]``.
162
+ kernel: The kernel :math:`\Psi`, ``[2N_k - 1, L, 2L - 1]``.
163
+ s_out (int): Output spin :math:`s'`.
164
+
165
+ Returns:
166
+ jnp.ndarray: Spin-:math:`s'` harmonic coefficients, ``[..., L, 2L - 1]``.
167
+
168
+ """
169
+ v = kernels.to_row(kernel, s_out, azimuthal_bandlimit(flmn))
170
+ return group_to_spin_weights(flmn, v, s_out)
171
+
172
+
173
+ def group_to_group(flmn, kernel, N_out: int = None) -> jnp.ndarray:
174
+ r"""
175
+ :math:`\mathcal{C}_{G\to G}[\Psi] = \Psi \star\,\cdot\,`, the group convolution (Eqs. 20, 45).
176
+
177
+ Args:
178
+ flmn: Wigner coefficients, ``[..., 2N_in - 1, L, 2L - 1]``.
179
+ kernel: The kernel :math:`\Psi`, ``[2N_k - 1, L, 2L - 1]``.
180
+ N_out (int, optional): Azimuthal band-limit of the output. Defaults to :math:`L`.
181
+
182
+ Returns:
183
+ jnp.ndarray: Wigner coefficients, ``[..., 2N_out - 1, L, 2L - 1]``.
184
+
185
+ """
186
+ N_out = flmn.shape[-2] if N_out is None else N_out
187
+ B = kernels.to_block(kernel, azimuthal_bandlimit(flmn), N_out)
188
+ return group_to_group_weights(flmn, B)
189
+
190
+
191
+ group_convolution = group_to_group
192
+
193
+
194
+ def lifted_convolution(x, kernel, s_in, s_out, N_out: int = None) -> jnp.ndarray:
195
+ r"""
196
+ The generalised lifted convolution :math:`\mathcal{C}_{s\to s'}[\Psi]` of Definition V.1.
197
+
198
+ Args:
199
+ x: Spin-:math:`s` harmonic coefficients ``[..., L, 2L - 1]`` or, if
200
+ ``s_in == "G"``, Wigner coefficients ``[..., 2N - 1, L, 2L - 1]``.
201
+ kernel: The kernel :math:`\Psi`, ``[2N_k - 1, L, 2L - 1]``.
202
+ s_in (int or "G"): Input spin, or ``"G"`` for an input on the rotation group.
203
+ s_out (int or "G"): Output spin, or ``"G"`` for an output on the rotation group.
204
+ N_out (int, optional): Azimuthal band-limit of an output on the rotation group.
205
+ Defaults to :math:`L`.
206
+
207
+ Returns:
208
+ jnp.ndarray: Harmonic or Wigner coefficients of the output.
209
+
210
+ """
211
+ if s_in == "G":
212
+ return (
213
+ group_to_group(x, kernel, N_out)
214
+ if s_out == "G"
215
+ else group_to_spin(x, kernel, s_out)
216
+ )
217
+ return (
218
+ spin_to_group(x, kernel, s_in, N_out)
219
+ if s_out == "G"
220
+ else spin_to_spin(x, kernel, s_in, s_out)
221
+ )