jax-healpy 0.2.1__tar.gz
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.
- jax_healpy-0.2.1/.github/workflows/release.yml +34 -0
- jax_healpy-0.2.1/.gitignore +10 -0
- jax_healpy-0.2.1/.pre-commit-config.yaml +27 -0
- jax_healpy-0.2.1/PKG-INFO +66 -0
- jax_healpy-0.2.1/README.md +32 -0
- jax_healpy-0.2.1/benchmarks/__main__.py +286 -0
- jax_healpy-0.2.1/benchmarks/script_sphtfunc.py +53 -0
- jax_healpy-0.2.1/docs/benchmarks/chart-darkbackground-n10000000.png +0 -0
- jax_healpy-0.2.1/docs/benchmarks/chart-default-n10000000.png +0 -0
- jax_healpy-0.2.1/jax_healpy/__init__.py +75 -0
- jax_healpy-0.2.1/jax_healpy/pixelfunc.py +1585 -0
- jax_healpy-0.2.1/jax_healpy/sphtfunc.py +344 -0
- jax_healpy-0.2.1/jax_healpy.egg-info/PKG-INFO +66 -0
- jax_healpy-0.2.1/jax_healpy.egg-info/SOURCES.txt +28 -0
- jax_healpy-0.2.1/jax_healpy.egg-info/dependency_links.txt +1 -0
- jax_healpy-0.2.1/jax_healpy.egg-info/requires.txt +17 -0
- jax_healpy-0.2.1/jax_healpy.egg-info/top_level.txt +1 -0
- jax_healpy-0.2.1/pyproject.toml +97 -0
- jax_healpy-0.2.1/setup.cfg +4 -0
- jax_healpy-0.2.1/tests/conftest.py +15 -0
- jax_healpy-0.2.1/tests/data/cl_wmap_band_iqumap_r9_7yr_W_v4_udgraded32_II_lmax64_rmmono_3iter.fits +1 -0
- jax_healpy-0.2.1/tests/pixelfunc/conftest.py +36 -0
- jax_healpy-0.2.1/tests/pixelfunc/test.py +179 -0
- jax_healpy-0.2.1/tests/pixelfunc/test_ang_pix.py +386 -0
- jax_healpy-0.2.1/tests/pixelfunc/test_ang_vec.py +77 -0
- jax_healpy-0.2.1/tests/pixelfunc/test_ring_nest.py +189 -0
- jax_healpy-0.2.1/tests/pixelfunc/test_vec_pix.py +365 -0
- jax_healpy-0.2.1/tests/sphtfunc/__init__.py +0 -0
- jax_healpy-0.2.1/tests/sphtfunc/conftest.py +24 -0
- jax_healpy-0.2.1/tests/sphtfunc/test_map_alm.py +120 -0
|
@@ -0,0 +1,34 @@
|
|
|
1
|
+
name: Release to PyPI
|
|
2
|
+
|
|
3
|
+
on:
|
|
4
|
+
release:
|
|
5
|
+
types: [created]
|
|
6
|
+
|
|
7
|
+
jobs:
|
|
8
|
+
publish:
|
|
9
|
+
name: Publish package to PyPI
|
|
10
|
+
runs-on: ubuntu-latest
|
|
11
|
+
|
|
12
|
+
permissions:
|
|
13
|
+
id-token: write
|
|
14
|
+
|
|
15
|
+
steps:
|
|
16
|
+
- name: Checkout repository
|
|
17
|
+
uses: actions/checkout@v4
|
|
18
|
+
|
|
19
|
+
- name: Set up Python
|
|
20
|
+
uses: actions/setup-python@v4
|
|
21
|
+
with:
|
|
22
|
+
python-version: '3.x'
|
|
23
|
+
|
|
24
|
+
- name: Install dependencies
|
|
25
|
+
run: |
|
|
26
|
+
python -m pip install --upgrade pip
|
|
27
|
+
pip install build twine
|
|
28
|
+
|
|
29
|
+
- name: Build package
|
|
30
|
+
run: |
|
|
31
|
+
python -m build
|
|
32
|
+
|
|
33
|
+
- name: Publish to PyPI
|
|
34
|
+
uses: pypa/gh-action-pypi-publish@release/v1
|
|
@@ -0,0 +1,27 @@
|
|
|
1
|
+
repos:
|
|
2
|
+
- repo: https://github.com/hadialqattan/pycln
|
|
3
|
+
rev: "v2.4.0"
|
|
4
|
+
hooks:
|
|
5
|
+
- id: pycln
|
|
6
|
+
args:
|
|
7
|
+
- --all
|
|
8
|
+
|
|
9
|
+
- repo: https://github.com/astral-sh/ruff-pre-commit
|
|
10
|
+
rev: v0.8.0
|
|
11
|
+
hooks:
|
|
12
|
+
- id: ruff
|
|
13
|
+
- id: ruff-format
|
|
14
|
+
|
|
15
|
+
- repo: https://github.com/pre-commit/pre-commit-hooks
|
|
16
|
+
rev: 'v5.0.0'
|
|
17
|
+
hooks:
|
|
18
|
+
- id: trailing-whitespace
|
|
19
|
+
- id: end-of-file-fixer
|
|
20
|
+
- id: check-yaml
|
|
21
|
+
- id: check-merge-conflict
|
|
22
|
+
|
|
23
|
+
- repo: https://github.com/PyCQA/bandit
|
|
24
|
+
rev: '1.8.0'
|
|
25
|
+
hooks:
|
|
26
|
+
- id: bandit
|
|
27
|
+
files: ^jax_healpy/
|
|
@@ -0,0 +1,66 @@
|
|
|
1
|
+
Metadata-Version: 2.2
|
|
2
|
+
Name: jax-healpy
|
|
3
|
+
Version: 0.2.1
|
|
4
|
+
Summary: Healpix JAX implementation.
|
|
5
|
+
Author-email: Pierre Chanial <pierre.chanial@gmail.com>, Simon Biquard <sbiquard@gmail.com>
|
|
6
|
+
Maintainer-email: Pierre Chanial <pierre.chanial@gmail.com>
|
|
7
|
+
Project-URL: homepage, https://pchanial.github.io/jax-healpy
|
|
8
|
+
Project-URL: repository, https://github.com/pchanial/jax-healpy
|
|
9
|
+
Keywords: scientific computing
|
|
10
|
+
Classifier: Programming Language :: Python
|
|
11
|
+
Classifier: Programming Language :: Python :: 3.9
|
|
12
|
+
Classifier: Programming Language :: Python :: 3.10
|
|
13
|
+
Classifier: Programming Language :: Python :: 3.11
|
|
14
|
+
Classifier: Intended Audience :: Science/Research
|
|
15
|
+
Classifier: Operating System :: OS Independent
|
|
16
|
+
Classifier: Topic :: Scientific/Engineering
|
|
17
|
+
Requires-Python: >=3.8
|
|
18
|
+
Description-Content-Type: text/markdown
|
|
19
|
+
Requires-Dist: jax
|
|
20
|
+
Requires-Dist: jaxtyping
|
|
21
|
+
Provides-Extra: dev
|
|
22
|
+
Requires-Dist: healpy; extra == "dev"
|
|
23
|
+
Requires-Dist: matplotlib; extra == "dev"
|
|
24
|
+
Requires-Dist: mypy; extra == "dev"
|
|
25
|
+
Requires-Dist: pandas; extra == "dev"
|
|
26
|
+
Requires-Dist: pytest; extra == "dev"
|
|
27
|
+
Requires-Dist: pytest-cov; extra == "dev"
|
|
28
|
+
Requires-Dist: pytest-mock; extra == "dev"
|
|
29
|
+
Requires-Dist: setuptools_scm; extra == "dev"
|
|
30
|
+
Requires-Dist: typer; extra == "dev"
|
|
31
|
+
Requires-Dist: PyYAML; extra == "dev"
|
|
32
|
+
Provides-Extra: recommended
|
|
33
|
+
Requires-Dist: s2fft; extra == "recommended"
|
|
34
|
+
|
|
35
|
+
# Healpy with JAX
|
|
36
|
+
|
|
37
|
+
This project intends to assess the interest of implementing healpy functions using JAX.
|
|
38
|
+
|
|
39
|
+
_WARNING: BETA STAGE!!!_
|
|
40
|
+
|
|
41
|
+
## Installation
|
|
42
|
+
|
|
43
|
+
First install JAX following the [documentation](https://jax.readthedocs.io/en/latest/installation.html).
|
|
44
|
+
|
|
45
|
+
Then install the package with:
|
|
46
|
+
|
|
47
|
+
```bash
|
|
48
|
+
pip install jax-healpy
|
|
49
|
+
```
|
|
50
|
+
|
|
51
|
+
To use the spherical harmonics functions,
|
|
52
|
+
you will need [s2fft](https://astro-informatics.github.io/s2fft/),
|
|
53
|
+
part of the recommended dependencies:
|
|
54
|
+
|
|
55
|
+
```bash
|
|
56
|
+
pip install jax-healpy[recommended]
|
|
57
|
+
```
|
|
58
|
+
|
|
59
|
+
## Benchmarks
|
|
60
|
+
|
|
61
|
+
Execution time measured on the [Jean Zay supercomputer](http://www.idris.fr/jean-zay/cpu/jean-zay-cpu-hw.html).
|
|
62
|
+
|
|
63
|
+
- CPU: Intel(R) Xeon(R) Gold 2648 @ 2.50GHz
|
|
64
|
+
- GPU: NVIDIA Tesla V100-SXM2-16GB
|
|
65
|
+
|
|
66
|
+

|
|
@@ -0,0 +1,32 @@
|
|
|
1
|
+
# Healpy with JAX
|
|
2
|
+
|
|
3
|
+
This project intends to assess the interest of implementing healpy functions using JAX.
|
|
4
|
+
|
|
5
|
+
_WARNING: BETA STAGE!!!_
|
|
6
|
+
|
|
7
|
+
## Installation
|
|
8
|
+
|
|
9
|
+
First install JAX following the [documentation](https://jax.readthedocs.io/en/latest/installation.html).
|
|
10
|
+
|
|
11
|
+
Then install the package with:
|
|
12
|
+
|
|
13
|
+
```bash
|
|
14
|
+
pip install jax-healpy
|
|
15
|
+
```
|
|
16
|
+
|
|
17
|
+
To use the spherical harmonics functions,
|
|
18
|
+
you will need [s2fft](https://astro-informatics.github.io/s2fft/),
|
|
19
|
+
part of the recommended dependencies:
|
|
20
|
+
|
|
21
|
+
```bash
|
|
22
|
+
pip install jax-healpy[recommended]
|
|
23
|
+
```
|
|
24
|
+
|
|
25
|
+
## Benchmarks
|
|
26
|
+
|
|
27
|
+
Execution time measured on the [Jean Zay supercomputer](http://www.idris.fr/jean-zay/cpu/jean-zay-cpu-hw.html).
|
|
28
|
+
|
|
29
|
+
- CPU: Intel(R) Xeon(R) Gold 2648 @ 2.50GHz
|
|
30
|
+
- GPU: NVIDIA Tesla V100-SXM2-16GB
|
|
31
|
+
|
|
32
|
+

|
|
@@ -0,0 +1,286 @@
|
|
|
1
|
+
import os
|
|
2
|
+
import platform
|
|
3
|
+
import re
|
|
4
|
+
import subprocess
|
|
5
|
+
import timeit
|
|
6
|
+
from collections.abc import Callable
|
|
7
|
+
from dataclasses import asdict, dataclass
|
|
8
|
+
from pathlib import Path
|
|
9
|
+
from pprint import pprint
|
|
10
|
+
from typing import Any, Iterable, Literal
|
|
11
|
+
|
|
12
|
+
import healpy as hp
|
|
13
|
+
import jax
|
|
14
|
+
import matplotlib.pyplot as plt
|
|
15
|
+
import numpy as np
|
|
16
|
+
import pandas as pd
|
|
17
|
+
import typer
|
|
18
|
+
import yaml
|
|
19
|
+
from jaxtyping import ArrayLike
|
|
20
|
+
from matplotlib.ticker import ScalarFormatter
|
|
21
|
+
|
|
22
|
+
BENCH_PATH = Path(__file__).parent / 'results'
|
|
23
|
+
BENCHMARKED_FUNCS = [
|
|
24
|
+
'ang2vec',
|
|
25
|
+
'vec2ang',
|
|
26
|
+
'ang2pix',
|
|
27
|
+
'pix2ang',
|
|
28
|
+
'vec2pix',
|
|
29
|
+
'pix2vec',
|
|
30
|
+
'ring2nest',
|
|
31
|
+
'nest2ring',
|
|
32
|
+
'pix2xyf',
|
|
33
|
+
'xyf2pix',
|
|
34
|
+
'reorder',
|
|
35
|
+
]
|
|
36
|
+
CHART_PATH_NAME = 'chart-{style}-n{n}.png'
|
|
37
|
+
|
|
38
|
+
# TODO: use those when typer supports Literals
|
|
39
|
+
LibraryType = Literal['healpy', 'jax-healpy']
|
|
40
|
+
PrecisionType = Literal['32', '64']
|
|
41
|
+
|
|
42
|
+
app = typer.Typer()
|
|
43
|
+
|
|
44
|
+
|
|
45
|
+
@dataclass(frozen=True)
|
|
46
|
+
class BenchmarkResult:
|
|
47
|
+
library: str
|
|
48
|
+
version: str
|
|
49
|
+
processor: str
|
|
50
|
+
processor_type: str
|
|
51
|
+
precision: int
|
|
52
|
+
n: int
|
|
53
|
+
execution_times_s: dict[str, float]
|
|
54
|
+
|
|
55
|
+
|
|
56
|
+
def bench_it(
|
|
57
|
+
library: str,
|
|
58
|
+
func_name: str,
|
|
59
|
+
nside: int,
|
|
60
|
+
n: int,
|
|
61
|
+
precision: str,
|
|
62
|
+
rng: np.random.Generator,
|
|
63
|
+
) -> float:
|
|
64
|
+
if precision == '32':
|
|
65
|
+
dtype = np.dtype('float32')
|
|
66
|
+
elif precision == '64':
|
|
67
|
+
dtype = np.dtype('float64')
|
|
68
|
+
else:
|
|
69
|
+
raise ValueError(f'Invalid precision {precision}')
|
|
70
|
+
|
|
71
|
+
args = _get_args(library, func_name, nside, n, dtype, rng)
|
|
72
|
+
func = _get_func(library, func_name, *args)
|
|
73
|
+
with jax.experimental.enable_x64(precision == '64'):
|
|
74
|
+
return time_it(func)
|
|
75
|
+
|
|
76
|
+
|
|
77
|
+
def _get_args(
|
|
78
|
+
library: str,
|
|
79
|
+
func_name: str,
|
|
80
|
+
nside: int,
|
|
81
|
+
n: int,
|
|
82
|
+
dtype: np.dtype,
|
|
83
|
+
rng: np.random.Generator,
|
|
84
|
+
) -> tuple[ArrayLike]:
|
|
85
|
+
if func_name in {'ang2vec', 'ang2pix'}:
|
|
86
|
+
theta = rng.uniform(0, np.pi, size=n).astype(dtype)
|
|
87
|
+
phi = rng.uniform(0, 2 * np.pi, size=n).astype(dtype)
|
|
88
|
+
if func_name == 'ang2vec':
|
|
89
|
+
args = (theta, phi)
|
|
90
|
+
else:
|
|
91
|
+
args = (nside, theta, phi)
|
|
92
|
+
|
|
93
|
+
elif func_name in {'vec2ang', 'vec2pix'}:
|
|
94
|
+
vec = rng.uniform(-np.pi, np.pi, size=(3, n))
|
|
95
|
+
vec /= np.sqrt(np.sum(vec**2, axis=0))
|
|
96
|
+
vec = vec.astype(dtype)
|
|
97
|
+
if func_name == 'vec2ang':
|
|
98
|
+
args = (vec.T.copy(),)
|
|
99
|
+
else:
|
|
100
|
+
args = (nside, vec[0], vec[1], vec[2])
|
|
101
|
+
|
|
102
|
+
elif func_name in {'pix2ang', 'pix2vec', 'pix2xyf', 'ring2nest', 'nest2ring'}:
|
|
103
|
+
pixels = rng.uniform(0, hp.nside2npix(nside), size=n).astype(int)
|
|
104
|
+
args = (nside, pixels)
|
|
105
|
+
|
|
106
|
+
elif func_name == 'xyf2pix':
|
|
107
|
+
x = rng.uniform(0, nside, size=n).astype(int)
|
|
108
|
+
y = rng.uniform(0, nside, size=n).astype(int)
|
|
109
|
+
f = rng.uniform(0, 12, size=n).astype(int)
|
|
110
|
+
args = (nside, x, y, f)
|
|
111
|
+
|
|
112
|
+
elif func_name == 'reorder':
|
|
113
|
+
npix = hp.nside2npix(nside)
|
|
114
|
+
map_in = rng.uniform(size=npix)
|
|
115
|
+
args = (map_in,)
|
|
116
|
+
|
|
117
|
+
else:
|
|
118
|
+
raise NotImplementedError
|
|
119
|
+
|
|
120
|
+
if library == 'jax-healpy':
|
|
121
|
+
args = tuple(jax.device_put(_) if isinstance(_, np.ndarray) else _ for _ in args)
|
|
122
|
+
|
|
123
|
+
return args
|
|
124
|
+
|
|
125
|
+
|
|
126
|
+
def _get_func(library: str, func_name: str, *args: Any):
|
|
127
|
+
if library == 'jax-healpy':
|
|
128
|
+
import jax_healpy as module
|
|
129
|
+
else:
|
|
130
|
+
import healpy as module
|
|
131
|
+
|
|
132
|
+
func = getattr(module, func_name)
|
|
133
|
+
if library == 'healpy':
|
|
134
|
+
if func_name == 'reorder':
|
|
135
|
+
func_ = lambda: func(*args, r2n=True) # noqa: E731
|
|
136
|
+
else:
|
|
137
|
+
func_ = lambda: func(*args) # noqa: E731
|
|
138
|
+
else:
|
|
139
|
+
if func_name in {'pix2ang', 'vec2ang'}:
|
|
140
|
+
|
|
141
|
+
def func_() -> None:
|
|
142
|
+
theta, phi = func(*args)
|
|
143
|
+
theta.block_until_ready()
|
|
144
|
+
phi.block_until_ready()
|
|
145
|
+
|
|
146
|
+
elif func_name == 'pix2xyf':
|
|
147
|
+
|
|
148
|
+
def func_() -> None:
|
|
149
|
+
x, y, f = func(*args)
|
|
150
|
+
x.block_until_ready()
|
|
151
|
+
y.block_until_ready()
|
|
152
|
+
f.block_until_ready()
|
|
153
|
+
|
|
154
|
+
elif func_name == 'reorder':
|
|
155
|
+
|
|
156
|
+
def func_() -> None:
|
|
157
|
+
func(*args, r2n=True).block_until_ready()
|
|
158
|
+
|
|
159
|
+
else:
|
|
160
|
+
|
|
161
|
+
def func_() -> None:
|
|
162
|
+
func(*args).block_until_ready()
|
|
163
|
+
|
|
164
|
+
func_() # discard first call, which includes compilation
|
|
165
|
+
|
|
166
|
+
return func_
|
|
167
|
+
|
|
168
|
+
|
|
169
|
+
def time_it(func: Callable[[], None]) -> float:
|
|
170
|
+
timer = timeit.Timer(func)
|
|
171
|
+
number, _ = timer.autorange()
|
|
172
|
+
execution_time = min(timer.repeat(number=number)) / number
|
|
173
|
+
return execution_time
|
|
174
|
+
|
|
175
|
+
|
|
176
|
+
@app.command()
|
|
177
|
+
def run(
|
|
178
|
+
library: str,
|
|
179
|
+
nside: int = 512,
|
|
180
|
+
n: int = 10_000_000,
|
|
181
|
+
precision: str = '64',
|
|
182
|
+
) -> None:
|
|
183
|
+
if library == 'jax-healpy':
|
|
184
|
+
version = f'jax({jax.__version__})'
|
|
185
|
+
device = jax.devices()[0]
|
|
186
|
+
processor = device.device_kind
|
|
187
|
+
if processor == 'cpu':
|
|
188
|
+
processor = _get_cpu_name()
|
|
189
|
+
processor_type = 'cpu'
|
|
190
|
+
else:
|
|
191
|
+
processor_type = 'gpu'
|
|
192
|
+
elif library == 'healpy':
|
|
193
|
+
version = hp.__version__
|
|
194
|
+
processor = _get_cpu_name()
|
|
195
|
+
processor_type = 'cpu'
|
|
196
|
+
else:
|
|
197
|
+
raise ValueError(f'Invalid library {library}')
|
|
198
|
+
|
|
199
|
+
rng = np.random.default_rng(0)
|
|
200
|
+
execution_times = {}
|
|
201
|
+
|
|
202
|
+
print(f'Running {library}...')
|
|
203
|
+
for func_name in BENCHMARKED_FUNCS:
|
|
204
|
+
execution_time = bench_it(library, func_name, nside, n, precision, rng)
|
|
205
|
+
execution_times[func_name] = execution_time
|
|
206
|
+
|
|
207
|
+
result = BenchmarkResult(
|
|
208
|
+
library=library,
|
|
209
|
+
version=version,
|
|
210
|
+
processor=processor,
|
|
211
|
+
processor_type=processor_type,
|
|
212
|
+
precision=int(precision),
|
|
213
|
+
n=n,
|
|
214
|
+
execution_times_s=execution_times,
|
|
215
|
+
)
|
|
216
|
+
pprint(result)
|
|
217
|
+
|
|
218
|
+
chars = '[] '
|
|
219
|
+
formatted_processor = processor.lower().rstrip(chars).replace('(c)', '').replace('(tm)', '')
|
|
220
|
+
filename = f'{library}-{formatted_processor}-precision{precision}-n{n}.yaml'
|
|
221
|
+
for char in chars:
|
|
222
|
+
filename = filename.replace(char, '-')
|
|
223
|
+
BENCH_PATH.mkdir(parents=True, exist_ok=True)
|
|
224
|
+
with open(BENCH_PATH / filename, 'w') as f:
|
|
225
|
+
yaml.dump(asdict(result), f)
|
|
226
|
+
|
|
227
|
+
|
|
228
|
+
def _get_cpu_name():
|
|
229
|
+
if platform.system() == 'Darwin':
|
|
230
|
+
os.environ['PATH'] = os.environ['PATH'] + os.pathsep + '/usr/sbin'
|
|
231
|
+
command = 'sysctl -n machdep.cpu.brand_string'
|
|
232
|
+
return subprocess.check_output(command).strip()
|
|
233
|
+
elif platform.system() == 'Linux':
|
|
234
|
+
command = 'cat /proc/cpuinfo'
|
|
235
|
+
all_info = subprocess.check_output(command, shell=True).decode().strip()
|
|
236
|
+
for line in all_info.split('\n'):
|
|
237
|
+
if 'model name' in line:
|
|
238
|
+
return re.sub('.*model name.*:', '', line, 1).strip()
|
|
239
|
+
return platform.processor()
|
|
240
|
+
|
|
241
|
+
|
|
242
|
+
@app.command()
|
|
243
|
+
def collect():
|
|
244
|
+
files = BENCH_PATH.glob('*-n*.yaml')
|
|
245
|
+
ns = {}
|
|
246
|
+
for file in files:
|
|
247
|
+
result_as_dict = yaml.safe_load(file.read_text())
|
|
248
|
+
result = BenchmarkResult(**result_as_dict)
|
|
249
|
+
ns.setdefault(result.n, []).append(result)
|
|
250
|
+
|
|
251
|
+
for n, results in ns.items():
|
|
252
|
+
results = sorted(
|
|
253
|
+
results,
|
|
254
|
+
key=lambda _: sum(_.execution_times_s.values()) / len(BENCHMARKED_FUNCS),
|
|
255
|
+
)
|
|
256
|
+
export_chart(n, results, 'default')
|
|
257
|
+
export_chart(n, results, 'dark_background')
|
|
258
|
+
|
|
259
|
+
|
|
260
|
+
def export_chart(n: int, results: Iterable[BenchmarkResult], style: str):
|
|
261
|
+
chart_path = BENCH_PATH / CHART_PATH_NAME.format(style=style.replace('_', ''), n=n)
|
|
262
|
+
plt.style.use(style)
|
|
263
|
+
# I've used pandas because I was not able to do what I wanted with matplotlib alone
|
|
264
|
+
df = pd.DataFrame(
|
|
265
|
+
[
|
|
266
|
+
[f'{res.library}[{res.processor_type}{res.precision}]']
|
|
267
|
+
+ [res.execution_times_s[_] for _ in BENCHMARKED_FUNCS]
|
|
268
|
+
for res in results
|
|
269
|
+
],
|
|
270
|
+
columns=['library'] + BENCHMARKED_FUNCS,
|
|
271
|
+
)
|
|
272
|
+
for func in BENCHMARKED_FUNCS:
|
|
273
|
+
df[func] /= df[func][0]
|
|
274
|
+
pprint(df)
|
|
275
|
+
|
|
276
|
+
ax = df.plot.bar(x='library', title='Benchmark (lower is better)', rot=60)
|
|
277
|
+
ax.legend(framealpha=0, loc='upper left')
|
|
278
|
+
plt.yscale('log')
|
|
279
|
+
ax.yaxis.set_major_formatter(ScalarFormatter())
|
|
280
|
+
plt.xlabel('')
|
|
281
|
+
plt.tight_layout()
|
|
282
|
+
plt.savefig(chart_path, transparent=True)
|
|
283
|
+
|
|
284
|
+
|
|
285
|
+
if __name__ == '__main__':
|
|
286
|
+
app()
|
|
@@ -0,0 +1,53 @@
|
|
|
1
|
+
import timeit
|
|
2
|
+
from pathlib import Path
|
|
3
|
+
from typing import Callable
|
|
4
|
+
|
|
5
|
+
import healpy as hp
|
|
6
|
+
import jax
|
|
7
|
+
import numpy as np
|
|
8
|
+
|
|
9
|
+
import jax_healpy as jhp
|
|
10
|
+
|
|
11
|
+
data_path = Path(jhp.__file__).parent.parent / 'tests/data'
|
|
12
|
+
|
|
13
|
+
|
|
14
|
+
def cla(data_path: Path) -> np.ndarray:
|
|
15
|
+
return hp.read_cl(data_path / 'cl_wmap_band_iqumap_r9_7yr_W_v4_udgraded32_II_lmax64_rmmono_3iter.fits')
|
|
16
|
+
|
|
17
|
+
|
|
18
|
+
nside = 256
|
|
19
|
+
lmax = 2 * nside
|
|
20
|
+
fwhm_deg = 7.0
|
|
21
|
+
|
|
22
|
+
|
|
23
|
+
orig = hp.synfast(
|
|
24
|
+
cla(data_path),
|
|
25
|
+
nside,
|
|
26
|
+
lmax=lmax,
|
|
27
|
+
pixwin=False,
|
|
28
|
+
fwhm=np.radians(fwhm_deg),
|
|
29
|
+
new=False,
|
|
30
|
+
)
|
|
31
|
+
orig_d = jax.device_put(orig)
|
|
32
|
+
|
|
33
|
+
|
|
34
|
+
def map2alm_hp():
|
|
35
|
+
hp.map2alm(orig, iter=0, lmax=2 * nside - 1, use_weights=False)
|
|
36
|
+
|
|
37
|
+
|
|
38
|
+
def map2alm_jhp():
|
|
39
|
+
jhp.map2alm(orig_d).block_until_ready()
|
|
40
|
+
|
|
41
|
+
|
|
42
|
+
def time_it(func: Callable[[], None]) -> float:
|
|
43
|
+
timer = timeit.Timer(func)
|
|
44
|
+
number, _ = timer.autorange()
|
|
45
|
+
execution_time = min(timer.repeat(number=number)) / number
|
|
46
|
+
return execution_time
|
|
47
|
+
|
|
48
|
+
|
|
49
|
+
print(time_it(map2alm_hp))
|
|
50
|
+
# 0.0038599122400046326
|
|
51
|
+
|
|
52
|
+
print(time_it(map2alm_jhp))
|
|
53
|
+
# 0.437624677999338
|
|
Binary file
|
|
Binary file
|
|
@@ -0,0 +1,75 @@
|
|
|
1
|
+
from jax import config as _config
|
|
2
|
+
|
|
3
|
+
from .pixelfunc import (
|
|
4
|
+
UNSEEN,
|
|
5
|
+
ang2pix,
|
|
6
|
+
ang2vec,
|
|
7
|
+
isnpixok,
|
|
8
|
+
isnsideok,
|
|
9
|
+
maptype,
|
|
10
|
+
nest2ring,
|
|
11
|
+
npix2nside,
|
|
12
|
+
npix2order,
|
|
13
|
+
nside2npix,
|
|
14
|
+
nside2order,
|
|
15
|
+
nside2pixarea,
|
|
16
|
+
nside2resol,
|
|
17
|
+
order2npix,
|
|
18
|
+
order2nside,
|
|
19
|
+
pix2ang,
|
|
20
|
+
pix2vec,
|
|
21
|
+
pix2xyf,
|
|
22
|
+
reorder,
|
|
23
|
+
ring2nest,
|
|
24
|
+
vec2ang,
|
|
25
|
+
vec2pix,
|
|
26
|
+
xyf2pix,
|
|
27
|
+
)
|
|
28
|
+
from .sphtfunc import alm2map, map2alm
|
|
29
|
+
|
|
30
|
+
__all__ = [
|
|
31
|
+
'UNSEEN',
|
|
32
|
+
'pix2ang',
|
|
33
|
+
'ang2pix',
|
|
34
|
+
'pix2xyf',
|
|
35
|
+
'xyf2pix',
|
|
36
|
+
'pix2vec',
|
|
37
|
+
'vec2pix',
|
|
38
|
+
'ang2vec',
|
|
39
|
+
'vec2ang',
|
|
40
|
+
# 'get_interp_weights',
|
|
41
|
+
# 'get_interp_val',
|
|
42
|
+
# 'get_all_neighbours',
|
|
43
|
+
# 'max_pixrad',
|
|
44
|
+
'nest2ring',
|
|
45
|
+
'ring2nest',
|
|
46
|
+
'reorder',
|
|
47
|
+
# 'ud_grade',
|
|
48
|
+
# 'UNSEEN',
|
|
49
|
+
# 'mask_good',
|
|
50
|
+
# 'mask_bad',
|
|
51
|
+
# 'ma',
|
|
52
|
+
# 'fit_dipole',
|
|
53
|
+
# 'remove_dipole',
|
|
54
|
+
# 'fit_monopole',
|
|
55
|
+
# 'remove_monopole',
|
|
56
|
+
'nside2npix',
|
|
57
|
+
'npix2nside',
|
|
58
|
+
'nside2order',
|
|
59
|
+
'order2nside',
|
|
60
|
+
'order2npix',
|
|
61
|
+
'npix2order',
|
|
62
|
+
'nside2resol',
|
|
63
|
+
'nside2pixarea',
|
|
64
|
+
'isnsideok',
|
|
65
|
+
'isnpixok',
|
|
66
|
+
# 'get_map_size',
|
|
67
|
+
# 'get_min_valid_nside',
|
|
68
|
+
# 'get_nside',
|
|
69
|
+
'maptype',
|
|
70
|
+
# 'ma_to_array',
|
|
71
|
+
'alm2map',
|
|
72
|
+
'map2alm',
|
|
73
|
+
]
|
|
74
|
+
|
|
75
|
+
_config.update('jax_enable_x64', True)
|