jaxFMM 0.0.1__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- jaxfmm/__init__.py +5 -0
- jaxfmm/debug_helpers.py +134 -0
- jaxfmm/fmm.py +542 -0
- jaxfmm-0.0.1.dist-info/METADATA +98 -0
- jaxfmm-0.0.1.dist-info/RECORD +7 -0
- jaxfmm-0.0.1.dist-info/WHEEL +4 -0
- jaxfmm-0.0.1.dist-info/licenses/LICENSE +21 -0
jaxfmm/__init__.py
ADDED
jaxfmm/debug_helpers.py
ADDED
|
@@ -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
|
jaxfmm/fmm.py
ADDED
|
@@ -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,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.
|
|
@@ -0,0 +1,7 @@
|
|
|
1
|
+
jaxfmm/__init__.py,sha256=sU4Oilae3k66sTlncV2Ihn1f03pbGr8T2i_g17WO4SQ,119
|
|
2
|
+
jaxfmm/debug_helpers.py,sha256=4_9PgyCiqBZb2Ex5Y_ONT2a9LcoMoYOVFSVLuItoDXo,6005
|
|
3
|
+
jaxfmm/fmm.py,sha256=aDMmbUFqqqCPyMFEZFaf7aSOPZ8b8vhUzuDgSDgKAyo,28604
|
|
4
|
+
jaxfmm-0.0.1.dist-info/METADATA,sha256=flDcq2FmGsQ2uWARNhrL8btmadCei-lWbUm_PByf4r4,5524
|
|
5
|
+
jaxfmm-0.0.1.dist-info/WHEEL,sha256=qtCwoSJWgHk21S1Kb4ihdzI2rlJ1ZKaIurTj_ngOhyQ,87
|
|
6
|
+
jaxfmm-0.0.1.dist-info/licenses/LICENSE,sha256=vr-EfJ9VOmC_zudSft97RP5lNIl3YsBwBGZxw2MB2Xk,1069
|
|
7
|
+
jaxfmm-0.0.1.dist-info/RECORD,,
|
|
@@ -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.
|