How Doc2Graph Handles Class Imbalance with the balance_edges Method
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 (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:
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, the method creates a boolean mask and extracts indices for all edges matching the target class:
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:
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:
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 (lines 47-64). The DataLoader.balance method handles class name resolution and iterates over the entire dataset:
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:
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:
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_edgesmethod indoc2graph/data/graph_builder.pyrandomly downsamples dominant edge classes to match the combined count of minority classes. - Mathematical Balance: The retention formula
(edge_targets != cls).sum() / 2ensures 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_edgesfor single graphs, or batch processing throughDataLoader.balanceindoc2graph/data/dataloader.py. - DGL Integration: Uses native
dgl.remove_edgesfor 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 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 at lines 57-94 within the GraphBuilder class. The high-level wrapper balance is defined in 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.
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 →