ASNBase(in_dim, hid_dim, hid_dim_vae, num_classes, num_layers=3, act=F.relu, dropout=0.1, adv_dim=40, **kwargs)

Bases: Module

Base class for ASN.

Parameters:
  • in_dim (int) –

    Input feature dimension.

  • hid_dim (int) –

    Hidden dimension.

  • hid_dim_vae (int) –

    VAE hidden dimension.

  • num_classes (int) –

    Number of classes.

  • num_layers (int, default: 3 ) –

    Number of layers. Default: 3.

  • act (callable, default: relu ) –

    Activation function. Default: F.relu.

  • dropout (float, default: 0.1 ) –

    Dropout rate. Default: 0.1.

  • adv_dim (int, default: 40 ) –

    Adversarial module dimension. Default: 40.

Notes

Architecture components:

  1. Private encoders (local and global) for source/target
  2. Shared encoders (local and global)
  3. Decoders for reconstruction
  4. Domain discriminator
  5. Attention mechanisms
adj_label_for_reconstruction(data)

Prepare adjacency matrix for reconstruction.

Parameters:
  • data (Data) –

    Input graph data.

Returns:
  • tuple

    Contains: - adj_label : Processed adjacency matrix - pos_weight : Positive class weight - norm : Normalization factor

recon_loss(preds, labels, mu, logvar, n_nodes, norm, pos_weight)

Compute reconstruction loss with KL divergence.

Parameters:
  • preds (Tensor) –

    Predicted adjacency matrix.

  • labels (Tensor) –

    True adjacency matrix.

  • mu (Tensor) –

    Mean of latent distribution.

  • logvar (Tensor) –

    Log variance of latent distribution.

  • n_nodes (int) –

    Number of nodes.

  • norm (float) –

    Normalization factor.

  • pos_weight (Tensor) –

    Positive class weight.

Returns:
  • Tensor

    Combined reconstruction and KL loss.

DiffLoss()

Bases: Module

Difference loss for enforcing feature separation.

Notes
  • Computes normalized feature differences
  • Used to encourage orthogonality between domains
  • L2 normalization for stability
forward(input1, input2)

Compute difference loss between two feature sets.

Parameters:
  • input1 (Tensor) –

    First feature set.

  • input2 (Tensor) –

    Second feature set.

Returns:
  • Tensor

    Difference loss value.

Notes
  • L2 normalization of inputs
  • Computes mean squared inner product
  • Encourages orthogonality
GNNVAE(in_dim, hid_dim, num_classes, gnn_type='gcn', num_layers=3, base_model=None, act=F.relu, **kwargs)

Bases: Module

Graph Neural Network Variational Autoencoder.

Parameters:
  • in_dim (int) –

    Input feature dimension.

  • hid_dim (int) –

    Hidden dimension.

  • num_classes (int) –

    Number of output classes.

  • gnn_type (str, default: 'gcn' ) –

    Type of GNN layer ('gcn' or 'ppmi'). Default: 'gcn'.

  • num_layers (int, default: 3 ) –

    Number of GNN layers. Default: 3.

  • base_model (Module, default: None ) –

    Base model for weight initialization.

  • act (callable, default: relu ) –

    Activation function. Default: F.relu.

Notes
  • Implements variational encoding
  • Supports weight sharing
  • Multiple GNN layer types
forward(x, edge_index)

Forward pass of GNN-VAE.

Parameters:
  • x (Tensor) –

    Node features.

  • edge_index (Tensor) –

    Edge indices.

Returns:
  • tuple

    Contains: - z : Sampled latent vectors - mu : Mean of latent distribution - logvar : Log variance of latent distribution

reparameterize(mu, logvar)

Perform reparameterization trick.

Parameters:
  • mu (Tensor) –

    Mean of latent distribution.

  • logvar (Tensor) –

    Log variance of latent distribution.

Returns:
  • Tensor

    Sampled latent vectors.

InnerProductDecoder(dropout, act=torch.sigmoid)

Bases: Module

Decoder module using inner product for graph reconstruction.

Parameters:
  • dropout (float) –

    Dropout rate.

  • act (callable, default: sigmoid ) –

    Activation function. Default: torch.sigmoid.

Notes
  • Used for reconstructing adjacency matrices
  • Applies dropout for regularization
  • Configurable activation function
forward(z)

Decode latent representations to adjacency matrix.

Parameters:
  • z (Tensor) –

    Latent node representations.

Returns:
  • Tensor

    Reconstructed adjacency matrix.