CWGCNBase(in_dim, hid_dim, num_classes, num_layers=2, dropout=0.1, act=F.relu, gnn='gcn', mode='node', **kwargs)

Bases: Module

Base class for CWGCN.

Parameters:
  • in_dim (int) –

    Input dimension of model.

  • hid_dim (int) –

    Hidden dimension of model.

  • num_classes (int) –

    Number of classes.

  • num_layers (int, default: 2 ) –

    Total number of layers in model. Default: 2.

  • dropout (float, default: 0.1 ) –

    Dropout rate. Default: 0.1.

  • act (callable, default: relu ) –

    Activation function. Default: torch.nn.functional.relu.

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

    The backbone GNN model type ('gcn', 'sage', 'gat', 'gin'). Default: 'gcn'.

  • mode (str, default: 'node' ) –

    Task mode ('node' or 'graph'). Default: 'node'.

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

    Additional parameters for the backbone.

Notes

Currently supports only 2-layer architectures with various GNN backbones.

c_loss(preds, labels)

Compute correntropy-induced loss.

Parameters:
  • preds (Tensor) –

    Model predictions.

  • labels (Tensor) –

    Ground truth labels.

Returns:
  • tuple

    Contains: - torch.Tensor: Correntropy loss value - torch.Tensor: Sample weights based on prediction-label distance

feat_bottleneck(x, edge_index, edge_weight=None, batch=None)

Feature extraction through GNN layers.

Parameters:
  • x (Tensor) –

    Node feature matrix.

  • edge_index (Tensor) –

    Edge indices.

  • edge_weight (Tensor, default: None ) –

    Edge weights. Default: None.

  • batch (Tensor, default: None ) –

    Batch vector for graph-level tasks. Default: None.

Returns:
  • tuple

    Contains: - torch.Tensor: Final bottleneck features - list: Layer-wise features for domain adaptation

feat_classifier(x, edge_index, edge_weight=None)

Classification layer after feature extraction.

Parameters:
  • x (Tensor) –

    Input features.

  • edge_index (Tensor) –

    Edge indices.

  • edge_weight (Tensor, default: None ) –

    Edge weights. Default: None.

Returns:
  • Tensor

    Classification logits.

forward(x, edge_index, edge_weight=None, batch=None)

Forward pass of the model.

Parameters:
  • x (Tensor) –

    Node feature matrix.

  • edge_index (Tensor) –

    Edge indices.

  • edge_weight (Tensor, default: None ) –

    Edge weights. Default: None.

  • batch (Tensor, default: None ) –

    Batch vector for graph-level tasks. Default: None.

Returns:
  • tuple

    Contains: - torch.Tensor: Final node/graph representations - list: Intermediate representations from each layer