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,110 @@
|
|
|
1
|
+
from typing import Callable, Union, Tuple
|
|
2
|
+
from keras import layers, ops
|
|
3
|
+
|
|
4
|
+
from k3_node.layers.conv.message_passing import MessagePassing
|
|
5
|
+
|
|
6
|
+
|
|
7
|
+
class NNConv(MessagePassing):
|
|
8
|
+
r"""The continuous kernel-based convolutional operator from the
|
|
9
|
+
`"Neural Message Passing for Quantum Chemistry"
|
|
10
|
+
<https://arxiv.org/abs/1704.01212>`_ paper.
|
|
11
|
+
|
|
12
|
+
Args:
|
|
13
|
+
in_channels: Size of each input sample, or a tuple for bipartite graphs.
|
|
14
|
+
out_channels: Size of each output sample.
|
|
15
|
+
nn: A neural network :math:`h_{\mathbf{\Theta}}` that maps edge features
|
|
16
|
+
to shape :obj:`[-1, in_channels * out_channels]`.
|
|
17
|
+
aggr: The aggregation scheme to use (``"add"``, ``"mean"``, ``"max"``).
|
|
18
|
+
(default: ``"add"``)
|
|
19
|
+
root_weight: If set to :obj:`False`, the layer will not add the
|
|
20
|
+
transformed root node features. (default: ``True``)
|
|
21
|
+
bias: If set to :obj:`False`, the layer will not learn an additive bias.
|
|
22
|
+
(default: ``True``)
|
|
23
|
+
|
|
24
|
+
Example:
|
|
25
|
+
```python
|
|
26
|
+
import numpy as np
|
|
27
|
+
import keras
|
|
28
|
+
from k3_node.layers import NNConv
|
|
29
|
+
|
|
30
|
+
x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
|
|
31
|
+
edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
|
|
32
|
+
edge_attr = np.random.rand(30, 3).astype("float32") # 3 features per edge
|
|
33
|
+
|
|
34
|
+
# `nn` maps each edge's features to an [in_channels * out_channels] weight matrix
|
|
35
|
+
layer = NNConv(in_channels=8, out_channels=16, nn=keras.layers.Dense(8 * 16))
|
|
36
|
+
out = layer(x, edge_index, edge_attr)
|
|
37
|
+
print(tuple(out.shape)) # (10, 16)
|
|
38
|
+
```
|
|
39
|
+
"""
|
|
40
|
+
|
|
41
|
+
def __init__(
|
|
42
|
+
self,
|
|
43
|
+
in_channels: Union[int, Tuple[int, int]],
|
|
44
|
+
out_channels: int,
|
|
45
|
+
nn: Callable,
|
|
46
|
+
aggr: str = "add",
|
|
47
|
+
root_weight: bool = True,
|
|
48
|
+
bias: bool = True,
|
|
49
|
+
**kwargs,
|
|
50
|
+
):
|
|
51
|
+
super().__init__(aggr=aggr, **kwargs)
|
|
52
|
+
self.in_channels = in_channels
|
|
53
|
+
self.out_channels = out_channels
|
|
54
|
+
self.nn = nn
|
|
55
|
+
self.root_weight = root_weight
|
|
56
|
+
self.use_bias = bias
|
|
57
|
+
|
|
58
|
+
if isinstance(in_channels, int):
|
|
59
|
+
self.in_channels_src = in_channels
|
|
60
|
+
self.in_channels_dst = in_channels
|
|
61
|
+
else:
|
|
62
|
+
self.in_channels_src = in_channels[0]
|
|
63
|
+
self.in_channels_dst = in_channels[1]
|
|
64
|
+
|
|
65
|
+
if root_weight:
|
|
66
|
+
self.lin_root = layers.Dense(out_channels, use_bias=False)
|
|
67
|
+
else:
|
|
68
|
+
self.lin_root = None
|
|
69
|
+
|
|
70
|
+
def build(self, input_shape):
|
|
71
|
+
if self.lin_root is not None:
|
|
72
|
+
self.lin_root.build((None, self.in_channels_dst))
|
|
73
|
+
|
|
74
|
+
if self.use_bias:
|
|
75
|
+
self.bias = self.add_weight(
|
|
76
|
+
shape=(self.out_channels,),
|
|
77
|
+
initializer="zeros",
|
|
78
|
+
name="bias",
|
|
79
|
+
)
|
|
80
|
+
else:
|
|
81
|
+
self.bias = None
|
|
82
|
+
self.built = True
|
|
83
|
+
|
|
84
|
+
def call(self, x, edge_index=None, edge_attr=None, size=None, **kwargs):
|
|
85
|
+
if edge_index is None and isinstance(x, (tuple, list)):
|
|
86
|
+
x, edge_index = x[0], x[1]
|
|
87
|
+
|
|
88
|
+
if not isinstance(x, (tuple, list)):
|
|
89
|
+
x_src, x_dst = x, x
|
|
90
|
+
else:
|
|
91
|
+
x_src, x_dst = x[0], x[1]
|
|
92
|
+
|
|
93
|
+
weight = self.nn(edge_attr)
|
|
94
|
+
weight = ops.reshape(weight, (-1, self.in_channels_src, self.out_channels))
|
|
95
|
+
|
|
96
|
+
out = self.propagate(edge_index, x=x_src, weight=weight, size=size)
|
|
97
|
+
|
|
98
|
+
if self.root_weight and self.lin_root is not None and x_dst is not None:
|
|
99
|
+
out = out + self.lin_root(x_dst)
|
|
100
|
+
|
|
101
|
+
if self.bias is not None:
|
|
102
|
+
out = out + self.bias
|
|
103
|
+
return out
|
|
104
|
+
|
|
105
|
+
def message(self, x_j, weight):
|
|
106
|
+
# x_j: [E, in_channels_src], weight: [E, in_channels_src, out_channels]
|
|
107
|
+
x_j = ops.expand_dims(x_j, axis=1) # [E, 1, in_channels_src]
|
|
108
|
+
msg = ops.matmul(x_j, weight) # [E, 1, out_channels]
|
|
109
|
+
return ops.squeeze(msg, axis=1) # [E, out_channels]
|
|
110
|
+
|
|
@@ -0,0 +1,100 @@
|
|
|
1
|
+
from keras import ops
|
|
2
|
+
from keras.layers import Dense
|
|
3
|
+
|
|
4
|
+
from k3_node.layers.conv.message_passing import MessagePassing
|
|
5
|
+
from k3_node.ops.creation import scatter
|
|
6
|
+
|
|
7
|
+
|
|
8
|
+
class PANConv(MessagePassing):
|
|
9
|
+
r"""The path integral based convolution operator from the
|
|
10
|
+
`"Path Integral Based Convolution and Pooling for Graph Neural Networks"
|
|
11
|
+
<https://arxiv.org/abs/2004.14805>`_ paper.
|
|
12
|
+
|
|
13
|
+
Example:
|
|
14
|
+
```python
|
|
15
|
+
import numpy as np
|
|
16
|
+
from k3_node.layers import PANConv
|
|
17
|
+
|
|
18
|
+
x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
|
|
19
|
+
edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
|
|
20
|
+
|
|
21
|
+
layer = PANConv(in_channels=8, out_channels=16, filter_size=2)
|
|
22
|
+
out, weights = layer(x, edge_index) # also returns the learned path weights
|
|
23
|
+
print(tuple(out.shape)) # (10, 16)
|
|
24
|
+
```
|
|
25
|
+
"""
|
|
26
|
+
def __init__(
|
|
27
|
+
self,
|
|
28
|
+
in_channels: int,
|
|
29
|
+
out_channels: int,
|
|
30
|
+
filter_size: int,
|
|
31
|
+
**kwargs,
|
|
32
|
+
):
|
|
33
|
+
kwargs.setdefault("aggr", "add")
|
|
34
|
+
super().__init__(**kwargs)
|
|
35
|
+
|
|
36
|
+
self.in_channels = in_channels
|
|
37
|
+
self.out_channels = out_channels
|
|
38
|
+
self.filter_size = filter_size
|
|
39
|
+
|
|
40
|
+
self.lin = Dense(out_channels, use_bias=True)
|
|
41
|
+
self.weight = self.add_weight(
|
|
42
|
+
shape=(filter_size + 1,),
|
|
43
|
+
initializer="ones",
|
|
44
|
+
name="weight",
|
|
45
|
+
)
|
|
46
|
+
|
|
47
|
+
def build(self, input_shape=None):
|
|
48
|
+
self.lin.build((None, self.in_channels))
|
|
49
|
+
self.built = True
|
|
50
|
+
|
|
51
|
+
def call(self, inputs, edge_index=None, **kwargs):
|
|
52
|
+
if edge_index is None:
|
|
53
|
+
if isinstance(inputs, (list, tuple)) and len(inputs) == 2:
|
|
54
|
+
x, edge_index = inputs
|
|
55
|
+
else:
|
|
56
|
+
raise ValueError("Expected (x, edge_index) or x and edge_index")
|
|
57
|
+
else:
|
|
58
|
+
x = inputs
|
|
59
|
+
|
|
60
|
+
if not self.built:
|
|
61
|
+
self.build()
|
|
62
|
+
|
|
63
|
+
num_nodes = ops.shape(x)[0]
|
|
64
|
+
|
|
65
|
+
# Construct adjacency matrix
|
|
66
|
+
if hasattr(edge_index, "shape") and len(edge_index.shape) == 2 and edge_index.shape[0] == 2:
|
|
67
|
+
row, col = edge_index[0], edge_index[1]
|
|
68
|
+
adj = scatter(
|
|
69
|
+
ops.stack([row, col], axis=-1),
|
|
70
|
+
ops.ones((ops.shape(row)[0],), dtype=x.dtype),
|
|
71
|
+
shape=(num_nodes, num_nodes),
|
|
72
|
+
)
|
|
73
|
+
else:
|
|
74
|
+
adj = ops.cast(edge_index, x.dtype)
|
|
75
|
+
|
|
76
|
+
# PAN entropy / path calculation
|
|
77
|
+
# M = sum_{k=0}^filter_size weight[k] * A^k
|
|
78
|
+
# M = sum_{k=0}^filter_size exp(-E(k)/T) * A^k
|
|
79
|
+
w = ops.softplus(self.weight)
|
|
80
|
+
eye = ops.eye(num_nodes, dtype=x.dtype)
|
|
81
|
+
M = self.weight[0] * eye
|
|
82
|
+
M = w[0] * eye
|
|
83
|
+
curr_adj = eye
|
|
84
|
+
for k in range(1, self.filter_size + 1):
|
|
85
|
+
curr_adj = ops.matmul(curr_adj, adj)
|
|
86
|
+
M = M + self.weight[k] * curr_adj
|
|
87
|
+
M = M + w[k] * curr_adj
|
|
88
|
+
|
|
89
|
+
deg = ops.sum(M, axis=1)
|
|
90
|
+
deg_inv_sqrt = ops.power(ops.maximum(deg, 1e-12), -0.5)
|
|
91
|
+
# Avoid inf / nan
|
|
92
|
+
deg_inv_sqrt = ops.where(ops.isfinite(deg_inv_sqrt), deg_inv_sqrt, 0.0)
|
|
93
|
+
|
|
94
|
+
M_norm = ops.expand_dims(deg_inv_sqrt, 0) * M * ops.expand_dims(deg_inv_sqrt, 1)
|
|
95
|
+
M_norm = ops.expand_dims(deg_inv_sqrt, 1) * M * ops.expand_dims(deg_inv_sqrt, 0)
|
|
96
|
+
|
|
97
|
+
out = ops.matmul(M_norm, x)
|
|
98
|
+
out = self.lin(out)
|
|
99
|
+
|
|
100
|
+
return out, M_norm
|
|
@@ -0,0 +1,109 @@
|
|
|
1
|
+
from keras import ops
|
|
2
|
+
from keras.layers import Dense
|
|
3
|
+
|
|
4
|
+
from k3_node.layers.conv.message_passing import MessagePassing
|
|
5
|
+
from k3_node.layers.conv.utils import gcn_norm
|
|
6
|
+
|
|
7
|
+
|
|
8
|
+
class PDNConv(MessagePassing):
|
|
9
|
+
r"""The pathfinder discovery network convolutional operator from the
|
|
10
|
+
`"Pathfinder Discovery Networks for Neural Message Passing"
|
|
11
|
+
<https://arxiv.org/abs/2010.12878>`_ paper.
|
|
12
|
+
|
|
13
|
+
Example:
|
|
14
|
+
```python
|
|
15
|
+
import numpy as np
|
|
16
|
+
from k3_node.layers import PDNConv
|
|
17
|
+
|
|
18
|
+
x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
|
|
19
|
+
edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
|
|
20
|
+
edge_attr = np.random.rand(30, 3).astype("float32") # 3 features per edge
|
|
21
|
+
|
|
22
|
+
layer = PDNConv(in_channels=8, out_channels=16, edge_dim=3, hidden_channels=16)
|
|
23
|
+
out = layer(x, edge_index, edge_attr)
|
|
24
|
+
print(tuple(out.shape)) # (10, 16)
|
|
25
|
+
```
|
|
26
|
+
"""
|
|
27
|
+
def __init__(
|
|
28
|
+
self,
|
|
29
|
+
in_channels: int,
|
|
30
|
+
out_channels: int,
|
|
31
|
+
edge_dim: int,
|
|
32
|
+
hidden_channels: int = 16,
|
|
33
|
+
add_self_loops: bool = True,
|
|
34
|
+
normalize: bool = True,
|
|
35
|
+
bias: bool = True,
|
|
36
|
+
**kwargs,
|
|
37
|
+
):
|
|
38
|
+
kwargs.setdefault("aggr", "add")
|
|
39
|
+
super().__init__(**kwargs)
|
|
40
|
+
|
|
41
|
+
self.in_channels = in_channels
|
|
42
|
+
self.out_channels = out_channels
|
|
43
|
+
self.edge_dim = edge_dim
|
|
44
|
+
self.hidden_channels = hidden_channels
|
|
45
|
+
self.add_self_loops = add_self_loops
|
|
46
|
+
self.normalize = normalize
|
|
47
|
+
self.use_bias = bias
|
|
48
|
+
|
|
49
|
+
self.mlp_1 = Dense(hidden_channels, activation="relu")
|
|
50
|
+
self.mlp_2 = Dense(1, activation="sigmoid")
|
|
51
|
+
self.lin = Dense(out_channels, use_bias=False)
|
|
52
|
+
|
|
53
|
+
if bias:
|
|
54
|
+
self.bias = self.add_weight(
|
|
55
|
+
shape=(out_channels,),
|
|
56
|
+
initializer="zeros",
|
|
57
|
+
name="bias",
|
|
58
|
+
)
|
|
59
|
+
else:
|
|
60
|
+
self.bias = None
|
|
61
|
+
|
|
62
|
+
def build(self, input_shape=None):
|
|
63
|
+
self.mlp_1.build((None, self.edge_dim))
|
|
64
|
+
self.mlp_2.build((None, self.hidden_channels))
|
|
65
|
+
self.lin.build((None, self.in_channels))
|
|
66
|
+
self.built = True
|
|
67
|
+
|
|
68
|
+
def call(self, inputs, edge_index=None, edge_attr=None, **kwargs):
|
|
69
|
+
if edge_index is None:
|
|
70
|
+
if isinstance(inputs, (list, tuple)):
|
|
71
|
+
if len(inputs) == 3:
|
|
72
|
+
x, edge_index, edge_attr = inputs
|
|
73
|
+
elif len(inputs) == 2:
|
|
74
|
+
x, edge_index = inputs
|
|
75
|
+
else:
|
|
76
|
+
raise ValueError(f"Unexpected input length {len(inputs)}")
|
|
77
|
+
else:
|
|
78
|
+
raise ValueError("Expected (x, edge_index) or x and edge_index")
|
|
79
|
+
else:
|
|
80
|
+
x = inputs
|
|
81
|
+
|
|
82
|
+
if not self.built:
|
|
83
|
+
self.build()
|
|
84
|
+
|
|
85
|
+
if edge_attr is not None:
|
|
86
|
+
edge_attr = self.mlp_1(edge_attr)
|
|
87
|
+
edge_attr = ops.squeeze(self.mlp_2(edge_attr), -1)
|
|
88
|
+
|
|
89
|
+
num_nodes = ops.shape(x)[0]
|
|
90
|
+
if self.normalize:
|
|
91
|
+
edge_index, edge_attr = gcn_norm(
|
|
92
|
+
edge_index,
|
|
93
|
+
edge_attr,
|
|
94
|
+
num_nodes=num_nodes,
|
|
95
|
+
add_self_loops=self.add_self_loops,
|
|
96
|
+
dtype=x.dtype,
|
|
97
|
+
)
|
|
98
|
+
|
|
99
|
+
x = self.lin(x)
|
|
100
|
+
out = self.propagate(edge_index, x=x, edge_weight=edge_attr)
|
|
101
|
+
|
|
102
|
+
if self.bias is not None:
|
|
103
|
+
out = out + self.bias
|
|
104
|
+
|
|
105
|
+
return out
|
|
106
|
+
|
|
107
|
+
def message(self, x_j, edge_weight=None):
|
|
108
|
+
return x_j if edge_weight is None else ops.expand_dims(edge_weight, -1) * x_j
|
|
109
|
+
|
|
@@ -0,0 +1,177 @@
|
|
|
1
|
+
from typing import Optional, Union, List, Callable
|
|
2
|
+
from keras import ops, activations
|
|
3
|
+
from keras.layers import Dense
|
|
4
|
+
|
|
5
|
+
from k3_node.layers.conv.message_passing import MessagePassing
|
|
6
|
+
from k3_node.layers.aggr import DegreeScalerAggregation
|
|
7
|
+
from k3_node.ops.creation import repeat
|
|
8
|
+
|
|
9
|
+
|
|
10
|
+
class PNAConv(MessagePassing):
|
|
11
|
+
r"""The Principal Neighbourhood Aggregation graph convolutional operator
|
|
12
|
+
from the `"Principal Neighbourhood Aggregation for Graph Nets"
|
|
13
|
+
<https://arxiv.org/abs/2004.05718>`_ paper.
|
|
14
|
+
|
|
15
|
+
Example:
|
|
16
|
+
```python
|
|
17
|
+
import numpy as np
|
|
18
|
+
from k3_node.layers import PNAConv
|
|
19
|
+
|
|
20
|
+
x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
|
|
21
|
+
edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
|
|
22
|
+
|
|
23
|
+
deg = np.array([0, 2, 4, 3, 1]) # in-degree histogram of the training graphs
|
|
24
|
+
layer = PNAConv(in_channels=8, out_channels=16, aggregators=["mean", "max", "min", "std"],
|
|
25
|
+
scalers=["identity", "amplification"], deg=deg)
|
|
26
|
+
out = layer(x, edge_index)
|
|
27
|
+
print(tuple(out.shape)) # (10, 16)
|
|
28
|
+
```
|
|
29
|
+
"""
|
|
30
|
+
@staticmethod
|
|
31
|
+
def get_degree_histogram(graphs):
|
|
32
|
+
r"""Returns the histogram of in-degrees over ``graphs`` (a dataset, a list of graphs or a
|
|
33
|
+
loader), which PNA uses to normalize its degree scalers:
|
|
34
|
+
``deg[d]`` is the number of nodes with ``d`` incoming edges."""
|
|
35
|
+
import numpy as np
|
|
36
|
+
|
|
37
|
+
counts = []
|
|
38
|
+
for graph in graphs:
|
|
39
|
+
target = np.asarray(ops.convert_to_numpy(graph.edge_index))[1]
|
|
40
|
+
counts.append(np.bincount(np.bincount(target, minlength=graph.num_nodes)))
|
|
41
|
+
hist = np.zeros(max(len(c) for c in counts), dtype=np.int64)
|
|
42
|
+
for c in counts:
|
|
43
|
+
hist[: len(c)] += c
|
|
44
|
+
return hist
|
|
45
|
+
|
|
46
|
+
def __init__(
|
|
47
|
+
self,
|
|
48
|
+
in_channels: int,
|
|
49
|
+
out_channels: int,
|
|
50
|
+
aggregators: List[str],
|
|
51
|
+
scalers: List[str],
|
|
52
|
+
deg,
|
|
53
|
+
edge_dim: Optional[int] = None,
|
|
54
|
+
towers: int = 1,
|
|
55
|
+
pre_layers: int = 1,
|
|
56
|
+
post_layers: int = 1,
|
|
57
|
+
divide_input: bool = False,
|
|
58
|
+
act: Union[str, Callable, None] = "relu",
|
|
59
|
+
train_norm: bool = False,
|
|
60
|
+
**kwargs,
|
|
61
|
+
):
|
|
62
|
+
aggr = DegreeScalerAggregation(aggregators, scalers, deg, train_norm)
|
|
63
|
+
super().__init__(aggr=aggr, node_dim=0, **kwargs)
|
|
64
|
+
|
|
65
|
+
if divide_input:
|
|
66
|
+
assert in_channels % towers == 0
|
|
67
|
+
assert out_channels % towers == 0
|
|
68
|
+
|
|
69
|
+
self.in_channels = in_channels
|
|
70
|
+
self.out_channels = out_channels
|
|
71
|
+
self.aggregators = aggregators
|
|
72
|
+
self.scalers = scalers
|
|
73
|
+
self.edge_dim = edge_dim
|
|
74
|
+
self.towers = towers
|
|
75
|
+
self.divide_input = divide_input
|
|
76
|
+
self.act = activations.get(act) if act is not None else None
|
|
77
|
+
|
|
78
|
+
self.F_in = in_channels // towers if divide_input else in_channels
|
|
79
|
+
self.F_out = out_channels // towers
|
|
80
|
+
|
|
81
|
+
if edge_dim is not None:
|
|
82
|
+
self.edge_encoder = Dense(self.F_in, use_bias=False)
|
|
83
|
+
else:
|
|
84
|
+
self.edge_encoder = None
|
|
85
|
+
|
|
86
|
+
self.pre_nns = []
|
|
87
|
+
self.post_nns = []
|
|
88
|
+
for _ in range(towers):
|
|
89
|
+
pre_dim = (3 if edge_dim else 2) * self.F_in
|
|
90
|
+
pre_layers_list = [Dense(self.F_in, use_bias=True)]
|
|
91
|
+
for _ in range(pre_layers - 1):
|
|
92
|
+
pre_layers_list.append(Dense(self.F_in, activation=self.act, use_bias=True))
|
|
93
|
+
self.pre_nns.append(pre_layers_list)
|
|
94
|
+
|
|
95
|
+
post_in_dim = (len(aggregators) * len(scalers) + 1) * self.F_in
|
|
96
|
+
post_layers_list = [Dense(self.F_out, use_bias=True)]
|
|
97
|
+
for _ in range(post_layers - 1):
|
|
98
|
+
post_layers_list.append(Dense(self.F_out, activation=self.act, use_bias=True))
|
|
99
|
+
self.post_nns.append(post_layers_list)
|
|
100
|
+
|
|
101
|
+
self.lin = Dense(out_channels, use_bias=True)
|
|
102
|
+
|
|
103
|
+
def build(self, input_shape=None):
|
|
104
|
+
if self.edge_encoder is not None:
|
|
105
|
+
self.edge_encoder.build((None, self.edge_dim))
|
|
106
|
+
for pre_list in self.pre_nns:
|
|
107
|
+
dim = (3 if self.edge_dim else 2) * self.F_in
|
|
108
|
+
for l in pre_list:
|
|
109
|
+
l.build((None, dim))
|
|
110
|
+
dim = self.F_in
|
|
111
|
+
for post_list in self.post_nns:
|
|
112
|
+
dim = (len(self.aggregators) * len(self.scalers) + 1) * self.F_in
|
|
113
|
+
for l in post_list:
|
|
114
|
+
l.build((None, dim))
|
|
115
|
+
dim = self.F_out
|
|
116
|
+
self.lin.build((None, self.out_channels))
|
|
117
|
+
self.built = True
|
|
118
|
+
|
|
119
|
+
def call(self, inputs, edge_index=None, edge_attr=None, **kwargs):
|
|
120
|
+
if edge_index is None:
|
|
121
|
+
if isinstance(inputs, (list, tuple)):
|
|
122
|
+
if len(inputs) == 3:
|
|
123
|
+
x, edge_index, edge_attr = inputs
|
|
124
|
+
elif len(inputs) == 2:
|
|
125
|
+
x, edge_index = inputs
|
|
126
|
+
else:
|
|
127
|
+
raise ValueError(f"Unexpected input length {len(inputs)}")
|
|
128
|
+
else:
|
|
129
|
+
raise ValueError("Expected (x, edge_index) or x and edge_index")
|
|
130
|
+
else:
|
|
131
|
+
x = inputs
|
|
132
|
+
|
|
133
|
+
if not self.built:
|
|
134
|
+
self.build()
|
|
135
|
+
|
|
136
|
+
num_nodes = ops.shape(x)[0]
|
|
137
|
+
if self.divide_input:
|
|
138
|
+
x_towers = ops.reshape(x, (-1, self.towers, self.F_in))
|
|
139
|
+
else:
|
|
140
|
+
x_towers = repeat(ops.expand_dims(x, 1), self.towers, axis=1)
|
|
141
|
+
|
|
142
|
+
out = self.propagate(
|
|
143
|
+
edge_index,
|
|
144
|
+
x=x_towers,
|
|
145
|
+
edge_attr=edge_attr,
|
|
146
|
+
size=(num_nodes, num_nodes),
|
|
147
|
+
)
|
|
148
|
+
|
|
149
|
+
out = ops.concatenate([x_towers, out], axis=-1)
|
|
150
|
+
|
|
151
|
+
outs = []
|
|
152
|
+
for i, post_list in enumerate(self.post_nns):
|
|
153
|
+
h_i = out[:, i]
|
|
154
|
+
for l in post_list:
|
|
155
|
+
h_i = l(h_i)
|
|
156
|
+
outs.append(h_i)
|
|
157
|
+
|
|
158
|
+
out = ops.concatenate(outs, axis=-1)
|
|
159
|
+
return self.lin(out)
|
|
160
|
+
|
|
161
|
+
def message(self, x_i, x_j, edge_attr=None):
|
|
162
|
+
if edge_attr is not None and self.edge_encoder is not None:
|
|
163
|
+
edge_attr = self.edge_encoder(edge_attr)
|
|
164
|
+
edge_attr = repeat(ops.expand_dims(edge_attr, 1), self.towers, axis=1)
|
|
165
|
+
h = ops.concatenate([x_i, x_j, edge_attr], axis=-1)
|
|
166
|
+
else:
|
|
167
|
+
h = ops.concatenate([x_i, x_j], axis=-1)
|
|
168
|
+
|
|
169
|
+
hs = []
|
|
170
|
+
for i, pre_list in enumerate(self.pre_nns):
|
|
171
|
+
h_i = h[:, i]
|
|
172
|
+
for l in pre_list:
|
|
173
|
+
h_i = l(h_i)
|
|
174
|
+
hs.append(h_i)
|
|
175
|
+
|
|
176
|
+
return ops.stack(hs, axis=1)
|
|
177
|
+
|
|
@@ -0,0 +1,101 @@
|
|
|
1
|
+
from typing import Optional, Callable, Tuple
|
|
2
|
+
from keras import ops
|
|
3
|
+
|
|
4
|
+
from k3_node.layers.conv.message_passing import MessagePassing
|
|
5
|
+
from k3_node.layers.conv.utils import remove_self_loops, add_self_loops
|
|
6
|
+
|
|
7
|
+
|
|
8
|
+
class PointNetConv(MessagePassing):
|
|
9
|
+
r"""The PointNet set abstraction layer from the `"PointNet++: Deep
|
|
10
|
+
Hierarchical Feature Learning on Point Sets in a Metric Space"
|
|
11
|
+
<https://arxiv.org/abs/1706.02413>`_ paper.
|
|
12
|
+
|
|
13
|
+
Example:
|
|
14
|
+
```python
|
|
15
|
+
import numpy as np
|
|
16
|
+
import keras
|
|
17
|
+
from k3_node.layers import PointNetConv
|
|
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
|
+
pos = np.random.rand(10, 3).astype("float32") # 3D node positions
|
|
22
|
+
|
|
23
|
+
local_nn = keras.Sequential([keras.layers.Dense(16, activation="relu"), keras.layers.Dense(16)])
|
|
24
|
+
layer = PointNetConv(local_nn=local_nn)
|
|
25
|
+
out = layer(x, pos, edge_index)
|
|
26
|
+
print(tuple(out.shape)) # (10, 16)
|
|
27
|
+
```
|
|
28
|
+
"""
|
|
29
|
+
def __init__(
|
|
30
|
+
self,
|
|
31
|
+
local_nn: Optional[Callable] = None,
|
|
32
|
+
global_nn: Optional[Callable] = None,
|
|
33
|
+
add_self_loops: bool = True,
|
|
34
|
+
**kwargs,
|
|
35
|
+
):
|
|
36
|
+
kwargs.setdefault("aggr", "max")
|
|
37
|
+
super().__init__(**kwargs)
|
|
38
|
+
|
|
39
|
+
self.local_nn = local_nn
|
|
40
|
+
self.global_nn = global_nn
|
|
41
|
+
self.add_self_loops = add_self_loops
|
|
42
|
+
|
|
43
|
+
def build(self, input_shape=None):
|
|
44
|
+
self.built = True
|
|
45
|
+
|
|
46
|
+
def call(self, inputs, pos=None, edge_index=None, **kwargs):
|
|
47
|
+
if edge_index is None:
|
|
48
|
+
if isinstance(inputs, (list, tuple)):
|
|
49
|
+
if len(inputs) == 3:
|
|
50
|
+
x, pos, edge_index = inputs
|
|
51
|
+
elif len(inputs) == 2:
|
|
52
|
+
# x can be None or pos
|
|
53
|
+
x, edge_index = inputs
|
|
54
|
+
else:
|
|
55
|
+
raise ValueError(f"Unexpected input length {len(inputs)}")
|
|
56
|
+
else:
|
|
57
|
+
raise ValueError("Expected (x, pos, edge_index) or (x, edge_index)")
|
|
58
|
+
else:
|
|
59
|
+
x = inputs
|
|
60
|
+
|
|
61
|
+
if not self.built:
|
|
62
|
+
self.build()
|
|
63
|
+
|
|
64
|
+
if isinstance(pos, (list, tuple)):
|
|
65
|
+
pos_src, pos_dst = pos
|
|
66
|
+
else:
|
|
67
|
+
pos_src = pos_dst = pos
|
|
68
|
+
|
|
69
|
+
if isinstance(x, (list, tuple)):
|
|
70
|
+
x_src, x_dst = x
|
|
71
|
+
else:
|
|
72
|
+
x_src = x_dst = x
|
|
73
|
+
|
|
74
|
+
num_nodes = ops.shape(pos_dst)[0]
|
|
75
|
+
if self.add_self_loops:
|
|
76
|
+
edge_index, _ = remove_self_loops(edge_index)
|
|
77
|
+
edge_index, _ = add_self_loops(edge_index, num_nodes=num_nodes)
|
|
78
|
+
|
|
79
|
+
out = self.propagate(
|
|
80
|
+
edge_index,
|
|
81
|
+
x=(x_src, x_dst),
|
|
82
|
+
pos=(pos_src, pos_dst),
|
|
83
|
+
size=(ops.shape(pos_src)[0], num_nodes),
|
|
84
|
+
)
|
|
85
|
+
|
|
86
|
+
if self.global_nn is not None:
|
|
87
|
+
out = self.global_nn(out)
|
|
88
|
+
|
|
89
|
+
return out
|
|
90
|
+
|
|
91
|
+
def message(self, x_j=None, pos_i=None, pos_j=None):
|
|
92
|
+
msg = pos_j - pos_i
|
|
93
|
+
if x_j is not None:
|
|
94
|
+
msg = ops.concatenate([x_j, msg], axis=-1)
|
|
95
|
+
if self.local_nn is not None:
|
|
96
|
+
msg = self.local_nn(msg)
|
|
97
|
+
return msg
|
|
98
|
+
|
|
99
|
+
|
|
100
|
+
PointConv = PointNetConv
|
|
101
|
+
|
|
@@ -0,0 +1,90 @@
|
|
|
1
|
+
from typing import Callable
|
|
2
|
+
from keras import ops
|
|
3
|
+
|
|
4
|
+
from k3_node.layers.conv.message_passing import MessagePassing
|
|
5
|
+
|
|
6
|
+
|
|
7
|
+
class PointGNNConv(MessagePassing):
|
|
8
|
+
r"""The PointGNN graph convolutional operator from the
|
|
9
|
+
`"Point-GNN: Graph Neural Network for 3D Object Detection in a Point Cloud"
|
|
10
|
+
<https://arxiv.org/abs/2003.01251>`_ paper.
|
|
11
|
+
|
|
12
|
+
Example:
|
|
13
|
+
```python
|
|
14
|
+
import numpy as np
|
|
15
|
+
import keras
|
|
16
|
+
from k3_node.layers import PointGNNConv
|
|
17
|
+
|
|
18
|
+
x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
|
|
19
|
+
edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
|
|
20
|
+
pos = np.random.rand(10, 3).astype("float32") # 3D node positions
|
|
21
|
+
|
|
22
|
+
layer = PointGNNConv(
|
|
23
|
+
mlp_h=keras.layers.Dense(3), # predicts a position offset per node
|
|
24
|
+
mlp_f=keras.layers.Dense(16), # edge feature network
|
|
25
|
+
mlp_g=keras.layers.Dense(8), # node update network (output size = input features)
|
|
26
|
+
)
|
|
27
|
+
out = layer(x, pos, edge_index)
|
|
28
|
+
print(tuple(out.shape)) # (10, 8)
|
|
29
|
+
```
|
|
30
|
+
"""
|
|
31
|
+
def __init__(
|
|
32
|
+
self,
|
|
33
|
+
mlp_h: Callable,
|
|
34
|
+
mlp_f: Callable,
|
|
35
|
+
mlp_g: Callable,
|
|
36
|
+
**kwargs,
|
|
37
|
+
):
|
|
38
|
+
kwargs.setdefault("aggr", "max")
|
|
39
|
+
super().__init__(node_dim=0, **kwargs)
|
|
40
|
+
|
|
41
|
+
self.mlp_h = mlp_h
|
|
42
|
+
self.mlp_f = mlp_f
|
|
43
|
+
self.mlp_g = mlp_g
|
|
44
|
+
|
|
45
|
+
def build(self, input_shape=None):
|
|
46
|
+
self.built = True
|
|
47
|
+
|
|
48
|
+
def call(self, inputs, pos=None, edge_index=None, **kwargs):
|
|
49
|
+
if edge_index is None:
|
|
50
|
+
if isinstance(inputs, (list, tuple)):
|
|
51
|
+
if len(inputs) == 3:
|
|
52
|
+
x, pos, edge_index = inputs
|
|
53
|
+
elif len(inputs) == 2:
|
|
54
|
+
x, edge_index = inputs
|
|
55
|
+
else:
|
|
56
|
+
raise ValueError(f"Unexpected input length {len(inputs)}")
|
|
57
|
+
else:
|
|
58
|
+
raise ValueError("Expected (x, pos, edge_index)")
|
|
59
|
+
else:
|
|
60
|
+
x = inputs
|
|
61
|
+
|
|
62
|
+
if not self.built:
|
|
63
|
+
self.build()
|
|
64
|
+
|
|
65
|
+
if isinstance(x, (list, tuple)):
|
|
66
|
+
x_src, x_dst = x
|
|
67
|
+
else:
|
|
68
|
+
x_src = x_dst = x
|
|
69
|
+
|
|
70
|
+
if isinstance(pos, (list, tuple)):
|
|
71
|
+
pos_src, pos_dst = pos
|
|
72
|
+
else:
|
|
73
|
+
pos_src = pos_dst = pos
|
|
74
|
+
|
|
75
|
+
num_nodes = ops.shape(pos_dst)[0]
|
|
76
|
+
out = self.propagate(
|
|
77
|
+
edge_index,
|
|
78
|
+
x=(x_src, x_dst),
|
|
79
|
+
pos=(pos_src, pos_dst),
|
|
80
|
+
size=(ops.shape(pos_src)[0], num_nodes),
|
|
81
|
+
)
|
|
82
|
+
|
|
83
|
+
out = self.mlp_g(out)
|
|
84
|
+
return x_dst + out
|
|
85
|
+
|
|
86
|
+
def message(self, pos_j, pos_i, x_i, x_j):
|
|
87
|
+
delta = self.mlp_h(x_i)
|
|
88
|
+
e = ops.concatenate([pos_j - pos_i + delta, x_j], axis=-1)
|
|
89
|
+
return self.mlp_f(e)
|
|
90
|
+
|