A2GNN(in_dim, hid_dim, num_classes, mode='node', num_layers=3, dropout=0.0, act=F.relu, s_pnums=0, t_pnums=30, adv=False, weight=5, weight_decay=0.0, lr=0.004, epoch=200, device='cuda:0', batch_size=0, num_neigh=-1, verbose=2, **kwargs)
Bases: BaseGDA
Rethinking Propagation for Unsupervised Graph Domain Adaptation (AAAI-24).
| Parameters: |
|
|---|
fit(source_data, target_data)
Train the A2GNN model on source and target domain data.
| Parameters: |
|
|---|
Notes
Training process includes:
Data Preparation
- Configures loaders for node/graph level tasks
- Handles both full-batch and mini-batch scenarios
- Sets up appropriate batch processing
Training Loop
- Dynamic adaptation parameter scaling
- Asymmetric message propagation
-
Domain adaptation through either:
- Adversarial training
- MMD minimization
-
Comprehensive progress monitoring
Implementation Features
- Flexible task handling (node/graph)
- Efficient batch processing
- Adaptive learning mechanisms
forward_model(source_data, target_data, alpha)
Forward pass of the A2GNN model.
| Parameters: |
|
|---|
| Returns: |
|
|---|
Notes
Implements multiple components:
- Asymmetric propagation (s_pnums vs t_pnums)
- Classification loss on source domain
- Domain adaptation (adversarial or MMD)
- Feature bottleneck processing
init_model(**kwargs)
Initialize the A2GNN base model.
| Parameters: |
|
|---|
| Returns: |
|
|---|
Notes
Configures model with:
- Asymmetric propagation settings
- Domain adaptation components
- Task-specific architecture (node/graph)
- Optional adversarial training module
predict(data, source=False)
Make predictions on input data.
| Parameters: |
|
|---|
| Returns: |
|
|---|
Notes
- Uses different propagation steps for source/target
- Handles batch processing efficiently
- Concatenates results for full predictions
- Maintains evaluation mode consistency
process_graph(data)
Process the input graph data.
| Parameters: |
|
|---|
Notes
Placeholder method as preprocessing is handled through:
- Asymmetric propagation mechanisms
- Domain-specific feature processing
- Batch-wise data handling