# How Doc2Graph Handles Class Imbalance with the balance_edges Method

> Doc2Graph uses the balance_edges method to tackle class imbalance by randomly downsampling dominant-class edges, promoting balanced training data.

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

---

**Doc2Graph handles class imbalance by pruning dominant-class edges through random downsampling in the `balance_edges` method, ensuring balanced edge distributions for training.**

Doc2Graph is an open-source framework for graph-based document understanding that converts documents into graph structures using DGL (Deep Graph Library). When processing document datasets like FUNSD, edge-class imbalance frequently occurs when certain relationship types—such as "table" edges—dominate the graph connectivity. The `balance_edges` method, implemented in [`doc2graph/data/graph_builder.py`](https://github.com/andreagemelli/doc2graph/blob/main/doc2graph/data/graph_builder.py) (lines 57-94), solves this by intelligently downsampling over-represented edge classes to create balanced training data.

## The balance_edges Algorithm Architecture

### Input Graph and Target Class Identification

The `balance_edges` method operates on a `dgl.DGLGraph` object where edge labels are stored in `g.edata["label"]`. The method signature requires an integer `cls` parameter that specifies which edge class to downsample:

```python
def balance_edges(self, g: dgl.DGLGraph, cls: int) -> dgl.DGLGraph:

```

This design allows precise targeting of specific dominant classes (for example, class `2` representing "table" relationships in FUNSD datasets) without affecting other edge types.

### Boolean Masking and Index Extraction

The implementation first isolates the target class edges using tensor operations. According to the source code in [`doc2graph/data/graph_builder.py`](https://github.com/andreagemelli/doc2graph/blob/main/doc2graph/data/graph_builder.py), the method creates a boolean mask and extracts indices for all edges matching the target class:

```python
edge_targets = g.edata["label"]
to_remove = edge_targets == cls                     # boolean mask of the target class

indices_to_remove = to_remove.nonzero().flatten().tolist()

```

At this stage, `indices_to_remove` contains the positions of all edges belonging to the over-represented class, which initially marks them for deletion.

### Random Retention Strategy

Rather than indiscriminately deleting edges, the method implements a stochastic retention strategy that preserves approximately half the number of non-target edges. This ensures the dominant class count matches the sum of all minority classes combined:

```python
for _ in range(int((edge_targets != cls).sum() / 2)):
    indeces_to_save = [random.choice(indices_to_remove)]
    for index in sorted(indeces_to_save, reverse=True):
        del indices_to_remove[indices_to_remove.index(index)]

```

The loop iterates `(edge_targets != cls).sum() / 2` times, calculating half the count of non-target edges. In each iteration, it randomly selects one target-class edge to preserve and removes it from the deletion list. This random selection prevents systematic bias while achieving numerical balance between the previously dominant class and all other edge types.

### Edge Removal via DGL

Finally, the method converts the remaining indices (those still marked for removal) into a tensor and applies DGL's native edge removal function:

```python
indices_to_remove = torch.flatten(
    torch.tensor(indices_to_remove, dtype=torch.int32)
)
g = dgl.remove_edges(g, indices_to_remove)

```

The returned graph contains a balanced distribution where the target class edge count approximately equals the total of all other classes, mitigating model bias toward majority-class predictions.

## DataLoader Integration for Batch Processing

While `balance_edges` operates on single graphs, the framework provides a high-level convenience wrapper in [`doc2graph/data/dataloader.py`](https://github.com/andreagemelli/doc2graph/blob/main/doc2graph/data/dataloader.py) (lines 47-64). The `DataLoader.balance` method handles class name resolution and iterates over the entire dataset:

```python
self.graphs[i] = self.GB.balance_edges(g, self.edge_num_classes, cls=cls)

```

This wrapper first maps human-readable class names (like `"table"`) to their corresponding integer IDs, then applies the balancing operation to every graph in the collection. This abstraction eliminates manual index management when preprocessing training batches.

## Practical Implementation Examples

### Balancing a Single Graph Manually

For fine-grained control over individual graphs, instantiate `GraphBuilder` and call `balance_edges` directly:

```python
from doc2graph.data.graph_builder import GraphBuilder
import dgl

# Build a graph from FUNSD data

gb = GraphBuilder()
graphs, _, _, _ = gb.__fromFUNSD("/path/to/funsd")
g = graphs[0]                     # select first document graph

# Downsample dominant class 2 ("table" relationships)

balanced_g = gb.balance_edges(g, cls=2)

print(f"Original edges: {g.num_edges()}")
print(f"Balanced edges: {balanced_g.num_edges()}")

```

This approach is ideal for debugging specific documents or integrating balancing into custom preprocessing pipelines.

### Batch Balancing with DataLoader

For standard training workflows, use the `DataLoader` interface to balance entire datasets automatically:

```python
from doc2graph.data.dataloader import DataLoader

# Initialize loader with FUNSD dataset

dl = DataLoader(
    name="funsd",
    src_path="/path/to/funsd",
    src_data="FUNSD",
    edge_type="fully",            # or "knn"

)

# Balance all graphs by downsampling "table" edges

dl.balance(cls="table")          # automatic string-to-int mapping

# Verify distribution on first graph

print(dl.graphs[0].edata["label"].bincount())

```

The `DataLoader.balance` method manages the iteration logic and class ID resolution, making it the preferred approach for dataset-wide preprocessing.

## Summary

- **Pruning Strategy**: The `balance_edges` method in [`doc2graph/data/graph_builder.py`](https://github.com/andreagemelli/doc2graph/blob/main/doc2graph/data/graph_builder.py) randomly downsamples dominant edge classes to match the combined count of minority classes.
- **Mathematical Balance**: The retention formula `(edge_targets != cls).sum() / 2` ensures the target class ends with approximately the same number of edges as all other classes combined.
- **Stochastic Selection**: Random selection of edges to preserve prevents systematic bias in the resulting graph structure.
- **Dual API Access**: Direct access via `GraphBuilder.balance_edges` for single graphs, or batch processing through `DataLoader.balance` in [`doc2graph/data/dataloader.py`](https://github.com/andreagemelli/doc2graph/blob/main/doc2graph/data/dataloader.py).
- **DGL Integration**: Uses native `dgl.remove_edges` for efficient tensor operations without breaking graph connectivity metadata.

## Frequently Asked Questions

### What causes edge-class imbalance in document graphs?

Document graphs often exhibit imbalance when structural relationships—such as table cells, headers, or list items—create significantly more edges than semantic relationships between text entities. In FUNSD datasets, "table" edges frequently outnumber "link" or "question-answer" edges by orders of magnitude, causing models to predict the majority class disproportionately during training.

### How does the balance_edges method determine retention counts?

The method calculates retention using the formula `(edge_targets != cls).sum() / 2`, which computes half the total count of non-target edges. By retaining exactly this many target-class edges, the final distribution achieves approximate parity between the previously dominant class and the sum of all other classes, preventing any single class from overwhelming the loss function.

### Can I balance multiple edge classes simultaneously with one call?

The current implementation in [`doc2graph/data/graph_builder.py`](https://github.com/andreagemelli/doc2graph/blob/main/doc2graph/data/graph_builder.py) processes one target class per invocation. To balance multiple classes, you must call `balance_edges` sequentially for each dominant class ID. Note that order matters—balancing one class affects the counts of others, so iterative application requires careful monitoring of class distributions between calls.

### Where is the balance_edges source code located?

The core logic resides in [`doc2graph/data/graph_builder.py`](https://github.com/andreagemelli/doc2graph/blob/main/doc2graph/data/graph_builder.py) at lines 57-94 within the `GraphBuilder` class. The high-level wrapper `balance` is defined in [`doc2graph/data/dataloader.py`](https://github.com/andreagemelli/doc2graph/blob/main/doc2graph/data/dataloader.py) at lines 47-64. Both files are part of the `andreagemelli/doc2graph` repository and require DGL (Deep Graph Library) as a dependency for the `dgl.remove_edges` operations.