Update v7.py

This commit is contained in:
PRAKUL HIREMATH
2026-04-10 19:33:20 +05:30
committed by GitHub
parent 74a233a810
commit 877fecdc0f
+11 -11
View File
@@ -443,7 +443,7 @@ def evaluation_metrics(deltas, sigma, tau, max_lag=MAX_LAG):
cov = compute_coverage(tau, sigma, max_lag) cov = compute_coverage(tau, sigma, max_lag)
return dict( return dict(
mean_delta = float(np.mean(deltas)), mean_delta = float(np.mean(deltas)),
pct_early = float(np.mean(valid)), # Precision precision = float(np.mean(valid)), # Precision
recall = float(cov), # Coverage = Recall recall = float(cov), # Coverage = Recall
mean_early = float(np.mean(deltas[valid])) if valid.any() else 0.0, mean_early = float(np.mean(deltas[valid])) if valid.any() else 0.0,
std_delta = float(np.std(deltas)), std_delta = float(np.std(deltas)),
@@ -503,7 +503,7 @@ def threshold_sweep(model, X_scaled, X_raw, sigma,
rows.append({ rows.append({
'pct' : pct, 'pct' : pct,
'mean_delta': m['mean_delta'], 'mean_delta': m['mean_delta'],
'precision' : m['pct_early'], 'precision' : m['precision'],
'recall' : m['coverage'], 'recall' : m['coverage'],
'n_tau' : m['n_tau'], 'n_tau' : m['n_tau'],
'n_early' : m['n_early'], 'n_early' : m['n_early'],
@@ -584,7 +584,7 @@ def run_robustness_grid(model_base, scaler_base,
rows.append({'delay': dname, 'noise': nname, 'strength': sname, rows.append({'delay': dname, 'noise': nname, 'strength': sname,
'rep': rep, 'rep': rep,
'mean_delta': m['mean_delta'], 'mean_delta': m['mean_delta'],
'precision' : m['pct_early'], 'precision' : m['precision'],
'recall' : m['coverage'], 'recall' : m['coverage'],
'n_tau' : m['n_tau'], 'n_tau' : m['n_tau'],
'n_sigma' : len(sig_r)}) 'n_sigma' : len(sig_r)})
@@ -614,14 +614,14 @@ def channel_lead_time_analysis(comps, sigma, Z_true, max_lag=MAX_LAG,
tau = apply_early_detection_constraint(tau, sigma, tau = apply_early_detection_constraint(tau, sigma,
min_lead=min_lead, max_lag=max_lag) min_lead=min_lead, max_lag=max_lag)
if len(tau) == 0: if len(tau) == 0:
results[key] = dict(mean_delta=np.nan, pct_early=0.0, results[key] = dict(mean_delta=np.nan, precision=0.0,
coverage=0.0, n_tau=0) coverage=0.0, n_tau=0)
continue continue
deltas = compute_lead_times(tau, sigma, max_lag) deltas = compute_lead_times(tau, sigma, max_lag)
m = evaluation_metrics(deltas, sigma, tau, max_lag) m = evaluation_metrics(deltas, sigma, tau, max_lag)
results[key] = dict( results[key] = dict(
mean_delta = m['mean_delta'], mean_delta = m['mean_delta'],
pct_early = m['pct_early'], precision = m['precision'],
coverage = m['coverage'], coverage = m['coverage'],
n_tau = m['n_tau'], n_tau = m['n_tau'],
) )
@@ -848,7 +848,7 @@ def plot_channel_diagnostics(comps, diag_results, winner_counts, lead_by_channel
fontsize=12, fontweight="bold") fontsize=12, fontweight="bold")
# Panel A: grouped bar # Panel A: grouped bar
metrics = ['mean_delta', 'pct_early', 'coverage'] metrics = ['mean_delta', 'precision', 'coverage']
m_labels = ['Mean Δ (scaled)', 'Precision', 'Coverage'] m_labels = ['Mean Δ (scaled)', 'Precision', 'Coverage']
x = np.arange(len(keys)) x = np.arange(len(keys))
width = 0.22 width = 0.22
@@ -1091,7 +1091,7 @@ def run_experiment():
"Detector" : name, "Detector" : name,
"Mean Δ" : f"{m['mean_delta']:+.2f}", "Mean Δ" : f"{m['mean_delta']:+.2f}",
"95% CI" : f"[{lo:+.2f}, {hi:+.2f}]", "95% CI" : f"[{lo:+.2f}, {hi:+.2f}]",
"Precision" : f"{m['pct_early']:.1%}", "Precision" : f"{m['precision']:.1%}",
"Coverage" : f"{m['coverage']:.1%}", "Coverage" : f"{m['coverage']:.1%}",
"Mean Δ|early" : f"{m['mean_early']:+.2f}", "Mean Δ|early" : f"{m['mean_early']:+.2f}",
"N(τ)" : m['n_tau'], "N(τ)" : m['n_tau'],
@@ -1118,7 +1118,7 @@ def run_experiment():
pct_range=SWEEP_PCTS, method='adaptive') pct_range=SWEEP_PCTS, method='adaptive')
sweep_mtr = threshold_sweep(model, X_scaled, X_raw, sigma, sweep_mtr = threshold_sweep(model, X_scaled, X_raw, sigma,
pct_range=SWEEP_PCTS, method='multitrigger') pct_range=SWEEP_PCTS, method='multitrigger')
print(f" Sweep complete. Standard: {sweep_std['pct_early'].notna().sum()} " print(f" Sweep complete. Standard: {sweep_std['precision'].notna().sum()} "
f"valid thresholds / {len(SWEEP_PCTS)}") f"valid thresholds / {len(SWEEP_PCTS)}")
# ── 7. Signal diagnostics ──────────────────────────────────── # ── 7. Signal diagnostics ────────────────────────────────────
@@ -1128,7 +1128,7 @@ def run_experiment():
print(" Channel standalone performance:") print(" Channel standalone performance:")
for key, res in diag_results.items(): for key, res in diag_results.items():
print(f" {key:15s}: mean Δ={res['mean_delta']:+.2f} " print(f" {key:15s}: mean Δ={res['mean_delta']:+.2f} "
f"precision={res['pct_early']:.1%} coverage={res['coverage']:.1%} " f"precision={res['precision']:.1%} coverage={res['coverage']:.1%} "
f"n_tau={res['n_tau']}") f"n_tau={res['n_tau']}")
print(" Earliest-trigger counts per channel:") print(" Earliest-trigger counts per channel:")
for key, cnt in sorted(winner_counts.items(), key=lambda x: -x[1]): for key, cnt in sorted(winner_counts.items(), key=lambda x: -x[1]):
@@ -1150,12 +1150,12 @@ def run_experiment():
m_std = evaluation_metrics(delta_dict['Model'], sigma, tau_std) m_std = evaluation_metrics(delta_dict['Model'], sigma, tau_std)
print(f"\n ── FINAL SUMMARY ─────────────────────────────────────────") print(f"\n ── FINAL SUMMARY ─────────────────────────────────────────")
print(f" Model Mean Δ : {m_std['mean_delta']:+.2f} steps") print(f" Model Mean Δ : {m_std['mean_delta']:+.2f} steps")
print(f" Model Precision : {m_std['pct_early']:.1%}") print(f" Model Precision : {m_std['precision']:.1%}")
print(f" Model Coverage : {m_std['coverage']:.1%}") print(f" Model Coverage : {m_std['coverage']:.1%}")
print(f" Model Mean Δ|early: {m_std['mean_early']:+.2f} steps") print(f" Model Mean Δ|early: {m_std['mean_early']:+.2f} steps")
checks = [ checks = [
(m_std['mean_delta'] > 0, "Positive mean lead-time"), (m_std['mean_delta'] > 0, "Positive mean lead-time"),
(m_std['pct_early'] > 0.60, "> 60% precision"), (m_std['precision'] > 0.60, "> 60% precision"),
(m_std['coverage'] > 0.30, "> 30% coverage"), (m_std['coverage'] > 0.30, "> 30% coverage"),
] ]
for ok, msg in checks: for ok, msg in checks: