Skip to content

API reference

The main public workflow consists of four functions:

  • compile_graph()
  • align_features_to_input_nodes()
  • customize_model()
  • interpret_model()

CompileArtifact is also exported because it is returned by compile_graph() and accepted by helper functions. Its stable user-facing fields are backend and feature_names. Additional fields such as graph, execution_plan, node_names_by_layer, input_nodes, output_nodes, hidden_nodes, and interpretation_sites expose compilation internals for inspection and debugging and may change across releases.

For target="nodes", interpret_model() returns one summary pandas.DataFrame by default. Use level="sites" to obtain per-site tables keyed by layer_1, layer_2, ... for feedforward models or step_1, step_2, ... for recurrent and graphnn models.

compile_graph

compile_graph(
    edgelist: DataFrame,
    backend: str = "feedforward",
    bias: bool = True,
    steps: int = 3,
    quiet: bool = False,
)

Compile an edgelist into a sparse PyTorch model and compilation artifact.

The edgelist defines the architecture graph. Each row describes one directed connection from source to target in the direction of computation. In other words, edges should point from input feature nodes toward hidden nodes and output nodes.

Input features are inferred as graph nodes with no incoming edges. Output nodes are inferred as graph nodes with no outgoing edges. The returned artifact stores the inferred input feature names in artifact.feature_names. Tensors passed to the compiled model must have columns in that exact order.

The edgelist may optionally include edge-level parameter metadata using the columns "initial_weight" and "constraint". These columns allow individual edges to define their initial effective weight and, where supported by the selected backend, constrain the trainable edge weight during optimization. Supported constraint values are "unconstrained", "positive", "negative", and "fixed". If omitted, edges use the backend's default trainable weight initialization and unconstrained weight behavior.

Graph-derived connectivity is enforced through masks on trainable edge weights. By default, compiled layers also include bias terms. Biases are node-level parameters, not graph edges, and are not constrained by the edge mask. Set bias=False to remove these offsets so node updates depend only on graph-defined weighted inputs.

The recurrent and graphnn backends apply a fixed number of graph state-update steps during each forward pass. The steps argument controls this update count. It is not a training epoch count and does not represent a sequence length in the input data.

Parameters:

Name Type Description Default
edgelist DataFrame

Edge table with required columns "source" and "target". Each row defines a directed connection from one named node to another, following the direction of computation. The table must include edges from input feature nodes into the rest of the architecture graph.

The table may also include optional columns "initial_weight" and "constraint". If provided, "initial_weight" defines the initial effective edge weight, and "constraint" defines how that edge weight is parameterized during training. Supported constraints are:

  • "unconstrained": the edge weight is trainable and may become positive or negative.
  • "positive": the edge weight is trainable and constrained to remain positive.
  • "negative": the edge weight is trainable and constrained to remain negative.
  • "fixed": the edge weight is fixed to "initial_weight" and is not trainable.

The optional columns are independent and row-wise sparse. Missing "initial_weight" values use the backend's default initialization for that edge. Missing "constraint" values are treated as "unconstrained" for that edge.

If both values are provided for an edge, positive-constrained edges must have positive initial weights and negative-constrained edges must have negative initial weights. Edges with constraint="fixed" must provide an "initial_weight" value in the same row, because fixed edges require an explicit constant value.

required
backend str

Backend to compile to. One of "feedforward", "recurrent", or "graphnn".

"feedforward"
bias bool

Whether compiled masked linear layers include bias terms. If True, each target node has a learned node-level offset in addition to its graph-defined weighted inputs. If False, node updates are computed only from graph-defined weighted inputs. Disabling bias gives the graph structure stricter control over node activations.

True
steps int

Number of recurrent/message-passing update steps for the recurrent and graphnn backends. Larger values allow information to propagate through longer graph paths and revisit cycles more times. This parameter is not used by the feedforward backend; non-default values with backend="feedforward" are rejected.

3
quiet bool

If False, emit informational notes during validation. If True, suppress informational notes.

False

Returns:

Type Description
tuple[Module, CompileArtifact]

Tuple (model, artifact). model is a PyTorch nn.Module compiled from the edgelist. artifact stores compilation metadata, including artifact.feature_names.

Raises:

Type Description
Edge2TorchError

If input validation, graph validation, or backend compilation fails.

Examples:

Compile a small feedforward architecture from an edgelist.

>>> import pandas as pd
>>> from edge2torch import compile_graph
>>>
>>> edgelist = pd.DataFrame(
...     {
...         "source": ["feature_a", "feature_b", "hidden_1"],
...         "target": ["hidden_1", "hidden_1", "prediction"],
...     }
... )
>>>
>>> model, artifact = compile_graph(
...     edgelist=edgelist,
...     backend="feedforward",
...     quiet=True,
... )
>>>
>>> artifact.feature_names
['feature_a', 'feature_b']

Compile a recurrent architecture with edge-level initial weights and constraints.

>>> edgelist = pd.DataFrame(
...     {
...         "source": ["feature_a", "feature_b", "hidden_1"],
...         "target": ["hidden_1", "hidden_1", "prediction"],
...         "initial_weight": [0.1, -0.2, 0.5],
...         "constraint": ["positive", "negative", "fixed"],
...     }
... )
>>>
>>> model, artifact = compile_graph(
...     edgelist=edgelist,
...     backend="recurrent",
...     quiet=True,
... )

Compilation metadata returned together with the compiled PyTorch model.

CompileArtifact is returned by compile_graph() and accepted by public helper functions such as align_features_to_input_nodes() and interpret_model(). It is exported for user-facing type hints and workflow integration.

The stable user-facing fields are backend and feature_names. Other fields expose compilation internals for inspection, testing, and debugging. They may change across releases and should not be treated as part of the stable public API.

Parameters:

Name Type Description Default
backend str

Backend used for compilation.

required
graph EdgeGraph

Internal edge2torch graph object used for compilation. The graph contains the normalized edge table and may include optional edge-level metadata such as initial_weight and constraint. This field is intended for inspection and debugging and may change across releases.

required
execution_plan object

Compiled execution plan used to build the model. This field exposes backend-specific internals and may change across releases.

required
node_names_by_layer dict[str, list[str]]

Mapping from layer name to node names in that layer. This field is populated for the feedforward backend and is primarily intended for inspection and feedforward-backend internals.

required
input_nodes list[str]

Names of graph input nodes inferred as nodes with no incoming edges.

required
output_nodes list[str]

Names of graph output nodes inferred as nodes with no outgoing edges.

required
hidden_nodes list[str]

Names of hidden graph nodes excluding inputs, outputs, and compiler pseudo nodes.

required
interpretation_sites dict[str, list[str]]

Mapping from interpretation site identifier to ordered node names for that site. Feedforward backends use layer_1, layer_2, and so on. Recurrent and graphnn backends use step_1, step_2, and so on.

required
feature_names list[str]

Names of the input features. This field defines the expected input column order for tensors passed to the compiled model.

required

align_features_to_input_nodes

align_features_to_input_nodes(
    data, artifact: CompileArtifact
) -> torch.Tensor

Align data features to the input-node order expected by a compiled model.

compile_graph() builds a sparse neural network from an edgelist. Input nodes are inferred from the graph structure and stored in artifact.feature_names. These names define the required column order for tensors passed to the compiled PyTorch model.

For named data containers, this function validates exact feature-name compatibility and reorders features by name:

  • pandas.DataFrame inputs are aligned using column names.
  • AnnData inputs are aligned using var_names if anndata is installed.

Named data containers must contain exactly the compiled model input-node features, although they may appear in any order. Missing or extra features raise an error.

torch.Tensor inputs do not contain feature names, so they are only validated by shape and are assumed to already follow artifact.feature_names order.

Parameters:

Name Type Description Default
data DataFrame | Tensor | AnnData

Input data to align. AnnData is supported when anndata is installed.

required
artifact CompileArtifact

Compilation artifact returned by compile_graph(). Its feature_names field defines the required input-node order.

required

Returns:

Type Description
Tensor

Float32 input tensor whose columns are ordered according to artifact.feature_names.

Raises:

Type Description
Edge2TorchError

If the input data type is unsupported, required features are missing, extra features are present in named data containers, non-numeric DataFrame columns are present, or tensor input has an incompatible shape.

Examples:

Align a DataFrame whose columns are named but not ordered like the compiled model input nodes.

>>> import pandas as pd
>>> import torch
>>> from edge2torch import align_features_to_input_nodes, compile_graph
>>>
>>> edgelist = pd.DataFrame(
...     {
...         "source": ["feature_a", "feature_b", "hidden"],
...         "target": ["hidden", "hidden", "prediction"],
...     }
... )
>>> model, artifact = compile_graph(edgelist, quiet=True)
>>>
>>> data = pd.DataFrame(
...     {
...         "feature_b": [2.0, 4.0],
...         "feature_a": [1.0, 3.0],
...     }
... )
>>>
>>> artifact.feature_names
['feature_a', 'feature_b']
>>>
>>> x = align_features_to_input_nodes(
...     data=data,
...     artifact=artifact,
... )
>>> x
tensor([[1., 2.],
        [3., 4.]])

Tensor inputs do not contain feature names, so they are only checked by shape and are assumed to already follow artifact.feature_names.

>>> x_tensor = torch.tensor(
...     [
...         [1.0, 2.0],
...         [3.0, 4.0],
...     ]
... )
>>> x_from_tensor = align_features_to_input_nodes(
...     data=x_tensor,
...     artifact=artifact,
... )
>>> torch.equal(x_from_tensor, x_tensor)
True

customize_model

customize_model(
    model: Module,
    activation: Module | None = None,
    dropout: float | int | None = None,
    head: Module | None = None,
) -> nn.Module

Wrap a compiled sparse neural network with optional PyTorch modules.

This function is a convenience layer for common post-compilation additions. It applies the requested components sequentially to the output of the compiled model. It does not modify the sparse graph structure, insert modules inside graph-derived layers, or replace ordinary PyTorch training and customization.

customize_model() wraps the provided model. Calling it repeatedly creates nested wrappers; it does not replace earlier customization modules. To change a customization, call customize_model() again on the original compiled model.

Parameters:

Name Type Description Default
model Module

PyTorch model returned by compile_graph().

required
activation Module | None

Optional PyTorch activation module applied after the compiled model. This should be an instantiated module such as nn.ReLU().

None
dropout float | int | None

Optional dropout probability applied after the activation. Must satisfy 0 <= dropout < 1.

None
head Module | None

Optional PyTorch module applied after dropout. This should be an instantiated module such as nn.Linear(...).

None

Returns:

Type Description
Module

Wrapped PyTorch model with the requested post-compilation modules.

Raises:

Type Description
Edge2TorchError

If any input is invalid.

Examples:

Add an activation function after the compiled sparse neural network.

>>> import pandas as pd
>>> from torch import nn
>>> from edge2torch import compile_graph, customize_model
>>>
>>> edgelist = pd.DataFrame(
...     {
...         "source": ["feature_a", "feature_b", "hidden"],
...         "target": ["hidden", "hidden", "prediction"],
...     }
... )
>>> model, artifact = compile_graph(edgelist, quiet=True)
>>>
>>> customized_model = customize_model(
...     model=model,
...     activation=nn.ReLU(),
... )

Add an activation, dropout, and task-specific prediction head.

>>> customized_model = customize_model(
...     model=model,
...     activation=nn.ReLU(),
...     dropout=0.2,
...     head=nn.Linear(1, 1),
... )

interpret_model

interpret_model(
    model: Any,
    artifact: Any,
    data: Any,
    target: str = "features",
    method: str = "IntegratedGradients",
    constructor_kwargs: dict[str, Any] | None = None,
    attribute_kwargs: dict[str, Any] | None = None,
    quiet: bool = False,
    level: str = "summary",
    nodes: str = "hidden",
    site_aggregation: str = "max_abs",
) -> Union[pd.DataFrame, dict[str, pd.DataFrame]]

Interpret a model compiled by edge2torch using a Captum attribution method.

Parameters:

Name Type Description Default
model Any

PyTorch model returned by compile_graph(), optionally customized and trained by the user.

required
artifact Any

Compilation artifact returned by compile_graph().

required
data DataFrame | AnnData | Tensor

Input data used for attribution.

required
target str

Interpretation target. Use "features" to attribute predictions to input features. Use "nodes" to attribute predictions to named graph nodes.

"features"
method str

Captum attribution method name. Method names follow Captum class names exactly and are case-sensitive, for example "IntegratedGradients", "Saliency", "DeepLift", "LayerConductance", or "LayerIntegratedGradients".

The selected method must be compatible with target and the compiled backend. If an unsupported method is provided, edge2torch raises an error listing the supported method names.

"IntegratedGradients"
constructor_kwargs dict[str, Any] | None

Optional keyword arguments passed directly to the constructor of the selected Captum attribution class. These arguments are passed through unchanged and are not interpreted, validated, or modified by edge2torch. Refer to the Captum documentation for the selected method to determine which constructor arguments are supported.

None
attribute_kwargs dict[str, Any] | None

Optional keyword arguments passed directly to the selected Captum method's attribute() call. These arguments are passed through unchanged and are not interpreted, validated, or modified by edge2torch. Refer to the Captum documentation for the selected method to determine which attribution arguments are supported.

None
quiet bool

If False, emit informational notes. If True, suppress informational notes.

False
level str

Detail level for target="nodes". Use "summary" to return one node-importance table per sample. Use "sites" to return one table per interpretation site such as layer_1 or step_2.

"summary"
nodes str

Node filter for target="nodes". Use "hidden" for internal graph nodes, "non_input" to include output nodes, or "all" for all visible graph nodes.

"hidden"
site_aggregation str

Aggregation rule used when target="nodes" and level="summary" for recurrent and graphnn backends. Ignored for feedforward summary results and for level="sites".

"max_abs"

Returns:

Type Description
DataFrame | dict[str, DataFrame]

If target="features", returns one DataFrame with rows as examples and columns as input feature names.

If target="nodes" and level="summary", returns one DataFrame with rows as examples and columns as named nodes selected by nodes.

If target="nodes" and level="sites", returns a dictionary mapping site identifiers to DataFrames. Each DataFrame has rows as examples and columns as named nodes for that site.

Notes

Feature interpretation is supported for all implemented backends.

Node interpretation is supported for the feedforward, recurrent, and graphnn backends. Node interpretation methods use Captum layer attribution classes.

For node-level interpretation, edge2torch must access the compiled model's internal interpretation sites. This works for raw models returned by compile_graph(), models returned by customize_model(), and manually wrapped PyTorch models if the compiled model remains registered as a PyTorch submodule. Highly custom wrappers that hide, replace, or bypass the compiled model may not support target="nodes".

interpret_model() temporarily switches the model to evaluation mode while computing attributions and restores the previous training/evaluation mode afterward.

constructor_kwargs and attribute_kwargs are passed through to Captum. Refer to the Captum documentation for method-specific arguments such as baselines, targets, additional forward arguments, or perturbation settings.

Raises:

Type Description
Edge2TorchError

If interpretation input validation fails, the requested target / method / backend combination is not supported, or Captum returns unsupported output.

Examples:

Compute feature-level attributions with integrated gradients.

>>> feature_attributions = interpret_model(
...     model=trained_model,
...     artifact=artifact,
...     data=data,
...     target="features",
...     method="IntegratedGradients",
...     quiet=True,
... )
>>> feature_attributions.head()

Compute summary node-level attributions.

>>> node_importance = interpret_model(
...     model=trained_model,
...     artifact=artifact,
...     data=data,
...     target="nodes",
...     method="LayerConductance",
...     quiet=True,
... )
>>> node_importance.head()

Compute per-site node-level attributions.

>>> node_attributions_by_site = interpret_model(
...     model=trained_model,
...     artifact=artifact,
...     data=data,
...     target="nodes",
...     level="sites",
...     nodes="non_input",
...     method="LayerConductance",
...     quiet=True,
... )
>>> node_attributions_by_site.keys()