How to Run Inference on Custom Documents Using Pretrained Doc2Graph Weights
You can run inference on custom documents using pretrained Doc2Graph weights by calling the inference() function in doc2graph/inference.py with a list of checkpoint files and document paths, or by using the CLI with the --inference, --weights, and --docs flags.
Doc2Graph converts documents (images or PDFs) into graph structures to extract key-value pairs using Graph Neural Networks (GNNs). The inference pipeline processes raw documents through graph construction, feature enrichment, and model prediction to output annotated visualizations and structured JSON files.
The Doc2Graph Inference Pipeline
The end-to-end workflow consists of five sequential steps implemented across the repository:
-
Document Parsing and Graph Construction. The
GraphBuilderclass indoc2graph/data/graph_builder.pyparses input files and instantiates adgl.DGLGraph. For images, it runs OCR; for datasets like FUNSD, it parses XML. Every detected bounding box becomes a node in the graph. -
Feature Enrichment. The
FeatureBuilderindoc2graph/data/feature_builder.pyadds geometric coordinates, visual embeddings, spaCy text embeddings, and edge polar features. These are stored ing.ndata["feat"]for nodes andg.edata["feat"]for edges. -
Model Instantiation. The
SetModelfactory indoc2graph/models/graphs.pyreads the configuration and creates the appropriate architecture (GCN, EDGE, or E2E). It loads the pretrained checkpoint specified in your arguments and moves the model to the target device. -
Forward Pass and Prediction. The model returns node logits (
n) and edge logits (e). The highest-probability edge class (epreds) determines the predicted key-value relationships. This logic resides indoc2graph/inference.py. -
Visualization and Export. The pipeline renders the original image with green circles for keys, red circles for values, and violet connecting lines. It also exports the extracted pairs as JSON files to the
inference/directory.
Running Inference via the Command Line
The CLI entry point in doc2graph/main.py provides the fastest way to run inference on custom documents using pretrained weights.
python -m doc2graph.main \
--inference \
--weights funsd-e2e.pt \
--docs path/to/doc1.png path/to/doc2.pdf \
--gpu 0
- The
--weightsargument must point to a file located underdoc2graph/models/checkpoints/. - The
--docsargument accepts a space-separated list of image or PDF paths. - Omit
--gpuor set it to-1to force CPU inference. - Results are written to
inference/<doc_name>.pngandinference/<doc_name>.json.
Running Inference via the Python API
For programmatic integration, import the inference function directly from doc2graph/inference.py.
from doc2graph.inference import inference
from doc2graph.utils import get_device
# Define inputs
weights = ["funsd-e2e.pt"] # Supports multiple checkpoints for experiments
docs = ["samples/invoice1.png", "samples/form2.pdf"]
device = get_device(0) # Use get_device(-1) for CPU
# Execute pipeline
inference(weights, docs, device)
This function automatically invokes GraphBuilder to create the graph, FeatureBuilder to compute features, and SetModel to load the E2E architecture (default for funsd-e2e.pt). All outputs are saved to the inference/ folder defined in doc2graph/paths.py.
Customizing Feature Extraction
You can control which feature channels are active by modifying the arguments namespace before calling inference. This directly affects the feature vectors stored in g.ndata["feat"].
# Configure features programmatically
args = type('Args', (object,), {
'add_geom': True,
'add_embs': False, # Disable text embeddings
'add_visual': True,
'add_hist': False,
'add_eweights': True,
'add_fudge': False,
'num_polar_bins': 8,
})()
# Pass args to the main routine or builders
Disabling unnecessary features reduces computation time when running inference on custom documents that lack certain modalities (e.g., documents where visual features are sufficient).
Key Source Files for Inference
Understanding these files helps debug and extend the inference pipeline:
-
doc2graph/inference.py: Contains the top-levelinference()function that orchestrates graph creation, feature addition, model loading, and result saving. -
doc2graph/data/graph_builder.py: ImplementsGraphBuilderand itsget_graphmethod to parse documents and build the initial DGL graph structure. -
doc2graph/data/feature_builder.py: ImplementsFeatureBuilderand theadd_featuresloop to enrich graph nodes and edges with geometric, textual, and visual data. -
doc2graph/models/graphs.py: Defines theSetModelfactory (get_model) and the GNN architectures (GCN, EDGE, E2E) used during inference. -
doc2graph/paths.py: Centralizes path constants includingCHECKPOINTSandINFERENCEdirectories used throughout the pipeline.
Summary
- Use the CLI flags
--inference --weights <file> --docs <paths>for quick testing from the terminal. - Call
inference(weights, docs, device)in Python to integrate Doc2Graph into larger applications. - The pipeline relies on
GraphBuilder(doc2graph/data/graph_builder.py) for document parsing andFeatureBuilder(doc2graph/data/feature_builder.py) for feature computation. SetModelindoc2graph/models/graphs.pyhandles architecture selection and checkpoint loading.- All inference outputs are saved to the
inference/directory as annotated PNG files and structured JSON files.
Frequently Asked Questions
Where should I place pretrained checkpoint files?
Place pretrained checkpoints (e.g., funsd-e2e.pt) inside the doc2graph/models/checkpoints/ directory. The SetModel factory in doc2graph/models/graphs.py resolves paths relative to this location when loading weights for inference.
What document formats are supported for custom inference?
Doc2Graph supports image files (PNG, JPG) and PDF documents. The GraphBuilder in doc2graph/data/graph_builder.py automatically applies OCR to images and extracts bounding boxes from PDFs or XML annotations when available.
How do I switch between CPU and GPU inference?
Pass --gpu -1 to the CLI or call get_device(-1) in Python to use CPU. Specify a GPU device index (e.g., 0) to use CUDA. The inference() function moves the model and data to the specified device before running the forward pass.
Can I disable text embeddings to speed up inference?
Yes. Set add_embs=False in the arguments namespace before calling inference(). This prevents FeatureBuilder from computing spaCy text embeddings, reducing the dimensionality of g.ndata["feat"] and decreasing processing time for documents where textual features are unnecessary.
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 →