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,231 @@
|
|
|
1
|
+
from typing import Optional, Union, Callable
|
|
2
|
+
from keras import ops, activations
|
|
3
|
+
from keras.layers import Dropout
|
|
4
|
+
|
|
5
|
+
from k3_node.layers.conv.message_passing import MessagePassing
|
|
6
|
+
from k3_node.layers.conv.utils import gcn_norm
|
|
7
|
+
|
|
8
|
+
|
|
9
|
+
class ARMAConv(MessagePassing):
|
|
10
|
+
r"""The ARMA graph convolutional operator from the `"Graph Neural Networks
|
|
11
|
+
with Convolutional ARMA Filters" <https://arxiv.org/abs/1901.01343>`_
|
|
12
|
+
paper.
|
|
13
|
+
|
|
14
|
+
Example:
|
|
15
|
+
```python
|
|
16
|
+
import numpy as np
|
|
17
|
+
from k3_node.layers import ARMAConv
|
|
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 = ARMAConv(in_channels=8, out_channels=16, num_stacks=1, num_layers=1)
|
|
23
|
+
out = layer(x, edge_index)
|
|
24
|
+
print(tuple(out.shape)) # (10, 16)
|
|
25
|
+
```
|
|
26
|
+
"""
|
|
27
|
+
def __init__(
|
|
28
|
+
self,
|
|
29
|
+
in_channels: int,
|
|
30
|
+
out_channels: Optional[int] = None,
|
|
31
|
+
num_stacks: int = 1,
|
|
32
|
+
num_layers: int = 1,
|
|
33
|
+
shared_weights: bool = False,
|
|
34
|
+
act: Union[str, Callable, None] = "relu",
|
|
35
|
+
dropout: float = 0.0,
|
|
36
|
+
bias: bool = True,
|
|
37
|
+
# Spektral compatibility arguments:
|
|
38
|
+
channels: Optional[int] = None,
|
|
39
|
+
order: Optional[int] = None,
|
|
40
|
+
iterations: Optional[int] = None,
|
|
41
|
+
share_weights: Optional[bool] = None,
|
|
42
|
+
gcn_activation: Optional[str] = None,
|
|
43
|
+
dropout_rate: Optional[float] = None,
|
|
44
|
+
activation: Optional[str] = None,
|
|
45
|
+
use_bias: Optional[bool] = None,
|
|
46
|
+
**kwargs,
|
|
47
|
+
):
|
|
48
|
+
kwargs.setdefault("aggr", "add")
|
|
49
|
+
super().__init__(node_dim=0, **kwargs)
|
|
50
|
+
|
|
51
|
+
if out_channels is None:
|
|
52
|
+
if channels is not None:
|
|
53
|
+
out_channels = channels
|
|
54
|
+
in_channels = -1
|
|
55
|
+
else:
|
|
56
|
+
out_channels = in_channels
|
|
57
|
+
in_channels = -1
|
|
58
|
+
|
|
59
|
+
if order is not None:
|
|
60
|
+
num_stacks = order
|
|
61
|
+
if iterations is not None:
|
|
62
|
+
num_layers = iterations
|
|
63
|
+
if share_weights is not None:
|
|
64
|
+
shared_weights = share_weights
|
|
65
|
+
if gcn_activation is not None:
|
|
66
|
+
act = gcn_activation
|
|
67
|
+
if dropout_rate is not None:
|
|
68
|
+
dropout = dropout_rate
|
|
69
|
+
if use_bias is not None:
|
|
70
|
+
bias = use_bias
|
|
71
|
+
|
|
72
|
+
self.in_channels = in_channels
|
|
73
|
+
self.out_channels = out_channels
|
|
74
|
+
self.num_stacks = num_stacks
|
|
75
|
+
self.num_layers = num_layers
|
|
76
|
+
self.shared_weights = shared_weights
|
|
77
|
+
self.act = activations.get(act) if act is not None else None
|
|
78
|
+
self.dropout_rate = dropout
|
|
79
|
+
self.use_bias = bias
|
|
80
|
+
|
|
81
|
+
K, T, F_in, F_out = num_stacks, num_layers, in_channels, out_channels
|
|
82
|
+
T_w = 1 if self.shared_weights else T
|
|
83
|
+
|
|
84
|
+
self.weight = self.add_weight(
|
|
85
|
+
shape=(max(1, T_w - 1), K, F_out, F_out),
|
|
86
|
+
initializer="glorot_uniform",
|
|
87
|
+
name="weight",
|
|
88
|
+
)
|
|
89
|
+
if in_channels is not None and in_channels != -1:
|
|
90
|
+
self.init_weight = self.add_weight(
|
|
91
|
+
shape=(K, F_in, F_out),
|
|
92
|
+
initializer="glorot_uniform",
|
|
93
|
+
name="init_weight",
|
|
94
|
+
)
|
|
95
|
+
self.root_weight = self.add_weight(
|
|
96
|
+
shape=(T_w, K, F_in, F_out),
|
|
97
|
+
initializer="glorot_uniform",
|
|
98
|
+
name="root_weight",
|
|
99
|
+
)
|
|
100
|
+
else:
|
|
101
|
+
self.init_weight = None
|
|
102
|
+
self.root_weight = None
|
|
103
|
+
|
|
104
|
+
if bias:
|
|
105
|
+
self.bias = self.add_weight(
|
|
106
|
+
shape=(T_w, K, 1, F_out),
|
|
107
|
+
initializer="zeros",
|
|
108
|
+
name="bias",
|
|
109
|
+
)
|
|
110
|
+
else:
|
|
111
|
+
self.bias = None
|
|
112
|
+
|
|
113
|
+
self.dropout = Dropout(dropout)
|
|
114
|
+
|
|
115
|
+
def build(self, input_shape=None):
|
|
116
|
+
if input_shape is not None:
|
|
117
|
+
if isinstance(input_shape, (list, tuple)) and len(input_shape) > 0 and isinstance(input_shape[0], (list, tuple)):
|
|
118
|
+
dim = input_shape[0][-1]
|
|
119
|
+
elif isinstance(input_shape, (list, tuple)) and len(input_shape) > 0 and input_shape[0] is not None and not isinstance(input_shape[0], (int, type(None))):
|
|
120
|
+
dim = getattr(input_shape[0], "shape", [None, None])[-1]
|
|
121
|
+
else:
|
|
122
|
+
dim = input_shape[-1]
|
|
123
|
+
if (self.in_channels is None or self.in_channels == -1) and dim is not None:
|
|
124
|
+
self.in_channels = dim
|
|
125
|
+
if self.in_channels is not None and self.in_channels != -1:
|
|
126
|
+
K, T, F_in, F_out = self.num_stacks, self.num_layers, self.in_channels, self.out_channels
|
|
127
|
+
T_w = 1 if self.shared_weights else T
|
|
128
|
+
if self.init_weight is None:
|
|
129
|
+
self.init_weight = self.add_weight(
|
|
130
|
+
shape=(K, F_in, F_out),
|
|
131
|
+
initializer="glorot_uniform",
|
|
132
|
+
name="init_weight",
|
|
133
|
+
)
|
|
134
|
+
if self.root_weight is None:
|
|
135
|
+
self.root_weight = self.add_weight(
|
|
136
|
+
shape=(T_w, K, F_in, F_out),
|
|
137
|
+
initializer="glorot_uniform",
|
|
138
|
+
name="root_weight",
|
|
139
|
+
)
|
|
140
|
+
self.built = True
|
|
141
|
+
|
|
142
|
+
def call(self, inputs, edge_index=None, edge_weight=None, training=None, **kwargs):
|
|
143
|
+
if edge_index is None:
|
|
144
|
+
if isinstance(inputs, (list, tuple)):
|
|
145
|
+
if len(inputs) == 3:
|
|
146
|
+
x, edge_index, edge_weight = inputs
|
|
147
|
+
elif len(inputs) == 2:
|
|
148
|
+
x, edge_index = inputs
|
|
149
|
+
else:
|
|
150
|
+
raise ValueError(f"Unexpected input length {len(inputs)}")
|
|
151
|
+
else:
|
|
152
|
+
raise ValueError("Expected (x, edge_index) or x and edge_index")
|
|
153
|
+
else:
|
|
154
|
+
x = inputs
|
|
155
|
+
|
|
156
|
+
if self.in_channels is None or self.in_channels == -1:
|
|
157
|
+
self.build((None, ops.shape(x)[-1]))
|
|
158
|
+
|
|
159
|
+
# Legacy adj check
|
|
160
|
+
is_legacy = False
|
|
161
|
+
if hasattr(edge_index, "shape") and len(edge_index.shape) == 2:
|
|
162
|
+
if edge_index.shape[0] != 2 and edge_index.shape[0] == edge_index.shape[1]:
|
|
163
|
+
is_legacy = True
|
|
164
|
+
elif not hasattr(edge_index, "shape"):
|
|
165
|
+
is_legacy = True
|
|
166
|
+
|
|
167
|
+
if is_legacy:
|
|
168
|
+
if hasattr(edge_index, "indices") and not callable(edge_index.indices):
|
|
169
|
+
edge_weight = edge_index.values
|
|
170
|
+
edge_index = ops.transpose(edge_index.indices)
|
|
171
|
+
else:
|
|
172
|
+
adj = edge_index
|
|
173
|
+
row, col = ops.where(adj > 0)
|
|
174
|
+
edge_index = ops.stack([row, col], axis=0)
|
|
175
|
+
edge_weight = ops.take(adj, row * ops.shape(adj)[1] + col)
|
|
176
|
+
|
|
177
|
+
num_nodes = ops.shape(x)[0]
|
|
178
|
+
edge_index, edge_weight = gcn_norm(
|
|
179
|
+
edge_index,
|
|
180
|
+
edge_weight,
|
|
181
|
+
num_nodes=num_nodes,
|
|
182
|
+
add_self_loops=False,
|
|
183
|
+
dtype=x.dtype,
|
|
184
|
+
)
|
|
185
|
+
|
|
186
|
+
# PyG: x = x.unsqueeze(-3) -> (1, N, F_in)
|
|
187
|
+
# out = x
|
|
188
|
+
# for t in range(num_layers):
|
|
189
|
+
# if t == 0: out = out @ init_weight (K, F_in, F_out) -> (K, N, F_out)
|
|
190
|
+
# else: out = out @ weight[t-1] (K, F_out, F_out) -> (K, N, F_out)
|
|
191
|
+
# out = propagate(edge_index, x=out, edge_weight=edge_weight)
|
|
192
|
+
# root = dropout(x) @ root_weight[t] (K, F_in, F_out) -> (K, N, F_out)
|
|
193
|
+
# out = out + root
|
|
194
|
+
# if bias: out = out + bias[t]
|
|
195
|
+
# if act: out = act(out)
|
|
196
|
+
# return out.mean(dim=-3)
|
|
197
|
+
K, T, F_in, F_out = self.num_stacks, self.num_layers, self.in_channels, self.out_channels
|
|
198
|
+
|
|
199
|
+
# out shape: (N, K, F_out)
|
|
200
|
+
out = None
|
|
201
|
+
for t in range(self.num_layers):
|
|
202
|
+
w_idx = 0 if self.shared_weights else t
|
|
203
|
+
if t == 0:
|
|
204
|
+
# x: (N, F_in), init_weight: (K, F_in, F_out)
|
|
205
|
+
# out: (K, N, F_out)
|
|
206
|
+
out = ops.einsum("nf,kfo->kno", x, self.init_weight)
|
|
207
|
+
else:
|
|
208
|
+
w = self.weight[0 if self.shared_weights else t - 1]
|
|
209
|
+
out = ops.einsum("kno,kof->knf", out, w)
|
|
210
|
+
|
|
211
|
+
# Transpose to (N, K, F_out) so node_dim=0
|
|
212
|
+
out_n = ops.transpose(out, (1, 0, 2))
|
|
213
|
+
out_prop = self.propagate(edge_index, x=out_n, edge_weight=edge_weight, size=(num_nodes, num_nodes))
|
|
214
|
+
out = ops.transpose(out_prop, (1, 0, 2)) # (K, N, F_out)
|
|
215
|
+
|
|
216
|
+
root_x = self.dropout(x, training=training)
|
|
217
|
+
root = ops.einsum("nf,kfo->kno", root_x, self.root_weight[w_idx])
|
|
218
|
+
out = out + root
|
|
219
|
+
|
|
220
|
+
if self.bias is not None:
|
|
221
|
+
# bias[w_idx]: (K, 1, F_out)
|
|
222
|
+
out = out + self.bias[w_idx]
|
|
223
|
+
|
|
224
|
+
if self.act is not None:
|
|
225
|
+
out = self.act(out)
|
|
226
|
+
|
|
227
|
+
# out: (K, N, F_out) -> mean over K -> (N, F_out)
|
|
228
|
+
return ops.mean(out, axis=0)
|
|
229
|
+
|
|
230
|
+
def message(self, x_j, edge_weight=None):
|
|
231
|
+
return x_j if edge_weight is None else ops.expand_dims(ops.expand_dims(edge_weight, -1), -1) * x_j
|
|
@@ -0,0 +1,92 @@
|
|
|
1
|
+
from typing import Union, Tuple
|
|
2
|
+
from keras import layers, ops
|
|
3
|
+
|
|
4
|
+
from k3_node.layers.conv.message_passing import MessagePassing
|
|
5
|
+
|
|
6
|
+
|
|
7
|
+
class CGConv(MessagePassing):
|
|
8
|
+
r"""The Crystal Graph Convolutional operator from the
|
|
9
|
+
`"Crystal Graph Convolutional Neural Networks for an Accurate and
|
|
10
|
+
Interpretable Prediction of Material Properties"
|
|
11
|
+
<https://arxiv.org/abs/1710.10324>`_ paper.
|
|
12
|
+
|
|
13
|
+
Args:
|
|
14
|
+
channels: Size of each input sample, or a tuple for bipartite graphs.
|
|
15
|
+
dim: Edge feature dimensionality. (default: ``0``)
|
|
16
|
+
aggr: The aggregation scheme to use (``"add"``, ``"mean"``, ``"max"``).
|
|
17
|
+
(default: ``"add"``)
|
|
18
|
+
batch_norm: If set to :obj:`True`, will apply batch normalization.
|
|
19
|
+
(default: ``False``)
|
|
20
|
+
bias: If set to :obj:`False`, the layer will not learn an additive bias.
|
|
21
|
+
(default: ``True``)
|
|
22
|
+
|
|
23
|
+
Example:
|
|
24
|
+
```python
|
|
25
|
+
import numpy as np
|
|
26
|
+
from k3_node.layers import CGConv
|
|
27
|
+
|
|
28
|
+
x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
|
|
29
|
+
edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
|
|
30
|
+
edge_attr = np.random.rand(30, 3).astype("float32") # 3 features per edge
|
|
31
|
+
|
|
32
|
+
layer = CGConv(channels=8, dim=3)
|
|
33
|
+
out = layer(x, edge_index, edge_attr)
|
|
34
|
+
print(tuple(out.shape)) # (10, 8)
|
|
35
|
+
```
|
|
36
|
+
"""
|
|
37
|
+
|
|
38
|
+
def __init__(
|
|
39
|
+
self,
|
|
40
|
+
channels: Union[int, Tuple[int, int]],
|
|
41
|
+
dim: int = 0,
|
|
42
|
+
aggr: str = "add",
|
|
43
|
+
batch_norm: bool = False,
|
|
44
|
+
bias: bool = True,
|
|
45
|
+
**kwargs,
|
|
46
|
+
):
|
|
47
|
+
super().__init__(aggr=aggr, **kwargs)
|
|
48
|
+
self.channels = channels
|
|
49
|
+
self.dim = dim
|
|
50
|
+
self.batch_norm = batch_norm
|
|
51
|
+
self.use_bias = bias
|
|
52
|
+
|
|
53
|
+
if isinstance(channels, int):
|
|
54
|
+
self.in_channels_src = channels
|
|
55
|
+
self.in_channels_dst = channels
|
|
56
|
+
else:
|
|
57
|
+
self.in_channels_src = channels[0]
|
|
58
|
+
self.in_channels_dst = channels[1]
|
|
59
|
+
|
|
60
|
+
in_dim = self.in_channels_src + self.in_channels_dst + dim
|
|
61
|
+
self.lin_f = layers.Dense(self.in_channels_dst, use_bias=bias)
|
|
62
|
+
self.lin_s = layers.Dense(self.in_channels_dst, use_bias=bias)
|
|
63
|
+
self.bn = layers.BatchNormalization(momentum=0.9, epsilon=1e-5) if batch_norm else None
|
|
64
|
+
|
|
65
|
+
def build(self, input_shape):
|
|
66
|
+
in_dim = self.in_channels_src + self.in_channels_dst + self.dim
|
|
67
|
+
self.lin_f.build((None, in_dim))
|
|
68
|
+
self.lin_s.build((None, in_dim))
|
|
69
|
+
self.built = True
|
|
70
|
+
|
|
71
|
+
def call(self, x, edge_index=None, edge_attr=None, **kwargs):
|
|
72
|
+
if edge_index is None and isinstance(x, (tuple, list)):
|
|
73
|
+
x, edge_index = x[0], x[1]
|
|
74
|
+
|
|
75
|
+
if not isinstance(x, (tuple, list)):
|
|
76
|
+
x_src, x_dst = x, x
|
|
77
|
+
else:
|
|
78
|
+
x_src, x_dst = x[0], x[1]
|
|
79
|
+
|
|
80
|
+
out = self.propagate(edge_index, x=(x_src, x_dst), edge_attr=edge_attr)
|
|
81
|
+
if self.bn is not None:
|
|
82
|
+
out = self.bn(out)
|
|
83
|
+
out = x_dst + out
|
|
84
|
+
return out
|
|
85
|
+
|
|
86
|
+
def message(self, x_i, x_j, edge_attr=None):
|
|
87
|
+
if edge_attr is None:
|
|
88
|
+
z = ops.concatenate([x_i, x_j], axis=-1)
|
|
89
|
+
else:
|
|
90
|
+
z = ops.concatenate([x_i, x_j, edge_attr], axis=-1)
|
|
91
|
+
return ops.sigmoid(self.lin_f(z)) * ops.softplus(self.lin_s(z))
|
|
92
|
+
|
|
@@ -0,0 +1,137 @@
|
|
|
1
|
+
from typing import Optional, List
|
|
2
|
+
from keras import layers, ops
|
|
3
|
+
|
|
4
|
+
from k3_node.layers.conv.message_passing import MessagePassing
|
|
5
|
+
from k3_node.layers.conv.utils import get_laplacian
|
|
6
|
+
|
|
7
|
+
|
|
8
|
+
class ChebConv(MessagePassing):
|
|
9
|
+
r"""The Chebyshev spectral graph convolutional operator from the
|
|
10
|
+
`"Convolutional Neural Networks on Graphs with Fast Localized Spectral
|
|
11
|
+
Filtering" <https://arxiv.org/abs/1606.09375>`_ paper.
|
|
12
|
+
|
|
13
|
+
Args:
|
|
14
|
+
in_channels: Size of each input sample.
|
|
15
|
+
out_channels: Size of each output sample.
|
|
16
|
+
K: Chebyshev filter size :math:`K`.
|
|
17
|
+
normalization: The normalization scheme for the graph
|
|
18
|
+
Laplacian (``"sym"``, ``"rw"`` or :obj:`None`). (default: ``"sym"``)
|
|
19
|
+
bias: If set to :obj:`False`, the layer will not learn
|
|
20
|
+
an additive bias. (default: ``"True"``)
|
|
21
|
+
|
|
22
|
+
Example:
|
|
23
|
+
```python
|
|
24
|
+
import numpy as np
|
|
25
|
+
from k3_node.layers import ChebConv
|
|
26
|
+
|
|
27
|
+
x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
|
|
28
|
+
edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
|
|
29
|
+
|
|
30
|
+
layer = ChebConv(in_channels=8, out_channels=16, K=2)
|
|
31
|
+
out = layer(x, edge_index)
|
|
32
|
+
print(tuple(out.shape)) # (10, 16)
|
|
33
|
+
```
|
|
34
|
+
"""
|
|
35
|
+
|
|
36
|
+
def __init__(
|
|
37
|
+
self,
|
|
38
|
+
in_channels: int,
|
|
39
|
+
out_channels: int,
|
|
40
|
+
K: int,
|
|
41
|
+
normalization: Optional[str] = "sym",
|
|
42
|
+
bias: bool = True,
|
|
43
|
+
**kwargs,
|
|
44
|
+
):
|
|
45
|
+
super().__init__(aggr="add", **kwargs)
|
|
46
|
+
if K <= 0:
|
|
47
|
+
raise ValueError(f"K must be a positive integer, got {K}")
|
|
48
|
+
|
|
49
|
+
self.in_channels = in_channels
|
|
50
|
+
self.out_channels = out_channels
|
|
51
|
+
self.K = K
|
|
52
|
+
self.normalization = normalization
|
|
53
|
+
self.use_bias = bias
|
|
54
|
+
|
|
55
|
+
self.lins = [layers.Dense(out_channels, use_bias=False) for _ in range(K)]
|
|
56
|
+
|
|
57
|
+
def build(self, input_shape):
|
|
58
|
+
if isinstance(input_shape, (tuple, list)) and len(input_shape) > 0 and isinstance(input_shape[0], (tuple, list)):
|
|
59
|
+
feat_shape = input_shape[0]
|
|
60
|
+
else:
|
|
61
|
+
feat_shape = input_shape
|
|
62
|
+
for lin in self.lins:
|
|
63
|
+
lin.build(feat_shape)
|
|
64
|
+
|
|
65
|
+
if self.use_bias:
|
|
66
|
+
self.bias = self.add_weight(
|
|
67
|
+
shape=(self.out_channels,),
|
|
68
|
+
initializer="zeros",
|
|
69
|
+
name="bias",
|
|
70
|
+
)
|
|
71
|
+
else:
|
|
72
|
+
self.bias = None
|
|
73
|
+
self.built = True
|
|
74
|
+
|
|
75
|
+
def __norm__(
|
|
76
|
+
self,
|
|
77
|
+
edge_index,
|
|
78
|
+
num_nodes: Optional[int],
|
|
79
|
+
edge_weight=None,
|
|
80
|
+
normalization: Optional[str] = "sym",
|
|
81
|
+
lambda_max=None,
|
|
82
|
+
dtype=None,
|
|
83
|
+
):
|
|
84
|
+
edge_index, edge_weight = get_laplacian(
|
|
85
|
+
edge_index, edge_weight, normalization, dtype, num_nodes
|
|
86
|
+
)
|
|
87
|
+
if lambda_max is None:
|
|
88
|
+
lambda_max = 2.0 * ops.max(edge_weight)
|
|
89
|
+
else:
|
|
90
|
+
lambda_max = ops.convert_to_tensor(lambda_max, dtype=edge_weight.dtype)
|
|
91
|
+
|
|
92
|
+
edge_weight = (2.0 * edge_weight) / lambda_max
|
|
93
|
+
edge_weight = ops.where(
|
|
94
|
+
ops.isinf(edge_weight) | ops.isnan(edge_weight), 0.0, edge_weight
|
|
95
|
+
)
|
|
96
|
+
|
|
97
|
+
loop_mask = edge_index[0] == edge_index[1]
|
|
98
|
+
edge_weight = ops.where(loop_mask, edge_weight - 1.0, edge_weight)
|
|
99
|
+
return edge_index, edge_weight
|
|
100
|
+
|
|
101
|
+
def call(self, x, edge_index=None, edge_weight=None, lambda_max=None, **kwargs):
|
|
102
|
+
if edge_index is None and isinstance(x, (tuple, list)):
|
|
103
|
+
x, edge_index = x[0], x[1]
|
|
104
|
+
|
|
105
|
+
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]
|
|
106
|
+
edge_index, norm = self.__norm__(
|
|
107
|
+
edge_index,
|
|
108
|
+
num_nodes,
|
|
109
|
+
edge_weight,
|
|
110
|
+
self.normalization,
|
|
111
|
+
lambda_max,
|
|
112
|
+
dtype=x.dtype,
|
|
113
|
+
)
|
|
114
|
+
|
|
115
|
+
Tx_0 = x
|
|
116
|
+
Tx_1 = x
|
|
117
|
+
out = self.lins[0](Tx_0)
|
|
118
|
+
|
|
119
|
+
if len(self.lins) > 1:
|
|
120
|
+
Tx_1 = self.propagate(edge_index, x=x, norm=norm)
|
|
121
|
+
out = out + self.lins[1](Tx_1)
|
|
122
|
+
|
|
123
|
+
for lin in self.lins[2:]:
|
|
124
|
+
Tx_2 = self.propagate(edge_index, x=Tx_1, norm=norm)
|
|
125
|
+
Tx_2 = 2.0 * Tx_2 - Tx_0
|
|
126
|
+
out = out + lin(Tx_2)
|
|
127
|
+
Tx_0, Tx_1 = Tx_1, Tx_2
|
|
128
|
+
|
|
129
|
+
if self.bias is not None:
|
|
130
|
+
out = out + self.bias
|
|
131
|
+
return out
|
|
132
|
+
|
|
133
|
+
def message(self, x_j, norm=None):
|
|
134
|
+
if norm is None:
|
|
135
|
+
return x_j
|
|
136
|
+
return ops.expand_dims(norm, -1) * x_j
|
|
137
|
+
|
|
@@ -0,0 +1,102 @@
|
|
|
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 (
|
|
5
|
+
add_self_loops,
|
|
6
|
+
degree,
|
|
7
|
+
extend_mask_for_self_loops,
|
|
8
|
+
remove_self_loops_masked,
|
|
9
|
+
)
|
|
10
|
+
from k3_node.ops.segment import segment_sum
|
|
11
|
+
|
|
12
|
+
|
|
13
|
+
class ClusterGCNConv(MessagePassing):
|
|
14
|
+
r"""The ClusterGCN graph convolutional operator from the
|
|
15
|
+
`"Cluster-GCN: An Efficient Algorithm for Training Deep and Large Graph
|
|
16
|
+
Convolutional Networks" <https://arxiv.org/abs/1905.07953>`_ paper.
|
|
17
|
+
|
|
18
|
+
Args:
|
|
19
|
+
in_channels: Size of each input sample.
|
|
20
|
+
out_channels: Size of each output sample.
|
|
21
|
+
diag_lambda: Diagonal enhancement coefficient :math:`\lambda`.
|
|
22
|
+
(default: ``0.0``)
|
|
23
|
+
add_self_loops: If set to :obj:`False`, will not add self-loops.
|
|
24
|
+
(default: ``True``)
|
|
25
|
+
bias: If set to :obj:`False`, the layer will not learn an additive bias.
|
|
26
|
+
(default: ``True``)
|
|
27
|
+
|
|
28
|
+
Example:
|
|
29
|
+
```python
|
|
30
|
+
import numpy as np
|
|
31
|
+
from k3_node.layers import ClusterGCNConv
|
|
32
|
+
|
|
33
|
+
x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
|
|
34
|
+
edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
|
|
35
|
+
|
|
36
|
+
layer = ClusterGCNConv(in_channels=8, out_channels=16)
|
|
37
|
+
out = layer(x, edge_index)
|
|
38
|
+
print(tuple(out.shape)) # (10, 16)
|
|
39
|
+
```
|
|
40
|
+
"""
|
|
41
|
+
|
|
42
|
+
weighted_sum_message = True
|
|
43
|
+
|
|
44
|
+
def __init__(
|
|
45
|
+
self,
|
|
46
|
+
in_channels: int,
|
|
47
|
+
out_channels: int,
|
|
48
|
+
diag_lambda: float = 0.0,
|
|
49
|
+
add_self_loops: bool = True,
|
|
50
|
+
bias: bool = True,
|
|
51
|
+
**kwargs,
|
|
52
|
+
):
|
|
53
|
+
super().__init__(aggr="add", **kwargs)
|
|
54
|
+
self.in_channels = in_channels
|
|
55
|
+
self.out_channels = out_channels
|
|
56
|
+
self.diag_lambda = diag_lambda
|
|
57
|
+
self.add_self_loops = add_self_loops
|
|
58
|
+
self.use_bias = bias
|
|
59
|
+
|
|
60
|
+
self.lin_out = layers.Dense(out_channels, use_bias=bias)
|
|
61
|
+
self.lin_root = layers.Dense(out_channels, use_bias=False)
|
|
62
|
+
|
|
63
|
+
def build(self, input_shape):
|
|
64
|
+
feat_shape = input_shape[0] if isinstance(input_shape, (tuple, list)) and isinstance(input_shape[0], (tuple, list)) else input_shape
|
|
65
|
+
self.lin_out.build(feat_shape)
|
|
66
|
+
self.lin_root.build(feat_shape)
|
|
67
|
+
self.built = True
|
|
68
|
+
|
|
69
|
+
def call(self, x, edge_index=None, edge_weight=None, **kwargs):
|
|
70
|
+
if edge_index is None and isinstance(x, (tuple, list)):
|
|
71
|
+
x, edge_index = x[0], x[1]
|
|
72
|
+
|
|
73
|
+
num_nodes = ops.shape(x)[self.node_dim]
|
|
74
|
+
|
|
75
|
+
keep_mask = None
|
|
76
|
+
if self.add_self_loops:
|
|
77
|
+
edge_index, _, keep_mask = remove_self_loops_masked(edge_index)
|
|
78
|
+
edge_index, _ = add_self_loops(edge_index, num_nodes=num_nodes)
|
|
79
|
+
keep_mask = extend_mask_for_self_loops(keep_mask, num_nodes)
|
|
80
|
+
|
|
81
|
+
row, col = edge_index[0], edge_index[1]
|
|
82
|
+
col_cast = ops.cast(col, "int32")
|
|
83
|
+
if keep_mask is None:
|
|
84
|
+
deg = degree(col_cast, num_nodes=num_nodes)
|
|
85
|
+
else:
|
|
86
|
+
deg = segment_sum(ops.cast(keep_mask, x.dtype), col_cast, num_segments=num_nodes)
|
|
87
|
+
deg_inv = 1.0 / ops.maximum(ops.cast(deg, x.dtype), 1.0)
|
|
88
|
+
|
|
89
|
+
edge_weight = ops.take(deg_inv, col_cast, axis=0)
|
|
90
|
+
loop_mask = ops.equal(row, col)
|
|
91
|
+
edge_weight = ops.where(loop_mask, edge_weight + self.diag_lambda * ops.take(deg_inv, col_cast, axis=0), edge_weight)
|
|
92
|
+
if keep_mask is not None:
|
|
93
|
+
edge_weight = edge_weight * ops.cast(keep_mask, edge_weight.dtype)
|
|
94
|
+
|
|
95
|
+
out = self.propagate(edge_index, x=x, edge_weight=edge_weight)
|
|
96
|
+
return self.lin_out(out) + self.lin_root(x)
|
|
97
|
+
|
|
98
|
+
def message(self, x_j, edge_weight=None):
|
|
99
|
+
if edge_weight is None:
|
|
100
|
+
return x_j
|
|
101
|
+
return ops.expand_dims(edge_weight, -1) * x_j
|
|
102
|
+
|
|
@@ -0,0 +1,100 @@
|
|
|
1
|
+
# ported from spektral
|
|
2
|
+
|
|
3
|
+
import warnings
|
|
4
|
+
from functools import wraps
|
|
5
|
+
|
|
6
|
+
from keras import ops, backend
|
|
7
|
+
from keras.layers import Layer
|
|
8
|
+
|
|
9
|
+
from k3_node.utils import (
|
|
10
|
+
is_keras_kwarg,
|
|
11
|
+
is_layer_kwarg,
|
|
12
|
+
deserialize_kwarg,
|
|
13
|
+
serialize_kwarg,
|
|
14
|
+
)
|
|
15
|
+
|
|
16
|
+
|
|
17
|
+
class Conv(Layer):
|
|
18
|
+
def __init__(self, **kwargs):
|
|
19
|
+
unknown = sorted(k for k in kwargs if not (is_keras_kwarg(k) or is_layer_kwarg(k)))
|
|
20
|
+
if unknown:
|
|
21
|
+
raise TypeError(f"{type(self).__name__}() got unexpected keyword argument(s): {', '.join(unknown)}")
|
|
22
|
+
super().__init__(**{k: v for k, v in kwargs.items() if is_keras_kwarg(k)})
|
|
23
|
+
self.supports_masking = True
|
|
24
|
+
self.kwargs_keys = []
|
|
25
|
+
for key in kwargs:
|
|
26
|
+
if is_layer_kwarg(key):
|
|
27
|
+
attr = kwargs[key]
|
|
28
|
+
attr = deserialize_kwarg(key, attr)
|
|
29
|
+
self.kwargs_keys.append(key)
|
|
30
|
+
setattr(self, key, attr)
|
|
31
|
+
self.call = check_dtypes_decorator(self.call)
|
|
32
|
+
|
|
33
|
+
def build(self, input_shape):
|
|
34
|
+
self.built = True
|
|
35
|
+
|
|
36
|
+
def call(self, inputs):
|
|
37
|
+
raise NotImplementedError
|
|
38
|
+
|
|
39
|
+
def get_config(self):
|
|
40
|
+
base_config = super().get_config()
|
|
41
|
+
keras_config = {}
|
|
42
|
+
for key in self.kwargs_keys:
|
|
43
|
+
keras_config[key] = serialize_kwarg(key, getattr(self, key))
|
|
44
|
+
return {**base_config, **keras_config, **self.config}
|
|
45
|
+
|
|
46
|
+
@property
|
|
47
|
+
def config(self):
|
|
48
|
+
return {}
|
|
49
|
+
|
|
50
|
+
@staticmethod
|
|
51
|
+
def preprocess(a):
|
|
52
|
+
return a
|
|
53
|
+
|
|
54
|
+
|
|
55
|
+
def check_dtypes_decorator(call):
|
|
56
|
+
@wraps(call)
|
|
57
|
+
def _inner_check_dtypes(*args, **kwargs):
|
|
58
|
+
if len(args) == 0:
|
|
59
|
+
return call(**kwargs)
|
|
60
|
+
elif len(args) == 1:
|
|
61
|
+
inputs = check_dtypes(args[0])
|
|
62
|
+
return call(inputs, **kwargs)
|
|
63
|
+
else:
|
|
64
|
+
checked = check_dtypes(list(args))
|
|
65
|
+
if isinstance(checked, (list, tuple)) and len(checked) == len(args):
|
|
66
|
+
return call(*checked, **kwargs)
|
|
67
|
+
return call(*args, **kwargs)
|
|
68
|
+
|
|
69
|
+
return _inner_check_dtypes
|
|
70
|
+
|
|
71
|
+
|
|
72
|
+
def check_dtypes(inputs):
|
|
73
|
+
if not isinstance(inputs, (list, tuple)):
|
|
74
|
+
return inputs
|
|
75
|
+
for value in inputs:
|
|
76
|
+
if not hasattr(value, "dtype"):
|
|
77
|
+
# It's not a valid tensor.
|
|
78
|
+
return inputs
|
|
79
|
+
|
|
80
|
+
if len(inputs) == 2:
|
|
81
|
+
x, a = inputs
|
|
82
|
+
e = None
|
|
83
|
+
elif len(inputs) == 3:
|
|
84
|
+
x, a, e = inputs
|
|
85
|
+
else:
|
|
86
|
+
return inputs
|
|
87
|
+
|
|
88
|
+
# If 'a' is an edge_index of shape (2, E), it must remain integer
|
|
89
|
+
if hasattr(a, "shape") and len(a.shape) == 2 and a.shape[0] == 2 and a.shape[1] != 2:
|
|
90
|
+
pass
|
|
91
|
+
elif backend.is_int_dtype(a.dtype) and backend.is_float_dtype(x.dtype):
|
|
92
|
+
warnings.warn(
|
|
93
|
+
f"The adjacency matrix of dtype {a.dtype} is incompatible with the dtype "
|
|
94
|
+
f"of the node features {x.dtype} and has been automatically cast to "
|
|
95
|
+
f"{x.dtype}."
|
|
96
|
+
)
|
|
97
|
+
a = ops.cast(a, x.dtype)
|
|
98
|
+
|
|
99
|
+
output = [_ for _ in [x, a, e] if _ is not None]
|
|
100
|
+
return output
|