Skip to content

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", "gat", "sage", "gin", "pna", "mlp", etc.) or a custom :class:keras.Model. (default: "gcn")

'gcn'
in_channels int

Size of input node features. If not specified, it is automatically inferred from the dataset during :meth:fit.

None
hidden_channels int

Dimensionality of hidden node features. (default: 64)

64
out_channels int

Number of target classes. If not specified, it is automatically inferred from the dataset during :meth:fit.

None
num_layers int

Number of message passing layers. (default: 2)

2
dropout float

Dropout probability. (default: 0.5)

0.5
multi_label bool

If :obj:True, treats the problem as multi-label binary classification using binary crossentropy. (default: False)

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", "gat", "sage", "gin", "pna", "mlp", etc.) or a custom :class:keras.Model. (default: "gcn")

'gcn'
in_channels int

Size of input node features.

None
hidden_channels int

Dimensionality of hidden node features. (default: 64)

64
out_channels int

Number of continuous target variables. (default: 1)

1
num_layers int

Number of message passing layers. (default: 2)

2
dropout float

Dropout probability. (default: 0.0)

0.0
loss str

Regression loss ("mse", "mae", or a Keras loss instance). (default: "mse")

'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", "gcn", "gat", "sage", "pna", etc.) or a custom :class:keras.Model. (default: "gin")

'gin'
in_channels int

Size of input node features.

None
hidden_channels int

Dimensionality of hidden node features. (default: 64)

64
num_classes int

Number of graph classes.

None
num_layers int

Number of GNN layers. (default: 3)

3
pooling str

Readout pooling ("mean", "add", "max"). (default: "mean")

'mean'
dropout float

Dropout probability. (default: 0.5)

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 ("schnet", "dimenet++", "attentive_fp", "pna", "gin", "gcn", etc.) or custom model. (default: "gin")

'gin'
in_channels int

Size of input node features.

None
hidden_channels int

Hidden feature dimension. (default: 64)

64
out_channels int

Number of continuous target variables. (default: 1)

1
num_layers int

Number of GNN layers. (default: 3)

3
pooling str

Readout pooling ("add", "mean", "max"). (default: "add")

'add'
loss str

Loss function ("mae" or "mse"). (default: "mae")

'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.

Bases: BaseTask

High-level estimator for link prediction tasks on graphs.

Parameters:

Name Type Description Default
backbone Union[str, Model]

Model architecture string ("gcn", "gat", "sage", "gin", "pna", "mlp", etc.) or custom :class:keras.Model. (default: "gcn")

'gcn'
in_channels int

Size of input node features. If not specified, it is automatically inferred from the dataset during :meth:fit.

None
hidden_channels int

Dimensionality of hidden node features. (default: 64)

64
out_channels int

Dimensionality of output node embeddings used for link scoring. (default: 64)

64
num_layers int

Number of message passing layers. (default: 2)

2
decoder str

Type of edge score decoder ("inner_product", "cosine", or "mlp"). (default: "inner_product")

'inner_product'
dropout float

Dropout probability. (default: 0.0)

0.0
**backbone_kwargs

Additional arguments forwarded to the backbone constructor.

{}

Computes latent node representations for the input graph.

Evaluates link prediction performance (AUC, AP, Accuracy).

Trains the link predictor on graph connectivity.

Predicts binary link presence (0 or 1) for edge pairs.

Predicts link existence probabilities for edge pairs.