# How to Run Inference on Custom Documents Using Pretrained Doc2Graph Weights

> Run inference on custom documents with pretrained Doc2Graph weights. Follow clear instructions for using the inference function or CLI for efficient document analysis.

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

---

**You can run inference on custom documents using pretrained Doc2Graph weights by calling the `inference()` function in [`doc2graph/inference.py`](https://github.com/andreagemelli/doc2graph/blob/main/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:

1. **Document Parsing and Graph Construction**. The `GraphBuilder` class in [`doc2graph/data/graph_builder.py`](https://github.com/andreagemelli/doc2graph/blob/main/doc2graph/data/graph_builder.py) parses input files and instantiates a `dgl.DGLGraph`. For images, it runs OCR; for datasets like FUNSD, it parses XML. Every detected bounding box becomes a node in the graph.

2. **Feature Enrichment**. The `FeatureBuilder` in [`doc2graph/data/feature_builder.py`](https://github.com/andreagemelli/doc2graph/blob/main/doc2graph/data/feature_builder.py) adds geometric coordinates, visual embeddings, spaCy text embeddings, and edge polar features. These are stored in `g.ndata["feat"]` for nodes and `g.edata["feat"]` for edges.

3. **Model Instantiation**. The `SetModel` factory in [`doc2graph/models/graphs.py`](https://github.com/andreagemelli/doc2graph/blob/main/doc2graph/models/graphs.py) reads 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.

4. **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 in [`doc2graph/inference.py`](https://github.com/andreagemelli/doc2graph/blob/main/doc2graph/inference.py).

5. **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`](https://github.com/andreagemelli/doc2graph/blob/main/doc2graph/main.py) provides the fastest way to run inference on custom documents using pretrained weights.

```bash
python -m doc2graph.main \
    --inference \
    --weights funsd-e2e.pt \
    --docs path/to/doc1.png path/to/doc2.pdf \
    --gpu 0

```

- The `--weights` argument must point to a file located under `doc2graph/models/checkpoints/`.
- The `--docs` argument accepts a space-separated list of image or PDF paths.
- Omit `--gpu` or set it to `-1` to force CPU inference.
- Results are written to `inference/<doc_name>.png` and `inference/<doc_name>.json`.

## Running Inference via the Python API

For programmatic integration, import the `inference` function directly from [`doc2graph/inference.py`](https://github.com/andreagemelli/doc2graph/blob/main/doc2graph/inference.py).

```python
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`](https://github.com/andreagemelli/doc2graph/blob/main/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"]`.

```python

# 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`](https://github.com/andreagemelli/doc2graph/blob/main/doc2graph/inference.py)**: Contains the top-level `inference()` function that orchestrates graph creation, feature addition, model loading, and result saving.

- **[`doc2graph/data/graph_builder.py`](https://github.com/andreagemelli/doc2graph/blob/main/doc2graph/data/graph_builder.py)**: Implements `GraphBuilder` and its `get_graph` method to parse documents and build the initial DGL graph structure.

- **[`doc2graph/data/feature_builder.py`](https://github.com/andreagemelli/doc2graph/blob/main/doc2graph/data/feature_builder.py)**: Implements `FeatureBuilder` and the `add_features` loop to enrich graph nodes and edges with geometric, textual, and visual data.

- **[`doc2graph/models/graphs.py`](https://github.com/andreagemelli/doc2graph/blob/main/doc2graph/models/graphs.py)**: Defines the `SetModel` factory (`get_model`) and the GNN architectures (GCN, EDGE, E2E) used during inference.

- **[`doc2graph/paths.py`](https://github.com/andreagemelli/doc2graph/blob/main/doc2graph/paths.py)**: Centralizes path constants including `CHECKPOINTS` and `INFERENCE` directories 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`](https://github.com/andreagemelli/doc2graph/blob/main/doc2graph/data/graph_builder.py)) for document parsing and `FeatureBuilder` ([`doc2graph/data/feature_builder.py`](https://github.com/andreagemelli/doc2graph/blob/main/doc2graph/data/feature_builder.py)) for feature computation.
- `SetModel` in [`doc2graph/models/graphs.py`](https://github.com/andreagemelli/doc2graph/blob/main/doc2graph/models/graphs.py) handles 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`](https://github.com/andreagemelli/doc2graph/blob/main/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`](https://github.com/andreagemelli/doc2graph/blob/main/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.