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,149 @@
|
|
|
1
|
+
from .message_passing import MessagePassing
|
|
2
|
+
from .simple_conv import SimpleConv
|
|
3
|
+
from .gcn_conv import GCNConv
|
|
4
|
+
from .cheb_conv import ChebConv
|
|
5
|
+
from .sage_conv import SAGEConv
|
|
6
|
+
from .graph_conv import GraphConv
|
|
7
|
+
from .gated_graph_conv import GatedGraphConv
|
|
8
|
+
from .res_gated_graph_conv import ResGatedGraphConv
|
|
9
|
+
from .gat_conv import GATConv, FusedGATConv
|
|
10
|
+
from .gatv2_conv import GATv2Conv
|
|
11
|
+
from .transformer_conv import TransformerConv
|
|
12
|
+
from .agnn_conv import AGNNConv
|
|
13
|
+
from .tag_conv import TAGConv
|
|
14
|
+
from .gin_conv import GINConv, GINEConv
|
|
15
|
+
from .arma_conv import ARMAConv
|
|
16
|
+
from .sg_conv import SGConv
|
|
17
|
+
from .ssg_conv import SSGConv
|
|
18
|
+
from .appnp import APPNP
|
|
19
|
+
from .appnp_conv import APPNPConv
|
|
20
|
+
from .mf_conv import MFConv
|
|
21
|
+
from .rgcn_conv import RGCNConv, FastRGCNConv, CuGraphRGCNConv
|
|
22
|
+
from .rgat_conv import RGATConv
|
|
23
|
+
from .signed_conv import SignedConv
|
|
24
|
+
from .dir_gnn_conv import DirGNNConv
|
|
25
|
+
from .antisymmetric_conv import AntiSymmetricConv
|
|
26
|
+
from .mixhop_conv import MixHopConv
|
|
27
|
+
from .pdn_conv import PDNConv
|
|
28
|
+
from .fa_conv import FAConv
|
|
29
|
+
from .film_conv import FiLMConv
|
|
30
|
+
from .supergat_conv import SuperGATConv
|
|
31
|
+
from .eg_conv import EGConv
|
|
32
|
+
from .pan_conv import PANConv
|
|
33
|
+
from .gen_conv import GENConv
|
|
34
|
+
from .pna_conv import PNAConv
|
|
35
|
+
from .le_conv import LEConv
|
|
36
|
+
from .cluster_gcn_conv import ClusterGCNConv
|
|
37
|
+
from .gcn2_conv import GCN2Conv
|
|
38
|
+
from .lg_conv import LGConv
|
|
39
|
+
from .nn_conv import NNConv
|
|
40
|
+
from .cg_conv import CGConv
|
|
41
|
+
from .edge_conv import EdgeConv, DynamicEdgeConv
|
|
42
|
+
from .general_conv import GeneralConv
|
|
43
|
+
from .point_conv import PointNetConv, PointConv
|
|
44
|
+
from .point_transformer_conv import PointTransformerConv
|
|
45
|
+
from .point_gnn_conv import PointGNNConv
|
|
46
|
+
from .ppf_conv import PPFConv
|
|
47
|
+
from .feast_conv import FeaStConv
|
|
48
|
+
from .gmm_conv import GMMConv
|
|
49
|
+
from .gravnet_conv import GravNetConv
|
|
50
|
+
from .meshcnn_conv import MeshCNNConv
|
|
51
|
+
from .x_conv import XConv
|
|
52
|
+
from .spline_conv import SplineConv
|
|
53
|
+
from .hetero_conv import HeteroConv
|
|
54
|
+
from .hgt_conv import HGTConv
|
|
55
|
+
from .han_conv import HANConv
|
|
56
|
+
from .heat_conv import HEATConv
|
|
57
|
+
from .hypergraph_conv import HypergraphConv
|
|
58
|
+
from .dna_conv import DNAConv
|
|
59
|
+
from .wl_conv import WLConv, WLConvContinuous
|
|
60
|
+
from .gps_conv import GPSConv
|
|
61
|
+
from .cugraph import CuGraphGATConv, CuGraphSAGEConv
|
|
62
|
+
|
|
63
|
+
ECConv = NNConv
|
|
64
|
+
|
|
65
|
+
# Legacy Spektral imports
|
|
66
|
+
from .crystal_conv import CrystalConv
|
|
67
|
+
from .diffusion_conv import DiffusionConv
|
|
68
|
+
from .gcn import GraphConvolution
|
|
69
|
+
from .graph_attention import GraphAttention
|
|
70
|
+
from .ppnp import PPNPPropagation
|
|
71
|
+
|
|
72
|
+
__all__ = [
|
|
73
|
+
"MessagePassing",
|
|
74
|
+
"SimpleConv",
|
|
75
|
+
"GCNConv",
|
|
76
|
+
"ChebConv",
|
|
77
|
+
"SAGEConv",
|
|
78
|
+
"GraphConv",
|
|
79
|
+
"GatedGraphConv",
|
|
80
|
+
"ResGatedGraphConv",
|
|
81
|
+
"GATConv",
|
|
82
|
+
"FusedGATConv",
|
|
83
|
+
"GATv2Conv",
|
|
84
|
+
"TransformerConv",
|
|
85
|
+
"AGNNConv",
|
|
86
|
+
"TAGConv",
|
|
87
|
+
"GINConv",
|
|
88
|
+
"GINEConv",
|
|
89
|
+
"ARMAConv",
|
|
90
|
+
"SGConv",
|
|
91
|
+
"SSGConv",
|
|
92
|
+
"APPNP",
|
|
93
|
+
"APPNPConv",
|
|
94
|
+
"MFConv",
|
|
95
|
+
"RGCNConv",
|
|
96
|
+
"FastRGCNConv",
|
|
97
|
+
"CuGraphRGCNConv",
|
|
98
|
+
"RGATConv",
|
|
99
|
+
"SignedConv",
|
|
100
|
+
"DirGNNConv",
|
|
101
|
+
"AntiSymmetricConv",
|
|
102
|
+
"MixHopConv",
|
|
103
|
+
"PDNConv",
|
|
104
|
+
"FAConv",
|
|
105
|
+
"FiLMConv",
|
|
106
|
+
"SuperGATConv",
|
|
107
|
+
"EGConv",
|
|
108
|
+
"PANConv",
|
|
109
|
+
"GENConv",
|
|
110
|
+
"PNAConv",
|
|
111
|
+
"LEConv",
|
|
112
|
+
"ClusterGCNConv",
|
|
113
|
+
"GCN2Conv",
|
|
114
|
+
"LGConv",
|
|
115
|
+
"NNConv",
|
|
116
|
+
"ECConv",
|
|
117
|
+
"CGConv",
|
|
118
|
+
"EdgeConv",
|
|
119
|
+
"DynamicEdgeConv",
|
|
120
|
+
"GeneralConv",
|
|
121
|
+
"PointNetConv",
|
|
122
|
+
"PointConv",
|
|
123
|
+
"PointTransformerConv",
|
|
124
|
+
"PointGNNConv",
|
|
125
|
+
"PPFConv",
|
|
126
|
+
"FeaStConv",
|
|
127
|
+
"GMMConv",
|
|
128
|
+
"GravNetConv",
|
|
129
|
+
"MeshCNNConv",
|
|
130
|
+
"XConv",
|
|
131
|
+
"SplineConv",
|
|
132
|
+
"HeteroConv",
|
|
133
|
+
"HGTConv",
|
|
134
|
+
"HANConv",
|
|
135
|
+
"HEATConv",
|
|
136
|
+
"HypergraphConv",
|
|
137
|
+
"DNAConv",
|
|
138
|
+
"WLConv",
|
|
139
|
+
"WLConvContinuous",
|
|
140
|
+
"GPSConv",
|
|
141
|
+
"CuGraphGATConv",
|
|
142
|
+
"CuGraphSAGEConv",
|
|
143
|
+
# Legacy Spektral
|
|
144
|
+
"CrystalConv",
|
|
145
|
+
"DiffusionConv",
|
|
146
|
+
"GraphConvolution",
|
|
147
|
+
"GraphAttention",
|
|
148
|
+
"PPNPPropagation",
|
|
149
|
+
]
|
|
@@ -0,0 +1,120 @@
|
|
|
1
|
+
from typing import Optional
|
|
2
|
+
from keras import ops
|
|
3
|
+
|
|
4
|
+
from k3_node.layers.conv.message_passing import MessagePassing
|
|
5
|
+
from k3_node.layers.conv.utils import (
|
|
6
|
+
add_self_loops,
|
|
7
|
+
extend_mask_for_self_loops,
|
|
8
|
+
mask_edge_logits,
|
|
9
|
+
remove_self_loops_masked,
|
|
10
|
+
softmax,
|
|
11
|
+
)
|
|
12
|
+
|
|
13
|
+
|
|
14
|
+
class AGNNConv(MessagePassing):
|
|
15
|
+
r"""The graph attentional propagation layer from the
|
|
16
|
+
`"Attention-based Graph Neural Network for Semi-Supervised Learning"
|
|
17
|
+
<https://arxiv.org/abs/1803.03735>`_ paper.
|
|
18
|
+
|
|
19
|
+
Example:
|
|
20
|
+
```python
|
|
21
|
+
import numpy as np
|
|
22
|
+
from k3_node.layers import AGNNConv
|
|
23
|
+
|
|
24
|
+
x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
|
|
25
|
+
edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
|
|
26
|
+
|
|
27
|
+
layer = AGNNConv(requires_grad=True)
|
|
28
|
+
out = layer(x, edge_index)
|
|
29
|
+
print(tuple(out.shape)) # (10, 8)
|
|
30
|
+
```
|
|
31
|
+
"""
|
|
32
|
+
def __init__(
|
|
33
|
+
self,
|
|
34
|
+
requires_grad: bool = True,
|
|
35
|
+
add_self_loops: bool = True,
|
|
36
|
+
trainable: Optional[bool] = None,
|
|
37
|
+
aggregate: str = "add",
|
|
38
|
+
activation=None,
|
|
39
|
+
**kwargs,
|
|
40
|
+
):
|
|
41
|
+
if trainable is not None:
|
|
42
|
+
requires_grad = trainable
|
|
43
|
+
kwargs.setdefault("aggr", aggregate)
|
|
44
|
+
super().__init__(activation=activation, **kwargs)
|
|
45
|
+
|
|
46
|
+
self.requires_grad = requires_grad
|
|
47
|
+
self.add_self_loops = add_self_loops
|
|
48
|
+
|
|
49
|
+
if requires_grad:
|
|
50
|
+
self.beta = self.add_weight(
|
|
51
|
+
shape=(1,),
|
|
52
|
+
initializer="ones",
|
|
53
|
+
name="beta",
|
|
54
|
+
)
|
|
55
|
+
else:
|
|
56
|
+
self.beta = None
|
|
57
|
+
|
|
58
|
+
def build(self, input_shape=None):
|
|
59
|
+
self.built = True
|
|
60
|
+
|
|
61
|
+
def call(self, inputs, edge_index=None, **kwargs):
|
|
62
|
+
if edge_index is None:
|
|
63
|
+
if isinstance(inputs, (list, tuple)) and len(inputs) == 2:
|
|
64
|
+
x, edge_index = inputs
|
|
65
|
+
else:
|
|
66
|
+
raise ValueError("Expected (x, edge_index) or x and edge_index")
|
|
67
|
+
else:
|
|
68
|
+
x = inputs
|
|
69
|
+
|
|
70
|
+
if not self.built:
|
|
71
|
+
self.build()
|
|
72
|
+
|
|
73
|
+
# Check for legacy 2D matrix
|
|
74
|
+
is_legacy = False
|
|
75
|
+
if hasattr(edge_index, "shape") and len(edge_index.shape) == 2:
|
|
76
|
+
if (
|
|
77
|
+
edge_index.shape[0] is not None
|
|
78
|
+
and edge_index.shape[1] is not None
|
|
79
|
+
and edge_index.shape[0] != 2
|
|
80
|
+
and edge_index.shape[0] == edge_index.shape[1]
|
|
81
|
+
):
|
|
82
|
+
is_legacy = True
|
|
83
|
+
elif not hasattr(edge_index, "shape"):
|
|
84
|
+
is_legacy = True
|
|
85
|
+
|
|
86
|
+
x_norm = x / (ops.norm(x, axis=-1, keepdims=True) + 1e-12)
|
|
87
|
+
|
|
88
|
+
if is_legacy:
|
|
89
|
+
out = self.propagate(x, edge_index, x_norm=x_norm)
|
|
90
|
+
else:
|
|
91
|
+
num_nodes = x.shape[0] if hasattr(x, "shape") and x.shape[0] is not None else ops.shape(x)[0]
|
|
92
|
+
keep_mask = None
|
|
93
|
+
if self.add_self_loops:
|
|
94
|
+
edge_index, _, keep_mask = remove_self_loops_masked(edge_index)
|
|
95
|
+
edge_index, _ = add_self_loops(edge_index, num_nodes=num_nodes)
|
|
96
|
+
keep_mask = extend_mask_for_self_loops(keep_mask, num_nodes)
|
|
97
|
+
out = self.propagate(edge_index, x=x, x_norm=x_norm, keep_mask=keep_mask, size=(num_nodes, num_nodes))
|
|
98
|
+
|
|
99
|
+
if self.activation is not None:
|
|
100
|
+
out = self.activation(out)
|
|
101
|
+
|
|
102
|
+
return out
|
|
103
|
+
|
|
104
|
+
def message(self, x=None, x_j=None, x_norm=None, x_norm_i=None, x_norm_j=None, index=None, size_i=None, keep_mask=None):
|
|
105
|
+
beta = self.beta if self.beta is not None else 1.0
|
|
106
|
+
|
|
107
|
+
# Legacy Spektral path
|
|
108
|
+
if x_j is None and x is not None:
|
|
109
|
+
x_j = self.get_sources(x)
|
|
110
|
+
x_norm_i = self.get_targets(x_norm)
|
|
111
|
+
x_norm_j = self.get_sources(x_norm)
|
|
112
|
+
alpha = beta * ops.sum(x_norm_i * x_norm_j, axis=-1)
|
|
113
|
+
alpha = softmax(alpha, self.index_targets, num_nodes=self.n_nodes, dim=0)
|
|
114
|
+
return ops.expand_dims(alpha, -1) * x_j
|
|
115
|
+
|
|
116
|
+
# PyG path
|
|
117
|
+
alpha = beta * ops.sum(x_norm_i * x_norm_j, axis=-1)
|
|
118
|
+
alpha = mask_edge_logits(alpha, keep_mask)
|
|
119
|
+
alpha = softmax(alpha, index, num_nodes=size_i, dim=0)
|
|
120
|
+
return x_j * ops.expand_dims(alpha, -1)
|
|
@@ -0,0 +1,94 @@
|
|
|
1
|
+
from typing import Optional, Union, Callable
|
|
2
|
+
from keras import ops, activations
|
|
3
|
+
from keras.layers import Layer
|
|
4
|
+
|
|
5
|
+
from k3_node.layers.conv.message_passing import MessagePassing
|
|
6
|
+
from k3_node.layers.conv.gcn_conv import GCNConv
|
|
7
|
+
|
|
8
|
+
|
|
9
|
+
class AntiSymmetricConv(Layer):
|
|
10
|
+
r"""The anti-symmetric graph convolutional operator from the
|
|
11
|
+
`"Anti-Symmetric DGN: a Continuous approach to Deep Graph Neural Networks"
|
|
12
|
+
<https://arxiv.org/abs/2202.13085>`_ paper.
|
|
13
|
+
|
|
14
|
+
Example:
|
|
15
|
+
```python
|
|
16
|
+
import numpy as np
|
|
17
|
+
from k3_node.layers import AntiSymmetricConv
|
|
18
|
+
|
|
19
|
+
x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
|
|
20
|
+
edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
|
|
21
|
+
|
|
22
|
+
layer = AntiSymmetricConv(in_channels=8)
|
|
23
|
+
out = layer(x, edge_index)
|
|
24
|
+
print(tuple(out.shape)) # (10, 8)
|
|
25
|
+
```
|
|
26
|
+
"""
|
|
27
|
+
def __init__(
|
|
28
|
+
self,
|
|
29
|
+
in_channels: int,
|
|
30
|
+
phi: Optional[MessagePassing] = None,
|
|
31
|
+
num_iters: int = 1,
|
|
32
|
+
epsilon: float = 0.1,
|
|
33
|
+
gamma: float = 0.1,
|
|
34
|
+
act: Union[str, Callable, None] = "tanh",
|
|
35
|
+
bias: bool = True,
|
|
36
|
+
**kwargs,
|
|
37
|
+
):
|
|
38
|
+
super().__init__(**kwargs)
|
|
39
|
+
|
|
40
|
+
self.in_channels = in_channels
|
|
41
|
+
self.num_iters = num_iters
|
|
42
|
+
self.gamma = gamma
|
|
43
|
+
self.epsilon = epsilon
|
|
44
|
+
self.act = activations.get(act) if act is not None else None
|
|
45
|
+
|
|
46
|
+
if phi is None:
|
|
47
|
+
phi = GCNConv(in_channels, in_channels, bias=False)
|
|
48
|
+
self.phi = phi
|
|
49
|
+
|
|
50
|
+
self.W = self.add_weight(
|
|
51
|
+
shape=(in_channels, in_channels),
|
|
52
|
+
initializer="glorot_uniform",
|
|
53
|
+
name="W",
|
|
54
|
+
)
|
|
55
|
+
|
|
56
|
+
if bias:
|
|
57
|
+
self.bias = self.add_weight(
|
|
58
|
+
shape=(in_channels,),
|
|
59
|
+
initializer="zeros",
|
|
60
|
+
name="bias",
|
|
61
|
+
)
|
|
62
|
+
else:
|
|
63
|
+
self.bias = None
|
|
64
|
+
|
|
65
|
+
def build(self, input_shape=None):
|
|
66
|
+
if hasattr(self.phi, "build"):
|
|
67
|
+
self.phi.build((None, self.in_channels))
|
|
68
|
+
self.built = True
|
|
69
|
+
|
|
70
|
+
def call(self, inputs, edge_index=None, **kwargs):
|
|
71
|
+
if edge_index is None:
|
|
72
|
+
if isinstance(inputs, (list, tuple)) and len(inputs) == 2:
|
|
73
|
+
x, edge_index = inputs
|
|
74
|
+
else:
|
|
75
|
+
raise ValueError("Expected (x, edge_index) or x and edge_index")
|
|
76
|
+
else:
|
|
77
|
+
x = inputs
|
|
78
|
+
|
|
79
|
+
eye = ops.eye(self.in_channels, dtype=self.W.dtype)
|
|
80
|
+
antisymmetric_W = self.W - ops.transpose(self.W) - self.gamma * eye
|
|
81
|
+
|
|
82
|
+
for _ in range(self.num_iters):
|
|
83
|
+
h = self.phi(x, edge_index)
|
|
84
|
+
h = ops.matmul(x, ops.transpose(antisymmetric_W)) + h
|
|
85
|
+
|
|
86
|
+
if self.bias is not None:
|
|
87
|
+
h = h + self.bias
|
|
88
|
+
|
|
89
|
+
if self.act is not None:
|
|
90
|
+
h = self.act(h)
|
|
91
|
+
|
|
92
|
+
x = x + self.epsilon * h
|
|
93
|
+
|
|
94
|
+
return x
|
|
@@ -0,0 +1,105 @@
|
|
|
1
|
+
from keras import layers, ops
|
|
2
|
+
|
|
3
|
+
from k3_node.layers.conv.message_passing import MessagePassing
|
|
4
|
+
from k3_node.layers.conv.utils import gcn_norm, is_tracing
|
|
5
|
+
|
|
6
|
+
|
|
7
|
+
class APPNP(MessagePassing):
|
|
8
|
+
r"""The approximate personalized propagation of neural predictions (APPNP)
|
|
9
|
+
operator from the `"Predict then Propagate: Combining Neural Networks with
|
|
10
|
+
Personalized PageRank for Classification on Graphs"
|
|
11
|
+
<https://arxiv.org/abs/1810.05997>`_ paper.
|
|
12
|
+
|
|
13
|
+
Args:
|
|
14
|
+
K: Number of iterations :math:`K`.
|
|
15
|
+
alpha: Teleport probability :math:`\alpha`.
|
|
16
|
+
dropout: Dropout probability of edges or features during propagation.
|
|
17
|
+
(default: ``0.0``)
|
|
18
|
+
cached: If set to :obj:`True`, the layer will cache the computation of
|
|
19
|
+
normalization coefficients. (default: ``False``)
|
|
20
|
+
add_self_loops: If set to :obj:`False`, will not add self-loops.
|
|
21
|
+
(default: ``True``)
|
|
22
|
+
normalize: Whether to apply symmetric normalization. (default: ``True``)
|
|
23
|
+
|
|
24
|
+
Example:
|
|
25
|
+
```python
|
|
26
|
+
import numpy as np
|
|
27
|
+
from k3_node.layers import APPNP
|
|
28
|
+
|
|
29
|
+
x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
|
|
30
|
+
edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
|
|
31
|
+
|
|
32
|
+
layer = APPNP(K=2, alpha=0.1)
|
|
33
|
+
out = layer(x, edge_index)
|
|
34
|
+
print(tuple(out.shape)) # (10, 8)
|
|
35
|
+
```
|
|
36
|
+
"""
|
|
37
|
+
|
|
38
|
+
weighted_sum_message = True
|
|
39
|
+
|
|
40
|
+
def __init__(
|
|
41
|
+
self,
|
|
42
|
+
K: int,
|
|
43
|
+
alpha: float,
|
|
44
|
+
dropout: float = 0.0,
|
|
45
|
+
cached: bool = False,
|
|
46
|
+
add_self_loops: bool = True,
|
|
47
|
+
normalize: bool = True,
|
|
48
|
+
**kwargs,
|
|
49
|
+
):
|
|
50
|
+
super().__init__(aggr="add", **kwargs)
|
|
51
|
+
self.K = K
|
|
52
|
+
self.alpha = alpha
|
|
53
|
+
self.dropout_rate = dropout
|
|
54
|
+
self.cached = cached
|
|
55
|
+
self.add_self_loops = add_self_loops
|
|
56
|
+
self.normalize = normalize
|
|
57
|
+
self._cached_edge_index = None
|
|
58
|
+
self._cached_norm = None
|
|
59
|
+
self.dropout = layers.Dropout(dropout) if dropout > 0.0 else None
|
|
60
|
+
|
|
61
|
+
def build(self, input_shape=None):
|
|
62
|
+
if self.dropout is not None and hasattr(self.dropout, "build"):
|
|
63
|
+
self.dropout.build(input_shape)
|
|
64
|
+
self.built = True
|
|
65
|
+
|
|
66
|
+
def call(self, x, edge_index=None, edge_weight=None, training=None, **kwargs):
|
|
67
|
+
if edge_index is None and isinstance(x, (tuple, list)):
|
|
68
|
+
x, edge_index = x[0], x[1]
|
|
69
|
+
|
|
70
|
+
if self.normalize:
|
|
71
|
+
if self.cached and self._cached_edge_index is not None:
|
|
72
|
+
edge_index = self._cached_edge_index
|
|
73
|
+
edge_weight = self._cached_norm
|
|
74
|
+
else:
|
|
75
|
+
num_nodes = x.shape[self.node_dim] if hasattr(x, "shape") and x.shape[self.node_dim] is not None else ops.shape(x)[self.node_dim]
|
|
76
|
+
edge_index, edge_weight = gcn_norm(
|
|
77
|
+
edge_index,
|
|
78
|
+
edge_weight,
|
|
79
|
+
num_nodes=num_nodes,
|
|
80
|
+
add_self_loops=self.add_self_loops,
|
|
81
|
+
flow=self.flow,
|
|
82
|
+
dtype=x.dtype,
|
|
83
|
+
)
|
|
84
|
+
if self.cached and not is_tracing(edge_index):
|
|
85
|
+
self._cached_edge_index = edge_index
|
|
86
|
+
self._cached_norm = edge_weight
|
|
87
|
+
|
|
88
|
+
h = x
|
|
89
|
+
for _ in range(self.K):
|
|
90
|
+
if self.dropout is not None:
|
|
91
|
+
h = self.dropout(h, training=training)
|
|
92
|
+
h = self.propagate(edge_index, x=h, edge_weight=edge_weight)
|
|
93
|
+
h = (1.0 - self.alpha) * h + self.alpha * x
|
|
94
|
+
|
|
95
|
+
return h
|
|
96
|
+
|
|
97
|
+
def message(self, x_j, edge_weight=None):
|
|
98
|
+
if edge_weight is None:
|
|
99
|
+
return x_j
|
|
100
|
+
return ops.expand_dims(edge_weight, -1) * x_j
|
|
101
|
+
|
|
102
|
+
|
|
103
|
+
# Backward-compatible alias
|
|
104
|
+
APPNPConv = APPNP
|
|
105
|
+
|
|
@@ -0,0 +1,157 @@
|
|
|
1
|
+
# ported from spektral
|
|
2
|
+
|
|
3
|
+
from keras import activations
|
|
4
|
+
from keras import activations, ops
|
|
5
|
+
from keras.layers import Dense, Dropout
|
|
6
|
+
from keras.models import Sequential
|
|
7
|
+
|
|
8
|
+
from k3_node.layers.conv.conv import Conv
|
|
9
|
+
from k3_node.ops import gcn_filter, modal_dot
|
|
10
|
+
|
|
11
|
+
|
|
12
|
+
class APPNPConv(Conv):
|
|
13
|
+
"""
|
|
14
|
+
`k3_node.layers.APPNPConv`
|
|
15
|
+
Implementation of Approximate Personalized Propagation of Neural Predictions
|
|
16
|
+
|
|
17
|
+
Args:
|
|
18
|
+
channels: The number of output channels.
|
|
19
|
+
alpha: The teleport probability.
|
|
20
|
+
propagations: The number of propagation steps.
|
|
21
|
+
mlp_hidden: A list of hidden channels for the MLP.
|
|
22
|
+
mlp_activation: The activation function to use in the MLP.
|
|
23
|
+
dropout_rate: The dropout rate for the MLP.
|
|
24
|
+
activation: The activation function to use in the layer.
|
|
25
|
+
use_bias: Whether to add a bias to the linear transformation.
|
|
26
|
+
kernel_initializer: Initializer for the `kernel` weights matrix.
|
|
27
|
+
bias_initializer: Initializer for the bias vector.
|
|
28
|
+
kernel_regularizer: Regularizer for the `kernel` weights matrix.
|
|
29
|
+
bias_regularizer: Regularizer for the bias vector.
|
|
30
|
+
activity_regularizer: Regularizer for the output.
|
|
31
|
+
kernel_constraint: Constraint for the `kernel` weights matrix.
|
|
32
|
+
bias_constraint: Constraint for the bias vector.
|
|
33
|
+
**kwargs: Additional keyword arguments.
|
|
34
|
+
|
|
35
|
+
Example:
|
|
36
|
+
```python
|
|
37
|
+
import numpy as np
|
|
38
|
+
from k3_node.layers import APPNPConv
|
|
39
|
+
|
|
40
|
+
x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
|
|
41
|
+
edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
|
|
42
|
+
|
|
43
|
+
layer = APPNPConv(channels=16, alpha=0.1, propagations=2)
|
|
44
|
+
out = layer(x, edge_index)
|
|
45
|
+
print(tuple(out.shape)) # (10, 16)
|
|
46
|
+
```
|
|
47
|
+
"""
|
|
48
|
+
def __init__(
|
|
49
|
+
self,
|
|
50
|
+
channels,
|
|
51
|
+
alpha=0.2,
|
|
52
|
+
propagations=1,
|
|
53
|
+
mlp_hidden=None,
|
|
54
|
+
mlp_activation="relu",
|
|
55
|
+
dropout_rate=0.0,
|
|
56
|
+
activation=None,
|
|
57
|
+
use_bias=True,
|
|
58
|
+
kernel_initializer="glorot_uniform",
|
|
59
|
+
bias_initializer="zeros",
|
|
60
|
+
kernel_regularizer=None,
|
|
61
|
+
bias_regularizer=None,
|
|
62
|
+
activity_regularizer=None,
|
|
63
|
+
kernel_constraint=None,
|
|
64
|
+
bias_constraint=None,
|
|
65
|
+
**kwargs,
|
|
66
|
+
):
|
|
67
|
+
|
|
68
|
+
super().__init__(
|
|
69
|
+
activation=activation,
|
|
70
|
+
use_bias=use_bias,
|
|
71
|
+
kernel_initializer=kernel_initializer,
|
|
72
|
+
bias_initializer=bias_initializer,
|
|
73
|
+
kernel_regularizer=kernel_regularizer,
|
|
74
|
+
bias_regularizer=bias_regularizer,
|
|
75
|
+
activity_regularizer=activity_regularizer,
|
|
76
|
+
kernel_constraint=kernel_constraint,
|
|
77
|
+
bias_constraint=bias_constraint,
|
|
78
|
+
**kwargs,
|
|
79
|
+
)
|
|
80
|
+
self.channels = channels
|
|
81
|
+
self.mlp_hidden = mlp_hidden if mlp_hidden else []
|
|
82
|
+
self.alpha = alpha
|
|
83
|
+
self.propagations = propagations
|
|
84
|
+
self.mlp_activation = activations.get(mlp_activation)
|
|
85
|
+
self.dropout_rate = dropout_rate
|
|
86
|
+
|
|
87
|
+
def build(self, input_shape=None):
|
|
88
|
+
layer_kwargs = dict(
|
|
89
|
+
kernel_initializer=self.kernel_initializer,
|
|
90
|
+
bias_initializer=self.bias_initializer,
|
|
91
|
+
kernel_regularizer=self.kernel_regularizer,
|
|
92
|
+
bias_regularizer=self.bias_regularizer,
|
|
93
|
+
kernel_constraint=self.kernel_constraint,
|
|
94
|
+
bias_constraint=self.bias_constraint,
|
|
95
|
+
dtype=self.dtype,
|
|
96
|
+
)
|
|
97
|
+
mlp_layers = []
|
|
98
|
+
for channels in self.mlp_hidden:
|
|
99
|
+
mlp_layers.extend(
|
|
100
|
+
[
|
|
101
|
+
Dropout(self.dropout_rate),
|
|
102
|
+
Dense(channels, self.mlp_activation, **layer_kwargs),
|
|
103
|
+
]
|
|
104
|
+
)
|
|
105
|
+
mlp_layers.append(Dense(self.channels, "linear", **layer_kwargs))
|
|
106
|
+
self.mlp = Sequential(mlp_layers)
|
|
107
|
+
if input_shape is not None:
|
|
108
|
+
feat_shape = input_shape[0] if isinstance(input_shape, (list, tuple)) else input_shape
|
|
109
|
+
self.mlp.build(feat_shape)
|
|
110
|
+
self.built = True
|
|
111
|
+
|
|
112
|
+
def call(self, inputs, mask=None):
|
|
113
|
+
x, a = inputs
|
|
114
|
+
def call(self, inputs, a=None, mask=None):
|
|
115
|
+
if a is not None:
|
|
116
|
+
x = inputs
|
|
117
|
+
elif isinstance(inputs, (list, tuple)) and len(inputs) == 2:
|
|
118
|
+
x, a = inputs
|
|
119
|
+
else:
|
|
120
|
+
x = inputs
|
|
121
|
+
a = None
|
|
122
|
+
|
|
123
|
+
if not self.built:
|
|
124
|
+
self.build(getattr(x, "shape", None))
|
|
125
|
+
|
|
126
|
+
if a is not None and hasattr(a, "shape") and len(a.shape) == 2 and a.shape[0] == 2 and a.shape[1] != 2:
|
|
127
|
+
num_nodes = ops.shape(x)[-2] if len(ops.shape(x)) >= 2 else ops.shape(x)[0]
|
|
128
|
+
row, col = ops.cast(a[0], "int32"), ops.cast(a[1], "int32")
|
|
129
|
+
idx = ops.stack([row, col], axis=-1)
|
|
130
|
+
zeros = ops.zeros((num_nodes, num_nodes), dtype=x.dtype)
|
|
131
|
+
a = ops.scatter_update(zeros, idx, ops.ones((ops.shape(a)[1],), dtype=x.dtype))
|
|
132
|
+
|
|
133
|
+
mlp_out = self.mlp(x)
|
|
134
|
+
output = mlp_out
|
|
135
|
+
if a is not None:
|
|
136
|
+
for _ in range(self.propagations):
|
|
137
|
+
output = (1 - self.alpha) * modal_dot(a, output) + self.alpha * mlp_out
|
|
138
|
+
if mask is not None and isinstance(mask, (list, tuple)) and len(mask) > 0 and mask[0] is not None:
|
|
139
|
+
output *= mask[0]
|
|
140
|
+
output = self.activation(output)
|
|
141
|
+
|
|
142
|
+
return output
|
|
143
|
+
|
|
144
|
+
@property
|
|
145
|
+
def config(self):
|
|
146
|
+
return {
|
|
147
|
+
"channels": self.channels,
|
|
148
|
+
"alpha": self.alpha,
|
|
149
|
+
"propagations": self.propagations,
|
|
150
|
+
"mlp_hidden": self.mlp_hidden,
|
|
151
|
+
"mlp_activation": activations.serialize(self.mlp_activation),
|
|
152
|
+
"dropout_rate": self.dropout_rate,
|
|
153
|
+
}
|
|
154
|
+
|
|
155
|
+
@staticmethod
|
|
156
|
+
def preprocess(a):
|
|
157
|
+
return gcn_filter(a)
|