# How the Doc2Graph E2E Model Architecture Works: End-to-End Entity and Relation Extraction

> Explore the Doc2Graph E2E model architecture, a unified graph neural network for simultaneous entity and relation extraction in a single forward pass. Understand its core implementation.

- Repository: [Andrea Gemelli/doc2graph](https://github.com/andreagemelli/doc2graph)
- Tags: architecture
- Published: 2026-02-24

---

**The Doc2Graph end-to-end (E2E) model architecture is a unified graph neural network that simultaneously predicts node classes (entity types) and edge classes (relations) through a single forward pass, implemented in the `E2E` class within [`doc2graph/models/graphs.py`](https://github.com/andreagemelli/doc2graph/blob/main/doc2graph/models/graphs.py).**

Doc2Graph transforms unstructured documents into structured knowledge graphs using a novel end-to-end approach that eliminates separate pipelines for entity detection and relation extraction. The **E2E model architecture** consolidates multimodal feature projection, message passing, and joint prediction into one differentiable graph neural network. This design enables simultaneous optimization of both node and edge classification tasks directly from document-level graph representations.

## Core Components of the E2E Architecture

The E2E model composes three specialized PyTorch modules to process document graphs built from visual, textual, and layout features.

### InputProjector

The `InputProjector` handles multimodal fusion by projecting heterogeneous node features into a unified hidden space. Located in [`doc2graph/models/graphs.py`](https://github.com/andreagemelli/doc2graph/blob/main/doc2graph/models/graphs.py), this component concatenates feature chunks produced by [`doc2graph/data/feature_builder.py`](https://github.com/andreagemelli/doc2graph/blob/main/doc2graph/data/feature_builder.py) and applies linear transformations to create consistent dimensional representations for subsequent graph operations.

### GcnSAGELayer

Message passing occurs through `GcnSAGELayer`, a Graph Convolutional Network implementation combining GraphSAGE mechanics with self-attention mechanisms. This layer propagates information across document nodes while optionally applying layer normalization, updating hidden states based on neighborhood aggregation.

### MLPPredictor_E2E

The `MLPPredictor_E2E` specializes in relation classification by consuming concatenated hidden representations of endpoint nodes, their predicted class probability distributions, and polar edge features. This architecture injects node classification context directly into edge prediction, improving relation extraction accuracy compared to edge-only features.

## The E2E Class Implementation

The `E2E` class orchestrates these components within [`doc2graph/models/graphs.py`](https://github.com/andreagemelli/doc2graph/blob/main/doc2graph/models/graphs.py). The constructor initializes the projection, message passing, and prediction modules:

```python
class E2E(nn.Module):
    def __init__(self,
                 node_classes,
                 edge_classes,
                 m_layers,
                 dropout,
                 in_chunks,
                 out_chunks,
                 hidden_dim,
                 device,
                 edge_pred_features,
                 doProject=True):
        super().__init__()

        # 1️⃣ Project multimodal node features

        self.projector = InputProjector(in_chunks, out_chunks, device, doProject)

        # 2️⃣ Message‑passing (single GCN layer – the former loop is commented out)

        m_hidden = self.projector.get_out_lenght()
        self.message_passing = GcnSAGELayer(m_hidden, m_hidden,
                                            F.relu, 0.0)

        # 3️⃣ Edge predictor specialised for the E2E pipeline

        self.edge_pred = MLPPredictor_E2E(
            m_hidden, hidden_dim, edge_classes, dropout,
            edge_pred_features)

        # 4️⃣ Node classifier (simple linear + LayerNorm)

        node_pred = [nn.Linear(m_hidden, node_classes),
                     nn.LayerNorm(node_classes)]
        self.node_pred = nn.Sequential(*node_pred)

```

## Forward Pass Execution

The `forward` method implements the end-to-end inference flow in four sequential stages:

```python
def forward(self, g, h):
    # a) Project raw node features → hidden space

    h = self.projector(h)

    # b) One round of GCN message passing

    h = self.message_passing(g, h)

    # c) Node logits

    n = self.node_pred(h)

    # d) Edge logits – the edge predictor receives the graph,

    #    hidden node states *and* the node class probabilities (soft‑max)

    e = self.edge_pred(g, h, n)

    return n, e

```

This unified flow ensures that node predictions inform edge classification through the softmax probability injection in step (d), creating a feedback mechanism between entity recognition and relation extraction.

## Key Architectural Characteristics

### Unified Multimodal Representation

Unlike pipelined approaches, the **E2E architecture** projects all modalities—visual features, textual embeddings, and spatial coordinates—once through `InputProjector`, avoiding separate transformation pipelines for nodes and edges.

### Single-Layer Message Passing

The current implementation utilizes a single `GcnSAGELayer` for graph convolution. The source code retains commented multi-layer loop structures for future extensibility, but the active configuration prioritizes computational efficiency over deep hierarchical aggregation.

### Context-Aware Edge Prediction

The `MLPPredictor_E2E` uniquely concatenates three signal sources: hidden vectors of connected nodes (`h_u`, `h_v`), their predicted class distributions (`cls_u`, `cls_v`), and polar edge features. This tri-modal input provides richer relational context than traditional edge predictors.

### Joint Optimization

The architecture enables **end-to-end training** by producing both node and edge logits simultaneously, allowing combined loss functions to optimize entity detection and relation extraction jointly rather than sequentially.

## Practical Implementation Examples

### Factory-Based Model Creation

Instantiate the E2E architecture through the `SetModel` factory, which resolves configuration parameters from training checkpoints:

```python
from doc2graph.models.graphs import SetModel

# Suppose the document yields 100 node classes and 2 edge classes

chunks = [...]            # list of feature chunk sizes produced by FeatureBuilder

model_factory = SetModel(name="E2E", device="cpu")
model = model_factory.get_model(
    nodes=100,                 # node class count

    edges=2,                   # edge class count

    chunks=chunks,
    verbatim=False
)

```

### Inference on Document Graphs

Execute predictions on DGL graphs constructed by [`doc2graph/data/graph_builder.py`](https://github.com/andreagemelli/doc2graph/blob/main/doc2graph/data/graph_builder.py):

```python
import dgl
import torch

# g = DGL graph built from a document; node features stored in g.ndata["feat"]

node_logits, edge_logits = model(g, g.ndata["feat"])

# Convert logits to predictions

node_preds = torch.argmax(torch.softmax(node_logits, dim=1), dim=1)
edge_preds = torch.argmax(torch.softmax(edge_logits, dim=1), dim=1)

```

### Checkpoint Loading and Evaluation

The inference pipeline in [`doc2graph/inference.py`](https://github.com/andreagemelli/doc2graph/blob/main/doc2graph/inference.py) demonstrates production deployment:

```python

# doc2graph/inference.py

sm = SetModel(name=model, device=device)      # model name resolved from checkpoint prefix

model = sm.get_model(info["node_num_classes"],
                     info["edge_num_classes"],
                     chunks, False)
model.load_state_dict(torch.load(CHECKPOINTS / weights[0],
                                 map_location=device))
model.eval()

# later:

n, e = model(graph.to(device), graph.ndata["feat"].to(device))

```

## Summary

- The **E2E model** unifies entity and relation prediction within a single graph neural network architecture defined in [`doc2graph/models/graphs.py`](https://github.com/andreagemelli/doc2graph/blob/main/doc2graph/models/graphs.py).
- Three core components—`InputProjector`, `GcnSAGELayer`, and `MLPPredictor_E2E`—handle feature projection, message passing, and joint prediction respectively.
- The architecture injects node classification probabilities into edge prediction, creating context-aware relation extraction.
- **End-to-end training** optimizes both node and edge losses simultaneously, eliminating the need for separate entity detection and relation classification pipelines.
- Production inference utilizes the `SetModel` factory and standard DGL graph operations to process documents into knowledge graphs.

## Frequently Asked Questions

### What is the primary advantage of the Doc2Graph E2E architecture?

The primary advantage is **joint optimization** of entity detection and relation extraction within a single differentiable graph neural network. Unlike pipelined systems that propagate errors from entity recognition to relation classification, the E2E model in [`doc2graph/models/graphs.py`](https://github.com/andreagemelli/doc2graph/blob/main/doc2graph/models/graphs.py) trains both tasks simultaneously using shared hidden representations, improving overall accuracy and reducing inference complexity.

### How does the edge predictor utilize node classification information?

The `MLPPredictor_E2E` consumes the softmax probability distributions from node classification (`n`) alongside hidden node features and polar edge attributes. By concatenating these three inputs—source node hidden state, target node hidden state, their class probabilities, and edge features—the predictor gains semantic context about entity types when determining relations, as implemented in the edge prediction logic of the `forward` method.

### Can the E2E model support deeper graph convolutional networks?

While the current `E2E` implementation in [`doc2graph/models/graphs.py`](https://github.com/andreagemelli/doc2graph/blob/main/doc2graph/models/graphs.py) uses a single `GcnSAGELayer`, the source code contains commented multi-layer loop structures indicating planned extensibility. Users can modify the `__init__` method to wrap `GcnSAGELayer` in a sequential container or loop structure, though the default configuration prioritizes computational efficiency for document-level graphs.

### Which files are essential for understanding the complete E2E pipeline?

Four critical files define the end-to-end workflow: [`doc2graph/models/graphs.py`](https://github.com/andreagemelli/doc2graph/blob/main/doc2graph/models/graphs.py) contains the `E2E` class and core neural components; [`doc2graph/data/feature_builder.py`](https://github.com/andreagemelli/doc2graph/blob/main/doc2graph/data/feature_builder.py) generates multimodal inputs for `InputProjector`; [`doc2graph/data/graph_builder.py`](https://github.com/andreagemelli/doc2graph/blob/main/doc2graph/data/graph_builder.py) constructs DGL graphs from documents; and [`doc2graph/inference.py`](https://github.com/andreagemelli/doc2graph/blob/main/doc2graph/inference.py) demonstrates checkpoint loading and prediction execution.