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,443 @@
|
|
|
1
|
+
import copy
|
|
2
|
+
import inspect
|
|
3
|
+
from typing import Any, Callable, Dict, List, Optional, Tuple, Union
|
|
4
|
+
|
|
5
|
+
import keras
|
|
6
|
+
from keras import ops
|
|
7
|
+
|
|
8
|
+
from k3_node.layers.conv import (
|
|
9
|
+
EdgeConv,
|
|
10
|
+
GATConv,
|
|
11
|
+
GATv2Conv,
|
|
12
|
+
GCNConv,
|
|
13
|
+
GINConv,
|
|
14
|
+
MessagePassing,
|
|
15
|
+
PNAConv,
|
|
16
|
+
SAGEConv,
|
|
17
|
+
)
|
|
18
|
+
from k3_node.models.mlp import MLP, _normalization_resolver
|
|
19
|
+
from k3_node.models.jumping_knowledge import JumpingKnowledge
|
|
20
|
+
from k3_node.hub.hub_mixin import K3NodeHubMixin
|
|
21
|
+
|
|
22
|
+
|
|
23
|
+
class BasicGNN(K3NodeHubMixin, keras.Model):
|
|
24
|
+
r"""An abstract base class for implementing basic GNN models.
|
|
25
|
+
|
|
26
|
+
Args:
|
|
27
|
+
in_channels (int or tuple): Size of each input sample.
|
|
28
|
+
hidden_channels (int): Size of each hidden sample.
|
|
29
|
+
num_layers (int): Number of message passing layers.
|
|
30
|
+
out_channels (int, optional): If not set to :obj:`None`, will apply a
|
|
31
|
+
final linear transformation to convert hidden node embeddings to
|
|
32
|
+
output size :obj:`out_channels`. (default: :obj:`None`)
|
|
33
|
+
dropout (float, optional): Dropout probability. (default: :obj:`0.`)
|
|
34
|
+
act (str or Callable, optional): The non-linear activation function to
|
|
35
|
+
use. (default: :obj:`"relu"`)
|
|
36
|
+
act_first (bool, optional): If set to :obj:`True`, activation is
|
|
37
|
+
applied before normalization. (default: :obj:`False`)
|
|
38
|
+
act_kwargs (Dict[str, Any], optional): Arguments passed to the
|
|
39
|
+
respective activation function defined by :obj:`act`.
|
|
40
|
+
(default: :obj:`None`)
|
|
41
|
+
norm (str or Callable, optional): The normalization function to
|
|
42
|
+
use. (default: :obj:`None`)
|
|
43
|
+
norm_kwargs (Dict[str, Any], optional): Arguments passed to the
|
|
44
|
+
respective normalization function defined by :obj:`norm`.
|
|
45
|
+
(default: :obj:`None`)
|
|
46
|
+
jk (str, optional): The Jumping Knowledge mode. If specified, the model
|
|
47
|
+
will additionally apply a final linear transformation to transform
|
|
48
|
+
node embeddings to the expected output feature dimensionality.
|
|
49
|
+
(:obj:`None`, :obj:`"last"`, :obj:`"cat"`, :obj:`"max"`,
|
|
50
|
+
:obj:`"lstm"`). (default: :obj:`None`)
|
|
51
|
+
**kwargs (optional): Additional arguments of the underlying
|
|
52
|
+
:class:`torch_geometric.nn.conv.MessagePassing` layers.
|
|
53
|
+
|
|
54
|
+
Call arguments: ``x``, ``edge_index`` and optionally ``edge_weight``, ``edge_attr`` and
|
|
55
|
+
``batch``. With mini-batches from :class:`~k3_node.loader.NeighborLoader`, also pass
|
|
56
|
+
``num_sampled_nodes_per_hop`` and ``num_sampled_edges_per_hop`` (the batch's
|
|
57
|
+
``num_sampled_nodes`` / ``num_sampled_edges``) to skip the nodes and edges that no longer
|
|
58
|
+
affect the seed nodes after each layer, as in PyG. This trimming needs concrete sizes, so
|
|
59
|
+
compile the model with ``run_eagerly=True`` when using it.
|
|
60
|
+
"""
|
|
61
|
+
supports_edge_weight: bool = False
|
|
62
|
+
supports_edge_attr: bool = False
|
|
63
|
+
supports_norm_batch: bool = False
|
|
64
|
+
|
|
65
|
+
def __init__(
|
|
66
|
+
self,
|
|
67
|
+
in_channels: int,
|
|
68
|
+
hidden_channels: int,
|
|
69
|
+
num_layers: int,
|
|
70
|
+
out_channels: Optional[int] = None,
|
|
71
|
+
dropout: float = 0.0,
|
|
72
|
+
act: Union[str, Callable, None] = "relu",
|
|
73
|
+
act_first: bool = False,
|
|
74
|
+
act_kwargs: Optional[Dict[str, Any]] = None,
|
|
75
|
+
norm: Union[str, Callable, None] = None,
|
|
76
|
+
norm_kwargs: Optional[Dict[str, Any]] = None,
|
|
77
|
+
jk: Optional[str] = None,
|
|
78
|
+
**kwargs,
|
|
79
|
+
):
|
|
80
|
+
super().__init__()
|
|
81
|
+
|
|
82
|
+
self.in_channels = in_channels
|
|
83
|
+
self.hidden_channels = hidden_channels
|
|
84
|
+
self.num_layers = num_layers
|
|
85
|
+
dropout = float(dropout) if dropout is not None else 0.0
|
|
86
|
+
self.dropout_p = dropout
|
|
87
|
+
self.dropout = keras.layers.Dropout(rate=dropout) if dropout > 0 else None
|
|
88
|
+
|
|
89
|
+
if isinstance(act, str):
|
|
90
|
+
self.act = keras.activations.get(act)
|
|
91
|
+
else:
|
|
92
|
+
self.act = act
|
|
93
|
+
|
|
94
|
+
self.jk_mode = jk
|
|
95
|
+
self.act_first = act_first
|
|
96
|
+
self.norm_query = norm
|
|
97
|
+
self.norm_kwargs = norm_kwargs or {}
|
|
98
|
+
|
|
99
|
+
if out_channels is not None:
|
|
100
|
+
self.out_channels = out_channels
|
|
101
|
+
else:
|
|
102
|
+
self.out_channels = hidden_channels
|
|
103
|
+
|
|
104
|
+
self.convs = []
|
|
105
|
+
curr_in = in_channels
|
|
106
|
+
if num_layers > 1:
|
|
107
|
+
self.convs.append(self.init_conv(curr_in, hidden_channels, **kwargs))
|
|
108
|
+
if isinstance(curr_in, (tuple, list)):
|
|
109
|
+
curr_in = (hidden_channels, hidden_channels)
|
|
110
|
+
else:
|
|
111
|
+
curr_in = hidden_channels
|
|
112
|
+
|
|
113
|
+
for _ in range(num_layers - 2):
|
|
114
|
+
self.convs.append(self.init_conv(curr_in, hidden_channels, **kwargs))
|
|
115
|
+
if isinstance(curr_in, (tuple, list)):
|
|
116
|
+
curr_in = (hidden_channels, hidden_channels)
|
|
117
|
+
else:
|
|
118
|
+
curr_in = hidden_channels
|
|
119
|
+
|
|
120
|
+
if out_channels is not None and jk is None:
|
|
121
|
+
self._is_conv_to_out = True
|
|
122
|
+
self.convs.append(self.init_conv(curr_in, out_channels, **kwargs))
|
|
123
|
+
else:
|
|
124
|
+
self.convs.append(self.init_conv(curr_in, hidden_channels, **kwargs))
|
|
125
|
+
|
|
126
|
+
self.norms = []
|
|
127
|
+
self.supports_norm_batch = False
|
|
128
|
+
|
|
129
|
+
for _ in range(num_layers - 1):
|
|
130
|
+
if norm is not None:
|
|
131
|
+
norm_layer = _normalization_resolver(norm, hidden_channels, **self.norm_kwargs)
|
|
132
|
+
self.norms.append(norm_layer)
|
|
133
|
+
if hasattr(norm_layer, "call"):
|
|
134
|
+
sig = inspect.signature(norm_layer.call).parameters
|
|
135
|
+
self.supports_norm_batch = "batch" in sig
|
|
136
|
+
else:
|
|
137
|
+
self.norms.append(None)
|
|
138
|
+
|
|
139
|
+
if jk is not None:
|
|
140
|
+
if norm is not None:
|
|
141
|
+
self.norms.append(_normalization_resolver(norm, hidden_channels, **self.norm_kwargs))
|
|
142
|
+
else:
|
|
143
|
+
self.norms.append(None)
|
|
144
|
+
else:
|
|
145
|
+
self.norms.append(None)
|
|
146
|
+
|
|
147
|
+
if jk is not None and jk != "last":
|
|
148
|
+
self.jk = JumpingKnowledge(jk, hidden_channels, num_layers)
|
|
149
|
+
|
|
150
|
+
if jk is not None:
|
|
151
|
+
if jk == "cat":
|
|
152
|
+
jk_in = num_layers * hidden_channels
|
|
153
|
+
else:
|
|
154
|
+
jk_in = hidden_channels
|
|
155
|
+
self.lin = keras.layers.Dense(self.out_channels)
|
|
156
|
+
|
|
157
|
+
def init_conv(self, in_channels: Union[int, Tuple[int, int]],
|
|
158
|
+
out_channels: int, **kwargs) -> MessagePassing:
|
|
159
|
+
raise NotImplementedError
|
|
160
|
+
|
|
161
|
+
def build(self, input_shape=None):
|
|
162
|
+
self.built = True
|
|
163
|
+
|
|
164
|
+
def reset_parameters(self):
|
|
165
|
+
r"""Resets all learnable parameters of the module."""
|
|
166
|
+
for conv in self.convs:
|
|
167
|
+
if hasattr(conv, "reset_parameters"):
|
|
168
|
+
conv.reset_parameters()
|
|
169
|
+
for norm in self.norms:
|
|
170
|
+
if norm is not None and hasattr(norm, "reset_parameters"):
|
|
171
|
+
norm.reset_parameters()
|
|
172
|
+
if hasattr(self, "jk") and hasattr(self.jk, "reset_parameters"):
|
|
173
|
+
self.jk.reset_parameters()
|
|
174
|
+
if hasattr(self, "lin") and hasattr(self.lin, "reset_parameters"):
|
|
175
|
+
self.lin.reset_parameters()
|
|
176
|
+
|
|
177
|
+
def call(
|
|
178
|
+
self,
|
|
179
|
+
x,
|
|
180
|
+
edge_index=None,
|
|
181
|
+
edge_weight=None,
|
|
182
|
+
edge_attr=None,
|
|
183
|
+
batch=None,
|
|
184
|
+
batch_size=None,
|
|
185
|
+
num_sampled_nodes_per_hop=None,
|
|
186
|
+
num_sampled_edges_per_hop=None,
|
|
187
|
+
training=None,
|
|
188
|
+
):
|
|
189
|
+
if hasattr(x, "edge_index") and edge_index is None:
|
|
190
|
+
edge_index = getattr(x, "edge_index", None)
|
|
191
|
+
edge_weight = getattr(x, "edge_weight", None) if edge_weight is None else edge_weight
|
|
192
|
+
edge_attr = getattr(x, "edge_attr", None) if edge_attr is None else edge_attr
|
|
193
|
+
batch = getattr(x, "batch", None) if batch is None else batch
|
|
194
|
+
x = x.x
|
|
195
|
+
elif isinstance(x, (tuple, list)) and edge_index is None:
|
|
196
|
+
if len(x) > 1:
|
|
197
|
+
edge_index = x[1]
|
|
198
|
+
if len(x) > 2:
|
|
199
|
+
edge_attr = x[2]
|
|
200
|
+
x = x[0]
|
|
201
|
+
|
|
202
|
+
trim = num_sampled_nodes_per_hop is not None and num_sampled_edges_per_hop is not None
|
|
203
|
+
if trim:
|
|
204
|
+
from k3_node.ops.host import to_numpy
|
|
205
|
+
|
|
206
|
+
nodes_per_hop = [int(n) for n in to_numpy(num_sampled_nodes_per_hop)]
|
|
207
|
+
edges_per_hop = [int(n) for n in to_numpy(num_sampled_edges_per_hop)]
|
|
208
|
+
|
|
209
|
+
xs: List = []
|
|
210
|
+
# `training` is forwarded explicitly: Keras does not propagate it to nested layers on JAX.
|
|
211
|
+
for i, (conv, norm) in enumerate(zip(self.convs, self.norms)):
|
|
212
|
+
if trim and i > 0:
|
|
213
|
+
# Hierarchical neighborhood sampling: the nodes and edges of the outermost hop no
|
|
214
|
+
# longer influence the seed nodes, so drop them (as PyG's `trim_to_layer`).
|
|
215
|
+
x = x[: x.shape[0] - nodes_per_hop[-i]]
|
|
216
|
+
num_edges = edge_index.shape[1] - edges_per_hop[-i]
|
|
217
|
+
edge_index = edge_index[:, :num_edges]
|
|
218
|
+
edge_weight = None if edge_weight is None else edge_weight[:num_edges]
|
|
219
|
+
edge_attr = None if edge_attr is None else edge_attr[:num_edges]
|
|
220
|
+
if self.supports_edge_weight and self.supports_edge_attr:
|
|
221
|
+
x = conv(x, edge_index, edge_weight=edge_weight, edge_attr=edge_attr, training=training)
|
|
222
|
+
elif self.supports_edge_weight:
|
|
223
|
+
x = conv(x, edge_index, edge_weight=edge_weight, training=training)
|
|
224
|
+
elif self.supports_edge_attr:
|
|
225
|
+
x = conv(x, edge_index, edge_attr=edge_attr, training=training)
|
|
226
|
+
else:
|
|
227
|
+
x = conv(x, edge_index, training=training)
|
|
228
|
+
|
|
229
|
+
if i < self.num_layers - 1 or self.jk_mode is not None:
|
|
230
|
+
if self.act is not None and self.act_first:
|
|
231
|
+
x = self.act(x)
|
|
232
|
+
if norm is not None:
|
|
233
|
+
if self.supports_norm_batch and batch is not None:
|
|
234
|
+
x = norm(x, batch=batch, training=training)
|
|
235
|
+
else:
|
|
236
|
+
x = norm(x, training=training)
|
|
237
|
+
if self.act is not None and not self.act_first:
|
|
238
|
+
x = self.act(x)
|
|
239
|
+
if self.dropout is not None:
|
|
240
|
+
x = self.dropout(x, training=training)
|
|
241
|
+
if hasattr(self, "jk"):
|
|
242
|
+
xs.append(x)
|
|
243
|
+
|
|
244
|
+
if hasattr(self, "jk"):
|
|
245
|
+
x = self.jk(xs)
|
|
246
|
+
if hasattr(self, "lin"):
|
|
247
|
+
x = self.lin(x)
|
|
248
|
+
|
|
249
|
+
return x
|
|
250
|
+
|
|
251
|
+
def __repr__(self) -> str:
|
|
252
|
+
return (f'{self.__class__.__name__}({self.in_channels}, '
|
|
253
|
+
f'{self.out_channels}, num_layers={self.num_layers})')
|
|
254
|
+
|
|
255
|
+
|
|
256
|
+
class GCN(BasicGNN):
|
|
257
|
+
r"""The Graph Neural Network from the `"Semi-supervised
|
|
258
|
+
Classification with Graph Convolutional Networks"
|
|
259
|
+
<https://arxiv.org/abs/1609.02907>`_ paper, using the
|
|
260
|
+
:class:`~k3_node.layers.conv.GCNConv` operator for message passing.
|
|
261
|
+
|
|
262
|
+
Example:
|
|
263
|
+
```python
|
|
264
|
+
import numpy as np
|
|
265
|
+
from k3_node.models import GCN
|
|
266
|
+
|
|
267
|
+
x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
|
|
268
|
+
edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
|
|
269
|
+
|
|
270
|
+
model = GCN(in_channels=8, hidden_channels=16, num_layers=2, out_channels=4)
|
|
271
|
+
out = model(x, edge_index) # e.g. logits for 4 classes per node
|
|
272
|
+
print(tuple(out.shape)) # (10, 4)
|
|
273
|
+
```
|
|
274
|
+
"""
|
|
275
|
+
supports_edge_weight: bool = True
|
|
276
|
+
supports_edge_attr: bool = False
|
|
277
|
+
|
|
278
|
+
def init_conv(self, in_channels: int, out_channels: int, **kwargs) -> MessagePassing:
|
|
279
|
+
return GCNConv(in_channels, out_channels, **kwargs)
|
|
280
|
+
|
|
281
|
+
|
|
282
|
+
class GraphSAGE(BasicGNN):
|
|
283
|
+
r"""The Graph Neural Network from the `"Inductive Representation Learning
|
|
284
|
+
on Large Graphs" <https://arxiv.org/abs/1706.02216>`_ paper, using the
|
|
285
|
+
:class:`~k3_node.layers.conv.SAGEConv` operator for message passing.
|
|
286
|
+
|
|
287
|
+
Example:
|
|
288
|
+
```python
|
|
289
|
+
import numpy as np
|
|
290
|
+
from k3_node.models import GraphSAGE
|
|
291
|
+
|
|
292
|
+
x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
|
|
293
|
+
edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
|
|
294
|
+
|
|
295
|
+
model = GraphSAGE(in_channels=8, hidden_channels=16, num_layers=2, out_channels=4)
|
|
296
|
+
out = model(x, edge_index) # e.g. logits for 4 classes per node
|
|
297
|
+
print(tuple(out.shape)) # (10, 4)
|
|
298
|
+
```
|
|
299
|
+
"""
|
|
300
|
+
supports_edge_weight: bool = False
|
|
301
|
+
supports_edge_attr: bool = False
|
|
302
|
+
|
|
303
|
+
def init_conv(self, in_channels: Union[int, Tuple[int, int]],
|
|
304
|
+
out_channels: int, **kwargs) -> MessagePassing:
|
|
305
|
+
return SAGEConv(in_channels, out_channels, **kwargs)
|
|
306
|
+
|
|
307
|
+
|
|
308
|
+
class GIN(BasicGNN):
|
|
309
|
+
r"""The Graph Neural Network from the `"How Powerful are Graph Neural
|
|
310
|
+
Networks?" <https://arxiv.org/abs/1810.00826>`_ paper, using the
|
|
311
|
+
:class:`~k3_node.layers.conv.GINConv` operator for message passing.
|
|
312
|
+
|
|
313
|
+
Example:
|
|
314
|
+
```python
|
|
315
|
+
import numpy as np
|
|
316
|
+
from k3_node.models import GIN
|
|
317
|
+
|
|
318
|
+
x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
|
|
319
|
+
edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
|
|
320
|
+
|
|
321
|
+
model = GIN(in_channels=8, hidden_channels=16, num_layers=2, out_channels=4)
|
|
322
|
+
out = model(x, edge_index) # e.g. logits for 4 classes per node
|
|
323
|
+
print(tuple(out.shape)) # (10, 4)
|
|
324
|
+
```
|
|
325
|
+
"""
|
|
326
|
+
supports_edge_weight: bool = False
|
|
327
|
+
supports_edge_attr: bool = False
|
|
328
|
+
|
|
329
|
+
def init_conv(self, in_channels: int, out_channels: int, **kwargs) -> MessagePassing:
|
|
330
|
+
mlp = MLP(
|
|
331
|
+
[in_channels, out_channels, out_channels],
|
|
332
|
+
act=self.act,
|
|
333
|
+
act_first=self.act_first,
|
|
334
|
+
norm=self.norm_query,
|
|
335
|
+
norm_kwargs=self.norm_kwargs,
|
|
336
|
+
)
|
|
337
|
+
return GINConv(mlp, **kwargs)
|
|
338
|
+
|
|
339
|
+
|
|
340
|
+
class GAT(BasicGNN):
|
|
341
|
+
r"""The Graph Neural Network from `"Graph Attention Networks"
|
|
342
|
+
<https://arxiv.org/abs/1710.10903>`_ or `"How Attentive are Graph Attention
|
|
343
|
+
Networks?" <https://arxiv.org/abs/2105.14491>`_ papers, using the
|
|
344
|
+
:class:`~k3_node.layers.conv.GATConv` or
|
|
345
|
+
:class:`~k3_node.layers.conv.GATv2Conv` operator for message passing.
|
|
346
|
+
|
|
347
|
+
Example:
|
|
348
|
+
```python
|
|
349
|
+
import numpy as np
|
|
350
|
+
from k3_node.models import GAT
|
|
351
|
+
|
|
352
|
+
x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
|
|
353
|
+
edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
|
|
354
|
+
|
|
355
|
+
model = GAT(in_channels=8, hidden_channels=16, num_layers=2, out_channels=4, heads=2)
|
|
356
|
+
out = model(x, edge_index) # e.g. logits for 4 classes per node
|
|
357
|
+
print(tuple(out.shape)) # (10, 4)
|
|
358
|
+
```
|
|
359
|
+
"""
|
|
360
|
+
supports_edge_weight: bool = False
|
|
361
|
+
supports_edge_attr: bool = True
|
|
362
|
+
|
|
363
|
+
def init_conv(self, in_channels: Union[int, Tuple[int, int]],
|
|
364
|
+
out_channels: int, **kwargs) -> MessagePassing:
|
|
365
|
+
v2 = kwargs.pop('v2', False)
|
|
366
|
+
heads = kwargs.pop('heads', 1)
|
|
367
|
+
concat = kwargs.pop('concat', True)
|
|
368
|
+
|
|
369
|
+
if getattr(self, '_is_conv_to_out', False):
|
|
370
|
+
concat = False
|
|
371
|
+
|
|
372
|
+
if concat and out_channels % heads != 0:
|
|
373
|
+
raise ValueError(f"Ensure that the number of output channels of "
|
|
374
|
+
f"'GATConv' (got '{out_channels}') is divisible "
|
|
375
|
+
f"by the number of heads (got '{heads}')")
|
|
376
|
+
|
|
377
|
+
if concat:
|
|
378
|
+
out_channels = out_channels // heads
|
|
379
|
+
|
|
380
|
+
Conv = GATConv if not v2 else GATv2Conv
|
|
381
|
+
return Conv(in_channels, out_channels, heads=heads, concat=concat,
|
|
382
|
+
dropout=self.dropout_p, **kwargs)
|
|
383
|
+
|
|
384
|
+
|
|
385
|
+
class PNA(BasicGNN):
|
|
386
|
+
r"""The Graph Neural Network from the `"Principal Neighbourhood Aggregation
|
|
387
|
+
for Graph Nets" <https://arxiv.org/abs/2004.05718>`_ paper, using the
|
|
388
|
+
:class:`~k3_node.layers.conv.PNAConv` operator for message passing.
|
|
389
|
+
|
|
390
|
+
Example:
|
|
391
|
+
```python
|
|
392
|
+
import numpy as np
|
|
393
|
+
from k3_node.models import PNA
|
|
394
|
+
|
|
395
|
+
x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
|
|
396
|
+
edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
|
|
397
|
+
|
|
398
|
+
deg = np.array([0, 2, 4, 3, 1]) # in-degree histogram of the training graphs
|
|
399
|
+
model = PNA(in_channels=8, hidden_channels=16, num_layers=2, out_channels=4,
|
|
400
|
+
aggregators=["mean", "min", "max", "std"],
|
|
401
|
+
scalers=["identity", "amplification", "attenuation"], deg=deg)
|
|
402
|
+
out = model(x, edge_index)
|
|
403
|
+
print(tuple(out.shape)) # (10, 4)
|
|
404
|
+
```
|
|
405
|
+
"""
|
|
406
|
+
supports_edge_weight: bool = False
|
|
407
|
+
supports_edge_attr: bool = True
|
|
408
|
+
|
|
409
|
+
def init_conv(self, in_channels: int, out_channels: int, **kwargs) -> MessagePassing:
|
|
410
|
+
return PNAConv(in_channels, out_channels, **kwargs)
|
|
411
|
+
|
|
412
|
+
|
|
413
|
+
class EdgeCNN(BasicGNN):
|
|
414
|
+
r"""The Graph Neural Network from the `"Dynamic Graph CNN for Learning on
|
|
415
|
+
Point Clouds" <https://arxiv.org/abs/1801.07829>`_ paper, using the
|
|
416
|
+
:class:`~k3_node.layers.conv.EdgeConv` operator for message passing.
|
|
417
|
+
|
|
418
|
+
Example:
|
|
419
|
+
```python
|
|
420
|
+
import numpy as np
|
|
421
|
+
from k3_node.models import EdgeCNN
|
|
422
|
+
|
|
423
|
+
x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
|
|
424
|
+
edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
|
|
425
|
+
|
|
426
|
+
model = EdgeCNN(in_channels=8, hidden_channels=16, num_layers=2, out_channels=4)
|
|
427
|
+
out = model(x, edge_index) # e.g. logits for 4 classes per node
|
|
428
|
+
print(tuple(out.shape)) # (10, 4)
|
|
429
|
+
```
|
|
430
|
+
"""
|
|
431
|
+
supports_edge_weight: bool = False
|
|
432
|
+
supports_edge_attr: bool = False
|
|
433
|
+
|
|
434
|
+
def init_conv(self, in_channels: int, out_channels: int, **kwargs) -> MessagePassing:
|
|
435
|
+
mlp = MLP(
|
|
436
|
+
[2 * in_channels, out_channels, out_channels],
|
|
437
|
+
act=self.act,
|
|
438
|
+
act_first=self.act_first,
|
|
439
|
+
norm=self.norm_query,
|
|
440
|
+
norm_kwargs=self.norm_kwargs,
|
|
441
|
+
)
|
|
442
|
+
return EdgeConv(mlp, **kwargs)
|
|
443
|
+
|
k3_node/models/captum.py
ADDED
|
@@ -0,0 +1,52 @@
|
|
|
1
|
+
r"""Captum model interpretability integration.
|
|
2
|
+
|
|
3
|
+
Note: Captum is a PyTorch-specific model interpretability library.
|
|
4
|
+
In k3-node (a multi-backend Keras 3 graph library supporting TensorFlow,
|
|
5
|
+
PyTorch, and JAX), Captum integrations are provided for PyTorch backend
|
|
6
|
+
compatibility or documented as PyTorch-exclusive.
|
|
7
|
+
"""
|
|
8
|
+
|
|
9
|
+
from typing import Optional, Union, Any
|
|
10
|
+
|
|
11
|
+
|
|
12
|
+
def to_captum_model(
|
|
13
|
+
model: Any,
|
|
14
|
+
mask_type: str = "edge",
|
|
15
|
+
output_idx: Optional[int] = None,
|
|
16
|
+
metadata: Optional[Any] = None,
|
|
17
|
+
):
|
|
18
|
+
r"""Converts a model into a Captum-compatible module.
|
|
19
|
+
|
|
20
|
+
.. note::
|
|
21
|
+
Captum is a PyTorch-exclusive library. This function requires
|
|
22
|
+
PyTorch backend and torch.nn.Module models.
|
|
23
|
+
"""
|
|
24
|
+
raise NotImplementedError(
|
|
25
|
+
"Captum integration is PyTorch-specific and requires native torch.nn.Module."
|
|
26
|
+
)
|
|
27
|
+
|
|
28
|
+
|
|
29
|
+
def to_captum_input(
|
|
30
|
+
x: Any,
|
|
31
|
+
edge_index: Any,
|
|
32
|
+
mask_type: str = "edge",
|
|
33
|
+
*args,
|
|
34
|
+
**kwargs,
|
|
35
|
+
):
|
|
36
|
+
r"""Converts graph inputs into Captum-compatible inputs."""
|
|
37
|
+
raise NotImplementedError(
|
|
38
|
+
"Captum integration is PyTorch-specific and requires native torch.nn.Module."
|
|
39
|
+
)
|
|
40
|
+
|
|
41
|
+
|
|
42
|
+
def captum_output_to_dicts(
|
|
43
|
+
attributions: Any,
|
|
44
|
+
mask_type: str = "edge",
|
|
45
|
+
*args,
|
|
46
|
+
**kwargs,
|
|
47
|
+
):
|
|
48
|
+
r"""Converts Captum attributions into dictionaries."""
|
|
49
|
+
raise NotImplementedError(
|
|
50
|
+
"Captum integration is PyTorch-specific and requires native torch.nn.Module."
|
|
51
|
+
)
|
|
52
|
+
|
|
@@ -0,0 +1,146 @@
|
|
|
1
|
+
|
|
2
|
+
import keras
|
|
3
|
+
from keras import ops
|
|
4
|
+
|
|
5
|
+
from k3_node.models.label_prop import LabelPropagation
|
|
6
|
+
|
|
7
|
+
|
|
8
|
+
class CorrectAndSmooth(keras.layers.Layer):
|
|
9
|
+
r"""The correct and smooth (C&S) post-processing model from the
|
|
10
|
+
`"Combining Label Propagation And Simple Models Out-performs Graph Neural
|
|
11
|
+
Networks" <https://arxiv.org/abs/2010.13993>`_ paper.
|
|
12
|
+
|
|
13
|
+
Args:
|
|
14
|
+
num_correction_layers (int): The number of propagations :math:`L_1`.
|
|
15
|
+
correction_alpha (float): The :math:`\alpha_1` coefficient.
|
|
16
|
+
num_smoothing_layers (int): The number of propagations :math:`L_2`.
|
|
17
|
+
smoothing_alpha (float): The :math:`\alpha_2` coefficient.
|
|
18
|
+
autoscale (bool, optional): If set to :obj:`True`, will automatically
|
|
19
|
+
determine the scaling factor :math:`\gamma`. (default: :obj:`True`)
|
|
20
|
+
scale (float, optional): The scaling factor :math:`\gamma`, in case
|
|
21
|
+
:obj:`autoscale = False`. (default: :obj:`1.0`)
|
|
22
|
+
|
|
23
|
+
Example:
|
|
24
|
+
```python
|
|
25
|
+
import numpy as np
|
|
26
|
+
from k3_node.models import CorrectAndSmooth
|
|
27
|
+
|
|
28
|
+
y_soft = np.random.rand(6, 3).astype("float32") # base model's class probabilities
|
|
29
|
+
y_true = np.array([1, 0, 0, 2, 1, 1])
|
|
30
|
+
train_mask = np.array([True, False, True, False, True, False])
|
|
31
|
+
edge_index = np.array([[0, 1, 1, 2, 4, 5], [1, 0, 2, 1, 5, 4]])
|
|
32
|
+
|
|
33
|
+
model = CorrectAndSmooth(num_correction_layers=2, correction_alpha=0.5,
|
|
34
|
+
num_smoothing_layers=2, smoothing_alpha=0.5)
|
|
35
|
+
y_soft = model.correct(y_soft, y_true[train_mask], train_mask, edge_index) # propagate residual errors
|
|
36
|
+
y_soft = model.smooth(y_soft, y_true[train_mask], train_mask, edge_index) # propagate predictions
|
|
37
|
+
print(tuple(y_soft.shape)) # (6, 3)
|
|
38
|
+
```
|
|
39
|
+
"""
|
|
40
|
+
def __init__(
|
|
41
|
+
self,
|
|
42
|
+
num_correction_layers: int,
|
|
43
|
+
correction_alpha: float,
|
|
44
|
+
num_smoothing_layers: int,
|
|
45
|
+
smoothing_alpha: float,
|
|
46
|
+
autoscale: bool = True,
|
|
47
|
+
scale: float = 1.0,
|
|
48
|
+
**kwargs,
|
|
49
|
+
):
|
|
50
|
+
super().__init__(**kwargs)
|
|
51
|
+
self.autoscale = autoscale
|
|
52
|
+
self.scale = scale
|
|
53
|
+
|
|
54
|
+
self.prop1 = LabelPropagation(num_correction_layers, correction_alpha)
|
|
55
|
+
self.prop2 = LabelPropagation(num_smoothing_layers, smoothing_alpha)
|
|
56
|
+
|
|
57
|
+
def build(self, input_shape=None):
|
|
58
|
+
self.built = True
|
|
59
|
+
|
|
60
|
+
def call(self, y_soft, *args, **kwargs):
|
|
61
|
+
y_soft = self.correct(y_soft, *args, **kwargs)
|
|
62
|
+
return self.smooth(y_soft, *args, **kwargs)
|
|
63
|
+
|
|
64
|
+
def correct(self, y_soft, y_true, mask, edge_index, edge_weight=None):
|
|
65
|
+
# Plain NumPy inputs cannot be mixed with backend tensors; convert them first.
|
|
66
|
+
y_soft, y_true, mask = ops.convert_to_tensor(y_soft), ops.convert_to_tensor(y_true), ops.convert_to_tensor(mask)
|
|
67
|
+
num_classes = ops.shape(y_soft)[-1]
|
|
68
|
+
y_true_shape = ops.shape(y_true)
|
|
69
|
+
if len(y_true_shape) == 1:
|
|
70
|
+
y_true = ops.one_hot(y_true, num_classes)
|
|
71
|
+
y_true = ops.cast(y_true, y_soft.dtype)
|
|
72
|
+
|
|
73
|
+
mask_shape = ops.shape(mask)
|
|
74
|
+
is_bool_mask = (len(mask_shape) == 1 and 'bool' in str(mask.dtype))
|
|
75
|
+
if is_bool_mask:
|
|
76
|
+
indices = ops.where(mask)[0]
|
|
77
|
+
else:
|
|
78
|
+
indices = mask
|
|
79
|
+
|
|
80
|
+
numel = float(ops.shape(indices)[0])
|
|
81
|
+
if ops.shape(y_true)[0] == ops.shape(y_soft)[0]:
|
|
82
|
+
y_true_sub = ops.take(y_true, indices, axis=0)
|
|
83
|
+
else:
|
|
84
|
+
y_true_sub = y_true
|
|
85
|
+
|
|
86
|
+
error_zeros = ops.zeros_like(y_soft)
|
|
87
|
+
error = ops.scatter_update(
|
|
88
|
+
error_zeros,
|
|
89
|
+
ops.expand_dims(indices, -1),
|
|
90
|
+
y_true_sub - ops.take(y_soft, indices, axis=0)
|
|
91
|
+
)
|
|
92
|
+
|
|
93
|
+
if self.autoscale:
|
|
94
|
+
smoothed_error = self.prop1(
|
|
95
|
+
error, edge_index, edge_weight=edge_weight,
|
|
96
|
+
post_step=lambda x: ops.clip(x, -1.0, 1.0)
|
|
97
|
+
)
|
|
98
|
+
|
|
99
|
+
error_masked = ops.take(error, indices, axis=0)
|
|
100
|
+
sigma = ops.sum(ops.abs(error_masked)) / max(numel, 1.0)
|
|
101
|
+
sum_smoothed = ops.sum(ops.abs(smoothed_error), axis=1, keepdims=True)
|
|
102
|
+
scale = sigma / ops.maximum(sum_smoothed, 1e-12)
|
|
103
|
+
scale = ops.where(scale > 1000.0, ops.ones_like(scale), scale)
|
|
104
|
+
return y_soft + scale * smoothed_error
|
|
105
|
+
else:
|
|
106
|
+
def fix_input(x):
|
|
107
|
+
return ops.scatter_update(x, ops.expand_dims(indices, -1), ops.take(error, indices, axis=0))
|
|
108
|
+
|
|
109
|
+
smoothed_error = self.prop1(
|
|
110
|
+
error, edge_index, edge_weight=edge_weight,
|
|
111
|
+
post_step=fix_input,
|
|
112
|
+
)
|
|
113
|
+
return y_soft + self.scale * smoothed_error
|
|
114
|
+
|
|
115
|
+
def smooth(self, y_soft, y_true, mask, edge_index, edge_weight=None):
|
|
116
|
+
# Plain NumPy inputs cannot be mixed with backend tensors; convert them first.
|
|
117
|
+
y_soft, y_true, mask = ops.convert_to_tensor(y_soft), ops.convert_to_tensor(y_true), ops.convert_to_tensor(mask)
|
|
118
|
+
num_classes = ops.shape(y_soft)[-1]
|
|
119
|
+
y_true_shape = ops.shape(y_true)
|
|
120
|
+
if len(y_true_shape) == 1:
|
|
121
|
+
y_true = ops.one_hot(y_true, num_classes)
|
|
122
|
+
y_true = ops.cast(y_true, y_soft.dtype)
|
|
123
|
+
|
|
124
|
+
mask_shape = ops.shape(mask)
|
|
125
|
+
is_bool_mask = (len(mask_shape) == 1 and 'bool' in str(mask.dtype))
|
|
126
|
+
if is_bool_mask:
|
|
127
|
+
indices = ops.where(mask)[0]
|
|
128
|
+
else:
|
|
129
|
+
indices = mask
|
|
130
|
+
|
|
131
|
+
if ops.shape(y_true)[0] == ops.shape(y_soft)[0]:
|
|
132
|
+
y_true_sub = ops.take(y_true, indices, axis=0)
|
|
133
|
+
else:
|
|
134
|
+
y_true_sub = y_true
|
|
135
|
+
|
|
136
|
+
y_soft = ops.scatter_update(y_soft, ops.expand_dims(indices, -1), y_true_sub)
|
|
137
|
+
return self.prop2(y_soft, edge_index, edge_weight=edge_weight)
|
|
138
|
+
|
|
139
|
+
def __repr__(self):
|
|
140
|
+
L1, alpha1 = self.prop1.num_layers, self.prop1.alpha
|
|
141
|
+
L2, alpha2 = self.prop2.num_layers, self.prop2.alpha
|
|
142
|
+
return (f'{self.__class__.__name__}(\n'
|
|
143
|
+
f' correct: num_layers={L1}, alpha={alpha1}\n'
|
|
144
|
+
f' smooth: num_layers={L2}, alpha={alpha2}\n'
|
|
145
|
+
f' autoscale={self.autoscale}, scale={self.scale}\n'
|
|
146
|
+
')')
|