Skip to content
Open
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
131 changes: 113 additions & 18 deletions src/nwb_benchmarks/database/_visualization.py
Original file line number Diff line number Diff line change
@@ -1,6 +1,7 @@
import shutil
import textwrap
import warnings
import zlib
from collections import OrderedDict
from pathlib import Path
from typing import Any, Dict, List, Optional
Expand Down Expand Up @@ -34,6 +35,25 @@
)
DEFAULT_BENCHMARK_ORDER = list(DEFAULT_BENCHMARK_LABELS.keys())

DEFAULT_METHOD_COLORS = {
"hdf5 h5py remfile no cache": "#0072B2",
"hdf5 h5py remfile with cache": "#56B4E9",
"hdf5 h5py fsspec https no cache": "#009E73",
"hdf5 h5py fsspec https with cache": "#CC79A7",
"hdf5 h5py fsspec s3 no cache": "#D55E00",
"hdf5 h5py fsspec s3 with cache": "#E69F00",
"hdf5 h5py ros3": "#000000",
"lindi h5py": "#999999",
"zarr s3": "#F0E442",
"zarr s3 force no consolidated": "#9467BD",
}
DEFAULT_METHOD_COLORS.update(
{
method.replace("h5py", "pynwb").replace("zarr", "zarr pynwb"): color
for method, color in list(DEFAULT_METHOD_COLORS.items())
}
)

NETWORK_METRIC_LABELS = {
"amount_downloaded_in_bytes": "Downloaded bytes",
"amount_downloaded_in_number_of_packets": "Downloaded packets",
Expand All @@ -47,6 +67,8 @@
"total_transfer_time_in_seconds": "Total transfer time (s)",
}

BYTE_NETWORK_METRICS = {metric for metric in NETWORK_METRIC_LABELS if metric.endswith("_in_bytes")}


class BenchmarkVisualizer:
"""Handles plotting and visualization of benchmark results."""
Expand Down Expand Up @@ -240,6 +262,67 @@ def _display_method_label(method: str) -> str:
label = BenchmarkVisualizer._display_method_labels([method])[0]
return label

@staticmethod
def _method_palette(methods: List[str]) -> Dict[str, Any]:
"""Return a stable method-to-color mapping for the requested benchmark methods."""
fallback_colors = sns.color_palette("husl", n_colors=256)
palette = {}
for method in methods:
palette[method] = (
DEFAULT_METHOD_COLORS[method]
if method in DEFAULT_METHOD_COLORS
else fallback_colors[zlib.crc32(method.encode("utf-8")) % len(fallback_colors)]
)
return palette

@classmethod
def _resolve_palette(cls, df, group: str, order: List[str], palette: Any) -> Any:
"""Use stable method colors when a categorical plot groups by benchmark method names."""
if order is None:
return palette
if group in df and any(method in DEFAULT_METHOD_COLORS for method in order):
return cls._method_palette(order)
return palette

@staticmethod
def _add_throughput_lines(ax, rates_mb_s=(0.1, 1, 10, 100)) -> None:
"""Draw lines of constant throughput (runtime = bytes / rate) on a log-log axis."""
xmin, xmax = ax.get_xlim()
ymin, ymax = ax.get_ylim()
if xmin <= 0 or xmax <= 0 or ymin <= 0 or ymax <= 0:
return
x = np.logspace(np.log10(xmin), np.log10(xmax), 200)
for rate in rates_mb_s:
y = x / (rate * 1e6)
ax.plot(x, y, color="0.75", lw=0.8, ls="--", zorder=0)
y_at_xmax = xmax / (rate * 1e6)
if ymin <= y_at_xmax <= ymax:
ax.text(
xmax,
y_at_xmax,
f" {rate:g} MB/s",
color="0.5",
fontsize=7,
va="center",
ha="left",
clip_on=False,
)
else:
x_at_ymax = ymax * rate * 1e6
if xmin <= x_at_ymax <= xmax:
ax.text(
x_at_ymax,
ymax,
f"{rate:g} MB/s",
color="0.5",
fontsize=7,
va="bottom",
ha="center",
clip_on=False,
)
ax.set_xlim(xmin, xmax)
ax.set_ylim(ymin, ymax)

def _write_summary_table(
self,
df: pd.DataFrame,
Expand Down Expand Up @@ -887,7 +970,7 @@ def _plot_slicing_network_runtime_trajectories(
modality_order = ["Ecephys", "Ophys", "Icephys"]
methods = [method for method in self.pynwb_read_order if method in set(plot_df[col_name])]
methods.extend([method for method in plot_df[col_name].dropna().unique() if method not in methods])
palette = dict(zip(methods, sns.color_palette("colorblind", n_colors=len(methods))))
palette = self._method_palette(methods)
fig, axes = plt.subplots(1, len(modality_order), figsize=(15, 5), sharex=False, sharey=False)
for ax, modality in zip(axes, modality_order):
modality_df = plot_df[plot_df["modality"].eq(modality)]
Expand Down Expand Up @@ -939,7 +1022,9 @@ def _plot_slicing_network_runtime_trajectories(
ax.set_ylim(y_min / 1.6, y_max * 1.6)
ax.set_xscale("log")
ax.set_yscale("log")
ax.set_title(modality)
if network_metric in BYTE_NETWORK_METRICS:
self._add_throughput_lines(ax, rates_mb_s=(1, 10, 100))
ax.set_title(modality, y=1.08)
ax.grid(True, which="both", linewidth=0.3, alpha=0.35)
ax.set_xlabel(NETWORK_METRIC_LABELS.get(network_metric, network_metric))
axes[0].set_ylabel("Runtime (s)")
Expand All @@ -954,10 +1039,16 @@ def _plot_slicing_network_runtime_trajectories(
error_bar_sentence = (
"Error bars show standard deviation across repeated benchmark runs. " if show_error_bars else ""
)
throughput_caption = (
"Dashed gray lines indicate constant effective throughput. "
if network_metric in BYTE_NETWORK_METRICS
else ""
)
caption = (
f"Remote slicing {metric_label.lower()} vs. runtime for increasing slice sizes ({preload_caption}). "
"Points show average benchmark measurements and connected lines follow increasing slice size within each method. "
"Text annotations mark the largest slice size shown for each method. "
f"{throughput_caption}"
f"{error_bar_sentence}Panels separate data modalities."
)
caption = self._add_environment_caption(caption)
Expand Down Expand Up @@ -1021,7 +1112,7 @@ def _plot_file_open_network_runtime_scatter(
plot_df.to_csv(self.summary_tables_directory / f"{filename.stem}_summary.csv", index=False)

methods = [method for method in default_order if method in set(plot_df[col_name])]
palette = dict(zip(methods, sns.color_palette("colorblind", n_colors=len(methods))))
palette = self._method_palette(methods)
modality_order = ["Ecephys", "Ophys", "Icephys"]
modality_markers = {"Ecephys": "o", "Ophys": "s", "Icephys": "^"}
fig, ax = plt.subplots(figsize=(8, 6))
Expand All @@ -1043,7 +1134,7 @@ def _plot_file_open_network_runtime_scatter(
title += " - PyNWB"
elif suffix:
title += suffix
ax.set_title(title)
ax.set_title(title, y=1.08)
for modality in modality_order:
modality_df = method_df[method_df["modality"].eq(modality)]
if modality_df.empty:
Expand All @@ -1069,6 +1160,9 @@ def _plot_file_open_network_runtime_scatter(
ax.plot(modality_df["network_mean"], modality_df["runtime_mean"], **plot_kwargs)
ax.set_xscale("log")
ax.set_yscale("log")
if network_metric in BYTE_NETWORK_METRICS:
throughput_rates = (0.01, 0.1, 1, 10) if suffix == "_pynwb" else (0.1, 1, 10)
self._add_throughput_lines(ax, rates_mb_s=throughput_rates)
ax.grid(True, which="both", linewidth=0.3, alpha=0.35)
ax.set_xlabel(NETWORK_METRIC_LABELS.get(network_metric, network_metric))
ax.set_ylabel("Runtime (s)")
Expand All @@ -1086,10 +1180,16 @@ def _plot_file_open_network_runtime_scatter(
error_bar_sentence = (
"Error bars show standard deviation across repeated benchmark runs. " if show_error_bars else ""
)
throughput_caption = (
"Dashed gray lines indicate constant effective throughput. "
if network_metric in BYTE_NETWORK_METRICS
else ""
)
caption = (
f"Network-tracking {output_family} {metric_label.lower()} vs. runtime. "
"Points show average benchmark measurements. "
"Connected lines link modalities within each method. "
f"{throughput_caption}"
f"{error_bar_sentence}Marker shapes distinguish modalities."
)
caption = self._add_environment_caption(caption)
Expand Down Expand Up @@ -1319,6 +1419,8 @@ def plot_benchmark_dist(
self._record_skipped_output("plot", filename, "No data available to plot.")
return

resolved_palette = self._resolve_palette(df=df, group=group, order=metric_order, palette=palette)

g = sns.catplot(
data=df,
x="value",
Expand All @@ -1327,7 +1429,7 @@ def plot_benchmark_dist(
row=row,
hue=group,
sharex=sharex,
palette=palette,
palette=resolved_palette,
kind=kind,
order=metric_order,
legend=False,
Expand Down Expand Up @@ -1385,6 +1487,8 @@ def plot_benchmark_slices_vs_time(
self._record_skipped_output("plot", filename, "No data available to plot.")
return

palette = self._resolve_palette(df=df, group=group, order=metric_order, palette="Paired")

g = sns.catplot(
data=df,
x="slice_number",
Expand All @@ -1394,7 +1498,7 @@ def plot_benchmark_slices_vs_time(
hue=group,
hue_order=metric_order,
sharex=sharex,
palette="Paired",
palette=palette,
sharey=False,
kind="point",
)
Expand Down Expand Up @@ -1839,19 +1943,9 @@ def plot_benchmark_slice_extrapolations(
benchmarks = sorted(collected_stream.select(group).unique().to_series().to_list())

# Create mappings for plot
hdf5_colors = iter(sns.color_palette("Greens", n_colors=len([b for b in benchmarks if b.startswith("hdf5")])))
zarr_colors = iter(sns.color_palette("Reds", n_colors=len([b for b in benchmarks if b.startswith("zarr")])))
lindi_colors = iter(sns.color_palette("Blues", n_colors=len([b for b in benchmarks if b.startswith("lindi")])))
modality_to_col = {mod: i for i, mod in enumerate(modalities)}
row_to_row = {rv: i for i, rv in enumerate(row_values)}
benchmarks_to_color = {}
for bm in benchmarks:
if bm.startswith("hdf5"):
benchmarks_to_color[bm] = next(hdf5_colors)
elif bm.startswith("zarr"):
benchmarks_to_color[bm] = next(zarr_colors)
else:
benchmarks_to_color[bm] = next(lindi_colors)
benchmarks_to_color = self._method_palette(benchmarks)

# Create subplots
fig, axes = plt.subplots(nrows=len(row_values), ncols=len(modalities), figsize=(15, 10), squeeze=False)
Expand Down Expand Up @@ -2110,6 +2204,7 @@ def plot_performance_across_versions(
)

method_order = self.pynwb_read_order if order is None else order
palette = self._method_palette(method_order)
g = sns.catplot(
data=df,
x="environment_timepoint",
Expand All @@ -2120,7 +2215,7 @@ def plot_performance_across_versions(
hue_order=method_order,
order=sorted(df["environment_timepoint"].unique()),
sharex=True,
palette="Paired",
palette=palette,
sharey=False,
kind="point",
)
Expand Down
Loading