diff --git a/Plans/person4-interpretability-plan.md b/Plans/person4-interpretability-plan.md new file mode 100644 index 0000000..3459b95 --- /dev/null +++ b/Plans/person4-interpretability-plan.md @@ -0,0 +1,58 @@ +# Person 4: Explainability & Interpretation Plan + +**Date:** 2026-08-10 +**Branch:** `person4-interpretability` +**Assignees:** Leah, Apollo +**Status:** In progress + +## Goal + +Build the explainability layer (`confostate/explain/`) that turns model +predictions and feature vectors into human-readable justifications with +literature citations. + +## Approach + +1. **Package location:** Move from root `explain/` stub into + `confostate/explain/` to match the workplan and package layout. +2. **Feature importance:** Support three paths (in order of preference): + - Explicit importances passed by the caller (from Person 3's pipeline) + - Permutation importance for sklearn-compatible models with `predict_proba` + - Heuristic importance from feature values vs. LeuT state profiles (stub + path while models are unavailable) +3. **Rendering:** Template-based natural language (no LLM dependency for v1). + Keeps explanations deterministic, testable, and free of GPU/API cost. +4. **Citations:** Curated DOI/PubMed mappings for LeuT states and key + structural features, pulled from the annotations reference column. +5. **Public API:** `explain(prediction) -> ExplanationResult` with text and + markdown output formats. + +## Key decisions + +| Decision | Choice | Rationale | +|----------|--------|-----------| +| LLM rendering | Deferred | Workplan mentions Llama/ASU tokens; template v1 unblocks Person 5 CLI | +| SHAP | Deferred | Heavy dependency; permutation + heuristics sufficient for v1 | +| sklearn | Optional import | Used when available; heuristic fallback keeps package lightweight | +| Synthetic data | Heuristic profiles | Workplan: "make up sh**t to move forward" until Person 3 delivers models | + +## Deliverables + +- [x] `confostate/explain/importance.py` +- [x] `confostate/explain/render.py` +- [x] `confostate/explain/citations.py` +- [x] `confostate/explain/__init__.py` with `explain()` +- [x] `tests/test_explain.py` +- [x] `docs/explainability.md` + +## Dependencies + +- **Person 2:** Feature names and semantics (`docs/features/features.md`) +- **Person 3:** Trained models and real permutation importances (future) +- **Blocks:** Person 5 CLI integration + +## Future work + +- SHAP values for tree/neural models +- LLM polish pass (local Llama at ASU) +- Multi-family citation and profile tables diff --git a/confostate/explain/__init__.py b/confostate/explain/__init__.py new file mode 100644 index 0000000..ebc335d --- /dev/null +++ b/confostate/explain/__init__.py @@ -0,0 +1,141 @@ +"""Explainability layer for conformational-state predictions.""" + +from __future__ import annotations + +from typing import Any, Literal, Optional, Sequence, Union, overload + +from confostate.explain._types import ( + Citation, + ExplanationResult, + FeatureImportance, + PredictionResult, +) +from confostate.explain.citations import get_citations, state_label +from confostate.explain.importance import ( + FEATURE_DISPLAY_NAMES, + display_name, + rank_importances, + resolve_importances, +) +from confostate.explain.render import render_markdown, render_text + +__all__ = [ + "Citation", + "ExplanationResult", + "FeatureImportance", + "PredictionResult", + "FEATURE_DISPLAY_NAMES", + "display_name", + "explain", + "state_label", +] + + +@overload +def explain( + prediction: Union[PredictionResult, dict[str, Any]], + *, + top_k: int = 5, + importances: Optional[Sequence[FeatureImportance]] = None, + model: Optional[Any] = None, + feature_matrix: Optional[Any] = None, + feature_names: Optional[Sequence[str]] = None, + coefficients: Optional[dict[str, float]] = None, + output_format: Literal["result"] = "result", +) -> ExplanationResult: ... + + +@overload +def explain( + prediction: Union[PredictionResult, dict[str, Any]], + *, + top_k: int = 5, + importances: Optional[Sequence[FeatureImportance]] = None, + model: Optional[Any] = None, + feature_matrix: Optional[Any] = None, + feature_names: Optional[Sequence[str]] = None, + coefficients: Optional[dict[str, float]] = None, + output_format: Literal["text", "markdown"], +) -> str: ... + + +def explain( + prediction: Union[PredictionResult, dict[str, Any]], + *, + top_k: int = 5, + importances: Optional[Sequence[FeatureImportance]] = None, + model: Optional[Any] = None, + feature_matrix: Optional[Any] = None, + feature_names: Optional[Sequence[str]] = None, + coefficients: Optional[dict[str, float]] = None, + output_format: Literal["result", "text", "markdown"] = "result", +) -> Union[ExplanationResult, str]: + """ + Generate a human-readable explanation for a classifier prediction. + + Parameters + ---------- + prediction + PredictionResult or dict with pdb_id, predicted_state, + probabilities, and features. + top_k + Number of top features to include in the explanation. + importances + Pre-computed feature importances (from Person 3's pipeline). + model + Optional sklearn-compatible model for permutation importance. + feature_matrix + Feature matrix for permutation importance (single row or batch). + feature_names + Feature names aligned with feature_matrix columns. + coefficients + Linear-model coefficients keyed by feature name. + output_format + ``"result"`` returns ExplanationResult; ``"text"`` or + ``"markdown"`` returns rendered strings. + + Returns + ------- + ExplanationResult or str + Full explanation object or rendered text. + """ + if not isinstance(prediction, PredictionResult): + prediction = PredictionResult.from_dict(prediction) + + resolved, method = resolve_importances( + features=prediction.features, + predicted_state=prediction.predicted_state, + family=prediction.family, + importances=importances, + model=model, + feature_matrix=feature_matrix, + feature_names=feature_names, + coefficients=coefficients, + ) + top_features = rank_importances(resolved, top_k=top_k) + citations = get_citations( + predicted_state=prediction.predicted_state, + feature_names=[item.feature_name for item in top_features], + family=prediction.family, + ) + + confidence = prediction.probabilities.get(prediction.predicted_state, 0.0) + text = render_text(prediction, top_features, citations, method) + markdown = render_markdown(prediction, top_features, citations, method) + + result = ExplanationResult( + pdb_id=prediction.pdb_id, + predicted_state=prediction.predicted_state, + confidence=confidence, + text=text, + markdown=markdown, + top_features=top_features, + citations=citations, + method=method, + ) + + if output_format == "text": + return result.text + if output_format == "markdown": + return result.markdown + return result diff --git a/confostate/explain/_types.py b/confostate/explain/_types.py new file mode 100644 index 0000000..ebf2822 --- /dev/null +++ b/confostate/explain/_types.py @@ -0,0 +1,66 @@ +"""Shared types for the explainability layer.""" + +from __future__ import annotations + +from dataclasses import dataclass, field +from typing import Any, Optional + + +@dataclass +class PredictionResult: + """Output of a conformational-state classifier.""" + + pdb_id: str + predicted_state: str + probabilities: dict[str, float] + features: dict[str, float] + family: str = "LeuT" + + @classmethod + def from_dict(cls, data: dict[str, Any]) -> PredictionResult: + """Build a prediction from a plain dict.""" + return cls( + pdb_id=str(data["pdb_id"]), + predicted_state=str(data["predicted_state"]), + probabilities={ + str(k): float(v) for k, v in data["probabilities"].items() + }, + features={str(k): float(v) for k, v in data["features"].items()}, + family=str(data.get("family", "LeuT")), + ) + + +@dataclass +class FeatureImportance: + """Importance of a single feature for a prediction.""" + + feature_name: str + display_name: str + importance: float + value: float + direction: str # "supports" or "opposes" + + +@dataclass +class Citation: + """Literature reference for a state or feature.""" + + key: str + label: str + reference: str + doi: Optional[str] = None + pubmed_id: Optional[str] = None + + +@dataclass +class ExplanationResult: + """Full explanation for a single prediction.""" + + pdb_id: str + predicted_state: str + confidence: float + text: str + markdown: str + top_features: list[FeatureImportance] = field(default_factory=list) + citations: list[Citation] = field(default_factory=list) + method: str = "heuristic" diff --git a/confostate/explain/citations.py b/confostate/explain/citations.py new file mode 100644 index 0000000..2eb03f9 --- /dev/null +++ b/confostate/explain/citations.py @@ -0,0 +1,148 @@ +"""Literature citations for explainability output.""" + +from __future__ import annotations + +from typing import Sequence + +from confostate.explain._types import Citation + +# Curated references for LeuT conformational states and structural motifs. +LEUT_STATE_CITATIONS: dict[str, Citation] = { + "IF_open": Citation( + key="IF_open", + label="Inward-facing open", + reference="Yamashita et al., Nature 2005", + doi="10.1038/nature03578", + pubmed_id="15880103", + ), + "OF_open": Citation( + key="OF_open", + label="Outward-facing open", + reference="Singh et al., Science 2008", + doi="10.1126/science.1159299", + pubmed_id="18403708", + ), + "Occluded": Citation( + key="Occluded", + label="Occluded", + reference="Shi et al., Nature 2008", + doi="10.1038/nature06932", + pubmed_id="18497808", + ), + "Intermediate": Citation( + key="Intermediate", + label="Intermediate", + reference="Claxton et al., J. Mol. Biol. 2010", + doi="10.1016/j.jmb.2010.09.022", + pubmed_id="20869387", + ), +} + +LEUT_FEATURE_CITATIONS: dict[str, Citation] = { + "domain_TM1_TM7_distance": Citation( + key="domain_TM1_TM7_distance", + label="TM1–TM7 gate distance", + reference="Singh et al., Science 2008 (outward-open gate)", + doi="10.1126/science.1159299", + pubmed_id="18403708", + ), + "domain_gate_TM1_TM6_distance": Citation( + key="domain_gate_TM1_TM6_distance", + label="TM1–TM6 gate opening", + reference="Yamashita et al., Nature 2005 (inward-open gate)", + doi="10.1038/nature03578", + pubmed_id="15880103", + ), + "cavity_volume": Citation( + key="cavity_volume", + label="Binding-site cavity", + reference="Yamashita et al., Nature 2005", + doi="10.1038/nature03578", + pubmed_id="15880103", + ), + "rmsd_IF_open": Citation( + key="rmsd_IF_open", + label="RMSD to inward-open reference", + reference="Yamashita et al., Nature 2005 (3F3A)", + doi="10.1038/nature03578", + pubmed_id="15880103", + ), + "rmsd_OF_open": Citation( + key="rmsd_OF_open", + label="RMSD to outward-open reference", + reference="Singh et al., Science 2008 (3F3E)", + doi="10.1126/science.1159299", + pubmed_id="18403708", + ), + "rmsd_Occluded": Citation( + key="rmsd_Occluded", + label="RMSD to occluded reference", + reference="Shi et al., Nature 2008 (3F4J)", + doi="10.1038/nature06932", + pubmed_id="18497808", + ), + "rmsd_Intermediate": Citation( + key="rmsd_Intermediate", + label="RMSD to intermediate reference", + reference="Claxton et al., J. Mol. Biol. 2010 (3USI)", + doi="10.1016/j.jmb.2010.09.022", + pubmed_id="20869387", + ), +} + +STATE_LABELS: dict[str, str] = { + "IF_open": "inward-facing open", + "OF_open": "outward-facing open", + "Occluded": "occluded", + "Intermediate": "intermediate", +} + + +def state_label(state: str) -> str: + """Return a human-readable state name.""" + return STATE_LABELS.get(state, state.replace("_", " ")) + + +def get_citations( + predicted_state: str, + feature_names: Sequence[str], + family: str = "LeuT", +) -> list[Citation]: + """ + Collect literature citations for a predicted state and its top features. + + Deduplicates by DOI when the same reference covers multiple items. + """ + if family != "LeuT": + return [] + + seen_dois: set[str] = set() + citations: list[Citation] = [] + + state_citation = LEUT_STATE_CITATIONS.get(predicted_state) + if state_citation is not None: + citations.append(state_citation) + if state_citation.doi: + seen_dois.add(state_citation.doi) + + for name in feature_names: + feature_citation = LEUT_FEATURE_CITATIONS.get(name) + if feature_citation is None: + continue + if feature_citation.doi and feature_citation.doi in seen_dois: + continue + citations.append(feature_citation) + if feature_citation.doi: + seen_dois.add(feature_citation.doi) + + return citations + + +def format_citation(citation: Citation) -> str: + """Format a single citation for plain-text output.""" + parts = [citation.reference] + if citation.doi: + parts.append(f"DOI: {citation.doi}") + if citation.pubmed_id: + parts.append(f"PubMed: {citation.pubmed_id}") + return " — ".join(parts) diff --git a/confostate/explain/importance.py b/confostate/explain/importance.py new file mode 100644 index 0000000..f589a02 --- /dev/null +++ b/confostate/explain/importance.py @@ -0,0 +1,290 @@ +"""Feature importance extraction for conformational-state predictions.""" + +from __future__ import annotations + +from typing import Any, Optional, Sequence + +import numpy as np + +from confostate.explain._types import FeatureImportance + +FEATURE_DISPLAY_NAMES: dict[str, str] = { + "cavity_volume": "binding-site cavity volume", + "cavity_accessibility_in": "inward solvent accessibility", + "cavity_accessibility_out": "outward solvent accessibility", + "domain_TM1_TM7_distance": "TM1–TM7 gate distance", + "domain_TM1_TM7_angle": "TM1–TM7 helix angle", + "domain_TM1_TM6_distance": "TM1–TM6 distance", + "domain_TM5_TM7_distance": "TM5–TM7 distance", + "domain_TM3_TM10_distance": "TM3–TM10 distance", + "domain_gate_TM1_TM6_distance": "TM1–TM6 gate-opening distance", + "rmsd_OF_open": "RMSD to outward-open reference (3F3E)", + "rmsd_IF_open": "RMSD to inward-open reference (3F3A)", + "rmsd_Occluded": "RMSD to occluded reference (3F4J)", + "rmsd_Intermediate": "RMSD to intermediate reference (3USI)", + "rmsd_min": "minimum RMSD across references", + "opm_tilt_angle": "membrane tilt angle", + "opm_rotation_angle": "membrane rotation angle", + "opm_depth": "membrane embedding depth", + "orientation_principal_axis_x": "principal axis (x)", + "orientation_principal_axis_y": "principal axis (y)", + "orientation_principal_axis_z": "principal axis (z)", +} + +# Typical feature values per LeuT state for heuristic importance (stub data). +# Values are approximate literature-informed placeholders until Person 3 +# provides real training statistics. +LEUT_STATE_PROFILES: dict[str, dict[str, float]] = { + "IF_open": { + "domain_TM1_TM7_distance": 18.0, + "domain_gate_TM1_TM6_distance": 22.0, + "rmsd_IF_open": 1.5, + "cavity_accessibility_in": 0.55, + "cavity_accessibility_out": 0.25, + }, + "OF_open": { + "domain_TM1_TM7_distance": 22.0, + "domain_gate_TM1_TM6_distance": 26.0, + "rmsd_OF_open": 1.5, + "cavity_accessibility_in": 0.25, + "cavity_accessibility_out": 0.55, + }, + "Occluded": { + "domain_TM1_TM7_distance": 16.0, + "domain_gate_TM1_TM6_distance": 18.0, + "rmsd_Occluded": 1.5, + "cavity_volume": 450.0, + }, + "Intermediate": { + "domain_TM1_TM7_distance": 19.0, + "domain_gate_TM1_TM6_distance": 20.0, + "rmsd_Intermediate": 1.5, + }, +} + +RMSD_STATE_MAP = { + "IF_open": "rmsd_IF_open", + "OF_open": "rmsd_OF_open", + "Occluded": "rmsd_Occluded", + "Intermediate": "rmsd_Intermediate", +} + + +def display_name(feature_name: str) -> str: + """Return a human-readable label for a feature.""" + return FEATURE_DISPLAY_NAMES.get( + feature_name, feature_name.replace("_", " ") + ) + + +def rank_importances( + importances: Sequence[FeatureImportance], top_k: int = 5 +) -> list[FeatureImportance]: + """Return the top-k features by absolute importance.""" + ranked = sorted( + importances, key=lambda item: abs(item.importance), reverse=True + ) + return ranked[:top_k] + + +def extract_coefficient_importance( + coefficients: dict[str, float], + feature_values: dict[str, float], + predicted_state: str, +) -> list[FeatureImportance]: + """ + Compute importance from linear-model coefficients. + + Importance is |coefficient * value| for the predicted state's class. + """ + results: list[FeatureImportance] = [] + for name, coef in coefficients.items(): + if name not in feature_values: + continue + value = feature_values[name] + score = coef * value + results.append( + FeatureImportance( + feature_name=name, + display_name=display_name(name), + importance=score, + value=value, + direction="supports" if score >= 0 else "opposes", + ) + ) + return results + + +def extract_permutation_importance( + model: Any, + feature_matrix: np.ndarray, + feature_names: Sequence[str], + baseline_proba: np.ndarray, + n_repeats: int = 5, + random_state: int = 0, +) -> list[FeatureImportance]: + """ + Model-agnostic permutation importance for a single prediction. + + Shuffles each feature column and measures the drop in predicted + probability for the baseline class. + """ + rng = np.random.default_rng(random_state) + baseline_score = float(np.max(baseline_proba)) + predicted_index = int(np.argmax(baseline_proba)) + row_idx = 0 + n_samples = feature_matrix.shape[0] + importances = np.zeros(len(feature_names)) + + for col_idx, name in enumerate(feature_names): + drops: list[float] = [] + original_value = float(feature_matrix[row_idx, col_idx]) + for _ in range(n_repeats): + permuted = feature_matrix.copy() + if n_samples > 1: + other_values = np.delete(feature_matrix[:, col_idx], row_idx) + permuted[row_idx, col_idx] = rng.choice(other_values) + else: + scale = max(abs(original_value), 1.0) + permuted[row_idx, col_idx] = original_value + rng.normal( + 0, 0.2 * scale + ) + proba = model.predict_proba(permuted)[row_idx] + drops.append(baseline_score - float(proba[predicted_index])) + importances[col_idx] = float(np.mean(drops)) + + results: list[FeatureImportance] = [] + for name, score in zip(feature_names, importances): + value = float(feature_matrix[row_idx, list(feature_names).index(name)]) + results.append( + FeatureImportance( + feature_name=name, + display_name=display_name(name), + importance=score, + value=value, + direction="supports" if score >= 0 else "opposes", + ) + ) + return results + + +def extract_heuristic_importance( + features: dict[str, float], + predicted_state: str, + family: str = "LeuT", +) -> list[FeatureImportance]: + """ + Estimate feature importance without a trained model. + + Compares each numeric feature to family-specific state profiles. + Lower RMSD to the matching reference and closer agreement with the + predicted-state profile yield higher scores. + """ + if family != "LeuT": + return _generic_heuristic_importance(features, predicted_state) + + profile = LEUT_STATE_PROFILES.get(predicted_state, {}) + results: list[FeatureImportance] = [] + + for name, value in features.items(): + if not isinstance(value, (int, float)) or np.isnan(value): + continue + + score = 0.0 + direction = "supports" + + if name in profile: + expected = profile[name] + rel_error = abs(value - expected) / max(abs(expected), 1e-6) + score = max(0.0, 1.0 - rel_error) + direction = "supports" if score > 0.3 else "opposes" + elif name.startswith("rmsd_"): + if name == RMSD_STATE_MAP.get(predicted_state): + score = max(0.0, 3.0 - value) + direction = "supports" + else: + score = max(0.0, value - 2.0) * 0.3 + direction = "opposes" + elif name.startswith("domain_") or name.startswith("cavity_"): + score = abs(value) * 0.01 + direction = "supports" + else: + score = abs(value) * 0.001 + direction = "supports" + + results.append( + FeatureImportance( + feature_name=name, + display_name=display_name(name), + importance=score, + value=float(value), + direction=direction, + ) + ) + + return results + + +def _generic_heuristic_importance( + features: dict[str, float], predicted_state: str +) -> list[FeatureImportance]: + """Fallback heuristic when no family profile is available.""" + results: list[FeatureImportance] = [] + for name, value in features.items(): + if not isinstance(value, (int, float)) or np.isnan(value): + continue + score = abs(float(value)) * 0.01 + results.append( + FeatureImportance( + feature_name=name, + display_name=display_name(name), + importance=score, + value=float(value), + direction="supports", + ) + ) + return results + + +def resolve_importances( + features: dict[str, float], + predicted_state: str, + family: str = "LeuT", + importances: Optional[Sequence[FeatureImportance]] = None, + model: Optional[Any] = None, + feature_matrix: Optional[np.ndarray] = None, + feature_names: Optional[Sequence[str]] = None, + coefficients: Optional[dict[str, float]] = None, +) -> tuple[list[FeatureImportance], str]: + """ + Resolve feature importances using the best available method. + + Returns (importances, method_name). + """ + if importances is not None: + return list(importances), "provided" + + if coefficients is not None: + return ( + extract_coefficient_importance( + coefficients, features, predicted_state + ), + "coefficient", + ) + + if model is not None and feature_matrix is not None and feature_names: + baseline_proba = model.predict_proba(feature_matrix)[0] + return ( + extract_permutation_importance( + model, + feature_matrix, + feature_names, + baseline_proba, + ), + "permutation", + ) + + return ( + extract_heuristic_importance(features, predicted_state, family), + "heuristic", + ) diff --git a/confostate/explain/render.py b/confostate/explain/render.py new file mode 100644 index 0000000..035f7b7 --- /dev/null +++ b/confostate/explain/render.py @@ -0,0 +1,116 @@ +"""Template-based explanation rendering.""" + +from __future__ import annotations + +from typing import Sequence + +from confostate.explain._types import ( + Citation, + FeatureImportance, + PredictionResult, +) +from confostate.explain.citations import format_citation, state_label + + +def _format_value(value: float) -> str: + if abs(value) >= 100: + return f"{value:.1f}" + if abs(value) >= 10: + return f"{value:.2f}" + return f"{value:.3f}" + + +def _feature_sentence(item: FeatureImportance, predicted_state: str) -> str: + value_str = _format_value(item.value) + state_name = state_label(predicted_state) + verb = "supports" if item.direction == "supports" else "is atypical for" + return ( + f"The {item.display_name} ({value_str}) {verb} " + f"an {state_name} conformation." + ) + + +def render_text( + prediction: PredictionResult, + top_features: Sequence[FeatureImportance], + citations: Sequence[Citation], + method: str, +) -> str: + """Render a plain-text explanation.""" + confidence = prediction.probabilities.get(prediction.predicted_state, 0.0) + state_name = state_label(prediction.predicted_state) + lines = [ + ( + f"Structure {prediction.pdb_id} is predicted to be " + f"{state_name} ({confidence:.0%} confidence)." + ), + "", + "Key structural evidence:", + ] + + if top_features: + for item in top_features: + sentence = _feature_sentence(item, prediction.predicted_state) + lines.append(f"- {sentence}") + else: + lines.append("- No distinguishing features were identified.") + + lines.extend(["", f"Importance method: {method}."]) + + if citations: + lines.extend(["", "References:"]) + for citation in citations: + lines.append(f"- {format_citation(citation)}") + + return "\n".join(lines) + + +def render_markdown( + prediction: PredictionResult, + top_features: Sequence[FeatureImportance], + citations: Sequence[Citation], + method: str, +) -> str: + """Render a Markdown explanation.""" + confidence = prediction.probabilities.get(prediction.predicted_state, 0.0) + state_name = state_label(prediction.predicted_state) + lines = [ + f"## Prediction: {prediction.pdb_id}", + "", + ( + f"**State:** {state_name} " + f"({prediction.predicted_state}) \n" + f"**Confidence:** {confidence:.0%} \n" + f"**Family:** {prediction.family}" + ), + "", + "### Key structural evidence", + "", + ] + + if top_features: + for item in top_features: + lines.append( + f"- {_feature_sentence(item, prediction.predicted_state)}" + ) + else: + lines.append("- No distinguishing features were identified.") + + lines.extend(["", f"*Importance method: {method}*", ""]) + + if citations: + lines.extend(["### References", ""]) + for citation in citations: + ref = citation.reference + links: list[str] = [] + if citation.doi: + links.append(f"[DOI](https://doi.org/{citation.doi})") + if citation.pubmed_id: + links.append( + f"[PubMed](https://pubmed.ncbi.nlm.nih.gov/" + f"{citation.pubmed_id}/)" + ) + suffix = f" ({', '.join(links)})" if links else "" + lines.append(f"- {ref}{suffix}") + + return "\n".join(lines) diff --git a/docs/explainability.md b/docs/explainability.md new file mode 100644 index 0000000..eaca083 --- /dev/null +++ b/docs/explainability.md @@ -0,0 +1,140 @@ +# Explainability + +ConfoState generates human-readable explanations for conformational-state +predictions. The explainability layer sits downstream of feature extraction +(Person 2) and model inference (Person 3). + +## Package layout + +``` +confostate/explain/ +├── __init__.py # explain() public API +├── _types.py # PredictionResult, ExplanationResult, etc. +├── importance.py # Feature importance extraction +├── render.py # Template-based text/markdown rendering +└── citations.py # Literature references (DOI, PubMed) +``` + +## Quick start + +```python +from confostate.explain import explain, PredictionResult + +prediction = PredictionResult( + pdb_id="3F3E", + predicted_state="OF_open", + probabilities={"IF_open": 0.05, "OF_open": 0.80, "Occluded": 0.10, + "Intermediate": 0.05}, + features={ + "domain_TM1_TM7_distance": 22.5, + "rmsd_OF_open": 1.2, + "cavity_accessibility_out": 0.52, + }, +) + +result = explain(prediction) +print(result.text) +``` + +Plain-text output: + +``` +Structure 3F3E is predicted to be outward-facing open (80% confidence). + +Key structural evidence: +- The TM1–TM7 gate distance (22.50) supports an outward-facing open conformation. +- The RMSD to outward-open reference (3F3E) (1.200) supports an outward-facing open conformation. +... + +Importance method: heuristic. + +References: +- Singh et al., Science 2008 — DOI: 10.1126/science.1159299 — PubMed: 18403708 +``` + +Markdown output is available via `explain(prediction, output_format="markdown")`. + +## Importance methods + +The `explain()` function picks the best available importance source: + +| Priority | Method | When used | +|----------|--------|-----------| +| 1 | `provided` | Caller passes pre-computed importances (Person 3 pipeline) | +| 2 | `coefficient` | Linear-model coefficients are supplied | +| 3 | `permutation` | sklearn-compatible model + feature matrix | +| 4 | `heuristic` | Default stub: compares features to LeuT state profiles | + +### Heuristic mode (current default) + +While Person 3's models are in development, heuristic importance compares +each feature to literature-informed LeuT state profiles in +`confostate/explain/importance.py`. RMSD features matching the predicted +state receive higher scores; domain and cavity features are scored by +agreement with the profile. + +### Permutation importance + +When a trained model with `predict_proba` is available: + +```python +import numpy as np +from confostate.explain import explain + +result = explain( + prediction, + model=trained_model, + feature_matrix=X, # shape (n_samples, n_features) + feature_names=feature_cols, +) +``` + +### Provided importances + +Person 3 can pass SHAP or other importances directly: + +```python +from confostate.explain import FeatureImportance, explain + +importances = [ + FeatureImportance( + feature_name="domain_TM1_TM7_distance", + display_name="TM1–TM7 gate distance", + importance=0.85, + value=22.5, + direction="supports", + ), +] +result = explain(prediction, importances=importances) +``` + +## Citations + +`confostate/explain/citations.py` maps LeuT states and key structural +features to curated DOI and PubMed references. Citations are included +automatically in rendered output. + +Supported states: `IF_open`, `OF_open`, `Occluded`, `Intermediate`. + +## Integration with Person 5 (CLI) + +The planned `Classifier.classify()` API will call `explain()` internally. +Expected flow: + +``` +PDB → extract_features() → model.predict() → explain() → PredictionResult +``` + +## Testing + +```bash +pip install -e ".[dev]" +pytest tests/test_explain.py -v +``` + +## Future work + +- SHAP values for tree and neural models +- LLM polish pass (local Llama at ASU) +- Multi-family citation and profile tables +- HTML explanation renderer for CLI `--output-format html` diff --git a/docs/index.md b/docs/index.md index 5e74051..75d0bbc 100644 --- a/docs/index.md +++ b/docs/index.md @@ -13,6 +13,7 @@ scripts/download_structures scripts/extract_feature_vectors examples/LeuT_descriptors features/features.md +explainability api ``` diff --git a/tests/test_explain.py b/tests/test_explain.py new file mode 100644 index 0000000..c64abfc --- /dev/null +++ b/tests/test_explain.py @@ -0,0 +1,186 @@ +"""Tests for confostate.explain package.""" + +from __future__ import annotations + +import numpy as np +import pytest + +from confostate.explain import ( + ExplanationResult, + FeatureImportance, + PredictionResult, + display_name, + explain, + state_label, +) +from confostate.explain.citations import format_citation, get_citations +from confostate.explain.importance import ( + extract_coefficient_importance, + extract_heuristic_importance, + extract_permutation_importance, + rank_importances, +) + + +@pytest.fixture +def sample_prediction() -> PredictionResult: + return PredictionResult( + pdb_id="3F3E", + predicted_state="OF_open", + probabilities={ + "IF_open": 0.05, + "OF_open": 0.80, + "Occluded": 0.10, + "Intermediate": 0.05, + }, + features={ + "domain_TM1_TM7_distance": 22.5, + "domain_gate_TM1_TM6_distance": 26.1, + "rmsd_OF_open": 1.2, + "rmsd_IF_open": 3.8, + "cavity_accessibility_out": 0.52, + "cavity_volume": 480.0, + }, + family="LeuT", + ) + + +def test_display_name(): + assert display_name("domain_TM1_TM7_distance") == "TM1–TM7 gate distance" + assert display_name("unknown_feature") == "unknown feature" + + +def test_state_label(): + assert state_label("OF_open") == "outward-facing open" + assert state_label("Custom_state") == "Custom state" + + +def test_heuristic_importance(sample_prediction): + importances = extract_heuristic_importance( + sample_prediction.features, + sample_prediction.predicted_state, + family="LeuT", + ) + assert len(importances) == len(sample_prediction.features) + names = {item.feature_name for item in importances} + assert "rmsd_OF_open" in names + top = rank_importances(importances, top_k=3) + assert len(top) == 3 + assert all(item.importance >= 0 for item in top) + + +def test_coefficient_importance(sample_prediction): + coefficients = { + "domain_TM1_TM7_distance": 0.5, + "rmsd_OF_open": -1.0, + "rmsd_IF_open": 0.3, + } + importances = extract_coefficient_importance( + coefficients, + sample_prediction.features, + sample_prediction.predicted_state, + ) + rmsd_of = next(i for i in importances if i.feature_name == "rmsd_OF_open") + assert rmsd_of.direction == "opposes" + assert rmsd_of.importance < 0 + + +class DummyModel: + """Minimal sklearn-like model for permutation-importance tests.""" + + def predict_proba(self, X: np.ndarray) -> np.ndarray: + # Higher domain_TM1_TM7_distance -> more OF_open probability. + scores = X[:, 0] if X.shape[1] >= 1 else np.zeros(len(X)) + of_prob = np.clip(scores / 30.0, 0.0, 1.0) + if_prob = 1.0 - of_prob + return np.column_stack([if_prob, of_prob]) + + +def test_permutation_importance(): + feature_names = ["domain_TM1_TM7_distance", "rmsd_OF_open"] + X = np.array( + [ + [22.0, 1.2], + [18.0, 2.5], + [20.0, 1.8], + ] + ) + model = DummyModel() + baseline_proba = model.predict_proba(X)[0] + importances = extract_permutation_importance( + model, + X, + feature_names, + baseline_proba, + random_state=42, + ) + assert len(importances) == 2 + assert importances[0].feature_name in feature_names + + +def test_get_citations(sample_prediction): + citations = get_citations( + predicted_state="OF_open", + feature_names=["domain_TM1_TM7_distance", "rmsd_OF_open"], + family="LeuT", + ) + # State and features share Singh et al. 2008 — deduplicated. + assert len(citations) == 1 + assert citations[0].key == "OF_open" + formatted = format_citation(citations[0]) + assert "DOI" in formatted + assert "PubMed" in formatted + + +def test_get_citations_deduplicates_shared_doi(): + citations = get_citations( + predicted_state="Occluded", + feature_names=["rmsd_Occluded", "cavity_volume"], + family="LeuT", + ) + dois = [c.doi for c in citations if c.doi] + assert len(dois) == len(set(dois)) + assert len(citations) == 2 + + +def test_explain_returns_result(sample_prediction): + result = explain(sample_prediction) + assert isinstance(result, ExplanationResult) + assert result.pdb_id == "3F3E" + assert result.predicted_state == "OF_open" + assert result.confidence == pytest.approx(0.80) + assert result.method == "heuristic" + assert len(result.top_features) <= 5 + assert "3F3E" in result.text + assert "outward-facing open" in result.text + assert "## Prediction" in result.markdown + + +def test_explain_text_and_markdown_formats(sample_prediction): + text = explain(sample_prediction, output_format="text") + markdown = explain(sample_prediction, output_format="markdown") + assert isinstance(text, str) + assert isinstance(markdown, str) + assert "References:" in text + assert "### References" in markdown + + +def test_explain_from_dict(sample_prediction): + result = explain(sample_prediction.__dict__) + assert result.pdb_id == "3F3E" + + +def test_explain_with_provided_importances(sample_prediction): + provided = [ + FeatureImportance( + feature_name="domain_TM1_TM7_distance", + display_name="TM1–TM7 gate distance", + importance=0.9, + value=22.5, + direction="supports", + ) + ] + result = explain(sample_prediction, importances=provided, top_k=1) + assert result.method == "provided" + assert len(result.top_features) == 1 + assert result.top_features[0].feature_name == "domain_TM1_TM7_distance"