stormlog.jax.visualizer

JAX Memory Visualization.

Classes

MemoryVisualizer([style, figure_size])

JAX memory visualization and dashboards.

class stormlog.jax.visualizer.MemoryVisualizer(style='default', figure_size=(12, 8))[source]

Bases: object

JAX memory visualization and dashboards.

Parameters:
  • style (str)

  • figure_size (Tuple[int, int])

plot_memory_timeline(results, interactive=False, save_path=None)[source]

Plot device memory usage timeline.

Parameters:
  • results (Any)

  • interactive (bool)

  • save_path (str | None)

Return type:

None

plot_function_comparison(function_profiles, save_path=None)[source]

Plot memory usage comparison for functions/contexts.

Parameters:
  • function_profiles (Dict[str, Dict[str, Any]])

  • save_path (str | None)

Return type:

None

create_memory_heatmap(results, save_path=None)[source]

Create a heatmap from available JAX device-memory samples.

Parameters:
  • results (Any)

  • save_path (str | None)

Return type:

None

export_data(results, output_path, format='csv')[source]

Export available JAX timeline samples as CSV or JSON.

Parameters:
  • results (Any)

  • output_path (str)

  • format (str)

Return type:

None

save_plots(results, output_dir='./plots/')[source]

Save the standard JAX timeline, comparison, and heatmap outputs.

Parameters:
  • results (Any)

  • output_dir (str)

Return type:

None

create_interactive_dashboard(results, port=8050)[source]

Serve an interactive JAX device-memory timeline when Dash is installed.

Parameters:
  • results (Any)

  • port (int)

Return type:

None