k3-node 1.0.0__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- k3_node/__init__.py +122 -0
- k3_node/applications/__init__.py +17 -0
- k3_node/applications/bio/__init__.py +21 -0
- k3_node/applications/chemistry/__init__.py +155 -0
- k3_node/applications/materials/__init__.py +127 -0
- k3_node/applications/materials/basis.py +449 -0
- k3_node/applications/materials/chgnet.py +360 -0
- k3_node/applications/materials/core.py +351 -0
- k3_node/applications/materials/grace.py +246 -0
- k3_node/applications/materials/io.py +230 -0
- k3_node/applications/materials/m3gnet.py +462 -0
- k3_node/applications/materials/megnet.py +395 -0
- k3_node/applications/materials/qet.py +220 -0
- k3_node/applications/materials/readout.py +235 -0
- k3_node/applications/materials/so3net.py +234 -0
- k3_node/applications/materials/tensornet.py +381 -0
- k3_node/applications/materials/test_materials.py +167 -0
- k3_node/applications/materials/wrappers.py +95 -0
- k3_node/data/__init__.py +47 -0
- k3_node/data/batch.py +102 -0
- k3_node/data/collate.py +282 -0
- k3_node/data/data.py +532 -0
- k3_node/data/database.py +154 -0
- k3_node/data/dataset.py +182 -0
- k3_node/data/download.py +49 -0
- k3_node/data/extract.py +45 -0
- k3_node/data/feature_store.py +70 -0
- k3_node/data/graph_store.py +92 -0
- k3_node/data/hetero_data.py +374 -0
- k3_node/data/hypergraph_data.py +59 -0
- k3_node/data/in_memory_dataset.py +177 -0
- k3_node/data/makedirs.py +7 -0
- k3_node/data/on_disk_dataset.py +77 -0
- k3_node/data/separate.py +115 -0
- k3_node/data/storage.py +593 -0
- k3_node/data/temporal.py +154 -0
- k3_node/data/test_batch.py +67 -0
- k3_node/data/test_data.py +68 -0
- k3_node/data/test_dataset_and_stores.py +111 -0
- k3_node/data/test_hetero_data.py +33 -0
- k3_node/data/test_temporal_and_hyper.py +32 -0
- k3_node/data/view.py +43 -0
- k3_node/datasets/__init__.py +88 -0
- k3_node/datasets/actor.py +101 -0
- k3_node/datasets/airports.py +84 -0
- k3_node/datasets/amazon.py +66 -0
- k3_node/datasets/ba2motif_dataset.py +73 -0
- k3_node/datasets/ba_shapes.py +81 -0
- k3_node/datasets/bitcoin_otc.py +77 -0
- k3_node/datasets/citation_full.py +81 -0
- k3_node/datasets/coauthor.py +66 -0
- k3_node/datasets/dblp.py +106 -0
- k3_node/datasets/digits.py +63 -0
- k3_node/datasets/email_eu_core.py +60 -0
- k3_node/datasets/entities.py +158 -0
- k3_node/datasets/explainer_dataset.py +101 -0
- k3_node/datasets/facebook.py +51 -0
- k3_node/datasets/fake.py +256 -0
- k3_node/datasets/freebase.py +90 -0
- k3_node/datasets/geometric_shapes.py +69 -0
- k3_node/datasets/github.py +51 -0
- k3_node/datasets/graph_generator/__init__.py +6 -0
- k3_node/datasets/graph_generator/ba_graph.py +20 -0
- k3_node/datasets/graph_generator/base.py +29 -0
- k3_node/datasets/graph_generator/er_graph.py +21 -0
- k3_node/datasets/icews.py +58 -0
- k3_node/datasets/imdb.py +96 -0
- k3_node/datasets/jodie.py +56 -0
- k3_node/datasets/karate.py +56 -0
- k3_node/datasets/lastfm_asia.py +51 -0
- k3_node/datasets/mesh_correspondence.py +50 -0
- k3_node/datasets/molecule_net.py +148 -0
- k3_node/datasets/motif_generator/__init__.py +7 -0
- k3_node/datasets/motif_generator/base.py +29 -0
- k3_node/datasets/motif_generator/custom.py +17 -0
- k3_node/datasets/motif_generator/cycle.py +25 -0
- k3_node/datasets/motif_generator/house.py +27 -0
- k3_node/datasets/movielens.py +55 -0
- k3_node/datasets/planetoid.py +137 -0
- k3_node/datasets/polblogs.py +63 -0
- k3_node/datasets/ppi.py +189 -0
- k3_node/datasets/qm7.py +65 -0
- k3_node/datasets/qm9.py +132 -0
- k3_node/datasets/reddit.py +121 -0
- k3_node/datasets/sbm_dataset.py +165 -0
- k3_node/datasets/seal.py +74 -0
- k3_node/datasets/shape_scenes.py +92 -0
- k3_node/datasets/test_datasets.py +322 -0
- k3_node/datasets/tu_dataset.py +131 -0
- k3_node/datasets/twitch.py +66 -0
- k3_node/datasets/webkb.py +102 -0
- k3_node/datasets/wikics.py +85 -0
- k3_node/datasets/word_net.py +184 -0
- k3_node/etl/__init__.py +37 -0
- k3_node/etl/encoders.py +248 -0
- k3_node/etl/graph_builders.py +270 -0
- k3_node/etl/relational_to_graph.py +201 -0
- k3_node/etl/table_to_graph.py +244 -0
- k3_node/etl/test_etl.py +318 -0
- k3_node/export/__init__.py +15 -0
- k3_node/export/cross_backend.py +172 -0
- k3_node/export/onnx_exporter.py +190 -0
- k3_node/export/runtime.py +254 -0
- k3_node/export/tensorrt_exporter.py +201 -0
- k3_node/export/test_export.py +337 -0
- k3_node/export/tflite_exporter.py +112 -0
- k3_node/hub/__init__.py +29 -0
- k3_node/hub/dataset_hub.py +242 -0
- k3_node/hub/hub_mixin.py +599 -0
- k3_node/hub/model_card.py +133 -0
- k3_node/hub/test_hub.py +419 -0
- k3_node/io/__init__.py +22 -0
- k3_node/io/fs.py +117 -0
- k3_node/io/npz.py +45 -0
- k3_node/io/off.py +29 -0
- k3_node/io/planetoid.py +98 -0
- k3_node/io/tu.py +137 -0
- k3_node/io/txt_array.py +58 -0
- k3_node/layers/__init__.py +14 -0
- k3_node/layers/aggr/__init__.py +70 -0
- k3_node/layers/aggr/attention.py +77 -0
- k3_node/layers/aggr/base.py +403 -0
- k3_node/layers/aggr/basic.py +412 -0
- k3_node/layers/aggr/deep_sets.py +65 -0
- k3_node/layers/aggr/deepsets.py +29 -0
- k3_node/layers/aggr/equilibrium.py +107 -0
- k3_node/layers/aggr/fused.py +43 -0
- k3_node/layers/aggr/gmt.py +89 -0
- k3_node/layers/aggr/gru.py +58 -0
- k3_node/layers/aggr/lcm.py +143 -0
- k3_node/layers/aggr/lstm.py +58 -0
- k3_node/layers/aggr/mlp.py +75 -0
- k3_node/layers/aggr/multi.py +154 -0
- k3_node/layers/aggr/patch_transformer.py +137 -0
- k3_node/layers/aggr/quantile.py +125 -0
- k3_node/layers/aggr/resolver.py +68 -0
- k3_node/layers/aggr/scaler.py +133 -0
- k3_node/layers/aggr/set2set.py +87 -0
- k3_node/layers/aggr/set_transformer.py +107 -0
- k3_node/layers/aggr/sort.py +68 -0
- k3_node/layers/aggr/test_aggr.py +337 -0
- k3_node/layers/aggr/utils.py +210 -0
- k3_node/layers/aggr/variance_preserving.py +54 -0
- k3_node/layers/attention/__init__.py +5 -0
- k3_node/layers/attention/pair_attention.py +448 -0
- k3_node/layers/attention/performer.py +187 -0
- k3_node/layers/attention/polynormer.py +160 -0
- k3_node/layers/attention/qformer.py +143 -0
- k3_node/layers/attention/sgformer.py +106 -0
- k3_node/layers/attention/test_attention.py +68 -0
- k3_node/layers/attention/test_pair_attention.py +91 -0
- k3_node/layers/conv/__init__.py +149 -0
- k3_node/layers/conv/agnn_conv.py +120 -0
- k3_node/layers/conv/antisymmetric_conv.py +94 -0
- k3_node/layers/conv/appnp.py +105 -0
- k3_node/layers/conv/appnp_conv.py +157 -0
- k3_node/layers/conv/arma_conv.py +231 -0
- k3_node/layers/conv/cg_conv.py +92 -0
- k3_node/layers/conv/cheb_conv.py +137 -0
- k3_node/layers/conv/cluster_gcn_conv.py +102 -0
- k3_node/layers/conv/conv.py +100 -0
- k3_node/layers/conv/crystal_conv.py +140 -0
- k3_node/layers/conv/cugraph.py +84 -0
- k3_node/layers/conv/diffusion_conv.py +144 -0
- k3_node/layers/conv/dir_gnn_conv.py +93 -0
- k3_node/layers/conv/dna_conv.py +192 -0
- k3_node/layers/conv/edge_conv.py +107 -0
- k3_node/layers/conv/eg_conv.py +155 -0
- k3_node/layers/conv/fa_conv.py +107 -0
- k3_node/layers/conv/feast_conv.py +126 -0
- k3_node/layers/conv/film_conv.py +143 -0
- k3_node/layers/conv/gat_conv.py +244 -0
- k3_node/layers/conv/gated_graph_conv.py +136 -0
- k3_node/layers/conv/gatv2_conv.py +205 -0
- k3_node/layers/conv/gcn.py +144 -0
- k3_node/layers/conv/gcn2_conv.py +126 -0
- k3_node/layers/conv/gcn_conv.py +135 -0
- k3_node/layers/conv/gen_conv.py +163 -0
- k3_node/layers/conv/general_conv.py +218 -0
- k3_node/layers/conv/gin_conv.py +218 -0
- k3_node/layers/conv/gmm_conv.py +172 -0
- k3_node/layers/conv/gps_conv.py +153 -0
- k3_node/layers/conv/graph_attention.py +262 -0
- k3_node/layers/conv/graph_conv.py +84 -0
- k3_node/layers/conv/gravnet_conv.py +93 -0
- k3_node/layers/conv/han_conv.py +175 -0
- k3_node/layers/conv/heat_conv.py +131 -0
- k3_node/layers/conv/hetero_conv.py +128 -0
- k3_node/layers/conv/hgt_conv.py +218 -0
- k3_node/layers/conv/hypergraph_conv.py +182 -0
- k3_node/layers/conv/le_conv.py +81 -0
- k3_node/layers/conv/lg_conv.py +58 -0
- k3_node/layers/conv/meshcnn_conv.py +84 -0
- k3_node/layers/conv/message_passing.py +451 -0
- k3_node/layers/conv/mf_conv.py +95 -0
- k3_node/layers/conv/mixhop_conv.py +108 -0
- k3_node/layers/conv/nn_conv.py +110 -0
- k3_node/layers/conv/pan_conv.py +100 -0
- k3_node/layers/conv/pdn_conv.py +109 -0
- k3_node/layers/conv/pna_conv.py +177 -0
- k3_node/layers/conv/point_conv.py +101 -0
- k3_node/layers/conv/point_gnn_conv.py +90 -0
- k3_node/layers/conv/point_transformer_conv.py +132 -0
- k3_node/layers/conv/ppf_conv.py +135 -0
- k3_node/layers/conv/ppnp.py +89 -0
- k3_node/layers/conv/res_gated_graph_conv.py +126 -0
- k3_node/layers/conv/rgat_conv.py +251 -0
- k3_node/layers/conv/rgcn_conv.py +321 -0
- k3_node/layers/conv/sage_conv.py +154 -0
- k3_node/layers/conv/sg_conv.py +96 -0
- k3_node/layers/conv/signed_conv.py +100 -0
- k3_node/layers/conv/simple_conv.py +75 -0
- k3_node/layers/conv/spline_conv.py +182 -0
- k3_node/layers/conv/ssg_conv.py +101 -0
- k3_node/layers/conv/supergat_conv.py +195 -0
- k3_node/layers/conv/tag_conv.py +98 -0
- k3_node/layers/conv/test_backend_consistency.py +164 -0
- k3_node/layers/conv/test_conv.py +176 -0
- k3_node/layers/conv/test_conv_pyg.py +566 -0
- k3_node/layers/conv/transformer_conv.py +168 -0
- k3_node/layers/conv/utils.py +403 -0
- k3_node/layers/conv/wl_conv.py +151 -0
- k3_node/layers/conv/x_conv.py +187 -0
- k3_node/layers/dense/__init__.py +40 -0
- k3_node/layers/dense/dense_gat_conv.py +149 -0
- k3_node/layers/dense/dense_gcn_conv.py +117 -0
- k3_node/layers/dense/dense_gin_conv.py +88 -0
- k3_node/layers/dense/dense_graph_conv.py +95 -0
- k3_node/layers/dense/dense_sage_conv.py +85 -0
- k3_node/layers/dense/diff_pool.py +76 -0
- k3_node/layers/dense/dmon_pool.py +223 -0
- k3_node/layers/dense/linear.py +327 -0
- k3_node/layers/dense/mincut_pool.py +92 -0
- k3_node/layers/dense/test_dense.py +377 -0
- k3_node/layers/functional/__init__.py +13 -0
- k3_node/layers/functional/bro.py +49 -0
- k3_node/layers/functional/edge_dropout.py +55 -0
- k3_node/layers/functional/gini.py +44 -0
- k3_node/layers/functional/test_functional.py +34 -0
- k3_node/layers/kge/__init__.py +17 -0
- k3_node/layers/kge/base.py +255 -0
- k3_node/layers/kge/complex.py +98 -0
- k3_node/layers/kge/distmult.py +79 -0
- k3_node/layers/kge/loader.py +50 -0
- k3_node/layers/kge/rotate.py +103 -0
- k3_node/layers/kge/test_kge.py +76 -0
- k3_node/layers/kge/transe.py +96 -0
- k3_node/layers/norm/__init__.py +23 -0
- k3_node/layers/norm/batch_norm.py +328 -0
- k3_node/layers/norm/diff_group_norm.py +141 -0
- k3_node/layers/norm/graph_norm.py +105 -0
- k3_node/layers/norm/graph_size_norm.py +57 -0
- k3_node/layers/norm/instance_norm.py +163 -0
- k3_node/layers/norm/layer_norm.py +245 -0
- k3_node/layers/norm/mean_subtraction_norm.py +57 -0
- k3_node/layers/norm/msg_norm.py +58 -0
- k3_node/layers/norm/pair_norm.py +94 -0
- k3_node/layers/norm/test_norm.py +275 -0
- k3_node/layers/pool/__init__.py +83 -0
- k3_node/layers/pool/approx_knn.py +101 -0
- k3_node/layers/pool/asap.py +173 -0
- k3_node/layers/pool/avg_pool.py +165 -0
- k3_node/layers/pool/cluster_pool.py +168 -0
- k3_node/layers/pool/connect/__init__.py +10 -0
- k3_node/layers/pool/connect/base.py +103 -0
- k3_node/layers/pool/connect/filter_edges.py +113 -0
- k3_node/layers/pool/consecutive.py +30 -0
- k3_node/layers/pool/decimation.py +48 -0
- k3_node/layers/pool/edge_pool.py +189 -0
- k3_node/layers/pool/glob.py +139 -0
- k3_node/layers/pool/graclus.py +66 -0
- k3_node/layers/pool/knn.py +253 -0
- k3_node/layers/pool/max_pool.py +159 -0
- k3_node/layers/pool/mem_pool.py +145 -0
- k3_node/layers/pool/pan_pool.py +144 -0
- k3_node/layers/pool/point_cloud.py +212 -0
- k3_node/layers/pool/pool.py +119 -0
- k3_node/layers/pool/sag_pool.py +174 -0
- k3_node/layers/pool/select/__init__.py +10 -0
- k3_node/layers/pool/select/base.py +112 -0
- k3_node/layers/pool/select/topk.py +206 -0
- k3_node/layers/pool/test_pool.py +456 -0
- k3_node/layers/pool/topk_pool.py +103 -0
- k3_node/layers/pool/voxel_grid.py +70 -0
- k3_node/layers/unpool/__init__.py +9 -0
- k3_node/layers/unpool/knn_interpolate.py +57 -0
- k3_node/layers/unpool/test_unpool.py +31 -0
- k3_node/loader/__init__.py +62 -0
- k3_node/loader/base.py +69 -0
- k3_node/loader/cache.py +68 -0
- k3_node/loader/cluster.py +127 -0
- k3_node/loader/data_list_loader.py +45 -0
- k3_node/loader/dataloader.py +117 -0
- k3_node/loader/dense_data_loader.py +62 -0
- k3_node/loader/dynamic_batch_sampler.py +93 -0
- k3_node/loader/graph_saint.py +188 -0
- k3_node/loader/hgt_loader.py +90 -0
- k3_node/loader/imbalanced_sampler.py +87 -0
- k3_node/loader/keras_dataset.py +334 -0
- k3_node/loader/link_loader.py +179 -0
- k3_node/loader/link_neighbor_loader.py +202 -0
- k3_node/loader/mixin.py +190 -0
- k3_node/loader/neighbor_loader.py +159 -0
- k3_node/loader/neighbor_sampler.py +167 -0
- k3_node/loader/node_loader.py +185 -0
- k3_node/loader/prefetch.py +115 -0
- k3_node/loader/random_node_loader.py +89 -0
- k3_node/loader/sampler_utils.py +499 -0
- k3_node/loader/shadow.py +115 -0
- k3_node/loader/temporal_dataloader.py +98 -0
- k3_node/loader/test_dataloader.py +113 -0
- k3_node/loader/test_keras_dataset.py +221 -0
- k3_node/loader/test_neighbor_loader.py +122 -0
- k3_node/loader/test_sampler_utils.py +82 -0
- k3_node/loader/test_samplers.py +96 -0
- k3_node/loader/test_subgraph_loaders.py +89 -0
- k3_node/loader/utils.py +232 -0
- k3_node/loader/zip_loader.py +88 -0
- k3_node/metrics.py +94 -0
- k3_node/models/__init__.py +424 -0
- k3_node/models/attentive_fp.py +232 -0
- k3_node/models/attract_repel.py +108 -0
- k3_node/models/autoencoder.py +318 -0
- k3_node/models/basic_gnn.py +443 -0
- k3_node/models/bio/__init__.py +4 -0
- k3_node/models/captum.py +52 -0
- k3_node/models/chemistry/__init__.py +4 -0
- k3_node/models/correct_and_smooth.py +146 -0
- k3_node/models/deep_graph_infomax.py +113 -0
- k3_node/models/deepgcn.py +121 -0
- k3_node/models/dimenet.py +737 -0
- k3_node/models/dimenet_utils.py +153 -0
- k3_node/models/gnnff.py +263 -0
- k3_node/models/gps_model.py +1122 -0
- k3_node/models/gpse.py +638 -0
- k3_node/models/graph_unet.py +199 -0
- k3_node/models/graphmae2.py +954 -0
- k3_node/models/graphormer.py +1258 -0
- k3_node/models/graphormer_3d.py +868 -0
- k3_node/models/grover.py +1066 -0
- k3_node/models/jumping_knowledge.py +200 -0
- k3_node/models/label_prop.py +110 -0
- k3_node/models/lightgcn.py +171 -0
- k3_node/models/linkx.py +181 -0
- k3_node/models/lpformer.py +404 -0
- k3_node/models/mask_label.py +114 -0
- k3_node/models/materials/__init__.py +33 -0
- k3_node/models/meta.py +133 -0
- k3_node/models/metapath2vec.py +234 -0
- k3_node/models/mlp.py +264 -0
- k3_node/models/mole_bert.py +379 -0
- k3_node/models/neural_fingerprint.py +95 -0
- k3_node/models/node2vec.py +213 -0
- k3_node/models/pmlp.py +157 -0
- k3_node/models/polynormer.py +229 -0
- k3_node/models/rect.py +93 -0
- k3_node/models/renet.py +221 -0
- k3_node/models/rev_gnn.py +128 -0
- k3_node/models/schnet.py +484 -0
- k3_node/models/sgformer.py +195 -0
- k3_node/models/signed_gcn.py +185 -0
- k3_node/models/test_attentive_fp.py +32 -0
- k3_node/models/test_attract_repel.py +33 -0
- k3_node/models/test_autoencoder.py +119 -0
- k3_node/models/test_basic_gnn.py +102 -0
- k3_node/models/test_correct_and_smooth.py +40 -0
- k3_node/models/test_deep_graph_infomax.py +68 -0
- k3_node/models/test_deepgcn.py +21 -0
- k3_node/models/test_dimenet.py +86 -0
- k3_node/models/test_domain_apis.py +138 -0
- k3_node/models/test_gnnff.py +24 -0
- k3_node/models/test_gps_model.py +271 -0
- k3_node/models/test_gpse.py +34 -0
- k3_node/models/test_graph_unet.py +26 -0
- k3_node/models/test_graphmae2.py +226 -0
- k3_node/models/test_graphormer.py +233 -0
- k3_node/models/test_graphormer3d.py +163 -0
- k3_node/models/test_grover.py +287 -0
- k3_node/models/test_jumping_knowledge.py +129 -0
- k3_node/models/test_label_prop.py +37 -0
- k3_node/models/test_lightgcn.py +38 -0
- k3_node/models/test_linkx.py +31 -0
- k3_node/models/test_lpformer.py +22 -0
- k3_node/models/test_mask_label.py +90 -0
- k3_node/models/test_meta.py +159 -0
- k3_node/models/test_metapath2vec.py +45 -0
- k3_node/models/test_mlp.py +62 -0
- k3_node/models/test_mole_bert.py +164 -0
- k3_node/models/test_neural_fingerprint.py +13 -0
- k3_node/models/test_node2vec.py +57 -0
- k3_node/models/test_pmlp.py +81 -0
- k3_node/models/test_polynormer.py +104 -0
- k3_node/models/test_rect.py +23 -0
- k3_node/models/test_renet.py +32 -0
- k3_node/models/test_rev_gnn.py +24 -0
- k3_node/models/test_schnet.py +43 -0
- k3_node/models/test_sgformer.py +48 -0
- k3_node/models/test_signed_gcn.py +28 -0
- k3_node/models/test_tgn.py +77 -0
- k3_node/models/test_unimol.py +179 -0
- k3_node/models/test_unimol2.py +114 -0
- k3_node/models/test_unimol_plus.py +131 -0
- k3_node/models/test_visnet.py +44 -0
- k3_node/models/tgn.py +382 -0
- k3_node/models/unimol.py +1156 -0
- k3_node/models/unimol2.py +616 -0
- k3_node/models/unimol_docking_v2.py +301 -0
- k3_node/models/unimol_plus.py +456 -0
- k3_node/models/utils.py +97 -0
- k3_node/models/visnet.py +759 -0
- k3_node/ops/__init__.py +4 -0
- k3_node/ops/conv.py +56 -0
- k3_node/ops/creation.py +43 -0
- k3_node/ops/graph.py +27 -0
- k3_node/ops/host.py +41 -0
- k3_node/ops/matmul.py +49 -0
- k3_node/ops/numpy.py +24 -0
- k3_node/ops/segment.py +54 -0
- k3_node/ops/sparse.py +51 -0
- k3_node/rag/__init__.py +49 -0
- k3_node/rag/encoders.py +312 -0
- k3_node/rag/pipeline.py +192 -0
- k3_node/rag/projector.py +184 -0
- k3_node/rag/subgraph.py +270 -0
- k3_node/rag/test_rag.py +347 -0
- k3_node/rag/verbalizer.py +162 -0
- k3_node/tasks/__init__.py +19 -0
- k3_node/tasks/backbone_resolver.py +125 -0
- k3_node/tasks/base.py +67 -0
- k3_node/tasks/graph_classification.py +270 -0
- k3_node/tasks/graph_regression.py +228 -0
- k3_node/tasks/link_prediction.py +306 -0
- k3_node/tasks/node_classification.py +194 -0
- k3_node/tasks/node_regression.py +138 -0
- k3_node/tasks/test_tasks.py +319 -0
- k3_node/test_docstring_examples.py +106 -0
- k3_node/test_training_forwarding.py +116 -0
- k3_node/training.py +115 -0
- k3_node/transforms/__init__.py +166 -0
- k3_node/transforms/base_transform.py +32 -0
- k3_node/transforms/compose.py +58 -0
- k3_node/transforms/general.py +676 -0
- k3_node/transforms/graph.py +1070 -0
- k3_node/transforms/spatial.py +797 -0
- k3_node/transforms/test_random_link_split.py +45 -0
- k3_node/transforms/test_spatial_transforms.py +65 -0
- k3_node/transforms/test_transforms.py +253 -0
- k3_node/transforms/utils.py +102 -0
- k3_node/utils/__init__.py +5 -0
- k3_node/utils/backend_import.py +12 -0
- k3_node/utils/graph.py +286 -0
- k3_node/utils/keras.py +94 -0
- k3_node/utils/random.py +103 -0
- k3_node/utils/smiles.py +235 -0
- k3_node-1.0.0.dist-info/METADATA +284 -0
- k3_node-1.0.0.dist-info/RECORD +459 -0
- k3_node-1.0.0.dist-info/WHEEL +5 -0
- k3_node-1.0.0.dist-info/licenses/LICENSE +21 -0
- k3_node-1.0.0.dist-info/top_level.txt +1 -0
|
@@ -0,0 +1,194 @@
|
|
|
1
|
+
"""High-level Node Classification Task."""
|
|
2
|
+
|
|
3
|
+
from typing import Any, Dict, List, Optional, Union
|
|
4
|
+
import keras
|
|
5
|
+
from keras import ops
|
|
6
|
+
|
|
7
|
+
from k3_node.tasks.base import BaseTask
|
|
8
|
+
from k3_node.tasks.backbone_resolver import resolve_backbone
|
|
9
|
+
|
|
10
|
+
|
|
11
|
+
class NodeClassifier(BaseTask):
|
|
12
|
+
r"""High-level estimator for node classification tasks.
|
|
13
|
+
|
|
14
|
+
Args:
|
|
15
|
+
backbone: Model architecture string (``"gcn"``, ``"gat"``, ``"sage"``,
|
|
16
|
+
``"gin"``, ``"pna"``, ``"mlp"``, etc.) or a custom :class:`keras.Model`.
|
|
17
|
+
(default: ``"gcn"``)
|
|
18
|
+
in_channels (int, optional): Size of input node features. If not specified,
|
|
19
|
+
it is automatically inferred from the dataset during :meth:`fit`.
|
|
20
|
+
hidden_channels (int, optional): Dimensionality of hidden node features.
|
|
21
|
+
(default: ``64``)
|
|
22
|
+
out_channels (int, optional): Number of target classes. If not specified,
|
|
23
|
+
it is automatically inferred from the dataset during :meth:`fit`.
|
|
24
|
+
num_layers (int, optional): Number of message passing layers. (default: ``2``)
|
|
25
|
+
dropout (float, optional): Dropout probability. (default: ``0.5``)
|
|
26
|
+
multi_label (bool, optional): If :obj:`True`, treats the problem as multi-label
|
|
27
|
+
binary classification using binary crossentropy. (default: ``False``)
|
|
28
|
+
**backbone_kwargs: Additional arguments forwarded to the backbone constructor.
|
|
29
|
+
"""
|
|
30
|
+
|
|
31
|
+
def __init__(
|
|
32
|
+
self,
|
|
33
|
+
backbone: Union[str, keras.Model] = "gcn",
|
|
34
|
+
in_channels: Optional[int] = None,
|
|
35
|
+
hidden_channels: int = 64,
|
|
36
|
+
out_channels: Optional[int] = None,
|
|
37
|
+
num_classes: Optional[int] = None,
|
|
38
|
+
num_layers: int = 2,
|
|
39
|
+
dropout: float = 0.5,
|
|
40
|
+
multi_label: bool = False,
|
|
41
|
+
**backbone_kwargs,
|
|
42
|
+
):
|
|
43
|
+
super().__init__()
|
|
44
|
+
self.backbone = backbone
|
|
45
|
+
self.in_channels = in_channels
|
|
46
|
+
self.hidden_channels = hidden_channels
|
|
47
|
+
self.out_channels = out_channels if out_channels is not None else num_classes
|
|
48
|
+
self.num_layers = num_layers
|
|
49
|
+
self.dropout = dropout
|
|
50
|
+
self.multi_label = multi_label
|
|
51
|
+
self.backbone_kwargs = backbone_kwargs
|
|
52
|
+
|
|
53
|
+
if isinstance(backbone, keras.Model):
|
|
54
|
+
self.model = backbone
|
|
55
|
+
|
|
56
|
+
def _init_model(self, data: Any):
|
|
57
|
+
r"""Infers missing dimensions and instantiates the backbone model."""
|
|
58
|
+
in_c = self.in_channels
|
|
59
|
+
if in_c is None:
|
|
60
|
+
if hasattr(data, "num_node_features") and data.num_node_features > 0:
|
|
61
|
+
in_c = data.num_node_features
|
|
62
|
+
elif hasattr(data, "num_features") and data.num_features > 0:
|
|
63
|
+
in_c = data.num_features
|
|
64
|
+
elif hasattr(data, "x") and data.x is not None:
|
|
65
|
+
in_c = int(ops.shape(data.x)[-1])
|
|
66
|
+
else:
|
|
67
|
+
raise ValueError("Could not automatically infer in_channels from data. Please specify in_channels.")
|
|
68
|
+
|
|
69
|
+
out_c = self.out_channels
|
|
70
|
+
if out_c is None:
|
|
71
|
+
if hasattr(data, "num_classes") and data.num_classes is not None:
|
|
72
|
+
out_c = data.num_classes
|
|
73
|
+
elif hasattr(data, "y") and data.y is not None:
|
|
74
|
+
y = data.y
|
|
75
|
+
if self.multi_label:
|
|
76
|
+
out_c = int(ops.shape(y)[-1])
|
|
77
|
+
else:
|
|
78
|
+
out_c = int(ops.convert_to_numpy(ops.max(y))) + 1
|
|
79
|
+
else:
|
|
80
|
+
raise ValueError("Could not automatically infer out_channels from data. Please specify out_channels.")
|
|
81
|
+
|
|
82
|
+
if not self.multi_label:
|
|
83
|
+
out_c = max(int(out_c), 2)
|
|
84
|
+
|
|
85
|
+
self.in_channels = in_c
|
|
86
|
+
self.out_channels = out_c
|
|
87
|
+
|
|
88
|
+
self.model = resolve_backbone(
|
|
89
|
+
self.backbone,
|
|
90
|
+
in_channels=in_c,
|
|
91
|
+
out_channels=out_c,
|
|
92
|
+
hidden_channels=self.hidden_channels,
|
|
93
|
+
num_layers=self.num_layers,
|
|
94
|
+
dropout=self.dropout,
|
|
95
|
+
**self.backbone_kwargs,
|
|
96
|
+
)
|
|
97
|
+
|
|
98
|
+
def fit(
|
|
99
|
+
self,
|
|
100
|
+
data: Any,
|
|
101
|
+
epochs: int = 20,
|
|
102
|
+
lr: float = 0.01,
|
|
103
|
+
weight_decay: float = 5e-4,
|
|
104
|
+
mask: Optional[str] = "train_mask",
|
|
105
|
+
val_mask: Optional[str] = "val_mask",
|
|
106
|
+
verbose: int = 1,
|
|
107
|
+
callbacks: Optional[List[Any]] = None,
|
|
108
|
+
):
|
|
109
|
+
r"""Trains the node classifier on the provided graph data."""
|
|
110
|
+
if self.model is None:
|
|
111
|
+
self._init_model(data)
|
|
112
|
+
|
|
113
|
+
if not self._is_compiled:
|
|
114
|
+
opt = keras.optimizers.Adam(learning_rate=lr, weight_decay=weight_decay)
|
|
115
|
+
if self.multi_label:
|
|
116
|
+
loss = keras.losses.BinaryCrossentropy(from_logits=True)
|
|
117
|
+
metrics = [keras.metrics.BinaryAccuracy(name="acc")]
|
|
118
|
+
else:
|
|
119
|
+
loss = keras.losses.SparseCategoricalCrossentropy(from_logits=True)
|
|
120
|
+
metrics = [keras.metrics.SparseCategoricalAccuracy(name="acc")]
|
|
121
|
+
|
|
122
|
+
self.model.compile(
|
|
123
|
+
optimizer=opt,
|
|
124
|
+
loss=loss,
|
|
125
|
+
weighted_metrics=metrics,
|
|
126
|
+
)
|
|
127
|
+
self._is_compiled = True
|
|
128
|
+
|
|
129
|
+
# Generate training batches
|
|
130
|
+
if hasattr(data, "to_generator"):
|
|
131
|
+
gen = data.to_generator(mask=mask)
|
|
132
|
+
else:
|
|
133
|
+
inputs = self._extract_inputs(data)
|
|
134
|
+
y = ops.convert_to_tensor(data.y)
|
|
135
|
+
m = getattr(data, mask) if mask and hasattr(data, mask) else None
|
|
136
|
+
sample_weight = ops.cast(m, "float32") if m is not None else None
|
|
137
|
+
if sample_weight is not None: # the loss is the mean over the masked nodes, as in PyG
|
|
138
|
+
sample_weight = sample_weight * (ops.cast(ops.size(sample_weight), "float32") / ops.maximum(ops.sum(sample_weight), 1.0))
|
|
139
|
+
|
|
140
|
+
def gen_fn():
|
|
141
|
+
while True:
|
|
142
|
+
if sample_weight is not None:
|
|
143
|
+
yield inputs, y, sample_weight
|
|
144
|
+
else:
|
|
145
|
+
yield inputs, y
|
|
146
|
+
|
|
147
|
+
gen = gen_fn()
|
|
148
|
+
|
|
149
|
+
return self.model.fit(
|
|
150
|
+
gen,
|
|
151
|
+
steps_per_epoch=1,
|
|
152
|
+
epochs=epochs,
|
|
153
|
+
shuffle=False, # one full-graph batch per step
|
|
154
|
+
verbose=verbose,
|
|
155
|
+
callbacks=callbacks,
|
|
156
|
+
)
|
|
157
|
+
|
|
158
|
+
def predict_proba(self, data: Any, mask: Optional[str] = None):
|
|
159
|
+
r"""Returns class probabilities for nodes."""
|
|
160
|
+
if self.model is None:
|
|
161
|
+
raise RuntimeError("Model is not initialized. Fit or load a model first.")
|
|
162
|
+
|
|
163
|
+
inputs = self._extract_inputs(data)
|
|
164
|
+
logits = self.model(inputs, training=False)
|
|
165
|
+
|
|
166
|
+
if self.multi_label:
|
|
167
|
+
probs = ops.sigmoid(logits)
|
|
168
|
+
else:
|
|
169
|
+
probs = ops.softmax(logits, axis=-1)
|
|
170
|
+
|
|
171
|
+
if mask is not None and hasattr(data, mask):
|
|
172
|
+
m = getattr(data, mask)
|
|
173
|
+
probs = probs[m]
|
|
174
|
+
return probs
|
|
175
|
+
|
|
176
|
+
def predict(self, data: Any, mask: Optional[str] = None):
|
|
177
|
+
r"""Predicts discrete class labels for nodes."""
|
|
178
|
+
probs = self.predict_proba(data, mask=mask)
|
|
179
|
+
if self.multi_label:
|
|
180
|
+
return ops.cast(probs > 0.5, "int64")
|
|
181
|
+
return ops.argmax(probs, axis=-1)
|
|
182
|
+
|
|
183
|
+
def evaluate(self, data: Any, mask: Optional[str] = "test_mask") -> Dict[str, float]:
|
|
184
|
+
r"""Evaluates classification accuracy on a given mask."""
|
|
185
|
+
pred = self.predict(data, mask=mask)
|
|
186
|
+
y = data.y
|
|
187
|
+
if mask is not None and hasattr(data, mask):
|
|
188
|
+
m = getattr(data, mask)
|
|
189
|
+
y = y[m]
|
|
190
|
+
|
|
191
|
+
y_cast = ops.cast(y, "int64")
|
|
192
|
+
pred_cast = ops.cast(pred, "int64")
|
|
193
|
+
acc = float(ops.convert_to_numpy(ops.mean(ops.cast(pred_cast == y_cast, "float32"))))
|
|
194
|
+
return {"accuracy": acc}
|
|
@@ -0,0 +1,138 @@
|
|
|
1
|
+
"""High-level Node Regression Task."""
|
|
2
|
+
|
|
3
|
+
from typing import Any, Dict, List, Optional, Union
|
|
4
|
+
import keras
|
|
5
|
+
from keras import ops
|
|
6
|
+
|
|
7
|
+
from k3_node.tasks.base import BaseTask
|
|
8
|
+
from k3_node.tasks.backbone_resolver import resolve_backbone
|
|
9
|
+
|
|
10
|
+
|
|
11
|
+
class NodeRegressor(BaseTask):
|
|
12
|
+
r"""High-level estimator for node regression tasks.
|
|
13
|
+
|
|
14
|
+
Args:
|
|
15
|
+
backbone: Model architecture string (``"gcn"``, ``"gat"``, ``"sage"``,
|
|
16
|
+
``"gin"``, ``"pna"``, ``"mlp"``, etc.) or a custom :class:`keras.Model`.
|
|
17
|
+
(default: ``"gcn"``)
|
|
18
|
+
in_channels (int, optional): Size of input node features.
|
|
19
|
+
hidden_channels (int, optional): Dimensionality of hidden node features. (default: ``64``)
|
|
20
|
+
out_channels (int, optional): Number of continuous target variables. (default: ``1``)
|
|
21
|
+
num_layers (int, optional): Number of message passing layers. (default: ``2``)
|
|
22
|
+
dropout (float, optional): Dropout probability. (default: ``0.0``)
|
|
23
|
+
loss: Regression loss (``"mse"``, ``"mae"``, or a Keras loss instance). (default: ``"mse"``)
|
|
24
|
+
**backbone_kwargs: Additional arguments forwarded to the backbone constructor.
|
|
25
|
+
"""
|
|
26
|
+
|
|
27
|
+
def __init__(
|
|
28
|
+
self,
|
|
29
|
+
backbone: Union[str, keras.Model] = "gcn",
|
|
30
|
+
in_channels: Optional[int] = None,
|
|
31
|
+
hidden_channels: int = 64,
|
|
32
|
+
out_channels: int = 1,
|
|
33
|
+
num_layers: int = 2,
|
|
34
|
+
dropout: float = 0.0,
|
|
35
|
+
loss: str = "mse",
|
|
36
|
+
**backbone_kwargs,
|
|
37
|
+
):
|
|
38
|
+
super().__init__()
|
|
39
|
+
self.backbone = backbone
|
|
40
|
+
self.in_channels = in_channels
|
|
41
|
+
self.hidden_channels = hidden_channels
|
|
42
|
+
self.out_channels = out_channels
|
|
43
|
+
self.num_layers = num_layers
|
|
44
|
+
self.dropout = dropout
|
|
45
|
+
self.loss_name = loss
|
|
46
|
+
self.backbone_kwargs = backbone_kwargs
|
|
47
|
+
|
|
48
|
+
if isinstance(backbone, keras.Model):
|
|
49
|
+
self.model = backbone
|
|
50
|
+
|
|
51
|
+
def _init_model(self, data: Any):
|
|
52
|
+
in_c = self.in_channels
|
|
53
|
+
if in_c is None:
|
|
54
|
+
if hasattr(data, "num_node_features") and data.num_node_features > 0:
|
|
55
|
+
in_c = data.num_node_features
|
|
56
|
+
elif hasattr(data, "num_features") and data.num_features > 0:
|
|
57
|
+
in_c = data.num_features
|
|
58
|
+
elif hasattr(data, "x") and data.x is not None:
|
|
59
|
+
in_c = int(ops.shape(data.x)[-1])
|
|
60
|
+
else:
|
|
61
|
+
raise ValueError("Could not automatically infer in_channels from data.")
|
|
62
|
+
|
|
63
|
+
self.in_channels = in_c
|
|
64
|
+
self.model = resolve_backbone(
|
|
65
|
+
self.backbone,
|
|
66
|
+
in_channels=in_c,
|
|
67
|
+
out_channels=self.out_channels,
|
|
68
|
+
hidden_channels=self.hidden_channels,
|
|
69
|
+
num_layers=self.num_layers,
|
|
70
|
+
dropout=self.dropout,
|
|
71
|
+
**self.backbone_kwargs,
|
|
72
|
+
)
|
|
73
|
+
|
|
74
|
+
def fit(
|
|
75
|
+
self,
|
|
76
|
+
data: Any,
|
|
77
|
+
epochs: int = 20,
|
|
78
|
+
lr: float = 0.01,
|
|
79
|
+
mask: Optional[str] = "train_mask",
|
|
80
|
+
verbose: int = 1,
|
|
81
|
+
callbacks: Optional[List[Any]] = None,
|
|
82
|
+
):
|
|
83
|
+
r"""Trains the node regressor on the provided graph data."""
|
|
84
|
+
if self.model is None:
|
|
85
|
+
self._init_model(data)
|
|
86
|
+
|
|
87
|
+
if not self._is_compiled:
|
|
88
|
+
loss = keras.losses.MeanSquaredError() if self.loss_name == "mse" else keras.losses.MeanAbsoluteError()
|
|
89
|
+
self.model.compile(
|
|
90
|
+
optimizer=keras.optimizers.Adam(learning_rate=lr),
|
|
91
|
+
loss=loss,
|
|
92
|
+
weighted_metrics=[keras.metrics.MeanAbsoluteError(name="mae")],
|
|
93
|
+
)
|
|
94
|
+
self._is_compiled = True
|
|
95
|
+
|
|
96
|
+
inputs = self._extract_inputs(data)
|
|
97
|
+
y = ops.cast(data.y, "float32")
|
|
98
|
+
m = getattr(data, mask) if mask and hasattr(data, mask) else None
|
|
99
|
+
sample_weight = ops.cast(m, "float32") if m is not None else None
|
|
100
|
+
if sample_weight is not None: # the loss is the mean over the masked nodes, as in PyG
|
|
101
|
+
sample_weight = sample_weight * (ops.cast(ops.size(sample_weight), "float32") / ops.maximum(ops.sum(sample_weight), 1.0))
|
|
102
|
+
|
|
103
|
+
def gen_fn():
|
|
104
|
+
while True:
|
|
105
|
+
if sample_weight is not None:
|
|
106
|
+
yield inputs, y, sample_weight
|
|
107
|
+
else:
|
|
108
|
+
yield inputs, y
|
|
109
|
+
|
|
110
|
+
return self.model.fit(
|
|
111
|
+
gen_fn(),
|
|
112
|
+
steps_per_epoch=1,
|
|
113
|
+
epochs=epochs,
|
|
114
|
+
shuffle=False, # one full-graph batch per step
|
|
115
|
+
verbose=verbose,
|
|
116
|
+
callbacks=callbacks,
|
|
117
|
+
)
|
|
118
|
+
|
|
119
|
+
def predict(self, data: Any, mask: Optional[str] = None):
|
|
120
|
+
r"""Returns continuous predictions for nodes."""
|
|
121
|
+
if self.model is None:
|
|
122
|
+
raise RuntimeError("Model is not initialized.")
|
|
123
|
+
inputs = self._extract_inputs(data)
|
|
124
|
+
pred = self.model(inputs, training=False)
|
|
125
|
+
if mask is not None and hasattr(data, mask):
|
|
126
|
+
pred = pred[getattr(data, mask)]
|
|
127
|
+
return pred
|
|
128
|
+
|
|
129
|
+
def evaluate(self, data: Any, mask: Optional[str] = "test_mask") -> Dict[str, float]:
|
|
130
|
+
r"""Evaluates Mean Absolute Error and Mean Squared Error."""
|
|
131
|
+
pred = self.predict(data, mask=mask)
|
|
132
|
+
y = data.y
|
|
133
|
+
if mask is not None and hasattr(data, mask):
|
|
134
|
+
y = y[getattr(data, mask)]
|
|
135
|
+
diff = ops.cast(pred, "float32") - ops.cast(y, "float32")
|
|
136
|
+
mae = float(ops.convert_to_numpy(ops.mean(ops.abs(diff))))
|
|
137
|
+
mse = float(ops.convert_to_numpy(ops.mean(ops.square(diff))))
|
|
138
|
+
return {"mae": mae, "mse": mse, "loss": mse}
|
|
@@ -0,0 +1,319 @@
|
|
|
1
|
+
"""Unit tests for K3-Node high-level Task APIs."""
|
|
2
|
+
|
|
3
|
+
import os
|
|
4
|
+
import tempfile
|
|
5
|
+
import pytest
|
|
6
|
+
import numpy as np
|
|
7
|
+
import keras
|
|
8
|
+
from keras import ops
|
|
9
|
+
|
|
10
|
+
import k3_node
|
|
11
|
+
from k3_node.data import Data
|
|
12
|
+
from k3_node.tasks import (
|
|
13
|
+
BaseTask,
|
|
14
|
+
resolve_backbone,
|
|
15
|
+
NodeClassifier,
|
|
16
|
+
NodeRegressor,
|
|
17
|
+
GraphClassifier,
|
|
18
|
+
GraphRegressor,
|
|
19
|
+
LinkPredictor,
|
|
20
|
+
)
|
|
21
|
+
|
|
22
|
+
|
|
23
|
+
def _create_synthetic_node_data(num_nodes=24, in_channels=8, num_classes=3, multi_label=False):
|
|
24
|
+
x = np.random.randn(num_nodes, in_channels).astype("float32")
|
|
25
|
+
edges_src = np.arange(num_nodes - 1, dtype="int64")
|
|
26
|
+
edges_dst = np.arange(1, num_nodes, dtype="int64")
|
|
27
|
+
edge_index = np.stack([edges_src, edges_dst], axis=0)
|
|
28
|
+
|
|
29
|
+
if multi_label:
|
|
30
|
+
y = np.random.randint(0, 2, size=(num_nodes, num_classes)).astype("float32")
|
|
31
|
+
else:
|
|
32
|
+
y = np.random.randint(0, num_classes, size=(num_nodes,)).astype("int64")
|
|
33
|
+
|
|
34
|
+
train_mask = np.zeros(num_nodes, dtype=bool)
|
|
35
|
+
train_mask[: num_nodes // 2] = True
|
|
36
|
+
val_mask = np.zeros(num_nodes, dtype=bool)
|
|
37
|
+
val_mask[num_nodes // 2 : 3 * num_nodes // 4] = True
|
|
38
|
+
test_mask = ~(train_mask | val_mask)
|
|
39
|
+
|
|
40
|
+
return Data(
|
|
41
|
+
x=x,
|
|
42
|
+
edge_index=edge_index,
|
|
43
|
+
y=y,
|
|
44
|
+
train_mask=train_mask,
|
|
45
|
+
val_mask=val_mask,
|
|
46
|
+
test_mask=test_mask,
|
|
47
|
+
)
|
|
48
|
+
|
|
49
|
+
|
|
50
|
+
def _create_synthetic_graph_dataset(num_graphs=12, nodes_per_graph=6, in_channels=8, is_regression=False):
|
|
51
|
+
dataset = []
|
|
52
|
+
for i in range(num_graphs):
|
|
53
|
+
x = np.random.randn(nodes_per_graph, in_channels).astype("float32")
|
|
54
|
+
edges_src = np.arange(nodes_per_graph - 1, dtype="int64")
|
|
55
|
+
edges_dst = np.arange(1, nodes_per_graph, dtype="int64")
|
|
56
|
+
edge_index = np.stack([edges_src, edges_dst], axis=0)
|
|
57
|
+
|
|
58
|
+
if is_regression:
|
|
59
|
+
y = np.array([float(np.mean(x))], dtype="float32")
|
|
60
|
+
else:
|
|
61
|
+
y = np.array(i % 2, dtype="int64")
|
|
62
|
+
|
|
63
|
+
dataset.append(Data(x=x, edge_index=edge_index, y=y))
|
|
64
|
+
return dataset
|
|
65
|
+
|
|
66
|
+
|
|
67
|
+
# ==============================================================================
|
|
68
|
+
# 1. Top-Level Imports & Base Task Tests
|
|
69
|
+
# ==============================================================================
|
|
70
|
+
|
|
71
|
+
def test_tasks_module_exports():
|
|
72
|
+
"""Verify tasks are exported at both k3_node.tasks and k3_node top level."""
|
|
73
|
+
assert hasattr(k3_node, "tasks")
|
|
74
|
+
assert hasattr(k3_node, "NodeClassifier")
|
|
75
|
+
assert hasattr(k3_node, "NodeRegressor")
|
|
76
|
+
assert hasattr(k3_node, "GraphClassifier")
|
|
77
|
+
assert hasattr(k3_node, "GraphRegressor")
|
|
78
|
+
assert hasattr(k3_node, "LinkPredictor")
|
|
79
|
+
|
|
80
|
+
assert k3_node.NodeClassifier is NodeClassifier
|
|
81
|
+
assert k3_node.NodeRegressor is NodeRegressor
|
|
82
|
+
assert k3_node.GraphClassifier is GraphClassifier
|
|
83
|
+
assert k3_node.GraphRegressor is GraphRegressor
|
|
84
|
+
assert k3_node.LinkPredictor is LinkPredictor
|
|
85
|
+
|
|
86
|
+
|
|
87
|
+
def test_base_task():
|
|
88
|
+
"""Test BaseTask compile, extract_inputs, save and load."""
|
|
89
|
+
dummy_model = keras.Sequential([keras.layers.Dense(4)])
|
|
90
|
+
task = BaseTask(model=dummy_model)
|
|
91
|
+
task.compile(optimizer="adam", loss="mse")
|
|
92
|
+
assert task._is_compiled
|
|
93
|
+
|
|
94
|
+
# extract_inputs
|
|
95
|
+
data = _create_synthetic_node_data(num_nodes=5, in_channels=4)
|
|
96
|
+
extracted = task._extract_inputs(data)
|
|
97
|
+
assert isinstance(extracted, tuple)
|
|
98
|
+
|
|
99
|
+
# save and load
|
|
100
|
+
with tempfile.TemporaryDirectory() as tmpdir:
|
|
101
|
+
filepath = os.path.join(tmpdir, "test_task.keras")
|
|
102
|
+
task.save(filepath)
|
|
103
|
+
loaded = BaseTask.load(filepath)
|
|
104
|
+
assert loaded.model is not None
|
|
105
|
+
|
|
106
|
+
|
|
107
|
+
def test_resolve_backbone():
|
|
108
|
+
"""Test backbone resolver with valid strings, custom models, and errors."""
|
|
109
|
+
for name in ["gcn", "gat", "sage", "gin", "mlp"]:
|
|
110
|
+
model = resolve_backbone(name, in_channels=8, hidden_channels=16, out_channels=4, num_layers=2)
|
|
111
|
+
assert isinstance(model, keras.Model)
|
|
112
|
+
|
|
113
|
+
# Custom model pass-through
|
|
114
|
+
custom_model = keras.Sequential([keras.layers.Dense(4)])
|
|
115
|
+
resolved_custom = resolve_backbone(custom_model, in_channels=8, out_channels=4)
|
|
116
|
+
assert resolved_custom is custom_model
|
|
117
|
+
|
|
118
|
+
# Unknown backbone
|
|
119
|
+
with pytest.raises(ValueError):
|
|
120
|
+
resolve_backbone("non_existent_backbone", in_channels=8, out_channels=4)
|
|
121
|
+
|
|
122
|
+
# Invalid type
|
|
123
|
+
with pytest.raises(TypeError):
|
|
124
|
+
resolve_backbone(12345, in_channels=8, out_channels=4)
|
|
125
|
+
|
|
126
|
+
|
|
127
|
+
# ==============================================================================
|
|
128
|
+
# 2. NodeClassifier Tests
|
|
129
|
+
# ==============================================================================
|
|
130
|
+
|
|
131
|
+
@pytest.mark.parametrize("backbone", ["gcn", "sage"])
|
|
132
|
+
def test_node_classifier_fit_predict_evaluate(backbone):
|
|
133
|
+
data = _create_synthetic_node_data(num_nodes=20, in_channels=8, num_classes=3)
|
|
134
|
+
clf = NodeClassifier(backbone=backbone, hidden_channels=16, num_layers=2, dropout=0.0)
|
|
135
|
+
|
|
136
|
+
# Fit
|
|
137
|
+
clf.fit(data, epochs=2, lr=0.01, verbose=0)
|
|
138
|
+
assert clf.model is not None
|
|
139
|
+
|
|
140
|
+
# Predict discrete labels
|
|
141
|
+
preds = clf.predict(data, mask="test_mask")
|
|
142
|
+
assert ops.shape(preds)[0] == int(ops.convert_to_numpy(ops.sum(ops.cast(data.test_mask, "int32"))))
|
|
143
|
+
|
|
144
|
+
# Predict probabilities
|
|
145
|
+
probs = clf.predict_proba(data, mask="test_mask")
|
|
146
|
+
assert ops.shape(probs)[-1] == 3
|
|
147
|
+
|
|
148
|
+
# Evaluate
|
|
149
|
+
metrics = clf.evaluate(data, mask="test_mask")
|
|
150
|
+
assert "accuracy" in metrics
|
|
151
|
+
assert 0.0 <= metrics["accuracy"] <= 1.0
|
|
152
|
+
|
|
153
|
+
|
|
154
|
+
def test_node_classifier_multi_label():
|
|
155
|
+
data = _create_synthetic_node_data(num_nodes=20, in_channels=8, num_classes=4, multi_label=True)
|
|
156
|
+
clf = NodeClassifier(backbone="gcn", hidden_channels=16, num_layers=2, multi_label=True)
|
|
157
|
+
clf.fit(data, epochs=2, verbose=0)
|
|
158
|
+
|
|
159
|
+
probs = clf.predict_proba(data, mask="test_mask")
|
|
160
|
+
assert ops.shape(probs)[-1] == 4
|
|
161
|
+
|
|
162
|
+
preds = clf.predict(data, mask="test_mask")
|
|
163
|
+
assert ops.shape(preds)[-1] == 4
|
|
164
|
+
|
|
165
|
+
|
|
166
|
+
# ==============================================================================
|
|
167
|
+
# 3. NodeRegressor Tests
|
|
168
|
+
# ==============================================================================
|
|
169
|
+
|
|
170
|
+
def test_node_regressor_fit_predict_evaluate():
|
|
171
|
+
data = _create_synthetic_node_data(num_nodes=20, in_channels=8, num_classes=1)
|
|
172
|
+
data.y = np.random.randn(20, 1).astype("float32")
|
|
173
|
+
|
|
174
|
+
reg = NodeRegressor(backbone="gat", hidden_channels=16, num_layers=2)
|
|
175
|
+
reg.fit(data, epochs=2, lr=0.01, verbose=0)
|
|
176
|
+
|
|
177
|
+
preds = reg.predict(data, mask="test_mask")
|
|
178
|
+
assert ops.shape(preds)[0] == int(ops.convert_to_numpy(ops.sum(ops.cast(data.test_mask, "int32"))))
|
|
179
|
+
|
|
180
|
+
metrics = reg.evaluate(data, mask="test_mask")
|
|
181
|
+
assert "mae" in metrics
|
|
182
|
+
assert "mse" in metrics
|
|
183
|
+
assert "loss" in metrics
|
|
184
|
+
assert metrics["mae"] >= 0.0
|
|
185
|
+
|
|
186
|
+
|
|
187
|
+
# ==============================================================================
|
|
188
|
+
# 4. GraphClassifier Tests
|
|
189
|
+
# ==============================================================================
|
|
190
|
+
|
|
191
|
+
@pytest.mark.parametrize("pooling", ["mean", "add", "max"])
|
|
192
|
+
def test_graph_classifier_fit_predict_evaluate(pooling):
|
|
193
|
+
dataset = _create_synthetic_graph_dataset(num_graphs=10, nodes_per_graph=5, in_channels=6)
|
|
194
|
+
clf = GraphClassifier(
|
|
195
|
+
backbone="gin",
|
|
196
|
+
hidden_channels=16,
|
|
197
|
+
num_layers=2,
|
|
198
|
+
pooling=pooling,
|
|
199
|
+
dropout=0.0,
|
|
200
|
+
)
|
|
201
|
+
|
|
202
|
+
clf.fit(dataset, epochs=2, batch_size=4, verbose=0)
|
|
203
|
+
|
|
204
|
+
# Predict
|
|
205
|
+
preds = clf.predict(dataset[:4], batch_size=4)
|
|
206
|
+
assert ops.shape(preds)[0] == 4
|
|
207
|
+
|
|
208
|
+
# Predict proba
|
|
209
|
+
probs = clf.predict_proba(dataset[:4], batch_size=4)
|
|
210
|
+
assert ops.shape(probs)[0] == 4
|
|
211
|
+
assert ops.shape(probs)[1] == 2
|
|
212
|
+
|
|
213
|
+
# Evaluate
|
|
214
|
+
metrics = clf.evaluate(dataset[:4], batch_size=4)
|
|
215
|
+
assert "accuracy" in metrics
|
|
216
|
+
assert 0.0 <= metrics["accuracy"] <= 1.0
|
|
217
|
+
|
|
218
|
+
|
|
219
|
+
# ==============================================================================
|
|
220
|
+
# 5. GraphRegressor Tests
|
|
221
|
+
# ==============================================================================
|
|
222
|
+
|
|
223
|
+
@pytest.mark.parametrize("loss_name", ["mae", "mse"])
|
|
224
|
+
def test_graph_regressor_fit_predict_evaluate(loss_name):
|
|
225
|
+
dataset = _create_synthetic_graph_dataset(num_graphs=10, nodes_per_graph=5, in_channels=6, is_regression=True)
|
|
226
|
+
reg = GraphRegressor(
|
|
227
|
+
backbone="sage",
|
|
228
|
+
hidden_channels=16,
|
|
229
|
+
num_layers=2,
|
|
230
|
+
pooling="mean",
|
|
231
|
+
loss=loss_name,
|
|
232
|
+
)
|
|
233
|
+
|
|
234
|
+
reg.fit(dataset, epochs=2, batch_size=4, verbose=0)
|
|
235
|
+
|
|
236
|
+
preds = reg.predict(dataset[:4], batch_size=4)
|
|
237
|
+
assert ops.shape(preds)[0] == 4
|
|
238
|
+
|
|
239
|
+
metrics = reg.evaluate(dataset[:4], batch_size=4)
|
|
240
|
+
assert "mae" in metrics
|
|
241
|
+
assert "mse" in metrics
|
|
242
|
+
assert "loss" in metrics
|
|
243
|
+
assert metrics["mae"] >= 0.0
|
|
244
|
+
|
|
245
|
+
|
|
246
|
+
# ==============================================================================
|
|
247
|
+
# 6. LinkPredictor Tests
|
|
248
|
+
# ==============================================================================
|
|
249
|
+
|
|
250
|
+
@pytest.mark.parametrize("decoder", ["inner_product", "cosine", "mlp"])
|
|
251
|
+
def test_link_predictor_decoders_and_fit(decoder):
|
|
252
|
+
data = _create_synthetic_node_data(num_nodes=16, in_channels=8)
|
|
253
|
+
lp = LinkPredictor(
|
|
254
|
+
backbone="gcn",
|
|
255
|
+
hidden_channels=16,
|
|
256
|
+
out_channels=16,
|
|
257
|
+
num_layers=2,
|
|
258
|
+
decoder=decoder,
|
|
259
|
+
)
|
|
260
|
+
|
|
261
|
+
# Dynamic negative sampling training
|
|
262
|
+
lp.fit(data, epochs=2, lr=0.01, neg_ratio=1.0, verbose=0)
|
|
263
|
+
|
|
264
|
+
# Encode node embeddings
|
|
265
|
+
z = lp.encode(data)
|
|
266
|
+
assert ops.shape(z) == (16, 16)
|
|
267
|
+
|
|
268
|
+
# Predict proba & binary labels
|
|
269
|
+
query_edges = np.array([[0, 1, 2], [1, 2, 3]], dtype="int32")
|
|
270
|
+
probs = lp.predict_proba(data, edge_label_index=query_edges)
|
|
271
|
+
assert ops.shape(probs) == (3,)
|
|
272
|
+
|
|
273
|
+
preds = lp.predict(data, edge_label_index=query_edges, threshold=0.5)
|
|
274
|
+
assert ops.shape(preds) == (3,)
|
|
275
|
+
|
|
276
|
+
# Evaluate
|
|
277
|
+
metrics = lp.evaluate(data, edge_label_index=query_edges, edge_label=[1, 1, 0])
|
|
278
|
+
assert "accuracy" in metrics
|
|
279
|
+
assert 0.0 <= metrics["accuracy"] <= 1.0
|
|
280
|
+
|
|
281
|
+
|
|
282
|
+
def test_link_predictor_explicit_labels():
|
|
283
|
+
data = _create_synthetic_node_data(num_nodes=16, in_channels=8)
|
|
284
|
+
edge_label_index = np.array([[0, 1, 2, 3], [1, 2, 3, 0]], dtype="int32")
|
|
285
|
+
edge_label = np.array([1, 1, 0, 0], dtype="float32")
|
|
286
|
+
|
|
287
|
+
lp = LinkPredictor(backbone="gcn", hidden_channels=16, out_channels=16)
|
|
288
|
+
lp.fit(data, edge_label_index=edge_label_index, edge_label=edge_label, epochs=2, verbose=0)
|
|
289
|
+
|
|
290
|
+
metrics = lp.evaluate(data, edge_label_index=edge_label_index, edge_label=edge_label)
|
|
291
|
+
assert "accuracy" in metrics
|
|
292
|
+
assert "auc" in metrics
|
|
293
|
+
assert "ap" in metrics
|
|
294
|
+
|
|
295
|
+
|
|
296
|
+
# ==============================================================================
|
|
297
|
+
# 7. Applications API Tests
|
|
298
|
+
# ==============================================================================
|
|
299
|
+
|
|
300
|
+
def test_applications_api_exports():
|
|
301
|
+
"""Verify k3_node.applications submodules and model exports."""
|
|
302
|
+
import k3_node.applications as apps
|
|
303
|
+
assert hasattr(apps, "chemistry")
|
|
304
|
+
assert hasattr(apps, "materials")
|
|
305
|
+
assert hasattr(apps, "bio")
|
|
306
|
+
|
|
307
|
+
# Chemistry exports
|
|
308
|
+
assert hasattr(apps.chemistry, "AttentiveFP")
|
|
309
|
+
assert hasattr(apps.chemistry, "SchNet")
|
|
310
|
+
assert hasattr(apps.chemistry, "DimeNetPlusPlus")
|
|
311
|
+
|
|
312
|
+
# Materials exports
|
|
313
|
+
assert hasattr(apps.materials, "MEGNet")
|
|
314
|
+
assert hasattr(apps.materials, "M3GNet")
|
|
315
|
+
assert hasattr(apps.materials, "CHGNet")
|
|
316
|
+
|
|
317
|
+
# Bio exports
|
|
318
|
+
assert hasattr(apps.bio, "UniMolDockingModel")
|
|
319
|
+
assert hasattr(apps.bio, "DockingPoseModelV2")
|