from functools import partial
import itertools
from pathlib import Path
import re
import time
import matplotlib.pyplot as plt
import numpy as np
import pandas as pd
import nomad.contact_estimation as contact
import nomad.filters as filters
import nomad.io.base as loader
[docs]
class AlgorithmRegistry:
"""Simple registry for parameterized stop-detection runs."""
def __init__(self):
self._algos = []
self._timings = []
self._family_counts = {}
def __len__(self):
return len(self._algos)
def __iter__(self):
return iter(self._algos)
def _next_auto_family(self, base_name):
count = self._family_counts.get(base_name, 0) + 1
self._family_counts[base_name] = count
if count == 1:
return base_name
return f"{base_name}_{count}"
@staticmethod
def _normalize_algorithm_name(name):
for suffix in ("_labels_per_user", "_per_user", "_labels"):
if name.endswith(suffix):
return name[: -len(suffix)]
return name
@staticmethod
def _param_codifier(value):
if isinstance(value, np.generic):
value = value.item()
if isinstance(value, float):
text = f"{value:.12g}"
elif isinstance(value, bool):
text = str(value).lower()
elif value is None:
text = "none"
else:
text = str(value)
text = re.sub(r"[^0-9A-Za-z_.-]+", "-", text).strip("-")
return text or "empty"
@classmethod
def _algorithm_identifier(cls, index, params):
if not params:
return f"{index:03d}"
param_parts = [
f"{cls._param_codifier(key)}-{cls._param_codifier(value)}"
for key, value in sorted(params.items())
]
return f"{index:03d}__{'__'.join(param_parts)}"
@staticmethod
def _expand_values(value, granularity):
if isinstance(value, tuple):
start, stop = value
return np.linspace(start, stop, int(granularity)).tolist()
if isinstance(value, np.ndarray):
return value.tolist()
if isinstance(value, list):
return value
return [value]
[docs]
def add_algorithm(self, fn, family=None, granularity=30, **param_specs):
algorithm_name = self._normalize_algorithm_name(fn.__name__)
family_name = family or self._next_auto_family(algorithm_name)
expanded = {
key: self._expand_values(value, granularity)
for key, value in param_specs.items()
}
keys = list(expanded.keys())
values = [expanded[k] for k in keys]
for index, combo in enumerate(itertools.product(*values), start=1):
params = dict(zip(keys, combo))
algorithm_id = self._algorithm_identifier(index, params)
self._algos.append(
{
"algorithm": algorithm_id,
"family": family_name,
"fn": fn,
"call": partial(fn, **params),
"params": params,
}
)
return family_name
[docs]
def iter_algorithms(self, family=None):
if family is None:
yield from self._algos
return
for algo in self._algos:
if algo["family"] == family:
yield algo
[docs]
def annotate_metrics(self, metrics, algo):
row = dict(metrics)
row["algorithm"] = algo["algorithm"]
row["family"] = algo["family"]
row.update(algo["params"])
return row
[docs]
def time_call(self, algo, *args, **kwargs):
t0 = time.perf_counter()
output = algo["call"](*args, **kwargs)
elapsed_s = time.perf_counter() - t0
ping_count = len(args[0]) if args else np.nan
self.record_timing(algo, elapsed_s, ping_count=ping_count)
return output
[docs]
def record_timing(self, algo, elapsed_s, ping_count=np.nan):
self._timings.append(
{
"algorithm": algo["algorithm"],
"family": algo["family"],
"elapsed_s": float(elapsed_s),
"ping_count": float(ping_count),
**algo["params"],
}
)
[docs]
def timing_frame(self):
return pd.DataFrame(self._timings)
[docs]
def family_timing_summary(self):
timing = self.timing_frame()
if timing.empty:
return pd.DataFrame(columns=["family", "n", "mean_s", "avg_pings"])
return (
timing.groupby("family", as_index=False)
.agg(n=("elapsed_s", "count"), mean_s=("elapsed_s", "mean"), avg_pings=("ping_count", "mean"))
.sort_values("mean_s")
)
[docs]
def compute_visitation_errors(overlaps, true_visits, traj_cols=None, right_traj_cols=None, **kwargs):
if right_traj_cols is None:
right_schema_input = traj_cols
right_kwargs = kwargs
else:
right_schema_input = right_traj_cols
right_kwargs = {}
temp_traj_cols = loader._parse_traj_cols(true_visits.columns, right_schema_input, right_kwargs, warn=False)
true_visits = true_visits.dropna()
t_name, _ = loader._fallback_time_cols_dt(true_visits.columns, right_schema_input, right_kwargs)
t_key = temp_traj_cols[t_name]
if true_visits[t_key].duplicated().any():
dup_ts = true_visits.loc[true_visits[t_key].duplicated(), t_key].unique()
raise ValueError(
"Ground-truth stops share the same start time(s), which violates the "
"per-stop key assumption. Duplicated timestamps: " + repr(dup_ts)
)
n_truth = len(true_visits)
temp_cols = loader._parse_traj_cols([], traj_cols, kwargs, warn=False)
loc_left = f"{temp_cols['location_id']}_left"
loc_right = f"{temp_traj_cols['location_id']}_right"
time_keys = ["datetime", "start_datetime", "timestamp", "start_timestamp"]
if "timestamp" in kwargs or "start_timestamp" in kwargs:
time_keys = ["timestamp", "start_timestamp", "datetime", "start_datetime"]
if "datetime" in kwargs or "start_datetime" in kwargs:
time_keys = ["datetime", "start_datetime", "timestamp", "start_timestamp"]
t_left = None
for key in time_keys:
col = f"{temp_cols[key]}_left"
if col in overlaps.columns:
t_left = col
break
time_keys = ["datetime", "start_datetime", "timestamp", "start_timestamp"]
if "timestamp" in right_kwargs or "start_timestamp" in right_kwargs:
time_keys = ["timestamp", "start_timestamp", "datetime", "start_datetime"]
if "datetime" in right_kwargs or "start_datetime" in right_kwargs:
time_keys = ["datetime", "start_datetime", "timestamp", "start_timestamp"]
t_right = None
for key in time_keys:
col = f"{temp_traj_cols[key]}_right"
if col in overlaps.columns:
t_right = col
break
if t_left is None or t_right is None:
raise ValueError("compute_visitation_errors: could not resolve the overlap start-time columns.")
for col in (loc_left, loc_right):
if col not in overlaps.columns:
raise ValueError(f"compute_visitation_errors: expected column '{col}' in overlaps but not found.")
overlaps = overlaps.fillna({loc_left: "Street"})
bad_ts = set(overlaps[t_right]) - set(true_visits[t_key])
if bad_ts:
raise ValueError(
"compute_visitation_errors: overlap rows reference start times that "
"do not exist in ground truth: " + repr(sorted(bad_ts)[:10])
)
same_loc = overlaps[loc_left] == overlaps[loc_right]
merged_rows = overlaps.groupby(t_left)[loc_right].transform("nunique").gt(1)
num_overlapped = overlaps[t_right].nunique()
missed_fraction = 1 - num_overlapped / n_truth
merged_fraction = overlaps.loc[merged_rows, t_right].nunique() / n_truth
split_fraction = overlaps[same_loc].groupby(t_right)[t_left].nunique().gt(1).sum() / n_truth
return {
"missed_fraction": missed_fraction,
"merged_fraction": merged_fraction,
"split_fraction": split_fraction,
}
[docs]
def compute_stop_detection_metrics(
stops,
truth,
user_id=None,
algorithm=None,
prf_only=True,
traj_cols=None,
right_traj_cols=None,
**kwargs,
):
if len(stops) == 0:
return {
"precision": 0.0,
"recall": 0.0,
"f1": 0.0,
"missed_fraction": 1.0,
"merged_fraction": 0.0,
"split_fraction": 0.0,
"user_id": user_id,
"algorithm": algorithm,
}
input_traj_cols = traj_cols
left_kwargs = dict(kwargs)
if right_traj_cols is None:
right_schema_input = traj_cols
right_kwargs = left_kwargs
else:
right_schema_input = right_traj_cols
right_kwargs = {}
right_schema_hint = ""
if right_traj_cols is None:
right_schema_hint = (
" The shared traj_cols/kwargs mapping appears not to fit the truth table. "
"Pass right_traj_cols for the truth table, or make both tables use the same relevant "
"column names for time, duration, and location."
)
traj_cols = loader._parse_traj_cols(stops.columns, traj_cols, left_kwargs, warn=False)
temp_traj_cols = loader._parse_traj_cols(truth.columns, right_schema_input, right_kwargs, warn=False)
left_loc = traj_cols["location_id"]
right_loc = temp_traj_cols["location_id"]
if left_loc not in stops.columns:
raise ValueError(
"Could not find the mapped location column in predicted stops. "
f"Expected '{left_loc}' in columns {list(stops.columns)}."
)
if right_loc not in truth.columns:
raise ValueError(
"Could not find the mapped location column in ground-truth stops. "
f"Expected '{right_loc}' in columns {list(truth.columns)}."
+ right_schema_hint
)
stops_clean = stops.copy()
truth_clean = truth.copy()
stops_clean[left_loc] = stops_clean[left_loc].fillna("Street")
truth_clean[right_loc] = truth_clean[right_loc].fillna("Street")
truth_buildings = truth[truth[right_loc].notna()].copy()
overlaps = contact.overlapping_visits(
left=stops_clean,
right=truth_clean,
match_location=True,
traj_cols=input_traj_cols,
right_traj_cols=right_traj_cols,
**kwargs,
)
left_t_name, left_use_datetime = loader._fallback_time_cols_dt(stops_clean.columns, input_traj_cols, left_kwargs)
left_t_key = traj_cols[left_t_name]
left_e_t_key = traj_cols["end_datetime" if left_use_datetime else "end_timestamp"]
left_end_col_present = loader._has_end_cols(stops_clean.columns, traj_cols)
left_duration_col_present = loader._has_duration_cols(stops_clean.columns, traj_cols)
if not (left_end_col_present or left_duration_col_present):
raise ValueError("Predicted stops must provide either an end time or a duration.")
if not left_end_col_present:
if left_use_datetime:
stops_clean[left_e_t_key] = stops_clean[left_t_key] + pd.to_timedelta(stops_clean[traj_cols["duration"]], unit="m")
else:
stops_clean[left_e_t_key] = stops_clean[left_t_key] + stops_clean[traj_cols["duration"]] * 60
right_t_name, right_use_datetime = loader._fallback_time_cols_dt(truth_clean.columns, right_schema_input, right_kwargs)
right_t_key = temp_traj_cols[right_t_name]
right_e_t_key = temp_traj_cols["end_datetime" if right_use_datetime else "end_timestamp"]
right_end_col_present = loader._has_end_cols(truth_clean.columns, temp_traj_cols)
right_duration_col_present = loader._has_duration_cols(truth_clean.columns, temp_traj_cols)
if not (right_end_col_present or right_duration_col_present):
raise ValueError("Ground-truth stops must provide either an end time or a duration." + right_schema_hint)
if not right_end_col_present:
if right_use_datetime:
truth_clean[right_e_t_key] = truth_clean[right_t_key] + pd.to_timedelta(truth_clean[temp_traj_cols["duration"]], unit="m")
else:
truth_clean[right_e_t_key] = truth_clean[right_t_key] + truth_clean[temp_traj_cols["duration"]] * 60
left_duration = traj_cols["duration"]
if left_duration in stops_clean.columns:
total_pred = stops_clean[left_duration].sum()
else:
if left_use_datetime:
total_pred = (filters.to_timestamp(stops_clean[left_e_t_key]) - filters.to_timestamp(stops_clean[left_t_key])).floordiv(60).sum()
else:
total_pred = ((stops_clean[left_e_t_key] - stops_clean[left_t_key]) // 60).sum()
right_duration = temp_traj_cols["duration"]
if right_duration in truth_clean.columns:
total_truth = truth_clean[right_duration].sum()
else:
if right_use_datetime:
total_truth = (filters.to_timestamp(truth_clean[right_e_t_key]) - filters.to_timestamp(truth_clean[right_t_key])).floordiv(60).sum()
else:
total_truth = ((truth_clean[right_e_t_key] - truth_clean[right_t_key]) // 60).sum()
tp = overlaps[left_duration].sum()
prf_metrics = contact.precision_recall_f1_from_minutes(total_pred, total_truth, tp)
if prf_only:
return {**prf_metrics, "user_id": user_id, "algorithm": algorithm}
if len(truth_buildings) > 0:
overlaps_err = contact.overlapping_visits(
left=stops_clean,
right=truth_buildings,
match_location=False,
traj_cols=input_traj_cols,
right_traj_cols=right_traj_cols,
**kwargs,
)
error_metrics = compute_visitation_errors(
overlaps_err,
truth_buildings,
traj_cols=input_traj_cols,
right_traj_cols=right_traj_cols,
**kwargs,
)
else:
error_metrics = {"missed_fraction": 0.0, "merged_fraction": 0.0, "split_fraction": 0.0}
return {**prf_metrics, **error_metrics, "user_id": user_id, "algorithm": algorithm}
def _desaturate_toward_white(color, amount=0.4):
rgb = np.array(color[:3])
white = np.array([1.0, 1.0, 1.0])
return tuple(rgb + (white - rgb) * amount)
def _blend_toward_white(color, amount=0.3):
rgb = np.array(color[:3])
white = np.array([1.0, 1.0, 1.0])
return tuple(rgb + (white - rgb) * amount)
def _group_palette(group_order, group_families=None, cmap="tab10"):
cmap_obj = plt.colormaps.get_cmap(cmap)
if group_families is None:
return {
group: cmap_obj(i / max(len(group_order) - 1, 1))
for i, group in enumerate(group_order)
}, None
family_order = list(dict.fromkeys(group_families[group] for group in group_order))
base_colors = {
family: cmap_obj(i / max(len(family_order) - 1, 1))
for i, family in enumerate(family_order)
}
colors = {}
for family in family_order:
family_groups = [group for group in group_order if group_families[group] == family]
shade_levels = np.linspace(0.0, 0.35, len(family_groups))
for group, shade in zip(family_groups, shade_levels):
colors[group] = _blend_toward_white(base_colors[family], shade)
return colors, base_colors
[docs]
def bootstrap_metric_summary(
data,
metrics,
group_col="algorithm",
unit_col="user",
n_boot=2000,
interval=(0.05, 0.95),
random_state=2025,
):
"""Bootstrap per-unit medians and summarize them as point estimates plus intervals."""
columns = [group_col, "metric", "estimate", "lower", "upper"]
if data.empty:
return pd.DataFrame(columns=columns)
per_unit = (
data.groupby([unit_col, group_col], as_index=False)[metrics]
.median()
)
point_estimates = (
per_unit.groupby(group_col, as_index=False)[metrics]
.median()
.melt(
id_vars=group_col,
value_vars=metrics,
var_name="metric",
value_name="estimate",
)
)
unit_ids = per_unit[unit_col].drop_duplicates().to_numpy()
per_unit = per_unit.set_index(unit_col)
rng = np.random.default_rng(random_state)
bootstrap_rows = []
for bootstrap_id in range(n_boot):
sampled_units = rng.choice(unit_ids, size=len(unit_ids), replace=True)
sampled = per_unit.loc[sampled_units].reset_index()
summary = sampled.groupby(group_col, as_index=False)[metrics].median()
summary["bootstrap_id"] = bootstrap_id
bootstrap_rows.append(summary)
bootstrap_df = pd.concat(bootstrap_rows, ignore_index=True)
lower_q, upper_q = interval
intervals = (
bootstrap_df.melt(
id_vars=[group_col, "bootstrap_id"],
value_vars=metrics,
var_name="metric",
value_name="value",
)
.groupby([group_col, "metric"])["value"]
.quantile([lower_q, upper_q])
.unstack()
.reset_index()
.rename(columns={lower_q: "lower", upper_q: "upper"})
)
return point_estimates.merge(intervals, on=[group_col, "metric"], how="left")
[docs]
def plot_metric_vs_param(
ax,
data,
x,
y,
algo_param,
title="",
xlabel="",
ylabel="",
show_max_mean_line=False,
show_band=False,
color="C0",
):
"""Simple sensitivity plot used by validation notebooks."""
if data.empty:
ax.set_title(title)
ax.set_xlabel(xlabel)
ax.set_ylabel(ylabel)
return ax
stats = (
data.groupby(algo_param)[y]
.agg(mean="mean", std="std")
.reset_index()
.sort_values(algo_param)
)
if "user_id" in data.columns:
per_user_max = (
data.loc[data.groupby("user_id")[y].idxmax(), ["user_id", algo_param, y]]
.sort_values(algo_param)
)
else:
per_user_max = pd.DataFrame(columns=[algo_param, y])
scatter_color = _desaturate_toward_white(plt.matplotlib.colors.to_rgba(color), amount=0.5)
if show_band:
ax.fill_between(
stats[algo_param],
stats["mean"] - stats["std"],
stats["mean"] + stats["std"],
color=color,
alpha=0.2,
label="Mean +/- 1 SD",
zorder=1,
)
ax.scatter(data[x], data[y], s=20, alpha=0.07, color=scatter_color, zorder=2)
if not per_user_max.empty:
ax.scatter(
per_user_max[algo_param],
per_user_max[y],
s=35,
alpha=0.9,
color="red",
marker="x",
label="Per-user max",
zorder=4,
)
ax.plot(stats[algo_param], stats["mean"], linewidth=2.5, color=color, label="Mean", zorder=3)
if show_max_mean_line:
max_mean = float(np.nanmax(stats["mean"].to_numpy()))
ax.axhline(max_mean, color="black", linestyle="--", linewidth=1.2, label=f"Max mean ({max_mean:.3g})")
ax.set_title(title)
ax.set_xlabel(xlabel)
ax.set_ylabel(ylabel)
ax.grid(True, linestyle=":", alpha=0.6)
ax.legend(loc="best", frameon=True)
return ax
[docs]
def plot_family_timing(summary_df, ax=None, title="Mean Time Per Parameterization"):
"""Bar plot for family-level runtime summaries."""
if ax is None:
_, ax = plt.subplots(figsize=(8, 4.5))
if summary_df is None or summary_df.empty:
ax.set_title(title)
ax.set_xlabel("Family")
ax.set_ylabel("Mean seconds")
return ax
ordered = summary_df.sort_values("mean_s")
ax.bar(ordered["family"], ordered["mean_s"], alpha=0.85)
ax.set_title(title)
ax.set_xlabel("Family")
ax.set_ylabel("Mean seconds")
ax.tick_params(axis="x", rotation=45)
ax.grid(True, axis="y", linestyle=":", alpha=0.5)
return ax
[docs]
def plot_metric_boxplots(
data,
metrics,
group_col="algorithm",
group_order=None,
group_families=None,
colors=None,
cmap="tab10",
figsize=None,
save_path=None,
metric_titles=None,
legend_title="Base Algorithm",
x_tick_rotation=35,
whis=(5, 95),
showfliers=False,
):
"""Plot per-group metric distributions as boxplots."""
if group_order is None:
group_order = data[group_col].drop_duplicates().tolist()
if colors is None:
colors, legend_colors = _group_palette(group_order, group_families=group_families, cmap=cmap)
else:
legend_colors = None
if figsize is None:
figsize = (max(4.2 * len(metrics), 10), 6.2)
fig, axes = plt.subplots(1, len(metrics), figsize=figsize, sharey=False)
if len(metrics) == 1:
axes = [axes]
for ax, metric in zip(axes, metrics):
metric_data = [data.loc[data[group_col] == group, metric].dropna() for group in group_order]
bp = ax.boxplot(
metric_data,
positions=np.arange(len(group_order)),
patch_artist=True,
widths=0.5,
whis=whis,
showfliers=showfliers,
medianprops={"color": "black", "linewidth": 0.9},
boxprops={"linewidth": 0.8},
whiskerprops={"linewidth": 0.8, "color": "black"},
capprops={"linewidth": 0.8, "color": "black"},
)
for box, group in zip(bp["boxes"], group_order):
box.set_facecolor(colors[group])
box.set_edgecolor("black")
ax.set_facecolor("#EAEAF2")
ax.grid(axis="y", color="darkgray", linestyle="--", linewidth=0.8, alpha=0.75)
ax.set_xticks(np.arange(len(group_order)))
ax.set_xticklabels(group_order, rotation=x_tick_rotation, ha="right")
if metric_titles is None:
title = metric.replace("_", " ").title()
else:
title = metric_titles.get(metric, metric.replace("_", " ").title())
ax.set_title(title, fontsize=16)
if legend_colors is not None:
family_order = list(legend_colors.keys())
handles = [plt.matplotlib.patches.Patch(color=legend_colors[family], label=family) for family in family_order]
fig.legend(
handles,
family_order,
loc="lower center",
ncol=len(family_order),
bbox_to_anchor=(0.5, -0.08),
fontsize=12,
title=legend_title,
title_fontsize=13,
frameon=True,
)
plt.subplots_adjust(bottom=0.34, top=0.92)
else:
plt.subplots_adjust(bottom=0.22, top=0.92)
if save_path is not None:
save_path = Path(save_path)
save_path.parent.mkdir(parents=True, exist_ok=True)
fig.savefig(save_path.with_suffix(".png"), dpi=300, bbox_inches="tight")
fig.savefig(save_path.with_suffix(".svg"), bbox_inches="tight")
return fig, axes
[docs]
def plot_metric_intervals(
summary_df,
metrics,
group_col="algorithm",
group_order=None,
group_families=None,
colors=None,
cmap="tab10",
figsize=None,
save_path=None,
metric_titles=None,
legend_title="Base Algorithm",
x_tick_rotation=35,
):
"""Plot precomputed point estimates with interval whiskers for each metric."""
if group_order is None:
group_order = summary_df[group_col].drop_duplicates().tolist()
if colors is None:
colors, legend_colors = _group_palette(group_order, group_families=group_families, cmap=cmap)
else:
legend_colors = None
if figsize is None:
figsize = (max(4.2 * len(metrics), 10), 6.2)
fig, axes = plt.subplots(1, len(metrics), figsize=figsize, sharey=False)
if len(metrics) == 1:
axes = [axes]
for ax, metric in zip(axes, metrics):
metric_summary = (
summary_df.loc[summary_df["metric"] == metric]
.set_index(group_col)
.reindex(group_order)
.reset_index()
)
x = np.arange(len(group_order))
y = metric_summary["estimate"].to_numpy()
yerr = np.vstack([
y - metric_summary["lower"].to_numpy(),
metric_summary["upper"].to_numpy() - y,
])
ax.set_facecolor("#EAEAF2")
for idx, group in enumerate(group_order):
ax.errorbar(
x[idx],
y[idx],
yerr=yerr[:, idx:idx + 1],
fmt="o",
color="black",
ecolor="black",
elinewidth=1.0,
capsize=4,
markersize=7,
markerfacecolor=colors[group],
markeredgecolor="black",
markeredgewidth=0.8,
zorder=4,
)
ax.grid(axis="y", color="darkgray", linestyle="--", linewidth=0.8, alpha=0.75)
ax.set_xticks(x)
ax.set_xticklabels(group_order, rotation=x_tick_rotation, ha="right")
if metric_titles is None:
title = metric.replace("_", " ").title()
else:
title = metric_titles.get(metric, metric.replace("_", " ").title())
ax.set_title(title, fontsize=16)
if legend_colors is not None:
family_order = list(legend_colors.keys())
handles = [plt.matplotlib.patches.Patch(color=legend_colors[family], label=family) for family in family_order]
fig.legend(
handles,
family_order,
loc="lower center",
ncol=len(family_order),
bbox_to_anchor=(0.5, -0.08),
fontsize=12,
title=legend_title,
title_fontsize=13,
frameon=True,
)
plt.subplots_adjust(bottom=0.34, top=0.92)
else:
plt.subplots_adjust(bottom=0.22, top=0.92)
if save_path is not None:
save_path = Path(save_path)
save_path.parent.mkdir(parents=True, exist_ok=True)
fig.savefig(save_path.with_suffix(".png"), dpi=300, bbox_inches="tight")
fig.savefig(save_path.with_suffix(".svg"), bbox_inches="tight")
return fig, axes
__all__ = [
"AlgorithmRegistry",
"bootstrap_metric_summary",
"compute_visitation_errors",
"compute_stop_detection_metrics",
"plot_metric_boxplots",
"plot_metric_intervals",
"plot_metric_vs_param",
"plot_family_timing",
]