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,404 @@
|
|
|
1
|
+
"""LPFormer, ported from PyG (``torch_geometric.nn.models.LPFormer``).
|
|
2
|
+
|
|
3
|
+
The choice of the nodes each candidate link attends to (common neighbors, one-hop neighbors and
|
|
4
|
+
nodes with a high personalized PageRank) depends on the graph only and is computed on the host
|
|
5
|
+
with SciPy sparse matrices; the learnable parts run with Keras ops. The model runs eagerly.
|
|
6
|
+
"""
|
|
7
|
+
import math
|
|
8
|
+
from typing import List, Optional
|
|
9
|
+
|
|
10
|
+
import keras
|
|
11
|
+
import numpy as np
|
|
12
|
+
from keras import ops
|
|
13
|
+
|
|
14
|
+
from k3_node.models.basic_gnn import GCN
|
|
15
|
+
from k3_node.ops.segment import segment_sum
|
|
16
|
+
|
|
17
|
+
|
|
18
|
+
def get_ppr(edge_index, num_nodes: int, alpha: float = 0.15, eps: float = 5e-5):
|
|
19
|
+
r"""Approximate personalized PageRank of every node (the Andersen push algorithm used by
|
|
20
|
+
PyG's ``get_ppr``). Returns a ``scipy.sparse.csr_matrix`` whose row ``i`` holds the PPR scores
|
|
21
|
+
of the nodes reachable from ``i``.
|
|
22
|
+
|
|
23
|
+
Example:
|
|
24
|
+
```python
|
|
25
|
+
import numpy as np
|
|
26
|
+
from k3_node.models.lpformer import get_ppr
|
|
27
|
+
|
|
28
|
+
edge_index = np.array([[0, 1, 1, 2], [1, 0, 2, 1]]) # a path 0 - 1 - 2
|
|
29
|
+
ppr = get_ppr(edge_index, num_nodes=3)
|
|
30
|
+
print(ppr.shape, round(float(ppr[0, 0]), 2)) # (3, 3) 0.19
|
|
31
|
+
```
|
|
32
|
+
"""
|
|
33
|
+
import scipy.sparse as sp
|
|
34
|
+
|
|
35
|
+
from k3_node.ops.host import to_numpy
|
|
36
|
+
|
|
37
|
+
edge_index = np.asarray(to_numpy(edge_index)).astype(np.int64)
|
|
38
|
+
order = np.lexsort((edge_index[1], edge_index[0])) # CSR with sorted columns, as PyG's EdgeIndex
|
|
39
|
+
col = edge_index[1][order]
|
|
40
|
+
rowptr = np.concatenate([[0], np.cumsum(np.bincount(edge_index[0], minlength=num_nodes))])
|
|
41
|
+
cols_list, vals_list = _ppr_push(rowptr, col, alpha, eps)
|
|
42
|
+
rows = np.repeat(np.arange(num_nodes), [len(c) for c in cols_list])
|
|
43
|
+
cols = np.concatenate(cols_list) if cols_list else np.zeros(0, np.int64)
|
|
44
|
+
vals = np.concatenate(vals_list) if vals_list else np.zeros(0)
|
|
45
|
+
return sp.csr_matrix((np.array(vals, dtype=np.float32), (rows, cols)), shape=(num_nodes, num_nodes))
|
|
46
|
+
|
|
47
|
+
|
|
48
|
+
def _ppr_push_python(rowptr, col, alpha, eps):
|
|
49
|
+
"""PyG's Andersen push (``torch_geometric.utils.ppr._get_ppr``), one source node at a time."""
|
|
50
|
+
alpha_eps = alpha * eps
|
|
51
|
+
cols_list, vals_list = [], []
|
|
52
|
+
for inode in range(len(rowptr) - 1):
|
|
53
|
+
p, r, q, in_q = {inode: 0.0}, {inode: alpha}, [inode], {inode}
|
|
54
|
+
while q:
|
|
55
|
+
unode = q.pop()
|
|
56
|
+
in_q.discard(unode)
|
|
57
|
+
res = r.get(unode, 0.0)
|
|
58
|
+
p[unode] = p.get(unode, 0.0) + res
|
|
59
|
+
r[unode] = 0.0
|
|
60
|
+
start, end = rowptr[unode], rowptr[unode + 1]
|
|
61
|
+
ucount = end - start
|
|
62
|
+
for vnode in col[start:end]:
|
|
63
|
+
vnode = int(vnode)
|
|
64
|
+
r[vnode] = r.get(vnode, 0.0) + (1 - alpha) * res / ucount
|
|
65
|
+
if r[vnode] >= alpha_eps * (rowptr[vnode + 1] - rowptr[vnode]) and vnode not in in_q:
|
|
66
|
+
q.append(vnode)
|
|
67
|
+
in_q.add(vnode)
|
|
68
|
+
cols_list.append(np.fromiter(p.keys(), dtype=np.int64, count=len(p)))
|
|
69
|
+
vals_list.append(np.fromiter(p.values(), dtype=np.float64, count=len(p)))
|
|
70
|
+
return cols_list, vals_list
|
|
71
|
+
|
|
72
|
+
|
|
73
|
+
_PPR_NUMBA = None
|
|
74
|
+
|
|
75
|
+
|
|
76
|
+
def _ppr_push(rowptr, col, alpha, eps):
|
|
77
|
+
"""Runs the push algorithm compiled with numba when it is installed (as PyG does)."""
|
|
78
|
+
global _PPR_NUMBA
|
|
79
|
+
try:
|
|
80
|
+
import numba
|
|
81
|
+
except ImportError:
|
|
82
|
+
return _ppr_push_python(rowptr, col, alpha, eps)
|
|
83
|
+
if _PPR_NUMBA is None:
|
|
84
|
+
def push(rowptr, col, alpha, eps):
|
|
85
|
+
num_nodes = len(rowptr) - 1
|
|
86
|
+
alpha_eps = alpha * eps
|
|
87
|
+
js = [[0]] * num_nodes
|
|
88
|
+
vals = [[0.0]] * num_nodes
|
|
89
|
+
for inode_uint in numba.prange(num_nodes):
|
|
90
|
+
inode = numba.int64(inode_uint)
|
|
91
|
+
p = {inode: 0.0}
|
|
92
|
+
r = {}
|
|
93
|
+
r[inode] = alpha
|
|
94
|
+
q = [inode]
|
|
95
|
+
while len(q) > 0:
|
|
96
|
+
unode = q.pop()
|
|
97
|
+
res = r[unode] if unode in r else 0
|
|
98
|
+
if unode in p:
|
|
99
|
+
p[unode] += res
|
|
100
|
+
else:
|
|
101
|
+
p[unode] = res
|
|
102
|
+
r[unode] = 0
|
|
103
|
+
start, end = rowptr[unode], rowptr[unode + 1]
|
|
104
|
+
ucount = end - start
|
|
105
|
+
for vnode in col[start:end]:
|
|
106
|
+
_val = (1 - alpha) * res / ucount
|
|
107
|
+
if vnode in r:
|
|
108
|
+
r[vnode] += _val
|
|
109
|
+
else:
|
|
110
|
+
r[vnode] = _val
|
|
111
|
+
res_vnode = r[vnode] if vnode in r else 0
|
|
112
|
+
vcount = rowptr[vnode + 1] - rowptr[vnode]
|
|
113
|
+
if res_vnode >= alpha_eps * vcount:
|
|
114
|
+
if vnode not in q:
|
|
115
|
+
q.append(vnode)
|
|
116
|
+
js[inode_uint] = list(p.keys())
|
|
117
|
+
vals[inode_uint] = list(p.values())
|
|
118
|
+
return js, vals
|
|
119
|
+
|
|
120
|
+
_PPR_NUMBA = numba.jit(nopython=True, parallel=True)(push)
|
|
121
|
+
js, vals = _PPR_NUMBA(rowptr, col, alpha, eps)
|
|
122
|
+
return [np.asarray(j, dtype=np.int64) for j in js], [np.asarray(v, dtype=np.float64) for v in vals]
|
|
123
|
+
|
|
124
|
+
|
|
125
|
+
def compute_ppr_matrix(edge_index, num_nodes: int, alpha: float = 0.15, eps: float = 5e-5):
|
|
126
|
+
r"""Alias of :func:`get_ppr`."""
|
|
127
|
+
return get_ppr(edge_index, num_nodes, alpha=alpha, eps=eps)
|
|
128
|
+
|
|
129
|
+
|
|
130
|
+
def _lookup(matrix, rows, cols):
|
|
131
|
+
"""``matrix[rows[k], cols[k]]`` for every ``k``, as a 1-D array (also when empty)."""
|
|
132
|
+
if len(rows) == 0:
|
|
133
|
+
return np.zeros(0, dtype=np.float32)
|
|
134
|
+
values = matrix[rows, cols]
|
|
135
|
+
values = values.toarray() if hasattr(values, "toarray") else values
|
|
136
|
+
return np.asarray(values, dtype=np.float32).ravel()
|
|
137
|
+
|
|
138
|
+
|
|
139
|
+
class MLP(keras.layers.Layer):
|
|
140
|
+
r"""The small MLP of LPFormer: linear layers with (layer) normalization, ReLU and dropout
|
|
141
|
+
between them; the last dimension is squeezed if it has size 1."""
|
|
142
|
+
|
|
143
|
+
def __init__(self, in_channels: int, hid_channels: int, out_channels: int, num_layers: int = 2,
|
|
144
|
+
drop: float = 0.0, norm: Optional[str] = "layer", **kwargs):
|
|
145
|
+
super().__init__(**kwargs)
|
|
146
|
+
self.linears = ([keras.layers.Dense(out_channels)] if num_layers == 1 else
|
|
147
|
+
[keras.layers.Dense(hid_channels) for _ in range(num_layers - 1)] + [keras.layers.Dense(out_channels)])
|
|
148
|
+
if norm == "batch":
|
|
149
|
+
self.norm = keras.layers.BatchNormalization(momentum=0.9, epsilon=1e-5)
|
|
150
|
+
elif norm == "layer":
|
|
151
|
+
self.norm = keras.layers.LayerNormalization(epsilon=1e-5)
|
|
152
|
+
else:
|
|
153
|
+
self.norm = None
|
|
154
|
+
self.dropout = keras.layers.Dropout(drop)
|
|
155
|
+
|
|
156
|
+
def call(self, x, training=False):
|
|
157
|
+
for lin in self.linears[:-1]:
|
|
158
|
+
x = lin(x)
|
|
159
|
+
x = self.norm(x, training=training) if self.norm is not None else x
|
|
160
|
+
x = self.dropout(ops.relu(x), training=training)
|
|
161
|
+
x = self.linears[-1](x)
|
|
162
|
+
return ops.squeeze(x, axis=-1) if x.shape[-1] == 1 else x
|
|
163
|
+
|
|
164
|
+
|
|
165
|
+
class LPAttLayer(keras.layers.Layer):
|
|
166
|
+
r"""Attention of every candidate link over its selected nodes (PyG's ``LPAttLayer``).
|
|
167
|
+
|
|
168
|
+
Used inside :class:`LPFormer`. ``edge_index[0]`` is the link and ``edge_index[1]`` a node it
|
|
169
|
+
attends to; ``edge_feats`` holds the two end-node features of every link side by side and
|
|
170
|
+
``ppr_rpes`` a relative positional encoding for every (link, node) pair.
|
|
171
|
+
|
|
172
|
+
Example:
|
|
173
|
+
```python
|
|
174
|
+
import numpy as np
|
|
175
|
+
from k3_node.models import LPAttLayer
|
|
176
|
+
|
|
177
|
+
layer = LPAttLayer(in_channels=8, out_channels=8, node_dim=None, num_heads=2, dropout=0.0)
|
|
178
|
+
edge_feats = np.random.rand(4, 16).astype("float32") # 4 links: [x_src | x_dst]
|
|
179
|
+
node_feats = np.random.rand(10, 8).astype("float32") # 10 nodes
|
|
180
|
+
edge_index = np.stack([np.repeat(np.arange(4), 3), np.random.randint(0, 10, size=12)])
|
|
181
|
+
ppr_rpes = np.random.rand(12, 8).astype("float32") # one encoding per (link, node) pair
|
|
182
|
+
out = layer(edge_index, edge_feats, node_feats, ppr_rpes)
|
|
183
|
+
print(tuple(out.shape)) # (4, 16)
|
|
184
|
+
```
|
|
185
|
+
"""
|
|
186
|
+
|
|
187
|
+
def __init__(self, in_channels: int, out_channels: int, node_dim: Optional[int], num_heads: int,
|
|
188
|
+
dropout: float, concat: bool = True, **kwargs):
|
|
189
|
+
super().__init__(**kwargs)
|
|
190
|
+
self.in_channels, self.out_channels, self.heads, self.concat = in_channels, out_channels, num_heads, concat
|
|
191
|
+
self.negative_slope = 0.2
|
|
192
|
+
self.lin_l = keras.layers.Dense(num_heads * out_channels, kernel_initializer="glorot_uniform")
|
|
193
|
+
self.lin_r = keras.layers.Dense(num_heads * out_channels, kernel_initializer="glorot_uniform")
|
|
194
|
+
self.att = self.add_weight(shape=(1, num_heads, out_channels), initializer="glorot_uniform", name="att")
|
|
195
|
+
self.bias = self.add_weight(shape=(num_heads * out_channels if concat else out_channels,),
|
|
196
|
+
initializer="zeros", name="bias")
|
|
197
|
+
self.post_att_norm = keras.layers.LayerNormalization(epsilon=1e-5)
|
|
198
|
+
self.dropout = keras.layers.Dropout(dropout)
|
|
199
|
+
|
|
200
|
+
def call(self, edge_index, edge_feats, node_feats, ppr_rpes, training=False):
|
|
201
|
+
H, C = self.heads, self.out_channels
|
|
202
|
+
pair, node = edge_index[0], edge_index[1] # "target_to_source": link i attends to node j
|
|
203
|
+
num_links = edge_feats.shape[0]
|
|
204
|
+
x_i = ops.take(edge_feats, pair, axis=0)
|
|
205
|
+
x_j = ops.concatenate([ops.take(node_feats, node, axis=0), ppr_rpes], axis=-1)
|
|
206
|
+
x_j = ops.reshape(self.lin_r(x_j), (-1, H, C))
|
|
207
|
+
e1, e2 = ops.split(x_i, 2, axis=-1)
|
|
208
|
+
x = ops.leaky_relu(x_j * (ops.reshape(self.lin_l(e1), (-1, H, C)) + ops.reshape(self.lin_l(e2), (-1, H, C))),
|
|
209
|
+
negative_slope=self.negative_slope)
|
|
210
|
+
alpha = ops.sum(x * self.att, axis=-1) # [K, H]
|
|
211
|
+
from k3_node.layers.conv.utils import softmax
|
|
212
|
+
|
|
213
|
+
alpha = softmax(alpha, pair, num_nodes=num_links)
|
|
214
|
+
out = segment_sum(x_j * ops.expand_dims(alpha, -1), pair, num_segments=num_links) # [B, H, C]
|
|
215
|
+
out = ops.reshape(out, (-1, H * C)) if self.concat else ops.mean(out, axis=1)
|
|
216
|
+
out = self.post_att_norm(out + self.bias)
|
|
217
|
+
return self.dropout(out, training=training)
|
|
218
|
+
|
|
219
|
+
|
|
220
|
+
class LPFormer(keras.Model):
|
|
221
|
+
r"""The LPFormer model from the `"LPFormer: An Adaptive Graph Transformer for Link Prediction"
|
|
222
|
+
<https://arxiv.org/abs/2310.11009>`_ paper, as in PyG.
|
|
223
|
+
|
|
224
|
+
For every candidate link it attends over the common neighbors, the one-hop neighbors and the
|
|
225
|
+
other nodes with a high personalized PageRank (PPR) from both endpoints (``ppr_thresholds``
|
|
226
|
+
for the three kinds), using their PPR scores as relative positional encodings, and combines
|
|
227
|
+
this with counts of each kind of node and with the GCN embeddings of the two endpoints.
|
|
228
|
+
|
|
229
|
+
Args:
|
|
230
|
+
in_channels (int): Input feature dimension.
|
|
231
|
+
hidden_channels (int): Hidden dimension.
|
|
232
|
+
num_gnn_layers (int, optional): Number of GCN layers. (default: ``2``)
|
|
233
|
+
gnn_dropout (float, optional): GCN dropout. (default: ``0.1``)
|
|
234
|
+
num_transformer_layers (int, optional): Number of attention layers. (default: ``1``)
|
|
235
|
+
num_heads (int, optional): Number of attention heads. (default: ``1``)
|
|
236
|
+
transformer_dropout (float, optional): Attention dropout; during training this share of
|
|
237
|
+
the selected nodes is also dropped. (default: ``0.1``)
|
|
238
|
+
ppr_thresholds (list, optional): Minimum PPR of common neighbors, one-hop neighbors and
|
|
239
|
+
other nodes. (default: ``[0, 1e-4, 1e-2]``)
|
|
240
|
+
|
|
241
|
+
Call arguments: ``batch`` (the ``[2, num_links]`` candidate links), ``x`` (node features),
|
|
242
|
+
``edge_index`` (the graph) and the keyword ``ppr_matrix`` (from :meth:`calc_sparse_ppr`). Returns one logit
|
|
243
|
+
per link. The node selection runs on the host, so train with ``run_eagerly=True`` or
|
|
244
|
+
:func:`~k3_node.training.gradient_step`.
|
|
245
|
+
|
|
246
|
+
Example:
|
|
247
|
+
```python
|
|
248
|
+
import numpy as np
|
|
249
|
+
from k3_node.models import LPFormer
|
|
250
|
+
|
|
251
|
+
x = np.random.rand(10, 16).astype("float32") # 10 nodes with 16 features each
|
|
252
|
+
edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
|
|
253
|
+
model = LPFormer(in_channels=16, hidden_channels=16)
|
|
254
|
+
ppr = model.calc_sparse_ppr(edge_index, num_nodes=10)
|
|
255
|
+
target_links = np.array([[0, 1], [2, 3]]) # the (source, target) pairs to score
|
|
256
|
+
print(tuple(model(target_links, x, edge_index, ppr_matrix=ppr).shape)) # (2,): one logit per link
|
|
257
|
+
```
|
|
258
|
+
"""
|
|
259
|
+
|
|
260
|
+
def __init__(self, in_channels: int, hidden_channels: int, num_gnn_layers: int = 2, gnn_dropout: float = 0.1,
|
|
261
|
+
num_transformer_layers: int = 1, num_heads: int = 1, transformer_dropout: float = 0.1,
|
|
262
|
+
ppr_thresholds: Optional[List[float]] = None, **kwargs):
|
|
263
|
+
super().__init__(**kwargs)
|
|
264
|
+
ppr_thresholds = [0, 1e-4, 1e-2] if ppr_thresholds is None else ppr_thresholds
|
|
265
|
+
if len(ppr_thresholds) != 3:
|
|
266
|
+
raise ValueError("Argument 'ppr_thresholds' must only be length 3!")
|
|
267
|
+
self.thresh_cn, self.thresh_1hop, self.thresh_non1hop = ppr_thresholds
|
|
268
|
+
self.in_dim, self.hid_dim = in_channels, hidden_channels
|
|
269
|
+
self.trans_drop = transformer_dropout
|
|
270
|
+
self.gnn = GCN(in_channels, hidden_channels, num_gnn_layers, dropout=gnn_dropout, norm="layer_norm")
|
|
271
|
+
self.gnn_norm = keras.layers.LayerNormalization(epsilon=1e-5)
|
|
272
|
+
self.input_dropout = keras.layers.Dropout(gnn_dropout)
|
|
273
|
+
self.att_layers = []
|
|
274
|
+
for il in range(num_transformer_layers):
|
|
275
|
+
if il == 0:
|
|
276
|
+
node_dim = None
|
|
277
|
+
self.out_dim = hidden_channels * 2 if num_transformer_layers > 1 else hidden_channels
|
|
278
|
+
else:
|
|
279
|
+
self.out_dim = node_dim = hidden_channels
|
|
280
|
+
self.att_layers.append(LPAttLayer(hidden_channels, self.out_dim, node_dim, num_heads, transformer_dropout))
|
|
281
|
+
self.elementwise_lin = MLP(hidden_channels, hidden_channels, hidden_channels)
|
|
282
|
+
self.ppr_encoder_cn = MLP(2, hidden_channels, hidden_channels)
|
|
283
|
+
self.ppr_encoder_onehop = MLP(2, hidden_channels, hidden_channels)
|
|
284
|
+
self.ppr_encoder_non1hop = MLP(2, hidden_channels, hidden_channels)
|
|
285
|
+
if self.thresh_non1hop == 1 and self.thresh_1hop == 1:
|
|
286
|
+
self.mask = "cn"
|
|
287
|
+
elif self.thresh_non1hop == 1 and self.thresh_1hop < 1:
|
|
288
|
+
self.mask = "1-hop"
|
|
289
|
+
else:
|
|
290
|
+
self.mask = "all"
|
|
291
|
+
pairwise_dim = hidden_channels * num_heads + 4
|
|
292
|
+
self.pairwise_lin = MLP(pairwise_dim, pairwise_dim, hidden_channels)
|
|
293
|
+
self.score_func = MLP(hidden_channels * 2, hidden_channels * 2, 1, norm=None)
|
|
294
|
+
|
|
295
|
+
@staticmethod
|
|
296
|
+
def calc_sparse_ppr(edge_index, num_nodes: int, alpha: float = 0.15, eps: float = 5e-5):
|
|
297
|
+
r"""The personalized PageRank matrix LPFormer needs (see :func:`get_ppr`)."""
|
|
298
|
+
return get_ppr(edge_index, num_nodes, alpha=alpha, eps=eps)
|
|
299
|
+
|
|
300
|
+
# ---- node selection on the host ----------------------------------------------------------
|
|
301
|
+
def _drop(self, rng, n):
|
|
302
|
+
keep = math.ceil(n * (1 - self.trans_drop))
|
|
303
|
+
return rng.permutation(n)[:keep]
|
|
304
|
+
|
|
305
|
+
def compute_node_mask(self, u, v, adj, ppr, training):
|
|
306
|
+
r"""For the links ``(u, v)``: ``(pair index, node, PPR from u, PPR from v)`` of their common
|
|
307
|
+
neighbors, one-hop neighbors and other high-PPR nodes."""
|
|
308
|
+
pair_adj = (adj[u] * adj[v]) if self.mask == "cn" else (adj[u] + adj[v])
|
|
309
|
+
pair_adj = pair_adj.tocoo()
|
|
310
|
+
order = np.lexsort((pair_adj.col, pair_adj.row))
|
|
311
|
+
rows, cols, node_type = pair_adj.row[order], pair_adj.col[order], pair_adj.data[order]
|
|
312
|
+
src_ppr = _lookup(ppr, u[rows], cols)
|
|
313
|
+
tgt_ppr = _lookup(ppr, v[rows], cols)
|
|
314
|
+
cn_cond = (src_ppr >= self.thresh_cn) & (tgt_ppr >= self.thresh_cn)
|
|
315
|
+
onehop_cond = (src_ppr >= self.thresh_1hop) & (tgt_ppr >= self.thresh_1hop)
|
|
316
|
+
keep = np.where(node_type == 1, onehop_cond, cn_cond) if self.mask != "cn" else np.where(node_type == 0, onehop_cond, cn_cond)
|
|
317
|
+
rows, cols, node_type, src_ppr, tgt_ppr = rows[keep], cols[keep], node_type[keep], src_ppr[keep], tgt_ppr[keep]
|
|
318
|
+
|
|
319
|
+
non1hop = None
|
|
320
|
+
if self.mask == "all":
|
|
321
|
+
non1hop = self._non_1hop(u, v, adj, ppr, training)
|
|
322
|
+
rng = np.random
|
|
323
|
+
if training and self.trans_drop > 0:
|
|
324
|
+
idx = self._drop(rng, len(rows))
|
|
325
|
+
rows, cols, node_type, src_ppr, tgt_ppr = rows[idx], cols[idx], node_type[idx], src_ppr[idx], tgt_ppr[idx]
|
|
326
|
+
if non1hop is not None:
|
|
327
|
+
idx = self._drop(rng, len(non1hop[0]))
|
|
328
|
+
non1hop = tuple(a[idx] for a in non1hop)
|
|
329
|
+
if self.mask == "cn":
|
|
330
|
+
return (rows, cols, src_ppr, tgt_ppr), None, None
|
|
331
|
+
cn = node_type == 2
|
|
332
|
+
one = node_type == 1
|
|
333
|
+
return ((rows[cn], cols[cn], src_ppr[cn], tgt_ppr[cn]), (rows[one], cols[one], src_ppr[one], tgt_ppr[one]),
|
|
334
|
+
non1hop)
|
|
335
|
+
|
|
336
|
+
def _non_1hop(self, u, v, adj, ppr, training):
|
|
337
|
+
import scipy.sparse as sp
|
|
338
|
+
|
|
339
|
+
adj2 = adj
|
|
340
|
+
if training: # the links being predicted are known edges during training
|
|
341
|
+
n = adj.shape[0]
|
|
342
|
+
links = sp.csr_matrix((np.ones(2 * len(u)), (np.concatenate([u, v]), np.concatenate([v, u]))), shape=(n, n))
|
|
343
|
+
adj2 = ((adj + links) > 0).astype(np.float32).tocsr()
|
|
344
|
+
neighbors = ((adj2[u] + adj2[v]) > 0).astype(np.float32)
|
|
345
|
+
src, tgt = ppr[u], ppr[v]
|
|
346
|
+
both = (src >= self.thresh_non1hop).astype(np.float32).multiply((tgt >= self.thresh_non1hop).astype(np.float32))
|
|
347
|
+
both = sp.csr_matrix(both - both.multiply(neighbors)) # high PPR from both ends, not a neighbor
|
|
348
|
+
both.eliminate_zeros()
|
|
349
|
+
both = both.tocoo()
|
|
350
|
+
order = np.lexsort((both.col, both.row))
|
|
351
|
+
rows, cols = both.row[order], both.col[order]
|
|
352
|
+
return rows, cols, _lookup(ppr, u[rows], cols), _lookup(ppr, v[rows], cols)
|
|
353
|
+
|
|
354
|
+
# ---- forward ------------------------------------------------------------------------------------
|
|
355
|
+
def _pos_encoding(self, encoder, s, t, training):
|
|
356
|
+
a = ops.convert_to_tensor(np.stack([s, t], axis=1).astype(np.float32))
|
|
357
|
+
b = ops.convert_to_tensor(np.stack([t, s], axis=1).astype(np.float32))
|
|
358
|
+
return encoder(a, training=training) + encoder(b, training=training)
|
|
359
|
+
|
|
360
|
+
def call(self, batch, x, edge_index, ppr_matrix=None, training=False):
|
|
361
|
+
import scipy.sparse as sp
|
|
362
|
+
|
|
363
|
+
from k3_node.ops.host import to_numpy
|
|
364
|
+
|
|
365
|
+
batch_np = np.asarray(to_numpy(batch)).astype(np.int64)
|
|
366
|
+
edge_np = np.asarray(to_numpy(edge_index)).astype(np.int64)
|
|
367
|
+
num_nodes = x.shape[0]
|
|
368
|
+
if ppr_matrix is None:
|
|
369
|
+
ppr_matrix = self.calc_sparse_ppr(edge_np, num_nodes)
|
|
370
|
+
ppr = sp.csr_matrix(ppr_matrix)
|
|
371
|
+
adj = sp.csr_matrix((np.ones(edge_np.shape[1], np.float32), (edge_np[0], edge_np[1])), shape=(num_nodes, num_nodes))
|
|
372
|
+
adj.data[:] = 1.0 # {0, 1} even with duplicate edges
|
|
373
|
+
|
|
374
|
+
X_node = self.gnn_norm(self.gnn(self.input_dropout(x, training=training), edge_index, training=training))
|
|
375
|
+
u, v = batch_np[0], batch_np[1]
|
|
376
|
+
x_i, x_j = ops.take(X_node, u, axis=0), ops.take(X_node, v, axis=0)
|
|
377
|
+
elementwise = self.elementwise_lin(x_i * x_j, training=training)
|
|
378
|
+
|
|
379
|
+
cn, onehop, non1hop = self.compute_node_mask(u, v, adj, ppr, training)
|
|
380
|
+
groups = [(cn, self.ppr_encoder_cn), (onehop, self.ppr_encoder_onehop), (non1hop, self.ppr_encoder_non1hop)]
|
|
381
|
+
groups = [(g, enc) for g, enc in groups if g is not None]
|
|
382
|
+
rows = np.concatenate([g[0] for g, _ in groups])
|
|
383
|
+
cols = np.concatenate([g[1] for g, _ in groups])
|
|
384
|
+
pes = ops.concatenate([self._pos_encoding(enc, g[2], g[3], training) for g, enc in groups], axis=0)
|
|
385
|
+
all_mask = ops.convert_to_tensor(np.stack([rows, cols]).astype(np.int32))
|
|
386
|
+
|
|
387
|
+
pairwise = ops.concatenate([x_i, x_j], axis=-1)
|
|
388
|
+
for layer in self.att_layers:
|
|
389
|
+
pairwise = layer(all_mask, pairwise, X_node, pes, training=training)
|
|
390
|
+
|
|
391
|
+
B = len(u)
|
|
392
|
+
counts = [np.bincount(cn[0], minlength=B)] # common neighbors (all pass thresh_cn)
|
|
393
|
+
if onehop is not None:
|
|
394
|
+
num_1hop = np.bincount(onehop[0][(onehop[2] >= self.thresh_1hop) & (onehop[3] >= self.thresh_1hop)], minlength=B)
|
|
395
|
+
num_ppr_ones = np.bincount(onehop[0], minlength=B)
|
|
396
|
+
counts += [num_1hop]
|
|
397
|
+
else:
|
|
398
|
+
num_ppr_ones = np.zeros(B)
|
|
399
|
+
counts += [np.zeros(B)]
|
|
400
|
+
counts += [np.bincount(non1hop[0], minlength=B) if non1hop is not None else np.zeros(B), counts[0] + num_ppr_ones]
|
|
401
|
+
counts = ops.convert_to_tensor(np.stack(counts, axis=1).astype(np.float32))
|
|
402
|
+
|
|
403
|
+
pairwise = self.pairwise_lin(ops.concatenate([pairwise, counts], axis=-1), training=training)
|
|
404
|
+
return self.score_func(ops.concatenate([elementwise, pairwise], axis=-1), training=training)
|
|
@@ -0,0 +1,114 @@
|
|
|
1
|
+
from typing import Optional
|
|
2
|
+
|
|
3
|
+
import keras
|
|
4
|
+
from keras import ops
|
|
5
|
+
import numpy as np
|
|
6
|
+
|
|
7
|
+
|
|
8
|
+
class MaskLabel(keras.layers.Layer):
|
|
9
|
+
r"""The label embedding and masking layer from the `"Masked Label
|
|
10
|
+
Prediction: Unified Message Passing Model for Semi-Supervised
|
|
11
|
+
Classification" <https://arxiv.org/abs/2009.03509>`_ paper.
|
|
12
|
+
|
|
13
|
+
Here, node labels :obj:`y` are merged to the initial node features :obj:`x`
|
|
14
|
+
for a subset of their nodes according to :obj:`mask`.
|
|
15
|
+
|
|
16
|
+
Args:
|
|
17
|
+
num_classes (int): The number of classes.
|
|
18
|
+
out_channels (int): Size of each output sample.
|
|
19
|
+
method (str, optional): If set to :obj:`"add"`, label embeddings are
|
|
20
|
+
added to the input. If set to :obj:`"concat"`, label embeddings are
|
|
21
|
+
concatenated. In case :obj:`method="add"`, then :obj:`out_channels`
|
|
22
|
+
needs to be identical to the input dimensionality of node features.
|
|
23
|
+
(default: :obj:`"add"`)
|
|
24
|
+
|
|
25
|
+
Example:
|
|
26
|
+
```python
|
|
27
|
+
import numpy as np
|
|
28
|
+
from k3_node.models import MaskLabel
|
|
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
|
+
|
|
33
|
+
y = np.random.randint(0, 3, size=(10,)) # node labels
|
|
34
|
+
known = np.random.rand(10) > 0.5 # labels visible to the model
|
|
35
|
+
model = MaskLabel(num_classes=3, out_channels=8) # label embedding added to x
|
|
36
|
+
out = model(x, y, known)
|
|
37
|
+
print(tuple(out.shape)) # (10, 8)
|
|
38
|
+
```
|
|
39
|
+
"""
|
|
40
|
+
|
|
41
|
+
def __init__(
|
|
42
|
+
self,
|
|
43
|
+
num_classes: int,
|
|
44
|
+
out_channels: int,
|
|
45
|
+
method: str = "add",
|
|
46
|
+
**kwargs,
|
|
47
|
+
):
|
|
48
|
+
super().__init__(**kwargs)
|
|
49
|
+
|
|
50
|
+
self.num_classes = num_classes
|
|
51
|
+
self.out_channels = out_channels
|
|
52
|
+
self.method = method
|
|
53
|
+
|
|
54
|
+
if method not in ["add", "concat"]:
|
|
55
|
+
raise ValueError(
|
|
56
|
+
f"'method' must be either 'add' or 'concat' (got '{method}')"
|
|
57
|
+
)
|
|
58
|
+
|
|
59
|
+
self.emb = keras.layers.Embedding(num_classes, out_channels)
|
|
60
|
+
|
|
61
|
+
def build(self, input_shape=None):
|
|
62
|
+
self.emb.build((None,))
|
|
63
|
+
self.built = True
|
|
64
|
+
|
|
65
|
+
def reset_parameters(self) -> None:
|
|
66
|
+
r"""Resets all learnable parameters of the module."""
|
|
67
|
+
if self.emb.built:
|
|
68
|
+
self.emb.embeddings.assign(
|
|
69
|
+
keras.initializers.Uniform()(self.emb.embeddings.shape)
|
|
70
|
+
)
|
|
71
|
+
|
|
72
|
+
def call(self, x, y, mask):
|
|
73
|
+
"""Forward pass.
|
|
74
|
+
|
|
75
|
+
Args:
|
|
76
|
+
x (Tensor): The input node features of shape ``[N, in_channels]``.
|
|
77
|
+
y (Tensor): The node labels of shape ``[N]`` (integer class indices).
|
|
78
|
+
mask (Tensor): Boolean tensor of shape ``[N]`` indicating which
|
|
79
|
+
nodes have ground-truth labels to embed.
|
|
80
|
+
"""
|
|
81
|
+
# Embed all labels; then zero-out non-masked entries
|
|
82
|
+
all_emb = self.emb(y) # [N, out_channels]
|
|
83
|
+
|
|
84
|
+
# Build a float mask: 1.0 where mask is True, 0.0 elsewhere
|
|
85
|
+
float_mask = ops.cast(mask, dtype=all_emb.dtype) # [N]
|
|
86
|
+
float_mask = ops.expand_dims(float_mask, axis=-1) # [N, 1]
|
|
87
|
+
masked_emb = all_emb * float_mask # [N, out_channels]
|
|
88
|
+
|
|
89
|
+
if self.method == "concat":
|
|
90
|
+
return ops.concatenate([x, masked_emb], axis=-1)
|
|
91
|
+
else:
|
|
92
|
+
return x + masked_emb
|
|
93
|
+
|
|
94
|
+
@staticmethod
|
|
95
|
+
def ratio_mask(mask, ratio: float):
|
|
96
|
+
r"""Modifies :obj:`mask` by setting :obj:`ratio` of :obj:`True`
|
|
97
|
+
entries to :obj:`False`. Does not operate in-place.
|
|
98
|
+
|
|
99
|
+
Args:
|
|
100
|
+
mask (Tensor): The boolean mask to re-mask.
|
|
101
|
+
ratio (float): The ratio of True entries to keep.
|
|
102
|
+
"""
|
|
103
|
+
mask_np = ops.convert_to_numpy(mask).astype(bool)
|
|
104
|
+
n = int(mask_np.sum())
|
|
105
|
+
out_np = mask_np.copy()
|
|
106
|
+
if n > 0:
|
|
107
|
+
keep = np.random.rand(n) < ratio
|
|
108
|
+
true_indices = np.where(mask_np)[0]
|
|
109
|
+
out_np[true_indices] = keep
|
|
110
|
+
return ops.convert_to_tensor(out_np, dtype="bool")
|
|
111
|
+
|
|
112
|
+
def __repr__(self) -> str:
|
|
113
|
+
return f'{self.__class__.__name__}()'
|
|
114
|
+
|
|
@@ -0,0 +1,33 @@
|
|
|
1
|
+
"""Materials and crystal models (aliased from k3_node.applications.materials)."""
|
|
2
|
+
|
|
3
|
+
import sys
|
|
4
|
+
from k3_node.applications.materials import *
|
|
5
|
+
from k3_node.applications.materials import (
|
|
6
|
+
basis,
|
|
7
|
+
core,
|
|
8
|
+
readout,
|
|
9
|
+
wrappers,
|
|
10
|
+
io,
|
|
11
|
+
megnet,
|
|
12
|
+
m3gnet,
|
|
13
|
+
tensornet,
|
|
14
|
+
chgnet,
|
|
15
|
+
so3net,
|
|
16
|
+
grace,
|
|
17
|
+
qet,
|
|
18
|
+
)
|
|
19
|
+
from k3_node.applications.materials import __all__
|
|
20
|
+
|
|
21
|
+
# Alias submodules in sys.modules for full backward compatibility
|
|
22
|
+
sys.modules[__name__ + ".basis"] = basis
|
|
23
|
+
sys.modules[__name__ + ".core"] = core
|
|
24
|
+
sys.modules[__name__ + ".readout"] = readout
|
|
25
|
+
sys.modules[__name__ + ".wrappers"] = wrappers
|
|
26
|
+
sys.modules[__name__ + ".io"] = io
|
|
27
|
+
sys.modules[__name__ + ".megnet"] = megnet
|
|
28
|
+
sys.modules[__name__ + ".m3gnet"] = m3gnet
|
|
29
|
+
sys.modules[__name__ + ".tensornet"] = tensornet
|
|
30
|
+
sys.modules[__name__ + ".chgnet"] = chgnet
|
|
31
|
+
sys.modules[__name__ + ".so3net"] = so3net
|
|
32
|
+
sys.modules[__name__ + ".grace"] = grace
|
|
33
|
+
sys.modules[__name__ + ".qet"] = qet
|
k3_node/models/meta.py
ADDED
|
@@ -0,0 +1,133 @@
|
|
|
1
|
+
from typing import Optional, Tuple
|
|
2
|
+
|
|
3
|
+
import keras
|
|
4
|
+
from keras import ops
|
|
5
|
+
|
|
6
|
+
|
|
7
|
+
class MetaLayer(keras.layers.Layer):
|
|
8
|
+
r"""A meta layer for building any kind of graph network, inspired by the
|
|
9
|
+
`"Relational Inductive Biases, Deep Learning, and Graph Networks"
|
|
10
|
+
<https://arxiv.org/abs/1806.01261>`_ paper.
|
|
11
|
+
|
|
12
|
+
A graph network takes a graph as input and returns an updated graph as
|
|
13
|
+
output (with same connectivity). The input graph has node features :obj:`x`,
|
|
14
|
+
edge features :obj:`edge_attr` as well as graph-level features :obj:`u`.
|
|
15
|
+
The output graph has the same structure, but updated features.
|
|
16
|
+
|
|
17
|
+
Edge features, node features as well as global features are updated by
|
|
18
|
+
calling the modules :obj:`edge_model`, :obj:`node_model` and
|
|
19
|
+
:obj:`global_model`, respectively.
|
|
20
|
+
|
|
21
|
+
To allow for batch-wise graph processing, all callable functions take an
|
|
22
|
+
additional argument :obj:`batch`, which determines the assignment of
|
|
23
|
+
edges or nodes to their specific graphs.
|
|
24
|
+
|
|
25
|
+
Args:
|
|
26
|
+
edge_model (callable, optional): A callable which updates a graph's
|
|
27
|
+
edge features based on its source and target node features, its
|
|
28
|
+
current edge features and its global features.
|
|
29
|
+
(default: :obj:`None`)
|
|
30
|
+
node_model (callable, optional): A callable which updates a graph's
|
|
31
|
+
node features based on its current node features, its graph
|
|
32
|
+
connectivity, its edge features and its global features.
|
|
33
|
+
(default: :obj:`None`)
|
|
34
|
+
global_model (callable, optional): A callable which updates a graph's
|
|
35
|
+
global features based on its node features, its graph connectivity,
|
|
36
|
+
its edge features and its current global features.
|
|
37
|
+
(default: :obj:`None`)
|
|
38
|
+
|
|
39
|
+
Example:
|
|
40
|
+
```python
|
|
41
|
+
import numpy as np
|
|
42
|
+
import keras
|
|
43
|
+
from k3_node.models import MetaLayer
|
|
44
|
+
|
|
45
|
+
x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
|
|
46
|
+
edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
|
|
47
|
+
edge_attr = np.random.rand(30, 4).astype("float32")
|
|
48
|
+
|
|
49
|
+
edge_mlp = keras.layers.Dense(8)
|
|
50
|
+
node_mlp = keras.layers.Dense(8)
|
|
51
|
+
|
|
52
|
+
def edge_model(src, dst, edge_attr, u, batch): # update every edge from its endpoints
|
|
53
|
+
return edge_mlp(keras.ops.concatenate([src, dst, edge_attr], axis=-1))
|
|
54
|
+
|
|
55
|
+
def node_model(x, edge_index, edge_attr, u, batch): # update nodes from incoming edges
|
|
56
|
+
incoming = keras.ops.segment_sum(edge_attr, edge_index[1], num_segments=10)
|
|
57
|
+
return node_mlp(keras.ops.concatenate([x, incoming], axis=-1))
|
|
58
|
+
|
|
59
|
+
model = MetaLayer(edge_model=edge_model, node_model=node_model)
|
|
60
|
+
x_out, edge_attr_out, u_out = model(x, edge_index, edge_attr=edge_attr)
|
|
61
|
+
print(tuple(x_out.shape), tuple(edge_attr_out.shape)) # (10, 8) (30, 8)
|
|
62
|
+
```
|
|
63
|
+
"""
|
|
64
|
+
|
|
65
|
+
def __init__(
|
|
66
|
+
self,
|
|
67
|
+
edge_model=None,
|
|
68
|
+
node_model=None,
|
|
69
|
+
global_model=None,
|
|
70
|
+
**kwargs,
|
|
71
|
+
):
|
|
72
|
+
super().__init__(**kwargs)
|
|
73
|
+
self.edge_model = edge_model
|
|
74
|
+
self.node_model = node_model
|
|
75
|
+
self.global_model = global_model
|
|
76
|
+
|
|
77
|
+
self.reset_parameters()
|
|
78
|
+
|
|
79
|
+
def reset_parameters(self) -> None:
|
|
80
|
+
r"""Resets all learnable parameters of the module."""
|
|
81
|
+
for item in [self.node_model, self.edge_model, self.global_model]:
|
|
82
|
+
if hasattr(item, 'reset_parameters'):
|
|
83
|
+
item.reset_parameters()
|
|
84
|
+
|
|
85
|
+
def call(
|
|
86
|
+
self,
|
|
87
|
+
x,
|
|
88
|
+
edge_index,
|
|
89
|
+
edge_attr=None,
|
|
90
|
+
u=None,
|
|
91
|
+
batch=None,
|
|
92
|
+
) -> Tuple:
|
|
93
|
+
r"""Forward pass.
|
|
94
|
+
|
|
95
|
+
Args:
|
|
96
|
+
x (Tensor): The node features of shape ``[N, F_x]``.
|
|
97
|
+
edge_index (Tensor): The edge indices of shape ``[2, E]``.
|
|
98
|
+
edge_attr (Tensor, optional): The edge features of shape
|
|
99
|
+
``[E, F_e]``. (default: :obj:`None`)
|
|
100
|
+
u (Tensor, optional): The global graph features of shape
|
|
101
|
+
``[B, F_u]``. (default: :obj:`None`)
|
|
102
|
+
batch (Tensor, optional): The batch vector
|
|
103
|
+
:math:`\mathbf{b} \in {\{ 0, \ldots, B-1\}}^N`.
|
|
104
|
+
(default: :obj:`None`)
|
|
105
|
+
"""
|
|
106
|
+
row = edge_index[0]
|
|
107
|
+
col = edge_index[1]
|
|
108
|
+
|
|
109
|
+
if self.edge_model is not None:
|
|
110
|
+
edge_batch = batch if batch is None else ops.take(batch, row, axis=0)
|
|
111
|
+
edge_attr = self.edge_model(
|
|
112
|
+
ops.take(x, row, axis=0),
|
|
113
|
+
ops.take(x, col, axis=0),
|
|
114
|
+
edge_attr,
|
|
115
|
+
u,
|
|
116
|
+
edge_batch,
|
|
117
|
+
)
|
|
118
|
+
|
|
119
|
+
if self.node_model is not None:
|
|
120
|
+
x = self.node_model(x, edge_index, edge_attr, u, batch)
|
|
121
|
+
|
|
122
|
+
if self.global_model is not None:
|
|
123
|
+
u = self.global_model(x, edge_index, edge_attr, u, batch)
|
|
124
|
+
|
|
125
|
+
return x, edge_attr, u
|
|
126
|
+
|
|
127
|
+
def __repr__(self) -> str:
|
|
128
|
+
return (f'{self.__class__.__name__}(\n'
|
|
129
|
+
f' edge_model={self.edge_model},\n'
|
|
130
|
+
f' node_model={self.node_model},\n'
|
|
131
|
+
f' global_model={self.global_model}\n'
|
|
132
|
+
f')')
|
|
133
|
+
|