Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
21 changes: 6 additions & 15 deletions jaccpot/runtime/_interaction_cache.py
Original file line number Diff line number Diff line change
Expand Up @@ -2075,22 +2075,13 @@ def _interaction_cache_key(
hasher.update(str(topology_key).encode("utf8"))
else:
try:
morton_codes = np.asarray(
jax.device_get(tree.morton_codes),
dtype=np.uint64,
)
node_ranges = np.asarray(
jax.device_get(tree.node_ranges),
dtype=np.int64,
)
bounds_min = np.asarray(
jax.device_get(tree.bounds_min),
dtype=np.float64,
)
bounds_max = np.asarray(
jax.device_get(tree.bounds_max),
dtype=np.float64,
morton_codes, node_ranges, bounds_min, bounds_max = jax.device_get(
(tree.morton_codes, tree.node_ranges, tree.bounds_min, tree.bounds_max)
)
morton_codes = np.asarray(morton_codes, dtype=np.uint64)
node_ranges = np.asarray(node_ranges, dtype=np.int64)
bounds_min = np.asarray(bounds_min, dtype=np.float64)
bounds_max = np.asarray(bounds_max, dtype=np.float64)
except Exception:
return None

Expand Down
4 changes: 2 additions & 2 deletions jaccpot/upward/solidfmm_complex_tree_expansions.py
Original file line number Diff line number Diff line change
Expand Up @@ -21,7 +21,7 @@
from yggdrax.dtypes import INDEX_DTYPE, as_index, complex_dtype_for_real
from yggdrax.geometry import TreeGeometry
from yggdrax.tree import Tree, get_level_offsets, get_nodes_by_level
from yggdrax.tree_moments import TreeMassMoments, compute_tree_mass_moments
from yggdrax.tree_moments import TreeMassMoments, compute_tree_mass_moments_jit

from jaccpot._env import env_flag
from jaccpot.operators.complex_harmonics import p2m_complex_batch
Expand Down Expand Up @@ -846,7 +846,7 @@ def _record_stage(name: str, start: float, value) -> None:
_record_stage("geometry", stage_t0, geometry)
_upward_diag("geometry done")
stage_t0 = time.perf_counter()
mass_moments = compute_tree_mass_moments(
mass_moments = compute_tree_mass_moments_jit(
tree,
positions_sorted,
masses_sorted,
Expand Down
Loading