bgc-setnet / source /src /bgc_retrieval /statistics.py
whiteh4t's picture
Release final BGC retrieval checkpoints and model card
c87881a verified
Raw
History Blame Contribute Delete
2.53 kB
"""Family-level uncertainty and multiple-comparison correction."""
from __future__ import annotations
from collections.abc import Iterable
import numpy as np
import pandas as pd
from scipy import stats
def aggregate_draws_by_group(results: pd.DataFrame, metric_columns: Iterable[str]) -> pd.DataFrame:
required = {"method", "group_id", "draw"}
if missing := required.difference(results.columns):
raise ValueError(f"Results missing columns: {sorted(missing)}")
columns = list(metric_columns)
return results.groupby(["method", "group_id"], as_index=False)[columns].mean()
def hierarchical_bootstrap_ci(
group_values: pd.Series,
samples: int = 10000,
confidence: float = 0.95,
seed: int = 0,
) -> tuple[float, float]:
values = group_values.dropna().to_numpy(dtype=float)
if len(values) < 2:
raise ValueError("At least two groups are required for a confidence interval")
random_state = np.random.RandomState(seed)
indices = random_state.randint(0, len(values), size=(samples, len(values)))
bootstrap_means = values[indices].mean(axis=1)
tail = (1.0 - confidence) / 2.0
return tuple(float(value) for value in np.quantile(bootstrap_means, [tail, 1.0 - tail]))
def paired_family_test(
aggregated: pd.DataFrame,
method: str,
baseline: str,
metric: str,
) -> dict[str, float | int | str]:
pivot = aggregated.pivot(index="group_id", columns="method", values=metric).dropna(
subset=[method, baseline]
)
if len(pivot) < 2:
raise ValueError("At least two paired groups are required")
differences = pivot[method] - pivot[baseline]
if np.allclose(differences, 0.0):
statistic, p_value = 0.0, 1.0
else:
statistic, p_value = stats.wilcoxon(
pivot[method], pivot[baseline], alternative="greater", zero_method="wilcox"
)
return {
"method": method,
"baseline": baseline,
"metric": metric,
"groups": len(pivot),
"mean_delta": float(differences.mean()),
"statistic": float(statistic),
"p_value": float(p_value),
}
def holm_adjust(p_values: Iterable[float]) -> list[float]:
values = np.asarray(list(p_values), dtype=float)
order = np.argsort(values)
adjusted = np.empty_like(values)
running = 0.0
count = len(values)
for rank, index in enumerate(order):
running = max(running, (count - rank) * values[index])
adjusted[index] = min(1.0, running)
return adjusted.tolist()