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,168 @@
|
|
|
1
|
+
import math
|
|
2
|
+
from typing import Optional, Union, Tuple
|
|
3
|
+
from keras import layers, ops
|
|
4
|
+
|
|
5
|
+
from k3_node.layers.conv.message_passing import MessagePassing
|
|
6
|
+
from k3_node.layers.conv.utils import softmax
|
|
7
|
+
from k3_node.ops.segment import segment_sum
|
|
8
|
+
|
|
9
|
+
|
|
10
|
+
class TransformerConv(MessagePassing):
|
|
11
|
+
r"""The graph transformer operator from the `"Masked Label Prediction:
|
|
12
|
+
Unified Meta-Learning on Graph Neural Networks"
|
|
13
|
+
<https://arxiv.org/abs/2009.03509>`_ paper.
|
|
14
|
+
|
|
15
|
+
Args:
|
|
16
|
+
in_channels: Size of each input sample, or a tuple for bipartite graphs.
|
|
17
|
+
out_channels: Size of each output sample.
|
|
18
|
+
heads: Number of multi-head-attentions. (default: ``1``)
|
|
19
|
+
concat: If set to :obj:`False`, the multi-head-attentions are averaged
|
|
20
|
+
instead of concatenated. (default: ``True``)
|
|
21
|
+
beta: If set to :obj:`True`, will use a gated residual connection.
|
|
22
|
+
(default: ``False``)
|
|
23
|
+
dropout: Dropout probability of the normalized attention coefficients.
|
|
24
|
+
(default: ``0.0``)
|
|
25
|
+
edge_dim: Edge feature dimensionality (in case there are any).
|
|
26
|
+
(default: :obj:`None`)
|
|
27
|
+
bias: If set to :obj:`False`, the layer will not learn an additive bias.
|
|
28
|
+
(default: ``True``)
|
|
29
|
+
root_weight: If set to :obj:`False`, the layer will not add the
|
|
30
|
+
transformed root node features. (default: ``True``)
|
|
31
|
+
|
|
32
|
+
Example:
|
|
33
|
+
```python
|
|
34
|
+
import numpy as np
|
|
35
|
+
from k3_node.layers import TransformerConv
|
|
36
|
+
|
|
37
|
+
x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
|
|
38
|
+
edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
|
|
39
|
+
|
|
40
|
+
layer = TransformerConv(in_channels=8, out_channels=16, heads=2)
|
|
41
|
+
out = layer(x, edge_index)
|
|
42
|
+
print(tuple(out.shape)) # (10, 32)
|
|
43
|
+
```
|
|
44
|
+
"""
|
|
45
|
+
|
|
46
|
+
def __init__(
|
|
47
|
+
self,
|
|
48
|
+
in_channels: Union[int, Tuple[int, int]],
|
|
49
|
+
out_channels: int,
|
|
50
|
+
heads: int = 1,
|
|
51
|
+
concat: bool = True,
|
|
52
|
+
beta: bool = False,
|
|
53
|
+
dropout: float = 0.0,
|
|
54
|
+
edge_dim: Optional[int] = None,
|
|
55
|
+
bias: bool = True,
|
|
56
|
+
root_weight: bool = True,
|
|
57
|
+
**kwargs,
|
|
58
|
+
):
|
|
59
|
+
super().__init__(node_dim=0, aggr="add", **kwargs)
|
|
60
|
+
self.in_channels = in_channels
|
|
61
|
+
self.out_channels = out_channels
|
|
62
|
+
self.heads = heads
|
|
63
|
+
self.concat = concat
|
|
64
|
+
self.beta = beta and root_weight
|
|
65
|
+
self.root_weight = root_weight
|
|
66
|
+
self.dropout_rate = dropout
|
|
67
|
+
self.edge_dim = edge_dim
|
|
68
|
+
self.use_bias = bias
|
|
69
|
+
|
|
70
|
+
total_out_channels = out_channels * (heads if concat else 1)
|
|
71
|
+
|
|
72
|
+
self.lin_key = layers.Dense(heads * out_channels, use_bias=bias)
|
|
73
|
+
self.lin_query = layers.Dense(heads * out_channels, use_bias=bias)
|
|
74
|
+
self.lin_value = layers.Dense(heads * out_channels, use_bias=bias)
|
|
75
|
+
|
|
76
|
+
if edge_dim is not None:
|
|
77
|
+
self.lin_edge = layers.Dense(heads * out_channels, use_bias=False)
|
|
78
|
+
else:
|
|
79
|
+
self.lin_edge = None
|
|
80
|
+
|
|
81
|
+
if root_weight:
|
|
82
|
+
self.lin_skip = layers.Dense(total_out_channels, use_bias=bias)
|
|
83
|
+
if self.beta:
|
|
84
|
+
self.lin_beta = layers.Dense(1, use_bias=False)
|
|
85
|
+
else:
|
|
86
|
+
self.lin_beta = None
|
|
87
|
+
else:
|
|
88
|
+
self.lin_skip = None
|
|
89
|
+
self.lin_beta = None
|
|
90
|
+
|
|
91
|
+
self.dropout = layers.Dropout(dropout) if dropout > 0.0 else None
|
|
92
|
+
|
|
93
|
+
def build(self, input_shape):
|
|
94
|
+
if isinstance(input_shape, (tuple, list)) and len(input_shape) > 0 and isinstance(input_shape[0], (tuple, list)):
|
|
95
|
+
in_channels_src = input_shape[0][-1]
|
|
96
|
+
in_channels_dst = input_shape[1][-1] if len(input_shape) > 1 and input_shape[1] is not None else in_channels_src
|
|
97
|
+
else:
|
|
98
|
+
in_channels_src = input_shape[-1]
|
|
99
|
+
in_channels_dst = input_shape[-1]
|
|
100
|
+
|
|
101
|
+
self.lin_key.build((None, in_channels_src))
|
|
102
|
+
self.lin_query.build((None, in_channels_dst))
|
|
103
|
+
self.lin_value.build((None, in_channels_src))
|
|
104
|
+
|
|
105
|
+
if self.lin_edge is not None:
|
|
106
|
+
self.lin_edge.build((None, self.edge_dim))
|
|
107
|
+
if self.lin_skip is not None:
|
|
108
|
+
self.lin_skip.build((None, in_channels_dst))
|
|
109
|
+
if self.lin_beta is not None:
|
|
110
|
+
total_out = self.out_channels * (self.heads if self.concat else 1)
|
|
111
|
+
self.lin_beta.build((None, 3 * total_out))
|
|
112
|
+
|
|
113
|
+
self.built = True
|
|
114
|
+
|
|
115
|
+
def call(self, x, edge_index=None, edge_attr=None, return_attention_weights=None, training=None, **kwargs):
|
|
116
|
+
if edge_index is None and isinstance(x, (tuple, list)):
|
|
117
|
+
x, edge_index = x[0], x[1]
|
|
118
|
+
|
|
119
|
+
H, C = self.heads, self.out_channels
|
|
120
|
+
if isinstance(x, (tuple, list)):
|
|
121
|
+
x_src, x_dst = x[0], x[1]
|
|
122
|
+
else:
|
|
123
|
+
x_src, x_dst = x, x
|
|
124
|
+
|
|
125
|
+
query = ops.reshape(self.lin_query(x_dst), (-1, H, C))
|
|
126
|
+
key = ops.reshape(self.lin_key(x_src), (-1, H, C))
|
|
127
|
+
value = ops.reshape(self.lin_value(x_src), (-1, H, C))
|
|
128
|
+
|
|
129
|
+
row, col = edge_index[0], edge_index[1]
|
|
130
|
+
row, col = ops.cast(row, "int32"), ops.cast(col, "int32")
|
|
131
|
+
|
|
132
|
+
query_i = ops.take(query, col, axis=0)
|
|
133
|
+
key_j = ops.take(key, row, axis=0)
|
|
134
|
+
value_j = ops.take(value, row, axis=0)
|
|
135
|
+
|
|
136
|
+
if self.lin_edge is not None and edge_attr is not None:
|
|
137
|
+
edge_attr_proj = ops.reshape(self.lin_edge(edge_attr), (-1, H, C))
|
|
138
|
+
key_j = key_j + edge_attr_proj
|
|
139
|
+
value_j = value_j + edge_attr_proj
|
|
140
|
+
|
|
141
|
+
alpha = ops.sum(query_i * key_j, axis=-1) / math.sqrt(C)
|
|
142
|
+
num_nodes_dst = ops.shape(x_dst)[0]
|
|
143
|
+
alpha = softmax(alpha, col, num_nodes=num_nodes_dst, dim=0)
|
|
144
|
+
|
|
145
|
+
if self.dropout is not None:
|
|
146
|
+
alpha = self.dropout(alpha, training=training)
|
|
147
|
+
|
|
148
|
+
out = ops.expand_dims(alpha, -1) * value_j
|
|
149
|
+
out = segment_sum(out, col, num_segments=num_nodes_dst)
|
|
150
|
+
|
|
151
|
+
if self.concat:
|
|
152
|
+
out = ops.reshape(out, (-1, H * C))
|
|
153
|
+
else:
|
|
154
|
+
out = ops.mean(out, axis=1)
|
|
155
|
+
|
|
156
|
+
if self.root_weight and self.lin_skip is not None:
|
|
157
|
+
x_r = self.lin_skip(x_dst)
|
|
158
|
+
if self.lin_beta is not None:
|
|
159
|
+
b_input = ops.concatenate([out, x_r, out - x_r], axis=-1)
|
|
160
|
+
b = ops.sigmoid(self.lin_beta(b_input))
|
|
161
|
+
out = b * x_r + (1.0 - b) * out
|
|
162
|
+
else:
|
|
163
|
+
out = out + x_r
|
|
164
|
+
|
|
165
|
+
if return_attention_weights:
|
|
166
|
+
return out, (edge_index, alpha)
|
|
167
|
+
return out
|
|
168
|
+
|
|
@@ -0,0 +1,403 @@
|
|
|
1
|
+
from typing import Any, Optional, Tuple, Union
|
|
2
|
+
from keras import ops
|
|
3
|
+
from k3_node.ops.segment import segment_max, segment_min, segment_sum
|
|
4
|
+
from k3_node.ops.creation import full
|
|
5
|
+
|
|
6
|
+
|
|
7
|
+
def is_tracing(x: Any) -> bool:
|
|
8
|
+
if x is None:
|
|
9
|
+
return False
|
|
10
|
+
try:
|
|
11
|
+
from keras.src.backend.common.symbolic_scope import in_symbolic_scope
|
|
12
|
+
|
|
13
|
+
if in_symbolic_scope():
|
|
14
|
+
return True
|
|
15
|
+
except Exception:
|
|
16
|
+
pass
|
|
17
|
+
name = type(x).__name__
|
|
18
|
+
if "Tracer" in name or "KerasTensor" in name or "SymbolicTensor" in name:
|
|
19
|
+
return True
|
|
20
|
+
if hasattr(x, "_trace"):
|
|
21
|
+
return True
|
|
22
|
+
try:
|
|
23
|
+
import jax
|
|
24
|
+
if isinstance(x, jax.core.Tracer):
|
|
25
|
+
return True
|
|
26
|
+
except Exception:
|
|
27
|
+
pass
|
|
28
|
+
try:
|
|
29
|
+
import tensorflow as tf
|
|
30
|
+
if hasattr(x, "graph") and getattr(x, "graph", None) is not None:
|
|
31
|
+
if not tf.executing_eagerly():
|
|
32
|
+
return True
|
|
33
|
+
except Exception:
|
|
34
|
+
pass
|
|
35
|
+
return False
|
|
36
|
+
|
|
37
|
+
|
|
38
|
+
def _is_compiled_trace(x) -> bool:
|
|
39
|
+
"""True inside compiled functions, where tensor values are unknown.
|
|
40
|
+
|
|
41
|
+
On JAX, autodiff tracers created by ``jax.grad`` in eager mode (``run_eagerly=True``)
|
|
42
|
+
still carry concrete values, so they do not count as compiled.
|
|
43
|
+
"""
|
|
44
|
+
import keras
|
|
45
|
+
|
|
46
|
+
if keras.config.backend() != "jax":
|
|
47
|
+
return is_tracing(x)
|
|
48
|
+
import jax
|
|
49
|
+
import jax.numpy as jnp
|
|
50
|
+
|
|
51
|
+
if not isinstance(x, jax.core.Tracer):
|
|
52
|
+
return False
|
|
53
|
+
try:
|
|
54
|
+
int(jnp.sum(jnp.ravel(x)[:1] * 0)) # concretizing fails only inside jit
|
|
55
|
+
return False
|
|
56
|
+
except (jax.errors.ConcretizationTypeError, jax.errors.TracerIntegerConversionError):
|
|
57
|
+
return True
|
|
58
|
+
|
|
59
|
+
|
|
60
|
+
def host_callback(fn, out_specs, *args):
|
|
61
|
+
"""Runs the NumPy function ``fn`` on the concrete values of ``args``.
|
|
62
|
+
|
|
63
|
+
``out_specs`` is a sequence of ``(shape, dtype)`` giving the fixed shapes of ``fn``'s outputs.
|
|
64
|
+
No gradient flows through the outputs. On JAX this uses ``jax.pure_callback``, which also
|
|
65
|
+
works under ``jax.grad``; elsewhere the arguments are converted to NumPy directly.
|
|
66
|
+
"""
|
|
67
|
+
import keras
|
|
68
|
+
import numpy as np
|
|
69
|
+
|
|
70
|
+
def run(*values):
|
|
71
|
+
outs = fn(*[np.asarray(v) for v in values])
|
|
72
|
+
return tuple(np.asarray(o, dtype=dtype).reshape(shape) for o, (shape, dtype) in zip(outs, out_specs))
|
|
73
|
+
|
|
74
|
+
if keras.config.backend() == "jax":
|
|
75
|
+
import jax
|
|
76
|
+
|
|
77
|
+
specs = tuple(jax.ShapeDtypeStruct(shape, dtype) for shape, dtype in out_specs)
|
|
78
|
+
return jax.pure_callback(run, specs, *[jax.lax.stop_gradient(a) for a in args])
|
|
79
|
+
outs = run(*[ops.convert_to_numpy(ops.stop_gradient(a)) for a in args])
|
|
80
|
+
return tuple(ops.convert_to_tensor(o) for o in outs)
|
|
81
|
+
|
|
82
|
+
|
|
83
|
+
def eager_only_placeholder(layer_name: str, *tensors) -> bool:
|
|
84
|
+
"""Guards host-side (NumPy) computations with data-dependent output sizes.
|
|
85
|
+
|
|
86
|
+
Returns ``True`` during Keras shape inference, where the caller should return a
|
|
87
|
+
placeholder result. Raises inside compiled functions (``tf.function`` / XLA /
|
|
88
|
+
``jax.jit``), where the computation cannot run. Returns ``False`` when eager.
|
|
89
|
+
"""
|
|
90
|
+
try:
|
|
91
|
+
from keras.src.backend.common.symbolic_scope import in_symbolic_scope
|
|
92
|
+
|
|
93
|
+
if in_symbolic_scope():
|
|
94
|
+
return True
|
|
95
|
+
except Exception:
|
|
96
|
+
pass
|
|
97
|
+
if any(_is_compiled_trace(t) for t in tensors):
|
|
98
|
+
raise RuntimeError(
|
|
99
|
+
f"{layer_name} computes a data-dependent number of clusters on the host, so it cannot run "
|
|
100
|
+
"inside a compiled function (tf.function, XLA or jax.jit). Compile the model with "
|
|
101
|
+
"`run_eagerly=True`."
|
|
102
|
+
)
|
|
103
|
+
return False
|
|
104
|
+
|
|
105
|
+
|
|
106
|
+
def degree(index, num_nodes: Optional[int] = None, dtype=None):
|
|
107
|
+
"""Computes the (in/out) degree of a given index tensor.
|
|
108
|
+
|
|
109
|
+
Args:
|
|
110
|
+
index: 1D tensor of node indices.
|
|
111
|
+
num_nodes: The number of nodes.
|
|
112
|
+
dtype: Output data type.
|
|
113
|
+
"""
|
|
114
|
+
index = ops.cast(index, "int32")
|
|
115
|
+
if num_nodes is None:
|
|
116
|
+
if is_tracing(index):
|
|
117
|
+
num_nodes = ops.shape(index)[0]
|
|
118
|
+
else:
|
|
119
|
+
num_nodes = int(ops.max(index)) + 1 if ops.shape(index)[0] > 0 else 0
|
|
120
|
+
try:
|
|
121
|
+
num_nodes = int(num_nodes)
|
|
122
|
+
except (TypeError, ValueError):
|
|
123
|
+
pass
|
|
124
|
+
ones = ops.ones((ops.shape(index)[0],), dtype=dtype or "float32")
|
|
125
|
+
deg = segment_sum(ones, index, num_segments=num_nodes)
|
|
126
|
+
if dtype is not None:
|
|
127
|
+
deg = ops.cast(deg, dtype)
|
|
128
|
+
return deg
|
|
129
|
+
|
|
130
|
+
|
|
131
|
+
def remove_self_loops(
|
|
132
|
+
edge_index,
|
|
133
|
+
edge_attr=None,
|
|
134
|
+
) -> Tuple:
|
|
135
|
+
"""Removes self-loops from `edge_index` and optional `edge_attr`."""
|
|
136
|
+
if is_tracing(edge_index):
|
|
137
|
+
return edge_index, edge_attr
|
|
138
|
+
edge_index = ops.convert_to_tensor(edge_index)
|
|
139
|
+
if edge_attr is not None:
|
|
140
|
+
edge_attr = ops.convert_to_tensor(edge_attr)
|
|
141
|
+
mask = edge_index[0] != edge_index[1]
|
|
142
|
+
where_mask = ops.where(mask)
|
|
143
|
+
indices = where_mask[0] if isinstance(where_mask, (list, tuple)) else where_mask
|
|
144
|
+
indices = ops.reshape(indices, (-1,))
|
|
145
|
+
edge_index = ops.take(edge_index, indices, axis=1)
|
|
146
|
+
if edge_attr is not None:
|
|
147
|
+
edge_attr = ops.take(edge_attr, indices, axis=0)
|
|
148
|
+
return edge_index, edge_attr
|
|
149
|
+
|
|
150
|
+
|
|
151
|
+
def remove_self_loops_masked(
|
|
152
|
+
edge_index,
|
|
153
|
+
edge_attr=None,
|
|
154
|
+
) -> Tuple:
|
|
155
|
+
"""Removes self-loops in a way that is also correct under static-shape tracing.
|
|
156
|
+
|
|
157
|
+
Eagerly, self-loops are dropped (as in :func:`remove_self_loops`) and the
|
|
158
|
+
returned mask is ``None``. Under tracing (XLA / ``jax.jit``) the edge count
|
|
159
|
+
must stay static, so all edges are kept and a boolean ``keep_mask`` of shape
|
|
160
|
+
``[E]`` is returned that is ``False`` at the original self-loops. Callers must
|
|
161
|
+
exclude masked edges from aggregation.
|
|
162
|
+
|
|
163
|
+
Returns:
|
|
164
|
+
``(edge_index, edge_attr, keep_mask)``
|
|
165
|
+
"""
|
|
166
|
+
if not is_tracing(edge_index):
|
|
167
|
+
edge_index, edge_attr = remove_self_loops(edge_index, edge_attr)
|
|
168
|
+
return edge_index, edge_attr, None
|
|
169
|
+
edge_index = ops.convert_to_tensor(edge_index)
|
|
170
|
+
return edge_index, edge_attr, ops.not_equal(edge_index[0], edge_index[1])
|
|
171
|
+
|
|
172
|
+
|
|
173
|
+
def extend_mask_for_self_loops(keep_mask, num_nodes):
|
|
174
|
+
"""Extends a ``keep_mask`` to cover the ``num_nodes`` loops appended by :func:`add_self_loops`."""
|
|
175
|
+
if keep_mask is None:
|
|
176
|
+
return None
|
|
177
|
+
return ops.concatenate([keep_mask, ops.ones((num_nodes,), dtype="bool")], axis=0)
|
|
178
|
+
|
|
179
|
+
|
|
180
|
+
def mask_edge_logits(alpha, keep_mask):
|
|
181
|
+
"""Sets the logits of masked edges to ``-inf`` so they receive zero softmax weight."""
|
|
182
|
+
if keep_mask is None:
|
|
183
|
+
return alpha
|
|
184
|
+
mask = ops.reshape(keep_mask, (-1,) + (1,) * (len(alpha.shape) - 1))
|
|
185
|
+
return ops.where(mask, alpha, float("-inf")) # a scalar also works on torch's meta device
|
|
186
|
+
|
|
187
|
+
|
|
188
|
+
def add_self_loops(
|
|
189
|
+
edge_index,
|
|
190
|
+
edge_attr=None,
|
|
191
|
+
fill_value: Union[float, str, None] = None,
|
|
192
|
+
num_nodes: Optional[int] = None,
|
|
193
|
+
) -> Tuple:
|
|
194
|
+
"""Adds self-loops to `edge_index` and optional `edge_attr`."""
|
|
195
|
+
if is_tracing(edge_index) or is_tracing(num_nodes):
|
|
196
|
+
if num_nodes is None:
|
|
197
|
+
return edge_index, edge_attr
|
|
198
|
+
edge_index = ops.convert_to_tensor(edge_index)
|
|
199
|
+
if num_nodes is None:
|
|
200
|
+
num_nodes = int(ops.max(edge_index)) + 1 if ops.shape(edge_index)[1] > 0 else 0
|
|
201
|
+
else:
|
|
202
|
+
try:
|
|
203
|
+
num_nodes = int(num_nodes)
|
|
204
|
+
except (TypeError, ValueError):
|
|
205
|
+
pass
|
|
206
|
+
|
|
207
|
+
loop_index = ops.arange(0, num_nodes, dtype=edge_index.dtype)
|
|
208
|
+
loop_index = ops.stack([loop_index, loop_index], axis=0)
|
|
209
|
+
edge_index = ops.concatenate([edge_index, loop_index], axis=1)
|
|
210
|
+
|
|
211
|
+
if edge_attr is not None:
|
|
212
|
+
edge_attr = ops.convert_to_tensor(edge_attr)
|
|
213
|
+
attr_shape = (num_nodes,) + tuple(edge_attr.shape[1:]) if hasattr(edge_attr, "shape") else (num_nodes,)
|
|
214
|
+
if fill_value is None:
|
|
215
|
+
loop_attr = ops.zeros(attr_shape, dtype=edge_attr.dtype)
|
|
216
|
+
elif isinstance(fill_value, (int, float)):
|
|
217
|
+
loop_attr = full(attr_shape, fill_value, dtype=edge_attr.dtype)
|
|
218
|
+
elif fill_value == "add" or fill_value == "mean":
|
|
219
|
+
loop_attr = ops.zeros(attr_shape, dtype=edge_attr.dtype)
|
|
220
|
+
else:
|
|
221
|
+
loop_attr = full(attr_shape, fill_value, dtype=edge_attr.dtype)
|
|
222
|
+
edge_attr = ops.concatenate([edge_attr, loop_attr], axis=0)
|
|
223
|
+
|
|
224
|
+
return edge_index, edge_attr
|
|
225
|
+
|
|
226
|
+
|
|
227
|
+
def gcn_norm(
|
|
228
|
+
edge_index,
|
|
229
|
+
edge_weight=None,
|
|
230
|
+
num_nodes: Optional[int] = None,
|
|
231
|
+
improved: bool = False,
|
|
232
|
+
add_self_loops: bool = True,
|
|
233
|
+
flow: str = "source_to_target",
|
|
234
|
+
dtype=None,
|
|
235
|
+
) -> Tuple:
|
|
236
|
+
"""Computes the GCN normalization coefficients."""
|
|
237
|
+
fill_value = 2.0 if improved else 1.0
|
|
238
|
+
edge_index = ops.convert_to_tensor(edge_index)
|
|
239
|
+
if edge_weight is not None:
|
|
240
|
+
edge_weight = ops.convert_to_tensor(edge_weight)
|
|
241
|
+
|
|
242
|
+
if num_nodes is None:
|
|
243
|
+
if is_tracing(edge_index):
|
|
244
|
+
num_nodes = edge_index.shape[1] if hasattr(edge_index, "shape") and edge_index.shape[1] is not None else ops.shape(edge_index)[1]
|
|
245
|
+
else:
|
|
246
|
+
num_nodes = int(ops.max(edge_index)) + 1 if ops.shape(edge_index)[1] > 0 else 0
|
|
247
|
+
try:
|
|
248
|
+
num_nodes = int(num_nodes)
|
|
249
|
+
except (TypeError, ValueError):
|
|
250
|
+
pass
|
|
251
|
+
|
|
252
|
+
if edge_weight is None:
|
|
253
|
+
num_edges = edge_index.shape[1] if hasattr(edge_index, "shape") and edge_index.shape[1] is not None else ops.shape(edge_index)[1]
|
|
254
|
+
edge_weight = ops.ones((num_edges,), dtype=dtype or "float32")
|
|
255
|
+
|
|
256
|
+
if add_self_loops:
|
|
257
|
+
edge_index, edge_weight = globals()["add_self_loops"](
|
|
258
|
+
edge_index, edge_weight, fill_value=fill_value, num_nodes=num_nodes
|
|
259
|
+
)
|
|
260
|
+
|
|
261
|
+
row, col = edge_index[0], edge_index[1]
|
|
262
|
+
idx = col if flow == "source_to_target" else row
|
|
263
|
+
row_cast = ops.cast(row, "int32")
|
|
264
|
+
col_cast = ops.cast(col, "int32")
|
|
265
|
+
idx_cast = ops.cast(idx, "int32")
|
|
266
|
+
|
|
267
|
+
deg = segment_sum(edge_weight, idx_cast, num_segments=num_nodes)
|
|
268
|
+
deg_inv_sqrt = ops.power(deg, -0.5)
|
|
269
|
+
deg_inv_sqrt = ops.where(
|
|
270
|
+
ops.isinf(deg_inv_sqrt) | ops.isnan(deg_inv_sqrt), 0.0, deg_inv_sqrt
|
|
271
|
+
)
|
|
272
|
+
|
|
273
|
+
norm = ops.take(deg_inv_sqrt, row_cast, axis=0) * edge_weight * ops.take(deg_inv_sqrt, col_cast, axis=0)
|
|
274
|
+
return edge_index, norm
|
|
275
|
+
|
|
276
|
+
|
|
277
|
+
def get_laplacian(
|
|
278
|
+
edge_index,
|
|
279
|
+
edge_weight=None,
|
|
280
|
+
normalization: Optional[str] = None,
|
|
281
|
+
dtype=None,
|
|
282
|
+
num_nodes: Optional[int] = None,
|
|
283
|
+
) -> Tuple:
|
|
284
|
+
"""Computes the graph Laplacian of the given graph."""
|
|
285
|
+
edge_index = ops.convert_to_tensor(edge_index)
|
|
286
|
+
if edge_weight is not None:
|
|
287
|
+
edge_weight = ops.convert_to_tensor(edge_weight)
|
|
288
|
+
|
|
289
|
+
if num_nodes is None:
|
|
290
|
+
if is_tracing(edge_index):
|
|
291
|
+
num_nodes = edge_index.shape[1] if hasattr(edge_index, "shape") and edge_index.shape[1] is not None else ops.shape(edge_index)[1]
|
|
292
|
+
else:
|
|
293
|
+
num_nodes = int(ops.max(edge_index)) + 1 if ops.shape(edge_index)[1] > 0 else 0
|
|
294
|
+
try:
|
|
295
|
+
num_nodes = int(num_nodes)
|
|
296
|
+
except (TypeError, ValueError):
|
|
297
|
+
pass
|
|
298
|
+
|
|
299
|
+
if edge_weight is None:
|
|
300
|
+
num_edges = edge_index.shape[1] if hasattr(edge_index, "shape") and edge_index.shape[1] is not None else ops.shape(edge_index)[1]
|
|
301
|
+
edge_weight = ops.ones((num_edges,), dtype=dtype or "float32")
|
|
302
|
+
|
|
303
|
+
row, col = ops.cast(edge_index[0], "int32"), ops.cast(edge_index[1], "int32")
|
|
304
|
+
deg = degree(row, num_nodes=num_nodes, dtype=edge_weight.dtype)
|
|
305
|
+
|
|
306
|
+
if normalization is None:
|
|
307
|
+
edge_index, _ = add_self_loops(edge_index, num_nodes=num_nodes)
|
|
308
|
+
edge_weight = ops.concatenate([-edge_weight, deg], axis=0)
|
|
309
|
+
elif normalization == "sym":
|
|
310
|
+
deg_inv_sqrt = ops.power(deg, -0.5)
|
|
311
|
+
deg_inv_sqrt = ops.where(
|
|
312
|
+
ops.isinf(deg_inv_sqrt) | ops.isnan(deg_inv_sqrt), 0.0, deg_inv_sqrt
|
|
313
|
+
)
|
|
314
|
+
edge_weight = (
|
|
315
|
+
ops.take(deg_inv_sqrt, row, axis=0)
|
|
316
|
+
* (-edge_weight)
|
|
317
|
+
* ops.take(deg_inv_sqrt, col, axis=0)
|
|
318
|
+
)
|
|
319
|
+
edge_index, _ = add_self_loops(edge_index, num_nodes=num_nodes)
|
|
320
|
+
edge_weight = ops.concatenate(
|
|
321
|
+
[edge_weight, ops.ones((num_nodes,), dtype=edge_weight.dtype)], axis=0
|
|
322
|
+
)
|
|
323
|
+
elif normalization == "rw":
|
|
324
|
+
deg_inv = 1.0 / deg
|
|
325
|
+
deg_inv = ops.where(
|
|
326
|
+
ops.isinf(deg_inv) | ops.isnan(deg_inv), 0.0, deg_inv
|
|
327
|
+
)
|
|
328
|
+
edge_weight = ops.take(deg_inv, row, axis=0) * (-edge_weight)
|
|
329
|
+
edge_index, _ = add_self_loops(edge_index, num_nodes=num_nodes)
|
|
330
|
+
edge_weight = ops.concatenate(
|
|
331
|
+
[edge_weight, ops.ones((num_nodes,), dtype=edge_weight.dtype)], axis=0
|
|
332
|
+
)
|
|
333
|
+
return edge_index, edge_weight
|
|
334
|
+
|
|
335
|
+
|
|
336
|
+
def _infer_dim_size(index, dim_size=None):
|
|
337
|
+
if dim_size is not None:
|
|
338
|
+
return dim_size
|
|
339
|
+
if hasattr(index, "is_meta") and index.is_meta:
|
|
340
|
+
return None
|
|
341
|
+
try:
|
|
342
|
+
if hasattr(index, "numpy") and not hasattr(index, "_has_symbolic_representation"):
|
|
343
|
+
# NumPy array or eager tensor with numpy()
|
|
344
|
+
import torch
|
|
345
|
+
if not isinstance(index, torch.Tensor):
|
|
346
|
+
return int(index.numpy().max()) + 1 if index.shape[0] > 0 else 0
|
|
347
|
+
if hasattr(index, "max"):
|
|
348
|
+
max_val = index.max()
|
|
349
|
+
if hasattr(max_val, "item"):
|
|
350
|
+
return int(max_val.item()) + 1
|
|
351
|
+
return int(max_val) + 1
|
|
352
|
+
return int(ops.convert_to_numpy(ops.max(index))) + 1 if ops.shape(index)[0] > 0 else 0
|
|
353
|
+
except Exception:
|
|
354
|
+
pass
|
|
355
|
+
try:
|
|
356
|
+
return ops.cast(ops.max(index), "int32") + 1
|
|
357
|
+
except Exception:
|
|
358
|
+
return None
|
|
359
|
+
|
|
360
|
+
|
|
361
|
+
def softmax(src, index, num_nodes: Optional[int] = None, dim: int = -2):
|
|
362
|
+
"""Computes a sparsely evaluated softmax over index."""
|
|
363
|
+
index = ops.cast(index, "int32")
|
|
364
|
+
num_nodes = _infer_dim_size(index, num_nodes)
|
|
365
|
+
|
|
366
|
+
max_val = segment_max(src, index, num_segments=num_nodes)
|
|
367
|
+
max_val = ops.take(max_val, index, axis=dim)
|
|
368
|
+
exp = ops.exp(src - max_val)
|
|
369
|
+
sum_val = segment_sum(exp, index, num_segments=num_nodes)
|
|
370
|
+
sum_val = ops.take(sum_val, index, axis=dim)
|
|
371
|
+
return exp / (sum_val + 1e-12)
|
|
372
|
+
|
|
373
|
+
|
|
374
|
+
def scatter(src, index, dim=0, dim_size=None, reduce="sum"):
|
|
375
|
+
"""Computes scatter / segment reduction."""
|
|
376
|
+
index = ops.cast(index, "int32")
|
|
377
|
+
dim_size = _infer_dim_size(index, dim_size)
|
|
378
|
+
|
|
379
|
+
if reduce in ("add", "sum"):
|
|
380
|
+
return segment_sum(src, index, num_segments=dim_size)
|
|
381
|
+
elif reduce == "mean":
|
|
382
|
+
sum_val = segment_sum(src, index, num_segments=dim_size)
|
|
383
|
+
ones = ops.ones_like(src)
|
|
384
|
+
count = segment_sum(ones, index, num_segments=dim_size)
|
|
385
|
+
count = ops.maximum(count, 1.0)
|
|
386
|
+
return sum_val / count
|
|
387
|
+
elif reduce == "max":
|
|
388
|
+
return segment_max(src, index, num_segments=dim_size)
|
|
389
|
+
val = segment_max(src, index, num_segments=dim_size)
|
|
390
|
+
ones = ops.ones((ops.shape(index)[0], 1), dtype=src.dtype)
|
|
391
|
+
count = segment_sum(ones, index, num_segments=dim_size)
|
|
392
|
+
return ops.where(ops.greater(count, 0), val, ops.zeros_like(val))
|
|
393
|
+
elif reduce == "min":
|
|
394
|
+
return segment_min(src, index, num_segments=dim_size)
|
|
395
|
+
val = segment_min(src, index, num_segments=dim_size)
|
|
396
|
+
ones = ops.ones((ops.shape(index)[0], 1), dtype=src.dtype)
|
|
397
|
+
count = segment_sum(ones, index, num_segments=dim_size)
|
|
398
|
+
return ops.where(ops.greater(count, 0), val, ops.zeros_like(val))
|
|
399
|
+
else:
|
|
400
|
+
raise ValueError(f"Unknown reduce operation: {reduce}")
|
|
401
|
+
|
|
402
|
+
|
|
403
|
+
|