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,374 @@
|
|
|
1
|
+
import copy
|
|
2
|
+
from collections import defaultdict
|
|
3
|
+
from typing import Any, Callable, Dict, Iterator, List, NamedTuple, Optional, Sequence, Tuple, Union
|
|
4
|
+
|
|
5
|
+
import numpy as np
|
|
6
|
+
from keras import ops
|
|
7
|
+
|
|
8
|
+
from k3_node.data.data import BaseData, Data, size_repr
|
|
9
|
+
from k3_node.data.storage import (
|
|
10
|
+
BaseStorage,
|
|
11
|
+
EdgeStorage,
|
|
12
|
+
NodeStorage,
|
|
13
|
+
get_shape,
|
|
14
|
+
is_tensor_like,
|
|
15
|
+
)
|
|
16
|
+
|
|
17
|
+
NodeType = str
|
|
18
|
+
EdgeType = Tuple[str, str, str]
|
|
19
|
+
|
|
20
|
+
|
|
21
|
+
class HeteroData(BaseData):
|
|
22
|
+
"""A data object describing a heterogeneous graph."""
|
|
23
|
+
|
|
24
|
+
def __init__(self, _mapping: Optional[Dict[str, Any]] = None, **kwargs):
|
|
25
|
+
self.__dict__["_node_store_dict"] = {}
|
|
26
|
+
self.__dict__["_edge_store_dict"] = {}
|
|
27
|
+
|
|
28
|
+
for key, value in (_mapping or {}).items():
|
|
29
|
+
setattr(self, key, value)
|
|
30
|
+
for key, value in kwargs.items():
|
|
31
|
+
setattr(self, key, value)
|
|
32
|
+
|
|
33
|
+
def _to_edge_type(self, key: Any) -> Optional[EdgeType]:
|
|
34
|
+
if isinstance(key, tuple):
|
|
35
|
+
if len(key) == 3:
|
|
36
|
+
return (str(key[0]), str(key[1]), str(key[2]))
|
|
37
|
+
if len(key) == 2:
|
|
38
|
+
# Find matching edge type
|
|
39
|
+
matches = [k for k in self.edge_types if k[0] == key[0] and k[-1] == key[1]]
|
|
40
|
+
if len(matches) == 1:
|
|
41
|
+
return matches[0]
|
|
42
|
+
elif len(matches) > 1:
|
|
43
|
+
raise KeyError(f"Ambiguous edge type for '{key}': {matches}")
|
|
44
|
+
return (str(key[0]), "to", str(key[1]))
|
|
45
|
+
elif isinstance(key, str):
|
|
46
|
+
matches = [k for k in self.edge_types if k[1] == key]
|
|
47
|
+
if len(matches) == 1:
|
|
48
|
+
return matches[0]
|
|
49
|
+
elif len(matches) > 1:
|
|
50
|
+
raise KeyError(f"Ambiguous edge type for rel '{key}': {matches}")
|
|
51
|
+
return None
|
|
52
|
+
|
|
53
|
+
def __getitem__(self, key: Any) -> Any:
|
|
54
|
+
edge_key = self._to_edge_type(key)
|
|
55
|
+
if edge_key is not None and (edge_key in self._edge_store_dict or isinstance(key, tuple)):
|
|
56
|
+
if edge_key not in self._edge_store_dict:
|
|
57
|
+
store = EdgeStorage(_parent=self)
|
|
58
|
+
store.__dict__["_key"] = edge_key
|
|
59
|
+
self._edge_store_dict[edge_key] = store
|
|
60
|
+
return self._edge_store_dict[edge_key]
|
|
61
|
+
|
|
62
|
+
if isinstance(key, str):
|
|
63
|
+
if key not in self._node_store_dict:
|
|
64
|
+
store = NodeStorage(_parent=self)
|
|
65
|
+
store.__dict__["_key"] = key
|
|
66
|
+
self._node_store_dict[key] = store
|
|
67
|
+
return self._node_store_dict[key]
|
|
68
|
+
|
|
69
|
+
raise KeyError(f"Invalid key '{key}' for {self.__class__.__name__}")
|
|
70
|
+
|
|
71
|
+
def __setitem__(self, key: Any, value: Any):
|
|
72
|
+
store = self[key]
|
|
73
|
+
if isinstance(value, BaseStorage):
|
|
74
|
+
for k, v in value.items():
|
|
75
|
+
store[k] = v
|
|
76
|
+
elif isinstance(value, dict):
|
|
77
|
+
for k, v in value.items():
|
|
78
|
+
store[k] = v
|
|
79
|
+
else:
|
|
80
|
+
raise ValueError(f"Value for key '{key}' must be a Storage or dict, got {type(value)}")
|
|
81
|
+
|
|
82
|
+
def __delitem__(self, key: Any):
|
|
83
|
+
edge_key = self._to_edge_type(key)
|
|
84
|
+
if edge_key is not None and edge_key in self._edge_store_dict:
|
|
85
|
+
del self._edge_store_dict[edge_key]
|
|
86
|
+
elif isinstance(key, str) and key in self._node_store_dict:
|
|
87
|
+
del self._node_store_dict[key]
|
|
88
|
+
else:
|
|
89
|
+
raise KeyError(key)
|
|
90
|
+
|
|
91
|
+
def __getattr__(self, key: str) -> Any:
|
|
92
|
+
if key in self.__dict__:
|
|
93
|
+
return self.__dict__[key]
|
|
94
|
+
if "_node_store_dict" in self.__dict__ and key in self._node_store_dict:
|
|
95
|
+
return self._node_store_dict[key]
|
|
96
|
+
if key.endswith("_dict") and "_node_store_dict" in self.__dict__:
|
|
97
|
+
return self.collect(key[:-5])
|
|
98
|
+
raise AttributeError(f"'{self.__class__.__name__}' object has no attribute '{key}'")
|
|
99
|
+
|
|
100
|
+
def collect(self, key: str) -> Dict[Any, Any]:
|
|
101
|
+
r"""Returns the attribute ``key`` of every node and edge type that has it, e.g.
|
|
102
|
+
``data.collect("x")`` (also available as ``data.x_dict``) or ``data.edge_index_dict``."""
|
|
103
|
+
out = {}
|
|
104
|
+
for stores in (self._node_store_dict, self._edge_store_dict):
|
|
105
|
+
for type_, store in stores.items():
|
|
106
|
+
value = store.get(key) if hasattr(store, "get") else getattr(store, key, None)
|
|
107
|
+
if value is not None:
|
|
108
|
+
out[type_] = value
|
|
109
|
+
return out
|
|
110
|
+
|
|
111
|
+
def __setattr__(self, key: str, value: Any):
|
|
112
|
+
if key in ("_node_store_dict", "_edge_store_dict"):
|
|
113
|
+
self.__dict__[key] = value
|
|
114
|
+
elif isinstance(value, BaseStorage):
|
|
115
|
+
if isinstance(value, NodeStorage):
|
|
116
|
+
self._node_store_dict[key] = value
|
|
117
|
+
elif isinstance(value, EdgeStorage):
|
|
118
|
+
edge_key = self._to_edge_type(key)
|
|
119
|
+
self._edge_store_dict[edge_key or key] = value
|
|
120
|
+
else:
|
|
121
|
+
self.__dict__[key] = value
|
|
122
|
+
|
|
123
|
+
@property
|
|
124
|
+
def node_types(self) -> List[NodeType]:
|
|
125
|
+
return list(self._node_store_dict.keys())
|
|
126
|
+
|
|
127
|
+
@property
|
|
128
|
+
def edge_types(self) -> List[EdgeType]:
|
|
129
|
+
return list(self._edge_store_dict.keys())
|
|
130
|
+
|
|
131
|
+
def metadata(self) -> Tuple[List[NodeType], List[EdgeType]]:
|
|
132
|
+
return self.node_types, self.edge_types
|
|
133
|
+
|
|
134
|
+
def get_node_store(self, key: NodeType) -> NodeStorage:
|
|
135
|
+
r"""Gets the NodeStorage object of a particular node type."""
|
|
136
|
+
out = self._node_store_dict.get(key, None)
|
|
137
|
+
if out is None:
|
|
138
|
+
out = NodeStorage(_parent=self)
|
|
139
|
+
out.__dict__["_key"] = key
|
|
140
|
+
self._node_store_dict[key] = out
|
|
141
|
+
return out
|
|
142
|
+
|
|
143
|
+
def get_edge_store(self, src: str, rel: str, dst: str) -> EdgeStorage:
|
|
144
|
+
r"""Gets the EdgeStorage object of a particular edge type given by (src, rel, dst)."""
|
|
145
|
+
key = (src, rel, dst)
|
|
146
|
+
out = self._edge_store_dict.get(key, None)
|
|
147
|
+
if out is None:
|
|
148
|
+
out = EdgeStorage(_parent=self)
|
|
149
|
+
out.__dict__["_key"] = key
|
|
150
|
+
self._edge_store_dict[key] = out
|
|
151
|
+
return out
|
|
152
|
+
|
|
153
|
+
def stores_as(self, data: "HeteroData") -> "HeteroData":
|
|
154
|
+
for node_type in data.node_types:
|
|
155
|
+
self.get_node_store(node_type)
|
|
156
|
+
for edge_type in data.edge_types:
|
|
157
|
+
self.get_edge_store(*edge_type)
|
|
158
|
+
return self
|
|
159
|
+
|
|
160
|
+
@property
|
|
161
|
+
def stores(self) -> List[BaseStorage]:
|
|
162
|
+
return list(self._node_store_dict.values()) + list(self._edge_store_dict.values())
|
|
163
|
+
|
|
164
|
+
@property
|
|
165
|
+
def node_stores(self) -> List[NodeStorage]:
|
|
166
|
+
return list(self._node_store_dict.values())
|
|
167
|
+
|
|
168
|
+
@property
|
|
169
|
+
def edge_stores(self) -> List[EdgeStorage]:
|
|
170
|
+
return list(self._edge_store_dict.values())
|
|
171
|
+
|
|
172
|
+
@property
|
|
173
|
+
def num_nodes_dict(self) -> Dict[NodeType, int]:
|
|
174
|
+
return {k: v.num_nodes for k, v in self._node_store_dict.items()}
|
|
175
|
+
|
|
176
|
+
@property
|
|
177
|
+
def num_edges_dict(self) -> Dict[EdgeType, int]:
|
|
178
|
+
return {k: v.num_edges for k, v in self._edge_store_dict.items()}
|
|
179
|
+
|
|
180
|
+
def set_value_dict(self, key: str, value_dict: Optional[Dict[Any, Any]]):
|
|
181
|
+
for k, v in (value_dict or {}).items():
|
|
182
|
+
self[k][key] = v
|
|
183
|
+
return self
|
|
184
|
+
|
|
185
|
+
def __copy__(self):
|
|
186
|
+
out = self.__class__.__new__(self.__class__)
|
|
187
|
+
out.__dict__["_node_store_dict"] = {}
|
|
188
|
+
out.__dict__["_edge_store_dict"] = {}
|
|
189
|
+
for k, v in self._node_store_dict.items():
|
|
190
|
+
store = copy.copy(v)
|
|
191
|
+
setattr(store, "_parent", out)
|
|
192
|
+
out._node_store_dict[k] = store
|
|
193
|
+
for k, v in self._edge_store_dict.items():
|
|
194
|
+
store = copy.copy(v)
|
|
195
|
+
setattr(store, "_parent", out)
|
|
196
|
+
out._edge_store_dict[k] = store
|
|
197
|
+
return out
|
|
198
|
+
|
|
199
|
+
def __deepcopy__(self, memo=None):
|
|
200
|
+
out = self.__class__.__new__(self.__class__)
|
|
201
|
+
out.__dict__["_node_store_dict"] = {}
|
|
202
|
+
out.__dict__["_edge_store_dict"] = {}
|
|
203
|
+
for k, v in self._node_store_dict.items():
|
|
204
|
+
store = copy.deepcopy(v, memo)
|
|
205
|
+
setattr(store, "_parent", out)
|
|
206
|
+
out._node_store_dict[k] = store
|
|
207
|
+
for k, v in self._edge_store_dict.items():
|
|
208
|
+
store = copy.deepcopy(v, memo)
|
|
209
|
+
setattr(store, "_parent", out)
|
|
210
|
+
out._edge_store_dict[k] = store
|
|
211
|
+
return out
|
|
212
|
+
|
|
213
|
+
def __getstate__(self) -> Dict[str, Any]:
|
|
214
|
+
return self.__dict__.copy()
|
|
215
|
+
|
|
216
|
+
def __setstate__(self, mapping: Dict[str, Any]):
|
|
217
|
+
import weakref
|
|
218
|
+
|
|
219
|
+
for key, value in mapping.items():
|
|
220
|
+
self.__dict__[key] = value
|
|
221
|
+
for store in self.stores:
|
|
222
|
+
store.__dict__["_parent"] = weakref.ref(self)
|
|
223
|
+
|
|
224
|
+
def clone(self) -> "HeteroData":
|
|
225
|
+
return copy.deepcopy(self)
|
|
226
|
+
|
|
227
|
+
def collect(self, key: str, allow_missing: bool = True) -> Dict[Any, Any]:
|
|
228
|
+
mapping = {}
|
|
229
|
+
for k, store in list(self._node_store_dict.items()) + list(self._edge_store_dict.items()):
|
|
230
|
+
if key in store:
|
|
231
|
+
mapping[k] = store[key]
|
|
232
|
+
elif not allow_missing:
|
|
233
|
+
raise KeyError(f"Key '{key}' not found in store '{k}'")
|
|
234
|
+
return mapping
|
|
235
|
+
|
|
236
|
+
def __inc__(self, key: str, value: Any, store: Optional[BaseStorage] = None, *args, **kwargs) -> Any:
|
|
237
|
+
if "batch" in key:
|
|
238
|
+
return int(value.max()) + 1 if is_tensor_like(value) and value.ndim > 0 and value.shape[0] > 0 else 0
|
|
239
|
+
if "index" in key and store is not None and isinstance(store, EdgeStorage):
|
|
240
|
+
edge_type = store._key
|
|
241
|
+
src, _, dst = edge_type
|
|
242
|
+
src_num = self[src].num_nodes
|
|
243
|
+
dst_num = self[dst].num_nodes
|
|
244
|
+
return np.array([[src_num], [dst_num]])
|
|
245
|
+
return 0
|
|
246
|
+
|
|
247
|
+
def __cat_dim__(self, key: str, value: Any, store: Optional[BaseStorage] = None, *args, **kwargs) -> int:
|
|
248
|
+
if key in ("edge_index", "adj_t"):
|
|
249
|
+
return -1
|
|
250
|
+
if is_tensor_like(value) and len(get_shape(value)) == 2 and get_shape(value)[0] == 2 and "index" in key:
|
|
251
|
+
return -1
|
|
252
|
+
return 0
|
|
253
|
+
|
|
254
|
+
def to_dict(self) -> Dict[str, Any]:
|
|
255
|
+
out = {}
|
|
256
|
+
for k, store in self._node_store_dict.items():
|
|
257
|
+
out[k] = store.to_dict()
|
|
258
|
+
for k, store in self._edge_store_dict.items():
|
|
259
|
+
out[k] = store.to_dict()
|
|
260
|
+
return out
|
|
261
|
+
|
|
262
|
+
def to_namedtuple(self) -> NamedTuple:
|
|
263
|
+
# Build nested namedtuple
|
|
264
|
+
node_fields = sorted(list(self._node_store_dict.keys()))
|
|
265
|
+
node_dict = {k: self._node_store_dict[k].to_dict() for k in node_fields}
|
|
266
|
+
edge_fields = [f"{k[0]}__{k[1]}__{k[2]}" for k in sorted(list(self._edge_store_dict.keys()))]
|
|
267
|
+
edge_dict = {f"{k[0]}__{k[1]}__{k[2]}": self._edge_store_dict[k].to_dict() for k in sorted(list(self._edge_store_dict.keys()))}
|
|
268
|
+
all_fields = node_fields + edge_fields
|
|
269
|
+
HeteroTuple = collections.namedtuple("HeteroTuple", all_fields)
|
|
270
|
+
return HeteroTuple(**node_dict, **edge_dict)
|
|
271
|
+
|
|
272
|
+
def edge_type_subgraph(self, edge_types: List[EdgeType]) -> "HeteroData":
|
|
273
|
+
out = copy.deepcopy(self)
|
|
274
|
+
for et in list(out._edge_store_dict.keys()):
|
|
275
|
+
if et not in edge_types:
|
|
276
|
+
del out._edge_store_dict[et]
|
|
277
|
+
return out
|
|
278
|
+
|
|
279
|
+
def subgraph(self, subset_dict: Dict[NodeType, Any]) -> "HeteroData":
|
|
280
|
+
out = copy.deepcopy(self)
|
|
281
|
+
for node_type, subset in subset_dict.items():
|
|
282
|
+
store = out[node_type]
|
|
283
|
+
subset_np = ops.convert_to_numpy(subset)
|
|
284
|
+
indices = np.where(subset_np)[0] if subset_np.dtype == bool else subset_np
|
|
285
|
+
for key in store.node_attrs():
|
|
286
|
+
val = store[key]
|
|
287
|
+
if is_tensor_like(val):
|
|
288
|
+
store[key] = ops.take(val, indices, axis=self.__cat_dim__(key, val, store))
|
|
289
|
+
return out
|
|
290
|
+
|
|
291
|
+
def to_homogeneous(
|
|
292
|
+
self,
|
|
293
|
+
node_attrs: Optional[List[str]] = None,
|
|
294
|
+
edge_attrs: Optional[List[str]] = None,
|
|
295
|
+
add_node_type: bool = True,
|
|
296
|
+
add_edge_type: bool = True,
|
|
297
|
+
dummy_values: bool = True,
|
|
298
|
+
) -> Data:
|
|
299
|
+
data = Data()
|
|
300
|
+
|
|
301
|
+
# Compute node offsets and slices
|
|
302
|
+
node_slices = {}
|
|
303
|
+
curr_offset = 0
|
|
304
|
+
node_type_list = []
|
|
305
|
+
for i, node_type in enumerate(self.node_types):
|
|
306
|
+
num_nodes = self[node_type].num_nodes
|
|
307
|
+
node_slices[node_type] = curr_offset
|
|
308
|
+
if add_node_type:
|
|
309
|
+
node_type_list.append(np.full((num_nodes,), i, dtype=np.int64))
|
|
310
|
+
curr_offset += num_nodes
|
|
311
|
+
|
|
312
|
+
data.num_nodes = curr_offset
|
|
313
|
+
if add_node_type and len(node_type_list) > 0:
|
|
314
|
+
node_type_arr = np.concatenate(node_type_list, axis=0)
|
|
315
|
+
data.node_type = ops.convert_to_tensor(node_type_arr, dtype="int64")
|
|
316
|
+
|
|
317
|
+
# Concat node features
|
|
318
|
+
if node_attrs is None:
|
|
319
|
+
# find common node attrs across node types
|
|
320
|
+
all_node_attrs = set()
|
|
321
|
+
for store in self.node_stores:
|
|
322
|
+
all_node_attrs.update(store.node_attrs())
|
|
323
|
+
node_attrs = list(all_node_attrs)
|
|
324
|
+
|
|
325
|
+
for attr in node_attrs:
|
|
326
|
+
attr_vals = []
|
|
327
|
+
for node_type in self.node_types:
|
|
328
|
+
val = self[node_type].get(attr)
|
|
329
|
+
if val is not None:
|
|
330
|
+
attr_vals.append(ops.convert_to_numpy(val))
|
|
331
|
+
elif dummy_values:
|
|
332
|
+
num_nodes = self[node_type].num_nodes
|
|
333
|
+
attr_vals.append(np.zeros((num_nodes, 0), dtype=np.float32))
|
|
334
|
+
if len(attr_vals) > 0:
|
|
335
|
+
concat_val = np.concatenate(attr_vals, axis=0)
|
|
336
|
+
data[attr] = ops.convert_to_tensor(concat_val)
|
|
337
|
+
|
|
338
|
+
# Offsetting edge indices
|
|
339
|
+
edge_indices = []
|
|
340
|
+
edge_type_list = []
|
|
341
|
+
for i, edge_type in enumerate(self.edge_types):
|
|
342
|
+
store = self[edge_type]
|
|
343
|
+
edge_index = store.get("edge_index")
|
|
344
|
+
if edge_index is not None:
|
|
345
|
+
ei_np = ops.convert_to_numpy(edge_index).copy()
|
|
346
|
+
src_offset = node_slices[edge_type[0]]
|
|
347
|
+
dst_offset = node_slices[edge_type[-1]]
|
|
348
|
+
ei_np[0] += src_offset
|
|
349
|
+
ei_np[1] += dst_offset
|
|
350
|
+
edge_indices.append(ei_np)
|
|
351
|
+
if add_edge_type:
|
|
352
|
+
edge_type_list.append(np.full((ei_np.shape[1],), i, dtype=np.int64))
|
|
353
|
+
|
|
354
|
+
if len(edge_indices) > 0:
|
|
355
|
+
concat_ei = np.concatenate(edge_indices, axis=1)
|
|
356
|
+
data.edge_index = ops.convert_to_tensor(concat_ei, dtype="int64")
|
|
357
|
+
if add_edge_type and len(edge_type_list) > 0:
|
|
358
|
+
concat_et = np.concatenate(edge_type_list, axis=0)
|
|
359
|
+
data.edge_type = ops.convert_to_tensor(concat_et, dtype="int64")
|
|
360
|
+
|
|
361
|
+
return data
|
|
362
|
+
|
|
363
|
+
def __repr__(self) -> str:
|
|
364
|
+
cls = self.__class__.__name__
|
|
365
|
+
info_lines = []
|
|
366
|
+
for k, store in self._node_store_dict.items():
|
|
367
|
+
attrs = [size_repr(attr, store[attr]) for attr in store.keys()]
|
|
368
|
+
info_lines.append(f" {k}={{{', '.join(attrs)}}}")
|
|
369
|
+
for k, store in self._edge_store_dict.items():
|
|
370
|
+
attrs = [size_repr(attr, store[attr]) for attr in store.keys()]
|
|
371
|
+
edge_name = f"('{k[0]}', '{k[1]}', '{k[2]}')"
|
|
372
|
+
info_lines.append(f" {edge_name}={{{', '.join(attrs)}}}")
|
|
373
|
+
info = ",\n".join(info_lines)
|
|
374
|
+
return f"{cls}(\n{info}\n)" if info else f"{cls}()"
|
|
@@ -0,0 +1,59 @@
|
|
|
1
|
+
from typing import Any, List, Optional
|
|
2
|
+
import numpy as np
|
|
3
|
+
from keras import ops
|
|
4
|
+
|
|
5
|
+
from k3_node.data.data import Data
|
|
6
|
+
from k3_node.data.storage import is_tensor_like
|
|
7
|
+
|
|
8
|
+
|
|
9
|
+
class HypergraphData(Data):
|
|
10
|
+
"""A data object describing a hypergraph."""
|
|
11
|
+
|
|
12
|
+
def __init__(
|
|
13
|
+
self,
|
|
14
|
+
x=None,
|
|
15
|
+
edge_index=None,
|
|
16
|
+
edge_attr=None,
|
|
17
|
+
y=None,
|
|
18
|
+
pos=None,
|
|
19
|
+
**kwargs,
|
|
20
|
+
):
|
|
21
|
+
super().__init__(
|
|
22
|
+
x=x,
|
|
23
|
+
edge_index=edge_index,
|
|
24
|
+
edge_attr=edge_attr,
|
|
25
|
+
y=y,
|
|
26
|
+
pos=pos,
|
|
27
|
+
**kwargs,
|
|
28
|
+
)
|
|
29
|
+
|
|
30
|
+
@property
|
|
31
|
+
def num_edges(self) -> int:
|
|
32
|
+
if self.edge_index is None:
|
|
33
|
+
return 0
|
|
34
|
+
ei_np = ops.convert_to_numpy(self.edge_index)
|
|
35
|
+
if ei_np.size == 0 or ei_np.shape[1] == 0:
|
|
36
|
+
return 0
|
|
37
|
+
return int(np.max(ei_np[1])) + 1
|
|
38
|
+
|
|
39
|
+
@property
|
|
40
|
+
def num_nodes(self) -> Optional[int]:
|
|
41
|
+
num = super().num_nodes
|
|
42
|
+
if self.edge_index is not None and num == self.num_edges:
|
|
43
|
+
ei_np = ops.convert_to_numpy(self.edge_index)
|
|
44
|
+
if ei_np.size > 0 and ei_np.shape[1] > 0:
|
|
45
|
+
return int(np.max(ei_np[0])) + 1
|
|
46
|
+
return num
|
|
47
|
+
|
|
48
|
+
@num_nodes.setter
|
|
49
|
+
def num_nodes(self, num_nodes: Optional[int]) -> None:
|
|
50
|
+
self._store.num_nodes = num_nodes
|
|
51
|
+
|
|
52
|
+
def __inc__(self, key: str, value: Any, *args, **kwargs) -> Any:
|
|
53
|
+
if key == "edge_index":
|
|
54
|
+
return np.array([[self.num_nodes], [self.num_edges]])
|
|
55
|
+
return super().__inc__(key, value, *args, **kwargs)
|
|
56
|
+
|
|
57
|
+
|
|
58
|
+
HyperGraphData = HypergraphData
|
|
59
|
+
|
|
@@ -0,0 +1,177 @@
|
|
|
1
|
+
import copy
|
|
2
|
+
import os
|
|
3
|
+
import os.path as osp
|
|
4
|
+
import pickle
|
|
5
|
+
from typing import Any, Callable, Dict, List, Optional, Tuple, Union
|
|
6
|
+
|
|
7
|
+
from k3_node.data.collate import collate
|
|
8
|
+
from k3_node.data.data import BaseData, Data
|
|
9
|
+
from k3_node.data.dataset import Dataset
|
|
10
|
+
from k3_node.data.separate import separate
|
|
11
|
+
|
|
12
|
+
|
|
13
|
+
class InMemoryDataset(Dataset):
|
|
14
|
+
"""Dataset base class for in-memory graph collections."""
|
|
15
|
+
|
|
16
|
+
def __init__(
|
|
17
|
+
self,
|
|
18
|
+
root: Optional[str] = None,
|
|
19
|
+
transform: Optional[Callable] = None,
|
|
20
|
+
pre_transform: Optional[Callable] = None,
|
|
21
|
+
pre_filter: Optional[Callable] = None,
|
|
22
|
+
log: bool = True,
|
|
23
|
+
force_reload: bool = False,
|
|
24
|
+
):
|
|
25
|
+
# Must be set before `super().__init__()`, which triggers
|
|
26
|
+
# `self.process()` for datasets processed for the first time --
|
|
27
|
+
# `process()` implementations (e.g. `TUDataset`) commonly call
|
|
28
|
+
# `len(self)` / `self.get(idx)`, both of which read these attributes.
|
|
29
|
+
self._data: Optional[BaseData] = None
|
|
30
|
+
self.slices: Optional[Dict[str, Any]] = None
|
|
31
|
+
self.sizes: Dict[str, Any] = {}
|
|
32
|
+
self._data_list: Optional[List[BaseData]] = None
|
|
33
|
+
super().__init__(root, transform, pre_transform, pre_filter, log, force_reload)
|
|
34
|
+
|
|
35
|
+
@property
|
|
36
|
+
def data(self) -> Optional[BaseData]:
|
|
37
|
+
return self._data
|
|
38
|
+
|
|
39
|
+
@data.setter
|
|
40
|
+
def data(self, value: Optional[BaseData]):
|
|
41
|
+
self._data = value
|
|
42
|
+
|
|
43
|
+
def len(self) -> int:
|
|
44
|
+
if self._data_list is not None:
|
|
45
|
+
return len(self._data_list)
|
|
46
|
+
if self.slices is None:
|
|
47
|
+
return 1 if self._data is not None else 0
|
|
48
|
+
for key, value in self.slices.items():
|
|
49
|
+
if isinstance(value, dict):
|
|
50
|
+
for _, sub_val in value.items():
|
|
51
|
+
return len(sub_val) - 1
|
|
52
|
+
return len(value) - 1
|
|
53
|
+
return 0
|
|
54
|
+
|
|
55
|
+
def get(self, idx: int) -> BaseData:
|
|
56
|
+
if self._data_list is not None:
|
|
57
|
+
return self._data_list[idx]
|
|
58
|
+
if self._data is None:
|
|
59
|
+
raise RuntimeError("Dataset does not contain data. Call 'load' first.")
|
|
60
|
+
if self.slices is None:
|
|
61
|
+
if idx == 0:
|
|
62
|
+
return copy.copy(self._data)
|
|
63
|
+
raise IndexError(f"Index {idx} out of bounds for single graph dataset.")
|
|
64
|
+
return separate(self._data.__class__, self._data, idx, self.slices)
|
|
65
|
+
|
|
66
|
+
@classmethod
|
|
67
|
+
def collate(cls, data_list: List[BaseData]) -> Tuple[BaseData, Optional[Dict[str, Any]]]:
|
|
68
|
+
if len(data_list) == 1:
|
|
69
|
+
return data_list[0], None
|
|
70
|
+
base_cls = data_list[0].__class__
|
|
71
|
+
data, slices, _ = collate(
|
|
72
|
+
base_cls,
|
|
73
|
+
data_list,
|
|
74
|
+
increment=False,
|
|
75
|
+
add_batch=False,
|
|
76
|
+
)
|
|
77
|
+
return data, slices
|
|
78
|
+
|
|
79
|
+
def save(self, data_list: List[BaseData], path: str):
|
|
80
|
+
os.makedirs(osp.dirname(path), exist_ok=True)
|
|
81
|
+
data, slices = self.collate(data_list)
|
|
82
|
+
if path.endswith((".pt", ".pth")):
|
|
83
|
+
try:
|
|
84
|
+
import torch
|
|
85
|
+
data_dict = data.to_dict() if hasattr(data, "to_dict") else dict(data)
|
|
86
|
+
torch.save((data_dict, slices), path)
|
|
87
|
+
return
|
|
88
|
+
except Exception:
|
|
89
|
+
pass
|
|
90
|
+
with open(path, "wb") as f:
|
|
91
|
+
pickle.dump((data, slices), f)
|
|
92
|
+
|
|
93
|
+
def load(self, path: str):
|
|
94
|
+
obj = None
|
|
95
|
+
if path.endswith((".pt", ".pth")):
|
|
96
|
+
try:
|
|
97
|
+
import torch
|
|
98
|
+
obj = torch.load(path, map_location="cpu", weights_only=False)
|
|
99
|
+
except Exception:
|
|
100
|
+
pass
|
|
101
|
+
|
|
102
|
+
if obj is None:
|
|
103
|
+
try:
|
|
104
|
+
with open(path, "rb") as f:
|
|
105
|
+
obj = pickle.load(f)
|
|
106
|
+
except Exception:
|
|
107
|
+
try:
|
|
108
|
+
import torch
|
|
109
|
+
obj = torch.load(path, map_location="cpu", weights_only=False)
|
|
110
|
+
except Exception:
|
|
111
|
+
pass
|
|
112
|
+
|
|
113
|
+
if obj is None:
|
|
114
|
+
if hasattr(self, "process") and callable(self.process):
|
|
115
|
+
self.process()
|
|
116
|
+
try:
|
|
117
|
+
with open(path, "rb") as f:
|
|
118
|
+
obj = pickle.load(f)
|
|
119
|
+
except Exception:
|
|
120
|
+
try:
|
|
121
|
+
import torch
|
|
122
|
+
obj = torch.load(path, map_location="cpu", weights_only=False)
|
|
123
|
+
except Exception:
|
|
124
|
+
pass
|
|
125
|
+
if obj is None:
|
|
126
|
+
raise RuntimeError(f"Cannot load dataset from {path}")
|
|
127
|
+
|
|
128
|
+
if isinstance(obj, tuple):
|
|
129
|
+
if len(obj) == 2:
|
|
130
|
+
data, self.slices = obj
|
|
131
|
+
elif len(obj) == 3:
|
|
132
|
+
data, self.slices, extra = obj
|
|
133
|
+
if isinstance(extra, dict):
|
|
134
|
+
self.sizes = extra
|
|
135
|
+
elif len(obj) >= 4:
|
|
136
|
+
data, self.slices = obj[0], obj[1]
|
|
137
|
+
if isinstance(obj[2], dict):
|
|
138
|
+
self.sizes = obj[2]
|
|
139
|
+
else:
|
|
140
|
+
data = obj[0]
|
|
141
|
+
if isinstance(data, dict):
|
|
142
|
+
data = _from_dict(data)
|
|
143
|
+
self._data = data
|
|
144
|
+
elif isinstance(obj, list):
|
|
145
|
+
self._data_list = obj
|
|
146
|
+
elif isinstance(obj, dict):
|
|
147
|
+
self._data = _from_dict(obj)
|
|
148
|
+
else:
|
|
149
|
+
self._data = obj
|
|
150
|
+
|
|
151
|
+
if self._data is not None and hasattr(self._data, "to_backend"):
|
|
152
|
+
try:
|
|
153
|
+
self._data.to_backend()
|
|
154
|
+
except Exception:
|
|
155
|
+
pass
|
|
156
|
+
|
|
157
|
+
if isinstance(self.slices, dict):
|
|
158
|
+
from k3_node.data.storage import is_tensor_like, to_numpy
|
|
159
|
+
self.slices = {k: to_numpy(v) if is_tensor_like(v) else v for k, v in self.slices.items()}
|
|
160
|
+
|
|
161
|
+
|
|
162
|
+
|
|
163
|
+
def _from_dict(mapping):
|
|
164
|
+
"""Rebuilds a saved graph: a ``HeteroData`` if the mapping holds per-type attribute
|
|
165
|
+
dictionaries (node types and ``(src, rel, dst)`` edge types), else a ``Data``."""
|
|
166
|
+
if any(isinstance(v, dict) or not isinstance(k, str) for k, v in mapping.items()):
|
|
167
|
+
from k3_node.data.hetero_data import HeteroData
|
|
168
|
+
|
|
169
|
+
data = HeteroData()
|
|
170
|
+
for key, value in mapping.items():
|
|
171
|
+
if isinstance(value, dict):
|
|
172
|
+
for attr, item in value.items():
|
|
173
|
+
setattr(data[key], attr, item)
|
|
174
|
+
else:
|
|
175
|
+
setattr(data, key, value)
|
|
176
|
+
return data
|
|
177
|
+
return Data(**mapping)
|
k3_node/data/makedirs.py
ADDED
|
@@ -0,0 +1,77 @@
|
|
|
1
|
+
import os
|
|
2
|
+
import os.path as osp
|
|
3
|
+
from typing import Any, Callable, List, Optional, Union
|
|
4
|
+
|
|
5
|
+
from k3_node.data.data import BaseData
|
|
6
|
+
from k3_node.data.database import Database, RocksDatabase, SQLiteDatabase, Schema
|
|
7
|
+
from k3_node.data.dataset import Dataset
|
|
8
|
+
|
|
9
|
+
|
|
10
|
+
class OnDiskDataset(Dataset):
|
|
11
|
+
"""Dataset base class for out-of-core graph datasets using a Database backend."""
|
|
12
|
+
|
|
13
|
+
BACKENDS = {
|
|
14
|
+
"sqlite": SQLiteDatabase,
|
|
15
|
+
"rocksdb": RocksDatabase,
|
|
16
|
+
}
|
|
17
|
+
|
|
18
|
+
def __init__(
|
|
19
|
+
self,
|
|
20
|
+
root: str,
|
|
21
|
+
transform: Optional[Callable] = None,
|
|
22
|
+
pre_filter: Optional[Callable] = None,
|
|
23
|
+
backend: str = "sqlite",
|
|
24
|
+
schema: Schema = object,
|
|
25
|
+
log: bool = True,
|
|
26
|
+
):
|
|
27
|
+
if backend not in self.BACKENDS:
|
|
28
|
+
raise ValueError(f"Database backend must be one of {set(self.BACKENDS.keys())}, got '{backend}'")
|
|
29
|
+
|
|
30
|
+
self.backend = backend
|
|
31
|
+
self.schema = schema
|
|
32
|
+
self._db: Optional[Database] = None
|
|
33
|
+
|
|
34
|
+
super().__init__(root, transform, pre_filter=pre_filter, log=log)
|
|
35
|
+
|
|
36
|
+
@property
|
|
37
|
+
def processed_file_names(self) -> str:
|
|
38
|
+
return f"{self.backend}.db"
|
|
39
|
+
|
|
40
|
+
@property
|
|
41
|
+
def db(self) -> Database:
|
|
42
|
+
if self._db is not None:
|
|
43
|
+
return self._db
|
|
44
|
+
|
|
45
|
+
cls = self.BACKENDS[self.backend]
|
|
46
|
+
os.makedirs(self.processed_dir, exist_ok=True)
|
|
47
|
+
path = osp.join(self.processed_dir, self.processed_file_names)
|
|
48
|
+
self._db = cls(path=path, schema=self.schema)
|
|
49
|
+
return self._db
|
|
50
|
+
|
|
51
|
+
def close(self):
|
|
52
|
+
if self._db is not None:
|
|
53
|
+
self._db.close()
|
|
54
|
+
self._db = None
|
|
55
|
+
|
|
56
|
+
def serialize(self, data: BaseData) -> Any:
|
|
57
|
+
return data
|
|
58
|
+
|
|
59
|
+
def deserialize(self, data: Any) -> BaseData:
|
|
60
|
+
return data
|
|
61
|
+
|
|
62
|
+
def len(self) -> int:
|
|
63
|
+
return len(self.db)
|
|
64
|
+
|
|
65
|
+
def get(self, idx: int) -> BaseData:
|
|
66
|
+
return self.deserialize(self.db[idx])
|
|
67
|
+
|
|
68
|
+
def append(self, data: BaseData):
|
|
69
|
+
idx = len(self)
|
|
70
|
+
self.db[idx] = self.serialize(data)
|
|
71
|
+
|
|
72
|
+
def extend(self, data_list: List[BaseData]):
|
|
73
|
+
start = len(self)
|
|
74
|
+
indices = list(range(start, start + len(data_list)))
|
|
75
|
+
serialized = [self.serialize(d) for d in data_list]
|
|
76
|
+
self.db.multi_insert(indices, serialized)
|
|
77
|
+
|