Skip to content
Merged
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
13 changes: 12 additions & 1 deletion cuequivariance_jax/cuequivariance_jax/benchmarking.py
Original file line number Diff line number Diff line change
Expand Up @@ -19,7 +19,6 @@
from functools import partial

import jax
from jax.experimental.mosaic.gpu.profiler import _event_elapsed, _event_record


def measure_clock_ticks(f, *args, **kwargs) -> tuple[float, float]:
Expand Down Expand Up @@ -49,6 +48,18 @@ def my_function(x, y):
"""
from cuequivariance_ops_jax import noop, sleep, synchronize

try:
from jax.experimental.mosaic.gpu.profiler import _event_elapsed, _event_record
except ImportError:
try:
from cuequivariance_ops_jax import event_elapsed as _event_elapsed
from cuequivariance_ops_jax import event_record as _event_record
except ImportError:
raise ImportError(
"_event_elapsed and/or _event_record not found in jax.experimental.mosaic.gpu.profiler\n"
"They are known to be available in jax>=0.4.36,<=0.7.1"
)

def run_func(state):
"""Wrapper function that calls the target function and ensures proper data flow."""
args, kwargs = state
Expand Down