How the Doc2Graph E2E Model Architecture Works: End-to-End Entity and Relation Extraction
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.
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, this component concatenates feature chunks produced by 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. The constructor initializes the projection, message passing, and prediction modules:
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:
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:
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:
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 demonstrates production deployment:
# 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. - Three core components—
InputProjector,GcnSAGELayer, andMLPPredictor_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
SetModelfactory 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 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 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 contains the E2E class and core neural components; doc2graph/data/feature_builder.py generates multimodal inputs for InputProjector; doc2graph/data/graph_builder.py constructs DGL graphs from documents; and doc2graph/inference.py demonstrates checkpoint loading and prediction execution.
Have a question about this repo?
These articles cover the highlights, but your codebase questions are specific. Give your agent direct access to the source. Share this with your agent to get started:
curl -s "https://instagit.com/install.md" Maintain an open-source project? Get it listed too →