Skip to content
Merged
Show file tree
Hide file tree
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
6 changes: 4 additions & 2 deletions src/cell_eval/_cli/_run.py
Original file line number Diff line number Diff line change
Expand Up @@ -102,8 +102,10 @@ def parse_args_run(parser: ap.ArgumentParser):
action="store_true",
help="Additionally compute a data ceiling: a real-data-only upper bound on "
"each metric, estimated by splitting the real data into two disjoint halves "
"of n/2 cells and applying the Spearman-Brown correction (2r/(1+r)) to map "
"each reliability metric back to full depth. Writes ceiling_results.csv / "
"of floor(n/2) cells and applying the Spearman-Brown correction (2r/(1+r)) "
"to map each reliability metric back to full depth. Recomputes DE on each "
"half, so --de-real/--de-pred do not apply to the ceiling, and roughly "
"doubles the run time. Writes ceiling_results.csv / "
"agg_ceiling_results.csv alongside the normal results.",
)
parser.add_argument(
Expand Down
77 changes: 65 additions & 12 deletions src/cell_eval/_evaluator.py
Original file line number Diff line number Diff line change
Expand Up @@ -196,8 +196,8 @@ def compute_ceiling(
"""Estimate a data ceiling: the maximum achievable score per metric.

Uses the real data only. Each perturbation's cells (and the control's) are
split into two *disjoint* halves of ``n/2`` cells - no cell in both - and
one half is treated as "real", the other as "prediction". Running the
split into two *disjoint* halves of ``floor(n/2)`` cells - no cell in both -
and one half is treated as "real", the other as "prediction". Running the
normal metric pipeline on that self-split measures each metric per
perturbation at half depth; averaging over perturbations and applying the
Spearman-Brown correction ``r' = 2r/(1+r)`` maps that per-context mean to
Expand All @@ -214,14 +214,24 @@ def compute_ceiling(

The correction is applied only to the reliability metrics listed in the
module-level ``SB_METRICS`` set (bounded, higher-is-better, and empirically
well-behaved under doubling). Every other metric - error metrics, unbounded
counts, and reliability metrics left off that list (``clustering_agreement``,
``pearson_edistance``) - is emitted as ``NaN`` (no defensible ceiling).
well-behaved under doubling), and only where the measured reliability is
``r > 0`` - below that the correction is a pole rather than a correction, so
it is reported as ``NaN`` (see :func:`_spearman_brown_correct`). Every other
metric - error metrics, unbounded counts, and reliability metrics left off
that list (``clustering_agreement``, ``pearson_edistance``) - is emitted as
``NaN`` (no defensible ceiling).

``ceiling_results.csv`` holds the raw per-perturbation self-split scores;
``agg_ceiling_results.csv`` holds the SB-corrected per-metric ceiling. The
self-split DE is computed in-memory and never written. The same
``pdex_kwargs`` and ``allow_discrete`` as the main evaluation are reused so
the ceiling is directly comparable.

Cost: the two halves are materialized as copies, so peak memory is roughly
``2x`` the real matrix on top of the already-loaded pair, and the self-split
DE is computed for both halves - so a ceiling run roughly doubles the wall
time. A precomputed ``de_real``/``de_pred`` does not carry over to the
ceiling: the halves are new data and need their own DE.
"""
logger.info(f"Computing data ceiling (seed={seed})")
half_real, half_pred = self._disjoint_halves(seed)
Expand Down Expand Up @@ -268,33 +278,55 @@ def compute_ceiling(
return results, agg_results

def _disjoint_halves(self, seed: int) -> tuple[ad.AnnData, ad.AnnData]:
"""Split the real data into two *disjoint* halves of ``n/2`` cells each.
"""Split the real data into two *disjoint* halves of ``floor(n/2)`` cells each.

Each perturbation's cells (including the control's) are shuffled and split
without replacement into two halves of ``floor(n/2)`` cells - so no cell
appears in both halves, giving the independence a bootstrap self-split
lacks. Perturbations with fewer than 2 cells cannot be split and are
dropped from both halves. The resulting half depth (``n/2``) is corrected
back to full depth by the Spearman-Brown doubling in
:meth:`compute_ceiling`.
lacks. When ``n`` is odd the one leftover cell is discarded (both halves
must be the same depth for the doubling to hold). Perturbations with fewer
than 2 cells cannot be split and are dropped from both halves. The
resulting half depth (``floor(n/2)``) is corrected back to full depth by the
Spearman-Brown doubling in :meth:`compute_ceiling`.
"""
real = self.anndata_pair.real
pert_col = self.anndata_pair.pert_col
rng = np.random.default_rng(seed)

a_idx: list[np.ndarray] = []
b_idx: list[np.ndarray] = []
dropped = 0
for _pert, idx in real.obs.groupby(pert_col, observed=True).indices.items():
perm = rng.permutation(np.asarray(idx))
h = perm.size // 2
if h < 1:
continue # < 2 cells: cannot form two disjoint halves
dropped += 1 # < 2 cells: cannot form two disjoint halves
continue
a_idx.append(perm[:h])
b_idx.append(perm[h : 2 * h])

if not a_idx:
raise ValueError(
"no perturbation has >= 2 cells to split for the data ceiling"
)
if dropped:
logger.warning(
f"Ceiling: dropped {dropped} perturbation(s) with < 2 cells "
f"(cannot be split); the ceiling is averaged over the remaining "
f"{len(a_idx)} perturbation(s), a different set than the main "
f"evaluation's aggregate."
)

# Disjoint split has no duplicate rows, so obs names stay unique.
half_real = real[np.concatenate(a_idx)].copy()
half_pred = real[np.concatenate(b_idx)].copy()

control = self.anndata_pair.control_pert
if control not in set(half_real.obs[pert_col].astype(str)):
raise ValueError(
f"control {control!r} has < 2 cells; cannot compute a "
f"disjoint-split data ceiling"
)
Comment thread
LeonHafner marked this conversation as resolved.
return half_real, half_pred

def _write_results(
Expand Down Expand Up @@ -331,14 +363,35 @@ def _spearman_brown_correct(results: pl.DataFrame) -> pl.DataFrame:
Every other column (error metrics, unbounded counts, and reliability metrics
not in that set) is emitted as ``NaN``, since a Spearman-Brown ceiling has no
defensible meaning there.

The correction is applied only where the measured reliability is ``r > 0``.
``2r/(1+r)`` is a reliability correction only on that side; at ``r <= 0`` it is
a pole, not a correction (``r = -0.9`` gives ``-18.0``, and ``r = -1`` divides
by zero - in polars a silent ``-inf`` rather than a raise). Three metrics in
``SB_METRICS`` are sign-unbounded and can land there on a small or degenerate
context: ``pearson_delta``, ``de_spearman_sig`` and ``de_spearman_lfc_sig``. A
non-positive split-half reliability means the halves do not agree at all, i.e.
there is no defensible ceiling, so it is reported as ``NaN`` - never a negative
"ceiling" worse than any achievable score. Null means (a metric that produced
no value) fall through the same branch.

Note this threshold is 0 for every metric, including ``pr_auc`` / ``roc_auc``
whose chance baseline is 0.5 rather than 0; a below-chance AUC is still passed
through the correction. That is deliberate - a 0.5 floor would be a stricter
rule than the one the ceiling was empirically validated under.
"""
nan = float("nan")
exprs: list[pl.Expr] = []
for col in results.columns:
if col == "perturbation":
continue
if col in SB_METRICS:
exprs.append((2.0 * pl.col(col) / (1.0 + pl.col(col))).alias(col))
exprs.append(
pl.when(pl.col(col) > 0.0)
.then(2.0 * pl.col(col) / (1.0 + pl.col(col)))
.otherwise(pl.lit(nan))
.alias(col)
)
else:
exprs.append(pl.lit(nan).alias(col))
return results.with_columns(exprs) if exprs else results
Expand Down
72 changes: 68 additions & 4 deletions tests/test_ceiling.py
Original file line number Diff line number Diff line change
Expand Up @@ -12,8 +12,8 @@


def test_spearman_brown_doubling_on_reliability_metrics():
"""SB doubling r' = 2r/(1+r) is applied to reliability (best_value == ONE)
metric columns."""
"""SB doubling r' = 2r/(1+r) is applied to the reliability metrics named in the
explicit SB_METRICS list."""
df = pl.DataFrame(
{
"perturbation": ["a", "b"],
Expand Down Expand Up @@ -58,6 +58,64 @@ def test_non_reliability_and_excluded_metrics_are_nan():
assert "pearson_edistance" not in SB_METRICS


def test_non_positive_reliability_is_nan():
"""2r/(1+r) is a pole for r <= 0, not a correction: r = -0.9 gives -18.0 and
r = -1 divides by zero (a silent -inf in polars). A non-positive split-half
reliability means the halves do not agree, so there is no defensible ceiling
-> NaN, and never a negative "ceiling" or an inf."""
df = pl.DataFrame(
{
"perturbation": ["a", "b", "c", "d", "e"],
# r = -1 would divide by zero; -0.9 -> -18.0; -0.5 -> -2.0;
# 0.0 is the boundary (the guard is > 0, not >= 0); 0.5 still works.
"pearson_delta": [-1.0, -0.9, -0.5, 0.0, 0.5],
}
)
out = _spearman_brown_correct(df)["pearson_delta"].to_list()
assert all(np.isnan(v) for v in out[:4]), out
assert out[4] == pytest.approx(2 / 3, abs=1e-6) # r > 0 still corrected

# a null mean (metric produced no value) also falls through to NaN
null_df = pl.DataFrame(
{"perturbation": ["a"], "pearson_delta": [None]},
schema={"perturbation": pl.Utf8, "pearson_delta": pl.Float64},
)
assert np.isnan(_spearman_brown_correct(null_df)["pearson_delta"][0])


def test_disjoint_halves_requires_two_cells_and_a_control():
"""The split fails loudly rather than with a bare numpy/pipeline error when
nothing can be split, or when the control specifically cannot be."""
adata_real = build_random_anndata()

def _evaluator(adata):
return MetricsEvaluator(
adata_pred=adata.copy(),
adata_real=adata,
control_pert=CONTROL_VAR,
pert_col=PERT_COL,
outdir=OUTDIR,
skip_de=True,
)

# one cell per perturbation -> nothing is splittable
obs = adata_real.obs
first_of_each = [
int(np.flatnonzero((obs[PERT_COL].astype(str) == p).to_numpy())[0])
for p in obs[PERT_COL].astype(str).unique()
]
with pytest.raises(ValueError, match="no perturbation has >= 2 cells"):
_evaluator(adata_real[first_of_each].copy())._disjoint_halves(seed=0)

# control has a single cell, other perturbations are fine
is_ctrl = (obs[PERT_COL].astype(str) == CONTROL_VAR).to_numpy()
keep = np.concatenate([np.flatnonzero(is_ctrl)[:1], np.flatnonzero(~is_ctrl)])
with pytest.raises(ValueError, match="cannot compute a disjoint-split"):
_evaluator(adata_real[np.sort(keep)].copy())._disjoint_halves(seed=0)

shutil.rmtree(OUTDIR, ignore_errors=True)


def test_disjoint_halves_share_no_cells():
adata_real = build_random_anndata()
evaluator = MetricsEvaluator(
Expand Down Expand Up @@ -103,8 +161,14 @@ def test_compute_ceiling_end_to_end():
# the ceiling is the SB-corrected aggregate (mean over perturbations) - one row
assert agg.height == 1
assert "pearson_delta" in agg.columns
cv = agg["pearson_delta"].to_numpy()
assert np.all(cv <= 1.0 + 1e-9) # SB doubling can never exceed 1
# Two-sided: after the r > 0 guard an SB ceiling is either NaN (non-positive
# reliability, no defensible ceiling) or in (0, 1] - doubling an r in (0, 1] can
# neither exceed 1 nor come out negative. A one-sided `<= 1.0` would silently
# admit a blown-up pole value (r = -0.9 -> -18.0, r = -1 -> -inf).
for col in agg.columns:
if col in SB_METRICS:
v = agg[col].to_numpy()
assert np.all(np.isnan(v) | ((v > 0.0) & (v <= 1.0 + 1e-9))), col

# error metrics and excluded reliabilities have no ceiling -> NaN in the aggregate
for col in ("mse", "mae", "clustering_agreement", "pearson_edistance"):
Expand Down
Loading