SpecReg(in_dim, hid_dim, num_classes, num_layers=3, dropout=0.0, act=F.relu, ppmi=True, adv_dim=40, reg_mode=True, gamma_adv=0.1, thr_smooth=-1, gamma_smooth=0.01, thr_mfr=-1, gamma_mfr=0.01, weight_decay=0.003, lr=0.004, epoch=100, device='cuda:0', batch_size=0, num_neigh=-1, verbose=2, **kwargs)
Bases: BaseGDA
Graph Domain Adaptation via Theory-Grounded Spectral Regularization (ICLR-23).
| Parameters: |
|
|---|
calculate_gradient_penalty(x_src, x_tgt)
Calculate gradient penalty for Wasserstein GAN training.
| Parameters: |
|
|---|
| Returns: |
|
|---|
Notes
Implements Wasserstein GAN gradient penalty by:
- Interpolating between source and target features
- Computing gradients of critic output
- Penalizing gradients that deviate from norm 1
fit(source_data, target_data)
Train the SpecReg model.
| Parameters: |
|
|---|
Notes
Training process includes:
- Setting up data loaders
- Initializing model, critic, and optimizers
- Alternating training between:
- Critic optimization (Wasserstein distance)
- Model optimization with spectral regularization
- Computing and logging training metrics
forward_model(source_data, target_data, alpha, epoch)
Forward pass of the model.
| Parameters: |
|
|---|
| Returns: |
|
|---|
Notes
Computes multiple loss terms:
- Classification loss on source domain
- Wasserstein distance with gradient penalty
- Spectral smoothness regularization (if reg_mode)
- Maximum Frequency Response regularization (if reg_mode)
- Entropy minimization on target domain
init_model(**kwargs)
Initialize the SpecReg model.
| Parameters: |
|
|---|
| Returns: |
|
|---|
predict(data, source=False)
Make predictions on given data.
| Parameters: |
|
|---|
| Returns: |
|
|---|
Notes
Uses appropriate encoder based on domain (source/target).
process_graph(data)
Process the input graph data.
| Parameters: |
|
|---|
Notes
Placeholder method for graph preprocessing.