dkapur17/irgnn

Interpretability Experiments on GNNs for Relational Deep Learning (CMU 10747 Course Project)

★ 0Forks 0Jupyter NotebookGitHub ↗Compare

README

Interpretable Relational GNN (IRGNN)

Final Course Project for CMU 10747 Neurosymbolic AI

A novel approach combining relational deep learning using Graph Neural Networks (GNNs) with LassoNet for feature sparsity, enabling interpretable machine learning on relational databases.

Table of Contents

Overview

IRGNN addresses the challenge of interpretability in relational deep learning by combining two powerful techniques:

  1. Relational GNNs: Learn from relational databases by converting them into heterogeneous graphs and applying message-passing neural networks
  2. LassoNet: Induce feature sparsity through hierarchical proximal operators, identifying the most important features for predictions

This approach enables:

  • Training on complex relational databases with multiple entity types and relationships
  • Automatic feature selection to identify critical predictive features
  • Interpretable models that can explain which database columns drive predictions
  • Validation that sparse feature sets maintain predictive performance

Key Features

  • Heterogeneous Graph Neural Networks: Process multiple node and edge types with type-specific encoders and message-passing layers
  • Feature Sparsity via LassoNet: Hierarchical proximal operators for principled feature selection
  • Lambda Path Algorithm: Incremental sparsity increase with performance tracking at each level
  • Column Selective Evaluation: Validate sparse models by retraining on pruned databases
  • Multiple Task Types: Support for both regression (MAE) and binary classification (ROC-AUC)
  • Text Embedding: Automatic handling of text features using Sentence Transformers
  • Comprehensive Logging: PyTorch Lightning integration with detailed metric tracking
  • RelBench Integration: Standardized evaluation on benchmark relational datasets

Architecture

Base Relational GNN

The base model consists of three main components:

  1. RelFeatureEncoder (src/rgnn_modules.py:28-74)

    • Heterogeneous feature encoding using PyTorch Frame
    • Type-specific encoders for categorical, numerical, timestamp, and embedding features
    • Per-node-type MLPs to combine features and project to hidden dimension to get initial node representation
  2. RelConv (src/rgnn_modules.py:122-181)

    • Heterogeneous graph convolution using PyTorch Geometric's HeteroConv
    • Configurable message-passing layers (SAGEConv, GATConv, SimpleConv)
    • Layer normalization and ReLU activation
  3. BaseRelGNN (src/rgnn_modules.py:184-234)

    • PyTorch Lightning module for training and evaluation
    • Task-agnostic: supports regression and classification
    • Automatic metric computation and logging

LassoNet-Enhanced Architecture

The LassoNet variant extends the base model with:

  1. LassoRelEncoder (src/lasso_modules.py:29-88)

    • Optional encoder freezing to focus optimization on feature selection
    • Skip connections from raw features to output
    • Returns both pre-MLP features (for skip) and post-MLP embeddings
  2. Hierarchical Proximal Operator (src/lasso_modules.py:263-375)

    • Enforces hierarchical sparsity constraint: $||W^{(1)}_j || \leq M * |\theta_j\vert$
    • Features are removed when skip weight $\theta_j \to 0$
    • Applied after each optimizer step during sparse training
  3. Lambda Path Following (src/lasso_modules.py:377-470)

    • Dense training phase (no regularization)
    • Automatic $\lambda_0$ computation through prox sparsity simulation
    • Incremental regularization: $\lambda_{k+1} = \lambda (1 + \epsilon)$
    • Train $B$ epochs at each $\lambda$ value with sparsity tracking

Quick Start

Train a Base RelGNN

cd src
python base_train.py \
  --data.dataset rel-f1 \
  --data.task driver-position \
  --trainer.epochs 20 \
  --trainer.lr 0.001

Train with LassoNet for Feature Selection

python lasso_train.py \
  --data.dataset rel-f1 \
  --data.task driver-dnf \
  --model.lasso.dense_epochs 15 \
  --model.lasso.path_epochs 3 \
  --model.lasso.M 10.0 \
  --model.lasso.epsilon 0.1

Evaluate Sparse Model

python selective_eval.py \
  --lasso_result_path lightning_logs/lasso_experiment/sparsity_results.json \
  --sparsity_levels 0.5 0.3 0.1

Usage

Training Base RelGNN

The base model trains a standard relational GNN without feature selection:

python base_train.py --help  # See all options

# Example: Train on rel-trial dataset
python base_train.py \
  --data.dataset rel-trial \
  --data.task site-sponsor-run \
  --model.hidden_dim 128 \
  --model.num_gnn_layers 2 \
  --trainer.epochs 50 \
  --trainer.batch_size 512

Key parameters:

  • --data.dataset: RelBench dataset name
  • --data.task: Prediction task within dataset
  • --model.hidden_dim: Hidden dimension for GNN layers
  • --model.num_gnn_layers: Number of message-passing layers
  • --trainer.epochs: Training epochs

Training with LassoNet

LassoNet training follows a two-phase approach:

  1. Dense Phase: Train without regularization
  2. Sparse Phase: Gradually increase regularization along λ-path
python lasso_train.py \
  --data.dataset rel-hm \
  --data.task user-churn \
  --model.lasso.dense_epochs 20 \
  --model.lasso.path_epochs 5 \
  --model.lasso.M 10.0 \
  --model.lasso.epsilon 0.05 \
  --model.lasso.max_active_features 50

Key parameters:

  • --model.lasso.dense_epochs: Epochs for dense training (no sparsity)
  • --model.lasso.path_epochs: Epochs per λ value during sparse phase
  • --model.lasso.M: Hierarchical penalty parameter
  • --model.lasso.epsilon: λ increment factor (λ *= 1 + ε)
  • --model.lasso.max_active_features: Stop when this many features remain
  • --model.lasso.lambda_start: Manual λ start (auto-computed if omitted)

Output:

  • Model checkpoint: lightning_logs/version_X/checkpoints/
  • Sparsity trajectory: lightning_logs/version_X/sparsity_results.json

The sparsity results JSON contains:

{
  "sparsity_trajectory": [
    {
      "lambda": 0.05,
      "active_features": 150,
      "active_columns": 45,
      "test_metric": 0.234,
      "selected_features": {...}
    },
    ...
  ]
}

Selective Evaluation

Validate that selected features are sufficient by retraining on pruned databases:

python selective_eval.py \
  --lasso_result_path lightning_logs/version_5/sparsity_results.json \
  --sparsity_levels 0.5 0.3 0.1 \
  --num_trials 5 \
  --epochs 30

This:

  1. Loads sparsity results from LassoNet training
  2. For each sparsity level, prunes the database to keep only selected features
  3. Trains BaseRelGNN from scratch on pruned data
  4. Compares performance: full model vs. sparse models

Multi-Run Evaluation

For statistical significance, run multiple trials:

python base_eval.py \
  --data.dataset rel-f1 \
  --data.task driver-position \
  --num_trials 5 \
  --epochs 30

Outputs mean and standard deviation of metrics across runs to base_metrics.json.

Project Structure

irgnn/
├── README.md                          # This file
├── requirements.txt                   # Python dependencies
├── ref/                              # Reference papers
│   ├── LassoNet A Neural Network with Feature Sparsity.pdf
│   ├── Relational Deep Learning Graph Representation Learning on Relational Databases.pdf
│   ├── RelBench A Benchmark for Deep Learning on Relational Databases.pdf
│   └── RelGNN Composite Message Passing for Relational Deep Learning.pdf
└── src/
    ├── rgnn_modules.py               # Base relational GNN components
    ├── lasso_modules.py              # LassoNet-enhanced GNN components
    ├── utils.py                      # Data loading and preprocessing
    ├── base_train.py                 # Training script for base RelGNN
    ├── lasso_train.py                # Training script for LassoNet variant
    ├── base_eval.py                  # Multi-run evaluation
    ├── selective_eval.py             # Evaluation on pruned databases
    ├── materialize.sh                # Dataset download script
    ├── data/                         # Cached datasets
    ├── lightning_logs/               # Training logs and checkpoints
    └── notebooks/                    # Jupyter notebooks
        ├── base.ipynb                # Base RelGNN experiments
        ├── lasso.ipynb               # LassoNet experiments
        ├── selective.ipynb           # Sparse model analysis
        └── visualize_metrics.ipynb   # Result visualization

Methodology

1. Graph Construction

Relational databases are converted to heterogeneous graphs:

  • Nodes: Each database row becomes a node
  • Node Types: One per database table
  • Edges: Foreign key → Primary key relationships
  • Edge Types: (source_table, relationship_name, target_table)

Implementation: utils.py:make_pkey_fkey_graph

2. Feature Encoding

Heterogeneous features are encoded using TensorFrame:

  • Categorical: Embedding lookup
  • Numerical: Linear projection
  • Timestamp: Positional encoding
  • Text: Sentence Transformer embeddings (pre-computed)

Each node type has its own encoder, then per-type MLPs project to common hidden dimension.

3. Message Passing

Heterogeneous message passing aggregates information from neighbors:

For each edge type (src, rel, dst):
  messages = SAGE_CONV(x_src, edge_index)
  x_dst = AGGREGATE(messages, aggr='sum')
  x_dst = LAYER_NORM(x_dst)
  x_dst = RELU(x_dst)

Multiple layers enable multi-hop reasoning over the relational graph.

4. LassoNet Feature Selection

Hierarchical Proximal Operator

For each feature $j$ with skip weight $\theta_j$ and first-layer weight $W^{(1)}_j \in \R^d$

Minimize: $\ell(\theta, W) + \lambda ||\theta||_1$ subject to $||W^{(1)}_j|| \leq M|\theta_j|$

When $\theta_j \to 0$, the corresponding feature is removed.

Lambda Path Algorithm

  1. Dense Training (epochs 0 to dense_epochs):

    • No regularization ($\lambda = 0$)
    • Train all parameters normally
  2. Compute $\lambda_0$:

    • Binary search for smallest $\lambda$ that zeros at least one feature
    • Ensures meaningful sparsity trajectory
  3. Sparse Training:

    for λ in [λ_start, λ_start(1+ε), λ_start(1+ε)², ...]:
        for epoch in range(path_epochs):
            for batch in dataloader:
                loss.backward()
                optimizer.step()
                apply_hierarchical_prox(θ, W1, λ, M)  # Project to feasible set
    
        if sparsity_changed:
            evaluate_on_test_set()
            save_active_features()
  4. Termination:

    • Stop when active_features ≤ max_active_features
    • Or when all features removed

5. Selective Evaluation

Validates that selected features maintain performance:

  1. Load sparsity trajectory from LassoNet training
  2. For target sparsity level (e.g., 0.3 = 30% features):
    • Identify active features at that sparsity
    • Create pruned database with only those columns
  3. Train BaseRelGNN from scratch on pruned data
  4. Compare test metrics: sparse vs. full model

Configuration

All training scripts use structured configuration with jsonargparse:

Data Configuration

@dataclass
class DataConfig:
    dataset: str = 'rel-f1'              # RelBench dataset
    task: str = 'driver-position'        # Prediction task
    cache_dir: str = 'data'              # Dataset cache directory

Text Embedding Configuration

@dataclass
class TextEmbedderConfig:
    model_name: str = 'average_word_embeddings_glove.6B.300d'
    batch_size: int = 256
    device: str = 'cpu'

Model Configuration

@dataclass
class ModelConfig:
    hidden_dim: int = 128                # Hidden dimension
    encoder_dim: int = 64                # Per-node-type encoder output dim
    num_gnn_layers: int = 2              # Message-passing layers
    gnn_type: str = 'SAGEConv'           # GNN layer type

LassoNet Configuration

@dataclass
class LassoConfig:
    M: float = 10.0                      # Hierarchical penalty parameter
    epsilon: float = 0.1                 # Lambda increment: λ *= (1 + ε)
    dense_epochs: int = 15               # Dense training epochs
    path_epochs: int = 3                 # Epochs per λ value
    lambda_start: Optional[float] = None # Auto-computed if None
    max_active_features: int = 50        # Stopping criterion
    freeze_encoders: bool = True         # Freeze feature encoders

Trainer Configuration

@dataclass
class TrainerConfig:
    lr: float = 0.001                    # Learning rate
    epochs: int = 30                     # Total epochs (base) or dense epochs (lasso)
    batch_size: int = 512                # Batch size
    num_neighbors: List[int] = (128, 64) # Neighbor sampling per layer
    num_workers: int = 0                 # DataLoader workers

Command-Line Usage

# View configuration schema
python base_train.py --print_config

# Save configuration to file
python base_train.py --print_config > config.yaml

# Load configuration from file
python base_train.py --config config.yaml

# Override specific parameters
python lasso_train.py \
  --config config.yaml \
  --model.lasso.M 15.0 \
  --trainer.lr 0.0005

Results and Metrics

Output Files

  1. Training Logs: lightning_logs/version_X/

    • TensorBoard events
    • Hyperparameters
    • Model checkpoints
  2. Sparsity Results: sparsity_results.json

    {
      "sparsity_trajectory": [
        {
          "lambda": 0.05,
          "active_features": 150,
          "active_columns": 45,
          "test_metric": 0.234,
          "selected_features": {
            "entity_table": [0, 5, 12, ...],
            "other_table": [2, 8, ...]
          },
          "selected_columns": {
            "entity_table": ["age", "income", ...],
            "other_table": ["category", ...]
          }
        }
      ]
    }
  3. Evaluation Metrics: base_metrics.json or selective_metrics.json

    {
      "mean": {
        "val_metric": 0.245,
        "test_metric": 0.251
      },
      "std": {
        "val_metric": 0.012,
        "test_metric": 0.015
      }
    }

Metric Types

  • Regression Tasks: Mean Absolute Error (MAE) - lower is better
  • Classification Tasks: ROC-AUC - higher is better

Interpreting Results

Sparsity vs. Performance Trade-off:

  • Plot test_metric vs. active_columns to visualize trade-off
  • Identify "elbow point" where performance drops significantly
  • Optimal sparsity level balances interpretability and accuracy

Feature Importance:

  • Features removed early (high $\lambda$) are less important
  • Features remaining at low sparsity are critical
  • Examine selected_columns at different sparsity levels

Validation:

  • Compare selective evaluation metrics to full model
  • If sparse model ≈ full model performance, features are sufficient
  • Large performance drop suggests important features excluded

Datasets

IRGNN uses the RelBench benchmark suite:

Dataset Description Tables Tasks
rel-f1 Formula 1 racing data 8 driver-position, driver-dnf
rel-hm H&M fashion retail 5 user-churn, product-sales
rel-trial Clinical trials 6 site-sponsor-run
rel-avito Online classifieds 7 ad-ctr
rel-stack Stack Overflow 4 user-badge, post-votes
rel-event Event management 6 user-repeat-attendance
rel-amazon Product reviews 3 product-churn, user-ltv

Dataset Materialization

Download and cache datasets:

cd src
bash materialize.sh

Or programmatically:

from relbench.datasets import get_dataset
dataset = get_dataset('rel-f1', download=True)

Datasets are cached in src/data/ by default.

References

This project builds on the following research:

  1. LassoNet: Lemhadri, I., Ruan, F., Abraham, L., & Tibshirani, R. (2021). LassoNet: A Neural Network with Feature Sparsity. JMLR.

  2. Relational Deep Learning: Fey, M., et al. (2024). Relational Deep Learning: Graph Representation Learning on Relational Databases. arXiv:2312.04615.

  3. RelBench: Fey, M., et al. (2024). RelBench: A Benchmark for Deep Learning on Relational Databases. NeurIPS Datasets and Benchmarks Track.

  4. RelGNN: Existing work on composite message passing for heterogeneous graphs in relational settings.

See ref/ directory for full papers.

Known Issues

MPS Compatibility

Some RelBench functions use float64 internally, which is incompatible with Apple MPS (Metal Performance Shaders):

Workaround: Modify library functions to use np.float32:

# In relbench/modeling/utils.py or similar
# Change:
arr = np.array(..., dtype=np.float64)
# To:
arr = np.array(..., dtype=np.float32)

Affected functions:

  • get_node_train_table_input
  • Potentially others in data loading pipeline

Recommendation: Use CUDA if available, or CPU for small datasets.

Memory Usage

Large datasets (rel-amazon, rel-stack) may require:

  • 16GB+ RAM for full graph materialization
  • GPU with 8GB+ VRAM for batch size 512
  • Reduce --trainer.batch_size or --loader.num_neighbors if OOM

LassoNet Convergence

If sparsity trajectory is unstable:

  • Decrease epsilon (slower $\lambda$ increase)
  • Increase path_epochs (more training per $\lambda$)
  • Adjust M (higher = more aggressive sparsity)
  • Check that lambda_start is reasonable (should be small, e.g., 1e-3 to 0.1)

Citation

If you use this code, please cite:

@misc{irgnn2025,
  title={Interpretable Relational GNN: Feature Sparsity in Relational Deep Learning with LassoNet},
  author={Dhruv Kapur, Priyanka Vijaybhaskar},
  year={2025},
  note={CMU 10747 Course Project}
}

Contributors

dkapur17

Issues