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,106 @@
|
|
|
1
|
+
"""Runs the ``Example:`` code in the docstrings of public layers and models.
|
|
2
|
+
|
|
3
|
+
Every example must run on the current backend, and each ``print(...) # expected`` line must print
|
|
4
|
+
the value in its comment, so the documented output shapes cannot drift from the code.
|
|
5
|
+
"""
|
|
6
|
+
import ast
|
|
7
|
+
import inspect
|
|
8
|
+
import re
|
|
9
|
+
import textwrap
|
|
10
|
+
|
|
11
|
+
import keras
|
|
12
|
+
import pytest
|
|
13
|
+
|
|
14
|
+
import k3_node.layers
|
|
15
|
+
import k3_node.models
|
|
16
|
+
|
|
17
|
+
# Public objects that intentionally have no forward-pass example.
|
|
18
|
+
NO_EXAMPLE = {
|
|
19
|
+
"Aggregation", # abstract base class; see the concrete aggregations
|
|
20
|
+
"KGEModel", # abstract base class; see TransE, DistMult, ComplEx, RotatE
|
|
21
|
+
"Connect", # abstract base class of the pooling "connect" step
|
|
22
|
+
"Select", # abstract base class of the pooling "select" step
|
|
23
|
+
"BasicGNN", # abstract base class; see GCN, GraphSAGE, GIN, GAT, PNA, EdgeCNN
|
|
24
|
+
}
|
|
25
|
+
|
|
26
|
+
_FENCE = re.compile(r"```python\n(.*?)```", re.S)
|
|
27
|
+
|
|
28
|
+
|
|
29
|
+
def _public_objects():
|
|
30
|
+
"""Unique public layer/model classes (and layer functions), keyed by public name."""
|
|
31
|
+
found, seen = {}, set()
|
|
32
|
+
for package in (k3_node.layers, k3_node.models):
|
|
33
|
+
names = getattr(package, "__all__", None) or [n for n in dir(package) if not n.startswith("_")]
|
|
34
|
+
for name in sorted(names):
|
|
35
|
+
obj = getattr(package, name, None)
|
|
36
|
+
is_layer = inspect.isclass(obj) and issubclass(obj, keras.layers.Layer)
|
|
37
|
+
is_layer_fn = (
|
|
38
|
+
inspect.isfunction(obj) and package is k3_node.layers and obj.__module__.startswith("k3_node.layers")
|
|
39
|
+
)
|
|
40
|
+
if (is_layer or is_layer_fn) and id(obj) not in seen and obj.__module__.startswith("k3_node"):
|
|
41
|
+
seen.add(id(obj))
|
|
42
|
+
found[f"{package.__name__}.{name}"] = obj
|
|
43
|
+
return found
|
|
44
|
+
|
|
45
|
+
|
|
46
|
+
def _other_documented_objects():
|
|
47
|
+
"""Helpers outside of layers/models whose docstring examples are checked too."""
|
|
48
|
+
from k3_node import metrics, training, utils
|
|
49
|
+
from k3_node.data import Data
|
|
50
|
+
from k3_node.datasets import Digits, SEALDataset
|
|
51
|
+
from k3_node.ops.sparse import spmm
|
|
52
|
+
from k3_node.transforms import RandomLinkSplit
|
|
53
|
+
|
|
54
|
+
objects = [utils.normalized_cut, utils.k_hop_subgraph, utils.drnl_node_labeling, Digits, SEALDataset,
|
|
55
|
+
Data.edge_subgraph, RandomLinkSplit, spmm, training.gradient_step, metrics.F1Score]
|
|
56
|
+
return {f"{obj.__module__}.{obj.__qualname__}": obj for obj in objects}
|
|
57
|
+
|
|
58
|
+
|
|
59
|
+
PUBLIC = {**_public_objects(), **_other_documented_objects()}
|
|
60
|
+
|
|
61
|
+
|
|
62
|
+
def _examples(obj):
|
|
63
|
+
doc = obj.__doc__ or ""
|
|
64
|
+
section = doc[doc.find("Example") :] if "Example" in doc else ""
|
|
65
|
+
return [textwrap.dedent(block) for block in _FENCE.findall(section)]
|
|
66
|
+
|
|
67
|
+
|
|
68
|
+
def _expected_prints(code):
|
|
69
|
+
"""Maps the line number of each top-level print() to the text of its trailing comment."""
|
|
70
|
+
expected = {}
|
|
71
|
+
lines = code.split("\n")
|
|
72
|
+
for node in ast.parse(code).body:
|
|
73
|
+
if isinstance(node, ast.Expr) and isinstance(node.value, ast.Call) and getattr(node.value.func, "id", "") == "print":
|
|
74
|
+
_, sep, comment = lines[node.lineno - 1].partition(" # ")
|
|
75
|
+
if sep:
|
|
76
|
+
expected[node.lineno] = comment
|
|
77
|
+
return expected
|
|
78
|
+
|
|
79
|
+
|
|
80
|
+
WITH_EXAMPLES = [name for name, obj in PUBLIC.items() if _examples(obj)]
|
|
81
|
+
|
|
82
|
+
|
|
83
|
+
@pytest.mark.parametrize("name", WITH_EXAMPLES)
|
|
84
|
+
def test_docstring_example_runs(name):
|
|
85
|
+
for code in _examples(PUBLIC[name]):
|
|
86
|
+
expected = _expected_prints(code)
|
|
87
|
+
printed = []
|
|
88
|
+
|
|
89
|
+
def capture(*args, **kwargs):
|
|
90
|
+
frame = inspect.currentframe().f_back
|
|
91
|
+
printed.append((frame.f_lineno, " ".join(str(a) for a in args)))
|
|
92
|
+
|
|
93
|
+
exec(compile(code, f"<{name} example>", "exec"), {"print": capture, "__name__": "__example__"})
|
|
94
|
+
for lineno, text in printed:
|
|
95
|
+
if lineno in expected:
|
|
96
|
+
assert expected[lineno].startswith(text), (
|
|
97
|
+
f"{name}: line {lineno} printed {text!r}, but the docstring says {expected[lineno]!r}"
|
|
98
|
+
)
|
|
99
|
+
|
|
100
|
+
|
|
101
|
+
def test_every_public_layer_and_model_has_an_example():
|
|
102
|
+
missing = sorted(
|
|
103
|
+
name for name, obj in PUBLIC.items()
|
|
104
|
+
if not _examples(obj) and name.rpartition(".")[2] not in NO_EXAMPLE and inspect.isclass(obj)
|
|
105
|
+
)
|
|
106
|
+
assert not missing, f"{len(missing)} public layers/models have no docstring example: {missing}"
|
|
@@ -0,0 +1,116 @@
|
|
|
1
|
+
"""Every k3 layer must forward `training` to its dropout / batch-norm sublayers.
|
|
2
|
+
|
|
3
|
+
On the JAX backend Keras does not propagate `training` from a model to nested layers during
|
|
4
|
+
`fit`, so a layer that calls its dropout or batch norm without passing `training` silently runs
|
|
5
|
+
it in inference mode. This test reproduces that: it disables Keras' propagation, enables dropout
|
|
6
|
+
in every constructor, runs all docstring examples with `training=True`, and reports any nested
|
|
7
|
+
call that did not receive `training` from its parent.
|
|
8
|
+
"""
|
|
9
|
+
import collections
|
|
10
|
+
import inspect
|
|
11
|
+
import re
|
|
12
|
+
import textwrap
|
|
13
|
+
|
|
14
|
+
import keras
|
|
15
|
+
import pytest
|
|
16
|
+
from keras.src.layers.layer import Layer
|
|
17
|
+
|
|
18
|
+
import k3_node.layers
|
|
19
|
+
import k3_node.models
|
|
20
|
+
|
|
21
|
+
_TRAINING_SENSITIVE = ("BatchNormalization", "BatchNorm", "InstanceNorm", "HeteroBatchNorm")
|
|
22
|
+
|
|
23
|
+
|
|
24
|
+
def _training_sensitive(layer):
|
|
25
|
+
for sub in layer._flatten_layers(include_self=True):
|
|
26
|
+
if type(sub).__name__ in _TRAINING_SENSITIVE:
|
|
27
|
+
return True
|
|
28
|
+
if isinstance(sub, keras.layers.Dropout) and float(sub.rate or 0) > 0:
|
|
29
|
+
return True
|
|
30
|
+
if isinstance(getattr(sub, "dropout", None), float) and sub.dropout > 0: # e.g. attention dropout
|
|
31
|
+
return True
|
|
32
|
+
return False
|
|
33
|
+
|
|
34
|
+
|
|
35
|
+
def _enable_dropout_defaults(cls):
|
|
36
|
+
"""Sets zero-valued dropout defaults of cls.__init__ to 0.1 (so dropout paths are exercised)."""
|
|
37
|
+
init = cls.__dict__.get("__init__")
|
|
38
|
+
if init is None:
|
|
39
|
+
return None
|
|
40
|
+
saved = (init.__defaults__, dict(init.__kwdefaults__ or {}))
|
|
41
|
+
params = [p for p in inspect.signature(init).parameters.values() if p.kind in (p.POSITIONAL_ONLY, p.POSITIONAL_OR_KEYWORD)]
|
|
42
|
+
if init.__defaults__:
|
|
43
|
+
defaults = list(init.__defaults__)
|
|
44
|
+
for i, p in enumerate(params[len(params) - len(defaults):]):
|
|
45
|
+
if "drop" in p.name and not isinstance(defaults[i], bool) and defaults[i] == 0:
|
|
46
|
+
defaults[i] = 0.1
|
|
47
|
+
init.__defaults__ = tuple(defaults)
|
|
48
|
+
for k, v in (init.__kwdefaults__ or {}).items():
|
|
49
|
+
if "drop" in k and not isinstance(v, bool) and v == 0:
|
|
50
|
+
init.__kwdefaults__[k] = 0.1
|
|
51
|
+
return init, saved
|
|
52
|
+
|
|
53
|
+
|
|
54
|
+
@pytest.fixture
|
|
55
|
+
def no_training_propagation(monkeypatch):
|
|
56
|
+
stack, gaps = [], collections.Counter()
|
|
57
|
+
original_resolve = Layer._resolve_and_populate_arg
|
|
58
|
+
original_call = Layer.__call__
|
|
59
|
+
|
|
60
|
+
def resolve(self, arg_name, call_spec, call_context, kwargs):
|
|
61
|
+
if arg_name != "training":
|
|
62
|
+
return original_resolve(self, arg_name, call_spec, call_context, kwargs)
|
|
63
|
+
passed = arg_name in call_spec.user_arguments_dict
|
|
64
|
+
value = call_spec.user_arguments_dict.get(arg_name) if passed else (True if len(stack) <= 1 else None)
|
|
65
|
+
if self._call_has_context_arg.get(arg_name, False) and value is not None:
|
|
66
|
+
kwargs[arg_name] = value
|
|
67
|
+
parent = stack[-2] if len(stack) > 1 else None
|
|
68
|
+
if (parent is not None and not passed and type(parent).__module__.startswith("k3_node")
|
|
69
|
+
and self._call_has_context_arg.get("training", False) and _training_sensitive(self)):
|
|
70
|
+
gaps[f"{type(parent).__module__}.{type(parent).__name__} -> {type(self).__name__}"] += 1
|
|
71
|
+
|
|
72
|
+
def call(self, *args, **kwargs):
|
|
73
|
+
stack.append(self)
|
|
74
|
+
try:
|
|
75
|
+
return original_call(self, *args, **kwargs)
|
|
76
|
+
finally:
|
|
77
|
+
stack.pop()
|
|
78
|
+
|
|
79
|
+
monkeypatch.setattr(Layer, "_resolve_and_populate_arg", resolve)
|
|
80
|
+
monkeypatch.setattr(Layer, "__call__", call)
|
|
81
|
+
return gaps
|
|
82
|
+
|
|
83
|
+
|
|
84
|
+
@pytest.mark.skipif(keras.backend.backend() != "torch", reason="checks the code, not a backend; torch is fastest")
|
|
85
|
+
def test_layers_forward_training_to_dropout_and_batch_norm(no_training_propagation):
|
|
86
|
+
import sys
|
|
87
|
+
|
|
88
|
+
patched = []
|
|
89
|
+
for module in list(sys.modules.values()):
|
|
90
|
+
if getattr(module, "__name__", "").startswith("k3_node"):
|
|
91
|
+
for obj in list(vars(module).values()):
|
|
92
|
+
if inspect.isclass(obj) and issubclass(obj, Layer) and obj.__module__.startswith("k3_node"):
|
|
93
|
+
result = _enable_dropout_defaults(obj)
|
|
94
|
+
if result:
|
|
95
|
+
patched.append(result)
|
|
96
|
+
try:
|
|
97
|
+
seen = set()
|
|
98
|
+
for package in (k3_node.layers, k3_node.models):
|
|
99
|
+
for name in dir(package):
|
|
100
|
+
obj = getattr(package, name)
|
|
101
|
+
if id(obj) in seen or not getattr(obj, "__module__", "").startswith("k3_node"):
|
|
102
|
+
continue
|
|
103
|
+
seen.add(id(obj))
|
|
104
|
+
doc = obj.__doc__ or ""
|
|
105
|
+
for code in re.findall(r"```python\n(.*?)```", doc[doc.find("Example"):] if "Example" in doc else "", re.S):
|
|
106
|
+
exec(compile(textwrap.dedent(code), name, "exec"), {"print": lambda *a, **k: None})
|
|
107
|
+
finally:
|
|
108
|
+
for init, (defaults, kwdefaults) in patched:
|
|
109
|
+
init.__defaults__ = defaults
|
|
110
|
+
if init.__kwdefaults__ is not None:
|
|
111
|
+
init.__kwdefaults__.clear()
|
|
112
|
+
init.__kwdefaults__.update(kwdefaults)
|
|
113
|
+
assert not no_training_propagation, (
|
|
114
|
+
"These layers call a dropout / batch-norm sublayer without forwarding `training` "
|
|
115
|
+
f"(it would never train on JAX): {sorted(no_training_propagation)}"
|
|
116
|
+
)
|
k3_node/training.py
ADDED
|
@@ -0,0 +1,115 @@
|
|
|
1
|
+
"""Backend-agnostic training helpers for models that don't fit ``keras.Model.fit``.
|
|
2
|
+
|
|
3
|
+
Examples are models with several optimizers (like the adversarial autoencoders) or losses that
|
|
4
|
+
are computed outside of a model's ``call``.
|
|
5
|
+
"""
|
|
6
|
+
from typing import Callable, Sequence
|
|
7
|
+
|
|
8
|
+
import keras
|
|
9
|
+
from keras import ops
|
|
10
|
+
|
|
11
|
+
|
|
12
|
+
def _seed_state():
|
|
13
|
+
from keras.src.random.seed_generator import global_seed_generator
|
|
14
|
+
|
|
15
|
+
return global_seed_generator().state
|
|
16
|
+
|
|
17
|
+
|
|
18
|
+
def _scalar(loss):
|
|
19
|
+
return ops.mean(loss) if len(ops.shape(loss)) > 0 else loss
|
|
20
|
+
|
|
21
|
+
|
|
22
|
+
def no_grad():
|
|
23
|
+
r"""A context manager for evaluating a model outside of ``fit`` / ``predict``, like PyTorch's
|
|
24
|
+
``torch.no_grad()``.
|
|
25
|
+
|
|
26
|
+
On the torch backend, calling a model records every intermediate result for a possible
|
|
27
|
+
backward pass, which can take far more memory than the model itself. Inside ``no_grad()``
|
|
28
|
+
nothing is recorded. TensorFlow and JAX only record gradients when asked to, so there it does
|
|
29
|
+
nothing.
|
|
30
|
+
|
|
31
|
+
Example:
|
|
32
|
+
```python
|
|
33
|
+
import keras
|
|
34
|
+
import numpy as np
|
|
35
|
+
from k3_node.training import no_grad
|
|
36
|
+
|
|
37
|
+
model = keras.layers.Dense(4)
|
|
38
|
+
with no_grad():
|
|
39
|
+
out = model(np.random.rand(10, 8).astype("float32"))
|
|
40
|
+
print(tuple(out.shape)) # (10, 4)
|
|
41
|
+
```
|
|
42
|
+
"""
|
|
43
|
+
if keras.config.backend() == "torch":
|
|
44
|
+
import torch
|
|
45
|
+
|
|
46
|
+
return torch.no_grad()
|
|
47
|
+
import contextlib
|
|
48
|
+
|
|
49
|
+
return contextlib.nullcontext()
|
|
50
|
+
|
|
51
|
+
|
|
52
|
+
def gradient_step(loss_fn: Callable, variables: Sequence, optimizer, state_variables: Sequence = ()):
|
|
53
|
+
r"""Runs ``loss_fn()``, then updates ``variables`` with one ``optimizer`` step on its gradients.
|
|
54
|
+
|
|
55
|
+
Works eagerly on every backend. Variables used by ``loss_fn`` but not listed in ``variables``
|
|
56
|
+
are not trained. ``state_variables`` are non-trainable variables that ``loss_fn`` may update
|
|
57
|
+
(for example the moving statistics of batch normalization); this matters on JAX only.
|
|
58
|
+
|
|
59
|
+
Args:
|
|
60
|
+
loss_fn (callable): A function without arguments that returns the loss (a scalar, or
|
|
61
|
+
per-example losses, which are averaged).
|
|
62
|
+
variables (list): The variables to train, e.g. ``model.trainable_variables``.
|
|
63
|
+
optimizer (keras.optimizers.Optimizer): The optimizer applying the update.
|
|
64
|
+
state_variables (list, optional): Non-trainable variables updated by ``loss_fn``.
|
|
65
|
+
|
|
66
|
+
Returns:
|
|
67
|
+
The loss as a Python float.
|
|
68
|
+
|
|
69
|
+
Example:
|
|
70
|
+
```python
|
|
71
|
+
import keras
|
|
72
|
+
from keras import ops
|
|
73
|
+
from k3_node.training import gradient_step
|
|
74
|
+
|
|
75
|
+
w = keras.Variable(3.0)
|
|
76
|
+
optimizer = keras.optimizers.SGD(learning_rate=0.25)
|
|
77
|
+
loss = gradient_step(lambda: ops.square(w), [w], optimizer) # gradient 2 * w = 6
|
|
78
|
+
print(loss, float(ops.convert_to_numpy(w))) # 9.0 1.5
|
|
79
|
+
```
|
|
80
|
+
"""
|
|
81
|
+
variables = list(variables)
|
|
82
|
+
backend = keras.config.backend()
|
|
83
|
+
|
|
84
|
+
if backend == "tensorflow":
|
|
85
|
+
import tensorflow as tf
|
|
86
|
+
|
|
87
|
+
with tf.GradientTape() as tape:
|
|
88
|
+
loss = _scalar(loss_fn())
|
|
89
|
+
grads = tape.gradient(loss, variables)
|
|
90
|
+
elif backend == "torch":
|
|
91
|
+
import torch
|
|
92
|
+
|
|
93
|
+
loss = _scalar(loss_fn())
|
|
94
|
+
grads = torch.autograd.grad(loss, [v.value for v in variables], allow_unused=True)
|
|
95
|
+
elif backend == "jax":
|
|
96
|
+
import jax
|
|
97
|
+
from keras.src.backend.common.stateless_scope import StatelessScope
|
|
98
|
+
|
|
99
|
+
tracked = list(state_variables) + [_seed_state()]
|
|
100
|
+
|
|
101
|
+
def compute(values):
|
|
102
|
+
with StatelessScope(state_mapping=list(zip(variables, values))) as scope:
|
|
103
|
+
out = _scalar(loss_fn())
|
|
104
|
+
return out, [scope.get_current_value(v) for v in tracked]
|
|
105
|
+
|
|
106
|
+
(loss, updates), grads = jax.value_and_grad(compute, has_aux=True)([v.value for v in variables])
|
|
107
|
+
for v, value in zip(tracked, updates):
|
|
108
|
+
if value is not None:
|
|
109
|
+
v.assign(value)
|
|
110
|
+
else:
|
|
111
|
+
raise NotImplementedError(f"gradient_step does not support the {backend} backend.")
|
|
112
|
+
|
|
113
|
+
grads = [ops.zeros_like(v) if g is None else g for g, v in zip(grads, variables)]
|
|
114
|
+
optimizer.apply_gradients(zip(grads, variables))
|
|
115
|
+
return float(ops.convert_to_numpy(loss))
|
|
@@ -0,0 +1,166 @@
|
|
|
1
|
+
from k3_node.transforms.base_transform import BaseTransform, functional_transform
|
|
2
|
+
from k3_node.transforms.compose import Compose, ComposeFilters
|
|
3
|
+
|
|
4
|
+
from k3_node.transforms.general import (
|
|
5
|
+
ToDevice,
|
|
6
|
+
ToSparseTensor,
|
|
7
|
+
Constant,
|
|
8
|
+
NormalizeFeatures,
|
|
9
|
+
SVDFeatureReduction,
|
|
10
|
+
RemoveTrainingClasses,
|
|
11
|
+
RandomNodeSplit,
|
|
12
|
+
RandomLinkSplit,
|
|
13
|
+
AttentiveFPFeatures,
|
|
14
|
+
CompleteGraph,
|
|
15
|
+
NodePropertySplit,
|
|
16
|
+
IndexToMask,
|
|
17
|
+
MaskToIndex,
|
|
18
|
+
Pad,
|
|
19
|
+
Padding,
|
|
20
|
+
UniformPadding,
|
|
21
|
+
MappingPadding,
|
|
22
|
+
)
|
|
23
|
+
|
|
24
|
+
from k3_node.transforms.graph import (
|
|
25
|
+
ToUndirected,
|
|
26
|
+
OneHotDegree,
|
|
27
|
+
TargetIndegree,
|
|
28
|
+
LocalDegreeProfile,
|
|
29
|
+
AddSelfLoops,
|
|
30
|
+
AddRemainingSelfLoops,
|
|
31
|
+
RemoveSelfLoops,
|
|
32
|
+
RemoveIsolatedNodes,
|
|
33
|
+
RemoveDuplicatedEdges,
|
|
34
|
+
KNNGraph,
|
|
35
|
+
RadiusGraph,
|
|
36
|
+
ToDense,
|
|
37
|
+
TwoHop,
|
|
38
|
+
LineGraph,
|
|
39
|
+
LaplacianLambdaMax,
|
|
40
|
+
GDC,
|
|
41
|
+
SIGN,
|
|
42
|
+
GCNNorm,
|
|
43
|
+
AddMetaPaths,
|
|
44
|
+
AddRandomMetaPaths,
|
|
45
|
+
RootedEgoNets,
|
|
46
|
+
RootedRWSubgraph,
|
|
47
|
+
LargestConnectedComponents,
|
|
48
|
+
VirtualNode,
|
|
49
|
+
AddLaplacianEigenvectorPE,
|
|
50
|
+
AddRandomWalkPE,
|
|
51
|
+
AddGPSE,
|
|
52
|
+
FeaturePropagation,
|
|
53
|
+
HalfHop,
|
|
54
|
+
)
|
|
55
|
+
|
|
56
|
+
from k3_node.transforms.spatial import (
|
|
57
|
+
Distance,
|
|
58
|
+
Cartesian,
|
|
59
|
+
LocalCartesian,
|
|
60
|
+
Polar,
|
|
61
|
+
Spherical,
|
|
62
|
+
PointPairFeatures,
|
|
63
|
+
Center,
|
|
64
|
+
NormalizeRotation,
|
|
65
|
+
NormalizeScale,
|
|
66
|
+
RandomJitter,
|
|
67
|
+
RandomFlip,
|
|
68
|
+
LinearTransformation,
|
|
69
|
+
RandomScale,
|
|
70
|
+
RandomRotate,
|
|
71
|
+
RandomShear,
|
|
72
|
+
FaceToEdge,
|
|
73
|
+
SamplePoints,
|
|
74
|
+
FixedPoints,
|
|
75
|
+
GenerateMeshNormals,
|
|
76
|
+
Delaunay,
|
|
77
|
+
ToSLIC,
|
|
78
|
+
GridSampling,
|
|
79
|
+
RandomTranslate,
|
|
80
|
+
)
|
|
81
|
+
|
|
82
|
+
general_transforms = [
|
|
83
|
+
'BaseTransform',
|
|
84
|
+
'Compose',
|
|
85
|
+
'ComposeFilters',
|
|
86
|
+
'ToDevice',
|
|
87
|
+
'ToSparseTensor',
|
|
88
|
+
'Constant',
|
|
89
|
+
'NormalizeFeatures',
|
|
90
|
+
'SVDFeatureReduction',
|
|
91
|
+
'RemoveTrainingClasses',
|
|
92
|
+
'RandomNodeSplit',
|
|
93
|
+
'RandomLinkSplit',
|
|
94
|
+
'AttentiveFPFeatures',
|
|
95
|
+
'CompleteGraph',
|
|
96
|
+
'NodePropertySplit',
|
|
97
|
+
'IndexToMask',
|
|
98
|
+
'MaskToIndex',
|
|
99
|
+
'Pad',
|
|
100
|
+
]
|
|
101
|
+
|
|
102
|
+
graph_transforms = [
|
|
103
|
+
'ToUndirected',
|
|
104
|
+
'OneHotDegree',
|
|
105
|
+
'TargetIndegree',
|
|
106
|
+
'LocalDegreeProfile',
|
|
107
|
+
'AddSelfLoops',
|
|
108
|
+
'AddRemainingSelfLoops',
|
|
109
|
+
'RemoveSelfLoops',
|
|
110
|
+
'RemoveIsolatedNodes',
|
|
111
|
+
'RemoveDuplicatedEdges',
|
|
112
|
+
'KNNGraph',
|
|
113
|
+
'RadiusGraph',
|
|
114
|
+
'ToDense',
|
|
115
|
+
'TwoHop',
|
|
116
|
+
'LineGraph',
|
|
117
|
+
'LaplacianLambdaMax',
|
|
118
|
+
'GDC',
|
|
119
|
+
'SIGN',
|
|
120
|
+
'GCNNorm',
|
|
121
|
+
'AddMetaPaths',
|
|
122
|
+
'AddRandomMetaPaths',
|
|
123
|
+
'RootedEgoNets',
|
|
124
|
+
'RootedRWSubgraph',
|
|
125
|
+
'LargestConnectedComponents',
|
|
126
|
+
'VirtualNode',
|
|
127
|
+
'AddLaplacianEigenvectorPE',
|
|
128
|
+
'AddRandomWalkPE',
|
|
129
|
+
'AddGPSE',
|
|
130
|
+
'FeaturePropagation',
|
|
131
|
+
'HalfHop',
|
|
132
|
+
]
|
|
133
|
+
|
|
134
|
+
vision_transforms = [
|
|
135
|
+
'Distance',
|
|
136
|
+
'Cartesian',
|
|
137
|
+
'LocalCartesian',
|
|
138
|
+
'Polar',
|
|
139
|
+
'Spherical',
|
|
140
|
+
'PointPairFeatures',
|
|
141
|
+
'Center',
|
|
142
|
+
'NormalizeRotation',
|
|
143
|
+
'NormalizeScale',
|
|
144
|
+
'RandomJitter',
|
|
145
|
+
'RandomFlip',
|
|
146
|
+
'LinearTransformation',
|
|
147
|
+
'RandomScale',
|
|
148
|
+
'RandomRotate',
|
|
149
|
+
'RandomShear',
|
|
150
|
+
'FaceToEdge',
|
|
151
|
+
'SamplePoints',
|
|
152
|
+
'FixedPoints',
|
|
153
|
+
'GenerateMeshNormals',
|
|
154
|
+
'Delaunay',
|
|
155
|
+
'ToSLIC',
|
|
156
|
+
'GridSampling',
|
|
157
|
+
]
|
|
158
|
+
|
|
159
|
+
__all__ = general_transforms + graph_transforms + vision_transforms + [
|
|
160
|
+
'RandomTranslate',
|
|
161
|
+
'Padding',
|
|
162
|
+
'UniformPadding',
|
|
163
|
+
'MappingPadding',
|
|
164
|
+
'functional_transform',
|
|
165
|
+
]
|
|
166
|
+
|
|
@@ -0,0 +1,32 @@
|
|
|
1
|
+
import copy
|
|
2
|
+
from abc import ABC, abstractmethod
|
|
3
|
+
from typing import Any, Callable
|
|
4
|
+
|
|
5
|
+
|
|
6
|
+
class BaseTransform(ABC):
|
|
7
|
+
r"""An abstract base class for writing transforms.
|
|
8
|
+
|
|
9
|
+
Transforms are a general way to modify and customize
|
|
10
|
+
:class:`~k3_node.data.Data` or :class:`~k3_node.data.HeteroData` objects.
|
|
11
|
+
"""
|
|
12
|
+
|
|
13
|
+
def __call__(self, data: Any) -> Any:
|
|
14
|
+
# Shallow-copy the data to prevent in-place modification of caller's object
|
|
15
|
+
return self.forward(copy.copy(data))
|
|
16
|
+
|
|
17
|
+
@abstractmethod
|
|
18
|
+
def forward(self, data: Any) -> Any:
|
|
19
|
+
pass
|
|
20
|
+
|
|
21
|
+
def __repr__(self) -> str:
|
|
22
|
+
return f"{self.__class__.__name__}()"
|
|
23
|
+
|
|
24
|
+
|
|
25
|
+
def functional_transform(name: str) -> Callable:
|
|
26
|
+
r"""Decorator for functional transforms."""
|
|
27
|
+
|
|
28
|
+
def wrapper(cls: Any) -> Any:
|
|
29
|
+
return cls
|
|
30
|
+
|
|
31
|
+
return wrapper
|
|
32
|
+
|
|
@@ -0,0 +1,58 @@
|
|
|
1
|
+
from typing import Callable, List, Union
|
|
2
|
+
|
|
3
|
+
from k3_node.data import Data, HeteroData
|
|
4
|
+
from k3_node.transforms.base_transform import BaseTransform
|
|
5
|
+
|
|
6
|
+
|
|
7
|
+
class Compose(BaseTransform):
|
|
8
|
+
r"""Composes several transforms together.
|
|
9
|
+
|
|
10
|
+
Args:
|
|
11
|
+
transforms (List[Callable]): List of transforms to compose.
|
|
12
|
+
"""
|
|
13
|
+
|
|
14
|
+
def __init__(self, transforms: List[Callable]):
|
|
15
|
+
self.transforms = transforms
|
|
16
|
+
|
|
17
|
+
def forward(
|
|
18
|
+
self,
|
|
19
|
+
data: Union[Data, HeteroData],
|
|
20
|
+
) -> Union[Data, HeteroData]:
|
|
21
|
+
for transform in self.transforms:
|
|
22
|
+
if isinstance(data, (list, tuple)):
|
|
23
|
+
data = [transform(d) for d in data]
|
|
24
|
+
else:
|
|
25
|
+
data = transform(data)
|
|
26
|
+
return data
|
|
27
|
+
|
|
28
|
+
def __repr__(self) -> str:
|
|
29
|
+
args = [f" {transform}" for transform in self.transforms]
|
|
30
|
+
return "{}([\n{}\n])".format(self.__class__.__name__, ",\n".join(args))
|
|
31
|
+
|
|
32
|
+
|
|
33
|
+
class ComposeFilters:
|
|
34
|
+
r"""Composes several filters together.
|
|
35
|
+
|
|
36
|
+
Args:
|
|
37
|
+
filters (List[Callable]): List of filters to compose.
|
|
38
|
+
"""
|
|
39
|
+
|
|
40
|
+
def __init__(self, filters: List[Callable]):
|
|
41
|
+
self.filters = filters
|
|
42
|
+
|
|
43
|
+
def __call__(
|
|
44
|
+
self,
|
|
45
|
+
data: Union[Data, HeteroData],
|
|
46
|
+
) -> bool:
|
|
47
|
+
for filter_fn in self.filters:
|
|
48
|
+
if isinstance(data, (list, tuple)):
|
|
49
|
+
if not all([filter_fn(d) for d in data]):
|
|
50
|
+
return False
|
|
51
|
+
elif not filter_fn(data):
|
|
52
|
+
return False
|
|
53
|
+
return True
|
|
54
|
+
|
|
55
|
+
def __repr__(self) -> str:
|
|
56
|
+
args = [f" {filter_fn}" for filter_fn in self.filters]
|
|
57
|
+
return "{}([\n{}\n])".format(self.__class__.__name__, ",\n".join(args))
|
|
58
|
+
|