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/data.py
ADDED
|
@@ -0,0 +1,532 @@
|
|
|
1
|
+
import collections
|
|
2
|
+
import copy
|
|
3
|
+
import warnings
|
|
4
|
+
from collections.abc import Mapping, Sequence
|
|
5
|
+
from itertools import chain
|
|
6
|
+
from typing import Any, Callable, Dict, Iterable, Iterator, List, NamedTuple, Optional, Tuple, Union
|
|
7
|
+
|
|
8
|
+
import numpy as np
|
|
9
|
+
from keras import ops
|
|
10
|
+
|
|
11
|
+
from k3_node.data.storage import (
|
|
12
|
+
BaseStorage,
|
|
13
|
+
EdgeStorage,
|
|
14
|
+
GlobalStorage,
|
|
15
|
+
NodeStorage,
|
|
16
|
+
get_shape,
|
|
17
|
+
is_tensor_like,
|
|
18
|
+
recursive_apply,
|
|
19
|
+
recursive_apply_,
|
|
20
|
+
)
|
|
21
|
+
from k3_node.utils.graph import coalesce, contains_isolated_nodes, has_self_loops, is_undirected, subgraph
|
|
22
|
+
|
|
23
|
+
|
|
24
|
+
def size_repr(key: Any, value: Any, indent: int = 0) -> str:
|
|
25
|
+
pad = " " * indent
|
|
26
|
+
if is_tensor_like(value):
|
|
27
|
+
shape = get_shape(value)
|
|
28
|
+
if len(shape) == 0:
|
|
29
|
+
out = str(ops.convert_to_numpy(value).item())
|
|
30
|
+
else:
|
|
31
|
+
out = str(list(shape))
|
|
32
|
+
elif isinstance(value, str):
|
|
33
|
+
out = f"'{value}'"
|
|
34
|
+
elif isinstance(value, (Sequence, set)) and not isinstance(value, str):
|
|
35
|
+
out = str([len(value)])
|
|
36
|
+
elif isinstance(value, Mapping) and len(value) == 0:
|
|
37
|
+
out = "{}"
|
|
38
|
+
elif isinstance(value, Mapping) and len(value) == 1 and not isinstance(list(value.values())[0], Mapping):
|
|
39
|
+
lines = [size_repr(k, v, 0) for k, v in value.items()]
|
|
40
|
+
out = "{ " + ", ".join(lines) + " }"
|
|
41
|
+
elif isinstance(value, Mapping):
|
|
42
|
+
lines = [size_repr(k, v, indent + 2) for k, v in value.items()]
|
|
43
|
+
out = "{\n" + ",\n".join(lines) + ",\n" + pad + "}"
|
|
44
|
+
else:
|
|
45
|
+
out = str(value)
|
|
46
|
+
|
|
47
|
+
key = str(key).replace("'", "")
|
|
48
|
+
return f"{pad}{key}={out}"
|
|
49
|
+
|
|
50
|
+
|
|
51
|
+
class BaseData:
|
|
52
|
+
def __getattr__(self, key: str) -> Any:
|
|
53
|
+
raise NotImplementedError
|
|
54
|
+
|
|
55
|
+
def __setattr__(self, key: str, value: Any):
|
|
56
|
+
raise NotImplementedError
|
|
57
|
+
|
|
58
|
+
def __delattr__(self, key: str):
|
|
59
|
+
raise NotImplementedError
|
|
60
|
+
|
|
61
|
+
def __getitem__(self, key: str) -> Any:
|
|
62
|
+
raise NotImplementedError
|
|
63
|
+
|
|
64
|
+
def __setitem__(self, key: str, value: Any):
|
|
65
|
+
raise NotImplementedError
|
|
66
|
+
|
|
67
|
+
def __delitem__(self, key: str):
|
|
68
|
+
raise NotImplementedError
|
|
69
|
+
|
|
70
|
+
def __copy__(self):
|
|
71
|
+
raise NotImplementedError
|
|
72
|
+
|
|
73
|
+
def __deepcopy__(self, memo=None):
|
|
74
|
+
raise NotImplementedError
|
|
75
|
+
|
|
76
|
+
def __repr__(self) -> str:
|
|
77
|
+
raise NotImplementedError
|
|
78
|
+
|
|
79
|
+
@property
|
|
80
|
+
def stores(self) -> List[BaseStorage]:
|
|
81
|
+
raise NotImplementedError
|
|
82
|
+
|
|
83
|
+
@property
|
|
84
|
+
def node_stores(self) -> List[NodeStorage]:
|
|
85
|
+
raise NotImplementedError
|
|
86
|
+
|
|
87
|
+
@property
|
|
88
|
+
def edge_stores(self) -> List[EdgeStorage]:
|
|
89
|
+
raise NotImplementedError
|
|
90
|
+
|
|
91
|
+
def stores_as(self, data: "BaseData") -> "BaseData":
|
|
92
|
+
raise NotImplementedError
|
|
93
|
+
|
|
94
|
+
def to_dict(self) -> Dict[str, Any]:
|
|
95
|
+
raise NotImplementedError
|
|
96
|
+
|
|
97
|
+
def to_namedtuple(self) -> NamedTuple:
|
|
98
|
+
raise NotImplementedError
|
|
99
|
+
|
|
100
|
+
def to_backend(self, backend: Optional[str] = None) -> "BaseData":
|
|
101
|
+
for store in self.stores:
|
|
102
|
+
store.to_backend(backend)
|
|
103
|
+
return self
|
|
104
|
+
|
|
105
|
+
def update(self, data: "BaseData") -> "BaseData":
|
|
106
|
+
for store, other_store in zip(self.stores, data.stores):
|
|
107
|
+
for key, value in other_store.items():
|
|
108
|
+
store[key] = value
|
|
109
|
+
return self
|
|
110
|
+
|
|
111
|
+
def __len__(self) -> int:
|
|
112
|
+
return len(self.keys())
|
|
113
|
+
|
|
114
|
+
def __contains__(self, key: str) -> bool:
|
|
115
|
+
return key in self.keys()
|
|
116
|
+
|
|
117
|
+
def keys(self, *args: str) -> List[str]:
|
|
118
|
+
out = []
|
|
119
|
+
for store in self.stores:
|
|
120
|
+
out.extend(list(store.keys(*args)))
|
|
121
|
+
return list(set(out))
|
|
122
|
+
|
|
123
|
+
def values(self, *args: str) -> List[Any]:
|
|
124
|
+
return [self[k] for k in self.keys(*args)]
|
|
125
|
+
|
|
126
|
+
def items(self, *args: str) -> List[Tuple[str, Any]]:
|
|
127
|
+
return [(k, self[k]) for k in self.keys(*args)]
|
|
128
|
+
|
|
129
|
+
@property
|
|
130
|
+
def num_nodes(self) -> Optional[int]:
|
|
131
|
+
try:
|
|
132
|
+
return sum([v.num_nodes for v in self.node_stores])
|
|
133
|
+
except TypeError:
|
|
134
|
+
return None
|
|
135
|
+
|
|
136
|
+
@property
|
|
137
|
+
def num_edges(self) -> int:
|
|
138
|
+
return sum([v.num_edges for v in self.edge_stores])
|
|
139
|
+
|
|
140
|
+
def node_attrs(self) -> List[str]:
|
|
141
|
+
return list(set(chain(*[s.node_attrs() for s in self.node_stores])))
|
|
142
|
+
|
|
143
|
+
def edge_attrs(self) -> List[str]:
|
|
144
|
+
return list(set(chain(*[s.edge_attrs() for s in self.edge_stores])))
|
|
145
|
+
|
|
146
|
+
|
|
147
|
+
class Data(BaseData):
|
|
148
|
+
"""A data object describing a homogeneous graph."""
|
|
149
|
+
|
|
150
|
+
def __init__(
|
|
151
|
+
self,
|
|
152
|
+
x=None,
|
|
153
|
+
edge_index=None,
|
|
154
|
+
edge_attr=None,
|
|
155
|
+
y=None,
|
|
156
|
+
pos=None,
|
|
157
|
+
**kwargs,
|
|
158
|
+
):
|
|
159
|
+
self.__dict__["_store"] = GlobalStorage(_parent=self)
|
|
160
|
+
if x is not None:
|
|
161
|
+
self.x = x
|
|
162
|
+
if edge_index is not None:
|
|
163
|
+
self.edge_index = edge_index
|
|
164
|
+
if edge_attr is not None:
|
|
165
|
+
self.edge_attr = edge_attr
|
|
166
|
+
if y is not None:
|
|
167
|
+
self.y = y
|
|
168
|
+
if pos is not None:
|
|
169
|
+
self.pos = pos
|
|
170
|
+
for key, value in kwargs.items():
|
|
171
|
+
setattr(self, key, value)
|
|
172
|
+
|
|
173
|
+
def __getattr__(self, key: str) -> Any:
|
|
174
|
+
if "_store" not in self.__dict__:
|
|
175
|
+
raise AttributeError(f"'{self.__class__.__name__}' object has no attribute '{key}'")
|
|
176
|
+
try:
|
|
177
|
+
return getattr(self._store, key)
|
|
178
|
+
except AttributeError:
|
|
179
|
+
if key in ('x', 'edge_index', 'edge_attr', 'edge_weight', 'y', 'pos', 'face', 'normal', 'batch'):
|
|
180
|
+
return None
|
|
181
|
+
raise AttributeError(f"'{self.__class__.__name__}' object has no attribute '{key}'") from None
|
|
182
|
+
|
|
183
|
+
def __setattr__(self, key: str, value: Any):
|
|
184
|
+
if key == "_store":
|
|
185
|
+
self.__dict__["_store"] = value
|
|
186
|
+
elif "_store" in self.__dict__:
|
|
187
|
+
setattr(self._store, key, value)
|
|
188
|
+
else:
|
|
189
|
+
self.__dict__[key] = value
|
|
190
|
+
|
|
191
|
+
def __delattr__(self, key: str):
|
|
192
|
+
if key == "_store":
|
|
193
|
+
del self.__dict__["_store"]
|
|
194
|
+
elif "_store" in self.__dict__:
|
|
195
|
+
delattr(self._store, key)
|
|
196
|
+
else:
|
|
197
|
+
del self.__dict__[key]
|
|
198
|
+
|
|
199
|
+
def __getitem__(self, key: str) -> Any:
|
|
200
|
+
return self._store[key]
|
|
201
|
+
|
|
202
|
+
def __setitem__(self, key: str, value: Any):
|
|
203
|
+
self._store[key] = value
|
|
204
|
+
|
|
205
|
+
def __delitem__(self, key: str):
|
|
206
|
+
del self._store[key]
|
|
207
|
+
|
|
208
|
+
def __copy__(self):
|
|
209
|
+
out = self.__class__.__new__(self.__class__)
|
|
210
|
+
for k, v in self.__dict__.items():
|
|
211
|
+
out.__dict__[k] = v
|
|
212
|
+
out._store = copy.copy(self._store)
|
|
213
|
+
out._store._parent = weakref_self = None
|
|
214
|
+
out.__dict__["_store"].__dict__["_parent"] = weakref_self
|
|
215
|
+
setattr(out._store, "_parent", out)
|
|
216
|
+
return out
|
|
217
|
+
|
|
218
|
+
def __deepcopy__(self, memo=None):
|
|
219
|
+
out = self.__class__.__new__(self.__class__)
|
|
220
|
+
for k, v in self.__dict__.items():
|
|
221
|
+
if k == "_store":
|
|
222
|
+
out.__dict__[k] = copy.deepcopy(v, memo)
|
|
223
|
+
else:
|
|
224
|
+
out.__dict__[k] = copy.deepcopy(v, memo)
|
|
225
|
+
setattr(out._store, "_parent", out)
|
|
226
|
+
return out
|
|
227
|
+
|
|
228
|
+
def __getstate__(self) -> Dict[str, Any]:
|
|
229
|
+
return self.__dict__.copy()
|
|
230
|
+
|
|
231
|
+
def __setstate__(self, mapping: Dict[str, Any]):
|
|
232
|
+
import weakref
|
|
233
|
+
|
|
234
|
+
for key, value in mapping.items():
|
|
235
|
+
self.__dict__[key] = value
|
|
236
|
+
if "_store" in self.__dict__ and self._store is not None:
|
|
237
|
+
self._store.__dict__["_parent"] = weakref.ref(self)
|
|
238
|
+
|
|
239
|
+
def clone(self) -> "Data":
|
|
240
|
+
return copy.deepcopy(self)
|
|
241
|
+
|
|
242
|
+
def stores_as(self, data: "Data") -> "Data":
|
|
243
|
+
return self
|
|
244
|
+
|
|
245
|
+
@property
|
|
246
|
+
def stores(self) -> List[BaseStorage]:
|
|
247
|
+
return [self._store]
|
|
248
|
+
|
|
249
|
+
@property
|
|
250
|
+
def node_stores(self) -> List[NodeStorage]:
|
|
251
|
+
return [self._store]
|
|
252
|
+
|
|
253
|
+
@property
|
|
254
|
+
def edge_stores(self) -> List[EdgeStorage]:
|
|
255
|
+
return [self._store]
|
|
256
|
+
|
|
257
|
+
def get(self, key: str, default: Any = None) -> Any:
|
|
258
|
+
return self._store.get(key, default)
|
|
259
|
+
|
|
260
|
+
def __call__(self, *args: str) -> Iterator[Tuple[str, Any]]:
|
|
261
|
+
yield from self._store.items(*args)
|
|
262
|
+
|
|
263
|
+
@property
|
|
264
|
+
def num_features(self) -> int:
|
|
265
|
+
return self._store.num_features
|
|
266
|
+
|
|
267
|
+
@property
|
|
268
|
+
def num_node_features(self) -> int:
|
|
269
|
+
return self._store.num_node_features
|
|
270
|
+
|
|
271
|
+
@property
|
|
272
|
+
def num_edge_features(self) -> int:
|
|
273
|
+
return self._store.num_edge_features
|
|
274
|
+
|
|
275
|
+
@property
|
|
276
|
+
def num_node_types(self) -> int:
|
|
277
|
+
node_type = self.get("node_type")
|
|
278
|
+
return int(np.max(ops.convert_to_numpy(node_type))) + 1 if is_tensor_like(node_type) else 1
|
|
279
|
+
|
|
280
|
+
@property
|
|
281
|
+
def num_edge_types(self) -> int:
|
|
282
|
+
edge_type = self.get("edge_type")
|
|
283
|
+
return int(np.max(ops.convert_to_numpy(edge_type))) + 1 if is_tensor_like(edge_type) else 1
|
|
284
|
+
|
|
285
|
+
@property
|
|
286
|
+
def num_classes(self) -> Optional[int]:
|
|
287
|
+
y = self.get("y")
|
|
288
|
+
if y is not None and is_tensor_like(y):
|
|
289
|
+
y_np = ops.convert_to_numpy(y)
|
|
290
|
+
if np.issubdtype(y_np.dtype, np.integer):
|
|
291
|
+
return int(np.max(y_np)) + 1
|
|
292
|
+
return None
|
|
293
|
+
|
|
294
|
+
def is_directed(self) -> bool:
|
|
295
|
+
return self._store.is_directed()
|
|
296
|
+
|
|
297
|
+
def is_undirected(self) -> bool:
|
|
298
|
+
return self._store.is_undirected()
|
|
299
|
+
|
|
300
|
+
def has_self_loops(self) -> bool:
|
|
301
|
+
return self._store.has_self_loops()
|
|
302
|
+
|
|
303
|
+
def has_isolated_nodes(self) -> bool:
|
|
304
|
+
return self._store.has_isolated_nodes()
|
|
305
|
+
|
|
306
|
+
def is_coalesced(self) -> bool:
|
|
307
|
+
return self._store.is_coalesced()
|
|
308
|
+
|
|
309
|
+
def coalesce(self, reduce: str = "add"):
|
|
310
|
+
self._store.coalesce(reduce=reduce)
|
|
311
|
+
return self
|
|
312
|
+
|
|
313
|
+
def __inc__(self, key: str, value: Any, *args, **kwargs) -> Any:
|
|
314
|
+
if "batch" in key:
|
|
315
|
+
return int(value.max()) + 1 if is_tensor_like(value) and value.ndim > 0 and value.shape[0] > 0 else 0
|
|
316
|
+
if "index" in key or "face" in key:
|
|
317
|
+
return self.num_nodes or 0
|
|
318
|
+
return 0
|
|
319
|
+
|
|
320
|
+
def __cat_dim__(self, key: str, value: Any, *args, **kwargs) -> int:
|
|
321
|
+
if key in ("edge_index", "adj_t", "face"): # [2 or 3, num_edges / num_faces]
|
|
322
|
+
return -1
|
|
323
|
+
if is_tensor_like(value) and len(get_shape(value)) == 2 and get_shape(value)[0] == 2 and "index" in key:
|
|
324
|
+
return -1
|
|
325
|
+
return 0
|
|
326
|
+
|
|
327
|
+
def to_dict(self) -> Dict[str, Any]:
|
|
328
|
+
return self._store.to_dict()
|
|
329
|
+
|
|
330
|
+
def to_namedtuple(self) -> NamedTuple:
|
|
331
|
+
fields = sorted(list(self.keys()))
|
|
332
|
+
DataTuple = collections.namedtuple("DataTuple", fields)
|
|
333
|
+
return DataTuple(**{f: self[f] for f in fields})
|
|
334
|
+
|
|
335
|
+
@classmethod
|
|
336
|
+
def from_dict(cls, mapping: Dict[str, Any]) -> "Data":
|
|
337
|
+
return cls(**mapping)
|
|
338
|
+
|
|
339
|
+
def apply(self, func: Callable, *keys: str) -> "Data":
|
|
340
|
+
self._store.apply(func, *keys)
|
|
341
|
+
return self
|
|
342
|
+
|
|
343
|
+
def apply_(self, func: Callable, *keys: str) -> "Data":
|
|
344
|
+
self._store.apply_(func, *keys)
|
|
345
|
+
return self
|
|
346
|
+
|
|
347
|
+
def to(self, *args, **kwargs) -> "Data":
|
|
348
|
+
self._store.to(*args, **kwargs)
|
|
349
|
+
return self
|
|
350
|
+
|
|
351
|
+
def to_backend(self, backend: Optional[str] = None) -> "Data":
|
|
352
|
+
self._store.to_backend(backend)
|
|
353
|
+
return self
|
|
354
|
+
|
|
355
|
+
def cpu(self) -> "Data":
|
|
356
|
+
self._store.cpu()
|
|
357
|
+
return self
|
|
358
|
+
|
|
359
|
+
def cuda(self) -> "Data":
|
|
360
|
+
self._store.cuda()
|
|
361
|
+
return self
|
|
362
|
+
|
|
363
|
+
def requires_grad_(self, *keys: str) -> "Data":
|
|
364
|
+
self._store.requires_grad_(*keys)
|
|
365
|
+
return self
|
|
366
|
+
|
|
367
|
+
def contiguous(self, *keys: str) -> "Data":
|
|
368
|
+
self._store.contiguous(*keys)
|
|
369
|
+
return self
|
|
370
|
+
|
|
371
|
+
def subgraph(self, subset) -> "Data":
|
|
372
|
+
"""Returns the induced subgraph for subset nodes."""
|
|
373
|
+
data = copy.copy(self)
|
|
374
|
+
num_nodes = self.num_nodes
|
|
375
|
+
sub_edge_index, sub_edge_attr = subgraph(
|
|
376
|
+
subset,
|
|
377
|
+
self.edge_index,
|
|
378
|
+
edge_attr=self.get("edge_attr"),
|
|
379
|
+
relabel_nodes=True,
|
|
380
|
+
num_nodes=num_nodes,
|
|
381
|
+
)
|
|
382
|
+
data.edge_index = sub_edge_index
|
|
383
|
+
if sub_edge_attr is not None:
|
|
384
|
+
data.edge_attr = sub_edge_attr
|
|
385
|
+
subset_np = ops.convert_to_numpy(subset)
|
|
386
|
+
if subset_np.dtype == bool:
|
|
387
|
+
indices = np.where(subset_np)[0]
|
|
388
|
+
else:
|
|
389
|
+
indices = subset_np
|
|
390
|
+
|
|
391
|
+
for key in self.node_attrs():
|
|
392
|
+
val = self[key]
|
|
393
|
+
if is_tensor_like(val):
|
|
394
|
+
data[key] = ops.take(val, indices, axis=self.__cat_dim__(key, val))
|
|
395
|
+
if "num_nodes" in self._store: # an explicitly stored node count must shrink too
|
|
396
|
+
data.num_nodes = int(len(indices))
|
|
397
|
+
return data
|
|
398
|
+
|
|
399
|
+
def edge_subgraph(self, subset) -> "Data":
|
|
400
|
+
"""Returns the graph with only the edges in ``subset`` (a boolean edge mask or edge
|
|
401
|
+
indices). All nodes are kept; every edge-level attribute is filtered.
|
|
402
|
+
|
|
403
|
+
Example:
|
|
404
|
+
```python
|
|
405
|
+
import numpy as np
|
|
406
|
+
from k3_node.data import Data
|
|
407
|
+
|
|
408
|
+
data = Data(edge_index=np.array([[0, 1, 2], [1, 2, 0]]), edge_type=np.array([0, 1, 0]), num_nodes=3)
|
|
409
|
+
train = data.edge_subgraph(np.array([True, False, True]))
|
|
410
|
+
print(tuple(train.edge_index.shape), train.num_nodes) # (2, 2) 3
|
|
411
|
+
```
|
|
412
|
+
"""
|
|
413
|
+
subset_np = np.asarray(ops.convert_to_numpy(subset))
|
|
414
|
+
indices = np.where(subset_np)[0] if subset_np.dtype == bool else subset_np
|
|
415
|
+
data = copy.copy(self)
|
|
416
|
+
for key in self.edge_attrs():
|
|
417
|
+
val = self[key]
|
|
418
|
+
if is_tensor_like(val):
|
|
419
|
+
data[key] = ops.take(val, indices, axis=self.__cat_dim__(key, val))
|
|
420
|
+
return data
|
|
421
|
+
|
|
422
|
+
def to_heterogeneous(self, node_type: str = "0", edge_type: Tuple[str, str, str] = ("0", "0", "0")):
|
|
423
|
+
from k3_node.data.hetero_data import HeteroData
|
|
424
|
+
|
|
425
|
+
hetero = HeteroData()
|
|
426
|
+
if hasattr(self, "edge_type") and self.edge_type is not None:
|
|
427
|
+
edge_type_np = ops.convert_to_numpy(self.edge_type)
|
|
428
|
+
unique_edge_types = np.unique(edge_type_np)
|
|
429
|
+
hetero[node_type].x = self.x
|
|
430
|
+
for et in unique_edge_types:
|
|
431
|
+
mask = edge_type_np == et
|
|
432
|
+
sub_edge_index = self.edge_index[:, mask]
|
|
433
|
+
hetero[node_type, str(et), node_type].edge_index = sub_edge_index
|
|
434
|
+
if "edge_attr" in self and self.edge_attr is not None:
|
|
435
|
+
hetero[node_type, str(et), node_type].edge_attr = self.edge_attr[mask]
|
|
436
|
+
else:
|
|
437
|
+
for k in self.node_attrs():
|
|
438
|
+
hetero[node_type][k] = self[k]
|
|
439
|
+
for k in self.edge_attrs():
|
|
440
|
+
hetero[edge_type][k] = self[k]
|
|
441
|
+
return hetero
|
|
442
|
+
|
|
443
|
+
def validate(self, raise_on_error: bool = True) -> bool:
|
|
444
|
+
num_nodes = self.num_nodes
|
|
445
|
+
if "edge_index" in self and self.edge_index is not None:
|
|
446
|
+
edge_shape = get_shape(self.edge_index)
|
|
447
|
+
if len(edge_shape) != 2 or edge_shape[0] != 2:
|
|
448
|
+
msg = f"'edge_index' must have shape [2, num_edges], got {edge_shape}"
|
|
449
|
+
if raise_on_error:
|
|
450
|
+
raise ValueError(msg)
|
|
451
|
+
warnings.warn(msg)
|
|
452
|
+
return False
|
|
453
|
+
if num_nodes is not None and edge_shape[1] > 0:
|
|
454
|
+
edge_max = int(np.max(ops.convert_to_numpy(self.edge_index)))
|
|
455
|
+
if edge_max >= num_nodes:
|
|
456
|
+
msg = f"'edge_index' references node {edge_max}, but num_nodes is {num_nodes}"
|
|
457
|
+
if raise_on_error:
|
|
458
|
+
raise ValueError(msg)
|
|
459
|
+
warnings.warn(msg)
|
|
460
|
+
return False
|
|
461
|
+
return True
|
|
462
|
+
|
|
463
|
+
@property
|
|
464
|
+
def inputs(self):
|
|
465
|
+
r"""Returns the tuple of input tensors ``(x, edge_index)`` or ``(x, edge_index, edge_attr)``."""
|
|
466
|
+
if hasattr(self, "edge_attr") and self.edge_attr is not None:
|
|
467
|
+
return (self.x, self.edge_index, self.edge_attr)
|
|
468
|
+
return (self.x, self.edge_index)
|
|
469
|
+
|
|
470
|
+
def to_generator(self, mask: Optional[str] = "train_mask", repeat: bool = True):
|
|
471
|
+
r"""Generates tuples of ((x, edge_index), y, mask) or ((x, edge_index, edge_attr), y, mask)
|
|
472
|
+
ready for training directly with Keras `model.fit()`.
|
|
473
|
+
|
|
474
|
+
Args:
|
|
475
|
+
mask (str, optional): The name of the mask attribute (e.g. ``'train_mask'``,
|
|
476
|
+
``'val_mask'``, ``'test_mask'``) to use as sample_weight for loss masking.
|
|
477
|
+
If :obj:`None`, no mask is applied. (default: ``'train_mask'``)
|
|
478
|
+
repeat (bool, optional): Whether to yield infinitely for Keras generator training.
|
|
479
|
+
(default: :obj:`True`)
|
|
480
|
+
"""
|
|
481
|
+
x = ops.convert_to_tensor(self.x, dtype="float32") if self.x is not None else None
|
|
482
|
+
edge_index = ops.convert_to_tensor(self.edge_index, dtype="int64") if self.edge_index is not None else None
|
|
483
|
+
|
|
484
|
+
inputs = (x, edge_index)
|
|
485
|
+
if hasattr(self, "edge_attr") and self.edge_attr is not None:
|
|
486
|
+
edge_attr = ops.convert_to_tensor(self.edge_attr, dtype="float32")
|
|
487
|
+
inputs = (x, edge_index, edge_attr)
|
|
488
|
+
|
|
489
|
+
y = ops.convert_to_tensor(self.y, dtype="int64") if self.y is not None else None
|
|
490
|
+
|
|
491
|
+
sample_weight = None
|
|
492
|
+
if mask is not None and hasattr(self, mask) and getattr(self, mask) is not None:
|
|
493
|
+
sample_weight = ops.cast(getattr(self, mask), "float32")
|
|
494
|
+
|
|
495
|
+
while True:
|
|
496
|
+
if sample_weight is not None:
|
|
497
|
+
yield inputs, y, sample_weight
|
|
498
|
+
elif y is not None:
|
|
499
|
+
yield inputs, y
|
|
500
|
+
else:
|
|
501
|
+
yield inputs
|
|
502
|
+
if not repeat:
|
|
503
|
+
break
|
|
504
|
+
|
|
505
|
+
def accuracy(self, logits_or_pred: Any, mask: Optional[str] = "test_mask") -> float:
|
|
506
|
+
r"""Convenience method to calculate classification accuracy.
|
|
507
|
+
|
|
508
|
+
Args:
|
|
509
|
+
logits_or_pred: Model prediction output (either class logits or predicted labels).
|
|
510
|
+
mask (str, optional): The mask attribute (e.g. ``'test_mask'``) to evaluate on.
|
|
511
|
+
If :obj:`None`, evaluates across all nodes. (default: ``'test_mask'``)
|
|
512
|
+
"""
|
|
513
|
+
pred = logits_or_pred
|
|
514
|
+
shape = ops.shape(pred)
|
|
515
|
+
if len(shape) > 1 and shape[-1] > 1:
|
|
516
|
+
pred = ops.argmax(pred, axis=-1)
|
|
517
|
+
|
|
518
|
+
y = self.y
|
|
519
|
+
if mask is not None and hasattr(self, mask) and getattr(self, mask) is not None:
|
|
520
|
+
m = getattr(self, mask)
|
|
521
|
+
pred = pred[m]
|
|
522
|
+
y = y[m]
|
|
523
|
+
|
|
524
|
+
correct = ops.cast(ops.cast(pred, "int64") == ops.cast(y, "int64"), "float32")
|
|
525
|
+
return float(ops.convert_to_numpy(ops.mean(correct)))
|
|
526
|
+
|
|
527
|
+
def __repr__(self) -> str:
|
|
528
|
+
cls = self.__class__.__name__
|
|
529
|
+
attrs = [size_repr(k, v) for k, v in self._store.items()]
|
|
530
|
+
info = ", ".join(attrs)
|
|
531
|
+
return f"{cls}({info})"
|
|
532
|
+
|
k3_node/data/database.py
ADDED
|
@@ -0,0 +1,154 @@
|
|
|
1
|
+
import io
|
|
2
|
+
import pickle
|
|
3
|
+
import sqlite3
|
|
4
|
+
from abc import ABC, abstractmethod
|
|
5
|
+
from typing import Any, Dict, List, Optional, Sequence, Union
|
|
6
|
+
|
|
7
|
+
Schema = Any
|
|
8
|
+
|
|
9
|
+
|
|
10
|
+
class Database(ABC):
|
|
11
|
+
"""Base class for key/value and index-based graph databases."""
|
|
12
|
+
|
|
13
|
+
def __init__(self, schema: Schema = object):
|
|
14
|
+
self.schema = schema
|
|
15
|
+
|
|
16
|
+
@abstractmethod
|
|
17
|
+
def connect(self):
|
|
18
|
+
pass
|
|
19
|
+
|
|
20
|
+
@abstractmethod
|
|
21
|
+
def close(self):
|
|
22
|
+
pass
|
|
23
|
+
|
|
24
|
+
@abstractmethod
|
|
25
|
+
def insert(self, index: int, data: Any):
|
|
26
|
+
pass
|
|
27
|
+
|
|
28
|
+
def multi_insert(self, indices: Sequence[int], data_list: Sequence[Any]):
|
|
29
|
+
for idx, d in zip(indices, data_list):
|
|
30
|
+
self.insert(idx, d)
|
|
31
|
+
|
|
32
|
+
@abstractmethod
|
|
33
|
+
def get(self, index: int) -> Any:
|
|
34
|
+
pass
|
|
35
|
+
|
|
36
|
+
def multi_get(self, indices: Sequence[int]) -> List[Any]:
|
|
37
|
+
return [self.get(idx) for idx in indices]
|
|
38
|
+
|
|
39
|
+
@abstractmethod
|
|
40
|
+
def __len__(self) -> int:
|
|
41
|
+
pass
|
|
42
|
+
|
|
43
|
+
def __getitem__(self, idx: Any) -> Any:
|
|
44
|
+
if isinstance(idx, int):
|
|
45
|
+
return self.get(idx)
|
|
46
|
+
elif isinstance(idx, slice):
|
|
47
|
+
start = idx.start or 0
|
|
48
|
+
stop = idx.stop or len(self)
|
|
49
|
+
step = idx.step or 1
|
|
50
|
+
return self.multi_get(range(start, stop, step))
|
|
51
|
+
elif isinstance(idx, (list, tuple)):
|
|
52
|
+
return self.multi_get(idx)
|
|
53
|
+
else:
|
|
54
|
+
return self.get(int(idx))
|
|
55
|
+
|
|
56
|
+
def __setitem__(self, idx: Any, value: Any):
|
|
57
|
+
if isinstance(idx, int):
|
|
58
|
+
self.insert(idx, value)
|
|
59
|
+
elif isinstance(idx, slice):
|
|
60
|
+
start = idx.start or 0
|
|
61
|
+
stop = idx.stop or len(self)
|
|
62
|
+
step = idx.step or 1
|
|
63
|
+
indices = list(range(start, stop, step))
|
|
64
|
+
self.multi_insert(indices, value)
|
|
65
|
+
elif isinstance(idx, (list, tuple)):
|
|
66
|
+
self.multi_insert(idx, value)
|
|
67
|
+
else:
|
|
68
|
+
self.insert(int(idx), value)
|
|
69
|
+
|
|
70
|
+
|
|
71
|
+
class SQLiteDatabase(Database):
|
|
72
|
+
"""SQLite-backed persistent database."""
|
|
73
|
+
|
|
74
|
+
def __init__(self, path: str, name: str = "data", schema: Schema = object):
|
|
75
|
+
super().__init__(schema)
|
|
76
|
+
self.path = path
|
|
77
|
+
self.name = name
|
|
78
|
+
self._conn: Optional[sqlite3.Connection] = None
|
|
79
|
+
self.connect()
|
|
80
|
+
|
|
81
|
+
def connect(self):
|
|
82
|
+
self._conn = sqlite3.connect(self.path)
|
|
83
|
+
with self._conn:
|
|
84
|
+
self._conn.execute(
|
|
85
|
+
f"CREATE TABLE IF NOT EXISTS {self.name} (id INTEGER PRIMARY KEY, val BLOB)"
|
|
86
|
+
)
|
|
87
|
+
|
|
88
|
+
def close(self):
|
|
89
|
+
if self._conn is not None:
|
|
90
|
+
self._conn.close()
|
|
91
|
+
self._conn = None
|
|
92
|
+
|
|
93
|
+
def insert(self, index: int, data: Any):
|
|
94
|
+
buf = io.BytesIO()
|
|
95
|
+
pickle.dump(data, buf)
|
|
96
|
+
raw = buf.getvalue()
|
|
97
|
+
with self._conn:
|
|
98
|
+
self._conn.execute(
|
|
99
|
+
f"INSERT OR REPLACE INTO {self.name} (id, val) VALUES (?, ?)",
|
|
100
|
+
(index, raw),
|
|
101
|
+
)
|
|
102
|
+
|
|
103
|
+
def multi_insert(self, indices: Sequence[int], data_list: Sequence[Any]):
|
|
104
|
+
rows = []
|
|
105
|
+
for idx, d in zip(indices, data_list):
|
|
106
|
+
buf = io.BytesIO()
|
|
107
|
+
pickle.dump(d, buf)
|
|
108
|
+
rows.append((idx, buf.getvalue()))
|
|
109
|
+
with self._conn:
|
|
110
|
+
self._conn.executemany(
|
|
111
|
+
f"INSERT OR REPLACE INTO {self.name} (id, val) VALUES (?, ?)",
|
|
112
|
+
rows,
|
|
113
|
+
)
|
|
114
|
+
|
|
115
|
+
def get(self, index: int) -> Any:
|
|
116
|
+
cursor = self._conn.execute(
|
|
117
|
+
f"SELECT val FROM {self.name} WHERE id = ?", (index,)
|
|
118
|
+
)
|
|
119
|
+
row = cursor.fetchone()
|
|
120
|
+
if row is None:
|
|
121
|
+
raise KeyError(f"Index {index} not found in database")
|
|
122
|
+
return pickle.loads(row[0])
|
|
123
|
+
|
|
124
|
+
def multi_get(self, indices: Sequence[int]) -> List[Any]:
|
|
125
|
+
return [self.get(idx) for idx in indices]
|
|
126
|
+
|
|
127
|
+
def __len__(self) -> int:
|
|
128
|
+
cursor = self._conn.execute(f"SELECT COUNT(*) FROM {self.name}")
|
|
129
|
+
return cursor.fetchone()[0]
|
|
130
|
+
|
|
131
|
+
|
|
132
|
+
class RocksDatabase(Database):
|
|
133
|
+
"""RocksDB-backed database stub."""
|
|
134
|
+
|
|
135
|
+
def __init__(self, path: str, schema: Schema = object):
|
|
136
|
+
super().__init__(schema)
|
|
137
|
+
self.path = path
|
|
138
|
+
raise NotImplementedError("RocksDatabase requires rocksdb C++ binding; use SQLiteDatabase instead.")
|
|
139
|
+
|
|
140
|
+
def connect(self):
|
|
141
|
+
pass
|
|
142
|
+
|
|
143
|
+
def close(self):
|
|
144
|
+
pass
|
|
145
|
+
|
|
146
|
+
def insert(self, index: int, data: Any):
|
|
147
|
+
pass
|
|
148
|
+
|
|
149
|
+
def get(self, index: int) -> Any:
|
|
150
|
+
pass
|
|
151
|
+
|
|
152
|
+
def __len__(self) -> int:
|
|
153
|
+
return 0
|
|
154
|
+
|