jaxFMM 0.0.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.
@@ -0,0 +1,22 @@
1
+ # Byte-compiled / optimized / DLL files
2
+ __pycache__/
3
+ *.py[cod]
4
+ *$py.class
5
+
6
+ # C extensions
7
+ *.so
8
+ *.cpp
9
+
10
+ # build
11
+ .pytest_cache/
12
+ build/
13
+ *.egg-info/
14
+
15
+ # vim
16
+ *.swp
17
+
18
+ # docs build files
19
+ docs/_build
20
+
21
+ # local dev folder
22
+ dev
@@ -0,0 +1,61 @@
1
+ default:
2
+ interruptible: true
3
+
4
+ .test_template:
5
+ stage: test
6
+ variables:
7
+ PYTEST_ADDOPTS: "--color=yes"
8
+ script: python -m pytest --verbose
9
+
10
+ .cpu_template:
11
+ extends: .test_template
12
+ image: mambaorg/micromamba:jammy
13
+ before_script:
14
+ - micromamba install -y -n base -c conda-forge "python=$PYTHON_VERSION"
15
+ - python --version
16
+ - pip install .[dev]
17
+
18
+ .gpu_template:
19
+ extends: .test_template
20
+ image: mambaorg/micromamba:jammy-cuda-12.3.1
21
+ before_script:
22
+ - micromamba install -y -n base -c conda-forge "python=$PYTHON_VERSION"
23
+ - python --version
24
+ - nvidia-smi
25
+ - pip install .[dev,cuda]
26
+ tags:
27
+ - saas-linux-medium-amd64-gpu-standard
28
+
29
+ # test_py3_11_gpu:
30
+ # extends: .gpu_template
31
+ # variables:
32
+ # PYTHON_VERSION: "3.11"
33
+
34
+ test_py3_11_cpu:
35
+ extends: .cpu_template
36
+ variables:
37
+ PYTHON_VERSION: "3.11"
38
+
39
+ deploy_pypi:
40
+ stage: deploy
41
+ image: !reference [.cpu_template, image]
42
+ before_script:
43
+ - micromamba install -y -n base -c conda-forge 'python=3.11' make pandoc
44
+ - python -m pip install --upgrade twine
45
+ - python -m pip install --upgrade build
46
+ variables:
47
+ TWINE_USERNAME: __token__
48
+ TWINE_PASSWORD: ${PYPI_TOKEN}
49
+ script:
50
+ - python -m build
51
+ - python -m twine upload dist/*
52
+ rules:
53
+ - if: $CI_COMMIT_TAG
54
+ interruptible: false
55
+
56
+ workflow:
57
+ rules:
58
+ - if: $CI_COMMIT_TAG
59
+ - if: $CI_COMMIT_REF_NAME == $CI_DEFAULT_BRANCH
60
+ - if: $CI_PIPELINE_SOURCE == "merge_request_event"
61
+ - if: $CI_PIPELINE_SOURCE == "web"
jaxfmm-0.0.1/LICENSE ADDED
@@ -0,0 +1,21 @@
1
+ MIT License
2
+
3
+ Copyright (c) 2025 Robert Kraft
4
+
5
+ Permission is hereby granted, free of charge, to any person obtaining a copy
6
+ of this software and associated documentation files (the "Software"), to deal
7
+ in the Software without restriction, including without limitation the rights
8
+ to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
9
+ copies of the Software, and to permit persons to whom the Software is
10
+ furnished to do so, subject to the following conditions:
11
+
12
+ The above copyright notice and this permission notice shall be included in all
13
+ copies or substantial portions of the Software.
14
+
15
+ THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
16
+ IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
17
+ FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
18
+ AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
19
+ LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
20
+ OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
21
+ SOFTWARE.
jaxfmm-0.0.1/PKG-INFO ADDED
@@ -0,0 +1,98 @@
1
+ Metadata-Version: 2.4
2
+ Name: jaxFMM
3
+ Version: 0.0.1
4
+ Summary: Adaptive Fast Multipole Method with Laplace kernel in JAX.
5
+ Project-URL: Homepage, https://gitlab.com/jaxfmm/jaxfmm
6
+ Author-email: Robert Kraft <robert.kraft@univie.ac.at>
7
+ License: MIT License
8
+
9
+ Copyright (c) 2025 Robert Kraft
10
+
11
+ Permission is hereby granted, free of charge, to any person obtaining a copy
12
+ of this software and associated documentation files (the "Software"), to deal
13
+ in the Software without restriction, including without limitation the rights
14
+ to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
15
+ copies of the Software, and to permit persons to whom the Software is
16
+ furnished to do so, subject to the following conditions:
17
+
18
+ The above copyright notice and this permission notice shall be included in all
19
+ copies or substantial portions of the Software.
20
+
21
+ THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
22
+ IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
23
+ FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
24
+ AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
25
+ LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
26
+ OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
27
+ SOFTWARE.
28
+ License-File: LICENSE
29
+ Keywords: FMM,N-body,jax,potential,treecode
30
+ Classifier: Development Status :: 3 - Alpha
31
+ Classifier: Intended Audience :: Science/Research
32
+ Classifier: License :: OSI Approved :: GNU General Public License v3 (GPLv3)
33
+ Classifier: Programming Language :: Python :: 3
34
+ Classifier: Programming Language :: Python :: 3 :: Only
35
+ Classifier: Programming Language :: Python :: 3.9
36
+ Classifier: Programming Language :: Python :: 3.10
37
+ Classifier: Programming Language :: Python :: 3.11
38
+ Classifier: Programming Language :: Python :: 3.12
39
+ Classifier: Programming Language :: Python :: 3.13
40
+ Classifier: Topic :: Scientific/Engineering :: Physics
41
+ Requires-Python: >=3.9
42
+ Requires-Dist: jax
43
+ Provides-Extra: cuda
44
+ Requires-Dist: jax[cuda]; extra == 'cuda'
45
+ Provides-Extra: dev
46
+ Requires-Dist: matplotlib; extra == 'dev'
47
+ Requires-Dist: pytest; extra == 'dev'
48
+ Description-Content-Type: text/markdown
49
+
50
+ # jaxFMM
51
+
52
+ jaxFMM is an open source implementation of the Fast Multipole Method in JAX. The goal is to offer an easily readable/maintainable FMM implementation with good performance that runs on CPU/GPU and supports autodiff. This is enabled through JAX's just-in-time compiler.
53
+
54
+ ## Installation and Usage
55
+
56
+ jaxFMM depends only on JAX and can be installed from pypi or by downloading the source as follows:
57
+
58
+ pip install jaxfmm
59
+
60
+ If you want to run jaxFMM on GPUs, the easiest way is to use NVIDIA CUDA and cuDNN from pip wheels by instead typing:
61
+
62
+ pip install jaxfmm[cuda]
63
+
64
+ Using a custom, self-installed CUDA with jax is [described in the JAX documentation](https://docs.jax.dev/en/latest/installation.html).
65
+
66
+ The [unitcube demo](/demos/unitcube.py) is a short and simple example demonstrating how to use jaxFMM.
67
+
68
+ ## Features
69
+
70
+ There are many flavors of FMM implementations. In short, jaxFMM currently:
71
+
72
+ - only supports the Laplacian kernel.
73
+ - only supports point charges.
74
+ - only supports evaluation positions that are the same as the source positions.
75
+ - uses real basis functions computed via recurrence relations.
76
+ - uses "nested sum" O(p^4) M2M/M2L/L2L transformations.
77
+ - uses a non-uniform 2^N-ary tree hierarchy (directly inspired by [this work of A. Goude and S. Engblom](https://link.springer.com/article/10.1007/s11227-012-0836-0)), allowing arbitrary shape of the boxes in the hierarchy and guaranteeing balanced trees but requiring storage of interaction lists.
78
+ - has jit-compiled functions and autodiff for every substep of the algorithm except for the generation of interaction lists.
79
+
80
+ In summary, jaxFMM in its current state can do adaptive point charge FMM for Laplace kernels with good performance for lower expansion orders (p <= 3) and reasonably homogenous distributions. Autodiff only works if the particle positions remain constant. A first benchmark of uniformly distributed charges in the unit cube, computed on Google Cloud [g2-standard-8](https://cloud.google.com/compute/docs/gpus#l4-gpus) (GPU timings) and [c3d-highmem-16](https://cloud.google.com/compute/docs/general-purpose-machines#c3d_series) (CPU timings) machines can be found below:
81
+
82
+ <img src="docs/images/jax_unitcube_benchmark_p3.png" alt="unitcube benchmark" width="500"/>
83
+
84
+ ## TODOs
85
+
86
+ jaxFMM is primarily developed for my PhD project, where I am working on a GPU parallel FMM stray field evaluation routine with autodiff for finite-element micromagnetics. Alongside the very early state that it is in, this explains the currently limited feature set and design decisions mentioned above.
87
+
88
+ Contributions are always welcome however, and I plan to still improve jaxFMM. Topics that come to mind here are:
89
+
90
+ - volume FMM.
91
+ - other kernels or even a kernel-independent formulation.
92
+ - independent evaluation and source positions.
93
+ - distributed parallelism via jax.sharding.
94
+ - faster M2M/M2L/L2L transformations.
95
+ - faster jit-compilation, especially for large systems, higher orders and gradient computations.
96
+ - jit-compilable (and differentiable) interaction list generation, if possible.
97
+ - various other performance improvements.
98
+ - some degree of autotuning for the parameters.
jaxfmm-0.0.1/README.md ADDED
@@ -0,0 +1,49 @@
1
+ # jaxFMM
2
+
3
+ jaxFMM is an open source implementation of the Fast Multipole Method in JAX. The goal is to offer an easily readable/maintainable FMM implementation with good performance that runs on CPU/GPU and supports autodiff. This is enabled through JAX's just-in-time compiler.
4
+
5
+ ## Installation and Usage
6
+
7
+ jaxFMM depends only on JAX and can be installed from pypi or by downloading the source as follows:
8
+
9
+ pip install jaxfmm
10
+
11
+ If you want to run jaxFMM on GPUs, the easiest way is to use NVIDIA CUDA and cuDNN from pip wheels by instead typing:
12
+
13
+ pip install jaxfmm[cuda]
14
+
15
+ Using a custom, self-installed CUDA with jax is [described in the JAX documentation](https://docs.jax.dev/en/latest/installation.html).
16
+
17
+ The [unitcube demo](/demos/unitcube.py) is a short and simple example demonstrating how to use jaxFMM.
18
+
19
+ ## Features
20
+
21
+ There are many flavors of FMM implementations. In short, jaxFMM currently:
22
+
23
+ - only supports the Laplacian kernel.
24
+ - only supports point charges.
25
+ - only supports evaluation positions that are the same as the source positions.
26
+ - uses real basis functions computed via recurrence relations.
27
+ - uses "nested sum" O(p^4) M2M/M2L/L2L transformations.
28
+ - uses a non-uniform 2^N-ary tree hierarchy (directly inspired by [this work of A. Goude and S. Engblom](https://link.springer.com/article/10.1007/s11227-012-0836-0)), allowing arbitrary shape of the boxes in the hierarchy and guaranteeing balanced trees but requiring storage of interaction lists.
29
+ - has jit-compiled functions and autodiff for every substep of the algorithm except for the generation of interaction lists.
30
+
31
+ In summary, jaxFMM in its current state can do adaptive point charge FMM for Laplace kernels with good performance for lower expansion orders (p <= 3) and reasonably homogenous distributions. Autodiff only works if the particle positions remain constant. A first benchmark of uniformly distributed charges in the unit cube, computed on Google Cloud [g2-standard-8](https://cloud.google.com/compute/docs/gpus#l4-gpus) (GPU timings) and [c3d-highmem-16](https://cloud.google.com/compute/docs/general-purpose-machines#c3d_series) (CPU timings) machines can be found below:
32
+
33
+ <img src="docs/images/jax_unitcube_benchmark_p3.png" alt="unitcube benchmark" width="500"/>
34
+
35
+ ## TODOs
36
+
37
+ jaxFMM is primarily developed for my PhD project, where I am working on a GPU parallel FMM stray field evaluation routine with autodiff for finite-element micromagnetics. Alongside the very early state that it is in, this explains the currently limited feature set and design decisions mentioned above.
38
+
39
+ Contributions are always welcome however, and I plan to still improve jaxFMM. Topics that come to mind here are:
40
+
41
+ - volume FMM.
42
+ - other kernels or even a kernel-independent formulation.
43
+ - independent evaluation and source positions.
44
+ - distributed parallelism via jax.sharding.
45
+ - faster M2M/M2L/L2L transformations.
46
+ - faster jit-compilation, especially for large systems, higher orders and gradient computations.
47
+ - jit-compilable (and differentiable) interaction list generation, if possible.
48
+ - various other performance improvements.
49
+ - some degree of autotuning for the parameters.
@@ -0,0 +1,23 @@
1
+ import jax.numpy as jnp
2
+ from jax.scipy.optimize import minimize
3
+ from jaxfmm import *
4
+
5
+ print("Assembling grid of charges.")
6
+ nside, sidelen = 15, 4*jnp.pi
7
+ pts = (jnp.mgrid[:nside,:nside,:nside].T/(nside-1) * sidelen - sidelen/2).reshape((-1,3))
8
+ chrgs = 0.01*jnp.ones(pts.shape[0])
9
+
10
+ print("Generating hierarchy.")
11
+ tree_info = gen_hierarchy(pts)
12
+
13
+ print("Computing desired potential.")
14
+ desired_pot = jnp.sin(jnp.linalg.norm(pts,axis=-1))
15
+ norm = jnp.linalg.norm(desired_pot)
16
+
17
+ def loss(chrgs):
18
+ return jnp.linalg.norm(desired_pot-eval_potential(*tree_info,chrgs))/norm
19
+
20
+ print("Initial loss: %.2e"%loss(chrgs))
21
+ print("Compiling + running minimizer.")
22
+ res = minimize(loss,chrgs,method="BFGS",options={"maxiter": 1000})
23
+ print("Final loss (%i iterations): %.2e"%(res.nit,res.fun))
@@ -0,0 +1,20 @@
1
+ import jax.numpy as jnp
2
+ from jax import random
3
+ from jaxfmm import *
4
+
5
+ print("Generating sample points and charges.")
6
+ N = 2**15
7
+ key = random.key(124)
8
+ pts = random.uniform(key,(N,3))
9
+ chrgs = random.uniform(key,N)
10
+
11
+ print("Generating FMM hierarchy.")
12
+ tree_info = gen_hierarchy(pts)
13
+
14
+ print("Compiling and computing FMM potential.")
15
+ pot_FMM = eval_potential(*tree_info,chrgs)
16
+
17
+ print("Compiling and computing analytic potential.")
18
+ pot_dir = eval_potential_direct(pts,chrgs)
19
+
20
+ print("FMM normwise relative error: %.2e"%(jnp.linalg.norm(pot_dir-pot_FMM)/jnp.linalg.norm(pot_dir)))
@@ -0,0 +1,5 @@
1
+ from jaxfmm.fmm import *
2
+ from jaxfmm.debug_helpers import *
3
+
4
+ __all__ = (fmm.__all__ +
5
+ debug_helpers.__all__)
@@ -0,0 +1,134 @@
1
+ import matplotlib.pyplot as plt
2
+ from mpl_toolkits.mplot3d.art3d import Poly3DCollection
3
+ import jax.numpy as jnp
4
+ import jax.scipy as jsp
5
+ import jaxfmm.fmm as fmm
6
+ from math import isqrt
7
+
8
+ __all__ = ["print_stats", "plot_fmm_boxes"]
9
+
10
+ def print_stats(pts, idcs, rev_idcs, boxcenters, mpl_cnct, dir_cnct, n_split):
11
+ r"""
12
+ Print some basic information about the hierarchy.
13
+ """
14
+ print("----------------------------------------FMM Hierarchy Stats----------------------------------------")
15
+ print("%i points, %i levels, %i children per box, %i charges per box"%(pts.shape[0],len(mpl_cnct)-1, 2**n_split, idcs.shape[1]))
16
+ for i in range(len(mpl_cnct)):
17
+ print("Total number of mpl interactions on level %i (with padding, fraction: %.2f): %i"%(i, mpl_cnct[i].size / (mpl_cnct[i]<mpl_cnct[i].shape[0]).sum(),mpl_cnct[i].size))
18
+ print("Total number of dir interactions on max lvl (with padding, fraction: %.2f): %i"%(dir_cnct.size / (dir_cnct < dir_cnct.shape[0]).sum(),dir_cnct.size))
19
+ print("Near field compression ratio (without padding):", (((dir_cnct<dir_cnct.shape[0]).sum()) * idcs.shape[1]**2) / (float(pts.shape[0])**2))
20
+ memory_usage = idcs.nbytes + rev_idcs.nbytes + dir_cnct.nbytes
21
+ for i in range(len(boxcenters)):
22
+ memory_usage += boxcenters[i].nbytes + mpl_cnct[i].nbytes
23
+ print("Total memory consumed by the hierarchy: %.2e Bytes"%memory_usage)
24
+ print("---------------------------------------------------------------------------------------------------")
25
+
26
+ def plot_box(center, L, ax, facecolor='cyan', edgecolor='k', alpha=0.25):
27
+ r"""
28
+ Plot a single box in the given ax.
29
+ """
30
+ faces = []
31
+ for perm in range(3): # for the 3 coordinate axes
32
+ faces.append([])
33
+ faces.append([])
34
+ for id in [(-1,-1), (1,-1), (1,1), (-1,1)]: # we generate 4 points on each side
35
+ shift = jnp.array([-1,id[0],id[1]])
36
+ shift2 = jnp.array([1,id[0],id[1]])
37
+
38
+ start = center + L/2 * shift[(jnp.arange(3,dtype=int)+perm)%3]
39
+ end = center + L/2 * shift2[(jnp.arange(3,dtype=int)+perm)%3]
40
+
41
+ faces[-2].append(start)
42
+ faces[-1].append(end)
43
+ faces[-2].append(faces[-2][0])
44
+ faces[-1].append(faces[-1][0])
45
+ faces = jnp.array(faces)
46
+ ax.add_collection3d(Poly3DCollection(faces, facecolors=facecolor, linewidths=0.5, edgecolors=edgecolor, alpha=alpha))
47
+
48
+ def plot_fmm_boxes(pts, boxcenters, boxlens, level, show_wellsep=None, mpl_cnct=None, dir_cnct=None, n_split=3, elev=45, azim=45, roll=0, fname=None, plot_all=False):
49
+ r"""
50
+ Plot the hierarchy.
51
+ """
52
+ fig = plt.figure()
53
+ ax = fig.add_subplot(projection='3d')
54
+ ax.set_box_aspect((jnp.ptp(pts[:,0]), jnp.ptp(pts[:,1]), jnp.ptp(pts[:,2]))) # setting equal aspect ratio so we see the true shape of objects
55
+
56
+ nboxs = (2**n_split)**level
57
+ colors = plt.cm.jet(jnp.linspace(0,1,nboxs))
58
+ plot_all = plot_all or show_wellsep is None # if there is no show_wellsep set, we always plot everything
59
+ for i in range(nboxs):
60
+ if(show_wellsep is not None):
61
+ if(i==show_wellsep):
62
+ color = 'red'
63
+ elif(show_wellsep in mpl_cnct[level][i]):
64
+ color = 'green'
65
+ elif(level==(len(boxcenters)-1) and dir_cnct is not None and show_wellsep in dir_cnct[i]):
66
+ color = 'blue'
67
+ else:
68
+ color = 'gray'
69
+ else:
70
+ color = colors[i]
71
+ if(plot_all or color == 'blue' or color == 'red' or color=="green"):
72
+ plot_box(boxcenters[level][i],boxlens[level][i],ax,facecolor=color,alpha=0.5)
73
+
74
+ plt.axis('off')
75
+ ax.view_init(elev=elev, azim=azim, roll=roll)
76
+ plt.tight_layout()
77
+ plt.show()
78
+ if(fname is not None):
79
+ fig.savefig(fname,transparent=True,bbox_inches='tight',dpi=300)
80
+ return fig, ax
81
+
82
+ def eval_multipole(coeff, boxcenter, eval_pts):
83
+ r"""
84
+ Evaluate a single multipole expansion.
85
+ """
86
+ p = get_deg(coeff.shape[-1])
87
+ sing = fmm.eval_singular_basis(eval_pts - boxcenter,p)
88
+ res = jnp.zeros(eval_pts.shape[0])
89
+ for n in range(p+1):
90
+ for m in range(-n,n+1):
91
+ if(m!=0):
92
+ res += (-1)**n * 2 * coeff[...,fmm.mpl_idx(m,n)] * sing[...,fmm.mpl_idx(-m,n)]
93
+ else:
94
+ res += (-1)**n * coeff[...,fmm.mpl_idx(m,n)] * sing[...,fmm.mpl_idx(-m,n)]
95
+ res /= (4*jnp.pi)
96
+ return res
97
+
98
+ def get_local_expansions(pts, chrgs, exp_centers, p):
99
+ r"""
100
+ Generate local expansions.
101
+ """
102
+ dist = pts[None,...] - exp_centers[:,None,:]
103
+ coeff = (fmm.eval_singular_basis(dist,p) * chrgs[None,:,None]).sum(axis=1)
104
+ return coeff
105
+
106
+ def binom(x, y):
107
+ return jnp.exp(jsp.special.gammaln(x + 1) - jsp.special.gammaln(y + 1) - jsp.special.gammaln(x - y + 1))
108
+
109
+ def gen_multipole_dist(m, n, eps = 0.5):
110
+ r"""
111
+ Generate a point charge distribution corresponding to a specific multipole moment (Majic, Matt. (2022). Point charge representations of multipoles. European Journal of Physics. 43. 10.1088/1361-6404/ac578b.)
112
+ """
113
+ if(m == 0): # axial
114
+ k = jnp.arange(-n, n+1, 2)
115
+ chrgs = (-1)**((n-k)/2) * binom(n, (n-k)/2.0) / (jsp.special.factorial(n) * (2*eps)**n)
116
+ pts = jnp.zeros((k.shape[0],3))
117
+ pts = pts.at[:,2].set(k*eps)
118
+ else: # (stacked) bracelet
119
+ rotate = m < 0
120
+ m = abs(m) # we work with the real basis and rotate later
121
+ knum = n-m+1
122
+ jnum = 2*m
123
+ j = jnp.tile(jnp.arange(jnum),knum)
124
+ k = jnp.repeat(jnp.arange(-n+m,n-m+1,2),jnum)
125
+ phi = (j-0.5) * jnp.pi/m if rotate else j * jnp.pi/m
126
+ pts = jnp.array([eps*jnp.cos(phi), eps*jnp.sin(phi), k*eps]).T
127
+ chrgs = 4**(m-1) * jsp.special.factorial(m-1) / ((2*eps)**n * jsp.special.factorial(n-m)) * (-1)**((n-m-k)/2 + j) * binom(n-m,(n-m-k)/2)
128
+ return pts, chrgs
129
+
130
+ def get_deg(N_coeff):
131
+ r"""
132
+ Get the degree p of a multipole expansion from the number of coefficients.
133
+ """
134
+ return isqrt(N_coeff) - 1
@@ -0,0 +1,542 @@
1
+ import jax
2
+ import jax.numpy as jnp
3
+ from functools import partial
4
+ from math import log2, ceil, sqrt
5
+
6
+ __all__ = ["gen_hierarchy", "eval_potential", "eval_potential_direct"]
7
+
8
+ def get_max_l(N_tot, N_max, n_split = 3): # need to get this outside of the jit compiled function - if the particle number changes too much, we must recompile...
9
+ r"""
10
+ Compute number of levels in the hierarchy.
11
+
12
+ :param N_tot: Total number of point charges.
13
+ :type N_tot: int
14
+ :param N_max: Maximum allowed number of point charges per box.
15
+ :type N_max: int
16
+ :param n_split: How many splits per level and box will be performed. Each box will have 2^n_split children.
17
+ :type n_split: int, optional
18
+
19
+ :return: The maximum level in the hierarchy.
20
+ :rtype: int
21
+ """
22
+ max_l = int(ceil(log2(N_tot/N_max)/n_split))
23
+ return 0 if max_l < 0 else max_l
24
+
25
+ @partial(jax.jit, static_argnames = ["max_l", "n_split"])
26
+ def balanced_tree(pts, max_l, n_split = 3):
27
+ r"""
28
+ Generate a balanced 2^n-tree hierarchy.
29
+
30
+ :param pts: Array of shape (N_tot,3) containing the positions of N_tot point charges.
31
+ :type pts: jnp.array
32
+ :param max_l: Maximum level, computed with get_max_l().
33
+ :type max_l: int
34
+ :param n_split: How many splits per level and box will be performed. Each box will have 2^n_split children.
35
+ :type n_split: int, optional
36
+
37
+ :return: Indices to resort pts into the highest level of the hierarchy (includes padding), indices to reverse sorting, centers of the boxes on all levels, sidelengths of the boxes on all levels.
38
+ :rtype: (jnp.array, jnp.array, list(jnp.array), list(jnp.array))
39
+ """
40
+ n_chi = 2**n_split
41
+ idcs = jnp.arange(pts.shape[0], dtype = jnp.int32)[None,:]
42
+
43
+ for l in range(max_l*n_split): # carry out n_split splits on max_l levels in total (we cannot make this a for_i loop as the shape constantly changes)
44
+ splitpos = idcs.shape[1]//2 # split position in the middle
45
+ needpad = idcs.shape[1]%2 # modulo tells us if padding must be inserted
46
+
47
+ pts_sorted = pts.at[idcs].get(mode="fill",fill_value=-jnp.nan**2) # padded values get converted to NaNs - NOTE: due to the argpartition implementation, we need to make sure that all the NaNs have the same sign
48
+ axis_to_split = jnp.argmax(jnp.nanmax(pts_sorted,axis=1) - jnp.nanmin(pts_sorted,axis=1),axis=1) # nanmax and -min to ignore NaNs
49
+ idcs = idcs[jnp.arange(idcs.shape[0], dtype = jnp.int32)[:,None],jnp.argpartition(pts_sorted[jnp.arange(axis_to_split.shape[0], dtype = jnp.int32),:,axis_to_split],splitpos,axis=1)] # splitting - the NaNs introduced below get transported to the beginning
50
+
51
+ # padding so the next array has the correct shape - this does not break JIT compiling because it can be computed from only the input shape
52
+ idcs = jax.lax.pad(idcs,jnp.int32(pts.shape[0]),[(0,0,0),(0,needpad,0)]) # pad at the end with out of range values
53
+ idcs = idcs.reshape((-1,idcs.shape[1]//2)) # now that we padded, we can safely reshape this
54
+
55
+ idcs = jnp.sort(idcs,axis=1) # sorting is good for locality but might be overkill TODO: swap only first and last positions instead of full sort?
56
+ rev_idcs = jnp.argsort(idcs.flatten())[:pts.shape[0]] # reverse sorting indices, to undo the sorting
57
+ pts_sorted = pts.at[idcs].get(mode="fill",fill_value=jnp.nan)
58
+ boxcenters, boxlens = [jnp.zeros((n_chi**l,3)) for l in range(max_l+1)], [jnp.zeros((n_chi**l,3)) for l in range(max_l+1)]
59
+
60
+ for l in range(max_l,0,-1):
61
+ minc, maxc = jnp.nanmin(pts_sorted,axis=1), jnp.nanmax(pts_sorted,axis=1) # nanmax and -min to ignore NaNs
62
+ boxlens[l] = maxc - minc # TODO: in principle we only need to save the norm of this, but it is nice to have for visualizations
63
+ boxcenters[l] = minc + boxlens[l]/2
64
+ pts_sorted = pts_sorted.reshape((-1,pts_sorted.shape[1]*n_chi,3))
65
+
66
+ minc, maxc = jnp.nanmin(pts_sorted,axis=1), jnp.nanmax(pts_sorted,axis=1)
67
+ boxlens[0] = maxc - minc # box sidelengths
68
+ boxcenters[0] = minc + boxlens[0]/2 # box center
69
+ return idcs, rev_idcs, boxcenters, boxlens
70
+
71
+ #@partial(jax.jit, static_argnames = ["theta", "n_split"])
72
+ def gen_connectivity(boxcenters, boxlens, theta = 0.75, n_split = 3): # TODO: is there a way to make this JIT compilable?
73
+ r"""
74
+ Compute connectivity information for a given hierarchy. To determine well-separatedness, we check for
75
+
76
+ R + theta * r <= theta * d
77
+
78
+ where R = max(r1,r2), r = min(r1,r2) and d is the distance between the centers of two boxes with radii r1 and r2. This should give a FMM error that scales as
79
+
80
+ theta^(p+1)
81
+
82
+ with the expansion order p.
83
+
84
+ :param boxcenters: List of boxcenters computed with balanced_tree().
85
+ :type boxcenters: list(jnp.array)
86
+ :param boxlens: List of box sidelengths computed with balanced_tree().
87
+ :type boxlens: list(jnp.array)
88
+ :param theta: Well-separatedness parameter, determines accuracy.
89
+ :type theta: float
90
+ :param n_split: How many splits per level and box have been performed. Each box has 2^n_split children.
91
+ :type n_split: int, optional
92
+
93
+ :return: List of (padded) M2L interaction partner index arrays for each box on each level, array of (padded) P2P interaction partner indices on the highest level.
94
+ :rtype: (list(jnp.array), jnp.array)
95
+ """
96
+ n_l = len(boxcenters) # number of levels
97
+ n_chi = 2**n_split # number of child boxes per split
98
+
99
+ mpl_cnct = []
100
+ non_wellseps = jnp.array([[0]], dtype = jnp.int32) # initial value
101
+
102
+ for l in range(n_l): # for every level
103
+ r1 = jnp.linalg.norm(boxlens[l],axis=1)/2 # L/2 is the radius of the box TODO: see comment on only storing the norm above
104
+ d = jnp.linalg.norm(boxcenters[l][:,None,:] - boxcenters[l].at[non_wellseps].get(mode="fill",fill_value=jnp.nan),axis=-1) # find the distances between boxcenters
105
+ R = r1[non_wellseps] # r2
106
+ r = R.copy()
107
+ tmp = (r1[:,None] > R)
108
+ R = jnp.where(tmp,R,r1[:,None]) # wherever r2 < r1, we replace r2 with r1 => R = max(r1,r2)
109
+ r = jnp.where(~tmp,r,r1[:,None]) # wherever r2 > r1, we replace r2 with r1 => r = min(r1,r2)
110
+
111
+ wellsep = (R + theta*r <= theta*d) # NOTE: NaNs always return false here - this is how we rid ourselves of the padding
112
+ non_wellsep = (R + theta*r > theta*d)
113
+ wellsep_nums = wellsep.sum(axis=1, dtype = jnp.int32) # number of well-separated boxes for each box
114
+ non_wellsep_nums = non_wellsep.sum(axis=1, dtype = jnp.int32) # number of non-well-separated boxes for each box
115
+
116
+ to_pad_wellsep = jnp.max(wellsep_nums) - wellsep_nums # how much padding per box must be inserted
117
+ wellsep_padding = jnp.repeat(jnp.cumsum(wellsep_nums),to_pad_wellsep) # the correct indices for the insert below NOTE: this has a dynamic size and therefore breaks JIT compilation
118
+ mpl_cnct.append(jnp.insert(non_wellseps[wellsep],wellsep_padding,n_chi**l).reshape(n_chi**l,-1)) # save result to the interaction lists
119
+
120
+ to_pad_non_wellsep = jnp.max(non_wellsep_nums) - non_wellsep_nums # how much padding per box must be inserted for non-well-separated boxes
121
+ non_wellsep_padding = jnp.repeat(jnp.cumsum(non_wellsep_nums),to_pad_non_wellsep) # the correct indices for the insert below
122
+ non_wellseps = jnp.insert(non_wellseps[non_wellsep],non_wellsep_padding,n_chi**l).reshape((n_chi**l,-1)) # we overwrite the old values here
123
+
124
+ if(l<n_l-1): # prepare for the computation on the next level by transforming to child indices
125
+ non_wellseps = jnp.repeat(jnp.repeat(non_wellseps,n_chi,axis=1)*n_chi + jnp.tile(jnp.arange(n_chi, dtype = jnp.int32),non_wellseps.shape[1]),n_chi,axis=0)
126
+
127
+ return mpl_cnct, non_wellseps
128
+
129
+ def mpl_idx(m,n):
130
+ r"""
131
+ Compute "flattened" array position of multipole coefficient C^m_n with order m and degree n.
132
+ """
133
+ return n**2 + (m+n)
134
+
135
+ def inv_mpl_idx(idx):
136
+ r"""
137
+ Compute order m and degree n of "flattened" array position idx.
138
+ """
139
+ n = int(sqrt(idx))
140
+ m = idx - n*(n+1)
141
+ return m, n
142
+
143
+ @partial(jax.jit, static_argnames=['p'])
144
+ def eval_regular_basis(rvec, p):
145
+ r"""
146
+ Evaluate real regular basis functions (Laplace kernel) with a recursion relation [Gumerov, N. A. et al. Fast multipole methods on graphics processors. J. Comp. Phys., B 227, 8290 (2008)].
147
+
148
+ :param rvec: Array of positions in cartesian coordinates to evaluate the basis at.
149
+ :type rvec: jnp.array
150
+ :param p: Maximum degree (degree of multipole expansion) for the evaluation.
151
+ :type p: int
152
+
153
+ :return: Regular basis evaluated until the given degree at the given locations.
154
+ :rtype: jnp.array
155
+ """
156
+ x, y, z = rvec[..., 0], rvec[...,1], rvec[...,2]
157
+ coeff = jnp.zeros((*rvec.shape[:-1], (p+1)**2))
158
+ coeff = coeff.at[...,mpl_idx(0,0)].set(1)
159
+
160
+ if(p>0):
161
+ coeff = coeff.at[...,mpl_idx(1,1)].set(-0.5*x)
162
+ coeff = coeff.at[...,mpl_idx(-1,1)].set(0.5*y)
163
+
164
+ for n in range(2,p+1): # first recursion: extreme values
165
+ coeff = coeff.at[...,mpl_idx(n,n)].set(-(x*coeff[...,mpl_idx(n-1,n-1)] + y*coeff[...,mpl_idx(1-n,n-1)])/(2*n))
166
+ coeff = coeff.at[...,mpl_idx(-n,n)].set((y*coeff[...,mpl_idx(n-1,n-1)] - x*coeff[...,mpl_idx(1-n,n-1)])/(2*n))
167
+
168
+ for n in range(0,p): # second recursion: neighbors of extreme values
169
+ coeff = coeff.at[...,mpl_idx(n,n+1)].set(-z*coeff[...,mpl_idx(n,n)])
170
+ coeff = coeff.at[...,mpl_idx(-n,n+1)].set(-z*coeff[...,mpl_idx(-n,n)])
171
+
172
+ for n in range(2,p+1): # third recursion: all values inbetween
173
+ for m in range(-n+2,n-1):
174
+ coeff = coeff.at[...,mpl_idx(m,n)].set(-((2*n-1)*z*coeff[...,mpl_idx(m,n-1)] + (x**2+y**2+z**2)*coeff[...,mpl_idx(m,n-2)])/((n-abs(m))*(n+abs(m))))
175
+
176
+ return coeff
177
+
178
+ @partial(jax.jit, static_argnames=['p'])
179
+ def eval_singular_basis(rvec, p): # NOTE: might have to include a factor (-1)^n, also use S_n^-m for evaluating
180
+ r"""
181
+ Evaluate real singular basis functions (Laplace kernel) with a recursion relation.
182
+
183
+ :param rvec: Array of positions in cartesian coordinates to evaluate the basis at.
184
+ :type rvec: jnp.array
185
+ :param p: Maximum degree (degree of multipole expansion) for the evaluation.
186
+ :type p: int
187
+
188
+ :return: Singular basis evaluated until the given degree at the given locations.
189
+ :rtype: jnp.array
190
+ """
191
+ x, y, z = rvec[..., 0], rvec[...,1], rvec[...,2]
192
+ r2 = x**2 + y**2 + z**2
193
+ coeff = jnp.zeros((*rvec.shape[:-1], (p+1)**2))
194
+ coeff = coeff.at[...,mpl_idx(0,0)].set(1/jnp.sqrt(r2))
195
+
196
+ if(p>0):
197
+ coeff = coeff.at[...,mpl_idx(1,1)].set(-coeff[...,mpl_idx(0,0)]*y/r2)
198
+ coeff = coeff.at[...,mpl_idx(-1,1)].set(coeff[...,mpl_idx(0,0)]*x/r2)
199
+
200
+ for n in range(2,p+1): # first recursion: extreme values
201
+ coeff = coeff.at[...,mpl_idx(n,n)].set((2*n-1)*(x*coeff[...,mpl_idx(n-1,n-1)] - y*coeff[...,mpl_idx(1-n,n-1)])/r2)
202
+ coeff = coeff.at[...,mpl_idx(-n,n)].set((2*n-1)*(y*coeff[...,mpl_idx(n-1,n-1)] + x*coeff[...,mpl_idx(1-n,n-1)])/r2)
203
+
204
+ for n in range(0,p): # second recursion: neighbors of extreme values
205
+ coeff = coeff.at[...,mpl_idx(n,n+1)].set((2*n+1)*z*coeff[...,mpl_idx(n,n)]/r2)
206
+ coeff = coeff.at[...,mpl_idx(-n,n+1)].set((2*n+1)*z*coeff[...,mpl_idx(-n,n)]/r2)
207
+
208
+ for n in range(2,p+1): # third recursion: all values inbetween
209
+ for m in range(-n+2,n-1):
210
+ coeff = coeff.at[...,mpl_idx(m,n)].set(((2*n-1)*z*coeff[...,mpl_idx(m,n-1)] - (n-1-m)*(n-1+m)*coeff[...,mpl_idx(m,n-2)])/r2)
211
+
212
+ return coeff
213
+
214
+ @partial(jax.jit, static_argnames=['p'])
215
+ def get_initial_mpls(padded_pts, padded_chrgs, boxcenters, p):
216
+ r"""
217
+ Get initial multipole expansions for each box on the highest level.
218
+
219
+ :param padded_pts: Array of point positions, sorted into the highest level of the hierarchy via balanced_tree().
220
+ :type padded_pts: jnp.array
221
+ :param padded_chrgs: Array of point charges, sorted into the highest level of the hierarchy via balanced_tree().
222
+ :type padded_chrgs: jnp.array
223
+ :param boxcenters: Array of box centers, generated via balanced_tree().
224
+ :type boxcenters: jnp.array
225
+ :param p: Multipole expansion order.
226
+ :type p: int
227
+
228
+ :return: Array of multipole coefficients of the boxes on the highest level.
229
+ :rtype: jnp.array
230
+ """
231
+ dist = padded_pts - boxcenters[:,None]
232
+ return (eval_regular_basis(dist,p) * padded_chrgs[...,None]).sum(axis=1)
233
+
234
+ @partial(jax.jit, static_argnames=['p', 'n_split'])
235
+ def M2M(coeff, oldboxdims, newboxdims, p, n_split):
236
+ r"""
237
+ Multipole-to-multipole transformation, merging "small" multipole expansions on higher levels into "large" multipole expansions on lower levels. This is the O(p^4) algorithm proposed in the original 3D FMM paper.
238
+
239
+ :param coeff: Array of multipole coefficients on level l.
240
+ :type coeff: jnp.array
241
+ :param oldboxdims: Array of box centers on level l.
242
+ :type oldboxdims: jnp.array
243
+ :param newboxdims: Array of box centers on level l-1.
244
+ :type newboxdims: jnp.array
245
+ :param p: Multipole expansion order.
246
+ :type p: int
247
+ :param n_split: How many splits per level and box have been performed. Each box has 2^n_split children which are merged into one in this function.
248
+ :type n_split: int
249
+
250
+ :return: Array of multipole coefficients of the boxes on level l-1.
251
+ :rtype: jnp.array
252
+ """
253
+ n_chi = 2**n_split
254
+ mpls = coeff.reshape((coeff.shape[0]//n_chi,n_chi,coeff.shape[1]))
255
+ new_mpls = jnp.zeros((mpls.shape[0],mpls.shape[2]))
256
+ oldboxdims = oldboxdims.reshape((-1,n_chi,3))
257
+ reg = eval_regular_basis(oldboxdims - newboxdims[:,None,:],p) # shift direction points from target to source
258
+
259
+ for j in range(p+1):
260
+ for k in range(0,j+1): # Real coeffs
261
+ for n in range(j+1):
262
+ for m in range(max(k+n-j,-n),min(k+j-n,n)+1):
263
+ new_mpls = new_mpls.at[...,mpl_idx(k,j)].set(new_mpls[...,mpl_idx(k,j)] + (-1)**((abs(k)-abs(m)-abs(k-m))//2) * (
264
+ reg[...,mpl_idx(abs(m),n)]*mpls[...,mpl_idx(abs(k-m),j-n)] -
265
+ jnp.sign(m)*jnp.sign(k-m)*reg[...,mpl_idx(-abs(m),n)]*mpls[...,mpl_idx(-abs(k-m),j-n)]).sum(axis=-1))
266
+ for k in range(-j,0): # Imag coeffs
267
+ for n in range(j+1):
268
+ for m in range(max(k+n-j,-n),min(k+j-n,n)+1):
269
+ new_mpls = new_mpls.at[...,mpl_idx(k,j)].set( new_mpls[...,mpl_idx(k,j)] - (-1)**((abs(k)-abs(m)-abs(k-m))//2) * (
270
+ jnp.sign(k-m)*reg[...,mpl_idx(abs(m),n)]*mpls[...,mpl_idx(-abs(k-m),j-n)] +
271
+ jnp.sign(m)*reg[...,mpl_idx(-abs(m),n)]*mpls[...,mpl_idx(abs(k-m),j-n)]).sum(axis=-1))
272
+ return new_mpls
273
+
274
+ @partial(jax.jit, static_argnames=['p', 'n_split'])
275
+ def L2L(locs, oldboxdims, newboxdims, p, n_split):
276
+ r"""
277
+ Local-to-local transformation, distributing "large" local expansions on lower levels to "small" local expansions on higher levels. This is the O(p^4) algorithm proposed in the original 3D FMM paper.
278
+
279
+ :param locs: Array of local coefficients on level l.
280
+ :type locs: jnp.array
281
+ :param oldboxdims: Array of box centers on level l.
282
+ :type oldboxdims: jnp.array
283
+ :param newboxdims: Array of box centers on level l+1.
284
+ :type newboxdims: jnp.array
285
+ :param p: Multipole expansion order.
286
+ :type p: int
287
+ :param n_split: How many splits per level and box have been performed. Each box has 2^n_split children which all receive a local expansion of their parent in this function.
288
+ :type n_split: int
289
+
290
+ :return: Array of local coefficients of the boxes on level l+1.
291
+ :rtype: jnp.array
292
+ """
293
+ n_chi = 2**n_split
294
+ new_locs = jnp.zeros((locs.shape[0], n_chi, locs.shape[1]))
295
+ newboxdims = newboxdims.reshape((-1,n_chi,3))
296
+ reg = eval_regular_basis(oldboxdims[:,None,:]-newboxdims,p) # shift direction points from target to source
297
+
298
+ for j in range(p+1):
299
+ for k in range(1,j+1): # -Imag coeffs!
300
+ for n in range(j,p+1):
301
+ for m in range(k+j-n,k+n-j+1):
302
+ new_locs = new_locs.at[...,mpl_idx(k,j)].set(new_locs[...,mpl_idx(k,j)] - (-1)**((abs(m)-abs(m-k)-abs(k))//2) * (
303
+ -jnp.sign(m)*reg[...,mpl_idx(abs(m-k),n-j)]*locs[...,None,mpl_idx(abs(m),n)] +
304
+ jnp.sign(m-k)*reg[...,mpl_idx(-abs(m-k),n-j)]*locs[...,None,mpl_idx(-abs(m),n)]))
305
+ for k in range(-j,1): # Real coeffs!
306
+ for n in range(j,p+1):
307
+ for m in range(k+j-n,k+n-j+1):
308
+ new_locs = new_locs.at[...,mpl_idx(k,j)].set(new_locs[...,mpl_idx(k,j)] + (-1)**((abs(m)-abs(m-k)-abs(k))//2) * (
309
+ reg[...,mpl_idx(abs(m-k),n-j)]*locs[...,None,mpl_idx(-abs(m),n)] +
310
+ jnp.sign(m-k)*jnp.sign(m)*reg[...,mpl_idx(-abs(m-k),n-j)]*locs[...,None,mpl_idx(abs(m),n)]))
311
+ return new_locs.reshape((-1,new_locs.shape[-1]))
312
+
313
+ @partial(jax.jit, static_argnames=['p'])
314
+ def M2L(mpls, locs, boxcenters, mpl_cnct, p): # TODO: try to write this as a jax.pallas kernel?
315
+ r"""
316
+ Multipole-to-local transformation, turning multipole expansions into local expansions on the same level. This is the O(p^4) algorithm proposed in the original 3D FMM paper.
317
+
318
+ :param mpls: Array of multipole coefficients.
319
+ :type mpls: jnp.array
320
+ :param locs: Array of local coefficients.
321
+ :type locs: jnp.array
322
+ :param boxcenters: Array of box centers.
323
+ :type boxcenters: jnp.array
324
+ :param mpl_cnct: Array containing interaction partner indices for each box.
325
+ :type mpl_cnct: jnp.array
326
+ :param p: Multipole expansion order.
327
+ :type p: int
328
+
329
+ :return: Updated array of local coefficients, now includes all the well-separated multipole expansions on this level.
330
+ :rtype: jnp.array
331
+ """
332
+ sing = eval_singular_basis(boxcenters.at[mpl_cnct].get(mode="fill",fill_value=123456789)-boxcenters[:,None,:],p) # shift direction points from target to source TODO: find a proper way of getting non-nan values here
333
+
334
+ for j in range(p+1):
335
+ for k in range(1,j+1): # -Imag coeffs!
336
+ for n in range(p+1-j):
337
+ for m in range(max(k-j-n,-n),min(j+n+k,n)+1):
338
+ locs = locs.at[:,mpl_idx(k,j)].set( locs[:,mpl_idx(k,j)] + (-1)**((abs(k-m)-abs(k)-abs(m))//2) * (
339
+ -jnp.sign(m-k)*sing[...,mpl_idx(abs(m-k),j+n)]*mpls.at[mpl_cnct,mpl_idx(abs(m),n)].get(mode="fill",fill_value=0) +
340
+ jnp.sign(m)*sing[...,mpl_idx(-abs(m-k),j+n)]*mpls.at[mpl_cnct,mpl_idx(-abs(m),n)].get(mode="fill",fill_value=0)).sum(axis=1))
341
+ for k in range(-j,1): # Real coeffs!
342
+ for n in range(p+1-j):
343
+ for m in range(max(k-j-n,-n),min(j+n+k,n)+1):
344
+ locs = locs.at[:,mpl_idx(k,j)].set(locs[:,mpl_idx(k,j)] + (-1)**((abs(k-m)-abs(k)-abs(m))//2) * (
345
+ sing[...,mpl_idx(-abs(m-k),j+n)]*mpls.at[mpl_cnct,mpl_idx(abs(m),n)].get(mode="fill",fill_value=0) +
346
+ jnp.sign(m)* jnp.sign(m-k) * sing[...,mpl_idx(abs(m-k),j+n)]*mpls.at[mpl_cnct,mpl_idx(-abs(m),n)].get(mode="fill",fill_value=0)).sum(axis=1))
347
+ return locs
348
+
349
+ @partial(jax.jit, static_argnames=['p', 'n_split'])
350
+ def go_up(coeff, boxcenters, p, n_split):
351
+ r"""
352
+ Using multipole-to-multipole transformation, descend the hierarchy.
353
+
354
+ :param coeff: Array of multipole coefficients on the highest level.
355
+ :type coeff: jnp.array
356
+ :param boxcenters: List containing box center arrays for every level.
357
+ :type boxcenters: list(jnp.array)
358
+ :param p: Multipole expansion order.
359
+ :type p: int
360
+ :param n_split: How many splits per level and box have been performed. Each box has 2^n_split children which are merged on every level for this step.
361
+ :type n_split: int
362
+
363
+ :return: List of multipole coefficient arrays, containing multipole expansion coefficients for every box on every level.
364
+ :rtype: list(jnp.array)
365
+ """
366
+ mpls = [coeff]
367
+ n_l = len(boxcenters)
368
+ for i in range(n_l-1):
369
+ mpls.append(M2M(mpls[-1],boxcenters[-(i+1)],boxcenters[-(i+2)], p, n_split))
370
+ return mpls
371
+
372
+ @partial(jax.jit, static_argnames=['p', 'n_split'])
373
+ def go_down(mpls, boxcenters, mpl_cnct, p, n_split):
374
+ r"""
375
+ Using multipole-to-local and local-to-local transformation, ascend the hierarchy.
376
+
377
+ :param mpls: List of multipole coefficient arrays on every level.
378
+ :type mpls: list(jnp.array)
379
+ :param boxcenters: List containing box center arrays for every level.
380
+ :type boxcenters: list(jnp.array)
381
+ :param mpl_cnct: List of interaction partner index arrays on every level.
382
+ :type mpl_cnct: list(jnp.array)
383
+ :param p: Multipole expansion order.
384
+ :type p: int
385
+ :param n_split: How many splits per level and box have been performed. Each box has 2^n_split children which receive local expansions from their parent in this step.
386
+ :type n_split: int
387
+
388
+ :return: List of local coefficient arrays, containing local expansion coefficients for every box on every level.
389
+ :rtype: list(jnp.array)
390
+ """
391
+ locs = [jnp.zeros((mpls[-1].shape))]
392
+ n_l = len(boxcenters)
393
+ for i in range(n_l-1): # for every level TODO: level 1 can be skipped
394
+ locs.append(L2L(locs[-1],boxcenters[i],boxcenters[i+1], p, n_split)) # get local expansions from parent boxes
395
+ locs[-1] = M2L(mpls[-(i+2)], locs[-1], boxcenters[i+1], mpl_cnct[i+1], p) # next, it is time to do the M2L shifts for this level
396
+ return locs
397
+
398
+ @partial(jax.jit, static_argnames=['p'])
399
+ def eval_local(locs, padded_pts, rev_idcs, boxcenters, p):
400
+ r"""
401
+ Evaluate local expansions on the highest level.
402
+
403
+ :param locs: Array of local coefficients on the highest level.
404
+ :type locs: jnp.array
405
+ :param padded_pts: Array containing point positions, sorted into the highest level of the hierarchy.
406
+ :type padded_pts: jnp.array
407
+ :param rev_idcs: Index array to remove padding and return to original sorting.
408
+ :type rev_idcs: jnp.array
409
+ :param boxcenters: Array of box center positions on the highest level.
410
+ :type boxcenters: jnp.array
411
+ :param p: Multipole expansion order.
412
+ :type p: int
413
+
414
+ :return: Far-field potential in the original sorting positions.
415
+ :rtype: jnp.array
416
+ """
417
+ padded_res = jnp.zeros(padded_pts.shape[:2])
418
+ reg = eval_regular_basis(padded_pts-boxcenters[:,None],p) # this evaluation needs to be relative to the boxcenter
419
+
420
+ for n in range(p+1):
421
+ for m in range(-n,n+1):
422
+ if(m!=0):
423
+ padded_res += (-1)**n * 2*locs[:,None,mpl_idx(-m,n)] * reg[...,mpl_idx(m,n)]
424
+ else:
425
+ padded_res += (-1)**n * locs[:,None,mpl_idx(-m,n)] * reg[...,mpl_idx(m,n)]
426
+ return padded_res.flatten()[rev_idcs] / (4*jnp.pi) # TODO: flatten and resort only once, after adding?
427
+
428
+ @jax.jit
429
+ def eval_direct(padded_pts, padded_chrgs, rev_idcs, direct_cnct): # TODO: replace some of these evaluations with multipole evaluations for better performance
430
+ r"""
431
+ Evaluate the near-field potential directly (P2P).
432
+
433
+ :param padded_pts: Array containing point positions, sorted into the highest level of the hierarchy.
434
+ :type padded_pts: jnp.array
435
+ :param padded_chrgs: Array containing point charges, sorted into the highest level of the hierarchy.
436
+ :type padded_chrgs: jnp.array
437
+ :param rev_idcs: Index array to remove padding and return to original sorting.
438
+ :type rev_idcs: jnp.array
439
+ :param direct_cnct: Index array of interaction partners for each box on the highest level.
440
+ :type direct_cnct: jnp.array
441
+
442
+ :return: Near-field potential in the original sorting positions.
443
+ :rtype: jnp.array
444
+ """
445
+ n = padded_pts.shape[1]
446
+ nrows = padded_pts.shape[0]
447
+
448
+ def i_body(i):
449
+ acc = jnp.zeros((n,))
450
+ xblk = padded_pts[i]
451
+ def k_body(k, acc):
452
+ partner = direct_cnct[i,k]
453
+ dists = jnp.linalg.norm(xblk[:,None] - padded_pts.at[partner].get(mode="fill",fill_value=jnp.inf),axis=-1)
454
+ dists = 1/jnp.where(dists==0,jnp.inf,dists)
455
+ chunk = padded_chrgs[partner]
456
+ return acc + dists.dot(chunk)
457
+ acc = jax.lax.fori_loop(0,direct_cnct.shape[1], k_body, acc)
458
+ return acc
459
+
460
+ accs = jax.vmap(i_body)(jnp.arange(nrows))
461
+ return accs.flatten()[rev_idcs]/(4*jnp.pi) # TODO: flatten and resort only once, after adding?
462
+
463
+ def gen_hierarchy(pts, N_max = 128, theta = 0.75, n_split = 3):
464
+ r"""
465
+ Generate the balanced tree and connectivity for the FMM.
466
+
467
+ :param pts: Array of shape (N,3) containing the positions of N point charges.
468
+ :type pts: jnp.array
469
+ :param N_max: Maximum allowed number of point charges per box.
470
+ :type N_max: int, optional
471
+ :param theta: Well-separatedness parameter, determines accuracy.
472
+ :type theta: float, optional
473
+ :param n_split: How many splits per level and box are been performed. Each box has 2^n_split children.
474
+ :type n_split: int, optional
475
+
476
+ :return: Tuple containing full hierarchy information.
477
+ :rtype: tuple
478
+ """
479
+ max_l = get_max_l(pts.shape[0], N_max, n_split)
480
+ idcs, rev_idcs, boxcenters, boxlens = balanced_tree(pts, max_l, n_split)
481
+ mpl_cnct, dir_cnct = gen_connectivity(boxcenters, boxlens, theta, n_split)
482
+ return pts, idcs, rev_idcs, boxcenters, mpl_cnct, dir_cnct, n_split
483
+
484
+ @partial(jax.jit, static_argnames=['p', 'n_split'])#,backend="cpu")
485
+ def eval_potential(pts, idcs, rev_idcs, boxcenters, mpl_cnct, direct_cnct, n_split, chrgs, p = 3):
486
+ r"""
487
+ Full FMM potential evaluation, does not include creation of the hierarchy (see gen_hierarchy(), which generates the first 7 parameters). Therefore, only chrgs and p can be safely changed without recomputing the hierarchy.
488
+
489
+ :param pts: Array containing point positions.
490
+ :type pts: jnp.array
491
+ :param idcs: Index array to sort points/charges into highest level of the hierarchy.
492
+ :type idcs: jnp.array
493
+ :param rev_idcs: Index array to remove padding and return to original sorting.
494
+ :type rev_idcs: jnp.array
495
+ :param boxcenters: List containing box center arrays for every level.
496
+ :type boxcenters: list(jnp.array)
497
+ :param mpl_cnct: List of interaction partner index arrays on every level.
498
+ :type mpl_cnct: list(jnp.array)
499
+ :param direct_cnct: Index array of interaction partners for each box on the highest level.
500
+ :type direct_cnct: jnp.array
501
+ :param n_split: How many splits per level and box have been performed. Each box has 2^n_split children.
502
+ :type n_split: int
503
+ :param chrgs: Array containing point charges.
504
+ :type chrgs: jnp.array
505
+ :param p: Multipole expansion order.
506
+ :type p: int, optional
507
+
508
+ :return: Electrostatic potential of the points and corresponding charges.
509
+ :rtype: jnp.array
510
+ """
511
+ padded_pts = pts.at[idcs].get(mode="fill",fill_value=0.0) # TODO: we could buffer this alternatively
512
+ padded_chrgs = chrgs.at[idcs].get(mode="fill",fill_value=0.0)
513
+ coeff = get_initial_mpls(padded_pts, padded_chrgs, boxcenters[-1], p)
514
+ mpls = go_up(coeff, boxcenters, p, n_split)
515
+ locs = go_down(mpls, boxcenters, mpl_cnct, p, n_split)
516
+ pot = eval_direct(padded_pts, padded_chrgs, rev_idcs, direct_cnct) + eval_local(locs[-1], padded_pts, rev_idcs, boxcenters[-1], p)
517
+ return pot
518
+
519
+ @jax.jit
520
+ def eval_potential_direct(pts, chrgs, eval_pts = None):
521
+ r"""
522
+ Evaluate the potential directly via pairwise sums.
523
+
524
+ :param pts: Array containing point positions.
525
+ :type padded_pts: jnp.array
526
+ :param chrgs: Array containing point charges.
527
+ :type chrgs: jnp.array
528
+ :param eval_pts: Array containing points to evaluate the potential at. Defaults to pts.
529
+ :type eval_pts: jnp.array, optional
530
+
531
+ :return: Electrostatic potential of the points and corresponding charges.
532
+ :rtype: jnp.array
533
+ """
534
+ if(eval_pts is None):
535
+ eval_pts = pts
536
+ res = jnp.zeros(eval_pts.shape[0])
537
+ def eval_direct_body(i, val):
538
+ inv_dists = jnp.linalg.norm(pts[:,:] - eval_pts[i,None,:],axis=-1)
539
+ inv_dists = 1/jnp.where(inv_dists==0,jnp.inf,inv_dists) # take out self-interaction
540
+ val = val.at[i].set((chrgs * inv_dists).sum())
541
+ return val
542
+ return jax.lax.fori_loop(0,eval_pts.shape[0],eval_direct_body,res)/(4*jnp.pi)
@@ -0,0 +1,44 @@
1
+ [build-system]
2
+ requires = ["hatchling"]
3
+ build-backend = "hatchling.build"
4
+
5
+ [project]
6
+ name = "jaxFMM"
7
+ version = "0.0.1"
8
+ description = "Adaptive Fast Multipole Method with Laplace kernel in JAX."
9
+ readme = "README.md"
10
+ requires-python = ">=3.9"
11
+ license = { file = "LICENSE" }
12
+ keywords = ["jax", "FMM", "potential", "treecode", "N-body"]
13
+ authors = [{ name = "Robert Kraft", email = "robert.kraft@univie.ac.at" }]
14
+ # maintainers = [
15
+ # { name = "A. Great Maintainer", email = "maintainer@example.com" },
16
+ # ]
17
+ classifiers = [
18
+ "Development Status :: 3 - Alpha", # 4 - Beta, 5 - Production
19
+ "Intended Audience :: Science/Research",
20
+ "Topic :: Scientific/Engineering :: Physics",
21
+ "License :: OSI Approved :: GNU General Public License v3 (GPLv3)",
22
+ "Programming Language :: Python :: 3",
23
+ "Programming Language :: Python :: 3.9",
24
+ "Programming Language :: Python :: 3.10",
25
+ "Programming Language :: Python :: 3.11",
26
+ "Programming Language :: Python :: 3.12",
27
+ "Programming Language :: Python :: 3.13",
28
+ "Programming Language :: Python :: 3 :: Only",
29
+ ]
30
+ dependencies = ["jax"]
31
+
32
+ [project.optional-dependencies]
33
+ cuda = ["jax[cuda]"]
34
+ dev = ["pytest","matplotlib"]
35
+
36
+ [project.urls]
37
+ "Homepage" = "https://gitlab.com/jaxfmm/jaxfmm"
38
+
39
+ [tool.pytest.ini_options]
40
+ minversion = "6.0"
41
+ # addopts = "-ra -q"
42
+ testpaths = [
43
+ "tests"
44
+ ]
@@ -0,0 +1,34 @@
1
+ import jax.numpy as jnp
2
+ import jaxfmm.fmm as fmm
3
+ import jaxfmm.debug_helpers as debug
4
+ from jax import random
5
+
6
+ ### TODO:
7
+ # - find similar tests for local coeffs and expansions
8
+
9
+ def test_mpl_coeffs(): # TODO: also check for correct normalization
10
+ for n in range(9):
11
+ for m in range(-n,n+1):
12
+ pts, chrgs = debug.gen_multipole_dist(m,n,eps=10.0) # special point charge distribution corresponding to multipole moments - set eps large to minimize error
13
+ tree_info = fmm.gen_hierarchy(pts)
14
+ coeff = fmm.get_initial_mpls(pts,chrgs,tree_info[3][0],n)[0,:]
15
+ test = jnp.where(jnp.abs(coeff)>1e-5)[0] # only the (m,n) coefficient should be nonzero
16
+ assert test.shape[0] == 1
17
+ assert test[0] == fmm.mpl_idx(m,n)
18
+
19
+ def test_mpl_eval(): # TODO: do a better test, maybe independent from computing coeffs (related to the above todo)
20
+ key = random.key(743)
21
+ pts = random.uniform(key,(128,3),minval=-1,maxval=1)
22
+ chrgs = random.uniform(key,128,minval=-1,maxval=1)
23
+
24
+ tree_info = fmm.gen_hierarchy(pts)
25
+ coeff = fmm.get_initial_mpls(pts,chrgs,tree_info[3][0],10)[0,:]
26
+
27
+ nside, sidelen = 100, 10
28
+ eval_pts = (jnp.mgrid[:nside+1,:nside+1,:nside+1].T/nside * sidelen - sidelen/2).reshape((-1,3))
29
+ eval_pts = eval_pts[jnp.linalg.norm(eval_pts,axis=-1)>2*jnp.sqrt(3)]
30
+
31
+ pot_fmm = debug.eval_multipole(coeff, tree_info[3][0], eval_pts)
32
+ pot_dir = fmm.eval_potential_direct(pts,chrgs,eval_pts)
33
+ max_err = jnp.max(jnp.abs(pot_fmm-pot_dir))
34
+ assert max_err < 5.75e-6
@@ -0,0 +1,20 @@
1
+ import jax.numpy as jnp
2
+ from jaxfmm import *
3
+ from jax import random
4
+
5
+ ### TODO:
6
+ # - check if the FMM error scaling roughly works out
7
+
8
+ def test_potential_unitcube():
9
+ N = 2**15
10
+ key = random.key(856)
11
+ pts = random.uniform(key,(N,3))
12
+ chrgs = random.uniform(key,N,minval=-1,maxval=1)
13
+
14
+ tree_info = gen_hierarchy(pts)
15
+ pot_FMM = eval_potential(*tree_info,chrgs)
16
+
17
+ pot_dir = eval_potential_direct(pts,chrgs)
18
+
19
+ err = jnp.linalg.norm(pot_dir-pot_FMM)/jnp.linalg.norm(pot_dir)
20
+ assert err < 3.75e-3
@@ -0,0 +1,21 @@
1
+ import jax.numpy as jnp
2
+ import jaxfmm.fmm as fmm
3
+ from jax import random
4
+
5
+ ### TODO:
6
+ # - find similar tests for L2L, M2L
7
+
8
+ def test_M2M():
9
+ key = random.key(825)
10
+ pts = random.uniform(key,(8*128,3),minval=-1,maxval=1)
11
+ chrgs = random.uniform(key,8*128,minval=-1,maxval=1)
12
+
13
+ tree_info = fmm.gen_hierarchy(pts)
14
+ padded_pts = pts.at[tree_info[1]].get(mode="fill",fill_value=0.0)
15
+ padded_chrgs = chrgs.at[tree_info[1]].get(mode="fill",fill_value=0.0)
16
+
17
+ coeff = fmm.get_initial_mpls(padded_pts,padded_chrgs,tree_info[3][1],10)
18
+ coeff_merged = fmm.M2M(coeff,tree_info[3][1], tree_info[3][0], 10, 3)[0,:]
19
+ coeff_dir = fmm.get_initial_mpls(pts,chrgs,tree_info[3][0],10)[0,:]
20
+
21
+ assert jnp.allclose(coeff_merged, coeff_dir, rtol=1e-3, atol=1e-6)
@@ -0,0 +1,24 @@
1
+ import jax.numpy as jnp
2
+ import jaxfmm.fmm as fmm
3
+
4
+ ### TODO:
5
+ # - find better tests
6
+
7
+ def test_tree_connectivity_unitcube():
8
+ nside, sidelen = 64, 10.0
9
+ pts = (jnp.mgrid[:nside,:nside,:nside].T/(nside-1) * sidelen - sidelen/2).reshape((-1,3))
10
+
11
+ max_l = fmm.get_max_l(pts.shape[0], 128)
12
+ assert max_l==4
13
+
14
+ idcs, rev_idcs, boxcenters, boxlens = fmm.balanced_tree(pts,max_l)
15
+ assert idcs.shape[1]==64
16
+ assert jnp.all(idcs < pts.shape[0])
17
+ assert jnp.allclose(pts, pts[idcs].reshape((-1,3))[rev_idcs])
18
+ for i in range(len(boxlens)):
19
+ assert jnp.allclose(boxlens[i], pts[nside//(2**i)-1,0] - pts[0,0])
20
+
21
+ mpl_cnct, dir_cnct = fmm.gen_connectivity(boxcenters, boxlens)
22
+
23
+ assert (mpl_cnct[-1]<mpl_cnct[-1].shape[0]).sum() == 611136
24
+ assert (dir_cnct<dir_cnct.shape[0]).sum() == 70336