falsesync 0.1.2__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
falsesync/features.py ADDED
@@ -0,0 +1,164 @@
1
+ """Feature extraction for regime classification (Phase 3).
2
+
3
+ Every feature is a rescaled, roughly O(1) quantity so the multinomial
4
+ logistic classifier is numerically stable. Features are documented in
5
+ diagnostic_specification.md §4 (s1..s4) plus Phase-3 additions.
6
+ """
7
+
8
+ from __future__ import annotations
9
+
10
+ import numpy as np
11
+
12
+ from .aggregation import weighted_aggregate
13
+ from .changepoints import aggregate_breakpoint, ls_breakpoint, multibreak
14
+ from .inference import deconvolve_timing, unit_break_uncertainties
15
+
16
+ FEATURE_NAMES = [
17
+ "agg_strength", "excess_sd_ratio", "n_modes", "mode_mass2",
18
+ "s1_width", "s3_alignment", "s4_weighting", "noise_floor_ratio",
19
+ "fwhm_ratio", "agg_skew", "log_n", "log_T", "frac_valid_units",
20
+ "binseg_n_breaks", "unit_strength_med",
21
+ ]
22
+
23
+
24
+ def _kde(samples, grid, bw):
25
+ u = (grid[:, None] - samples[None, :]) / bw
26
+ return np.exp(-0.5 * u * u).mean(axis=1) / (bw * np.sqrt(2 * np.pi))
27
+
28
+
29
+ def _mode_count(f, thr=0.3, min_sep_frac=0.08):
30
+ """Count modes with prominence threshold and minimum separation.
31
+
32
+ A mode must exceed `thr` * max(f) locally and be at least
33
+ `min_sep_frac` * len(f) grid points away from a higher mode
34
+ (greedy, highest peak first).
35
+ """
36
+ if f.size < 3 or f.max() <= 0:
37
+ return 0, []
38
+ cand = [i for i in range(1, len(f) - 1)
39
+ if f[i] >= f[i - 1] and f[i] > f[i + 1] and f[i] > thr * f.max()]
40
+ cand.sort(key=lambda i: -f[i])
41
+ sep = max(1, int(min_sep_frac * len(f)))
42
+ chosen = []
43
+ for i in cand:
44
+ if all(abs(i - j) >= sep for j in chosen):
45
+ chosen.append(i)
46
+ return len(chosen), sorted(chosen)
47
+
48
+
49
+ def extract_features(
50
+ t: np.ndarray,
51
+ values: np.ndarray,
52
+ weights: np.ndarray | None = None,
53
+ windows: np.ndarray | None = None,
54
+ aggregate: np.ndarray | None = None,
55
+ rng: np.random.Generator | None = None,
56
+ unit_boot: int = 30,
57
+ ) -> tuple[np.ndarray, dict]:
58
+ """Feature vector + aux dict used by the classifier and the result object."""
59
+ rng = rng or np.random.default_rng(0)
60
+ t = np.asarray(t, dtype=float)
61
+ n = values.shape[0]
62
+ if aggregate is None:
63
+ aggregate = weighted_aggregate(values, weights)
64
+ valid_u = np.isfinite(values).sum(axis=1) >= 8
65
+ agg_f = np.isfinite(aggregate)
66
+ agg_y = np.nan_to_num(aggregate, nan=np.nanmean(aggregate[agg_f])) if agg_f.any() else np.zeros(t.size)
67
+
68
+ # aggregate break
69
+ res_agg = ls_breakpoint(t, agg_y)
70
+ dm = np.gradient(agg_y, t)
71
+ half = np.max(dm) / 2
72
+ idx = np.where(dm >= half)[0]
73
+ w_agg = float(t[idx[-1]] - t[idx[0]]) if idx.size > 1 else np.inf
74
+
75
+ # unit breaks
76
+ taus = np.full(n, np.nan)
77
+ strengths = np.full(n, np.nan)
78
+ for i in np.where(valid_u)[0]:
79
+ y = values[i]
80
+ fill = np.nanmean(y[np.isfinite(y)])
81
+ try:
82
+ r = ls_breakpoint(t, np.nan_to_num(y, nan=fill))
83
+ taus[i] = r.location
84
+ strengths[i] = r.strength
85
+ except Exception:
86
+ pass
87
+ tv = taus[np.isfinite(taus)]
88
+ sv = strengths[np.isfinite(taus)]
89
+
90
+ s_i = unit_break_uncertainties(t, values[np.isfinite(taus)], tv, rng=rng, n_boot=unit_boot) if tv.size >= 5 else np.array([])
91
+ noise_floor = float(np.nanmedian(s_i)) if s_i.size and np.isfinite(s_i).any() else np.nan
92
+
93
+ tau_sd = float(np.std(tv)) if tv.size else np.nan
94
+ excess = np.sqrt(max(tau_sd**2 - (noise_floor if np.isfinite(noise_floor) else 0) ** 2, 0)) if np.isfinite(tau_sd) else np.nan
95
+ scale = np.ptp(t)
96
+ excess_ratio = float(excess / scale) if np.isfinite(excess) else np.nan
97
+ floor_ratio = float(noise_floor / max(tau_sd, 1e-9)) if np.isfinite(noise_floor) and np.isfinite(tau_sd) else np.nan
98
+
99
+ grid = np.linspace(t.min(), t.max(), 300)
100
+ bw = max(1.06 * (tau_sd if np.isfinite(tau_sd) else 1.0) * max(len(tv), 2) ** -0.2, 3 * np.median(np.diff(t)))
101
+ f_hat = _kde(tv, grid, bw) if tv.size >= 5 else np.zeros(grid.size)
102
+ f_dec_raw = deconvolve_timing(tv, s_i, grid) if tv.size >= 5 else np.zeros(grid.size)
103
+ # smooth deconvolved density before mode counting (deconvolution is noisy)
104
+ k = max(3, int(0.04 * grid.size) | 1)
105
+ kern = np.exp(-0.5 * (np.arange(k) - k // 2) ** 2 / (k / 4) ** 2)
106
+ kern /= kern.sum()
107
+ f_dec = np.convolve(np.pad(f_dec_raw, k // 2, mode="edge"), kern, mode="valid")
108
+ f_dec = np.clip(f_dec, 0, None)
109
+ # modes counted on the (smoother) KDE; deconvolved density is used for
110
+ # dispersion correction, where its noise matters less
111
+ n_modes, modes = _mode_count(f_hat, thr=0.3, min_sep_frac=0.08)
112
+ mode_mass2 = float(np.sort(f_dec)[::-1][:2].sum() / max(f_dec.sum(), 1e-12)) if f_dec.sum() > 0 else 0.0
113
+
114
+ s1 = float(tau_sd / w_agg) if np.isfinite(tau_sd) and np.isfinite(w_agg) and w_agg > 0 else np.nan
115
+ s3 = float(np.mean(np.abs(sv) > 2.0)) if sv.size else np.nan
116
+
117
+ s4 = 0.0
118
+ if weights is not None and tv.size >= 5:
119
+ wv = np.asarray(weights)
120
+ if wv.ndim == 2:
121
+ wv = wv.mean(axis=1)
122
+ try:
123
+ s4 = max(s4, abs(float(np.corrcoef(wv[np.isfinite(taus)], tv)[0, 1])))
124
+ except Exception:
125
+ pass
126
+ if windows is not None and tv.size >= 5:
127
+ span = np.asarray(windows)[np.isfinite(taus)]
128
+ span = span[:, 1] - span[:, 0]
129
+ try:
130
+ s4 = max(s4, abs(float(np.corrcoef(span, tv)[0, 1])))
131
+ except Exception:
132
+ pass
133
+
134
+ agg_skew = float(((agg_y - agg_y.mean()) ** 3).mean() / max(agg_y.std(), 1e-9) ** 3)
135
+ try:
136
+ nb = len(multibreak(t, agg_y, n_bkps=5).extras.get("all_break_indices", []))
137
+ except Exception:
138
+ nb = 1
139
+
140
+ feats = np.array([
141
+ res_agg.strength,
142
+ excess_ratio if np.isfinite(excess_ratio) else 0.0,
143
+ float(n_modes),
144
+ mode_mass2,
145
+ s1 if np.isfinite(s1) else 0.0,
146
+ s3 if np.isfinite(s3) else 0.0,
147
+ float(s4),
148
+ floor_ratio if np.isfinite(floor_ratio) else 1.0,
149
+ float(w_agg / scale) if np.isfinite(w_agg) else 1.0,
150
+ agg_skew,
151
+ float(np.log(max(n, 1))),
152
+ float(np.log(t.size)),
153
+ float(valid_u.mean()),
154
+ float(nb),
155
+ float(np.nanmedian(np.abs(sv))) if sv.size else 0.0,
156
+ ])
157
+ aux = {
158
+ "agg_break": res_agg.location, "agg_strength": res_agg.strength,
159
+ "taus": taus, "unit_strengths": strengths, "s_i": s_i,
160
+ "noise_floor": noise_floor, "tau_sd": tau_sd, "f_dec": (grid, f_dec),
161
+ "f_hat": (grid, f_hat), "w_agg": w_agg, "n_modes": n_modes,
162
+ "aggregate": agg_y,
163
+ }
164
+ return feats, aux
falsesync/inference.py ADDED
@@ -0,0 +1,90 @@
1
+ """Uncertainty: unit-level bootstrap and error-law deconvolution.
2
+
3
+ See resampling_plan.md and deconvolution_assessment.md. No asymptotic
4
+ coverage claims — intervals are bootstrap intervals.
5
+ """
6
+
7
+ from __future__ import annotations
8
+
9
+ import numpy as np
10
+
11
+
12
+ def unit_break_uncertainties(
13
+ t: np.ndarray,
14
+ values: np.ndarray,
15
+ taus: np.ndarray,
16
+ rng: np.random.Generator | None = None,
17
+ n_boot: int = 60,
18
+ ) -> np.ndarray:
19
+ """Residual bootstrap around each unit's two-segment fit -> s_i (the F_e law)."""
20
+ from .changepoints import ls_breakpoint
21
+
22
+ rng = rng or np.random.default_rng(0)
23
+ t = np.asarray(t, dtype=float)
24
+ n = values.shape[0]
25
+ out = np.full(n, np.nan)
26
+ for i in range(n):
27
+ y = np.asarray(values[i], dtype=float)
28
+ ok = np.isfinite(y)
29
+ if ok.sum() < 8:
30
+ continue
31
+ ti, yi = t[ok], y[ok]
32
+ res = ls_breakpoint(ti, yi)
33
+ c = res.index
34
+ fit = np.concatenate([np.full(c, yi[:c].mean()), np.full(yi.size - c, yi[c:].mean())])
35
+ resid = yi - fit
36
+ resid = resid - resid.mean()
37
+ boots = []
38
+ for _ in range(n_boot):
39
+ yb = fit + rng.choice(resid, resid.size, replace=True)
40
+ try:
41
+ boots.append(ls_breakpoint(ti, yb).location)
42
+ except Exception:
43
+ continue
44
+ if len(boots) >= 10:
45
+ out[i] = float(np.std(boots))
46
+ return out
47
+
48
+
49
+ def deconvolve_timing(
50
+ taus: np.ndarray, s_i: np.ndarray, grid: np.ndarray, damp: float = 0.02
51
+ ) -> np.ndarray:
52
+ """Recover f_tau from noisy unit break estimates (P8c machinery).
53
+
54
+ Gaussian errors with per-unit sd s_i: f_hat_tau = f_tau * phi_bar, where
55
+ phi_bar is the average error kernel. Damped Fourier deconvolution with
56
+ fixed Tikhonov factor `damp` (a priori, not tuned).
57
+ """
58
+ taus = np.asarray(taus, dtype=float)
59
+ s_i = np.asarray(s_i, dtype=float)
60
+ grid = np.asarray(grid, dtype=float)
61
+ dt = grid[1] - grid[0]
62
+ n_g = grid.size
63
+ hist, _ = np.histogram(taus, bins=np.append(grid, grid[-1] + dt), density=True)
64
+ sbar2 = float(np.nanmean(s_i**2)) if np.isfinite(s_i).any() else 0.0
65
+ F_hist = np.fft.rfft(hist)
66
+ freqs = np.fft.rfftfreq(n_g, dt)
67
+ phi_bar = np.exp(-2 * np.pi**2 * freqs**2 * sbar2) # E over units of error char. fn (approx: mean sd)
68
+ denom = phi_bar**2 + damp
69
+ F_f = F_hist * phi_bar / denom
70
+ f = np.fft.irfft(F_f, n_g).real
71
+ return np.clip(f, 0, None)
72
+
73
+
74
+ def cluster_bootstrap_band(
75
+ taus: np.ndarray,
76
+ grid: np.ndarray,
77
+ bw: float,
78
+ rng: np.random.Generator | None = None,
79
+ n_boot: int = 200,
80
+ ) -> tuple[np.ndarray, np.ndarray]:
81
+ """Cluster (unit-resample) bootstrap band for the KDE of f_tau."""
82
+ rng = rng or np.random.default_rng(0)
83
+ taus = np.asarray(taus, dtype=float)
84
+ dens = []
85
+ for _ in range(n_boot):
86
+ s = rng.choice(taus, taus.size, replace=True)
87
+ u = (grid[:, None] - s[None, :]) / bw
88
+ dens.append(np.exp(-0.5 * u * u).mean(axis=1) / (bw * np.sqrt(2 * np.pi)))
89
+ d = np.stack(dens)
90
+ return np.quantile(d, 0.025, axis=0), np.quantile(d, 0.975, axis=0)
falsesync/model.py ADDED
@@ -0,0 +1,221 @@
1
+ """FalseSynchronyModel — Phase-3 research-grade workflow (spec §17).
2
+
3
+ from falsesync.model import FalseSynchronyModel
4
+ model = FalseSynchronyModel(classifier=clf)
5
+ result = model.fit(t, values, weights=None, windows=None)
6
+
7
+ `classifier` is a fitted falsesync.regimes.RegimeClassifier. If None, the
8
+ result carries uncalibrated heuristic probabilities and a warning.
9
+ """
10
+
11
+ from __future__ import annotations
12
+
13
+ from dataclasses import dataclass, field
14
+
15
+ import numpy as np
16
+
17
+ from .aggregation import weighted_aggregate
18
+ from .alignment import event_time_realign, realign_summary
19
+ from .changepoints import aggregate_breakpoint, ls_breakpoint, multibreak
20
+ from .diagnostics import FalseSynchronyResult
21
+ from .features import extract_features
22
+ from .regimes import REGIME_CLASSES, RegimeClassifier
23
+
24
+
25
+ @dataclass
26
+ class FalseSynchronyModel:
27
+ classifier: RegimeClassifier | None = None
28
+ method: str = "ls"
29
+ operators: tuple = ("ls", "maxslope", "cusum", "binseg")
30
+ unit_boot: int = 30
31
+ agg_strength_crit: float = 3.0 # aggregate break evidence criterion
32
+ sync_prob_crit: float = 0.5 # calibrated synchrony-support threshold
33
+ seed: int = 0
34
+
35
+ def _trend_check(self, t, values, taus, rng) -> tuple[bool, dict]:
36
+ """Detrend each unit (per-unit linear fit), refit breaks, compare.
37
+
38
+ If unit break dispersion collapses after detrending, part of the
39
+ apparent timing heterogeneity is trend confounding (§7).
40
+ """
41
+ t = np.asarray(t, dtype=float)
42
+ tc = (t - t.mean()) / max(np.ptp(t), 1e-9)
43
+ taus_d = []
44
+ valid = np.isfinite(taus)
45
+ for i in np.where(valid)[0]:
46
+ y = values[i]
47
+ ok = np.isfinite(y)
48
+ if ok.sum() < 8:
49
+ continue
50
+ slope = np.polyfit(tc[ok], y[ok], 1)[0]
51
+ yd = y.copy()
52
+ yd[ok] = y[ok] - slope * tc[ok]
53
+ fill = np.nanmean(yd[ok])
54
+ try:
55
+ r = ls_breakpoint(t, np.nan_to_num(yd, nan=fill))
56
+ taus_d.append((i, r.location, abs(slope)))
57
+ except Exception:
58
+ pass
59
+ if not taus_d:
60
+ return False, {"detrended_sd": np.nan, "raw_sd": np.nan}
61
+ idx = np.array([d[0] for d in taus_d], dtype=int)
62
+ raw_sd = float(np.std(taus[idx]))
63
+ det_sd = float(np.std([d[1] for d in taus_d]))
64
+ med_slope = float(np.median([d[2] for d in taus_d]))
65
+ info = {"raw_sd": raw_sd, "detrended_sd": det_sd, "median_slope": med_slope}
66
+ # confounding: detrending materially shrinks spread AND slopes nonzero
67
+ confounded = (raw_sd > 0 and det_sd < 0.6 * raw_sd and med_slope > 0.05)
68
+ return confounded, info
69
+
70
+ def _operator_spread(self, t, agg_y) -> dict:
71
+ locs, strengths = {}, {}
72
+ for m in self.operators:
73
+ try:
74
+ r = aggregate_breakpoint(t, agg_y, m) if m != "binseg" else multibreak(t, agg_y, 3)
75
+ locs[m] = r.location
76
+ strengths[m] = r.strength
77
+ except Exception:
78
+ pass
79
+ vals = np.array(list(locs.values()))
80
+ spread = float(vals.max() - vals.min()) if vals.size > 1 else 0.0
81
+ consensus = float(1 - spread / max(np.ptp(t), 1e-9))
82
+ return {
83
+ "locations": locs, "strengths": strengths,
84
+ "operator_spread": spread,
85
+ "operator_consensus": consensus,
86
+ "operator_sensitivity_warning": spread > 0.2 * np.ptp(t),
87
+ }
88
+
89
+ def fit(
90
+ self,
91
+ t: np.ndarray,
92
+ values: np.ndarray,
93
+ weights: np.ndarray | None = None,
94
+ windows: np.ndarray | None = None,
95
+ aggregate: np.ndarray | None = None,
96
+ ) -> FalseSynchronyResult:
97
+ rng = np.random.default_rng(self.seed)
98
+ t = np.asarray(t, dtype=float)
99
+ if aggregate is None:
100
+ aggregate = weighted_aggregate(values, weights)
101
+ if self.method == "ls":
102
+ prim_loc, prim_str = None, None # taken from aux below
103
+ else:
104
+ prim = (multibreak(t, aggregate, 3) if self.method == "binseg"
105
+ else aggregate_breakpoint(t, aggregate, self.method))
106
+ prim_loc, prim_str = prim.location, prim.strength
107
+ feats, aux = extract_features(
108
+ t, values, weights=weights, windows=windows,
109
+ aggregate=aggregate, rng=rng, unit_boot=self.unit_boot,
110
+ )
111
+ warnings: list[str] = []
112
+ if self.classifier is not None and self.classifier.fitted:
113
+ probs = self.classifier.predict_proba(feats)[0]
114
+ regime_probabilities = dict(zip(REGIME_CLASSES, probs.tolist()))
115
+ else:
116
+ regime_probabilities = {k: np.nan for k in REGIME_CLASSES}
117
+ warnings.append("no calibrated classifier: regime probabilities not estimable")
118
+
119
+ ops = self._operator_spread(t, aux["aggregate"])
120
+ if ops["operator_sensitivity_warning"]:
121
+ warnings.append(
122
+ f"operator spread {ops['operator_spread']:.3g} > 20% of window: "
123
+ "no method-independent aggregate breakpoint (P7); report operator."
124
+ )
125
+
126
+ indeterminate = False
127
+ obs_corr = np.nan
128
+ if windows is not None:
129
+ win = np.asarray(windows, dtype=float)
130
+ if win.ndim == 2 and win.shape[1] == 2:
131
+ lo = np.nanmin(t)
132
+ trunc_start = win[:, 0] > lo + 0.02 * max(np.ptp(t), 1e-9)
133
+ m_ok = trunc_start & np.isfinite(aux["taus"])
134
+ if m_ok.sum() >= 5 and trunc_start.mean() >= 0.2:
135
+ obs_corr = float(np.corrcoef(win[m_ok, 0], aux["taus"][m_ok])[0, 1])
136
+ indeterminate = obs_corr >= 0.5
137
+
138
+ confounded, trend_info = self._trend_check(t, values, aux["taus"], rng)
139
+ if confounded:
140
+ warnings.append(
141
+ "trend confounding: detrending shrinks tau_hat dispersion "
142
+ f"{trend_info['raw_sd']:.3g} -> {trend_info['detrended_sd']:.3g}; "
143
+ "timing inference unreliable until trends modeled (§7)."
144
+ )
145
+
146
+ if indeterminate:
147
+ warnings.append(
148
+ "indeterminate observation process: unit entry times correlate "
149
+ f"with estimated break times (corr={obs_corr:.2f}); synchrony "
150
+ "inference is not supported under a possibly informative "
151
+ "observation process (Phase-4 F1 mitigation)."
152
+ )
153
+ agg_loc = aux["agg_break"] if self.method == "ls" else prim_loc
154
+ agg_str = aux["agg_strength"] if self.method == "ls" else prim_str
155
+ p_sync = regime_probabilities.get("SYNCHRONOUS", np.nan)
156
+ warn_fire = (
157
+ np.isfinite(agg_str) and agg_str >= self.agg_strength_crit
158
+ and np.isfinite(p_sync) and p_sync < self.sync_prob_crit
159
+ ) or indeterminate
160
+ if warn_fire:
161
+ warnings.append(
162
+ f"false common-event warning: aggregate strength "
163
+ f"{agg_str:.2f} >= {self.agg_strength_crit} but "
164
+ f"P(SYNCHRONOUS)={p_sync:.2f} < {self.sync_prob_crit}."
165
+ )
166
+
167
+ tv = aux["taus"][np.isfinite(aux["taus"])]
168
+ grid, f_dec = aux["f_dec"]
169
+ comps = {
170
+ "s1_width": feats[4], "s2_shape": feats[2],
171
+ "s3_alignment": feats[5], "s4_weighting": feats[6],
172
+ "noise_floor_ratio": feats[7],
173
+ }
174
+ scale = np.ptp(t)
175
+ a_sync = p_sync if np.isfinite(p_sync) else np.nan
176
+ synchrony_score = a_sync # calibrated: score IS P(synchrony)
177
+
178
+ aligned = event_time_realign(t, values[np.isfinite(aux["taus"])], tv) if tv.size else np.zeros((0, 0))
179
+ s_dt = float(np.median(np.diff(t)))
180
+ ea = realign_summary(aligned, s_dt) if aligned.size else {}
181
+ ea.update({"n_aligned": int(tv.size),
182
+ "tau_spread_before": float(np.std(tv)) if tv.size else np.nan,
183
+ "est_noise_floor": aux["noise_floor"],
184
+ "trend_confounded": confounded, "trend_info": trend_info})
185
+
186
+ timing_dispersion = {
187
+ "sd": float(np.std(tv)) if tv.size else np.nan,
188
+ "iqr": float(np.quantile(tv, 0.75) - np.quantile(tv, 0.25)) if tv.size else np.nan,
189
+ "sd_corrected": float(np.sqrt(max(np.std(tv) ** 2 - (aux["noise_floor"] or 0) ** 2, 0)))
190
+ if tv.size else np.nan,
191
+ }
192
+ cluster_structure = {
193
+ "n_components": aux["n_modes"],
194
+ "centers": [float(grid[i]) for i in np.argsort(f_dec)[::-1][: aux["n_modes"]]] if f_dec.size else [],
195
+ }
196
+ return FalseSynchronyResult(
197
+ aggregate_breakpoint=agg_loc,
198
+ aggregate_break_strength=agg_str,
199
+ unit_breakpoints=aux["taus"],
200
+ unit_breakpoint_uncertainty=aux["s_i"],
201
+ timing_distribution={"grid": grid, "density": aux["f_hat"][1], "deconvolved": f_dec},
202
+ timing_dispersion=timing_dispersion,
203
+ cluster_structure=cluster_structure,
204
+ synchrony_score=synchrony_score,
205
+ regime_probabilities=regime_probabilities,
206
+ event_aligned_summary=ea,
207
+ interpretation_warning=warnings,
208
+ method=self.method,
209
+ components={
210
+ **comps,
211
+ "operator_locations": ops["locations"],
212
+ "operator_strengths": ops["strengths"],
213
+ "operator_spread": ops["operator_spread"],
214
+ "operator_consensus": ops["operator_consensus"],
215
+ "operator_sensitivity_warning": ops["operator_sensitivity_warning"],
216
+ "trend_confounding_warning": confounded,
217
+ "indeterminate_observation_process": indeterminate,
218
+ "obs_process_corr": obs_corr,
219
+ "false_common_event_warning": warn_fire,
220
+ },
221
+ )
falsesync/plotting.py ADDED
@@ -0,0 +1,59 @@
1
+ """Figures for Phase 2 outputs (matplotlib, non-interactive)."""
2
+
3
+ from __future__ import annotations
4
+
5
+ import matplotlib
6
+
7
+ matplotlib.use("Agg")
8
+ import matplotlib.pyplot as plt
9
+ import numpy as np
10
+
11
+
12
+ def aggregate_curve_figure(t, m_pop, agg_obs, path, title="Aggregate curve"):
13
+ fig, ax = plt.subplots(figsize=(7, 4))
14
+ ax.plot(t, m_pop, label="population m(t)", lw=2)
15
+ if agg_obs is not None:
16
+ ax.plot(t, agg_obs, label="observed aggregate", alpha=0.7)
17
+ ax.set_xlabel("t")
18
+ ax.set_title(title)
19
+ ax.legend()
20
+ fig.tight_layout()
21
+ fig.savefig(path, dpi=140)
22
+ plt.close(fig)
23
+
24
+
25
+ def timing_distribution_figure(grid, densities: dict, path, title="Timing distribution"):
26
+ fig, ax = plt.subplots(figsize=(7, 4))
27
+ for label, d in densities.items():
28
+ ax.plot(grid, d, label=label)
29
+ ax.set_xlabel("tau")
30
+ ax.set_title(title)
31
+ ax.legend()
32
+ fig.tight_layout()
33
+ fig.savefig(path, dpi=140)
34
+ plt.close(fig)
35
+
36
+
37
+ def operator_comparison_figure(t, m, results: dict, path):
38
+ fig, ax = plt.subplots(figsize=(7, 4))
39
+ ax.plot(t, m, lw=1.5, label="m(t)")
40
+ for name, loc in results.items():
41
+ ax.axvline(loc, ls="--", label=f"{name}: {loc:.2f}")
42
+ ax.set_xlabel("t")
43
+ ax.legend(fontsize=8)
44
+ fig.tight_layout()
45
+ fig.savefig(path, dpi=140)
46
+ plt.close(fig)
47
+
48
+
49
+ def aligned_panel_figure(s_grid, aligned, path, title="Event-time aligned panel"):
50
+ fig, ax = plt.subplots(figsize=(7, 4))
51
+ for row in aligned:
52
+ ax.plot(s_grid, row, color="0.7", lw=0.5)
53
+ ax.plot(s_grid, np.nanmean(aligned, axis=0), color="C0", lw=2, label="aligned mean")
54
+ ax.set_xlabel("event time s = t - tau_i")
55
+ ax.set_title(title)
56
+ ax.legend()
57
+ fig.tight_layout()
58
+ fig.savefig(path, dpi=140)
59
+ plt.close(fig)
falsesync/regimes.py ADDED
@@ -0,0 +1,124 @@
1
+ """Calibrated regime classifier (Phase 3).
2
+
3
+ Multinomial logistic regression on simulation-grid features, probability-
4
+ calibrated by temperature scaling fitted on the calibration grid itself.
5
+ The locked evaluation grid is never touched during fitting.
6
+ """
7
+
8
+ from __future__ import annotations
9
+
10
+ from dataclasses import dataclass, field
11
+
12
+ import numpy as np
13
+
14
+ REGIME_CLASSES = [
15
+ "SYNCHRONOUS", "NEAR_SYNCHRONOUS", "DIFFUSE_ASYNCHRONOUS",
16
+ "CLUSTERED", "NO_TRANSITION",
17
+ ]
18
+
19
+
20
+ def regime_label(truth: str, taus: np.ndarray, cluster_gap: float = 1.5) -> str:
21
+ """Map engine truth + sampled taus to classifier classes."""
22
+ if truth == "no_transition":
23
+ return "NO_TRANSITION"
24
+ if truth == "synchronous":
25
+ return "SYNCHRONOUS"
26
+ if truth == "near_synchronous":
27
+ return "NEAR_SYNCHRONOUS"
28
+ if truth in ("two_cluster", "three_cluster"):
29
+ return "CLUSTERED"
30
+ if truth == "diffuse":
31
+ return "DIFFUSE_ASYNCHRONOUS"
32
+ return truth
33
+
34
+
35
+ @dataclass
36
+ class RegimeClassifier:
37
+ """Logistic + temperature-calibrated probabilities over REGIME_CLASSES."""
38
+
39
+ C: float = 1.0
40
+ max_iter: int = 2000
41
+ model: object | None = None
42
+ scaler_mean: np.ndarray | None = None
43
+ scaler_sd: np.ndarray | None = None
44
+ temperature: float = 1.0
45
+ fitted: bool = False
46
+ meta: dict = field(default_factory=dict)
47
+
48
+ def _fit_logreg(self, X, y):
49
+ from sklearn.linear_model import LogisticRegression
50
+ from sklearn.preprocessing import LabelEncoder
51
+
52
+ self._le = LabelEncoder().fit(y)
53
+ yv = self._le.transform(y)
54
+ self.scaler_mean = X.mean(axis=0)
55
+ self.scaler_sd = X.std(axis=0)
56
+ self.scaler_sd[self.scaler_sd == 0] = 1.0
57
+ Xs = (X - self.scaler_mean) / self.scaler_sd
58
+ self.model = LogisticRegression(
59
+ C=self.C, max_iter=self.max_iter, multi_class="multinomial",
60
+ class_weight="balanced",
61
+ )
62
+ self.model.fit(Xs, yv)
63
+ self.fitted = True
64
+ self._classes_in_model = list(self._le.classes_)
65
+
66
+ def fit(self, X: np.ndarray, y: np.ndarray) -> "RegimeClassifier":
67
+ self._fit_logreg(np.asarray(X), np.asarray(y))
68
+ # temperature scaling on the same (held-out-to-temp data should be
69
+ # supplied separately via calibrate() when available)
70
+ return self
71
+
72
+ def calibrate(self, X: np.ndarray, y: np.ndarray) -> float:
73
+ """Fit temperature on a held-out calibration split; returns T."""
74
+ logits = self._logits(X)
75
+ best_t, best_loss = 1.0, np.inf
76
+ for T_c in np.geomspace(0.2, 5.0, 60):
77
+ p = self._softmax(logits / T_c)
78
+ yi = self._le.transform(np.asarray(y))
79
+ loss = -np.log(np.clip(p[np.arange(len(yi)), yi], 1e-12, 1)).mean()
80
+ if loss < best_loss:
81
+ best_loss, best_t = loss, T_c
82
+ self.temperature = float(best_t)
83
+ return self.temperature
84
+
85
+ def _logits(self, X: np.ndarray) -> np.ndarray:
86
+ Xs = (np.asarray(X) - self.scaler_mean) / self.scaler_sd
87
+ return self.model.decision_function(Xs)
88
+
89
+ @staticmethod
90
+ def _softmax(z: np.ndarray) -> np.ndarray:
91
+ z = z - z.max(axis=1, keepdims=True)
92
+ e = np.exp(z)
93
+ return e / e.sum(axis=1, keepdims=True)
94
+
95
+ def predict_proba(self, X: np.ndarray) -> np.ndarray:
96
+ """(n, len(REGIME_CLASSES)) probability matrix, columns in class order."""
97
+ p = self._softmax(self._logits(np.atleast_2d(X)) / self.temperature)
98
+ out = np.zeros((p.shape[0], len(REGIME_CLASSES)))
99
+ for j, cls in enumerate(self._classes_in_model):
100
+ if cls in REGIME_CLASSES:
101
+ out[:, REGIME_CLASSES.index(cls)] = p[:, j]
102
+ return out
103
+
104
+ def predict(self, X: np.ndarray) -> np.ndarray:
105
+ p = self.predict_proba(X)
106
+ return np.array(REGIME_CLASSES)[p.argmax(axis=1)]
107
+
108
+
109
+ def multiclass_brier(P: np.ndarray, y_idx: np.ndarray, n_classes: int) -> float:
110
+ Y = np.zeros((len(y_idx), n_classes))
111
+ Y[np.arange(len(y_idx)), y_idx] = 1.0
112
+ return float(np.mean(np.sum((P - Y) ** 2, axis=1)))
113
+
114
+
115
+ def expected_calibration_error(P: np.ndarray, y_idx: np.ndarray, n_bins: int = 10) -> float:
116
+ conf = P.max(axis=1)
117
+ pred = P.argmax(axis=1)
118
+ ece = 0.0
119
+ for b in range(n_bins):
120
+ lo, hi = b / n_bins, (b + 1) / n_bins
121
+ m = (conf > lo) & (conf <= hi)
122
+ if m.any():
123
+ ece += m.mean() * abs((pred[m] == y_idx[m]).mean() - conf[m].mean())
124
+ return float(ece)