"""Frozen hierarchical Wiener first-passage residual for the public V05 planner. This module is inference-only. It estimates a battery's terminal drift from twelve non-overlapping weekly increments of the evaluator-exact smoothed median voltage, shrinks that drift toward a leave-target-building pool, and uses the analytic Wiener first-passage probability as a 15% rank residual. It never changes V05's probability multiset, RUL heads, quota, timing policy, or planner. """ from __future__ import annotations import math from dataclasses import dataclass, field import numpy as np import pandas as pd from scipy.stats import norm, rankdata from .competition_planner import CompetitionPlanner from .identity_ensemble import ( CausalHistoryCache, CausalHistoryView, planner_freshness_factors, ) SCHEMA_VERSION = 1 HORIZON_DAYS = 42.0 EOL_VOLTAGE = 2.40 TERMINAL_WEEKS = 12 MIN_WEEKLY_INCREMENTS = 4 MAX_STATE_STALENESS_DAYS = 7.0 RANK_RESIDUAL_WEIGHT = 0.15 @dataclass(frozen=True) class FPTSnapshotFeatures: last_smooth_voltage: np.ndarray smooth_staleness_days: np.ndarray weekly_increment_count: np.ndarray eb_drift_v_per_day: np.ndarray pooled_diffusion_v_per_sqrt_day: np.ndarray eb_device_weight: np.ndarray hit_probability_42d: np.ndarray reliable: np.ndarray def terminal_weekly_increments( series: pd.Series, cutoff: pd.Timestamp | str, ) -> tuple[float, float, np.ndarray]: """Return the last pre-cut state and twelve disjoint seven-day changes. The strict calendar-day cut matters because official scenarios start at midnight. A cutoff-day aggregate could otherwise include later readings from that same day when a full series is used for an offline audit. """ cutoff_day = pd.Timestamp(cutoff).normalize() if series.empty: return math.nan, math.inf, np.asarray([], dtype=float) work = pd.Series( pd.to_numeric(series, errors="coerce").to_numpy(float), index=pd.DatetimeIndex(series.index).normalize(), ).sort_index(kind="stable") if work.index.has_duplicates: raise ValueError("smoothed voltage series has duplicate calendar days") past = work.loc[work.index < cutoff_day].dropna() if past.empty: return math.nan, math.inf, np.asarray([], dtype=float) last_day = pd.Timestamp(past.index[-1]) anchors = pd.DatetimeIndex( [ last_day - pd.Timedelta(days=7 * offset) for offset in range(TERMINAL_WEEKS, -1, -1) ] ) values = work.reindex(anchors).to_numpy(float) increments = np.diff(values) increments = increments[np.isfinite(increments)] staleness = float((cutoff_day - last_day) / pd.Timedelta(days=1)) return float(past.iloc[-1]), staleness, increments def _robust_scale(values: np.ndarray) -> float: values = np.asarray(values, dtype=float) center = float(np.median(values)) return float(1.4826 * np.median(np.abs(values - center))) def wiener_hit_probability( initial_voltage: float, drift_v_per_day: float, diffusion_v_per_sqrt_day: float, horizon_days: float = HORIZON_DAYS, ) -> float: """Analytic probability that drifted Brownian voltage hits 2.40 V.""" distance = float(initial_voltage - EOL_VOLTAGE) if distance <= 0.0: return 1.0 diffusion = max(float(diffusion_v_per_sqrt_day), 1e-6) horizon = float(horizon_days) scale = diffusion * math.sqrt(horizon) drift = float(drift_v_per_day) first = norm.cdf((-drift * horizon - distance) / scale) log_second = ( -2.0 * drift * distance / (diffusion * diffusion) + norm.logcdf((drift * horizon - distance) / scale) ) second = math.exp(min(float(log_second), 0.0)) return float(np.clip(first + second, 0.0, 1.0)) def hierarchical_fpt_features( history: CausalHistoryView, batteries: np.ndarray, buildings: np.ndarray, cutoff: pd.Timestamp | str, ) -> FPTSnapshotFeatures: """Compute the exact frozen residual features for one causal landmark.""" batteries = np.asarray(batteries, dtype=object) buildings = np.asarray(buildings, dtype=object).astype(str) if len(batteries) != len(buildings): raise ValueError("battery and building arrays have different lengths") state = np.full(len(batteries), np.nan) staleness = np.full(len(batteries), np.inf) changes: list[np.ndarray] = [] for position, battery in enumerate(batteries): x0, gap, increments = terminal_weekly_increments( history.smooth_lookup.get(str(battery), pd.Series(dtype=float)), cutoff ) state[position], staleness[position] = x0, gap changes.append(increments) hit = np.full(len(batteries), np.nan) drift_hat = np.full(len(batteries), np.nan) diffusion_hat = np.full(len(batteries), np.nan) shrinkage = np.full(len(batteries), np.nan) for building in np.unique(buildings): target = buildings == building outer = ~target pooled_parts = [ changes[position] for position in np.flatnonzero(outer) if len(changes[position]) ] if not pooled_parts: continue pooled = np.concatenate(pooled_parts) pooled_drift = float(np.median(pooled) / 7.0) diffusion = max( _robust_scale(pooled - 7.0 * pooled_drift) / math.sqrt(7.0), 1e-6, ) outer_devices = np.asarray( [ position for position in np.flatnonzero(outer) if len(changes[position]) >= MIN_WEEKLY_INCREMENTS ], dtype=int, ) if outer_devices.size: raw_drifts = np.asarray( [np.median(changes[position]) / 7.0 for position in outer_devices] ) between_variance = _robust_scale(raw_drifts) ** 2 noise_variance = np.asarray( [ diffusion**2 / (7.0 * len(changes[position])) for position in outer_devices ] ) prior_variance = max( float(between_variance - np.median(noise_variance)), 0.0 ) else: prior_variance = 0.0 for position in np.flatnonzero(target): count = len(changes[position]) if ( count < MIN_WEEKLY_INCREMENTS or not np.isfinite(state[position]) or staleness[position] > MAX_STATE_STALENESS_DAYS ): continue raw_drift = float(np.median(changes[position]) / 7.0) observation_variance = diffusion**2 / (7.0 * count) weight = ( prior_variance / (prior_variance + observation_variance) if prior_variance > 0.0 else 0.0 ) drift = weight * raw_drift + (1.0 - weight) * pooled_drift hit[position] = wiener_hit_probability(state[position], drift, diffusion) drift_hat[position] = drift diffusion_hat[position] = diffusion shrinkage[position] = weight reliable = np.isfinite(hit) return FPTSnapshotFeatures( last_smooth_voltage=state, smooth_staleness_days=staleness, weekly_increment_count=np.asarray([len(value) for value in changes], dtype=int), eb_drift_v_per_day=drift_hat, pooled_diffusion_v_per_sqrt_day=diffusion_hat, eb_device_weight=shrinkage, hit_probability_42d=hit, reliable=reliable, ) def rerank_fpt_risk( baseline: np.ndarray, hit_probability: np.ndarray, reliable: np.ndarray, freshness: np.ndarray, batteries: np.ndarray, ) -> np.ndarray: """Apply the frozen 15% residual without changing either risk multiset.""" baseline = np.asarray(baseline, dtype=float) hit_probability = np.asarray(hit_probability, dtype=float) reliable = np.asarray(reliable, dtype=bool) freshness = np.asarray(freshness, dtype=float) batteries = np.asarray(batteries, dtype=object) lengths = {len(baseline), len(hit_probability), len(reliable), len(freshness), len(batteries)} if len(lengths) != 1: raise ValueError("FPT rerank inputs have different lengths") if not np.isfinite(baseline).all() or not np.isfinite(freshness).all(): raise ValueError("base risks and freshness factors must be finite") if not np.isfinite(hit_probability[reliable]).all(): raise ValueError("reliable FPT probabilities must be finite") treatment = baseline.copy() for factor in np.sort(np.unique(freshness)): positions = np.flatnonzero((freshness == factor) & reliable) if len(positions) < 2: continue base_rank = rankdata(baseline[positions], method="average") / len(positions) fpt_rank = rankdata(hit_probability[positions], method="average") / len(positions) score = ( (1.0 - RANK_RESIDUAL_WEIGHT) * base_rank + RANK_RESIDUAL_WEIGHT * fpt_rank ) order = np.lexsort((batteries[positions].astype(str), score)) treatment[positions[order]] = np.sort(baseline[positions]) if not np.array_equal(np.sort(treatment), np.sort(baseline)): raise AssertionError("FPT residual changed the raw risk multiset") if not np.array_equal( np.sort(treatment * freshness), np.sort(baseline * freshness) ): raise AssertionError("FPT residual changed planner-effective risk multiset") return treatment @dataclass class HierarchicalFPTPlanner: """Serializable wrapper around the exact public V05 planner artifact.""" base_planner: CompetitionPlanner base_artifact_sha256: str schema_version: int = SCHEMA_VERSION _history_cache: CausalHistoryCache | None = field( default=None, init=False, repr=False, compare=False ) def reset_split(self, split_id: str | None = None) -> None: self._history_cache = CausalHistoryCache(split_id=split_id) def plan_scenario( self, visible_history: pd.DataFrame, snapshot: pd.DataFrame, locations: pd.DataFrame, travel_costs: pd.DataFrame, settings, start_time: pd.Timestamp | str, ) -> pd.DataFrame: if self.schema_version != SCHEMA_VERSION: raise RuntimeError( f"hierarchical FPT schema {self.schema_version} != runtime {SCHEMA_VERSION}" ) if self._history_cache is None: self.reset_split(None) assert self._history_cache is not None start = pd.Timestamp(start_time) history = self._history_cache.update(visible_history, start) batteries = snapshot["battery"].astype(str).to_numpy(object) if not np.array_equal( batteries, locations["battery"].astype(str).to_numpy(object) ): raise AssertionError("snapshot and locations battery order differs") base_risk = self.base_planner.model.predict_event_risk(snapshot) predicted_rul = self.base_planner.model.predict_rul(snapshot) predicted_survivor = self.base_planner.model.predict_survivor_rul(snapshot) if predicted_survivor is None: predicted_survivor = np.full( len(snapshot), float(settings.planning_window_days) ) freshness = planner_freshness_factors( snapshot["data_gap_days"].to_numpy(float), self.base_planner.policy ) features = hierarchical_fpt_features( history, batteries, snapshot["building"].astype(str).to_numpy(object), start, ) treatment_risk = rerank_fpt_risk( base_risk, features.hit_probability_42d, features.reliable, freshness, batteries, ) return self.base_planner.plan_snapshot( snapshot, locations, travel_costs, settings, start, predicted_rul=predicted_rul, predicted_risk=treatment_risk, predicted_survivor_rul=predicted_survivor, )