From 288476978155cdc1f8ef93151cace54859b48482 Mon Sep 17 00:00:00 2001 From: Tommaso Tarchi Date: Thu, 3 Sep 2026 13:23:29 +0200 Subject: [PATCH 1/2] first optimizations --- examples/benchmark_gpu_n_ladder_production.py | 7 +++ examples/benchmark_gpu_radix_worker.py | 44 ++++++++++++------- examples/benchmark_utils.py | 16 +++++-- jaccpot/runtime/_interaction_cache.py | 21 +++------ .../solidfmm_complex_tree_expansions.py | 9 +++- 5 files changed, 60 insertions(+), 37 deletions(-) diff --git a/examples/benchmark_gpu_n_ladder_production.py b/examples/benchmark_gpu_n_ladder_production.py index e444ac16..ec592b2a 100644 --- a/examples/benchmark_gpu_n_ladder_production.py +++ b/examples/benchmark_gpu_n_ladder_production.py @@ -496,6 +496,12 @@ def run_worker_case( "--config-json", json.dumps(payload), ] + trace_dir = ( + OUTPUT_DIR + / "jax_traces" + / f"N{num_particles}_{_traversal_cfg_label(traversal_cfg)}_m2l{m2l_chunk_size}_nf{nearfield_edge_chunk_size}" + ) + cmd += ["--jax-trace-dir", str(trace_dir)] proc = subprocess.Popen( cmd, stdout=subprocess.PIPE, @@ -800,6 +806,7 @@ def main() -> None: print(f"Starting N-ladder production sweep run_id={ladder_run_id}") sweep_rows: list[dict[str, Any]] = [] + print(f"-> OUTPUT_DIR is {OUTPUT_DIR}") for num_particles in ARGS.particle_counts: guidance = traversal_guidance_for_num_particles(int(num_particles)) preferred_label = _traversal_cfg_label(guidance["recommended"]) diff --git a/examples/benchmark_gpu_radix_worker.py b/examples/benchmark_gpu_radix_worker.py index 416edfe5..e3a9dfbc 100644 --- a/examples/benchmark_gpu_radix_worker.py +++ b/examples/benchmark_gpu_radix_worker.py @@ -54,6 +54,7 @@ def _configure_worker_environment() -> None: import jax import jax.numpy as jnp +import jax.profiler REPO_ROOT = pathlib.Path(__file__).resolve().parents[1] if str(REPO_ROOT) not in sys.path: @@ -210,6 +211,7 @@ def _parse_args() -> argparse.Namespace: parser.add_argument("--autotune-cache", default=None) parser.add_argument("--emit-ready-marker", action="store_true") parser.add_argument("--config-json", required=True) + parser.add_argument("--jax-trace-dir", default=None) return parser.parse_args() @@ -1906,6 +1908,7 @@ def _run_sweep_case( cfg: dict[str, Any], fmm_kwargs: dict[str, Any], autotune_cache_path: Optional[str] = None, + jax_trace_dir: Optional[str] = None, ) -> dict[str, Any]: jaccpot_symbols = _load_jaccpot_symbols() FastMultipoleMethod = jaccpot_symbols["FastMultipoleMethod"] @@ -1943,23 +1946,31 @@ def _run_sweep_case( ) _emit_ready_marker() - prepare_once_timing = bench_utils.time_callable( - fmm.prepare_state, - positions, - masses, - leaf_size=int(leaf_size), - max_order=int(max_order), - warmup=int(warmup), - runs=int(runs), - ) - prepared_state = prepare_once_timing.result - eval_timing = bench_utils.time_callable( - fmm.evaluate_prepared_state, - prepared_state, - warmup=int(warmup), - runs=int(runs), - **_evaluate_prepared_kwargs(fmm), + trace_ctx = ( + jax.profiler.trace(jax_trace_dir, create_perfetto_link=False) + if jax_trace_dir + else contextlib.nullcontext() ) + with trace_ctx: + prepare_once_timing = bench_utils.time_callable( + fmm.prepare_state, + positions, + masses, + leaf_size=int(leaf_size), + max_order=int(max_order), + warmup=int(warmup), + runs=int(runs), + jax_trace_label="prepare_state", + ) + prepared_state = prepare_once_timing.result + eval_timing = bench_utils.time_callable( + fmm.evaluate_prepared_state, + prepared_state, + warmup=int(warmup), + runs=int(runs), + jax_trace_label="evaluate_prepared_state", + **_evaluate_prepared_kwargs(fmm), + ) if benchmark_scope == "steady_eval": full_mean = float(eval_timing.mean) @@ -2712,6 +2723,7 @@ def main() -> None: cfg=cfg, fmm_kwargs=fmm_kwargs, autotune_cache_path=autotune_cache_path, + jax_trace_dir=args.jax_trace_dir, ) elif args.mode == "audit": row = _run_audit_case( diff --git a/examples/benchmark_utils.py b/examples/benchmark_utils.py index fa3e8862..54923821 100644 --- a/examples/benchmark_utils.py +++ b/examples/benchmark_utils.py @@ -42,6 +42,7 @@ def time_callable( *args: Any, warmup: int = 1, runs: int = 5, + jax_trace_label: Optional[str] = None, **kwargs: Any, ) -> TimingResult: """Measure execution time for ``fn`` with optional warmup passes.""" @@ -54,13 +55,20 @@ def time_callable( for _ in range(warmup): _block_until_ready(fn(*args, **kwargs)) + trace_ctx = ( + jax.profiler.TraceAnnotation(jax_trace_label) + if jax_trace_label + else contextlib.nullcontext() + ) + samples = [] result: Any = None for _ in range(runs): - start = time.perf_counter() - result = fn(*args, **kwargs) - result = _block_until_ready(result) - end = time.perf_counter() + with trace_ctx: + start = time.perf_counter() + result = fn(*args, **kwargs) + result = _block_until_ready(result) + end = time.perf_counter() samples.append(end - start) wall_times = tuple(samples) diff --git a/jaccpot/runtime/_interaction_cache.py b/jaccpot/runtime/_interaction_cache.py index bcff1dab..a15128d8 100644 --- a/jaccpot/runtime/_interaction_cache.py +++ b/jaccpot/runtime/_interaction_cache.py @@ -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 diff --git a/jaccpot/upward/solidfmm_complex_tree_expansions.py b/jaccpot/upward/solidfmm_complex_tree_expansions.py index 5d4c21d7..00f93133 100644 --- a/jaccpot/upward/solidfmm_complex_tree_expansions.py +++ b/jaccpot/upward/solidfmm_complex_tree_expansions.py @@ -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, compute_tree_mass_moments_jit from jaccpot._env import env_flag from jaccpot.operators.complex_harmonics import p2m_complex_batch @@ -846,7 +846,12 @@ 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( + # tree, + # positions_sorted, + # masses_sorted, + #) + mass_moments = compute_tree_mass_moments_jit( tree, positions_sorted, masses_sorted, From 69e0c281fa416f452f259755e617c5e81e000be7 Mon Sep 17 00:00:00 2001 From: Tommaso Tarchi Date: Fri, 4 Sep 2026 13:14:48 +0200 Subject: [PATCH 2/2] cleanup --- examples/benchmark_gpu_n_ladder_production.py | 7 --- examples/benchmark_gpu_radix_worker.py | 44 +++++++------------ examples/benchmark_utils.py | 16 ++----- .../solidfmm_complex_tree_expansions.py | 7 +-- 4 files changed, 21 insertions(+), 53 deletions(-) diff --git a/examples/benchmark_gpu_n_ladder_production.py b/examples/benchmark_gpu_n_ladder_production.py index ec592b2a..e444ac16 100644 --- a/examples/benchmark_gpu_n_ladder_production.py +++ b/examples/benchmark_gpu_n_ladder_production.py @@ -496,12 +496,6 @@ def run_worker_case( "--config-json", json.dumps(payload), ] - trace_dir = ( - OUTPUT_DIR - / "jax_traces" - / f"N{num_particles}_{_traversal_cfg_label(traversal_cfg)}_m2l{m2l_chunk_size}_nf{nearfield_edge_chunk_size}" - ) - cmd += ["--jax-trace-dir", str(trace_dir)] proc = subprocess.Popen( cmd, stdout=subprocess.PIPE, @@ -806,7 +800,6 @@ def main() -> None: print(f"Starting N-ladder production sweep run_id={ladder_run_id}") sweep_rows: list[dict[str, Any]] = [] - print(f"-> OUTPUT_DIR is {OUTPUT_DIR}") for num_particles in ARGS.particle_counts: guidance = traversal_guidance_for_num_particles(int(num_particles)) preferred_label = _traversal_cfg_label(guidance["recommended"]) diff --git a/examples/benchmark_gpu_radix_worker.py b/examples/benchmark_gpu_radix_worker.py index e3a9dfbc..416edfe5 100644 --- a/examples/benchmark_gpu_radix_worker.py +++ b/examples/benchmark_gpu_radix_worker.py @@ -54,7 +54,6 @@ def _configure_worker_environment() -> None: import jax import jax.numpy as jnp -import jax.profiler REPO_ROOT = pathlib.Path(__file__).resolve().parents[1] if str(REPO_ROOT) not in sys.path: @@ -211,7 +210,6 @@ def _parse_args() -> argparse.Namespace: parser.add_argument("--autotune-cache", default=None) parser.add_argument("--emit-ready-marker", action="store_true") parser.add_argument("--config-json", required=True) - parser.add_argument("--jax-trace-dir", default=None) return parser.parse_args() @@ -1908,7 +1906,6 @@ def _run_sweep_case( cfg: dict[str, Any], fmm_kwargs: dict[str, Any], autotune_cache_path: Optional[str] = None, - jax_trace_dir: Optional[str] = None, ) -> dict[str, Any]: jaccpot_symbols = _load_jaccpot_symbols() FastMultipoleMethod = jaccpot_symbols["FastMultipoleMethod"] @@ -1946,31 +1943,23 @@ def _run_sweep_case( ) _emit_ready_marker() - trace_ctx = ( - jax.profiler.trace(jax_trace_dir, create_perfetto_link=False) - if jax_trace_dir - else contextlib.nullcontext() + prepare_once_timing = bench_utils.time_callable( + fmm.prepare_state, + positions, + masses, + leaf_size=int(leaf_size), + max_order=int(max_order), + warmup=int(warmup), + runs=int(runs), + ) + prepared_state = prepare_once_timing.result + eval_timing = bench_utils.time_callable( + fmm.evaluate_prepared_state, + prepared_state, + warmup=int(warmup), + runs=int(runs), + **_evaluate_prepared_kwargs(fmm), ) - with trace_ctx: - prepare_once_timing = bench_utils.time_callable( - fmm.prepare_state, - positions, - masses, - leaf_size=int(leaf_size), - max_order=int(max_order), - warmup=int(warmup), - runs=int(runs), - jax_trace_label="prepare_state", - ) - prepared_state = prepare_once_timing.result - eval_timing = bench_utils.time_callable( - fmm.evaluate_prepared_state, - prepared_state, - warmup=int(warmup), - runs=int(runs), - jax_trace_label="evaluate_prepared_state", - **_evaluate_prepared_kwargs(fmm), - ) if benchmark_scope == "steady_eval": full_mean = float(eval_timing.mean) @@ -2723,7 +2712,6 @@ def main() -> None: cfg=cfg, fmm_kwargs=fmm_kwargs, autotune_cache_path=autotune_cache_path, - jax_trace_dir=args.jax_trace_dir, ) elif args.mode == "audit": row = _run_audit_case( diff --git a/examples/benchmark_utils.py b/examples/benchmark_utils.py index 54923821..fa3e8862 100644 --- a/examples/benchmark_utils.py +++ b/examples/benchmark_utils.py @@ -42,7 +42,6 @@ def time_callable( *args: Any, warmup: int = 1, runs: int = 5, - jax_trace_label: Optional[str] = None, **kwargs: Any, ) -> TimingResult: """Measure execution time for ``fn`` with optional warmup passes.""" @@ -55,20 +54,13 @@ def time_callable( for _ in range(warmup): _block_until_ready(fn(*args, **kwargs)) - trace_ctx = ( - jax.profiler.TraceAnnotation(jax_trace_label) - if jax_trace_label - else contextlib.nullcontext() - ) - samples = [] result: Any = None for _ in range(runs): - with trace_ctx: - start = time.perf_counter() - result = fn(*args, **kwargs) - result = _block_until_ready(result) - end = time.perf_counter() + start = time.perf_counter() + result = fn(*args, **kwargs) + result = _block_until_ready(result) + end = time.perf_counter() samples.append(end - start) wall_times = tuple(samples) diff --git a/jaccpot/upward/solidfmm_complex_tree_expansions.py b/jaccpot/upward/solidfmm_complex_tree_expansions.py index 00f93133..0bf5803b 100644 --- a/jaccpot/upward/solidfmm_complex_tree_expansions.py +++ b/jaccpot/upward/solidfmm_complex_tree_expansions.py @@ -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, compute_tree_mass_moments_jit +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 @@ -846,11 +846,6 @@ 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( - # tree, - # positions_sorted, - # masses_sorted, - #) mass_moments = compute_tree_mass_moments_jit( tree, positions_sorted,