AdaGCN(in_dim, hid_dim, num_classes, mode='node', num_layers=3, dropout=0.0, act=F.relu, gnn_type='gcn', adv_dim=40, gp_weight=5, domain_weight=1, weight_decay=0.0, lr=0.004, epoch=100, device='cuda:0', batch_size=0, num_neigh=-1, verbose=2, **kwargs)
Bases: BaseGDA
Graph Transfer Learning via Adversarial Domain Adaptation with Graph Convolution (TKDE-22).
| Parameters: |
|
|---|
fit(source_data, target_data)
Train the AdaGCN model.
| Parameters: |
|
|---|
Notes
Training process includes:
- Setting up data loaders for both domains
- Initializing GNN and discriminator
- Training with adversarial learning
- Supporting both node and graph level tasks
forward_model(source_data, target_data)
Forward pass of the model.
| Parameters: |
|
|---|
| Returns: |
|
|---|
Notes
Performs adversarial training with:
- Discriminator optimization
- Gradient penalty computation
- Classification loss
- Domain adaptation loss
gradient_penalty(encoded_source, encoded_target)
Compute gradient penalty for Wasserstein GAN training.
| Parameters: |
|
|---|
| Returns: |
|
|---|
Notes
Implements Wasserstein GAN gradient penalty by:
- Interpolating between source and target features
- Computing gradients w.r.t. discriminator outputs
- Penalizing gradients that deviate from norm 1
- Handling different batch sizes between domains
init_model(**kwargs)
Initialize the AdaGCN model.
| Parameters: |
|
|---|
| Returns: |
|
|---|
predict(data, source=False)
Make predictions on given data.
| Parameters: |
|
|---|
| Returns: |
|
|---|
Notes
Handles predictions for both source and target domains using appropriate data loaders.
process_graph(data)
Process the input graph data.
| Parameters: |
|
|---|
Notes
Placeholder method for graph preprocessing.