Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
58 changes: 58 additions & 0 deletions Plans/person4-interpretability-plan.md
Original file line number Diff line number Diff line change
@@ -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
141 changes: 141 additions & 0 deletions confostate/explain/__init__.py
Original file line number Diff line number Diff line change
@@ -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
66 changes: 66 additions & 0 deletions confostate/explain/_types.py
Original file line number Diff line number Diff line change
@@ -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"
Loading
Loading