CausalModelEvaluation / model /causalmodelevaluation.py
zhangrenchao's picture
Publish CausalModelEvaluation engineering reproduction
20cdc88 verified
Raw
History Blame Contribute Delete
9.65 kB
"""Core statistical models and metrics for causal model evaluation (CME)."""
from __future__ import annotations
from dataclasses import dataclass
import numpy as np
import torch
from scipy import stats
from sklearn.gaussian_process import GaussianProcessRegressor
from sklearn.gaussian_process.kernels import ConstantKernel, RBF, WhiteKernel
FORMAT_VERSION = "cme_lagged_parcorr_v1"
def _as_series(data: np.ndarray | torch.Tensor, nodes: int) -> np.ndarray:
"""Normalize [T,N] or [B,T,N] input without joining sample boundaries."""
array = data.detach().cpu().numpy() if isinstance(data, torch.Tensor) else np.asarray(data)
if array.ndim == 2:
array = array[None, ...]
if array.ndim != 3 or array.shape[-1] != nodes:
raise ValueError(f"expected [T,{nodes}] or [B,T,{nodes}], got {array.shape}")
if not np.isfinite(array).all():
raise ValueError("node time series contains NaN or Inf")
return array.astype(np.float64, copy=False)
@dataclass
class CausalNetwork:
"""A directed signed lagged network with source-target-lag tensor layout."""
edges: torch.Tensor
pvalues: torch.Tensor
mci: torch.Tensor
def state_dict(self) -> dict[str, torch.Tensor]:
return {"edges": self.edges.cpu(), "pvalues": self.pvalues.cpu(), "mci": self.mci.cpu()}
@classmethod
def from_state_dict(cls, state: dict[str, torch.Tensor]) -> "CausalNetwork":
return cls(state["edges"].bool(), state["pvalues"].float(), state["mci"].float())
class LaggedPartialCorrelationCME:
"""Linear conditional-regression approximation to PCMCI's ParCorr MCI step.
Each target is conditioned on its own history through ``max_lag``. All
source-lag candidates are then residualized against that common condition
set and tested with a two-sided t test. This captures directed lagged
conditional dependence, but does not implement PCMCI's iterative PC stage.
"""
def __init__(self, nodes: int = 50, max_lag: int = 10, alpha: float = 0.02,
ridge: float = 1e-6, exclude_self_links: bool = True):
self.nodes = int(nodes)
self.max_lag = int(max_lag)
self.alpha = float(alpha)
self.ridge = float(ridge)
self.exclude_self_links = bool(exclude_self_links)
if self.nodes <= 1 or self.max_lag < 1 or not 0 < self.alpha < 1:
raise ValueError("nodes, max_lag, and alpha must define a valid test")
def fit(self, data: np.ndarray | torch.Tensor) -> CausalNetwork:
samples = _as_series(data, self.nodes)
if samples.shape[1] <= self.max_lag + 2:
raise ValueError("time dimension is too short for requested maximum lag")
current = np.concatenate([x[self.max_lag:] for x in samples], axis=0)
lagged = np.concatenate([
np.stack([x[self.max_lag - lag:-lag] for lag in range(1, self.max_lag + 1)], axis=2)
for x in samples
], axis=0) # [observations, source, lag]
candidates = lagged.reshape(len(current), -1)
candidates -= candidates.mean(axis=0, keepdims=True)
pvalues = np.ones((self.nodes, self.nodes, self.max_lag), dtype=np.float32)
mci = np.zeros_like(pvalues)
for target in range(self.nodes):
controls = lagged[:, target, :]
controls = np.column_stack([np.ones(len(controls)), controls])
gram = controls.T @ controls + self.ridge * np.eye(controls.shape[1])
projection = np.linalg.solve(gram, controls.T)
residual_x = candidates - controls @ (projection @ candidates)
y = current[:, target]
residual_y = y - controls @ (projection @ y)
numerator = residual_x.T @ residual_y
denominator = np.sqrt(np.sum(residual_x ** 2, axis=0) * np.sum(residual_y ** 2))
correlation = np.divide(numerator, denominator, out=np.zeros_like(numerator), where=denominator > 1e-12)
correlation = np.clip(correlation, -0.999999, 0.999999)
dof = max(len(y) - controls.shape[1] - 1, 1)
statistic = np.abs(correlation) * np.sqrt(dof / np.maximum(1.0 - correlation ** 2, 1e-12))
probability = 2.0 * stats.t.sf(statistic, dof)
mci[:, target, :] = correlation.reshape(self.nodes, self.max_lag)
pvalues[:, target, :] = probability.reshape(self.nodes, self.max_lag)
if self.exclude_self_links:
diagonal = np.arange(self.nodes)
pvalues[diagonal, diagonal, :] = 1.0
mci[diagonal, diagonal, :] = 0.0
edges = pvalues < self.alpha
return CausalNetwork(torch.from_numpy(edges), torch.from_numpy(pvalues), torch.from_numpy(mci))
def config(self) -> dict:
return {"nodes": self.nodes, "max_lag": self.max_lag, "alpha": self.alpha,
"ridge": self.ridge, "exclude_self_links": self.exclude_self_links}
def asymmetric_f1(reference: CausalNetwork, candidate: CausalNetwork, lag_tolerance: int = 2) -> dict[str, float | int]:
"""Compare direction and sign while allowing candidate lag error of +/- tolerance."""
ref_edges, pred_edges = reference.edges.numpy(), candidate.edges.numpy()
ref_sign, pred_sign = np.sign(reference.mci.numpy()), np.sign(candidate.mci.numpy())
def matched(edges_a, sign_a, edges_b, sign_b):
hits = 0
for source, target, lag in np.argwhere(edges_a):
lo, hi = max(0, lag - lag_tolerance), min(edges_b.shape[2], lag + lag_tolerance + 1)
hits += int(np.any(edges_b[source, target, lo:hi] &
(sign_b[source, target, lo:hi] == sign_a[source, target, lag])))
return hits
ref_count, pred_count = int(ref_edges.sum()), int(pred_edges.sum())
recall_hits = matched(ref_edges, ref_sign, pred_edges, pred_sign)
precision_hits = matched(pred_edges, pred_sign, ref_edges, ref_sign)
recall = recall_hits / ref_count if ref_count else float(pred_count == 0)
precision = precision_hits / pred_count if pred_count else float(ref_count == 0)
f1 = 2 * precision * recall / (precision + recall) if precision + recall else 0.0
return {"f1": float(f1), "precision": float(precision), "recall": float(recall),
"reference_edges": ref_count, "candidate_edges": pred_count,
"precision_matches": precision_hits, "recall_matches": recall_hits,
"lag_tolerance": int(lag_tolerance)}
def pattern_correlation(reference: np.ndarray, model: np.ndarray, weights: np.ndarray | None = None) -> float:
reference, model = np.asarray(reference, float).ravel(), np.asarray(model, float).ravel()
weights = np.ones_like(reference) if weights is None else np.asarray(weights, float).ravel()
weights = weights / weights.sum()
ref_centered = reference - np.sum(weights * reference)
model_centered = model - np.sum(weights * model)
denominator = np.sqrt(np.sum(weights * ref_centered ** 2) * np.sum(weights * model_centered ** 2))
return float(np.sum(weights * ref_centered * model_centered) / denominator) if denominator > 0 else 0.0
def taylor_s_score(reference: np.ndarray, model: np.ndarray, weights: np.ndarray | None = None) -> dict[str, float]:
"""Paper Eq. 5: (1+R)^4 / (4*(SDR + 1/SDR)^2)."""
reference, model = np.asarray(reference, float).ravel(), np.asarray(model, float).ravel()
weights = np.ones_like(reference) if weights is None else np.asarray(weights, float).ravel()
weights = weights / weights.sum()
correlation = pattern_correlation(reference, model, weights)
ref_std = np.sqrt(np.sum(weights * (reference - np.sum(weights * reference)) ** 2))
model_std = np.sqrt(np.sum(weights * (model - np.sum(weights * model)) ** 2))
ratio = model_std / max(ref_std, 1e-12)
score = (1.0 + correlation) ** 4 / (4.0 * (ratio + 1.0 / max(ratio, 1e-12)) ** 2)
return {"s_score": float(score), "pattern_correlation": correlation, "standard_deviation_ratio": float(ratio)}
class PrecipitationConstraintGP:
"""RBF plus white-noise GP for F1-to-delta-precipitation constraints."""
def __init__(self, random_state: int = 42, restarts: int = 2):
kernel = ConstantKernel(1.0, (1e-3, 1e3)) * RBF(0.15, (1e-2, 10.0)) + WhiteKernel(0.01, (1e-6, 1.0))
self.model = GaussianProcessRegressor(kernel=kernel, normalize_y=True,
n_restarts_optimizer=int(restarts), random_state=random_state)
def fit(self, f1_scores: np.ndarray, delta_precipitation: np.ndarray) -> "PrecipitationConstraintGP":
self.x_train = np.asarray(f1_scores, float).reshape(-1, 1)
self.y_train = np.asarray(delta_precipitation, float)
self.model.fit(self.x_train, self.y_train)
return self
def predict(self, f1_scores: np.ndarray) -> tuple[np.ndarray, np.ndarray, np.ndarray]:
mean, std = self.model.predict(np.asarray(f1_scores, float).reshape(-1, 1), return_std=True)
return mean, mean - 1.96 * std, mean + 1.96 * std
def state_dict(self) -> dict:
return {"x_train": self.x_train.astype(np.float64), "y_train": self.y_train.astype(np.float64),
"kernel_theta": self.model.kernel_.theta.astype(np.float64)}
@classmethod
def from_state_dict(cls, state: dict, random_state: int = 42) -> "PrecipitationConstraintGP":
instance = cls(random_state=random_state, restarts=0)
instance.model.kernel.theta = np.asarray(state["kernel_theta"])
instance.model.optimizer = None
instance.model.fit(np.asarray(state["x_train"]), np.asarray(state["y_train"]))
return instance