GCNPooling(in_dim, hid_dim, device, sparse)

Bases: Module

GCN-based hierarchical pooling layer.

Parameters:
  • in_dim (int) –

    Input feature dimension.

  • hid_dim (int) –

    Number of nodes in pooled graph.

  • device (device) –

    Device to use.

  • sparse (bool) –

    Whether to use sparse assignment matrix.

forward(X_old, edge_index, edge_weight, A_old, Y_old, Z, use_sparse=False)

Forward pass of pooling layer.

Parameters:
  • X_old (Tensor) –

    Node features.

  • edge_index (Tensor) –

    Edge indices.

  • edge_weight (Tensor) –

    Edge weights.

  • A_old (Tensor) –

    Adjacency matrix.

  • Y_old (Tensor) –

    Node labels.

  • Z (Tensor) –

    Node embeddings.

  • use_sparse (bool, default: False ) –

    Whether to use sparse operations.

Returns:
  • tuple

    Contains:

    • S: Assignment matrix
    • X_new: Pooled features
    • A_new: Pooled adjacency
    • Y_new: Pooled labels
    • Y_new_prob: Label probabilities
to_onehot(label_matrix, num_classes)

Convert labels to one-hot encoding.

Parameters:
  • label_matrix (Tensor) –

    Label indices.

  • num_classes (int) –

    Number of classes.

Returns:
  • Tensor

    One-hot encoded labels.

JHGDABase(in_dim, hid_dim, num_classes, device, pool_ratio, num_s, num_t, num_layers=3, dropout=0.1, act=F.relu, share=False, classwise=False, sparse=False, **kwargs)

Bases: Module

Base class for JHGDA.

Parameters:
  • in_dim (int) –

    Input dimension of model.

  • hid_dim (int) –

    Hidden dimension of model.

  • num_classes (int) –

    Number of classes.

  • device (str) –

    GPU or CPU.

  • num_layers (int, default: 3 ) –

    Total number of layers in model. Default: 4.

  • dropout (float, default: 0.1 ) –

    Dropout rate. Default: 0..

  • act (callable activation function or None, default: relu ) –

    Activation function if not None. Default: torch.nn.functional.relu.

  • share

    Share the diffpool module or not. Default: False.

  • sparse

    Diffpool module sparse or not. Default: False.

  • classwise

    Classwise conditional shift or not. Default: True.

  • **kwargs (optional, default: {} ) –

    Other parameters for the backbone.

adj2coo(A)

Convert dense adjacency matrix to COO format.

Parameters:
  • A (Tensor) –

    Dense adjacency matrix, shape (n_nodes, n_nodes).

Returns:
  • tuple

    Contains:

    • torch.Tensor: Edge indices in COO format, shape (2, n_edges)
    • torch.Tensor: Edge weights, shape (n_edges,)
classwise_simple_mmd(source, target, src_y, tgt_y)

Compute class-wise Maximum Mean Discrepancy.

Parameters:
  • source (Tensor) –

    Source domain features.

  • target (Tensor) –

    Target domain features.

  • src_y (Tensor) –

    Source domain labels (one-hot).

  • tgt_y (Tensor) –

    Target domain labels (one-hot).

Returns:
  • float

    Sum of class-wise MMD values.

entropy(x, reduction='mean')

Compute entropy of probability distribution.

Parameters:
  • x (Tensor) –

    Probability distribution.

  • reduction (str, default: 'mean' ) –

    Reduction method. Default: 'mean'.

Returns:
  • Tensor

    Entropy value.

forward(x_s, edge_index_s, y_s, x_t, edge_index_t, y_t)

Forward pass of JHGDA model.

Parameters:
  • x_s (Tensor) –

    Source node features.

  • edge_index_s (Tensor) –

    Source edge indices.

  • y_s (Tensor) –

    Source labels.

  • x_t (Tensor) –

    Target node features.

  • edge_index_t (Tensor) –

    Target edge indices.

  • y_t (Tensor) –

    Target labels.

Returns:
  • tuple

    Contains:

    • embeddings: List of source/target embeddings
    • pred: Source/target predictions
    • pooling_loss: Dictionary of pooling losses
    • y: List of source/target labels
inference(data)

Perform inference on input data.

Parameters:
  • data (Data) –

    Input graph data containing: - x: Node features - edge_index: Edge indices

Returns:
  • Tensor

    Model predictions.

Notes

Simplified forward pass for inference:

  1. Single GNN layer
  2. Classification
label_matching(S, Y_old, Y_new)

Compute label consistency loss.

Parameters:
  • S (Tensor) –

    Assignment matrix.

  • Y_old (Tensor) –

    Original labels.

  • Y_new (Tensor) –

    New labels.

Returns:
  • Tensor

    Label matching loss value.

label_stable(S, Y_old, Y_new)

Compute label stability loss.

Parameters:
  • S (Tensor) –

    Assignment matrix.

  • Y_old (Tensor) –

    Original labels.

  • Y_new (Tensor) –

    New labels.

Returns:
  • Tensor

    Label stability loss value.

proximity_loss(A, S, adj_hop=1)

Compute graph structure preservation loss.

Parameters:
  • A (Tensor) –

    Original adjacency matrix.

  • S (Tensor) –

    Assignment matrix.

  • adj_hop (int, default: 1 ) –

    Number of hops. Default: 1.

Returns:
  • Tensor

    Proximity loss value.

pseudo_label(z_s, y_s, z_t, y_t, edge_index_t, edge_weight_t)

Generate pseudo-labels for target domain.

Parameters:
  • z_s (Tensor) –

    Source embeddings.

  • y_s (Tensor) –

    Source labels.

  • z_t (Tensor) –

    Target embeddings.

  • y_t (Tensor) –

    Target labels.

  • edge_index_t (Tensor) –

    Target edge indices.

  • edge_weight_t (Tensor) –

    Target edge weights.

Returns:
  • Tensor

    Pseudo-labels for target domain.

simple_mmd(source, target)

Compute simple Maximum Mean Discrepancy.

Parameters:
  • source (Tensor) –

    Source domain features.

  • target (Tensor) –

    Target domain features.

Returns:
  • Tensor

    L2 distance between mean feature vectors.

simple_mmd_kernel(source, target)

Compute kernel-based MMD with RBF kernel.

Parameters:
  • source (Tensor) –

    Source domain features.

  • target (Tensor) –

    Target domain features.

Returns:
  • Tensor

    RBF kernel value between mean feature vectors.

to_onehot(label_matrix, num_classes)

Convert label indices to one-hot encoding.

Parameters:
  • label_matrix (Tensor) –

    Label indices.

  • num_classes (int) –

    Number of classes.

Returns:
  • Tensor

    One-hot encoded labels.