Skip to content

voiage.backends.performance_profiler.JaxPerformanceProfiler

Profile and optimize JAX computations.

profile_function([positional or keyword] self: None = None, [positional or keyword] func: Callable[..., object] = None) -> Callable[..., object]

Profile function execution time.

Parameters:

  • self
  • func Callable[..., object]

Returns: Callable[..., object]

compare_implementations([positional or keyword] self: None = None, [positional or keyword] numpy_func: Callable[..., object] = None, [positional or keyword] jax_func: Callable[..., object] = None, [positional or keyword] test_data: tuple[object, ...] = None, [positional or keyword] n_runs: int = 10) -> dict[str, list[float]]

Compare NumPy vs JAX implementations.

Parameters:

  • self
  • numpy_func Callable[..., object]
  • jax_func Callable[..., object]
  • test_data tuple[object, ...]
  • n_runs int (default: 10)

Returns: dict[str, list[float]]

memory_usage_analysis([positional or keyword] self: None = None, [positional or keyword] func: Callable[..., object] = None, [variadic positional] args: object = (), [variadic keyword] kwargs: object = {}) -> dict[str, object]

Analyze memory usage of a function.

Parameters:

  • self
  • func Callable[..., object]
  • args object (default: ())
  • kwargs object (default: {})

Returns: dict[str, object]

get_performance_report([positional or keyword] self: None = None) -> dict[str, dict[str, object]]

Generate performance report.

Parameters:

  • self

Returns: dict[str, dict[str, object]]