# How InputProjector Handles Multi-Modal Feature Fusion in Doc2Graph

> Learn how InputProjector achieves multi-modal feature fusion by projecting and concatenating diverse document slices into a unified vector for GNN processing.

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

---

**InputProjector fuses heterogeneous document features by applying independent three-layer neural projections to each modality slice and concatenating the outputs into a unified vector before GNN processing.**

Doc2Graph processes documents containing heterogeneous signals—text embeddings, spatial coordinates, and visual features—that must align into a common representation before graph neural network layers can process them. The `InputProjector` class in [`doc2graph/models/graphs.py`](https://github.com/andreagemelli/doc2graph/blob/main/doc2graph/models/graphs.py) serves as the dedicated fusion module that transforms concatenated multi-modal inputs into a unified hidden space. This architecture ensures each modality receives specialized linear transformations while maintaining an efficient concatenation-based aggregation strategy.

## Architecture of InputProjector Multi-Modal Fusion

### Per-Modality Projection Networks

During initialization in `InputProjector.__init__`, the module constructs independent sub-networks for each input modality defined by the `in_chunks` parameter. For every chunk size in the list, the code instantiates a three-layer stack stored in `self.modalities` ([lines 40-46](https://github.com/andreagemelli/doc2graph/blob/master/doc2graph/models/graphs.py#L40-L46)):

- `nn.Linear(chunk, out_chunks)` projects the modality-specific features to the target dimension
- `nn.LayerNorm(out_chunks)` applies feature-wise normalization for training stability
- `nn.ReLU()` introduces non-linearity between the projection and downstream layers

All sub-networks reside within an `nn.Sequential` container, enabling modular processing of each data type.

### Slice Index Management

To efficiently extract modalities from the flattened input vector, the constructor prepends a zero to the `in_chunks` list stored as `self.chunks` ([lines 48-49](https://github.com/andreagemelli/doc2graph/blob/master/doc2graph/models/graphs.py#L48-L49)). This modification enables simple arithmetic for calculating start and end indices when slicing the concatenated tensor during the forward pass.

### Forward Pass and Concatenation

The `forward` method implements the core multi-modal fusion logic ([lines 59-69](https://github.com/andreagemelli/doc2graph/blob/master/doc2graph/models/graphs.py#L59-L69)). For each modality index, the method extracts the corresponding slice from input tensor `h` using the pre-computed cumulative indices, passes the slice through its dedicated projector module, and accumulates the results. The final operation `torch.cat(mid, dim=1)` concatenates all projected modality vectors along the feature dimension, producing an output tensor of shape `(batch_size, len(in_chunks) * out_chunks)`.

## Integration with Doc2Graph Classification Models

### NodeClassifier and EdgeClassifier

Both `NodeClassifier` and `EdgeClassifier` instantiate `InputProjector` during their initialization phases ([lines 21-24](https://github.com/andreagemelli/doc2graph/blob/master/doc2graph/models/graphs.py#L21-L24) and [lines 81-84](https://github.com/andreagemelli/doc2graph/blob/master/doc2graph/models/graphs.py#L81-L84)). The fused features feed directly into subsequent GCN layers for node classification and edge prediction tasks. The `E2E` (end-to-end) model also reuses this projector ([lines 25-27](https://github.com/andreagemelli/doc2graph/blob/master/doc2graph/models/graphs.py#L25-L27)), demonstrating its plug-and-play design across different graph-based architectures in the repository.

## Configuration and Bypass Mode

The projector supports an optional bypass mechanism controlled by the `doIt` parameter and the `doProject` configuration flag ([lines 29-33](https://github.com/andreagemelli/doc2graph/blob/master/doc2graph/models/graphs.py#L29-L33)). When `doIt=False`, the `forward` method returns the original input tensor unchanged, enabling ablation studies that measure the specific contribution of multi-modal feature fusion to model performance.

## Practical Implementation Examples

### Fusing Two Modalities

```python
import torch
from doc2graph.models.graphs import InputProjector

# Define modalities: text embeddings (12-dim) and layout features (20-dim)

in_chunks = [12, 20]
out_chunks = 64
device = "cpu"

proj = InputProjector(in_chunks, out_chunks, device)

# Batch of 5 samples with concatenated features (32-dim input)

x = torch.randn(5, sum(in_chunks))

# Fuse modalities: output contains 2 * 64 = 128 features

fused = proj(x)
print(fused.shape)  # torch.Size([5, 128])

```

### Integration in NodeClassifier

```python
from doc2graph.models.graphs import NodeClassifier
import torch

# Three modalities (text, layout, visual) each with 64 dimensions

in_chunks = [64, 64, 64]
out_chunks = 128

classifier = NodeClassifier(
    in_chunks=in_chunks,
    out_chunks=out_chunks,
    n_classes=10,
    n_layers=3,
    activation=torch.nn.ReLU(),
    dropout=0.2,
    device="cuda:0"
)

# Forward pass: g is the DGL graph, h is the concatenated node features

# logits = classifier(g, h)

```

### Disabling Projection for Ablation

```python

# Bypass the projector to use raw concatenated features

proj_bypass = InputProjector(in_chunks, out_chunks, device, doIt=False)
unchanged = proj_bypass(x)  # Returns x with original shape intact

```

## Summary

- **InputProjector** resides in [`doc2graph/models/graphs.py`](https://github.com/andreagemelli/doc2graph/blob/main/doc2graph/models/graphs.py) and specializes in multi-modal feature fusion before GNN processing.
- Each modality defined in `in_chunks` receives an independent three-layer projection (Linear → LayerNorm → ReLU) mapping to a uniform `out_chunks` dimension.
- Fusion occurs through concatenation of projected modality vectors along the feature dimension, creating a unified representation for downstream graph layers.
- The module integrates seamlessly with `NodeClassifier`, `EdgeClassifier`, and `E2E` models throughout the Doc2Graph architecture.
- A bypass mode (`doIt=False`) enables controlled experiments without feature projection by returning inputs unchanged.

## Frequently Asked Questions

### What input format does InputProjector expect?

The module expects a PyTorch tensor where multi-modal features are concatenated along the last dimension in the order specified by `in_chunks`. For a configuration like `in_chunks=[12, 20]`, the input tensor must have shape `(batch_size, 32)`, where the first 12 features correspond to modality A and the subsequent 20 features correspond to modality B.

### Can I use different output dimensions for each modality?

No, the current implementation requires a single `out_chunks` parameter that applies uniformly to all modalities. Each projector outputs the same dimensionality, and the final fused representation size equals `len(in_chunks) * out_chunks`. This design choice simplifies the concatenation logic but requires all modalities to share the same target projection dimension.

### How does the bypass mode affect model performance?

When initialized with `doIt=False` or when the configuration flag `doProject` disables the projector, the module returns raw concatenated features without transformation. This typically reduces model capacity and may decrease accuracy on heterogeneous documents, but it provides a critical baseline for quantifying the impact of learned multi-modal fusion versus simple concatenation.

### Where is the fusion logic implemented in the source code?

The core fusion logic resides in the `forward` method of the `InputProjector` class at [lines 59-69](https://github.com/andreagemelli/doc2graph/blob/master/doc2graph/models/graphs.py#L59-L69) of [`doc2graph/models/graphs.py`](https://github.com/andreagemelli/doc2graph/blob/main/doc2graph/models/graphs.py). Specifically, the slice extraction loop iterates over `self.chunks` to separate modalities, and the `torch.cat` operation on line 68 combines the projected vectors into the final fused representation.