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
k3_node/data/separate.py
ADDED
|
@@ -0,0 +1,115 @@
|
|
|
1
|
+
from typing import Any, Type, TypeVar
|
|
2
|
+
|
|
3
|
+
from keras import ops
|
|
4
|
+
|
|
5
|
+
from k3_node.data.data import BaseData, Data
|
|
6
|
+
from k3_node.data.hetero_data import HeteroData
|
|
7
|
+
from k3_node.data.storage import BaseStorage, get_shape, is_tensor_like
|
|
8
|
+
|
|
9
|
+
T = TypeVar("T")
|
|
10
|
+
|
|
11
|
+
|
|
12
|
+
def narrow(tensor, dim: int, start: int, length: int):
|
|
13
|
+
shape = list(get_shape(tensor))
|
|
14
|
+
if dim < 0:
|
|
15
|
+
dim = len(shape) + dim
|
|
16
|
+
slices = [slice(None)] * len(shape)
|
|
17
|
+
slices[dim] = slice(start, start + length)
|
|
18
|
+
return tensor[tuple(slices)]
|
|
19
|
+
|
|
20
|
+
|
|
21
|
+
def separate(
|
|
22
|
+
cls: Type[T],
|
|
23
|
+
batch: Any,
|
|
24
|
+
idx: int,
|
|
25
|
+
slice_dict: Any,
|
|
26
|
+
inc_dict: Any = None,
|
|
27
|
+
decrement: bool = True,
|
|
28
|
+
) -> T:
|
|
29
|
+
is_hetero = isinstance(batch, HeteroData)
|
|
30
|
+
data = HeteroData() if is_hetero else Data()
|
|
31
|
+
|
|
32
|
+
if not is_hetero:
|
|
33
|
+
batch_store = batch._store
|
|
34
|
+
data_store = data._store
|
|
35
|
+
|
|
36
|
+
for attr in slice_dict.keys():
|
|
37
|
+
if attr not in batch_store:
|
|
38
|
+
continue
|
|
39
|
+
slices = slice_dict[attr]
|
|
40
|
+
incs = inc_dict[attr] if (decrement and inc_dict is not None and attr in inc_dict) else None
|
|
41
|
+
|
|
42
|
+
val = batch_store[attr]
|
|
43
|
+
if is_tensor_like(val):
|
|
44
|
+
start = int(slices[idx])
|
|
45
|
+
end = int(slices[idx + 1])
|
|
46
|
+
length = end - start
|
|
47
|
+
cat_dim = batch.__cat_dim__(attr, val, batch_store)
|
|
48
|
+
sub_val = narrow(val, cat_dim, start, length)
|
|
49
|
+
if decrement and incs is not None and idx < len(incs):
|
|
50
|
+
inc_val = incs[idx]
|
|
51
|
+
if inc_val != 0:
|
|
52
|
+
inc_t = ops.convert_to_tensor(inc_val, dtype=sub_val.dtype)
|
|
53
|
+
sub_val = sub_val - inc_t
|
|
54
|
+
data_store[attr] = sub_val
|
|
55
|
+
elif isinstance(val, (list, tuple)) and len(val) > idx:
|
|
56
|
+
data_store[attr] = val[idx]
|
|
57
|
+
else:
|
|
58
|
+
data_store[attr] = val
|
|
59
|
+
|
|
60
|
+
if hasattr(batch_store, "_num_nodes") and idx < len(batch_store._num_nodes):
|
|
61
|
+
data_store.num_nodes = batch_store._num_nodes[idx]
|
|
62
|
+
|
|
63
|
+
else:
|
|
64
|
+
# Heterogeneous separate
|
|
65
|
+
for node_type in batch.node_types:
|
|
66
|
+
batch_store = batch[node_type]
|
|
67
|
+
data_store = data[node_type]
|
|
68
|
+
store_slice_dict = slice_dict.get(node_type, {})
|
|
69
|
+
store_inc_dict = inc_dict.get(node_type, {}) if decrement and inc_dict else {}
|
|
70
|
+
|
|
71
|
+
for attr in store_slice_dict.keys():
|
|
72
|
+
if attr not in batch_store:
|
|
73
|
+
continue
|
|
74
|
+
slices = store_slice_dict[attr]
|
|
75
|
+
val = batch_store[attr]
|
|
76
|
+
if is_tensor_like(val):
|
|
77
|
+
start = int(slices[idx])
|
|
78
|
+
end = int(slices[idx + 1])
|
|
79
|
+
length = end - start
|
|
80
|
+
cat_dim = batch.__cat_dim__(attr, val, batch_store)
|
|
81
|
+
sub_val = narrow(val, cat_dim, start, length)
|
|
82
|
+
data_store[attr] = sub_val
|
|
83
|
+
elif isinstance(val, (list, tuple)) and len(val) > idx:
|
|
84
|
+
data_store[attr] = val[idx]
|
|
85
|
+
|
|
86
|
+
for edge_type in batch.edge_types:
|
|
87
|
+
batch_store = batch[edge_type]
|
|
88
|
+
data_store = data[edge_type]
|
|
89
|
+
store_slice_dict = slice_dict.get(edge_type, {})
|
|
90
|
+
store_inc_dict = inc_dict.get(edge_type, {}) if decrement and inc_dict else {}
|
|
91
|
+
|
|
92
|
+
for attr in store_slice_dict.keys():
|
|
93
|
+
if attr not in batch_store:
|
|
94
|
+
continue
|
|
95
|
+
slices = store_slice_dict[attr]
|
|
96
|
+
incs = store_inc_dict.get(attr) if decrement else None
|
|
97
|
+
val = batch_store[attr]
|
|
98
|
+
if is_tensor_like(val):
|
|
99
|
+
start = int(slices[idx])
|
|
100
|
+
end = int(slices[idx + 1])
|
|
101
|
+
length = end - start
|
|
102
|
+
cat_dim = batch.__cat_dim__(attr, val, batch_store)
|
|
103
|
+
sub_val = narrow(val, cat_dim, start, length)
|
|
104
|
+
if decrement and incs is not None and idx < len(incs):
|
|
105
|
+
inc_val = incs[idx]
|
|
106
|
+
if is_tensor_like(inc_val) or (isinstance(inc_val, np.ndarray) and np.any(inc_val != 0)):
|
|
107
|
+
if hasattr(inc_val, "ndim") and inc_val.ndim == 1:
|
|
108
|
+
inc_val = inc_val[:, None]
|
|
109
|
+
inc_t = ops.convert_to_tensor(inc_val, dtype=sub_val.dtype)
|
|
110
|
+
sub_val = sub_val - inc_t
|
|
111
|
+
data_store[attr] = sub_val
|
|
112
|
+
elif isinstance(val, (list, tuple)) and len(val) > idx:
|
|
113
|
+
data_store[attr] = val[idx]
|
|
114
|
+
|
|
115
|
+
return data
|
k3_node/data/storage.py
ADDED
|
@@ -0,0 +1,593 @@
|
|
|
1
|
+
import copy
|
|
2
|
+
import weakref
|
|
3
|
+
from collections import defaultdict
|
|
4
|
+
from collections.abc import Mapping, MutableMapping, Sequence
|
|
5
|
+
from enum import Enum
|
|
6
|
+
from typing import Any, Callable, Dict, Iterator, List, Optional, Set, Tuple, Union
|
|
7
|
+
|
|
8
|
+
import numpy as np
|
|
9
|
+
from keras import ops
|
|
10
|
+
|
|
11
|
+
from k3_node.data.view import ItemsView, KeysView, ValuesView
|
|
12
|
+
from k3_node.utils.graph import coalesce, contains_isolated_nodes, has_self_loops, is_undirected
|
|
13
|
+
|
|
14
|
+
N_KEYS = {"x", "feat", "pos", "batch", "node_type", "n_id", "tf"}
|
|
15
|
+
E_KEYS = {"edge_index", "edge_weight", "edge_attr", "edge_type", "e_id"}
|
|
16
|
+
|
|
17
|
+
|
|
18
|
+
def is_tensor_like(x: Any) -> bool:
|
|
19
|
+
if isinstance(x, np.ndarray):
|
|
20
|
+
return True
|
|
21
|
+
return hasattr(x, "shape") and hasattr(x, "dtype")
|
|
22
|
+
|
|
23
|
+
|
|
24
|
+
def to_numpy(x: Any) -> Any:
|
|
25
|
+
if x is None:
|
|
26
|
+
return None
|
|
27
|
+
if hasattr(x, "detach"):
|
|
28
|
+
x = x.detach()
|
|
29
|
+
if hasattr(x, "numpy"):
|
|
30
|
+
try:
|
|
31
|
+
return x.numpy()
|
|
32
|
+
except TypeError:
|
|
33
|
+
if hasattr(x, "cpu"):
|
|
34
|
+
return x.cpu().numpy()
|
|
35
|
+
raise
|
|
36
|
+
if hasattr(x, "cpu"):
|
|
37
|
+
x = x.cpu()
|
|
38
|
+
if hasattr(x, "numpy"):
|
|
39
|
+
return x.numpy()
|
|
40
|
+
if hasattr(x, "_numpy"):
|
|
41
|
+
return x._numpy()
|
|
42
|
+
return np.asarray(x)
|
|
43
|
+
|
|
44
|
+
|
|
45
|
+
def get_shape(x: Any) -> Tuple[int, ...]:
|
|
46
|
+
if hasattr(x, "shape"):
|
|
47
|
+
return tuple(int(s) if s is not None else 0 for s in x.shape)
|
|
48
|
+
if isinstance(x, (list, tuple)):
|
|
49
|
+
return (len(x),)
|
|
50
|
+
return ()
|
|
51
|
+
|
|
52
|
+
|
|
53
|
+
def recursive_apply(data: Any, func: Callable) -> Any:
|
|
54
|
+
if is_tensor_like(data):
|
|
55
|
+
return func(data)
|
|
56
|
+
elif isinstance(data, tuple) and hasattr(data, "_fields"):
|
|
57
|
+
return type(data)(*(recursive_apply(d, func) for d in data))
|
|
58
|
+
elif isinstance(data, Sequence) and not isinstance(data, str):
|
|
59
|
+
return [recursive_apply(d, func) for d in data]
|
|
60
|
+
elif isinstance(data, Mapping):
|
|
61
|
+
return {key: recursive_apply(data[key], func) for key in data}
|
|
62
|
+
else:
|
|
63
|
+
try:
|
|
64
|
+
return func(data)
|
|
65
|
+
except Exception:
|
|
66
|
+
return data
|
|
67
|
+
|
|
68
|
+
|
|
69
|
+
def recursive_apply_(data: Any, func: Callable) -> Any:
|
|
70
|
+
if is_tensor_like(data):
|
|
71
|
+
try:
|
|
72
|
+
func(data)
|
|
73
|
+
except Exception:
|
|
74
|
+
pass
|
|
75
|
+
elif isinstance(data, tuple) and hasattr(data, "_fields"):
|
|
76
|
+
for value in data:
|
|
77
|
+
recursive_apply_(value, func)
|
|
78
|
+
elif isinstance(data, Sequence) and not isinstance(data, str):
|
|
79
|
+
for value in data:
|
|
80
|
+
recursive_apply_(value, func)
|
|
81
|
+
elif isinstance(data, Mapping):
|
|
82
|
+
for value in data.values():
|
|
83
|
+
recursive_apply_(value, func)
|
|
84
|
+
else:
|
|
85
|
+
try:
|
|
86
|
+
func(data)
|
|
87
|
+
except Exception:
|
|
88
|
+
pass
|
|
89
|
+
|
|
90
|
+
|
|
91
|
+
class AttrType(Enum):
|
|
92
|
+
NODE = "NODE"
|
|
93
|
+
EDGE = "EDGE"
|
|
94
|
+
OTHER = "OTHER"
|
|
95
|
+
|
|
96
|
+
|
|
97
|
+
class BaseStorage(MutableMapping):
|
|
98
|
+
def __init__(self, _mapping: Optional[Dict[str, Any]] = None, **kwargs: Any) -> None:
|
|
99
|
+
super().__init__()
|
|
100
|
+
self._mapping: Dict[str, Any] = {}
|
|
101
|
+
for key, value in (_mapping or {}).items():
|
|
102
|
+
setattr(self, key, value)
|
|
103
|
+
for key, value in kwargs.items():
|
|
104
|
+
setattr(self, key, value)
|
|
105
|
+
|
|
106
|
+
@property
|
|
107
|
+
def _key(self) -> Any:
|
|
108
|
+
return None
|
|
109
|
+
|
|
110
|
+
def _pop_cache(self, key: str) -> None:
|
|
111
|
+
for cache in getattr(self, "_cached_attr", {}).values():
|
|
112
|
+
cache.discard(key)
|
|
113
|
+
|
|
114
|
+
def __len__(self) -> int:
|
|
115
|
+
return len(self._mapping)
|
|
116
|
+
|
|
117
|
+
def __getattr__(self, key: str) -> Any:
|
|
118
|
+
if key == "_mapping":
|
|
119
|
+
self._mapping = {}
|
|
120
|
+
return self._mapping
|
|
121
|
+
try:
|
|
122
|
+
return self[key]
|
|
123
|
+
except KeyError:
|
|
124
|
+
raise AttributeError(f"'{self.__class__.__name__}' object has no attribute '{key}'") from None
|
|
125
|
+
|
|
126
|
+
def __setattr__(self, key: str, value: Any) -> None:
|
|
127
|
+
propobj = getattr(self.__class__, key, None)
|
|
128
|
+
if propobj is not None and getattr(propobj, "fset", None) is not None:
|
|
129
|
+
propobj.fset(self, value)
|
|
130
|
+
elif key == "_parent":
|
|
131
|
+
self.__dict__[key] = weakref.ref(value) if value is not None else None
|
|
132
|
+
elif key[:1] == "_":
|
|
133
|
+
self.__dict__[key] = value
|
|
134
|
+
else:
|
|
135
|
+
self[key] = value
|
|
136
|
+
|
|
137
|
+
def __delattr__(self, key: str) -> None:
|
|
138
|
+
if key[:1] == "_":
|
|
139
|
+
if key in self.__dict__:
|
|
140
|
+
del self.__dict__[key]
|
|
141
|
+
else:
|
|
142
|
+
del self[key]
|
|
143
|
+
|
|
144
|
+
def __getitem__(self, key: str) -> Any:
|
|
145
|
+
return self._mapping[key]
|
|
146
|
+
|
|
147
|
+
def __setitem__(self, key: str, value: Any) -> None:
|
|
148
|
+
self._pop_cache(key)
|
|
149
|
+
if value is None and key in self._mapping:
|
|
150
|
+
del self._mapping[key]
|
|
151
|
+
elif value is not None:
|
|
152
|
+
self._mapping[key] = value
|
|
153
|
+
|
|
154
|
+
def __delitem__(self, key: str) -> None:
|
|
155
|
+
if key in self._mapping:
|
|
156
|
+
self._pop_cache(key)
|
|
157
|
+
del self._mapping[key]
|
|
158
|
+
|
|
159
|
+
def __iter__(self) -> Iterator[Any]:
|
|
160
|
+
return iter(self._mapping)
|
|
161
|
+
|
|
162
|
+
def __copy__(self):
|
|
163
|
+
out = self.__class__.__new__(self.__class__)
|
|
164
|
+
for key, value in self.__dict__.items():
|
|
165
|
+
if key != "_cached_attr":
|
|
166
|
+
out.__dict__[key] = value
|
|
167
|
+
out._mapping = copy.copy(out._mapping)
|
|
168
|
+
return out
|
|
169
|
+
|
|
170
|
+
def __deepcopy__(self, memo=None):
|
|
171
|
+
out = self.__class__.__new__(self.__class__)
|
|
172
|
+
for key, value in self.__dict__.items():
|
|
173
|
+
if key == "_parent":
|
|
174
|
+
out.__dict__[key] = self.__dict__[key]
|
|
175
|
+
elif key != "_cached_attr":
|
|
176
|
+
out.__dict__[key] = copy.deepcopy(value, memo)
|
|
177
|
+
out._mapping = copy.deepcopy(out._mapping, memo)
|
|
178
|
+
return out
|
|
179
|
+
|
|
180
|
+
def __getstate__(self) -> Dict[str, Any]:
|
|
181
|
+
out = self.__dict__.copy()
|
|
182
|
+
_parent = out.get("_parent", None)
|
|
183
|
+
if _parent is not None:
|
|
184
|
+
out["_parent"] = _parent()
|
|
185
|
+
return out
|
|
186
|
+
|
|
187
|
+
def __setstate__(self, mapping: Dict[str, Any]) -> None:
|
|
188
|
+
for key, value in mapping.items():
|
|
189
|
+
self.__dict__[key] = value
|
|
190
|
+
_parent = self.__dict__.get("_parent", None)
|
|
191
|
+
if _parent is not None:
|
|
192
|
+
self.__dict__["_parent"] = weakref.ref(_parent)
|
|
193
|
+
|
|
194
|
+
def __repr__(self) -> str:
|
|
195
|
+
return repr(self._mapping)
|
|
196
|
+
|
|
197
|
+
def _parent(self):
|
|
198
|
+
parent_ref = self.__dict__.get("_parent", None)
|
|
199
|
+
return parent_ref() if parent_ref is not None else None
|
|
200
|
+
|
|
201
|
+
def keys(self, *args: str) -> KeysView:
|
|
202
|
+
return KeysView(self._mapping, *args)
|
|
203
|
+
|
|
204
|
+
def values(self, *args: str) -> ValuesView:
|
|
205
|
+
return ValuesView(self._mapping, *args)
|
|
206
|
+
|
|
207
|
+
def items(self, *args: str) -> ItemsView:
|
|
208
|
+
return ItemsView(self._mapping, *args)
|
|
209
|
+
|
|
210
|
+
def to_dict(self) -> Dict[str, Any]:
|
|
211
|
+
return copy.copy(self._mapping)
|
|
212
|
+
|
|
213
|
+
def apply_(self, func: Callable, *args: str):
|
|
214
|
+
for value in self.values(*args):
|
|
215
|
+
recursive_apply_(value, func)
|
|
216
|
+
return self
|
|
217
|
+
|
|
218
|
+
def apply(self, func: Callable, *args: str):
|
|
219
|
+
for key, value in self.items(*args):
|
|
220
|
+
self[key] = recursive_apply(value, func)
|
|
221
|
+
return self
|
|
222
|
+
|
|
223
|
+
def to(self, *args, **kwargs):
|
|
224
|
+
def _to(x):
|
|
225
|
+
if hasattr(x, "to"):
|
|
226
|
+
return x.to(*args, **kwargs)
|
|
227
|
+
return x
|
|
228
|
+
|
|
229
|
+
return self.apply(_to)
|
|
230
|
+
|
|
231
|
+
def to_backend(self, backend: Optional[str] = None):
|
|
232
|
+
for key in list(self.keys()):
|
|
233
|
+
val = self[key]
|
|
234
|
+
if is_tensor_like(val):
|
|
235
|
+
v_np = to_numpy(val)
|
|
236
|
+
if v_np.dtype == np.bool_:
|
|
237
|
+
self[key] = ops.convert_to_tensor(v_np, dtype="bool")
|
|
238
|
+
elif np.issubdtype(v_np.dtype, np.integer):
|
|
239
|
+
self[key] = ops.convert_to_tensor(v_np, dtype="int64")
|
|
240
|
+
elif np.issubdtype(v_np.dtype, np.floating):
|
|
241
|
+
self[key] = ops.convert_to_tensor(v_np, dtype="float32")
|
|
242
|
+
else:
|
|
243
|
+
self[key] = ops.convert_to_tensor(v_np)
|
|
244
|
+
return self
|
|
245
|
+
|
|
246
|
+
def cpu(self):
|
|
247
|
+
return self.to("cpu")
|
|
248
|
+
|
|
249
|
+
def cuda(self):
|
|
250
|
+
return self.to("cuda")
|
|
251
|
+
|
|
252
|
+
def requires_grad_(self, *args: str):
|
|
253
|
+
def _req(x):
|
|
254
|
+
if hasattr(x, "requires_grad_"):
|
|
255
|
+
x.requires_grad_()
|
|
256
|
+
elif hasattr(x, "requires_grad"):
|
|
257
|
+
x.requires_grad = True
|
|
258
|
+
|
|
259
|
+
return self.apply_(_req, *args)
|
|
260
|
+
|
|
261
|
+
def contiguous(self, *args: str):
|
|
262
|
+
def _cont(x):
|
|
263
|
+
if hasattr(x, "contiguous"):
|
|
264
|
+
return x.contiguous()
|
|
265
|
+
return x
|
|
266
|
+
|
|
267
|
+
return self.apply(_cont, *args)
|
|
268
|
+
|
|
269
|
+
|
|
270
|
+
class NodeStorage(BaseStorage):
|
|
271
|
+
@property
|
|
272
|
+
def _key(self) -> Any:
|
|
273
|
+
return self.__dict__.get("_key", None)
|
|
274
|
+
|
|
275
|
+
@property
|
|
276
|
+
def num_nodes(self) -> int:
|
|
277
|
+
if "num_nodes" in self:
|
|
278
|
+
return int(self["num_nodes"])
|
|
279
|
+
parent = self._parent()
|
|
280
|
+
for key, value in self.items():
|
|
281
|
+
if is_tensor_like(value) and key in N_KEYS:
|
|
282
|
+
cat_dim = parent.__cat_dim__(key, value, self) if parent is not None else 0
|
|
283
|
+
return get_shape(value)[cat_dim]
|
|
284
|
+
for key, value in self.items():
|
|
285
|
+
if is_tensor_like(value) and "node" in key:
|
|
286
|
+
cat_dim = parent.__cat_dim__(key, value, self) if parent is not None else 0
|
|
287
|
+
return get_shape(value)[cat_dim]
|
|
288
|
+
edge_index = self.get("edge_index")
|
|
289
|
+
if is_tensor_like(edge_index) and get_shape(edge_index)[-1] > 0:
|
|
290
|
+
# As in PyG: without node-level attributes, infer the count from the edges
|
|
291
|
+
return int(np.asarray(to_numpy(edge_index)).max()) + 1
|
|
292
|
+
return 0
|
|
293
|
+
|
|
294
|
+
@property
|
|
295
|
+
def num_node_features(self) -> int:
|
|
296
|
+
x = self.get("x")
|
|
297
|
+
if x is not None and is_tensor_like(x):
|
|
298
|
+
shape = get_shape(x)
|
|
299
|
+
return 1 if len(shape) == 1 else shape[-1]
|
|
300
|
+
return 0
|
|
301
|
+
|
|
302
|
+
@property
|
|
303
|
+
def num_features(self) -> int:
|
|
304
|
+
return self.num_node_features
|
|
305
|
+
|
|
306
|
+
def is_node_attr(self, key: str) -> bool:
|
|
307
|
+
if "_cached_attr" not in self.__dict__:
|
|
308
|
+
self._cached_attr: Dict[AttrType, Set[str]] = defaultdict(set)
|
|
309
|
+
|
|
310
|
+
if key in self._cached_attr[AttrType.NODE]:
|
|
311
|
+
return True
|
|
312
|
+
if key in self._cached_attr[AttrType.OTHER]:
|
|
313
|
+
return False
|
|
314
|
+
|
|
315
|
+
value = self.get(key)
|
|
316
|
+
if value is None:
|
|
317
|
+
return False
|
|
318
|
+
|
|
319
|
+
if isinstance(value, (list, tuple)) and len(value) == self.num_nodes:
|
|
320
|
+
self._cached_attr[AttrType.NODE].add(key)
|
|
321
|
+
return True
|
|
322
|
+
|
|
323
|
+
if not is_tensor_like(value):
|
|
324
|
+
self._cached_attr[AttrType.OTHER].add(key)
|
|
325
|
+
return False
|
|
326
|
+
|
|
327
|
+
shape = get_shape(value)
|
|
328
|
+
if len(shape) == 0:
|
|
329
|
+
self._cached_attr[AttrType.OTHER].add(key)
|
|
330
|
+
return False
|
|
331
|
+
|
|
332
|
+
parent = self._parent()
|
|
333
|
+
cat_dim = parent.__cat_dim__(key, value, self) if parent is not None else 0
|
|
334
|
+
if shape[cat_dim] != self.num_nodes:
|
|
335
|
+
self._cached_attr[AttrType.OTHER].add(key)
|
|
336
|
+
return False
|
|
337
|
+
|
|
338
|
+
self._cached_attr[AttrType.NODE].add(key)
|
|
339
|
+
return True
|
|
340
|
+
|
|
341
|
+
def is_edge_attr(self, key: str) -> bool:
|
|
342
|
+
return False
|
|
343
|
+
|
|
344
|
+
def node_attrs(self) -> List[str]:
|
|
345
|
+
return [key for key in self.keys() if self.is_node_attr(key)]
|
|
346
|
+
|
|
347
|
+
|
|
348
|
+
class EdgeStorage(BaseStorage):
|
|
349
|
+
@property
|
|
350
|
+
def _key(self) -> Any:
|
|
351
|
+
return self.__dict__.get("_key", None)
|
|
352
|
+
|
|
353
|
+
@property
|
|
354
|
+
def edge_index(self):
|
|
355
|
+
if "edge_index" in self:
|
|
356
|
+
return self["edge_index"]
|
|
357
|
+
raise AttributeError(f"'{self.__class__.__name__}' object has no attribute 'edge_index'")
|
|
358
|
+
|
|
359
|
+
@edge_index.setter
|
|
360
|
+
def edge_index(self, edge_index) -> None:
|
|
361
|
+
self["edge_index"] = edge_index
|
|
362
|
+
|
|
363
|
+
@property
|
|
364
|
+
def num_edges(self) -> int:
|
|
365
|
+
if "num_edges" in self:
|
|
366
|
+
return int(self["num_edges"])
|
|
367
|
+
parent = self._parent()
|
|
368
|
+
for key, value in self.items():
|
|
369
|
+
if is_tensor_like(value) and key in E_KEYS:
|
|
370
|
+
cat_dim = parent.__cat_dim__(key, value, self) if parent is not None else -1
|
|
371
|
+
return get_shape(value)[cat_dim]
|
|
372
|
+
for key, value in self.items():
|
|
373
|
+
if is_tensor_like(value) and "edge" in key:
|
|
374
|
+
cat_dim = parent.__cat_dim__(key, value, self) if parent is not None else -1
|
|
375
|
+
return get_shape(value)[cat_dim]
|
|
376
|
+
return 0
|
|
377
|
+
|
|
378
|
+
@property
|
|
379
|
+
def num_edge_features(self) -> int:
|
|
380
|
+
edge_attr = self.get("edge_attr")
|
|
381
|
+
if edge_attr is not None and is_tensor_like(edge_attr):
|
|
382
|
+
shape = get_shape(edge_attr)
|
|
383
|
+
return 1 if len(shape) == 1 else shape[-1]
|
|
384
|
+
return 0
|
|
385
|
+
|
|
386
|
+
@property
|
|
387
|
+
def num_features(self) -> int:
|
|
388
|
+
return self.num_edge_features
|
|
389
|
+
|
|
390
|
+
def size(self, dim: Optional[int] = None) -> Union[Tuple[Optional[int], Optional[int]], Optional[int]]:
|
|
391
|
+
parent = self._parent()
|
|
392
|
+
if self._key is None or parent is None:
|
|
393
|
+
num = self.num_edges
|
|
394
|
+
res = (num, num)
|
|
395
|
+
return res if dim is None else res[dim]
|
|
396
|
+
size = (parent[self._key[0]].num_nodes, parent[self._key[-1]].num_nodes)
|
|
397
|
+
return size if dim is None else size[dim]
|
|
398
|
+
|
|
399
|
+
def is_node_attr(self, key: str) -> bool:
|
|
400
|
+
return False
|
|
401
|
+
|
|
402
|
+
def is_edge_attr(self, key: str) -> bool:
|
|
403
|
+
if "_cached_attr" not in self.__dict__:
|
|
404
|
+
self._cached_attr: Dict[AttrType, Set[str]] = defaultdict(set)
|
|
405
|
+
|
|
406
|
+
if key in self._cached_attr[AttrType.EDGE]:
|
|
407
|
+
return True
|
|
408
|
+
if key in self._cached_attr[AttrType.OTHER]:
|
|
409
|
+
return False
|
|
410
|
+
|
|
411
|
+
value = self.get(key)
|
|
412
|
+
if value is None:
|
|
413
|
+
return False
|
|
414
|
+
|
|
415
|
+
if isinstance(value, (list, tuple)) and len(value) == self.num_edges:
|
|
416
|
+
self._cached_attr[AttrType.EDGE].add(key)
|
|
417
|
+
return True
|
|
418
|
+
|
|
419
|
+
if not is_tensor_like(value):
|
|
420
|
+
self._cached_attr[AttrType.OTHER].add(key)
|
|
421
|
+
return False
|
|
422
|
+
|
|
423
|
+
shape = get_shape(value)
|
|
424
|
+
if len(shape) == 0:
|
|
425
|
+
self._cached_attr[AttrType.OTHER].add(key)
|
|
426
|
+
return False
|
|
427
|
+
|
|
428
|
+
parent = self._parent()
|
|
429
|
+
cat_dim = parent.__cat_dim__(key, value, self) if parent is not None else -1
|
|
430
|
+
if shape[cat_dim] != self.num_edges:
|
|
431
|
+
self._cached_attr[AttrType.OTHER].add(key)
|
|
432
|
+
return False
|
|
433
|
+
|
|
434
|
+
self._cached_attr[AttrType.EDGE].add(key)
|
|
435
|
+
return True
|
|
436
|
+
|
|
437
|
+
def edge_attrs(self) -> List[str]:
|
|
438
|
+
return [key for key in self.keys() if self.is_edge_attr(key)]
|
|
439
|
+
|
|
440
|
+
def is_coalesced(self) -> bool:
|
|
441
|
+
if "edge_index" in self:
|
|
442
|
+
edge_index = self.edge_index
|
|
443
|
+
new_edge_index, _ = coalesce(edge_index)
|
|
444
|
+
orig_np = ops.convert_to_numpy(edge_index)
|
|
445
|
+
new_np = ops.convert_to_numpy(new_edge_index)
|
|
446
|
+
return orig_np.shape == new_np.shape and np.array_equal(orig_np, new_np)
|
|
447
|
+
return True
|
|
448
|
+
|
|
449
|
+
def coalesce(self, reduce: str = "add"):
|
|
450
|
+
if "edge_index" in self:
|
|
451
|
+
self.edge_index, self.edge_attr = coalesce(
|
|
452
|
+
self.edge_index,
|
|
453
|
+
edge_attr=self.get("edge_attr"),
|
|
454
|
+
reduce=reduce,
|
|
455
|
+
)
|
|
456
|
+
return self
|
|
457
|
+
|
|
458
|
+
def has_self_loops(self) -> bool:
|
|
459
|
+
if self.is_bipartite() or "edge_index" not in self:
|
|
460
|
+
return False
|
|
461
|
+
return has_self_loops(self.edge_index)
|
|
462
|
+
|
|
463
|
+
def has_isolated_nodes(self) -> bool:
|
|
464
|
+
if "edge_index" not in self:
|
|
465
|
+
return False
|
|
466
|
+
parent = self._parent()
|
|
467
|
+
num_nodes = parent[self._key[-1]].num_nodes if parent and self._key else None
|
|
468
|
+
return contains_isolated_nodes(self.edge_index, num_nodes=num_nodes)
|
|
469
|
+
|
|
470
|
+
def is_undirected(self) -> bool:
|
|
471
|
+
if self.is_bipartite() or "edge_index" not in self:
|
|
472
|
+
return False
|
|
473
|
+
return is_undirected(self.edge_index, edge_attr=self.get("edge_attr"))
|
|
474
|
+
|
|
475
|
+
def is_directed(self) -> bool:
|
|
476
|
+
return not self.is_undirected()
|
|
477
|
+
|
|
478
|
+
def is_bipartite(self) -> bool:
|
|
479
|
+
return self._key is not None and isinstance(self._key, tuple) and self._key[0] != self._key[-1]
|
|
480
|
+
|
|
481
|
+
|
|
482
|
+
class GlobalStorage(NodeStorage, EdgeStorage):
|
|
483
|
+
@property
|
|
484
|
+
def _key(self) -> Any:
|
|
485
|
+
return None
|
|
486
|
+
|
|
487
|
+
@property
|
|
488
|
+
def num_features(self) -> int:
|
|
489
|
+
return self.num_node_features
|
|
490
|
+
|
|
491
|
+
def size(self, dim: Optional[int] = None) -> Union[Tuple[Optional[int], Optional[int]], Optional[int]]:
|
|
492
|
+
size = (self.num_nodes, self.num_nodes)
|
|
493
|
+
return size if dim is None else size[dim]
|
|
494
|
+
|
|
495
|
+
def is_node_attr(self, key: str) -> bool:
|
|
496
|
+
if "_cached_attr" not in self.__dict__:
|
|
497
|
+
self._cached_attr: Dict[AttrType, Set[str]] = defaultdict(set)
|
|
498
|
+
|
|
499
|
+
if key in self._cached_attr[AttrType.NODE]:
|
|
500
|
+
return True
|
|
501
|
+
if key in self._cached_attr[AttrType.EDGE] or key in self._cached_attr[AttrType.OTHER]:
|
|
502
|
+
return False
|
|
503
|
+
|
|
504
|
+
value = self.get(key)
|
|
505
|
+
if value is None:
|
|
506
|
+
return False
|
|
507
|
+
|
|
508
|
+
if isinstance(value, (list, tuple)) and len(value) == self.num_nodes:
|
|
509
|
+
self._cached_attr[AttrType.NODE].add(key)
|
|
510
|
+
return True
|
|
511
|
+
|
|
512
|
+
if not is_tensor_like(value):
|
|
513
|
+
return False
|
|
514
|
+
|
|
515
|
+
shape = get_shape(value)
|
|
516
|
+
if len(shape) == 0:
|
|
517
|
+
self._cached_attr[AttrType.OTHER].add(key)
|
|
518
|
+
return False
|
|
519
|
+
|
|
520
|
+
parent = self._parent()
|
|
521
|
+
cat_dim = parent.__cat_dim__(key, value, self) if parent is not None else 0
|
|
522
|
+
if not isinstance(cat_dim, int):
|
|
523
|
+
return False
|
|
524
|
+
|
|
525
|
+
num_nodes, num_edges = self.num_nodes, self.num_edges
|
|
526
|
+
|
|
527
|
+
if shape[cat_dim] != num_nodes:
|
|
528
|
+
if shape[cat_dim] == num_edges:
|
|
529
|
+
self._cached_attr[AttrType.EDGE].add(key)
|
|
530
|
+
else:
|
|
531
|
+
self._cached_attr[AttrType.OTHER].add(key)
|
|
532
|
+
return False
|
|
533
|
+
|
|
534
|
+
if num_nodes != num_edges:
|
|
535
|
+
self._cached_attr[AttrType.NODE].add(key)
|
|
536
|
+
return True
|
|
537
|
+
|
|
538
|
+
if "edge" not in key:
|
|
539
|
+
self._cached_attr[AttrType.NODE].add(key)
|
|
540
|
+
return True
|
|
541
|
+
else:
|
|
542
|
+
self._cached_attr[AttrType.EDGE].add(key)
|
|
543
|
+
return False
|
|
544
|
+
|
|
545
|
+
def is_edge_attr(self, key: str) -> bool:
|
|
546
|
+
if "_cached_attr" not in self.__dict__:
|
|
547
|
+
self._cached_attr = defaultdict(set)
|
|
548
|
+
|
|
549
|
+
if key in self._cached_attr[AttrType.EDGE]:
|
|
550
|
+
return True
|
|
551
|
+
if key in self._cached_attr[AttrType.NODE] or key in self._cached_attr[AttrType.OTHER]:
|
|
552
|
+
return False
|
|
553
|
+
|
|
554
|
+
value = self.get(key)
|
|
555
|
+
if value is None:
|
|
556
|
+
return False
|
|
557
|
+
|
|
558
|
+
if isinstance(value, (list, tuple)) and len(value) == self.num_edges:
|
|
559
|
+
self._cached_attr[AttrType.EDGE].add(key)
|
|
560
|
+
return True
|
|
561
|
+
|
|
562
|
+
if not is_tensor_like(value):
|
|
563
|
+
return False
|
|
564
|
+
|
|
565
|
+
shape = get_shape(value)
|
|
566
|
+
if len(shape) == 0:
|
|
567
|
+
self._cached_attr[AttrType.OTHER].add(key)
|
|
568
|
+
return False
|
|
569
|
+
|
|
570
|
+
parent = self._parent()
|
|
571
|
+
cat_dim = parent.__cat_dim__(key, value, self) if parent is not None else -1
|
|
572
|
+
if not isinstance(cat_dim, int):
|
|
573
|
+
return False
|
|
574
|
+
|
|
575
|
+
num_nodes, num_edges = self.num_nodes, self.num_edges
|
|
576
|
+
|
|
577
|
+
if shape[cat_dim] != num_edges:
|
|
578
|
+
if shape[cat_dim] == num_nodes:
|
|
579
|
+
self._cached_attr[AttrType.NODE].add(key)
|
|
580
|
+
else:
|
|
581
|
+
self._cached_attr[AttrType.OTHER].add(key)
|
|
582
|
+
return False
|
|
583
|
+
|
|
584
|
+
if num_edges != num_nodes:
|
|
585
|
+
self._cached_attr[AttrType.EDGE].add(key)
|
|
586
|
+
return True
|
|
587
|
+
|
|
588
|
+
if "edge" in key:
|
|
589
|
+
self._cached_attr[AttrType.EDGE].add(key)
|
|
590
|
+
return True
|
|
591
|
+
else:
|
|
592
|
+
self._cached_attr[AttrType.NODE].add(key)
|
|
593
|
+
return False
|