Skip to content

Commit 1549fe5

Browse files
committed
Fix NaN propagation bug
1 parent fbe69b1 commit 1549fe5

3 files changed

Lines changed: 14 additions & 5 deletions

File tree

bartab/models/base.py

Lines changed: 12 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -147,12 +147,17 @@ def fit(
147147
valid,
148148
)):
149149
this_base_result = base_result | {"i": i}
150+
error_suffix = f"{_valid.mean()=}, {finite_y.mean()=}, {finite_X.mean()=}"
150151
this_valid = _valid & finite_y & finite_X
151152
if w is not None:
152-
this_valid = this_valid & np.isfinite(w) & (w > 0.)
153+
finite_w = np.isfinite(w)
154+
nonzero_w = (w > 0.)
155+
this_valid = this_valid & finite_w & nonzero_w
156+
error_suffix += f", {finite_w.mean()=}, {nonzero_w.mean()=}"
153157

154-
if this_valid.sum() < min_obs:
155-
results.append(this_base_result)
158+
n_valid_obs = this_valid.sum()
159+
if n_valid_obs < min_obs:
160+
results.append(this_base_result | {"fit_status": f"fail:{n_valid_obs=} < {min_obs=}; {error_suffix}"})
156161
_preds.append((np.array([]), np.full(x.shape, np.nan), np.full(x.shape, np.nan)))
157162
continue
158163
if groups is None:
@@ -188,6 +193,10 @@ def fit(
188193
np.stack([_y for _, _y, _ in _preds]),
189194
np.stack([_p for _, _, _p in _preds]),
190195
)
196+
all_failed = all(d["fit_status"].startswith("fail") for d in results)
197+
if all_failed:
198+
fail_modes = sorted(set(d["fit_status"].removeprefix("fail") for d in results if d["fit_status"].startswith("fail")))
199+
raise ValueError(f"All fits failed: {fail_modes}")
191200
return results, _preds
192201

193202

bartab/models/linear.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -32,7 +32,7 @@ def _delta_method_weights(
3232
True
3333
3434
"""
35-
ref_counts = raw[control_mask, :].sum(axis=0) # (n_samples,)
35+
ref_counts = np.nansum(raw[control_mask, :], axis=0) # (n_samples,)
3636
ref_disp = _estimate_dispersion_mom(ref_counts[None], groups) # scalar
3737
var_y = (
3838
1. / raw + dispersion[:, None] # (n_strains, n_samples)

pyproject.toml

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1,6 +1,6 @@
11
[project]
22
name = "bartab"
3-
version = "0.0.10"
3+
version = "0.0.10.post1"
44
authors = [
55
{ name="Eachan Johnson", email="eachan.johnson@crick.ac.uk" },
66
]

0 commit comments

Comments
 (0)