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,101 @@
|
|
|
1
|
+
from keras import layers, ops
|
|
2
|
+
|
|
3
|
+
from k3_node.layers.conv.message_passing import MessagePassing
|
|
4
|
+
from k3_node.layers.conv.utils import gcn_norm, is_tracing
|
|
5
|
+
|
|
6
|
+
|
|
7
|
+
class SSGConv(MessagePassing):
|
|
8
|
+
r"""The simple spectral graph convolutional operator from the
|
|
9
|
+
`"Simple Spectral Graph Convolution" <https://arxiv.org/abs/2109.07191>`_ paper.
|
|
10
|
+
|
|
11
|
+
Args:
|
|
12
|
+
in_channels: Size of each input sample.
|
|
13
|
+
out_channels: Size of each output sample.
|
|
14
|
+
alpha: Teleport probability :math:`\alpha`.
|
|
15
|
+
K: Number of hops :math:`K`. (default: ``1``)
|
|
16
|
+
cached: If set to :obj:`True`, the layer will cache normalization coefficients.
|
|
17
|
+
(default: ``False``)
|
|
18
|
+
add_self_loops: If set to :obj:`False`, will not add self-loops.
|
|
19
|
+
(default: ``True``)
|
|
20
|
+
bias: If set to :obj:`False`, the layer will not learn an additive bias.
|
|
21
|
+
(default: ``True``)
|
|
22
|
+
|
|
23
|
+
Example:
|
|
24
|
+
```python
|
|
25
|
+
import numpy as np
|
|
26
|
+
from k3_node.layers import SSGConv
|
|
27
|
+
|
|
28
|
+
x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
|
|
29
|
+
edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
|
|
30
|
+
|
|
31
|
+
layer = SSGConv(in_channels=8, out_channels=16, alpha=0.1, K=2)
|
|
32
|
+
out = layer(x, edge_index)
|
|
33
|
+
print(tuple(out.shape)) # (10, 16)
|
|
34
|
+
```
|
|
35
|
+
"""
|
|
36
|
+
|
|
37
|
+
weighted_sum_message = True
|
|
38
|
+
|
|
39
|
+
def __init__(
|
|
40
|
+
self,
|
|
41
|
+
in_channels: int,
|
|
42
|
+
out_channels: int,
|
|
43
|
+
alpha: float,
|
|
44
|
+
K: int = 1,
|
|
45
|
+
cached: bool = False,
|
|
46
|
+
add_self_loops: bool = True,
|
|
47
|
+
bias: bool = True,
|
|
48
|
+
**kwargs,
|
|
49
|
+
):
|
|
50
|
+
super().__init__(aggr="add", **kwargs)
|
|
51
|
+
self.in_channels = in_channels
|
|
52
|
+
self.out_channels = out_channels
|
|
53
|
+
self.alpha = alpha
|
|
54
|
+
self.K = K
|
|
55
|
+
self.cached = cached
|
|
56
|
+
self.add_self_loops = add_self_loops
|
|
57
|
+
self.use_bias = bias
|
|
58
|
+
|
|
59
|
+
self.lin = layers.Dense(out_channels, use_bias=bias)
|
|
60
|
+
self._cached_edge_index = None
|
|
61
|
+
self._cached_norm = None
|
|
62
|
+
|
|
63
|
+
def build(self, input_shape):
|
|
64
|
+
feat_shape = input_shape[0] if isinstance(input_shape, (tuple, list)) and isinstance(input_shape[0], (tuple, list)) else input_shape
|
|
65
|
+
self.lin.build(feat_shape)
|
|
66
|
+
self.built = True
|
|
67
|
+
|
|
68
|
+
def call(self, x, edge_index=None, edge_weight=None, **kwargs):
|
|
69
|
+
if edge_index is None and isinstance(x, (tuple, list)):
|
|
70
|
+
x, edge_index = x[0], x[1]
|
|
71
|
+
|
|
72
|
+
if self.cached and self._cached_edge_index is not None:
|
|
73
|
+
edge_index = self._cached_edge_index
|
|
74
|
+
edge_weight = self._cached_norm
|
|
75
|
+
else:
|
|
76
|
+
num_nodes = x.shape[self.node_dim] if hasattr(x, "shape") and x.shape[self.node_dim] is not None else ops.shape(x)[self.node_dim]
|
|
77
|
+
edge_index, edge_weight = gcn_norm(
|
|
78
|
+
edge_index,
|
|
79
|
+
edge_weight,
|
|
80
|
+
num_nodes=num_nodes,
|
|
81
|
+
add_self_loops=self.add_self_loops,
|
|
82
|
+
flow=self.flow,
|
|
83
|
+
dtype=x.dtype,
|
|
84
|
+
)
|
|
85
|
+
if self.cached and not is_tracing(edge_index):
|
|
86
|
+
self._cached_edge_index = edge_index
|
|
87
|
+
self._cached_norm = edge_weight
|
|
88
|
+
|
|
89
|
+
out = self.alpha * x
|
|
90
|
+
h = x
|
|
91
|
+
for _ in range(self.K):
|
|
92
|
+
h = self.propagate(edge_index, x=h, edge_weight=edge_weight)
|
|
93
|
+
out = out + ((1.0 - self.alpha) / self.K) * h
|
|
94
|
+
|
|
95
|
+
return self.lin(out)
|
|
96
|
+
|
|
97
|
+
def message(self, x_j, edge_weight=None):
|
|
98
|
+
if edge_weight is None:
|
|
99
|
+
return x_j
|
|
100
|
+
return ops.expand_dims(edge_weight, -1) * x_j
|
|
101
|
+
|
|
@@ -0,0 +1,195 @@
|
|
|
1
|
+
import math
|
|
2
|
+
import keras
|
|
3
|
+
from keras import ops
|
|
4
|
+
from keras.layers import Dense, Dropout
|
|
5
|
+
|
|
6
|
+
from k3_node.layers.conv.message_passing import MessagePassing
|
|
7
|
+
from k3_node.layers.conv.utils import (
|
|
8
|
+
add_self_loops,
|
|
9
|
+
extend_mask_for_self_loops,
|
|
10
|
+
mask_edge_logits,
|
|
11
|
+
remove_self_loops_masked,
|
|
12
|
+
softmax,
|
|
13
|
+
)
|
|
14
|
+
from k3_node.ops.creation import full
|
|
15
|
+
|
|
16
|
+
|
|
17
|
+
class SuperGATConv(MessagePassing):
|
|
18
|
+
r"""The self-supervised graph attentional operator from the
|
|
19
|
+
`"How to Find Your Friendly Neighborhood: Graph Attention Design with Self-Supervision"
|
|
20
|
+
<https://openreview.net/forum?id=Wi5KUNlqWty>`_ paper.
|
|
21
|
+
|
|
22
|
+
Args:
|
|
23
|
+
attention_type (str): ``"MX"`` (mixed GO/DP) or ``"SD"`` (scaled dot-product).
|
|
24
|
+
neg_sample_ratio (float): Negative (random) pairs per positive edge in the attention loss.
|
|
25
|
+
edge_sample_ratio (float): Fraction of edges used as positives in the attention loss.
|
|
26
|
+
attention_loss_weight (float): If positive, the self-supervised attention loss is added to
|
|
27
|
+
the model loss (``model.losses``) with this weight while training, so ``fit`` optimizes
|
|
28
|
+
it automatically. PyG's example uses ``4.0``. (default: ``0.0``)
|
|
29
|
+
|
|
30
|
+
Example:
|
|
31
|
+
```python
|
|
32
|
+
import numpy as np
|
|
33
|
+
from k3_node.layers import SuperGATConv
|
|
34
|
+
|
|
35
|
+
x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
|
|
36
|
+
edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
|
|
37
|
+
|
|
38
|
+
layer = SuperGATConv(in_channels=8, out_channels=16, heads=2)
|
|
39
|
+
out = layer(x, edge_index)
|
|
40
|
+
print(tuple(out.shape)) # (10, 32)
|
|
41
|
+
```
|
|
42
|
+
"""
|
|
43
|
+
def __init__(
|
|
44
|
+
self,
|
|
45
|
+
in_channels: int,
|
|
46
|
+
out_channels: int,
|
|
47
|
+
heads: int = 1,
|
|
48
|
+
concat: bool = True,
|
|
49
|
+
negative_slope: float = 0.2,
|
|
50
|
+
dropout: float = 0.0,
|
|
51
|
+
add_self_loops: bool = True,
|
|
52
|
+
bias: bool = True,
|
|
53
|
+
attention_type: str = "MX",
|
|
54
|
+
neg_sample_ratio: float = 0.5,
|
|
55
|
+
edge_sample_ratio: float = 1.0,
|
|
56
|
+
is_undirected: bool = False,
|
|
57
|
+
attention_loss_weight: float = 0.0,
|
|
58
|
+
**kwargs,
|
|
59
|
+
):
|
|
60
|
+
kwargs.setdefault("aggr", "add")
|
|
61
|
+
super().__init__(node_dim=0, **kwargs)
|
|
62
|
+
|
|
63
|
+
assert attention_type in ["MX", "SD"]
|
|
64
|
+
|
|
65
|
+
self.in_channels = in_channels
|
|
66
|
+
self.out_channels = out_channels
|
|
67
|
+
self.heads = heads
|
|
68
|
+
self.concat = concat
|
|
69
|
+
self.negative_slope = negative_slope
|
|
70
|
+
self.dropout_rate = dropout
|
|
71
|
+
self.add_self_loops = add_self_loops
|
|
72
|
+
self.attention_type = attention_type
|
|
73
|
+
self.neg_sample_ratio = neg_sample_ratio
|
|
74
|
+
self.edge_sample_ratio = edge_sample_ratio
|
|
75
|
+
self.is_undirected = is_undirected
|
|
76
|
+
self.use_bias = bias
|
|
77
|
+
|
|
78
|
+
self.attention_loss_weight = attention_loss_weight
|
|
79
|
+
self.lin = Dense(heads * out_channels, use_bias=False)
|
|
80
|
+
self.dropout = Dropout(dropout)
|
|
81
|
+
self.seed_generator = keras.random.SeedGenerator()
|
|
82
|
+
self._last_attention_loss = None
|
|
83
|
+
|
|
84
|
+
if self.attention_type == "MX":
|
|
85
|
+
self.att_l = self.add_weight(
|
|
86
|
+
shape=(1, heads, out_channels),
|
|
87
|
+
initializer="glorot_uniform",
|
|
88
|
+
name="att_l",
|
|
89
|
+
)
|
|
90
|
+
self.att_r = self.add_weight(
|
|
91
|
+
shape=(1, heads, out_channels),
|
|
92
|
+
initializer="glorot_uniform",
|
|
93
|
+
name="att_r",
|
|
94
|
+
)
|
|
95
|
+
else:
|
|
96
|
+
self.att_l = None
|
|
97
|
+
self.att_r = None
|
|
98
|
+
|
|
99
|
+
if bias:
|
|
100
|
+
out_dim = heads * out_channels if concat else out_channels
|
|
101
|
+
self.bias = self.add_weight(
|
|
102
|
+
shape=(out_dim,),
|
|
103
|
+
initializer="zeros",
|
|
104
|
+
name="bias",
|
|
105
|
+
)
|
|
106
|
+
else:
|
|
107
|
+
self.bias = None
|
|
108
|
+
|
|
109
|
+
def build(self, input_shape=None):
|
|
110
|
+
self.lin.build((None, self.in_channels))
|
|
111
|
+
self.built = True
|
|
112
|
+
|
|
113
|
+
def call(self, inputs, edge_index=None, training=None, **kwargs):
|
|
114
|
+
if edge_index is None:
|
|
115
|
+
if isinstance(inputs, (list, tuple)) and len(inputs) == 2:
|
|
116
|
+
x, edge_index = inputs
|
|
117
|
+
else:
|
|
118
|
+
raise ValueError("Expected (x, edge_index) or x and edge_index")
|
|
119
|
+
else:
|
|
120
|
+
x = inputs
|
|
121
|
+
|
|
122
|
+
if not self.built:
|
|
123
|
+
self.build()
|
|
124
|
+
|
|
125
|
+
num_nodes = ops.shape(x)[0]
|
|
126
|
+
keep_mask = None
|
|
127
|
+
if self.add_self_loops:
|
|
128
|
+
edge_index, _, keep_mask = remove_self_loops_masked(edge_index)
|
|
129
|
+
edge_index, _ = add_self_loops(edge_index, num_nodes=num_nodes)
|
|
130
|
+
keep_mask = extend_mask_for_self_loops(keep_mask, num_nodes)
|
|
131
|
+
|
|
132
|
+
x = self.lin(x)
|
|
133
|
+
x = ops.reshape(x, (-1, self.heads, self.out_channels))
|
|
134
|
+
|
|
135
|
+
out = self.propagate(edge_index, x=x, keep_mask=keep_mask, training=training, size=(num_nodes, num_nodes))
|
|
136
|
+
|
|
137
|
+
if training:
|
|
138
|
+
loss = self._attention_loss(x, edge_index, num_nodes)
|
|
139
|
+
if self.attention_loss_weight:
|
|
140
|
+
self.add_loss(self.attention_loss_weight * loss)
|
|
141
|
+
from k3_node.layers.conv.utils import is_tracing
|
|
142
|
+
self._last_attention_loss = None if is_tracing(loss) else loss
|
|
143
|
+
|
|
144
|
+
if self.concat:
|
|
145
|
+
out = ops.reshape(out, (-1, self.heads * self.out_channels))
|
|
146
|
+
else:
|
|
147
|
+
out = ops.mean(out, axis=1)
|
|
148
|
+
|
|
149
|
+
if self.bias is not None:
|
|
150
|
+
out = out + self.bias
|
|
151
|
+
|
|
152
|
+
return out
|
|
153
|
+
|
|
154
|
+
def message(self, x_i, x_j, index=None, size_i=None, keep_mask=None, training=None):
|
|
155
|
+
if self.attention_type == "MX":
|
|
156
|
+
logits = ops.sum(x_i * x_j, axis=-1)
|
|
157
|
+
alpha = ops.sum(x_j * self.att_l, axis=-1) + ops.sum(x_i * self.att_r, axis=-1)
|
|
158
|
+
alpha = alpha * ops.sigmoid(logits)
|
|
159
|
+
else: # SD
|
|
160
|
+
alpha = ops.sum(x_i * x_j, axis=-1) / math.sqrt(self.out_channels)
|
|
161
|
+
|
|
162
|
+
alpha = ops.leaky_relu(alpha, negative_slope=self.negative_slope)
|
|
163
|
+
alpha = mask_edge_logits(alpha, keep_mask)
|
|
164
|
+
alpha = softmax(alpha, index, num_nodes=size_i, dim=0)
|
|
165
|
+
alpha = self.dropout(alpha, training=training)
|
|
166
|
+
return x_j * ops.expand_dims(alpha, -1)
|
|
167
|
+
|
|
168
|
+
def _attention_logits(self, x_i, x_j):
|
|
169
|
+
logits = ops.sum(x_i * x_j, axis=-1)
|
|
170
|
+
if self.attention_type == "SD":
|
|
171
|
+
logits = logits / math.sqrt(self.out_channels)
|
|
172
|
+
return logits
|
|
173
|
+
|
|
174
|
+
def _attention_loss(self, x, edge_index, num_nodes):
|
|
175
|
+
# Self-supervised attention loss: attention logits should separate real edges (label 1)
|
|
176
|
+
# from random node pairs (label 0). Edges are kept with probability `edge_sample_ratio` and
|
|
177
|
+
# random pairs with probability `neg_sample_ratio * edge_sample_ratio` (the expected counts
|
|
178
|
+
# used by PyG); weighting instead of slicing keeps all shapes static for XLA / jax.jit.
|
|
179
|
+
edge_index = ops.cast(edge_index, "int32")
|
|
180
|
+
num_edges = ops.shape(edge_index)[1]
|
|
181
|
+
neg = keras.random.randint(ops.shape(edge_index), 0, num_nodes, seed=self.seed_generator, dtype="int32")
|
|
182
|
+
pairs = ops.concatenate([edge_index, neg], axis=1)
|
|
183
|
+
logits = ops.mean(self._attention_logits(ops.take(x, pairs[1], axis=0), ops.take(x, pairs[0], axis=0)), axis=-1)
|
|
184
|
+
labels = ops.concatenate([ops.ones((num_edges,)), ops.zeros((num_edges,))])
|
|
185
|
+
keep = keras.random.uniform(ops.shape(labels), seed=self.seed_generator) < ops.concatenate([
|
|
186
|
+
full((num_edges,), self.edge_sample_ratio),
|
|
187
|
+
full((num_edges,), self.neg_sample_ratio * self.edge_sample_ratio),
|
|
188
|
+
])
|
|
189
|
+
weights = ops.cast(keep, "float32")
|
|
190
|
+
losses = ops.binary_crossentropy(labels, logits, from_logits=True)
|
|
191
|
+
return ops.sum(losses * weights) / ops.maximum(ops.sum(weights), 1.0)
|
|
192
|
+
|
|
193
|
+
def get_attention_loss(self):
|
|
194
|
+
r"""The self-supervised attention loss of the last training call (as in PyG)."""
|
|
195
|
+
return self._last_attention_loss
|
|
@@ -0,0 +1,98 @@
|
|
|
1
|
+
from keras import layers, ops
|
|
2
|
+
|
|
3
|
+
from k3_node.layers.conv.message_passing import MessagePassing
|
|
4
|
+
from k3_node.layers.conv.utils import gcn_norm
|
|
5
|
+
|
|
6
|
+
|
|
7
|
+
class TAGConv(MessagePassing):
|
|
8
|
+
r"""The topology adaptive graph convolutional operator from the
|
|
9
|
+
`"Topology Adaptive Graph Convolutional Networks"
|
|
10
|
+
<https://arxiv.org/abs/1710.10370>`_ paper.
|
|
11
|
+
|
|
12
|
+
Args:
|
|
13
|
+
in_channels: Size of each input sample.
|
|
14
|
+
out_channels: Size of each output sample.
|
|
15
|
+
K: Number of hops :math:`K`. (default: ``3``)
|
|
16
|
+
bias: If set to :obj:`False`, the layer will not learn an additive bias.
|
|
17
|
+
(default: ``True``)
|
|
18
|
+
normalize: Whether to apply symmetric normalization. (default: ``True``)
|
|
19
|
+
|
|
20
|
+
Example:
|
|
21
|
+
```python
|
|
22
|
+
import numpy as np
|
|
23
|
+
from k3_node.layers import TAGConv
|
|
24
|
+
|
|
25
|
+
x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
|
|
26
|
+
edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
|
|
27
|
+
|
|
28
|
+
layer = TAGConv(in_channels=8, out_channels=16, K=2)
|
|
29
|
+
out = layer(x, edge_index)
|
|
30
|
+
print(tuple(out.shape)) # (10, 16)
|
|
31
|
+
```
|
|
32
|
+
"""
|
|
33
|
+
|
|
34
|
+
weighted_sum_message = True
|
|
35
|
+
|
|
36
|
+
def __init__(
|
|
37
|
+
self,
|
|
38
|
+
in_channels: int,
|
|
39
|
+
out_channels: int,
|
|
40
|
+
K: int = 3,
|
|
41
|
+
bias: bool = True,
|
|
42
|
+
normalize: bool = True,
|
|
43
|
+
**kwargs,
|
|
44
|
+
):
|
|
45
|
+
super().__init__(aggr="add", **kwargs)
|
|
46
|
+
self.in_channels = in_channels
|
|
47
|
+
self.out_channels = out_channels
|
|
48
|
+
self.K = K
|
|
49
|
+
self.normalize = normalize
|
|
50
|
+
self.use_bias = bias
|
|
51
|
+
|
|
52
|
+
self.lins = [layers.Dense(out_channels, use_bias=False) for _ in range(K + 1)]
|
|
53
|
+
|
|
54
|
+
def build(self, input_shape):
|
|
55
|
+
feat_shape = input_shape[0] if isinstance(input_shape, (tuple, list)) and isinstance(input_shape[0], (tuple, list)) else input_shape
|
|
56
|
+
for lin in self.lins:
|
|
57
|
+
lin.build(feat_shape)
|
|
58
|
+
|
|
59
|
+
if self.use_bias:
|
|
60
|
+
self.bias = self.add_weight(
|
|
61
|
+
shape=(self.out_channels,),
|
|
62
|
+
initializer="zeros",
|
|
63
|
+
name="bias",
|
|
64
|
+
)
|
|
65
|
+
else:
|
|
66
|
+
self.bias = None
|
|
67
|
+
self.built = True
|
|
68
|
+
|
|
69
|
+
def call(self, x, edge_index=None, edge_weight=None, **kwargs):
|
|
70
|
+
if edge_index is None and isinstance(x, (tuple, list)):
|
|
71
|
+
x, edge_index = x[0], x[1]
|
|
72
|
+
|
|
73
|
+
if self.normalize:
|
|
74
|
+
num_nodes = x.shape[self.node_dim] if hasattr(x, "shape") and x.shape[self.node_dim] is not None else ops.shape(x)[self.node_dim]
|
|
75
|
+
edge_index, edge_weight = gcn_norm(
|
|
76
|
+
edge_index,
|
|
77
|
+
edge_weight,
|
|
78
|
+
num_nodes=num_nodes,
|
|
79
|
+
add_self_loops=False,
|
|
80
|
+
flow=self.flow,
|
|
81
|
+
dtype=x.dtype,
|
|
82
|
+
)
|
|
83
|
+
|
|
84
|
+
out = self.lins[0](x)
|
|
85
|
+
h = x
|
|
86
|
+
for k in range(1, self.K + 1):
|
|
87
|
+
h = self.propagate(edge_index, x=h, edge_weight=edge_weight)
|
|
88
|
+
out = out + self.lins[k](h)
|
|
89
|
+
|
|
90
|
+
if self.bias is not None:
|
|
91
|
+
out = out + self.bias
|
|
92
|
+
return out
|
|
93
|
+
|
|
94
|
+
def message(self, x_j, edge_weight=None):
|
|
95
|
+
if edge_weight is None:
|
|
96
|
+
return x_j
|
|
97
|
+
return ops.expand_dims(edge_weight, -1) * x_j
|
|
98
|
+
|
|
@@ -0,0 +1,164 @@
|
|
|
1
|
+
"""Cross-backend and compiled-mode consistency tests for convolution layers.
|
|
2
|
+
|
|
3
|
+
Each layer is run with identical seeded inputs and weights, and its eager output is
|
|
4
|
+
compared against golden outputs recorded with the torch backend (the backend that
|
|
5
|
+
``tests_reference`` validates against PyG). On TensorFlow and JAX the output under
|
|
6
|
+
``jit_compile=True`` must also match the eager output.
|
|
7
|
+
|
|
8
|
+
The input graph deliberately contains self-loops: layers that remove and re-add
|
|
9
|
+
self-loops must handle them without dynamic shapes when compiled.
|
|
10
|
+
|
|
11
|
+
To regenerate the golden file after an intentional numerical change::
|
|
12
|
+
|
|
13
|
+
KERAS_BACKEND=torch python k3_node/layers/conv/test_backend_consistency.py
|
|
14
|
+
"""
|
|
15
|
+
import os
|
|
16
|
+
import os.path as osp
|
|
17
|
+
|
|
18
|
+
import numpy as np
|
|
19
|
+
import pytest
|
|
20
|
+
import keras
|
|
21
|
+
from keras import layers
|
|
22
|
+
|
|
23
|
+
import k3_node.layers as L
|
|
24
|
+
|
|
25
|
+
GOLDEN_PATH = osp.join(osp.dirname(__file__), "testdata", "backend_consistency_golden.npz")
|
|
26
|
+
|
|
27
|
+
N, E, C, O = 12, 30, 8, 5
|
|
28
|
+
|
|
29
|
+
|
|
30
|
+
def _graph():
|
|
31
|
+
rng = np.random.default_rng(0)
|
|
32
|
+
edge_index = rng.integers(0, N, size=(2, E))
|
|
33
|
+
loops = np.array([[0, 3, 7], [0, 3, 7]]) # guarantee pre-existing self-loops
|
|
34
|
+
return {
|
|
35
|
+
"x": rng.standard_normal((N, C)).astype("float32"),
|
|
36
|
+
"pos": rng.standard_normal((N, 3)).astype("float32"),
|
|
37
|
+
"normal": rng.standard_normal((N, 3)).astype("float32"),
|
|
38
|
+
"edge_index": np.concatenate([edge_index, loops], axis=1).astype("int32"),
|
|
39
|
+
}
|
|
40
|
+
|
|
41
|
+
|
|
42
|
+
def _mlp(units):
|
|
43
|
+
return keras.Sequential([layers.Dense(units, activation="relu"), layers.Dense(units)])
|
|
44
|
+
|
|
45
|
+
|
|
46
|
+
def _xe(layer, g):
|
|
47
|
+
return layer(g["x"], g["edge_index"])
|
|
48
|
+
|
|
49
|
+
|
|
50
|
+
# name -> (layer factory, call function)
|
|
51
|
+
CASES = {
|
|
52
|
+
"GCNConv": (lambda: L.GCNConv(C, O), _xe),
|
|
53
|
+
"GATConv": (lambda: L.GATConv(C, O, heads=2), _xe),
|
|
54
|
+
"GATv2Conv": (lambda: L.GATv2Conv(C, O, heads=2), _xe),
|
|
55
|
+
"SuperGATConv": (lambda: L.SuperGATConv(C, O, heads=2), _xe),
|
|
56
|
+
"ClusterGCNConv": (lambda: L.ClusterGCNConv(C, O), _xe),
|
|
57
|
+
"AGNNConv": (lambda: L.AGNNConv(), _xe),
|
|
58
|
+
"FeaStConv": (lambda: L.FeaStConv(C, O, heads=2), _xe),
|
|
59
|
+
"PointNetConv": (lambda: L.PointNetConv(local_nn=_mlp(O)), lambda l, g: l(g["x"], g["pos"], g["edge_index"])),
|
|
60
|
+
"PointTransformerConv": (
|
|
61
|
+
lambda: L.PointTransformerConv(C, O),
|
|
62
|
+
lambda l, g: l(g["x"], g["pos"], g["edge_index"]),
|
|
63
|
+
),
|
|
64
|
+
"PPFConv": (
|
|
65
|
+
lambda: L.PPFConv(local_nn=_mlp(O)),
|
|
66
|
+
lambda l, g: l(g["x"], g["pos"], g["normal"], g["edge_index"]),
|
|
67
|
+
),
|
|
68
|
+
"SAGEConv": (lambda: L.SAGEConv(C, O), _xe),
|
|
69
|
+
"GraphConv": (lambda: L.GraphConv(C, O), _xe),
|
|
70
|
+
"TransformerConv": (lambda: L.TransformerConv(C, O, heads=2), _xe),
|
|
71
|
+
"ChebConv": (lambda: L.ChebConv(C, O, K=3), _xe),
|
|
72
|
+
"TAGConv": (lambda: L.TAGConv(C, O, K=2), _xe),
|
|
73
|
+
"SGConv": (lambda: L.SGConv(C, O, K=2), _xe),
|
|
74
|
+
"LEConv": (lambda: L.LEConv(C, O), _xe),
|
|
75
|
+
"ResGatedGraphConv": (lambda: L.ResGatedGraphConv(C, O), _xe),
|
|
76
|
+
"GENConv": (lambda: L.GENConv(C, O), _xe),
|
|
77
|
+
"GINConv": (lambda: L.GINConv(_mlp(O)), _xe),
|
|
78
|
+
"EdgeConv": (lambda: L.EdgeConv(keras.Sequential([layers.Dense(O)])), _xe),
|
|
79
|
+
"MFConv": (lambda: L.MFConv(C, O), _xe),
|
|
80
|
+
"FiLMConv": (lambda: L.FiLMConv(C, O), _xe),
|
|
81
|
+
"GeneralConv": (lambda: L.GeneralConv(C, O), _xe),
|
|
82
|
+
"PNAConv": (
|
|
83
|
+
lambda: L.PNAConv(
|
|
84
|
+
C,
|
|
85
|
+
O,
|
|
86
|
+
aggregators=["mean", "max", "min", "std"],
|
|
87
|
+
scalers=["identity", "amplification"],
|
|
88
|
+
deg=np.array([0, 2, 4, 3, 2, 1], "int32"),
|
|
89
|
+
),
|
|
90
|
+
_xe,
|
|
91
|
+
),
|
|
92
|
+
}
|
|
93
|
+
|
|
94
|
+
|
|
95
|
+
class ConsistencyWrapper(keras.Model):
|
|
96
|
+
def __init__(self, layer, call_fn):
|
|
97
|
+
super().__init__()
|
|
98
|
+
self.layer = layer
|
|
99
|
+
self.call_fn = call_fn
|
|
100
|
+
|
|
101
|
+
def call(self, inputs):
|
|
102
|
+
out = self.call_fn(self.layer, inputs)
|
|
103
|
+
return out[0] if isinstance(out, (tuple, list)) else out
|
|
104
|
+
|
|
105
|
+
|
|
106
|
+
def _build(name):
|
|
107
|
+
factory, call_fn = CASES[name]
|
|
108
|
+
g = _graph()
|
|
109
|
+
model = ConsistencyWrapper(factory(), call_fn)
|
|
110
|
+
model(g)
|
|
111
|
+
rng = np.random.default_rng(123)
|
|
112
|
+
weights = []
|
|
113
|
+
for v in model.weights:
|
|
114
|
+
w = rng.standard_normal(v.shape) * 0.3
|
|
115
|
+
if "moving_variance" in v.path:
|
|
116
|
+
w = np.abs(w) + 0.5
|
|
117
|
+
weights.append(w.astype(v.dtype))
|
|
118
|
+
model.set_weights(weights)
|
|
119
|
+
shapes = np.array([str(tuple(v.shape)) for v in model.weights])
|
|
120
|
+
return model, g, shapes
|
|
121
|
+
|
|
122
|
+
|
|
123
|
+
def _eager(model, g):
|
|
124
|
+
return keras.ops.convert_to_numpy(model(g))
|
|
125
|
+
|
|
126
|
+
|
|
127
|
+
@pytest.fixture(scope="module")
|
|
128
|
+
def golden():
|
|
129
|
+
if not osp.exists(GOLDEN_PATH):
|
|
130
|
+
pytest.fail(f"Golden file missing; regenerate it (see module docstring): {GOLDEN_PATH}")
|
|
131
|
+
return np.load(GOLDEN_PATH)
|
|
132
|
+
|
|
133
|
+
|
|
134
|
+
@pytest.mark.parametrize("name", sorted(CASES))
|
|
135
|
+
def test_eager_matches_torch_golden(name, golden):
|
|
136
|
+
model, g, shapes = _build(name)
|
|
137
|
+
assert list(shapes) == list(golden[f"{name}/shapes"]), (
|
|
138
|
+
f"{name}: weight layout differs from the torch backend, so weights cannot be shared"
|
|
139
|
+
)
|
|
140
|
+
np.testing.assert_allclose(_eager(model, g), golden[f"{name}/out"], rtol=1e-4, atol=1e-5)
|
|
141
|
+
|
|
142
|
+
|
|
143
|
+
@pytest.mark.skipif(
|
|
144
|
+
keras.backend.backend() not in ("tensorflow", "jax"),
|
|
145
|
+
reason="jit_compile=True means XLA only on the TensorFlow and JAX backends",
|
|
146
|
+
)
|
|
147
|
+
@pytest.mark.parametrize("name", sorted(CASES))
|
|
148
|
+
def test_jit_matches_eager(name):
|
|
149
|
+
model, g, _ = _build(name)
|
|
150
|
+
expected = _eager(model, g)
|
|
151
|
+
model.compile(jit_compile=True)
|
|
152
|
+
np.testing.assert_allclose(np.asarray(model.predict_on_batch(g)), expected, rtol=1e-4, atol=1e-5)
|
|
153
|
+
|
|
154
|
+
|
|
155
|
+
if __name__ == "__main__":
|
|
156
|
+
assert keras.backend.backend() == "torch", "Golden outputs must be generated with KERAS_BACKEND=torch"
|
|
157
|
+
arrays = {}
|
|
158
|
+
for case in sorted(CASES):
|
|
159
|
+
m, graph, layer_shapes = _build(case)
|
|
160
|
+
arrays[f"{case}/out"] = _eager(m, graph)
|
|
161
|
+
arrays[f"{case}/shapes"] = layer_shapes
|
|
162
|
+
os.makedirs(osp.dirname(GOLDEN_PATH), exist_ok=True)
|
|
163
|
+
np.savez(GOLDEN_PATH, **arrays)
|
|
164
|
+
print(f"Wrote {len(CASES)} golden outputs to {GOLDEN_PATH}")
|