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.
- Overview
- Key Features
- Architecture
- Installation
- Quick Start
- Usage
- Project Structure
- Methodology
- Configuration
- Results and Metrics
- Datasets
- References
- Known Issues
IRGNN addresses the challenge of interpretability in relational deep learning by combining two powerful techniques:
- Relational GNNs: Learn from relational databases by converting them into heterogeneous graphs and applying message-passing neural networks
- 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
- 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
The base model consists of three main components:
-
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
-
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
-
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
The LassoNet variant extends the base model with:
-
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
-
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
- Enforces hierarchical sparsity constraint:
-
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
cd src
python base_train.py \
--data.dataset rel-f1 \
--data.task driver-position \
--trainer.epochs 20 \
--trainer.lr 0.001python 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.1python selective_eval.py \
--lasso_result_path lightning_logs/lasso_experiment/sparsity_results.json \
--sparsity_levels 0.5 0.3 0.1The 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 512Key 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
LassoNet training follows a two-phase approach:
- Dense Phase: Train without regularization
- 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 50Key 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": {...}
},
...
]
}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 30This:
- Loads sparsity results from LassoNet training
- For each sparsity level, prunes the database to keep only selected features
- Trains BaseRelGNN from scratch on pruned data
- Compares performance: full model vs. sparse models
For statistical significance, run multiple trials:
python base_eval.py \
--data.dataset rel-f1 \
--data.task driver-position \
--num_trials 5 \
--epochs 30Outputs mean and standard deviation of metrics across runs to base_metrics.json.
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
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
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.
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.
For each feature
Minimize:
When
-
Dense Training (epochs 0 to
dense_epochs):- No regularization (
$\lambda = 0$ ) - Train all parameters normally
- No regularization (
-
Compute
$\lambda_0$ :- Binary search for smallest
$\lambda$ that zeros at least one feature - Ensures meaningful sparsity trajectory
- Binary search for smallest
-
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()
-
Termination:
- Stop when
active_features≤max_active_features - Or when all features removed
- Stop when
Validates that selected features maintain performance:
- Load sparsity trajectory from LassoNet training
- For target sparsity level (e.g., 0.3 = 30% features):
- Identify active features at that sparsity
- Create pruned database with only those columns
- Train BaseRelGNN from scratch on pruned data
- Compare test metrics: sparse vs. full model
All training scripts use structured configuration with jsonargparse:
@dataclass
class DataConfig:
dataset: str = 'rel-f1' # RelBench dataset
task: str = 'driver-position' # Prediction task
cache_dir: str = 'data' # Dataset cache directory@dataclass
class TextEmbedderConfig:
model_name: str = 'average_word_embeddings_glove.6B.300d'
batch_size: int = 256
device: str = 'cpu'@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@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@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# 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-
Training Logs:
lightning_logs/version_X/- TensorBoard events
- Hyperparameters
- Model checkpoints
-
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", ...] } } ] } -
Evaluation Metrics:
base_metrics.jsonorselective_metrics.json{ "mean": { "val_metric": 0.245, "test_metric": 0.251 }, "std": { "val_metric": 0.012, "test_metric": 0.015 } }
- Regression Tasks: Mean Absolute Error (MAE) - lower is better
- Classification Tasks: ROC-AUC - higher is better
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_columnsat 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
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 |
Download and cache datasets:
cd src
bash materialize.shOr programmatically:
from relbench.datasets import get_dataset
dataset = get_dataset('rel-f1', download=True)Datasets are cached in src/data/ by default.
This project builds on the following research:
-
LassoNet: Lemhadri, I., Ruan, F., Abraham, L., & Tibshirani, R. (2021). LassoNet: A Neural Network with Feature Sparsity. JMLR.
-
Relational Deep Learning: Fey, M., et al. (2024). Relational Deep Learning: Graph Representation Learning on Relational Databases. arXiv:2312.04615.
-
RelBench: Fey, M., et al. (2024). RelBench: A Benchmark for Deep Learning on Relational Databases. NeurIPS Datasets and Benchmarks Track.
-
RelGNN: Existing work on composite message passing for heterogeneous graphs in relational settings.
See ref/ directory for full papers.
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.
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_sizeor--loader.num_neighborsif OOM
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_startis reasonable (should be small, e.g., 1e-3 to 0.1)
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}
}