GraphCTA(in_dim, hid_dim, num_classes, num_layers=3, dropout=0.0, act=F.relu, loop_model=3, loop_adj=1, loop_feat=4, ratio=0.1, K=5, tau=0.2, lamb=0.2, momentum=0.9, make_undirected=True, weight_decay=0.0, lr=0.0001, epoch=500, gnn='gcn', device='cuda:0', batch_size=0, num_neigh=-1, verbose=2, **kwargs)
Bases: BaseGDA
Collaborate to Adapt: Source-Free Graph Domain Adaptation via Bi-directional Adaptation (WWW-24).
| Parameters: |
|
|---|
entropy(input_)
Calculate entropy of probability distribution.
| Parameters: |
|
|---|
| Returns: |
|
|---|
Notes
- Handles numerical stability with epsilon
- Used for confidence-based pseudo-labeling
fit(source_data, target_data)
Train the GraphCTA model with collaborative adaptation.
| Parameters: |
|
|---|
Notes
Implementation consists of three main phases:
Initialization
- Sets up data loaders
- Initializes model and optimizers
- Creates memory banks for features and classes
- Prepares feature and structure perturbation variables
Source Pretraining
- Trains model on source domain
- Uses standard cross-entropy loss
- Prepares model for adaptation
Target Adaptation (Iterative)
-
Model Update Loop:
- Updates model parameters
- Computes prototype-based alignment
- Updates memory banks with momentum
- Combines local and contrastive losses
-
Feature Optimization Loop:
- Optimizes feature perturbations
- Uses test-time adaptation loss
- Maintains feature consistency
-
Structure Optimization Loop:
- Modifies edge weights
- Ensures budget constraints
- Preserves graph properties
Implementation Features:
- Memory-based prototype learning
- Gradient checkpointing for efficiency
- Momentum updates for stability
- Budget-constrained modifications
- Multiple optimization objectives
- Collaborative feature-structure adaptation
forward_model(data, **kwargs)
Forward pass placeholder for GraphCTA model.
| Parameters: |
|
|---|
Notes
Placeholder method as GraphCTA implements custom forward logic through:
- Source domain training in train_source()
- Collaborative adaptation in fit()
- Feature and structure optimization
- Memory-based prototype learning
init_model(**kwargs)
Initialize the GraphCTA base model.
| Parameters: |
|
|---|
| Returns: |
|
|---|
Notes
Configures base GNN model with:
- Input and hidden dimensions
- Number of classes and layers
- Dropout rate
- Specified GNN backbone type
- Device placement
instance_proto_alignment(feat, center, pred)
Compute instance-prototype alignment loss.
| Parameters: |
|
|---|
| Returns: |
|
|---|
Notes
- Implements temperature-scaled contrastive loss
- Handles both instance-prototype and instance-instance relations
- Uses cosine similarity for feature comparison
predict(data)
Make predictions using the adapted model.
| Parameters: |
|
|---|
| Returns: |
|
|---|
Notes
- Uses transformed features (self.new_feat)
- Uses optimized edge structure (self.edge_index, self.edge_weight)
- Evaluates model in inference mode
process_graph(data)
Process input graph data.
| Parameters: |
|
|---|
Notes
Placeholder method as graph processing is handled through:
- Feature perturbation optimization
- Structure modification
- Memory bank updates
- Prototype-based alignment
sample_final_edges(n_perturbations, perturbed_edge_weight, data, modified_edge_index, n, mem_fea, mem_cls)
Sample final edge structure based on learned weights.
| Parameters: |
|
|---|
| Returns: |
|
|---|
Notes
- Uses iterative sampling strategy
- Maintains best performing structure
- Ensures perturbation budget constraints
- Handles undirected graph requirements
test_time_loss(feat, edge_index, edge_weight, mem_fea, mem_cls)
Compute test-time adaptation loss.
| Parameters: |
|
|---|
| Returns: |
|
|---|
Notes
Loss components include:
- Pseudo-label based classification
- Feature similarity with memory bank
- Class-wise prototype alignment
- Confidence-based sample selection
train_source(optimizer)
Train the model on source domain data.
| Parameters: |
|
|---|
Notes
Training process includes:
Per-epoch Operations:
- Tracks cumulative loss
- Maintains logits and labels
- Computes performance metrics
Batch Processing:
- Moves data to device
- Computes model predictions
- Applies negative log-likelihood loss
- Updates model parameters
Monitoring:
- Computes micro-F1 score
- Logs training progress
- Tracks timing information
Implementation Features:
- Supports batch processing
- Uses softmax with log probabilities
- Accumulates predictions for full evaluation
- Comprehensive logging
update_edge_weights(gradient, optimizer_adj, perturbed_edge_weight)
Update edge weights during structure optimization.
| Parameters: |
|
|---|
Notes
- Applies gradient updates to edge weights
- Maintains minimum weight threshold
- Uses Adam optimizer for updates