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.
Files changed (30) hide show
  1. jax_healpy-0.2.1/.github/workflows/release.yml +34 -0
  2. jax_healpy-0.2.1/.gitignore +10 -0
  3. jax_healpy-0.2.1/.pre-commit-config.yaml +27 -0
  4. jax_healpy-0.2.1/PKG-INFO +66 -0
  5. jax_healpy-0.2.1/README.md +32 -0
  6. jax_healpy-0.2.1/benchmarks/__main__.py +286 -0
  7. jax_healpy-0.2.1/benchmarks/script_sphtfunc.py +53 -0
  8. jax_healpy-0.2.1/docs/benchmarks/chart-darkbackground-n10000000.png +0 -0
  9. jax_healpy-0.2.1/docs/benchmarks/chart-default-n10000000.png +0 -0
  10. jax_healpy-0.2.1/jax_healpy/__init__.py +75 -0
  11. jax_healpy-0.2.1/jax_healpy/pixelfunc.py +1585 -0
  12. jax_healpy-0.2.1/jax_healpy/sphtfunc.py +344 -0
  13. jax_healpy-0.2.1/jax_healpy.egg-info/PKG-INFO +66 -0
  14. jax_healpy-0.2.1/jax_healpy.egg-info/SOURCES.txt +28 -0
  15. jax_healpy-0.2.1/jax_healpy.egg-info/dependency_links.txt +1 -0
  16. jax_healpy-0.2.1/jax_healpy.egg-info/requires.txt +17 -0
  17. jax_healpy-0.2.1/jax_healpy.egg-info/top_level.txt +1 -0
  18. jax_healpy-0.2.1/pyproject.toml +97 -0
  19. jax_healpy-0.2.1/setup.cfg +4 -0
  20. jax_healpy-0.2.1/tests/conftest.py +15 -0
  21. jax_healpy-0.2.1/tests/data/cl_wmap_band_iqumap_r9_7yr_W_v4_udgraded32_II_lmax64_rmmono_3iter.fits +1 -0
  22. jax_healpy-0.2.1/tests/pixelfunc/conftest.py +36 -0
  23. jax_healpy-0.2.1/tests/pixelfunc/test.py +179 -0
  24. jax_healpy-0.2.1/tests/pixelfunc/test_ang_pix.py +386 -0
  25. jax_healpy-0.2.1/tests/pixelfunc/test_ang_vec.py +77 -0
  26. jax_healpy-0.2.1/tests/pixelfunc/test_ring_nest.py +189 -0
  27. jax_healpy-0.2.1/tests/pixelfunc/test_vec_pix.py +365 -0
  28. jax_healpy-0.2.1/tests/sphtfunc/__init__.py +0 -0
  29. jax_healpy-0.2.1/tests/sphtfunc/conftest.py +24 -0
  30. 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,10 @@
1
+ .coverage
2
+ MANIFEST
3
+ build
4
+ dist
5
+ __pycache__
6
+ *.egg-info
7
+ /.idea/
8
+ /.vscode/
9
+ /venv*/
10
+ /benchmarks/results/
@@ -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
+ ![Benchmark](/docs/benchmarks/chart-darkbackground-n10000000.png)
@@ -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
+ ![Benchmark](/docs/benchmarks/chart-darkbackground-n10000000.png)
@@ -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
@@ -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)