High-Level Task APIs Reference
k3_node.tasks.base.BaseTask
Bases: K3NodeHubMixin
Abstract base task estimator providing common training, evaluation, and serialization workflows.
compile(optimizer=None, loss=None, metrics=None, **kwargs)
Configures the task model for training.
load(filepath, **kwargs)
classmethod
Loads a saved task model from disk.
save(filepath)
Saves the underlying model weights.
summary()
Prints a string summary of the underlying neural network.
k3_node.tasks.node_classification.NodeClassifier
Bases: BaseTask
High-level estimator for node classification tasks.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
backbone
|
Union[str, Model]
|
Model architecture string ( |
'gcn'
|
in_channels
|
int
|
Size of input node features. If not specified,
it is automatically inferred from the dataset during :meth: |
None
|
hidden_channels
|
int
|
Dimensionality of hidden node features.
(default: |
64
|
out_channels
|
int
|
Number of target classes. If not specified,
it is automatically inferred from the dataset during :meth: |
None
|
num_layers
|
int
|
Number of message passing layers. (default: |
2
|
dropout
|
float
|
Dropout probability. (default: |
0.5
|
multi_label
|
bool
|
If :obj: |
False
|
**backbone_kwargs
|
Additional arguments forwarded to the backbone constructor. |
{}
|
evaluate(data, mask='test_mask')
Evaluates classification accuracy on a given mask.
fit(data, epochs=20, lr=0.01, weight_decay=0.0005, mask='train_mask', val_mask='val_mask', verbose=1, callbacks=None)
Trains the node classifier on the provided graph data.
predict(data, mask=None)
Predicts discrete class labels for nodes.
predict_proba(data, mask=None)
Returns class probabilities for nodes.
k3_node.tasks.node_regression.NodeRegressor
Bases: BaseTask
High-level estimator for node regression tasks.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
backbone
|
Union[str, Model]
|
Model architecture string ( |
'gcn'
|
in_channels
|
int
|
Size of input node features. |
None
|
hidden_channels
|
int
|
Dimensionality of hidden node features. (default: |
64
|
out_channels
|
int
|
Number of continuous target variables. (default: |
1
|
num_layers
|
int
|
Number of message passing layers. (default: |
2
|
dropout
|
float
|
Dropout probability. (default: |
0.0
|
loss
|
str
|
Regression loss ( |
'mse'
|
**backbone_kwargs
|
Additional arguments forwarded to the backbone constructor. |
{}
|
evaluate(data, mask='test_mask')
Evaluates Mean Absolute Error and Mean Squared Error.
fit(data, epochs=20, lr=0.01, mask='train_mask', verbose=1, callbacks=None)
Trains the node regressor on the provided graph data.
predict(data, mask=None)
Returns continuous predictions for nodes.
k3_node.tasks.graph_classification.GraphClassifier
Bases: BaseTask
High-level estimator for graph classification tasks (e.g., molecular property, bioinformatics, social graph classification).
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
backbone
|
Union[str, Model]
|
GNN architecture ( |
'gin'
|
in_channels
|
int
|
Size of input node features. |
None
|
hidden_channels
|
int
|
Dimensionality of hidden node features. (default: |
64
|
num_classes
|
int
|
Number of graph classes. |
None
|
num_layers
|
int
|
Number of GNN layers. (default: |
3
|
pooling
|
str
|
Readout pooling ( |
'mean'
|
dropout
|
float
|
Dropout probability. (default: |
0.5
|
**backbone_kwargs
|
Additional arguments forwarded to the backbone constructor. |
{}
|
evaluate(dataset_or_loader, batch_size=32)
Evaluates classification accuracy on the dataset.
fit(dataset, epochs=20, lr=0.01, batch_size=32, shuffle=True, verbose=1, callbacks=None)
Trains the graph classifier.
predict(dataset_or_loader, batch_size=32)
Predicts discrete class labels for graphs.
predict_proba(dataset_or_loader, batch_size=32)
Predicts class probabilities for graphs.
k3_node.tasks.graph_regression.GraphRegressor
Bases: BaseTask
High-level estimator for graph and molecular property regression tasks.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
backbone
|
Union[str, Model]
|
Architecture ( |
'gin'
|
in_channels
|
int
|
Size of input node features. |
None
|
hidden_channels
|
int
|
Hidden feature dimension. (default: |
64
|
out_channels
|
int
|
Number of continuous target variables. (default: |
1
|
num_layers
|
int
|
Number of GNN layers. (default: |
3
|
pooling
|
str
|
Readout pooling ( |
'add'
|
loss
|
str
|
Loss function ( |
'mae'
|
**backbone_kwargs
|
Additional arguments forwarded to the backbone constructor. |
{}
|
evaluate(dataset_or_loader, batch_size=32)
Evaluates MAE on the dataset.
fit(dataset, epochs=20, lr=0.001, batch_size=32, shuffle=True, verbose=1, callbacks=None)
Trains the graph regressor.
predict(dataset_or_loader, batch_size=32)
Returns continuous predictions for graphs.
k3_node.tasks.link_prediction.LinkPredictor
Bases: BaseTask
High-level estimator for link prediction tasks on graphs.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
backbone
|
Union[str, Model]
|
Model architecture string ( |
'gcn'
|
in_channels
|
int
|
Size of input node features. If not specified,
it is automatically inferred from the dataset during :meth: |
None
|
hidden_channels
|
int
|
Dimensionality of hidden node features.
(default: |
64
|
out_channels
|
int
|
Dimensionality of output node embeddings
used for link scoring. (default: |
64
|
num_layers
|
int
|
Number of message passing layers. (default: |
2
|
decoder
|
str
|
Type of edge score decoder ( |
'inner_product'
|
dropout
|
float
|
Dropout probability. (default: |
0.0
|
**backbone_kwargs
|
Additional arguments forwarded to the backbone constructor. |
{}
|
encode(data)
Computes latent node representations for the input graph.
evaluate(data, edge_label_index=None, edge_label=None)
Evaluates link prediction performance (AUC, AP, Accuracy).
fit(data, edge_label_index=None, edge_label=None, epochs=20, lr=0.01, weight_decay=0.0, neg_ratio=1.0, verbose=1, callbacks=None)
Trains the link predictor on graph connectivity.
predict(data, edge_label_index=None, threshold=0.5)
Predicts binary link presence (0 or 1) for edge pairs.
predict_proba(data, edge_label_index=None)
Predicts link existence probabilities for edge pairs.