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,188 @@
|
|
|
1
|
+
from typing import Optional
|
|
2
|
+
|
|
3
|
+
import numpy as np
|
|
4
|
+
from keras import ops
|
|
5
|
+
|
|
6
|
+
try:
|
|
7
|
+
import torch
|
|
8
|
+
import torch.utils.data
|
|
9
|
+
BaseDataLoader = torch.utils.data.DataLoader
|
|
10
|
+
except ImportError:
|
|
11
|
+
torch = None
|
|
12
|
+
BaseDataLoader = object
|
|
13
|
+
|
|
14
|
+
from k3_node.data import Data
|
|
15
|
+
from k3_node.loader.keras_dataset import loader_bases
|
|
16
|
+
|
|
17
|
+
|
|
18
|
+
def _np(x):
|
|
19
|
+
return np.asarray(ops.convert_to_numpy(x))
|
|
20
|
+
|
|
21
|
+
|
|
22
|
+
class _SamplingSteps:
|
|
23
|
+
"""The sampler's "dataset": item ``i`` is a freshly sampled ``(node_idx, edge_idx)`` pair."""
|
|
24
|
+
|
|
25
|
+
def __init__(self, sampler):
|
|
26
|
+
self.sampler = sampler
|
|
27
|
+
|
|
28
|
+
def __len__(self):
|
|
29
|
+
return self.sampler.num_steps
|
|
30
|
+
|
|
31
|
+
def __getitem__(self, idx):
|
|
32
|
+
return self.sampler._sample_subgraph()
|
|
33
|
+
|
|
34
|
+
|
|
35
|
+
class GraphSAINTSampler(*loader_bases(BaseDataLoader)):
|
|
36
|
+
r"""The GraphSAINT sampler base class from the `"GraphSAINT: Graph Sampling Based Inductive
|
|
37
|
+
Learning Method" <https://arxiv.org/abs/1907.04931>`_ paper. Every step samples a set of
|
|
38
|
+
nodes and yields the subgraph they induce.
|
|
39
|
+
|
|
40
|
+
With ``sample_coverage > 0``, normalization statistics are estimated beforehand by sampling
|
|
41
|
+
about ``sample_coverage`` times every node: ``node_norm`` (to weight each node's loss) and
|
|
42
|
+
``edge_norm`` (to weight each edge's message) are added to every subgraph.
|
|
43
|
+
|
|
44
|
+
Args:
|
|
45
|
+
data (Data): The graph data object.
|
|
46
|
+
batch_size (int): The approximate number of samples per batch (see the subclasses).
|
|
47
|
+
num_steps (int, optional): The number of iterations per epoch. (default: ``1``)
|
|
48
|
+
sample_coverage (int): How many samples per node to compute the normalization
|
|
49
|
+
statistics with; ``0`` skips them. (default: ``0``)
|
|
50
|
+
save_dir (str, optional): Unused; kept for API compatibility with PyG.
|
|
51
|
+
log (bool, optional): Unused; kept for API compatibility with PyG.
|
|
52
|
+
**kwargs (optional): Additional arguments of :class:`torch.utils.data.DataLoader`.
|
|
53
|
+
"""
|
|
54
|
+
|
|
55
|
+
def __init__(self, data: Data, batch_size: int, num_steps: int = 1, sample_coverage: int = 0,
|
|
56
|
+
save_dir: Optional[str] = None, log: bool = True, **kwargs):
|
|
57
|
+
kwargs.pop('dataset', None)
|
|
58
|
+
kwargs.pop('collate_fn', None)
|
|
59
|
+
kwargs.pop('shuffle', None)
|
|
60
|
+
assert data.edge_index is not None
|
|
61
|
+
|
|
62
|
+
self.num_steps = num_steps
|
|
63
|
+
self._batch_size = batch_size
|
|
64
|
+
self.sample_coverage = sample_coverage
|
|
65
|
+
self.save_dir = save_dir
|
|
66
|
+
self.log = log
|
|
67
|
+
self.data = data
|
|
68
|
+
self.N = data.num_nodes
|
|
69
|
+
edge_index = _np(data.edge_index).astype(np.int64)
|
|
70
|
+
self.E = edge_index.shape[1]
|
|
71
|
+
# CSR by source node: the edges leaving node i are perm[rowptr[i]:rowptr[i+1]]
|
|
72
|
+
self._row, self._col = edge_index
|
|
73
|
+
self._perm = np.argsort(self._row, kind="stable")
|
|
74
|
+
self._rowptr = np.concatenate([[0], np.cumsum(np.bincount(self._row, minlength=self.N))])
|
|
75
|
+
|
|
76
|
+
steps = _SamplingSteps(self)
|
|
77
|
+
if torch is not None:
|
|
78
|
+
super().__init__(steps, batch_size=1, collate_fn=self._collate, **kwargs)
|
|
79
|
+
else:
|
|
80
|
+
self.dataset = steps
|
|
81
|
+
self.batch_size = 1
|
|
82
|
+
self.collate_fn = self._collate
|
|
83
|
+
|
|
84
|
+
if self.sample_coverage > 0:
|
|
85
|
+
self.node_norm, self.edge_norm = self._compute_norm()
|
|
86
|
+
|
|
87
|
+
def _sample_nodes(self, batch_size: int) -> np.ndarray:
|
|
88
|
+
raise NotImplementedError
|
|
89
|
+
|
|
90
|
+
def _sample_subgraph(self):
|
|
91
|
+
node_idx = np.unique(self._sample_nodes(self._batch_size))
|
|
92
|
+
in_sample = np.zeros(self.N, dtype=bool)
|
|
93
|
+
in_sample[node_idx] = True
|
|
94
|
+
edge_idx = np.nonzero(in_sample[self._row] & in_sample[self._col])[0]
|
|
95
|
+
return node_idx, edge_idx
|
|
96
|
+
|
|
97
|
+
def _collate(self, data_list):
|
|
98
|
+
node_idx, edge_idx = data_list[0]
|
|
99
|
+
new_id = np.full(self.N, -1, dtype=np.int64)
|
|
100
|
+
new_id[node_idx] = np.arange(len(node_idx))
|
|
101
|
+
|
|
102
|
+
data = Data()
|
|
103
|
+
data.num_nodes = len(node_idx)
|
|
104
|
+
data.edge_index = ops.convert_to_tensor(
|
|
105
|
+
np.stack([new_id[self._row[edge_idx]], new_id[self._col[edge_idx]]]), dtype="int64")
|
|
106
|
+
for key, item in self.data.items():
|
|
107
|
+
if key in ('edge_index', 'num_nodes'):
|
|
108
|
+
continue
|
|
109
|
+
shape = getattr(item, 'shape', None)
|
|
110
|
+
if shape is not None and len(shape) > 0 and shape[0] == self.N:
|
|
111
|
+
data[key] = ops.take(item, node_idx, axis=0)
|
|
112
|
+
elif shape is not None and len(shape) > 0 and shape[0] == self.E:
|
|
113
|
+
data[key] = ops.take(item, edge_idx, axis=0)
|
|
114
|
+
else:
|
|
115
|
+
data[key] = item
|
|
116
|
+
if self.sample_coverage > 0:
|
|
117
|
+
data.node_norm = ops.convert_to_tensor(self.node_norm[node_idx])
|
|
118
|
+
data.edge_norm = ops.convert_to_tensor(self.edge_norm[edge_idx])
|
|
119
|
+
return data
|
|
120
|
+
|
|
121
|
+
def _compute_norm(self):
|
|
122
|
+
node_count = np.zeros(self.N, dtype=np.float32)
|
|
123
|
+
edge_count = np.zeros(self.E, dtype=np.float32)
|
|
124
|
+
num_samples = total_sampled_nodes = 0
|
|
125
|
+
while total_sampled_nodes < self.N * self.sample_coverage:
|
|
126
|
+
for _ in range(self.num_steps):
|
|
127
|
+
node_idx, edge_idx = self._sample_subgraph()
|
|
128
|
+
node_count[node_idx] += 1
|
|
129
|
+
edge_count[edge_idx] += 1
|
|
130
|
+
total_sampled_nodes += len(node_idx)
|
|
131
|
+
num_samples += self.num_steps
|
|
132
|
+
|
|
133
|
+
with np.errstate(divide='ignore', invalid='ignore'):
|
|
134
|
+
edge_norm = np.clip(node_count[self._row] / edge_count, 0, 1e4)
|
|
135
|
+
edge_norm[np.isnan(edge_norm)] = 0.1
|
|
136
|
+
node_count[node_count == 0] = 0.1
|
|
137
|
+
node_norm = num_samples / node_count / self.N
|
|
138
|
+
return node_norm.astype(np.float32), edge_norm.astype(np.float32)
|
|
139
|
+
|
|
140
|
+
def _random_walk(self, start: np.ndarray, walk_length: int) -> np.ndarray:
|
|
141
|
+
walks, cur = [start], start
|
|
142
|
+
for _ in range(walk_length):
|
|
143
|
+
deg = self._rowptr[cur + 1] - self._rowptr[cur]
|
|
144
|
+
offset = np.floor(np.random.rand(len(cur)) * np.maximum(deg, 1)).astype(np.int64)
|
|
145
|
+
nxt = self._col[self._perm[np.minimum(self._rowptr[cur] + offset, self.E - 1)]]
|
|
146
|
+
cur = np.where(deg > 0, nxt, cur) # nodes without neighbors stay put
|
|
147
|
+
walks.append(cur)
|
|
148
|
+
return np.stack(walks, axis=1)
|
|
149
|
+
|
|
150
|
+
|
|
151
|
+
class GraphSAINTNodeSampler(GraphSAINTSampler):
|
|
152
|
+
r"""The GraphSAINT node sampler: samples ``batch_size`` nodes, each with probability
|
|
153
|
+
proportional to its out-degree."""
|
|
154
|
+
|
|
155
|
+
def _sample_nodes(self, batch_size: int) -> np.ndarray:
|
|
156
|
+
return self._row[np.random.randint(0, self.E, size=batch_size)]
|
|
157
|
+
|
|
158
|
+
|
|
159
|
+
class GraphSAINTEdgeSampler(GraphSAINTSampler):
|
|
160
|
+
r"""The GraphSAINT edge sampler: samples ``batch_size`` edges, each with probability
|
|
161
|
+
proportional to :math:`1 / \deg(u) + 1 / \deg(v)`, and keeps their endpoints."""
|
|
162
|
+
|
|
163
|
+
def _sample_nodes(self, batch_size: int) -> np.ndarray:
|
|
164
|
+
out_deg = np.maximum(np.bincount(self._row, minlength=self.N), 1)
|
|
165
|
+
in_deg = np.maximum(np.bincount(self._col, minlength=self.N), 1)
|
|
166
|
+
prob = 1.0 / in_deg[self._row] + 1.0 / out_deg[self._col]
|
|
167
|
+
# Weighted sampling without replacement (exponential keys, as in PyG)
|
|
168
|
+
keys = np.log(np.random.rand(self.E)) / (prob + 1e-10)
|
|
169
|
+
edge_sample = np.argsort(-keys)[:batch_size]
|
|
170
|
+
return np.concatenate([self._col[edge_sample], self._row[edge_sample]])
|
|
171
|
+
|
|
172
|
+
|
|
173
|
+
class GraphSAINTRandomWalkSampler(GraphSAINTSampler):
|
|
174
|
+
r"""The GraphSAINT random walk sampler: starts ``batch_size`` random walks of length
|
|
175
|
+
``walk_length`` and keeps the visited nodes.
|
|
176
|
+
|
|
177
|
+
Args:
|
|
178
|
+
walk_length (int): Length of each random walk.
|
|
179
|
+
"""
|
|
180
|
+
|
|
181
|
+
def __init__(self, data: Data, batch_size: int, walk_length: int, num_steps: int = 1,
|
|
182
|
+
sample_coverage: int = 0, save_dir: Optional[str] = None, log: bool = True, **kwargs):
|
|
183
|
+
self.walk_length = walk_length
|
|
184
|
+
super().__init__(data, batch_size, num_steps, sample_coverage, save_dir, log, **kwargs)
|
|
185
|
+
|
|
186
|
+
def _sample_nodes(self, batch_size: int) -> np.ndarray:
|
|
187
|
+
start = np.random.randint(0, self.N, size=batch_size)
|
|
188
|
+
return self._random_walk(start, self.walk_length).reshape(-1)
|
|
@@ -0,0 +1,90 @@
|
|
|
1
|
+
from typing import Any, Callable, Dict, List, Optional, Tuple, Union
|
|
2
|
+
|
|
3
|
+
|
|
4
|
+
from k3_node.data import HeteroData
|
|
5
|
+
from k3_node.loader.node_loader import HeteroSamplerOutput, NodeLoader, NodeSamplerInput
|
|
6
|
+
from k3_node.loader.sampler_utils import sample_neighbors_hetero
|
|
7
|
+
|
|
8
|
+
|
|
9
|
+
class InternalHGTSampler:
|
|
10
|
+
r"""HGT balanced neighborhood sampling engine."""
|
|
11
|
+
def __init__(
|
|
12
|
+
self,
|
|
13
|
+
data: HeteroData,
|
|
14
|
+
num_samples: Union[List[int], Dict[str, List[int]]],
|
|
15
|
+
):
|
|
16
|
+
self.data = data
|
|
17
|
+
self.num_samples = num_samples
|
|
18
|
+
self.edge_permutation = None
|
|
19
|
+
|
|
20
|
+
def sample_from_nodes(self, input_data: NodeSamplerInput) -> HeteroSamplerOutput:
|
|
21
|
+
edge_index_dict = {}
|
|
22
|
+
for edge_type in self.data.edge_types:
|
|
23
|
+
canonical = self.data._to_canonical(*edge_type) if hasattr(self.data, '_to_canonical') else edge_type
|
|
24
|
+
edge_index_dict[canonical] = self.data[edge_type].edge_index
|
|
25
|
+
|
|
26
|
+
node_type = input_data.input_type or self.data.node_types[0]
|
|
27
|
+
seed_dict = {k: None for k in self.data.node_types}
|
|
28
|
+
seed_dict[node_type] = input_data.node
|
|
29
|
+
|
|
30
|
+
# Build num_neighbors dict for hetero sampling
|
|
31
|
+
if isinstance(self.num_samples, dict):
|
|
32
|
+
num_neighbors = {}
|
|
33
|
+
for e in self.data.edge_types:
|
|
34
|
+
dst = e[2]
|
|
35
|
+
can = self.data._to_canonical(*e) if hasattr(self.data, '_to_canonical') else e
|
|
36
|
+
num_neighbors[can] = self.num_samples.get(dst, [10])
|
|
37
|
+
else:
|
|
38
|
+
num_neighbors = self.num_samples
|
|
39
|
+
|
|
40
|
+
node_dict, row_dict, col_dict, edge_dict, n_counts, e_counts = sample_neighbors_hetero(
|
|
41
|
+
edge_index_dict=edge_index_dict,
|
|
42
|
+
seed_nodes_dict=seed_dict,
|
|
43
|
+
num_neighbors=num_neighbors,
|
|
44
|
+
replace=True,
|
|
45
|
+
subgraph_type='directional',
|
|
46
|
+
)
|
|
47
|
+
|
|
48
|
+
return HeteroSamplerOutput(
|
|
49
|
+
node=node_dict,
|
|
50
|
+
row=row_dict,
|
|
51
|
+
col=col_dict,
|
|
52
|
+
edge=edge_dict,
|
|
53
|
+
num_sampled_nodes=n_counts,
|
|
54
|
+
num_sampled_edges=e_counts,
|
|
55
|
+
metadata=(input_data.input_id, input_data.time),
|
|
56
|
+
)
|
|
57
|
+
|
|
58
|
+
|
|
59
|
+
class HGTLoader(NodeLoader):
|
|
60
|
+
r"""The Heterogeneous Graph Sampler from the "Heterogeneous Graph Transformer" paper.
|
|
61
|
+
|
|
62
|
+
Args:
|
|
63
|
+
data (HeteroData): The heterogeneous graph data object.
|
|
64
|
+
num_samples (List[int] or Dict[str, List[int]]): The number of nodes to sample per iteration.
|
|
65
|
+
input_nodes (str or Tuple[str, Tensor]): Seed node type and indices.
|
|
66
|
+
**kwargs (optional): Additional arguments of :class:`torch.utils.data.DataLoader`.
|
|
67
|
+
"""
|
|
68
|
+
def __init__(
|
|
69
|
+
self,
|
|
70
|
+
data: HeteroData,
|
|
71
|
+
num_samples: Union[List[int], Dict[str, List[int]]],
|
|
72
|
+
input_nodes: Union[str, Tuple[str, Optional[Any]]],
|
|
73
|
+
is_sorted: bool = False,
|
|
74
|
+
transform: Optional[Callable] = None,
|
|
75
|
+
transform_sampler_output: Optional[Callable] = None,
|
|
76
|
+
filter_per_worker: Optional[bool] = None,
|
|
77
|
+
**kwargs,
|
|
78
|
+
):
|
|
79
|
+
hgt_sampler = InternalHGTSampler(data, num_samples=num_samples)
|
|
80
|
+
|
|
81
|
+
super().__init__(
|
|
82
|
+
data=data,
|
|
83
|
+
node_sampler=hgt_sampler,
|
|
84
|
+
input_nodes=input_nodes,
|
|
85
|
+
transform=transform,
|
|
86
|
+
transform_sampler_output=transform_sampler_output,
|
|
87
|
+
filter_per_worker=filter_per_worker,
|
|
88
|
+
**kwargs,
|
|
89
|
+
)
|
|
90
|
+
|
|
@@ -0,0 +1,87 @@
|
|
|
1
|
+
from typing import Any, List, Optional, Union
|
|
2
|
+
|
|
3
|
+
import numpy as np
|
|
4
|
+
|
|
5
|
+
try:
|
|
6
|
+
import torch
|
|
7
|
+
from torch import Tensor
|
|
8
|
+
BaseWeightedRandomSampler = torch.utils.data.WeightedRandomSampler
|
|
9
|
+
except ImportError:
|
|
10
|
+
torch = None
|
|
11
|
+
Tensor = type(None)
|
|
12
|
+
BaseWeightedRandomSampler = object
|
|
13
|
+
|
|
14
|
+
from k3_node.data import Data, Dataset, InMemoryDataset
|
|
15
|
+
|
|
16
|
+
|
|
17
|
+
class ImbalancedSampler(BaseWeightedRandomSampler):
|
|
18
|
+
r"""A weighted random sampler that randomly samples elements according to class distribution.
|
|
19
|
+
|
|
20
|
+
Args:
|
|
21
|
+
dataset (Dataset or Data or Tensor): The dataset or class distribution from which to sample.
|
|
22
|
+
input_nodes (Tensor, optional): The indices of nodes used by the corresponding loader. (default: :obj:`None`)
|
|
23
|
+
num_samples (int, optional): The number of samples to draw for a single epoch. (default: :obj:`None`)
|
|
24
|
+
"""
|
|
25
|
+
def __init__(
|
|
26
|
+
self,
|
|
27
|
+
dataset: Union[Dataset, Data, List[Data], Any],
|
|
28
|
+
input_nodes: Optional[Any] = None,
|
|
29
|
+
num_samples: Optional[int] = None,
|
|
30
|
+
):
|
|
31
|
+
if isinstance(dataset, Data):
|
|
32
|
+
y = dataset.y
|
|
33
|
+
if hasattr(y, 'view'):
|
|
34
|
+
y = y.view(-1)
|
|
35
|
+
else:
|
|
36
|
+
y = np.asarray(y).reshape(-1)
|
|
37
|
+
if input_nodes is not None:
|
|
38
|
+
y = y[input_nodes]
|
|
39
|
+
|
|
40
|
+
elif torch is not None and isinstance(dataset, Tensor):
|
|
41
|
+
y = dataset.view(-1)
|
|
42
|
+
if input_nodes is not None:
|
|
43
|
+
y = y[input_nodes]
|
|
44
|
+
|
|
45
|
+
elif isinstance(dataset, InMemoryDataset):
|
|
46
|
+
y = dataset.y
|
|
47
|
+
if hasattr(y, 'view'):
|
|
48
|
+
y = y.view(-1)
|
|
49
|
+
else:
|
|
50
|
+
y = np.asarray(y).reshape(-1)
|
|
51
|
+
|
|
52
|
+
elif isinstance(dataset, (list, tuple)):
|
|
53
|
+
ys = [data.y for data in dataset]
|
|
54
|
+
if torch is not None and isinstance(ys[0], Tensor):
|
|
55
|
+
y = torch.cat(ys, dim=0).view(-1)
|
|
56
|
+
else:
|
|
57
|
+
y = np.concatenate([np.asarray(x).reshape(-1) for x in ys], axis=0)
|
|
58
|
+
else:
|
|
59
|
+
y = np.asarray(dataset).reshape(-1)
|
|
60
|
+
|
|
61
|
+
if torch is not None and not isinstance(y, Tensor):
|
|
62
|
+
y = torch.as_tensor(y, dtype=torch.long)
|
|
63
|
+
|
|
64
|
+
num_samples = (y.numel() if hasattr(y, 'numel') else len(y)) if num_samples is None else num_samples
|
|
65
|
+
|
|
66
|
+
if torch is not None and isinstance(y, Tensor):
|
|
67
|
+
bincount = y.bincount().float()
|
|
68
|
+
class_weight = 1.0 / bincount
|
|
69
|
+
weight = class_weight[y]
|
|
70
|
+
super().__init__(weight, num_samples, replacement=True)
|
|
71
|
+
else:
|
|
72
|
+
classes, counts = np.unique(y, return_counts=True)
|
|
73
|
+
class_weight = {c: 1.0 / count for c, count in zip(classes, counts)}
|
|
74
|
+
weight = np.array([class_weight[int(val)] for val in y], dtype=np.float64)
|
|
75
|
+
weight = weight / weight.sum()
|
|
76
|
+
self.weight = weight
|
|
77
|
+
self.num_samples = num_samples
|
|
78
|
+
self.replacement = True
|
|
79
|
+
|
|
80
|
+
def __iter__(self):
|
|
81
|
+
if torch is not None and hasattr(super(), '__iter__'):
|
|
82
|
+
return super().__iter__()
|
|
83
|
+
indices = np.random.choice(len(self.weight), size=self.num_samples, replace=True, p=self.weight)
|
|
84
|
+
return iter(indices.tolist())
|
|
85
|
+
|
|
86
|
+
def __len__(self):
|
|
87
|
+
return self.num_samples
|