PoreGCN / xai_engine.py
catenate's picture
xai_engine.py: load property_means from per-ensemble JSON; classify_scenario uses manuscript dual criteria (ratio>=0.70 OR mean signed direction)
6840926 verified
Raw
History Blame Contribute Delete
25 kB
"""
xai_engine.py - Ensemble inference and XAI attribution for PoreGCN HF Space.
Frontend contract (app.py imports):
from xai_engine import (
load_ensemble,
ensemble_predict,
compute_attributions,
classify_scenario,
substructure_breakdown,
)
Design:
- load_ensemble() caches models per dataset at module level so repeated
Gradio calls do not reload from disk.
- ensemble_predict() runs all 5 models (consensus and uncertainty).
- compute_attributions() uses the best model only (compute-efficient XAI).
- classify_scenario() implements the Scenario A/B/C/D trust framework.
- substructure_breakdown() splits attributions into metal, linker, and pore.
XAI method: signed occlusion. For each atom (or pore):
attribution_i = base_prediction - prediction_with_feature_i_zeroed
Positive attribution: node drives prediction higher than baseline.
Negative attribution: node drives prediction lower than baseline.
Scenario framework (mirrors PoreGCN_unified/xai.py):
Consensus = ensemble CV < CV_THRESHOLD (10%)
Agreement = (>= AGREEMENT_THRESHOLD fraction of atoms in expected direction)
OR (mean signed attribution in expected direction)
A: consensus AND agreement (Trustworthy)
B: consensus AND NOT agreement (Overconfident)
C: NOT consensus AND agreement (Underconfident)
D: NOT consensus AND NOT agreement (Unreliable)
"""
from __future__ import annotations
import glob
import json
import logging
import os
from typing import Dict, List, Optional, Tuple
import numpy as np
import torch
from build_graph import graph_to_tensors
from config import (
AGREEMENT_THRESHOLD,
CV_THRESHOLD,
DATASETS,
DEVICE,
MODELS_DIR,
MOF_METALS,
)
from model import create_inference_model
logger = logging.getLogger(__name__)
# Inline Normalizer class (copied from data_module.py to avoid import chain)
# The normalizer is stored in each checkpoint; we just need load_state_dict.
class _Normalizer:
def __init__(self):
self.mean = 0.0
self.std = 1.0
def denorm(self, x: torch.Tensor) -> torch.Tensor:
return x * self.std + self.mean
def load_state_dict(self, d: Dict):
self.mean = float(d['mean'])
self.std = float(d['std'])
class _MultiNormalizer:
"""Per-property normalizer loaded from checkpoint state_dict."""
def __init__(self):
self.property_names: List[str] = []
self.normalizers: Dict[str, _Normalizer] = {}
def denorm(self, normed: torch.Tensor) -> torch.Tensor:
if normed.dim() == 1:
normed = normed.unsqueeze(0)
out = torch.zeros_like(normed)
for i, name in enumerate(self.property_names):
if name in self.normalizers and i < normed.shape[1]:
n = self.normalizers[name]
out[:, i] = normed[:, i] * n.std + n.mean
return out
def load_state_dict(self, d: Dict):
self.property_names = d.get('property_names', [])
self.normalizers = {}
for k, v in d.get('normalizers', {}).items():
norm = _Normalizer()
norm.load_state_dict(v)
self.normalizers[k] = norm
# =============================================================================
# Module-level cache: one entry per dataset
# =============================================================================
_CACHE: Dict[str, Dict] = {}
# Format: {dataset: {models, best_model, normalizers, best_normalizer,
# property_names, property_means}}
# =============================================================================
# load_ensemble
# =============================================================================
def load_ensemble(
dataset: str,
device: str = DEVICE,
) -> Tuple[List, object, object, List[str]]:
"""
Load 5-model ensemble + best model for a given dataset.
Checks module-level cache first; loads from disk only on the first call
for each dataset. This is safe for Gradio's multi-request environment
because model weights are read-only after loading.
Ensemble directory structure expected:
models/<dataset>_ensemble/
fold_0_best.ckpt (or similar pattern)
fold_1_best.ckpt
...
property_means.json
Each checkpoint must be a dict with either:
- 'model_state' + 'model_config' + 'normalizer' (CV format)
- 'model_state_dict' + 'config' + 'normalizer_state_dict' (final format)
Args:
dataset: One of DATASETS ('hmof_gas', 'core_mof', 'hmof_geometric')
device: Target device string (always 'cpu' for HF free tier)
Returns:
(all_models, best_model, best_normalizer, property_names)
all_models: list of 5 PoreGCN instances in eval mode
best_model: the model with lowest val_loss (used for XAI)
best_normalizer: _MultiNormalizer for the best model's checkpoint
property_names: list of property strings in prediction order
"""
global _CACHE
if dataset in _CACHE:
entry = _CACHE[dataset]
return (
entry['models'],
entry['best_model'],
entry['best_normalizer'],
entry['property_names'],
)
ensemble_dir = os.path.join(MODELS_DIR, f'{dataset}_ensemble')
if not os.path.isdir(ensemble_dir):
raise FileNotFoundError(
f'Ensemble directory not found: {ensemble_dir}. '
f'Expected: models/{dataset}_ensemble/ with checkpoint .ckpt or .pt files.'
)
# Gather checkpoint files
ckpt_files = sorted(
glob.glob(os.path.join(ensemble_dir, '*.ckpt')) +
glob.glob(os.path.join(ensemble_dir, '*.pt'))
)
if not ckpt_files:
raise FileNotFoundError(
f'No checkpoint files (.ckpt or .pt) found in {ensemble_dir}.'
)
# Load all checkpoints and rank by val_loss
loaded: List[Tuple[float, Dict, str]] = []
for path in ckpt_files:
try:
ckpt = torch.load(path, map_location='cpu', weights_only=False)
val_loss = float(ckpt.get('val_loss', ckpt.get('val_mae', float('inf'))))
loaded.append((val_loss, ckpt, path))
except Exception as exc:
logger.warning('Failed to load checkpoint %s: %s', path, exc)
if not loaded:
raise RuntimeError(f'Could not load any checkpoints from {ensemble_dir}.')
loaded.sort(key=lambda x: x[0])
top5 = loaded[:5]
logger.info(
'Loaded %d checkpoints for %s, using top %d.',
len(loaded), dataset, len(top5),
)
# Instantiate models
models: List = []
normalizers: List[_MultiNormalizer] = []
property_names_from_ckpt: Optional[List[str]] = None
for val_loss, ckpt, path in top5:
try:
model = create_inference_model(ckpt, device=device)
models.append(model)
# Load normalizer
normalizer = _MultiNormalizer()
norm_sd = ckpt.get('normalizer', ckpt.get('normalizer_state_dict', None))
if norm_sd is not None:
normalizer.load_state_dict(norm_sd)
normalizers.append(normalizer)
# Capture property names from first successful checkpoint
if property_names_from_ckpt is None:
cfg = ckpt.get('model_config', ckpt.get('config', {}))
property_names_from_ckpt = cfg.get('property_names', None)
if property_names_from_ckpt is None:
property_names_from_ckpt = normalizer.property_names or []
except Exception as exc:
logger.warning('Skipping checkpoint %s: %s', path, exc)
if not models:
raise RuntimeError(f'All checkpoints in {ensemble_dir} failed to load.')
best_model = models[0]
best_normalizer = normalizers[0]
property_names = property_names_from_ckpt or []
# Load property means (for XAI direction logic)
means_path = os.path.join(ensemble_dir, 'property_means.json')
property_means: Dict[str, float] = {}
if os.path.exists(means_path):
with open(means_path) as f:
property_means = json.load(f)
else:
logger.warning(
'property_means.json not found in %s. '
'XAI direction checks will use 0.0 as default mean.',
ensemble_dir,
)
_CACHE[dataset] = {
'models': models,
'normalizers': normalizers,
'best_model': best_model,
'best_normalizer': best_normalizer,
'property_names': property_names,
'property_means': property_means,
}
logger.info(
'Ensemble for %s ready: %d models, %d properties.',
dataset, len(models), len(property_names),
)
return models, best_model, best_normalizer, property_names
# =============================================================================
# ensemble_predict
# =============================================================================
def ensemble_predict(
graph: Dict,
models: List,
normalizer,
device: str = DEVICE,
) -> Dict:
"""
Run all ensemble models on a single graph.
Args:
graph: Graph dict returned by cif_to_graph().
models: List of PoreGCN instances (from load_ensemble()).
normalizer: _MultiNormalizer for the best model; used to denorm outputs.
All models in the ensemble share the same normalizer because
they were trained on the same data split.
device: Target device (always 'cpu').
Returns:
Dict keyed by property name:
{prop: {'mean': float, 'std': float, 'cv': float}}
Values are in physical (denormalized) units.
"""
tensors = graph_to_tensors(graph, device=device)
all_preds: List[np.ndarray] = []
for model in models:
model.eval()
with torch.no_grad():
raw = model(
tensors['atom_fea'],
tensors['bond_types'],
tensors['nbr_fea_idx'],
tensors['crystal_atom_idx'],
tensors['pore_fea'],
tensors['atom_pore_edges'],
tensors['crystal_pore_idx'],
) # [1, n_props]
denormed = normalizer.denorm(raw.cpu()) # [1, n_props]
all_preds.append(denormed[0].numpy())
all_preds_arr = np.array(all_preds) # [n_models, n_props]
means = all_preds_arr.mean(axis=0)
stds = all_preds_arr.std(axis=0)
property_names = normalizer.property_names
results: Dict[str, Dict[str, float]] = {}
for i, prop in enumerate(property_names):
mean_val = float(means[i])
std_val = float(stds[i])
cv = abs(std_val / mean_val) if abs(mean_val) > 1e-8 else float(std_val)
results[prop] = {'mean': mean_val, 'std': std_val, 'cv': cv}
return results
# =============================================================================
# compute_attributions
# =============================================================================
def compute_attributions(
graph: Dict,
best_model,
normalizer,
property_means: Dict[str, float],
target_prop: str,
device: str = DEVICE,
xai_method: str = 'fast',
progress_callback=None,
) -> Dict:
"""
Signed attribution per atom AND per pore for one target property.
Two methods, selectable via `xai_method`:
'fast' (default, ~5 s) — gradient x input. One forward + one backward
pass; per-atom score = sum_f (atom_fea[i,f] * d pred / d atom_fea[i,f]).
Approximates signed occlusion in sign for most atoms.
'full' (~3 min for 200 atoms) — signed occlusion. For each atom_i,
zero atom_fea[i] and rerun the model; for each pore_j, zero
pore_fea[j] and remove its edges; signed attribution =
base_pred - masked_pred. This is the exact method reported in
the manuscript and used to generate xcomprehensive_xai_results/.
`progress_callback(step, total, message)` is called periodically when
xai_method='full' so the UI can render a progress bar; ignored when
xai_method='fast'.
Args:
graph: Graph dict from cif_to_graph().
best_model: PoreGCN instance (lowest val_loss).
normalizer: _MultiNormalizer for denormalization.
property_means: Dict {prop_name: training_mean} loaded from
property_means.json in the ensemble dir.
target_prop: Property name string to explain.
device: Always 'cpu'.
Returns:
Dict:
per_atom np.ndarray [N_atoms] signed attributions
per_pore np.ndarray [N_pores] signed attributions
expected_direction '+' if prediction > mean, '-' otherwise
agreement_frac float in [0, 1] (fraction of attrs with expected sign)
"""
tensors = graph_to_tensors(graph, device=device)
property_names = normalizer.property_names
if target_prop not in property_names:
raise ValueError(
f'Property "{target_prop}" not in model property list: {property_names}'
)
target_idx = property_names.index(target_prop)
best_model.eval()
n_atoms = graph['n_atoms']
n_pores = graph['n_pores']
def _forward() -> float:
with torch.no_grad():
raw = best_model(
tensors['atom_fea'],
tensors['bond_types'],
tensors['nbr_fea_idx'],
tensors['crystal_atom_idx'],
tensors['pore_fea'],
tensors['atom_pore_edges'],
tensors['crystal_pore_idx'],
)
return float(normalizer.denorm(raw.cpu())[0, target_idx].item())
base_pred = _forward()
# ----- xai_method='full' (signed occlusion, manuscript method) -----
if xai_method == 'full':
atom_signed = np.zeros(n_atoms, dtype=np.float64)
# Atom occlusion: mask one atom at a time, measure prediction change
orig_atom = tensors['atom_fea'].clone()
for i in range(n_atoms):
masked = orig_atom.clone()
masked[i] = 0.0
tensors['atom_fea'] = masked
masked_pred = _forward()
atom_signed[i] = base_pred - masked_pred
if progress_callback is not None and (i + 1) % 10 == 0:
progress_callback(i + 1, n_atoms + n_pores,
f'Occlusion: atom {i+1}/{n_atoms}')
tensors['atom_fea'] = orig_atom
# Pore occlusion: zero pore feature and remove its edges
pore_signed = np.zeros(n_pores, dtype=np.float64)
if n_pores > 0:
orig_pore = tensors['pore_fea'].clone()
orig_edges = tensors['atom_pore_edges']
for j in range(n_pores):
masked_pore = orig_pore.clone()
masked_pore[j] = 0.0
if orig_edges.numel() > 0:
keep = orig_edges[1] != j
masked_edges = orig_edges[:, keep]
else:
masked_edges = orig_edges
tensors['pore_fea'] = masked_pore
tensors['atom_pore_edges'] = masked_edges
masked_pred = _forward()
pore_signed[j] = base_pred - masked_pred
if progress_callback is not None and (j + 1) % 10 == 0:
progress_callback(n_atoms + j + 1, n_atoms + n_pores,
f'Occlusion: pore {j+1}/{n_pores}')
tensors['pore_fea'] = orig_pore
tensors['atom_pore_edges'] = orig_edges
prop_mean = property_means.get(target_prop, 0.0)
expected_positive = base_pred > prop_mean
expected_direction = '+' if expected_positive else '-'
all_signed = np.concatenate([atom_signed, pore_signed]) if n_pores > 0 else atom_signed
n_total = len(all_signed)
if n_total > 0:
n_matching = int(np.sum(all_signed > 0)) if expected_positive \
else int(np.sum(all_signed < 0))
agreement_frac = n_matching / n_total
mean_signed = float(np.mean(all_signed))
else:
agreement_frac = 0.0
mean_signed = 0.0
return {
'per_atom': atom_signed.astype(np.float32),
'per_pore': pore_signed.astype(np.float32),
'expected_direction': expected_direction,
'agreement_frac': float(agreement_frac),
'mean_signed': mean_signed,
'expected_positive': bool(expected_positive),
}
# ----- xai_method='fast' (gradient x input surrogate, default) -----
# One forward + one backward pass.
# Per-atom score = sum_f (atom_fea_{i,f} * d pred / d atom_fea_{i,f}).
# Sign convention preserved: positive = atom drives prediction up.
orig_atom_fea = tensors['atom_fea'].clone()
atom_fea_grad = orig_atom_fea.detach().clone().requires_grad_(True)
if n_pores > 0:
pore_fea_grad = tensors['pore_fea'].detach().clone().requires_grad_(True)
else:
pore_fea_grad = tensors['pore_fea']
# Forward with grad enabled
with torch.enable_grad():
raw_pred = best_model(
atom_fea_grad,
tensors['bond_types'],
tensors['nbr_fea_idx'],
tensors['crystal_atom_idx'],
pore_fea_grad,
tensors['atom_pore_edges'],
tensors['crystal_pore_idx'],
)
# Use raw (normalised) target as the scalar to backprop through. Denormalisation
# is a constant scale + shift so it does not affect the relative attribution sign/rank.
target_scalar = raw_pred[0, target_idx]
target_scalar.backward()
atom_signed = (atom_fea_grad.detach() * atom_fea_grad.grad).sum(dim=1).cpu().numpy().astype(np.float64)
if n_pores > 0 and pore_fea_grad.grad is not None:
pore_signed_grad = (pore_fea_grad.detach() * pore_fea_grad.grad).sum(dim=1).cpu().numpy().astype(np.float64)
else:
pore_signed_grad = None
# Restore
tensors['atom_fea'] = orig_atom_fea
# --- Pore attributions (also via gradient x input from the same backward pass) ---
pore_signed = np.zeros(n_pores, dtype=np.float64)
if n_pores > 0 and pore_signed_grad is not None:
pore_signed = pore_signed_grad
# Agreement check: fraction of attributions in expected direction
prop_mean = property_means.get(target_prop, 0.0)
expected_positive = base_pred > prop_mean
expected_direction = '+' if expected_positive else '-'
all_signed = np.concatenate([atom_signed, pore_signed]) if n_pores > 0 else atom_signed
n_total = len(all_signed)
if n_total > 0:
if expected_positive:
n_matching = int(np.sum(all_signed > 0))
else:
n_matching = int(np.sum(all_signed < 0))
agreement_frac = n_matching / n_total
mean_signed = float(np.mean(all_signed))
else:
agreement_frac = 0.0
mean_signed = 0.0
return {
'per_atom': atom_signed.astype(np.float32),
'per_pore': pore_signed.astype(np.float32),
'expected_direction': expected_direction,
'agreement_frac': float(agreement_frac),
'mean_signed': mean_signed,
'expected_positive': bool(expected_positive),
}
# =============================================================================
# classify_scenario
# =============================================================================
def classify_scenario(
prediction_mean: float,
prediction_std: float,
agreement_frac: float,
property_mean: float,
mean_signed: float = 0.0,
expected_positive: bool = None,
) -> Tuple[str, str]:
"""
Classify prediction trustworthiness into Scenario A/B/C/D.
Mirrors the dual-criteria framework in PoreGCN/xai.py exactly:
Consensus: CV < CV_THRESHOLD (10%)
Agreement (EITHER satisfies):
Criterion 1: agreement_frac >= AGREEMENT_THRESHOLD (70%)
Criterion 2: mean signed attribution has the expected direction
Earlier versions of this function used the relaxed fallback
`agreement_frac >= 0.5` for criterion 2, which silently flipped
Scenario A/B/C/D classifications relative to the manuscript on
structures whose attributions are bipolar (e.g. signed-occlusion
output where some atoms increase the prediction and others decrease
it). Fixed to use the actual mean-signed value.
Args:
prediction_mean: Ensemble mean prediction (physical units).
prediction_std: Ensemble standard deviation.
agreement_frac: Fraction of attributions in the expected direction.
property_mean: Training-set mean of this property.
mean_signed: Mean of the signed attributions (for criterion 2).
expected_positive: If None, derived from prediction_mean > property_mean.
Returns:
(scenario_letter, human_label)
"""
cv = abs(prediction_std / prediction_mean) if abs(prediction_mean) > 1e-8 else abs(prediction_std)
high_consensus = cv < CV_THRESHOLD
if expected_positive is None:
expected_positive = prediction_mean > property_mean
# Manuscript dual criteria, exactly as in PoreGCN/xai.py check_xai_agreement
ratio_ok = agreement_frac >= AGREEMENT_THRESHOLD
if expected_positive:
mean_dir_ok = mean_signed > 0.0
else:
mean_dir_ok = mean_signed < 0.0
is_agreed = ratio_ok or mean_dir_ok
scenario_labels = {
'A': 'Trustworthy (high consensus, XAI agrees)',
'B': 'Overconfident (high consensus, XAI disagrees)',
'C': 'Underconfident (low consensus, XAI agrees)',
'D': 'Unreliable (low consensus, XAI disagrees)',
}
if high_consensus and is_agreed:
scenario = 'A'
elif high_consensus and not is_agreed:
scenario = 'B'
elif not high_consensus and is_agreed:
scenario = 'C'
else:
scenario = 'D'
return scenario, scenario_labels[scenario]
# =============================================================================
# substructure_breakdown
# =============================================================================
def substructure_breakdown(
graph: Dict,
per_atom_attrs: np.ndarray,
per_pore_attrs: np.ndarray,
) -> Dict:
"""
Split total attribution into metal, linker, and pore fractions.
Metal/linker classification uses element symbols decoded from the
one-hot portion of atom_fea (first 92 dimensions = element 1..92).
MOF_METALS set from config.py covers all common MOF metal centres.
Args:
graph: Graph dict from cif_to_graph() (must contain 'structure').
per_atom_attrs: np.ndarray [N_atoms] signed attributions.
per_pore_attrs: np.ndarray [N_pores] signed attributions.
Returns:
Dict:
metal_frac float |metal attributions| / total |attributions|
linker_frac float |linker attributions| / total |attributions|
pore_frac float |pore attributions| / total |attributions|
"""
structure = graph.get('structure')
n_atoms = graph['n_atoms']
if structure is not None and len(structure) == n_atoms:
elements = [str(site.specie) for site in structure]
else:
# Fallback: decode from one-hot features
ELEMENT_NAMES = [
'H', 'He', 'Li', 'Be', 'B', 'C', 'N', 'O', 'F', 'Ne',
'Na', 'Mg', 'Al', 'Si', 'P', 'S', 'Cl', 'Ar', 'K', 'Ca',
'Sc', 'Ti', 'V', 'Cr', 'Mn', 'Fe', 'Co', 'Ni', 'Cu', 'Zn',
'Ga', 'Ge', 'As', 'Se', 'Br', 'Kr', 'Rb', 'Sr', 'Y', 'Zr',
'Nb', 'Mo', 'Tc', 'Ru', 'Rh', 'Pd', 'Ag', 'Cd', 'In', 'Sn',
'Sb', 'Te', 'I', 'Xe', 'Cs', 'Ba', 'La', 'Ce', 'Pr', 'Nd',
'Pm', 'Sm', 'Eu', 'Gd', 'Tb', 'Dy', 'Ho', 'Er', 'Tm', 'Yb',
'Lu', 'Hf', 'Ta', 'W', 'Re', 'Os', 'Ir', 'Pt', 'Au', 'Hg',
'Tl', 'Pb', 'Bi', 'Po', 'At', 'Rn', 'Fr', 'Ra', 'Ac', 'Th',
'Pa', 'U',
]
atom_fea = graph['atom_fea'] # [N_atoms, 120]
elements = []
for i in range(n_atoms):
one_hot = atom_fea[i, :92]
idx = int(np.argmax(one_hot))
elements.append(ELEMENT_NAMES[idx] if idx < len(ELEMENT_NAMES) else 'X')
abs_atom = np.abs(per_atom_attrs)
abs_pore = np.abs(per_pore_attrs)
metal_sum = sum(abs_atom[i] for i, e in enumerate(elements) if e in MOF_METALS)
linker_sum = sum(abs_atom[i] for i, e in enumerate(elements) if e not in MOF_METALS)
pore_sum = float(np.sum(abs_pore))
total = metal_sum + linker_sum + pore_sum + 1e-12
return {
'metal_frac': float(metal_sum / total),
'linker_frac': float(linker_sum / total),
'pore_frac': float(pore_sum / total),
}