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 +35 -0
- s2conv/_version.py +24 -0
- s2conv/convolutions/__init__.py +40 -0
- s2conv/convolutions/classical.py +175 -0
- s2conv/convolutions/diagonal.py +41 -0
- s2conv/convolutions/lifted.py +221 -0
- s2conv/convolutions/physics.py +185 -0
- s2conv/convolutions/sampled.py +99 -0
- s2conv/convolutions/section.py +176 -0
- s2conv/kernels.py +371 -0
- s2conv/lifting.py +177 -0
- s2conv/transforms.py +321 -0
- s2conv/utils/__init__.py +32 -0
- s2conv/utils/equivariance.py +386 -0
- s2conv/utils/evaluation.py +113 -0
- s2conv/utils/indexing.py +160 -0
- s2conv/utils/rotation.py +199 -0
- s2conv-0.1.0.dist-info/METADATA +265 -0
- s2conv-0.1.0.dist-info/RECORD +22 -0
- s2conv-0.1.0.dist-info/WHEEL +5 -0
- s2conv-0.1.0.dist-info/licenses/LICENCE.txt +21 -0
- s2conv-0.1.0.dist-info/top_level.txt +1 -0
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
|
+
)
|