"""Causal, deployable identity residuals for the BatterySwapAI planner. The frozen base model supplies calibrated event risk and both RUL heads. This module changes only which battery receives each already-calibrated risk value. Every permutation is confined to one exact planner-freshness stratum, so both the raw and planner-consumed risk multisets remain unchanged. Only raw history supplied by ``iterate_scenarios`` may enter ``plan_scenario``. The runtime cache is an incremental acceleration of that visible prefix; it is never trained, serialized with observations, or populated from the full hidden split. """ from __future__ import annotations import math from dataclasses import dataclass, field, replace from typing import Protocol import numpy as np import pandas as pd from scipy.optimize import minimize from scipy.special import log_ndtr, ndtr from scipy.stats import rankdata from .competition_planner import CompetitionPlanner, PlannerPolicy SCHEMA_VERSION = 1 HORIZON_DAYS = 42.0 EOL_VOLTAGE = 2.40 FP_WINDOWS_DAYS = (30, 60, 90, 180) MIN_SMOOTHED_POINTS = 8 MIN_WINDOW_SPAN_FRACTION = 0.5 MIN_DEGRADATION_RATE = 1e-5 MAX_EXTRAPOLATION_DAYS = 730.0 MAX_SMOOTHED_STALENESS_DAYS = 7.0 MAX_LIFETIME_MAD_DAYS = 90.0 MIN_RELIABLE_WINDOWS = 3 NEIGHBORS = 8 PREFIX_LAGS_DAYS = tuple(range(0, 337, 7)) RECENCY_HALF_LIFE_DAYS = 42.0 BASELINE_VALID_POINTS = 30 MIN_PAIRED_LAGS = 8 MIN_WEIGHTED_COVERAGE = 0.50 MIN_EFFECTIVE_NEIGHBORS = 4.0 MIN_UNIQUE_DONOR_BUILDINGS = NEIGHBORS MIN_EFFECTIVE_DONOR_BUILDINGS = 4.0 MAX_LOGIT_RESIDUAL = 1.0 INVERSE_DISTANCE_EPSILON = 1e-6 DISTANCE_BATCH_SIZE = 64 PREFIX_WEIGHTS = np.power( 2.0, -np.asarray(PREFIX_LAGS_DAYS, dtype=float) / RECENCY_HALF_LIFE_DAYS, ) FROZEN_TEMPERATURE_BETA_V_PER_C = 0.00549 REFERENCE_TEMPERATURE_C = 20.0 TEMPERATURE_MAX_PREDICTED_MIN_VOLTAGE = 2.50 SEASONAL_LAG_DAYS = 364 DRIFT_WINDOW_DAYS = 90 MIN_DRIFT_POINTS = 30 MIN_DRIFT_SPAN_DAYS = 60 MIN_ANALOG_DAYS = 28 MAX_HEALTH_STALENESS_DAYS = 7 V07_EMERGENCY_OPERATIONAL_SCALE = 0.75 V07_MAX_MEAN_DEDICATED_TRIP_HOURS = 8.0 def _as_bool(values: pd.Series | np.ndarray) -> np.ndarray: array = np.asarray(values) if np.issubdtype(array.dtype, np.bool_): return array.astype(bool, copy=False) return np.asarray( [str(value).strip().lower() in {"1", "true", "yes"} for value in array], dtype=bool, ) def _rank(values: np.ndarray) -> np.ndarray: values = np.asarray(values, dtype=float) if not len(values) or not np.isfinite(values).all(): raise ValueError("rank input must be finite and non-empty") return rankdata(values, method="average") / len(values) def _logit(values: np.ndarray) -> np.ndarray: clipped = np.clip(np.asarray(values, dtype=float), 1e-6, 1.0 - 1e-6) return np.log(clipped / (1.0 - clipped)) def assign_multiset( baseline: np.ndarray, score: np.ndarray, batteries: np.ndarray, eligible: np.ndarray, ) -> np.ndarray: """Give high scores high baseline values, with deterministic battery ties.""" baseline = np.asarray(baseline, dtype=float) score = np.asarray(score, dtype=float) batteries = np.asarray(batteries, dtype=object) eligible = np.asarray(eligible, dtype=bool) if not (len(baseline) == len(score) == len(batteries) == len(eligible)): raise ValueError("multiset assignment inputs have different lengths") out = baseline.copy() positions = np.flatnonzero(eligible) if len(positions) < 2: return out if not np.isfinite(score[positions]).all(): raise ValueError("eligible multiset scores must be finite") order = np.lexsort((batteries[positions].astype(str), score[positions])) out[positions[order]] = np.sort(baseline[positions]) if not np.array_equal(out[~eligible], baseline[~eligible]): raise AssertionError("ineligible rows changed during multiset assignment") if not np.array_equal(np.sort(out), np.sort(baseline)): raise AssertionError("risk multiset changed during assignment") return out def planner_freshness_factors( data_gap_days: np.ndarray, policy: PlannerPolicy ) -> np.ndarray: gaps = np.asarray(data_gap_days, dtype=float) factors = np.ones(len(gaps), dtype=float) recent = (gaps > 0.0) & (gaps <= policy.stale_risk_cutoff_days) stale = (gaps > policy.stale_risk_cutoff_days) | ~np.isfinite(gaps) factors[recent] = policy.recent_gap_risk_factor factors[stale] = policy.stale_risk_factor return factors def mean_dedicated_trip_hours( snapshot: pd.DataFrame, travel_costs: pd.DataFrame, settings ) -> float: """Frozen V07 inference-only geometry gate.""" distances = travel_costs.set_index(["from", "to"])["hours"] if distances.index.has_duplicates: raise ValueError("travel matrix paths must be unique") base = str(settings.base_location) base_room = str(settings.base_room) hours: list[float] = [] for row in snapshot[["building", "room"]].itertuples(index=False): building = str(row.building) room = str(row.room) try: outbound = 0.0 if building == base else float(distances.loc[(base, building)]) inbound = float(distances.loc[(building, base)]) except KeyError as error: raise ValueError(f"travel matrix lacks {base}<->{building}") from error hours.append( outbound + inbound + (building != base) * float(settings.time_per_building_change_hours) + (room != base_room) * float(settings.time_per_room_change_hours) + float(settings.time_per_battery_hours) ) return float(np.mean(hours)) if hours else math.inf def v07_emergency_scale( snapshot: pd.DataFrame, travel_costs: pd.DataFrame, settings ) -> float: mean_hours = mean_dedicated_trip_hours(snapshot, travel_costs, settings) return ( V07_EMERGENCY_OPERATIONAL_SCALE if mean_hours <= V07_MAX_MEAN_DEDICATED_TRIP_HOURS else 0.0 ) def _raw_columns(timeseries: pd.DataFrame) -> pd.DataFrame: frame = timeseries if "device_id" not in frame.columns or "end_time" not in frame.columns: frame = frame.reset_index() required = {"device_id", "end_time", "voltage", "temperature"} missing = required - set(frame.columns) if missing: raise ValueError(f"battery time series lacks columns: {sorted(missing)}") frame = frame[["device_id", "end_time", "voltage", "temperature"]].copy() frame["device_id"] = frame["device_id"].astype(str) frame["end_time"] = pd.to_datetime(frame["end_time"], errors="raise") frame["voltage"] = pd.to_numeric(frame["voltage"], errors="coerce") frame["temperature"] = pd.to_numeric(frame["temperature"], errors="coerce") return frame def exact_smoothed_voltage(timeseries: pd.DataFrame) -> pd.DataFrame: """Reproduce ``batteryswap_public==0.3.4`` smoothing exactly.""" frame = _raw_columns(timeseries).dropna( subset=["device_id", "end_time", "voltage", "temperature"] ) stable = frame[ frame["temperature"].gt(10.0) & frame["temperature"].lt(30.0) ].reset_index(drop=True) daily_parts: list[pd.DataFrame] = [] for device_id, group in stable.groupby("device_id", observed=True, sort=True): indexed = group.set_index("end_time").sort_index() resample = indexed[["voltage"]].resample("1D") daily_quantile = resample.quantile(0.5) daily_quantile = daily_quantile[resample.count() >= 5] daily_quantile["device_id"] = str(device_id) daily_parts.append(daily_quantile.reset_index()) if not daily_parts: return pd.DataFrame(columns=["device_id", "end_time", "smooth_voltage"]) daily = pd.concat(daily_parts, ignore_index=True).sort_values( ["device_id", "end_time"], kind="stable" ) rolled = ( daily.set_index("end_time") .groupby("device_id", observed=True)[["voltage"]] .rolling(window=7, min_periods=3) .quantile(0.5) .reset_index() .rename(columns={"voltage": "smooth_voltage"}) ) return rolled.sort_values(["device_id", "end_time"], kind="stable").reset_index( drop=True ) def _daily_aggregates(timeseries: pd.DataFrame, beta_v_per_c: float) -> pd.DataFrame: frame = _raw_columns(timeseries).dropna( subset=["device_id", "end_time", "voltage", "temperature"] ) frame = frame[ frame["temperature"].gt(10.0) & frame["temperature"].lt(30.0) ].copy() if frame.empty: return pd.DataFrame( columns=[ "device_id", "day", "smooth_voltage_day", "raw_voltage", "temperature", "health", "count", ] ) frame["day"] = frame["end_time"].dt.normalize() frame["health"] = frame["voltage"] - beta_v_per_c * ( frame["temperature"] - REFERENCE_TEMPERATURE_C ) grouped = frame.groupby(["device_id", "day"], observed=True, sort=True) daily = grouped.agg( raw_voltage=("voltage", "median"), temperature=("temperature", "median"), health=("health", "median"), count=("voltage", "size"), ) daily["smooth_voltage_day"] = grouped["voltage"].quantile(0.5) daily = daily.reset_index() invalid = daily["count"].lt(5) daily.loc[ invalid, ["smooth_voltage_day", "raw_voltage", "temperature", "health"], ] = np.nan return daily.sort_values(["device_id", "day"], kind="stable").reset_index(drop=True) @dataclass class CausalHistoryView: smooth_lookup: dict[str, pd.Series] devices: np.ndarray device_index: dict[str, int] day0: pd.Timestamp | None raw_voltage: np.ndarray temperature: np.ndarray smooth_health: np.ndarray @classmethod def from_daily(cls, daily: pd.DataFrame) -> "CausalHistoryView": if daily.empty: empty = np.zeros((0, 1), dtype=np.float32) return cls({}, np.asarray([], dtype=object), {}, None, empty, empty, empty) lookup: dict[str, pd.Series] = {} for device_id, group in daily.groupby("device_id", observed=True, sort=True): group = group.sort_values("day", kind="stable") index = pd.date_range(group["day"].min(), group["day"].max(), freq="1D") index.name = "end_time" raw = group.set_index("day")["smooth_voltage_day"].reindex(index) lookup[str(device_id)] = raw.rolling(7, min_periods=3).quantile(0.5) devices = np.asarray(sorted(daily["device_id"].astype(str).unique()), dtype=object) device_index = {device: position for position, device in enumerate(devices)} day0 = pd.Timestamp(daily["day"].min()).normalize() day_n = pd.Timestamp(daily["day"].max()).normalize() width = int((day_n - day0) / pd.Timedelta(days=1)) + 1 shape = (len(devices), width) raw_voltage = np.full(shape, np.nan, dtype=np.float32) temperature = np.full(shape, np.nan, dtype=np.float32) health = np.full(shape, np.nan, dtype=np.float32) rows = daily["device_id"].astype(str).map(device_index).to_numpy(dtype=int) columns = ((pd.to_datetime(daily["day"]) - day0) / pd.Timedelta(days=1)).to_numpy( dtype=int ) raw_voltage[rows, columns] = daily["raw_voltage"].to_numpy(dtype=np.float32) temperature[rows, columns] = daily["temperature"].to_numpy(dtype=np.float32) health[rows, columns] = daily["health"].to_numpy(dtype=np.float32) smooth_health = np.full(shape, np.nan, dtype=np.float32) for row in range(len(devices)): smooth_health[row] = ( pd.Series(health[row]) .rolling(7, min_periods=3) .median() .to_numpy(dtype=np.float32) ) return cls( lookup, devices, device_index, day0, raw_voltage, temperature, smooth_health, ) @dataclass class CausalHistoryCache: """Incremental daily aggregation of scenario-visible raw histories.""" beta_v_per_c: float = FROZEN_TEMPERATURE_BETA_V_PER_C split_id: str | None = None last_cutoff: pd.Timestamp | None = None daily: pd.DataFrame = field(default_factory=pd.DataFrame) seen_devices: set[str] = field(default_factory=set) def reset(self, split_id: str | None = None) -> None: self.split_id = split_id self.last_cutoff = None self.daily = pd.DataFrame() self.seen_devices = set() def update( self, visible_history: pd.DataFrame, cutoff: pd.Timestamp | str ) -> CausalHistoryView: cutoff = pd.Timestamp(cutoff) flat = _raw_columns(visible_history) if not flat.empty and flat["end_time"].max() > cutoff: raise AssertionError( f"visible history reaches {flat['end_time'].max()} after cutoff {cutoff}" ) rebuild = self.last_cutoff is None or cutoff <= self.last_cutoff if rebuild: self.daily = _daily_aggregates(flat, self.beta_v_per_c) self.seen_devices = set(flat["device_id"].astype(str)) else: overlap_day = self.last_cutoff.normalize() current_devices = set(flat["device_id"].astype(str)) new_devices = current_devices - self.seen_devices use = flat["end_time"].ge(overlap_day) | flat["device_id"].isin(new_devices) replacement = _daily_aggregates(flat.loc[use], self.beta_v_per_c) if not self.daily.empty: replace_devices = current_devices | new_devices keep = ~( self.daily["device_id"].astype(str).isin(replace_devices) & pd.to_datetime(self.daily["day"]).ge(overlap_day) ) self.daily = self.daily.loc[keep] self.daily = pd.concat([self.daily, replacement], ignore_index=True) if not self.daily.empty: self.daily = self.daily.sort_values( ["device_id", "day"], kind="stable" ).reset_index(drop=True) if self.daily.duplicated(["device_id", "day"]).any(): raise AssertionError("incremental daily cache contains duplicate device-days") self.seen_devices.update(current_devices) self.last_cutoff = cutoff return CausalHistoryView.from_daily(self.daily) @dataclass(frozen=True) class AFTParameters: beta0: float beta1: float sigma: float def fit_aft(frame: pd.DataFrame) -> tuple[AFTParameters, dict[str, int | float]]: required = { "battery", "fp_reliable", "fp_lifetime_days", "landmark_age_days", "outcome_lifetime_days", "event_observed", } missing = required - set(frame.columns) if missing: raise ValueError(f"AFT training frame lacks columns: {sorted(missing)}") reliable = _as_bool(frame["fp_reliable"]) age_all = pd.to_numeric(frame["landmark_age_days"], errors="coerce").to_numpy(float) lifetime_all = pd.to_numeric( frame["outcome_lifetime_days"], errors="coerce" ).to_numpy(float) fp_all = pd.to_numeric(frame["fp_lifetime_days"], errors="coerce").to_numpy(float) eligible = ( reliable & np.isfinite(fp_all) & np.isfinite(age_all) & np.isfinite(lifetime_all) & (lifetime_all > age_all) & (age_all > 0.0) ) work = frame.loc[eligible].copy() if work.empty: raise ValueError("AFT fit has no reliable likelihood-eligible rows") event = _as_bool(work["event_observed"]) event_devices = int(work.loc[event, "battery"].nunique()) devices = int(work["battery"].nunique()) if event_devices < 10 or devices < 20: raise ValueError( "AFT fit is not identifiable: " f"event_devices={event_devices}, devices={devices}" ) counts = work.groupby("battery", observed=True)["battery"].transform("size") weights = 1.0 / counts.to_numpy(dtype=float) x = np.log(pd.to_numeric(work["fp_lifetime_days"]).to_numpy(dtype=float)) age = pd.to_numeric(work["landmark_age_days"]).to_numpy(dtype=float) outcome = pd.to_numeric(work["outcome_lifetime_days"]).to_numpy(dtype=float) observed_x = x[event] observed_y = np.log(outcome[event]) observed_w = weights[event] design = np.column_stack([np.ones(len(observed_x)), observed_x]) initial_beta = np.linalg.lstsq( design * np.sqrt(observed_w)[:, None], observed_y * np.sqrt(observed_w), rcond=None, )[0] initial_beta[0] = np.clip(initial_beta[0], -20.0, 20.0) initial_beta[1] = np.clip(initial_beta[1], 0.0, 3.0) residual = observed_y - (initial_beta[0] + initial_beta[1] * observed_x) initial_sigma = float( np.clip(np.sqrt(np.average(residual**2, weights=observed_w)), 0.1, 1.5) ) def objective(raw: np.ndarray) -> float: beta0, beta1, log_sigma = raw sigma = float(np.exp(log_sigma)) mu = beta0 + beta1 * x z_age = (np.log(age) - mu) / sigma log_survival_age = log_ndtr(-z_age) z_outcome = (np.log(outcome) - mu) / sigma likelihood = np.empty(len(work), dtype=float) likelihood[event] = ( -np.log(outcome[event]) - log_sigma - 0.5 * math.log(2.0 * math.pi) - 0.5 * z_outcome[event] ** 2 - log_survival_age[event] ) likelihood[~event] = log_ndtr(-z_outcome[~event]) - log_survival_age[~event] if not np.isfinite(likelihood).all(): return 1e100 return float(-np.sum(weights * likelihood)) fitted = minimize( objective, np.asarray( [initial_beta[0], initial_beta[1], math.log(initial_sigma)], dtype=float ), method="L-BFGS-B", bounds=( (-20.0, 20.0), (0.0, 3.0), (math.log(0.05), math.log(2.0)), ), options={"maxiter": 1_000, "ftol": 1e-12, "gtol": 1e-8}, ) if not fitted.success or not np.isfinite(fitted.fun): raise RuntimeError(f"three-parameter AFT optimization failed: {fitted.message}") parameters = AFTParameters( beta0=float(fitted.x[0]), beta1=float(fitted.x[1]), sigma=float(np.exp(fitted.x[2])), ) diagnostics: dict[str, int | float] = { "eligible_rows": int(len(work)), "devices": devices, "event_devices": event_devices, "censored_devices": int(work.loc[~event, "battery"].nunique()), "objective": float(fitted.fun), "iterations": int(fitted.nit), } return parameters, diagnostics def _theil_sen_slope(days: np.ndarray, values: np.ndarray) -> float: x = days.astype("datetime64[ns]").astype(np.int64) / 86_400_000_000_000.0 parts: list[np.ndarray] = [] for offset in range(1, len(x)): delta_x = x[offset:] - x[:-offset] usable = delta_x > 0.0 if usable.any(): parts.append((values[offset:] - values[:-offset])[usable] / delta_x[usable]) return float(np.median(np.concatenate(parts))) if parts else math.nan def _first_passage_lifetime( device_id: str, cutoff: pd.Timestamp, installation_start: pd.Timestamp, lookup: dict[str, pd.Series], ) -> tuple[float, bool]: series = lookup.get(str(device_id)) if series is None: return math.nan, False finite = series.dropna() finite = finite[finite.index <= pd.Timestamp(cutoff).normalize()] if finite.empty: return math.nan, False days = finite.index.to_numpy(dtype="datetime64[ns]") values = finite.to_numpy(dtype=float) last_day = pd.Timestamp(days[-1]) last_value = float(values[-1]) staleness = float((pd.Timestamp(cutoff) - last_day) / pd.Timedelta(days=1)) cutoff64 = np.datetime64(pd.Timestamp(cutoff).to_datetime64(), "ns") lifetimes: list[float] = [] for window in FP_WINDOWS_DAYS: low = cutoff64 - np.timedelta64(window - 1, "D") use = days >= low window_days = days[use] window_values = values[use] if len(window_values) < MIN_SMOOTHED_POINTS: continue span = float( (pd.Timestamp(window_days[-1]) - pd.Timestamp(window_days[0])) / pd.Timedelta(days=1) ) if span < window * MIN_WINDOW_SPAN_FRACTION: continue slope = _theil_sen_slope(window_days, window_values) if not np.isfinite(slope) or slope >= -MIN_DEGRADATION_RATE: continue eta = float( np.clip((last_value - EOL_VOLTAGE) / -slope, 0.0, MAX_EXTRAPOLATION_DAYS) ) lifetime = float( (last_day - installation_start) / pd.Timedelta(days=1) ) + eta if np.isfinite(lifetime) and lifetime > 0.0: lifetimes.append(lifetime) if not lifetimes: return math.nan, False median = float(np.median(lifetimes)) mad = float(np.median(np.abs(np.asarray(lifetimes) - median))) reliable = bool( len(lifetimes) >= MIN_RELIABLE_WINDOWS and 0.0 <= staleness <= MAX_SMOOTHED_STALENESS_DAYS and mad <= MAX_LIFETIME_MAD_DAYS ) return median, reliable @dataclass(frozen=True) class ResidualSignal: values: np.ndarray reliable: np.ndarray @dataclass(frozen=True) class ResidualContext: history: CausalHistoryView batteries: np.ndarray locations: pd.DataFrame cutoff: pd.Timestamp class IdentityRankResidual(Protocol): name: str def predict(self, context: ResidualContext) -> ResidualSignal: ... def component_risk( self, baseline: np.ndarray, batteries: np.ndarray, signal: ResidualSignal, ) -> np.ndarray: ... @dataclass(frozen=True) class LongTermAFTResidual: parameters: AFTParameters name: str = "lt_fp_aft" base_weight: float = 0.90 aft_weight: float = 0.10 def predict(self, context: ResidualContext) -> ResidualSignal: location_index = context.locations.assign( battery=context.locations["battery"].astype(str) ).set_index("battery", drop=False) values = np.full(len(context.batteries), np.nan) reliable = np.zeros(len(context.batteries), dtype=bool) for position, battery in enumerate(context.batteries.astype(str)): if battery not in location_index.index: continue installation = pd.Timestamp(location_index.loc[battery, "start_time"]) lifetime, is_reliable = _first_passage_lifetime( battery, context.cutoff, installation, context.history.smooth_lookup, ) age = float((context.cutoff - installation) / pd.Timedelta(days=1)) if not is_reliable or not np.isfinite(lifetime) or lifetime <= 0.0 or age <= 0.0: continue mu = self.parameters.beta0 + self.parameters.beta1 * math.log(lifetime) z_now = (math.log(age) - mu) / self.parameters.sigma z_horizon = (math.log(age + HORIZON_DAYS) - mu) / self.parameters.sigma survival_now = float(ndtr(-z_now)) probability = (float(ndtr(z_horizon)) - float(ndtr(z_now))) / max( survival_now, np.finfo(float).tiny ) values[position] = float(np.clip(probability, 0.0, 1.0)) reliable[position] = True return ResidualSignal(values, reliable) def component_risk( self, baseline: np.ndarray, batteries: np.ndarray, signal: ResidualSignal, ) -> np.ndarray: positions = np.flatnonzero(signal.reliable) if len(positions) < 2: return np.asarray(baseline, dtype=float).copy() score = np.zeros(len(baseline), dtype=float) score[positions] = self.base_weight * _rank( np.asarray(baseline)[positions] ) + self.aft_weight * _rank(signal.values[positions]) return assign_multiset(baseline, score, batteries, signal.reliable) @dataclass(frozen=True) class SimilarityDonorLibrary: values: np.ndarray masks: np.ndarray donor_ids: np.ndarray donor_buildings: np.ndarray donor_codes: np.ndarray endpoint_days: np.ndarray residual_days: np.ndarray endpoint_groups: tuple[np.ndarray, ...] device_ids_by_code: np.ndarray device_buildings_by_code: np.ndarray building_groups: tuple[np.ndarray, ...] @property def donor_count(self) -> int: return len(self.endpoint_groups) @property def endpoint_count(self) -> int: return len(self.values) @property def building_count(self) -> int: return len(self.building_groups) def _series_lookup(smoothed: pd.DataFrame) -> dict[str, pd.Series]: lookup: dict[str, pd.Series] = {} for device_id, group in smoothed.groupby("device_id", observed=True, sort=False): ordered = group.sort_values("end_time", kind="stable") days = pd.DatetimeIndex(ordered["end_time"]).normalize() if days.has_duplicates: raise AssertionError(f"smoothed grid has duplicate days for {device_id}") lookup[str(device_id)] = pd.Series( ordered["smooth_voltage"].to_numpy(dtype=float), index=days ) return lookup def _normalised_prefix( series: pd.Series, anchor: pd.Timestamp, baseline_voltage: float ) -> np.ndarray: target_days = pd.Timestamp(anchor).normalize() - pd.to_timedelta( np.asarray(PREFIX_LAGS_DAYS), unit="D" ) voltage = series.reindex(target_days).to_numpy(dtype=float) margin = float(baseline_voltage) - EOL_VOLTAGE if not np.isfinite(margin) or margin <= 1e-6: return np.full(len(PREFIX_LAGS_DAYS), np.nan) return (voltage - EOL_VOLTAGE) / margin def build_similarity_donor_library( full_smoothed: pd.DataFrame, eol_times: pd.Series, battery_building: dict[str, str], ) -> SimilarityDonorLibrary: lookup = _series_lookup(full_smoothed) values: list[np.ndarray] = [] donor_ids: list[str] = [] donor_buildings: list[str] = [] endpoint_days: list[np.datetime64] = [] residual_days: list[float] = [] observed = pd.to_datetime(eol_times.dropna(), errors="raise").sort_index() for raw_device_id, raw_eol in observed.items(): device_id = str(raw_device_id) building = str(battery_building.get(device_id, "")) if not building: continue series = lookup.get(device_id) if series is None: continue eol_day = pd.Timestamp(raw_eol).normalize() pre_eol = series[series.index <= eol_day] finite = pre_eol.dropna() if len(finite) < BASELINE_VALID_POINTS: continue baseline = float(finite.iloc[:BASELINE_VALID_POINTS].median()) if not np.isfinite(baseline) or baseline <= EOL_VOLTAGE + 1e-6: continue baseline_ready = pd.Timestamp(finite.index[BASELINE_VALID_POINTS - 1]).normalize() maximum_residual = int((eol_day - baseline_ready) / pd.Timedelta(days=1)) for residual in range(1, maximum_residual + 1, 7): endpoint = eol_day - pd.Timedelta(days=residual) vector = _normalised_prefix(pre_eol, endpoint, baseline) if not np.isfinite(vector[0]) or np.isfinite(vector).sum() < MIN_PAIRED_LAGS: continue values.append(vector) donor_ids.append(device_id) donor_buildings.append(building) endpoint_days.append(np.datetime64(endpoint, "ns")) residual_days.append(float(residual)) if not values: raise ValueError("full training split has no usable similarity donor endpoints") donor_order = sorted(set(donor_ids)) donor_to_code = {device: code for code, device in enumerate(donor_order)} donor_codes = np.asarray([donor_to_code[device] for device in donor_ids], dtype=int) endpoint_groups = tuple( np.flatnonzero(donor_codes == code) for code in range(len(donor_order)) ) device_buildings = np.asarray( [donor_buildings[int(group[0])] for group in endpoint_groups], dtype=object ) building_groups = tuple( np.flatnonzero(device_buildings.astype(str) == building) for building in sorted(set(device_buildings.astype(str))) ) if len(endpoint_groups) < NEIGHBORS or len(building_groups) < NEIGHBORS: raise ValueError("full donor library cannot supply K=8 distinct buildings") value_array = np.asarray(values, dtype=float) return SimilarityDonorLibrary( values=value_array, masks=np.isfinite(value_array), donor_ids=np.asarray(donor_ids, dtype=object), donor_buildings=np.asarray(donor_buildings, dtype=object), donor_codes=donor_codes, endpoint_days=np.asarray(endpoint_days, dtype="datetime64[ns]"), residual_days=np.asarray(residual_days, dtype=float), endpoint_groups=endpoint_groups, device_ids_by_code=np.asarray(donor_order, dtype=object), device_buildings_by_code=device_buildings, building_groups=building_groups, ) def _query_similarity( batteries: np.ndarray, cutoff: pd.Timestamp, lookup: dict[str, pd.Series], ) -> tuple[np.ndarray, np.ndarray, np.ndarray]: vectors = np.full((len(batteries), len(PREFIX_LAGS_DAYS)), np.nan) ready = np.zeros(len(batteries), dtype=bool) staleness = np.full(len(batteries), np.nan) for position, battery in enumerate(batteries.astype(str)): series = lookup.get(battery) if series is None: continue finite = series.dropna() finite = finite[finite.index <= cutoff.normalize()] if finite.empty: continue anchor = pd.Timestamp(finite.index[-1]).normalize() local_staleness = float((cutoff - anchor) / pd.Timedelta(days=1)) baseline = ( float(finite.iloc[:BASELINE_VALID_POINTS].median()) if len(finite) >= BASELINE_VALID_POINTS else math.nan ) vector = _normalised_prefix(series, anchor, baseline) mask = np.isfinite(vector) coverage = float(PREFIX_WEIGHTS[mask].sum() / PREFIX_WEIGHTS.sum()) vectors[position] = vector staleness[position] = local_staleness ready[position] = bool( len(finite) >= BASELINE_VALID_POINTS and np.isfinite(baseline) and baseline > EOL_VOLTAGE + 1e-6 and 0.0 <= local_staleness <= MAX_SMOOTHED_STALENESS_DAYS and int(mask.sum()) >= MIN_PAIRED_LAGS and coverage >= MIN_WEIGHTED_COVERAGE ) return vectors, ready, staleness def score_similarity( query_values: np.ndarray, query_ready: np.ndarray, query_staleness: np.ndarray, library: SimilarityDonorLibrary, ) -> ResidualSignal: row_count = len(query_values) risk = np.full(row_count, np.nan) reliable_out = np.zeros(row_count, dtype=bool) donor_values = np.nan_to_num(library.values, nan=0.0) donor_mask = library.masks.astype(float) donor_value_mask = donor_values * donor_mask donor_square_mask = donor_values**2 * donor_mask ready_positions = np.flatnonzero(np.asarray(query_ready, dtype=bool)) for begin in range(0, len(ready_positions), DISTANCE_BATCH_SIZE): positions = ready_positions[begin : begin + DISTANCE_BATCH_SIZE] raw_query = np.asarray(query_values[positions], dtype=float) query_mask = np.isfinite(raw_query).astype(float) query = np.nan_to_num(raw_query, nan=0.0) weighted_mask = query_mask * PREFIX_WEIGHTS[None, :] overlap_weight = weighted_mask @ donor_mask.T paired_lags = query_mask @ donor_mask.T query_weight = weighted_mask.sum(axis=1, keepdims=True) coverage = overlap_weight / np.maximum(query_weight, np.finfo(float).tiny) first = (weighted_mask * query**2) @ donor_mask.T second = weighted_mask @ donor_square_mask.T cross = (weighted_mask * query) @ donor_value_mask.T numerator = np.maximum(first + second - 2.0 * cross, 0.0) with np.errstate(divide="ignore", invalid="ignore"): distance = np.sqrt(numerator / overlap_weight) / np.sqrt(coverage) usable = ( (paired_lags >= MIN_PAIRED_LAGS) & (coverage >= MIN_WEIGHTED_COVERAGE) & np.isfinite(distance) ) distance = np.where(usable, distance, np.inf) block_size = len(positions) best_distance = np.full((block_size, library.donor_count), np.inf) best_endpoint = np.full((block_size, library.donor_count), -1, dtype=int) for donor_code, endpoint_index in enumerate(library.endpoint_groups): local = distance[:, endpoint_index] choice = np.argmin(local, axis=1) best_distance[:, donor_code] = local[np.arange(block_size), choice] best_endpoint[:, donor_code] = endpoint_index[choice] best_building_distance = np.full((block_size, library.building_count), np.inf) best_building_endpoint = np.full( (block_size, library.building_count), -1, dtype=int ) for building_code, donor_codes in enumerate(library.building_groups): local = best_distance[:, donor_codes] choice = np.argmin(local, axis=1) chosen_donor = donor_codes[choice] best_building_distance[:, building_code] = local[ np.arange(block_size), choice ] best_building_endpoint[:, building_code] = best_endpoint[ np.arange(block_size), chosen_donor ] top_buildings = np.argsort( best_building_distance, axis=1, kind="stable" )[:, :NEIGHBORS] top_distance = np.take_along_axis( best_building_distance, top_buildings, axis=1 ) top_endpoint = np.take_along_axis( best_building_endpoint, top_buildings, axis=1 ) has_k = np.isfinite(top_distance).all(axis=1) & (top_endpoint >= 0).all(axis=1) safe_endpoint = np.maximum(top_endpoint, 0) top_rul = np.maximum( library.residual_days[safe_endpoint] - query_staleness[positions, None], 0.0, ) inverse = 1.0 / (top_distance + INVERSE_DISTANCE_EPSILON) inverse = np.where(has_k[:, None], inverse, 0.0) weights = inverse / np.maximum( inverse.sum(axis=1, keepdims=True), np.finfo(float).tiny ) neighbor_risk = np.sum(weights * (top_rul <= HORIZON_DAYS), axis=1) effective_neighbors = 1.0 / np.maximum( np.sum(weights**2, axis=1), np.finfo(float).tiny ) selected_buildings = library.donor_buildings[safe_endpoint] unique_buildings = np.zeros(block_size) effective_buildings = np.zeros(block_size) for row in range(block_size): by_building: dict[str, float] = {} for building, weight in zip( selected_buildings[row], weights[row], strict=True ): name = str(building) by_building[name] = by_building.get(name, 0.0) + float(weight) unique_buildings[row] = len(by_building) effective_buildings[row] = 1.0 / max( sum(weight**2 for weight in by_building.values()), np.finfo(float).tiny, ) reliable = ( has_k & (effective_neighbors >= MIN_EFFECTIVE_NEIGHBORS) & (unique_buildings >= MIN_UNIQUE_DONOR_BUILDINGS) & (effective_buildings >= MIN_EFFECTIVE_DONOR_BUILDINGS) ) risk[positions] = np.where(has_k, neighbor_risk, np.nan) reliable_out[positions] = reliable return ResidualSignal(risk, reliable_out) @dataclass(frozen=True) class OriginalSimilarityResidual: library: SimilarityDonorLibrary name: str = "original_similarity_eol" def predict(self, context: ResidualContext) -> ResidualSignal: query, ready, staleness = _query_similarity( context.batteries, context.cutoff, context.history.smooth_lookup ) return score_similarity(query, ready, staleness, self.library) def component_risk( self, baseline: np.ndarray, batteries: np.ndarray, signal: ResidualSignal, ) -> np.ndarray: del batteries # Reference experiment deliberately used stable row-order ties. baseline = np.asarray(baseline, dtype=float) out = baseline.copy() positions = np.flatnonzero(signal.reliable) if len(positions) < 2: return out local_risk = signal.values[positions] if not np.isfinite(local_risk).all(): raise ValueError("reliable similarity rows have non-finite risk") score = _logit(baseline) score[positions] += np.clip( 2.0 * local_risk - 1.0, -MAX_LOGIT_RESIDUAL, MAX_LOGIT_RESIDUAL ) order = np.argsort(score[positions], kind="stable") reassigned = np.empty(len(positions), dtype=float) reassigned[order] = np.sort(baseline[positions]) out[positions] = reassigned if not np.array_equal(out[~signal.reliable], baseline[~signal.reliable]): raise AssertionError("unreliable similarity rows changed") if not np.array_equal(np.sort(out), np.sort(baseline)): raise AssertionError("similarity component changed risk multiset") return out def fit_temperature_beta(timeseries: pd.DataFrame) -> tuple[float, dict[str, int | float]]: frame = _raw_columns(timeseries).dropna(subset=["voltage", "temperature"]) frame = frame[ frame["temperature"].gt(10.0) & frame["temperature"].lt(30.0) ].copy() frame["day"] = frame["end_time"].dt.normalize() grouped = frame.groupby(["device_id", "day"], observed=True, sort=False) count = grouped["voltage"].transform("size") usable = count.ge(5) centered_voltage = frame["voltage"] - grouped["voltage"].transform("mean") centered_temperature = frame["temperature"] - grouped["temperature"].transform( "mean" ) x = centered_temperature[usable].to_numpy(dtype=float) y = centered_voltage[usable].to_numpy(dtype=float) denominator = float(np.dot(x, x)) if denominator <= 0.0: raise ValueError("temperature coefficient has no within-device-day variation") raw_beta = float(np.dot(x, y) / denominator) return max(raw_beta, 0.0), { "raw_beta_v_per_c": raw_beta, "clipped_beta_v_per_c": max(raw_beta, 0.0), "training_readings": int(usable.sum()), "training_device_days": int( frame.loc[usable, ["device_id", "day"]].drop_duplicates().shape[0] ), } def _seasonal_physics_signal( history: CausalHistoryView, batteries: np.ndarray, cutoff: pd.Timestamp, beta_v_per_c: float, maximum_predicted_min_voltage: float | None = None, ) -> ResidualSignal: count = len(batteries) reliable = np.zeros(count, dtype=bool) urgency = np.full(count, np.nan) if history.day0 is None: return ResidualSignal(urgency, reliable) cut_column = int((cutoff.normalize() - history.day0) / pd.Timedelta(days=1)) forecast_offsets = np.arange(1, int(HORIZON_DAYS) + 1, dtype=int) for position, battery in enumerate(batteries.astype(str)): row = history.device_index.get(battery) if row is None or cut_column <= 0: continue history_start = max(0, cut_column - DRIFT_WINDOW_DAYS) health = history.smooth_health[row, history_start:cut_column].astype(float) finite = np.flatnonzero(np.isfinite(health)) if len(finite) < MIN_DRIFT_POINTS: continue if finite[-1] - finite[0] < MIN_DRIFT_SPAN_DAYS: continue last_column = history_start + int(finite[-1]) staleness = cut_column - 1 - last_column if staleness > MAX_HEALTH_STALENESS_DAYS: continue x = finite.astype(float) y = health[finite] centered = x - x.mean() denominator = float(np.dot(centered, centered)) if denominator <= 0.0: continue slope = min(float(np.dot(centered, y - y.mean()) / denominator), 0.0) current_health = float(y[-1]) analog_columns = cut_column + forecast_offsets - SEASONAL_LAG_DAYS valid_columns = (analog_columns >= 0) & ( analog_columns < history.temperature.shape[1] ) analog = np.full(int(HORIZON_DAYS), np.nan) analog[valid_columns] = history.temperature[ row, analog_columns[valid_columns] ].astype(float) finite_temperature = np.isfinite(analog) if int(finite_temperature.sum()) < MIN_ANALOG_DAYS: continue analog_filled = np.interp( forecast_offsets.astype(float), forecast_offsets[finite_temperature].astype(float), analog[finite_temperature], ) forecast_health = current_health + slope * forecast_offsets forecast_raw = forecast_health + beta_v_per_c * ( analog_filled - REFERENCE_TEMPERATURE_C ) trailing = history.raw_voltage[ row, max(0, cut_column - 6) : cut_column ].astype(float) if len(trailing) < 6: trailing = np.pad(trailing, (6 - len(trailing), 0), constant_values=np.nan) combined = np.concatenate([trailing[-6:], forecast_raw]) forecast_smooth = ( pd.Series(combined).rolling(7, min_periods=3).median().to_numpy()[6:] ) if not np.isfinite(forecast_smooth).all(): continue predicted_min_voltage = float(np.min(forecast_smooth)) urgency[position] = -predicted_min_voltage reliable[position] = bool( maximum_predicted_min_voltage is None or predicted_min_voltage <= maximum_predicted_min_voltage ) return ResidualSignal(urgency, reliable) class PostRankResidual(Protocol): name: str def predict(self, context: ResidualContext) -> ResidualSignal: ... def rerank( self, baseline: np.ndarray, batteries: np.ndarray, signal: ResidualSignal, ) -> np.ndarray: ... @dataclass(frozen=True) class SeasonalTemperatureResidual: beta_v_per_c: float = FROZEN_TEMPERATURE_BETA_V_PER_C fitted_raw_beta_v_per_c: float = FROZEN_TEMPERATURE_BETA_V_PER_C training_readings: int = 0 maximum_predicted_min_voltage: float = TEMPERATURE_MAX_PREDICTED_MIN_VOLTAGE name: str = "seasonal_temperature_below_250_logit1" def predict(self, context: ResidualContext) -> ResidualSignal: return _seasonal_physics_signal( context.history, context.batteries, context.cutoff, self.beta_v_per_c, self.maximum_predicted_min_voltage, ) def rerank( self, baseline: np.ndarray, batteries: np.ndarray, signal: ResidualSignal, ) -> np.ndarray: positions = np.flatnonzero(signal.reliable) if len(positions) < 2: return np.asarray(baseline, dtype=float).copy() score = _logit(baseline) score[positions] += 2.0 * _rank(signal.values[positions]) - 1.0 return assign_multiset(baseline, score, batteries, signal.reliable) @dataclass(frozen=True) class WeightedIdentityResidual: residual: IdentityRankResidual weight: float @dataclass(frozen=True) class IdentityEnsembleModel: """Composable rank residuals followed by optional bounded post-residuals.""" identity_residuals: tuple[WeightedIdentityResidual, ...] post_residuals: tuple[PostRankResidual, ...] = () history_temperature_beta_v_per_c: float = FROZEN_TEMPERATURE_BETA_V_PER_C def predict_risk( self, baseline: np.ndarray, freshness: np.ndarray, context: ResidualContext, ) -> np.ndarray: identity_signals = tuple( item.residual.predict(context) for item in self.identity_residuals ) post_signals = tuple(residual.predict(context) for residual in self.post_residuals) return self.rerank_from_signals( baseline, freshness, context.batteries, identity_signals, post_signals, ) def rerank_from_signals( self, baseline: np.ndarray, freshness: np.ndarray, batteries: np.ndarray, identity_signals: tuple[ResidualSignal, ...], post_signals: tuple[ResidualSignal, ...] = (), ) -> np.ndarray: """Pure parity surface for already-computed causal component signals.""" baseline = np.asarray(baseline, dtype=float) freshness = np.asarray(freshness, dtype=float) batteries = np.asarray(batteries, dtype=object) if not (len(baseline) == len(freshness) == len(batteries)): raise ValueError("ensemble 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 self.identity_residuals: identity = baseline.copy() else: if len(identity_signals) != len(self.identity_residuals): raise ValueError("identity signal count differs from configured residuals") identity = baseline.copy() for factor in np.sort(np.unique(freshness)): positions = np.flatnonzero(freshness == factor) local_base = baseline[positions] local_batteries = batteries[positions] components: list[np.ndarray] = [] union = np.zeros(len(positions), dtype=bool) total_weight = 0.0 borda = np.zeros(len(positions), dtype=float) for item, signal in zip( self.identity_residuals, identity_signals, strict=True ): local_signal = ResidualSignal( signal.values[positions], signal.reliable[positions] ) component = item.residual.component_risk( local_base, local_batteries, local_signal ) components.append(component) union |= local_signal.reliable if item.weight > 0.0: borda += item.weight * _rank(component) total_weight += item.weight if total_weight <= 0.0: raise ValueError("identity residual weights must contain a positive value") borda /= total_weight identity[positions] = assign_multiset( local_base, borda, local_batteries, union ) if not np.array_equal( np.sort(local_base * factor), np.sort(identity[positions] * factor), ): raise AssertionError("identity changed planner-effective risk multiset") treatment = identity if len(post_signals) != len(self.post_residuals): raise ValueError("post signal count differs from configured residuals") for residual, signal in zip(self.post_residuals, post_signals, strict=True): updated = treatment.copy() for factor in np.sort(np.unique(freshness)): positions = np.flatnonzero(freshness == factor) local_signal = ResidualSignal( signal.values[positions], signal.reliable[positions] ) updated[positions] = residual.rerank( treatment[positions], batteries[positions], local_signal ) if not np.array_equal( np.sort(treatment[positions] * factor), np.sort(updated[positions] * factor), ): raise AssertionError( f"{residual.name} changed planner-effective risk multiset" ) treatment = updated if not np.array_equal(np.sort(treatment), np.sort(baseline)): raise AssertionError("ensemble changed scenario raw risk multiset") if not np.array_equal( np.sort(treatment * freshness), np.sort(baseline * freshness) ): raise AssertionError("ensemble changed scenario effective risk multiset") return treatment @dataclass class IdentityEnsemblePlanner: """Deployable wrapper that preserves the frozen base planner and its heads.""" base_planner: CompetitionPlanner ensemble: IdentityEnsembleModel 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( beta_v_per_c=self.ensemble.history_temperature_beta_v_per_c, 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"identity artifact 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() if not np.array_equal(batteries, locations["battery"].astype(str).to_numpy()): 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(dtype=float), self.base_planner.policy ) context = ResidualContext(history, batteries, locations, start) treatment_risk = self.ensemble.predict_risk(base_risk, freshness, context) scale = v07_emergency_scale(snapshot, travel_costs, settings) policy = replace( self.base_planner.policy, emergency_operational_scale=float(scale) ) planner = CompetitionPlanner(None, policy) return planner.plan_snapshot( snapshot, locations, travel_costs, settings, start, predicted_rul=predicted_rul, predicted_risk=treatment_risk, predicted_survivor_rul=predicted_survivor, )