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,31 @@
|
|
|
1
|
+
from keras import ops
|
|
2
|
+
|
|
3
|
+
from k3_node.layers.unpool import knn_interpolate
|
|
4
|
+
|
|
5
|
+
|
|
6
|
+
def test_knn_interpolate():
|
|
7
|
+
x = ops.convert_to_tensor(
|
|
8
|
+
[[1.0], [10.0], [100.0], [-1.0], [-10.0], [-100.0]], dtype="float32"
|
|
9
|
+
)
|
|
10
|
+
pos_x = ops.convert_to_tensor([
|
|
11
|
+
[-1.0, 0.0], [0.0, 0.0], [1.0, 0.0],
|
|
12
|
+
[-2.0, 0.0], [0.0, 0.0], [2.0, 0.0],
|
|
13
|
+
], dtype="float32")
|
|
14
|
+
pos_y = ops.convert_to_tensor([
|
|
15
|
+
[-1.0, -1.0], [1.0, 1.0], [-2.0, -2.0], [2.0, 2.0],
|
|
16
|
+
], dtype="float32")
|
|
17
|
+
batch_x = ops.convert_to_tensor([0, 0, 0, 1, 1, 1], dtype="int64")
|
|
18
|
+
batch_y = ops.convert_to_tensor([0, 0, 1, 1], dtype="int64")
|
|
19
|
+
|
|
20
|
+
y = knn_interpolate(x, pos_x, pos_y, batch_x, batch_y, k=2)
|
|
21
|
+
assert ops.shape(y) == (4, 1)
|
|
22
|
+
assert ops.convert_to_numpy(y).tolist() == [[4.0], [70.0], [-4.0], [-70.0]]
|
|
23
|
+
|
|
24
|
+
|
|
25
|
+
def test_knn_interpolate_no_batch():
|
|
26
|
+
x = ops.convert_to_tensor([[1.0], [10.0], [100.0]], dtype="float32")
|
|
27
|
+
pos_x = ops.convert_to_tensor([[-1.0, 0.0], [0.0, 0.0], [1.0, 0.0]], dtype="float32")
|
|
28
|
+
pos_y = ops.convert_to_tensor([[-1.0, -1.0], [1.0, 1.0]], dtype="float32")
|
|
29
|
+
|
|
30
|
+
y = knn_interpolate(x, pos_x, pos_y, k=2)
|
|
31
|
+
assert ops.shape(y) == (2, 1)
|
|
@@ -0,0 +1,62 @@
|
|
|
1
|
+
from .base import DataLoaderIterator
|
|
2
|
+
from .cache import CachedLoader
|
|
3
|
+
from .cluster import ClusterData, ClusterLoader
|
|
4
|
+
from .data_list_loader import DataListLoader
|
|
5
|
+
from .dataloader import Collater, DataLoader
|
|
6
|
+
from .dense_data_loader import DenseDataLoader
|
|
7
|
+
from .dynamic_batch_sampler import DynamicBatchSampler
|
|
8
|
+
from .keras_dataset import FullGraphDataset
|
|
9
|
+
from .graph_saint import (
|
|
10
|
+
GraphSAINTEdgeSampler,
|
|
11
|
+
GraphSAINTNodeSampler,
|
|
12
|
+
GraphSAINTRandomWalkSampler,
|
|
13
|
+
GraphSAINTSampler,
|
|
14
|
+
)
|
|
15
|
+
from .hgt_loader import HGTLoader
|
|
16
|
+
from .imbalanced_sampler import ImbalancedSampler
|
|
17
|
+
from .link_loader import EdgeSamplerInput, LinkLoader
|
|
18
|
+
from .link_neighbor_loader import LinkNeighborLoader
|
|
19
|
+
from .mixin import AffinityMixin, LogMemoryMixin, MultithreadingMixin
|
|
20
|
+
from .neighbor_loader import NeighborLoader
|
|
21
|
+
from .neighbor_sampler import Adj, EdgeIndex, NeighborSampler
|
|
22
|
+
from .node_loader import HeteroSamplerOutput, NodeLoader, NodeSamplerInput, SamplerOutput
|
|
23
|
+
from .prefetch import DeviceHelper, PrefetchLoader
|
|
24
|
+
from .random_node_loader import RandomNodeLoader
|
|
25
|
+
from .shadow import ShaDowKHopSampler
|
|
26
|
+
from .temporal_dataloader import TemporalDataLoader
|
|
27
|
+
from .zip_loader import ZipLoader
|
|
28
|
+
from .utils import to_numpy
|
|
29
|
+
|
|
30
|
+
__all__ = [
|
|
31
|
+
'to_numpy',
|
|
32
|
+
'DataLoader',
|
|
33
|
+
'NodeLoader',
|
|
34
|
+
'LinkLoader',
|
|
35
|
+
'NeighborLoader',
|
|
36
|
+
'LinkNeighborLoader',
|
|
37
|
+
'HGTLoader',
|
|
38
|
+
'ClusterData',
|
|
39
|
+
'ClusterLoader',
|
|
40
|
+
'GraphSAINTSampler',
|
|
41
|
+
'GraphSAINTNodeSampler',
|
|
42
|
+
'GraphSAINTEdgeSampler',
|
|
43
|
+
'GraphSAINTRandomWalkSampler',
|
|
44
|
+
'ShaDowKHopSampler',
|
|
45
|
+
'RandomNodeLoader',
|
|
46
|
+
'ZipLoader',
|
|
47
|
+
'DataListLoader',
|
|
48
|
+
'DenseDataLoader',
|
|
49
|
+
'TemporalDataLoader',
|
|
50
|
+
'NeighborSampler',
|
|
51
|
+
'ImbalancedSampler',
|
|
52
|
+
'DynamicBatchSampler',
|
|
53
|
+
'PrefetchLoader',
|
|
54
|
+
'CachedLoader',
|
|
55
|
+
'AffinityMixin',
|
|
56
|
+
'MultithreadingMixin',
|
|
57
|
+
'LogMemoryMixin',
|
|
58
|
+
'Collater',
|
|
59
|
+
'DataLoaderIterator',
|
|
60
|
+
'FullGraphDataset',
|
|
61
|
+
]
|
|
62
|
+
|
k3_node/loader/base.py
ADDED
|
@@ -0,0 +1,69 @@
|
|
|
1
|
+
from typing import Any, Callable
|
|
2
|
+
|
|
3
|
+
try:
|
|
4
|
+
import torch
|
|
5
|
+
BaseDataLoader = torch.utils.data.DataLoader
|
|
6
|
+
except ImportError:
|
|
7
|
+
class BaseDataLoader:
|
|
8
|
+
r"""Fallback BaseDataLoader when PyTorch is not installed."""
|
|
9
|
+
def __init__(self, dataset=None, batch_size=1, shuffle=False, **kwargs):
|
|
10
|
+
self.dataset = dataset
|
|
11
|
+
self.batch_size = batch_size
|
|
12
|
+
self.shuffle = shuffle
|
|
13
|
+
|
|
14
|
+
def __iter__(self):
|
|
15
|
+
collate_fn = getattr(self, 'collate_fn', None) or (lambda x: x)
|
|
16
|
+
dataset = getattr(self, 'dataset', [])
|
|
17
|
+
batch_size = getattr(self, 'batch_size', 1)
|
|
18
|
+
shuffle = getattr(self, 'shuffle', False)
|
|
19
|
+
indices = list(range(len(dataset)))
|
|
20
|
+
if shuffle:
|
|
21
|
+
import random
|
|
22
|
+
random.shuffle(indices)
|
|
23
|
+
for i in range(0, len(indices), batch_size):
|
|
24
|
+
batch_indices = indices[i:i + batch_size]
|
|
25
|
+
batch = [dataset[idx] for idx in batch_indices]
|
|
26
|
+
yield collate_fn(batch)
|
|
27
|
+
|
|
28
|
+
def __len__(self):
|
|
29
|
+
dataset = getattr(self, 'dataset', [])
|
|
30
|
+
batch_size = getattr(self, 'batch_size', 1)
|
|
31
|
+
return (len(dataset) + batch_size - 1) // batch_size if len(dataset) > 0 else 0
|
|
32
|
+
|
|
33
|
+
try:
|
|
34
|
+
from torch.utils.data.dataloader import (
|
|
35
|
+
_BaseDataLoaderIter,
|
|
36
|
+
_MultiProcessingDataLoaderIter,
|
|
37
|
+
)
|
|
38
|
+
except ImportError:
|
|
39
|
+
class _BaseDataLoaderIter:
|
|
40
|
+
pass
|
|
41
|
+
|
|
42
|
+
class _MultiProcessingDataLoaderIter:
|
|
43
|
+
pass
|
|
44
|
+
|
|
45
|
+
|
|
46
|
+
class DataLoaderIterator:
|
|
47
|
+
r"""A data loader iterator extended by a post transformation function
|
|
48
|
+
:meth:`transform_fn`.
|
|
49
|
+
"""
|
|
50
|
+
def __init__(self, iterator: Any, transform_fn: Callable):
|
|
51
|
+
self.iterator = iterator
|
|
52
|
+
self.transform_fn = transform_fn
|
|
53
|
+
|
|
54
|
+
def __iter__(self) -> 'DataLoaderIterator':
|
|
55
|
+
return self
|
|
56
|
+
|
|
57
|
+
def _reset(self, loader: Any, first_iter: bool = False):
|
|
58
|
+
if hasattr(self.iterator, '_reset'):
|
|
59
|
+
self.iterator._reset(loader, first_iter)
|
|
60
|
+
|
|
61
|
+
def __len__(self) -> int:
|
|
62
|
+
return len(self.iterator)
|
|
63
|
+
|
|
64
|
+
def __next__(self) -> Any:
|
|
65
|
+
return self.transform_fn(next(self.iterator))
|
|
66
|
+
|
|
67
|
+
def __del__(self) -> Any:
|
|
68
|
+
if isinstance(self.iterator, _MultiProcessingDataLoaderIter):
|
|
69
|
+
self.iterator.__del__()
|
k3_node/loader/cache.py
ADDED
|
@@ -0,0 +1,68 @@
|
|
|
1
|
+
from collections.abc import Mapping, Sequence
|
|
2
|
+
from typing import Any, Callable, List, Optional
|
|
3
|
+
|
|
4
|
+
try:
|
|
5
|
+
import torch
|
|
6
|
+
from torch.utils.data import DataLoader
|
|
7
|
+
except ImportError:
|
|
8
|
+
torch = None
|
|
9
|
+
DataLoader = object
|
|
10
|
+
|
|
11
|
+
|
|
12
|
+
def to_device(inputs: Any, device: Optional[Any] = None) -> Any:
|
|
13
|
+
if device is None:
|
|
14
|
+
return inputs
|
|
15
|
+
if hasattr(inputs, 'to'):
|
|
16
|
+
return inputs.to(device)
|
|
17
|
+
elif isinstance(inputs, Mapping):
|
|
18
|
+
return {key: to_device(value, device) for key, value in inputs.items()}
|
|
19
|
+
elif isinstance(inputs, tuple) and hasattr(inputs, '_fields'):
|
|
20
|
+
return type(inputs)(*(to_device(s, device) for s in zip(*inputs)))
|
|
21
|
+
elif isinstance(inputs, Sequence) and not isinstance(inputs, str):
|
|
22
|
+
return [to_device(s, device) for s in zip(*inputs)]
|
|
23
|
+
return inputs
|
|
24
|
+
|
|
25
|
+
|
|
26
|
+
class CachedLoader:
|
|
27
|
+
r"""A loader to cache mini-batch outputs in memory across epochs.
|
|
28
|
+
|
|
29
|
+
Args:
|
|
30
|
+
loader (DataLoader): The data loader.
|
|
31
|
+
device (torch.device, optional): The device to load the data to. (default: :obj:`None`)
|
|
32
|
+
transform (callable, optional): A function that takes in a sampled mini-batch and returns a transformed version. (default: :obj:`None`)
|
|
33
|
+
"""
|
|
34
|
+
def __init__(
|
|
35
|
+
self,
|
|
36
|
+
loader: Any,
|
|
37
|
+
device: Optional[Any] = None,
|
|
38
|
+
transform: Optional[Callable] = None,
|
|
39
|
+
):
|
|
40
|
+
self.loader = loader
|
|
41
|
+
self.device = device
|
|
42
|
+
self.transform = transform
|
|
43
|
+
self._cache: List[Any] = []
|
|
44
|
+
|
|
45
|
+
def clear(self):
|
|
46
|
+
r"""Clears the cache."""
|
|
47
|
+
self._cache = []
|
|
48
|
+
|
|
49
|
+
def __iter__(self) -> Any:
|
|
50
|
+
if len(self._cache) > 0:
|
|
51
|
+
for batch in self._cache:
|
|
52
|
+
yield batch
|
|
53
|
+
return
|
|
54
|
+
|
|
55
|
+
for batch in self.loader:
|
|
56
|
+
if self.transform is not None:
|
|
57
|
+
batch = self.transform(batch)
|
|
58
|
+
|
|
59
|
+
batch = to_device(batch, self.device)
|
|
60
|
+
self._cache.append(batch)
|
|
61
|
+
yield batch
|
|
62
|
+
|
|
63
|
+
def __len__(self) -> int:
|
|
64
|
+
return len(self.loader)
|
|
65
|
+
|
|
66
|
+
def __repr__(self) -> str:
|
|
67
|
+
return f'{self.__class__.__name__}({self.loader})'
|
|
68
|
+
|
|
@@ -0,0 +1,127 @@
|
|
|
1
|
+
import copy
|
|
2
|
+
from typing import Any, List, Optional, Union
|
|
3
|
+
|
|
4
|
+
import numpy as np
|
|
5
|
+
|
|
6
|
+
try:
|
|
7
|
+
import torch
|
|
8
|
+
import torch.utils.data
|
|
9
|
+
from torch import Tensor
|
|
10
|
+
BaseDataset = torch.utils.data.Dataset
|
|
11
|
+
BaseDataLoader = torch.utils.data.DataLoader
|
|
12
|
+
except ImportError:
|
|
13
|
+
torch = None
|
|
14
|
+
Tensor = type(None)
|
|
15
|
+
BaseDataset = object
|
|
16
|
+
BaseDataLoader = object
|
|
17
|
+
|
|
18
|
+
from k3_node.data import Data
|
|
19
|
+
from k3_node.loader.sampler_utils import partition_graph
|
|
20
|
+
from k3_node.loader.keras_dataset import loader_bases
|
|
21
|
+
|
|
22
|
+
|
|
23
|
+
class ClusterData(BaseDataset):
|
|
24
|
+
r"""Clusters/partitions a graph data object into multiple subgraphs, as
|
|
25
|
+
motivated by the "Cluster-GCN" paper.
|
|
26
|
+
|
|
27
|
+
Args:
|
|
28
|
+
data (Data): The graph data object.
|
|
29
|
+
num_parts (int): The number of partitions.
|
|
30
|
+
recursive (bool, optional): Multilevel recursive bisection if True. (default: :obj:`False`)
|
|
31
|
+
save_dir (str, optional): Directory to save partitioned data. (default: :obj:`None`)
|
|
32
|
+
log (bool, optional): If set to :obj:`False`, will not log. (default: :obj:`True`)
|
|
33
|
+
keep_inter_cluster_edges (bool, optional): Keep inter-cluster connections. (default: :obj:`False`)
|
|
34
|
+
"""
|
|
35
|
+
def __init__(
|
|
36
|
+
self,
|
|
37
|
+
data: Data,
|
|
38
|
+
num_parts: int,
|
|
39
|
+
recursive: bool = False,
|
|
40
|
+
save_dir: Optional[str] = None,
|
|
41
|
+
filename: Optional[str] = None,
|
|
42
|
+
log: bool = True,
|
|
43
|
+
keep_inter_cluster_edges: bool = False,
|
|
44
|
+
sparse_format: str = 'csr',
|
|
45
|
+
):
|
|
46
|
+
assert data.edge_index is not None
|
|
47
|
+
|
|
48
|
+
self.num_parts = num_parts
|
|
49
|
+
self.recursive = recursive
|
|
50
|
+
self.keep_inter_cluster_edges = keep_inter_cluster_edges
|
|
51
|
+
self.sparse_format = sparse_format
|
|
52
|
+
self.data = data
|
|
53
|
+
|
|
54
|
+
loaded = False
|
|
55
|
+
if save_dir is not None:
|
|
56
|
+
import os.path as osp
|
|
57
|
+
recursive_str = '_recursive' if recursive else ''
|
|
58
|
+
root_dir = osp.join(save_dir, f'part_{num_parts}{recursive_str}')
|
|
59
|
+
path = osp.join(root_dir, filename or 'metis.pt')
|
|
60
|
+
if osp.exists(path):
|
|
61
|
+
try:
|
|
62
|
+
import torch
|
|
63
|
+
part = torch.load(path, map_location="cpu", weights_only=False)
|
|
64
|
+
if hasattr(part, "partptr") and hasattr(part, "node_perm"):
|
|
65
|
+
partptr = np.asarray(part.partptr)
|
|
66
|
+
node_perm = np.asarray(part.node_perm)
|
|
67
|
+
self.part_nodes = [node_perm[partptr[i]:partptr[i+1]] for i in range(num_parts)]
|
|
68
|
+
self.cluster = np.zeros(data.num_nodes, dtype=np.int64)
|
|
69
|
+
for i in range(num_parts):
|
|
70
|
+
self.cluster[self.part_nodes[i]] = i
|
|
71
|
+
loaded = True
|
|
72
|
+
except Exception:
|
|
73
|
+
pass
|
|
74
|
+
|
|
75
|
+
if not loaded:
|
|
76
|
+
self.cluster = partition_graph(data.edge_index, data.num_nodes, num_parts)
|
|
77
|
+
from k3_node.loader.utils import to_numpy
|
|
78
|
+
cluster_np = to_numpy(self.cluster)
|
|
79
|
+
sort_idx = np.argsort(cluster_np)
|
|
80
|
+
sorted_cluster = cluster_np[sort_idx]
|
|
81
|
+
split_idx = np.searchsorted(sorted_cluster, np.arange(num_parts + 1))
|
|
82
|
+
self.part_nodes = [sort_idx[split_idx[i]:split_idx[i+1]] for i in range(num_parts)]
|
|
83
|
+
|
|
84
|
+
def __len__(self) -> int:
|
|
85
|
+
return self.num_parts
|
|
86
|
+
|
|
87
|
+
def __getitem__(self, idx: int) -> Data:
|
|
88
|
+
nodes = self.part_nodes[idx]
|
|
89
|
+
return self.data.subgraph(nodes)
|
|
90
|
+
|
|
91
|
+
def __repr__(self) -> str:
|
|
92
|
+
return f'{self.__class__.__name__}({self.num_parts})'
|
|
93
|
+
|
|
94
|
+
|
|
95
|
+
class ClusterLoader(*loader_bases(BaseDataLoader)):
|
|
96
|
+
r"""The data loader scheme from Cluster-GCN which merges partitioned
|
|
97
|
+
subgraphs to form a mini-batch.
|
|
98
|
+
|
|
99
|
+
Args:
|
|
100
|
+
cluster_data (ClusterData): The already partitioned data object.
|
|
101
|
+
**kwargs (optional): Additional arguments of :class:`torch.utils.data.DataLoader`.
|
|
102
|
+
"""
|
|
103
|
+
def __init__(self, cluster_data: ClusterData, **kwargs):
|
|
104
|
+
self.cluster_data = cluster_data
|
|
105
|
+
kwargs.pop('collate_fn', None)
|
|
106
|
+
iterator = range(len(cluster_data))
|
|
107
|
+
|
|
108
|
+
if torch is not None:
|
|
109
|
+
super().__init__(iterator, collate_fn=self._collate, **kwargs)
|
|
110
|
+
else:
|
|
111
|
+
self.dataset = iterator
|
|
112
|
+
self.collate_fn = self._collate
|
|
113
|
+
|
|
114
|
+
def _collate(self, batch: List[int]) -> Data:
|
|
115
|
+
all_nodes = []
|
|
116
|
+
is_torch = torch is not None and isinstance(self.cluster_data.cluster, Tensor)
|
|
117
|
+
|
|
118
|
+
for part_id in batch:
|
|
119
|
+
all_nodes.append(self.cluster_data.part_nodes[part_id])
|
|
120
|
+
|
|
121
|
+
if is_torch and all_nodes and isinstance(all_nodes[0], Tensor):
|
|
122
|
+
nodes = torch.cat(all_nodes, dim=0)
|
|
123
|
+
else:
|
|
124
|
+
nodes = np.concatenate([np.asarray(x) for x in all_nodes], axis=0)
|
|
125
|
+
|
|
126
|
+
return self.cluster_data.data.subgraph(nodes)
|
|
127
|
+
|
|
@@ -0,0 +1,45 @@
|
|
|
1
|
+
from typing import List, Union
|
|
2
|
+
|
|
3
|
+
try:
|
|
4
|
+
import torch
|
|
5
|
+
import torch.utils.data
|
|
6
|
+
BaseDataLoader = torch.utils.data.DataLoader
|
|
7
|
+
except ImportError:
|
|
8
|
+
torch = None
|
|
9
|
+
BaseDataLoader = object
|
|
10
|
+
|
|
11
|
+
from k3_node.data import Dataset
|
|
12
|
+
from k3_node.data.data import BaseData
|
|
13
|
+
|
|
14
|
+
|
|
15
|
+
def collate_fn(data_list):
|
|
16
|
+
return data_list
|
|
17
|
+
|
|
18
|
+
|
|
19
|
+
class DataListLoader(BaseDataLoader):
|
|
20
|
+
r"""A data loader which batches data objects from a
|
|
21
|
+
:class:`k3_node.data.Dataset` to a Python list without batch collation.
|
|
22
|
+
"""
|
|
23
|
+
def __init__(
|
|
24
|
+
self,
|
|
25
|
+
dataset: Union[Dataset, List[BaseData]],
|
|
26
|
+
batch_size: int = 1,
|
|
27
|
+
shuffle: bool = False,
|
|
28
|
+
**kwargs,
|
|
29
|
+
):
|
|
30
|
+
kwargs.pop('collate_fn', None)
|
|
31
|
+
|
|
32
|
+
if torch is not None:
|
|
33
|
+
super().__init__(
|
|
34
|
+
dataset,
|
|
35
|
+
batch_size=batch_size,
|
|
36
|
+
shuffle=shuffle,
|
|
37
|
+
collate_fn=collate_fn,
|
|
38
|
+
**kwargs,
|
|
39
|
+
)
|
|
40
|
+
else:
|
|
41
|
+
self.dataset = dataset
|
|
42
|
+
self.batch_size = batch_size
|
|
43
|
+
self.shuffle = shuffle
|
|
44
|
+
self.collate_fn = collate_fn
|
|
45
|
+
|
|
@@ -0,0 +1,117 @@
|
|
|
1
|
+
from collections.abc import Mapping, Sequence
|
|
2
|
+
from typing import Any, List, Optional, Union
|
|
3
|
+
|
|
4
|
+
import numpy as np
|
|
5
|
+
|
|
6
|
+
try:
|
|
7
|
+
import torch
|
|
8
|
+
import torch.utils.data
|
|
9
|
+
from torch.utils.data.dataloader import default_collate
|
|
10
|
+
except ImportError:
|
|
11
|
+
torch = None
|
|
12
|
+
default_collate = None
|
|
13
|
+
|
|
14
|
+
from k3_node.data import Batch
|
|
15
|
+
from k3_node.data.data import BaseData
|
|
16
|
+
from k3_node.data.dataset import Dataset
|
|
17
|
+
from k3_node.loader.keras_dataset import loader_bases
|
|
18
|
+
|
|
19
|
+
|
|
20
|
+
class Collater:
|
|
21
|
+
r"""Collates a list of graph data objects or primitives into a mini-batch."""
|
|
22
|
+
def __init__(
|
|
23
|
+
self,
|
|
24
|
+
dataset: Optional[Union[Dataset, Sequence[BaseData]]] = None,
|
|
25
|
+
follow_batch: Optional[List[str]] = None,
|
|
26
|
+
exclude_keys: Optional[List[str]] = None,
|
|
27
|
+
):
|
|
28
|
+
self.dataset = dataset
|
|
29
|
+
self.follow_batch = follow_batch
|
|
30
|
+
self.exclude_keys = exclude_keys
|
|
31
|
+
|
|
32
|
+
def __call__(self, batch: List[Any]) -> Any:
|
|
33
|
+
elem = batch[0]
|
|
34
|
+
if isinstance(elem, BaseData):
|
|
35
|
+
return Batch.from_data_list(
|
|
36
|
+
batch,
|
|
37
|
+
follow_batch=self.follow_batch,
|
|
38
|
+
exclude_keys=self.exclude_keys,
|
|
39
|
+
)
|
|
40
|
+
elif torch is not None and isinstance(elem, torch.Tensor):
|
|
41
|
+
return default_collate(batch)
|
|
42
|
+
elif isinstance(elem, np.ndarray):
|
|
43
|
+
return np.stack(batch, axis=0)
|
|
44
|
+
elif hasattr(elem, '__array__') and not isinstance(elem, (str, bytes)):
|
|
45
|
+
import keras
|
|
46
|
+
np_batch = np.stack([np.asarray(x) for x in batch], axis=0)
|
|
47
|
+
return keras.ops.convert_to_tensor(np_batch)
|
|
48
|
+
elif isinstance(elem, float):
|
|
49
|
+
if torch is not None:
|
|
50
|
+
return torch.tensor(batch, dtype=torch.float)
|
|
51
|
+
return np.array(batch, dtype=np.float32)
|
|
52
|
+
elif isinstance(elem, int):
|
|
53
|
+
if torch is not None:
|
|
54
|
+
return torch.tensor(batch, dtype=torch.long)
|
|
55
|
+
return np.array(batch, dtype=np.int64)
|
|
56
|
+
elif isinstance(elem, str):
|
|
57
|
+
return batch
|
|
58
|
+
elif isinstance(elem, Mapping):
|
|
59
|
+
return {key: self([data[key] for data in batch]) for key in elem}
|
|
60
|
+
elif isinstance(elem, tuple) and hasattr(elem, '_fields'):
|
|
61
|
+
return type(elem)(*(self(s) for s in zip(*batch)))
|
|
62
|
+
elif isinstance(elem, Sequence) and not isinstance(elem, str):
|
|
63
|
+
return [self(s) for s in zip(*batch)]
|
|
64
|
+
|
|
65
|
+
raise TypeError(f"DataLoader found invalid type: '{type(elem)}'")
|
|
66
|
+
|
|
67
|
+
|
|
68
|
+
BaseDataLoader = torch.utils.data.DataLoader if torch is not None else object
|
|
69
|
+
|
|
70
|
+
|
|
71
|
+
class DataLoader(*loader_bases(BaseDataLoader)):
|
|
72
|
+
r"""A data loader which merges data objects from a
|
|
73
|
+
:class:`k3_node.data.Dataset` to a mini-batch.
|
|
74
|
+
Data objects can be either of type :class:`~k3_node.data.Data` or
|
|
75
|
+
:class:`~k3_node.data.HeteroData`.
|
|
76
|
+
|
|
77
|
+
Args:
|
|
78
|
+
dataset (Dataset): The dataset from which to load the data.
|
|
79
|
+
batch_size (int, optional): How many samples per batch to load.
|
|
80
|
+
(default: :obj:`1`)
|
|
81
|
+
shuffle (bool, optional): If set to :obj:`True`, the data will be
|
|
82
|
+
reshuffled at every epoch. (default: :obj:`False`)
|
|
83
|
+
follow_batch (List[str], optional): Creates assignment batch
|
|
84
|
+
vectors for each key in the list. (default: :obj:`None`)
|
|
85
|
+
exclude_keys (List[str], optional): Will exclude each key in the
|
|
86
|
+
list. (default: :obj:`None`)
|
|
87
|
+
**kwargs (optional): Additional arguments of
|
|
88
|
+
:class:`torch.utils.data.DataLoader`.
|
|
89
|
+
"""
|
|
90
|
+
def __init__(
|
|
91
|
+
self,
|
|
92
|
+
dataset: Union[Dataset, Sequence[BaseData]],
|
|
93
|
+
batch_size: int = 1,
|
|
94
|
+
shuffle: bool = False,
|
|
95
|
+
follow_batch: Optional[List[str]] = None,
|
|
96
|
+
exclude_keys: Optional[List[str]] = None,
|
|
97
|
+
**kwargs,
|
|
98
|
+
):
|
|
99
|
+
kwargs.pop('collate_fn', None)
|
|
100
|
+
|
|
101
|
+
self.follow_batch = follow_batch
|
|
102
|
+
self.exclude_keys = exclude_keys
|
|
103
|
+
|
|
104
|
+
if torch is not None:
|
|
105
|
+
super().__init__(
|
|
106
|
+
dataset,
|
|
107
|
+
batch_size=batch_size,
|
|
108
|
+
shuffle=shuffle,
|
|
109
|
+
collate_fn=Collater(dataset, follow_batch, exclude_keys),
|
|
110
|
+
**kwargs,
|
|
111
|
+
)
|
|
112
|
+
else:
|
|
113
|
+
self.dataset = dataset
|
|
114
|
+
self.batch_size = batch_size
|
|
115
|
+
self.shuffle = shuffle
|
|
116
|
+
self.collate_fn = Collater(dataset, follow_batch, exclude_keys)
|
|
117
|
+
|
|
@@ -0,0 +1,62 @@
|
|
|
1
|
+
from typing import List, Union
|
|
2
|
+
|
|
3
|
+
import numpy as np
|
|
4
|
+
|
|
5
|
+
try:
|
|
6
|
+
import torch
|
|
7
|
+
import torch.utils.data
|
|
8
|
+
from torch.utils.data.dataloader import default_collate
|
|
9
|
+
BaseDataLoader = torch.utils.data.DataLoader
|
|
10
|
+
except ImportError:
|
|
11
|
+
torch = None
|
|
12
|
+
default_collate = None
|
|
13
|
+
BaseDataLoader = object
|
|
14
|
+
|
|
15
|
+
from k3_node.data import Batch, Data, Dataset
|
|
16
|
+
from k3_node.loader.keras_dataset import loader_bases
|
|
17
|
+
|
|
18
|
+
|
|
19
|
+
def collate_fn(data_list: List[Data]) -> Batch:
|
|
20
|
+
batch = Batch()
|
|
21
|
+
for key in data_list[0].keys():
|
|
22
|
+
vals = [data[key] for data in data_list]
|
|
23
|
+
if default_collate is not None and isinstance(vals[0], torch.Tensor):
|
|
24
|
+
batch[key] = default_collate(vals)
|
|
25
|
+
elif isinstance(vals[0], np.ndarray):
|
|
26
|
+
batch[key] = np.stack(vals, axis=0)
|
|
27
|
+
elif hasattr(vals[0], '__array__'):
|
|
28
|
+
import keras
|
|
29
|
+
batch[key] = keras.ops.stack(vals, axis=0)
|
|
30
|
+
else:
|
|
31
|
+
batch[key] = vals
|
|
32
|
+
return batch
|
|
33
|
+
|
|
34
|
+
|
|
35
|
+
class DenseDataLoader(*loader_bases(BaseDataLoader)):
|
|
36
|
+
r"""A data loader which batches data objects from a
|
|
37
|
+
:class:`k3_node.data.Dataset` to a :class:`k3_node.data.Batch`
|
|
38
|
+
object by stacking all attributes in a new dimension.
|
|
39
|
+
"""
|
|
40
|
+
def __init__(
|
|
41
|
+
self,
|
|
42
|
+
dataset: Union[Dataset, List[Data]],
|
|
43
|
+
batch_size: int = 1,
|
|
44
|
+
shuffle: bool = False,
|
|
45
|
+
**kwargs,
|
|
46
|
+
):
|
|
47
|
+
kwargs.pop('collate_fn', None)
|
|
48
|
+
|
|
49
|
+
if torch is not None:
|
|
50
|
+
super().__init__(
|
|
51
|
+
dataset,
|
|
52
|
+
batch_size=batch_size,
|
|
53
|
+
shuffle=shuffle,
|
|
54
|
+
collate_fn=collate_fn,
|
|
55
|
+
**kwargs,
|
|
56
|
+
)
|
|
57
|
+
else:
|
|
58
|
+
self.dataset = dataset
|
|
59
|
+
self.batch_size = batch_size
|
|
60
|
+
self.shuffle = shuffle
|
|
61
|
+
self.collate_fn = collate_fn
|
|
62
|
+
|
|
@@ -0,0 +1,93 @@
|
|
|
1
|
+
from typing import Iterator, List, Optional
|
|
2
|
+
|
|
3
|
+
import numpy as np
|
|
4
|
+
|
|
5
|
+
try:
|
|
6
|
+
import torch
|
|
7
|
+
import torch.utils.data.sampler
|
|
8
|
+
BaseSampler = torch.utils.data.sampler.Sampler
|
|
9
|
+
except ImportError:
|
|
10
|
+
torch = None
|
|
11
|
+
BaseSampler = object
|
|
12
|
+
|
|
13
|
+
from k3_node.data import Dataset
|
|
14
|
+
|
|
15
|
+
|
|
16
|
+
class DynamicBatchSampler(BaseSampler):
|
|
17
|
+
r"""Dynamically adds samples to a mini-batch up to a maximum size (either
|
|
18
|
+
based on number of nodes or number of edges).
|
|
19
|
+
|
|
20
|
+
Args:
|
|
21
|
+
dataset (Dataset): Dataset to sample from.
|
|
22
|
+
max_num (int): Size of mini-batch to aim for in number of nodes or edges.
|
|
23
|
+
mode (str, optional): :obj:`"node"` or :obj:`"edge"` to measure batch size. (default: :obj:`"node"`)
|
|
24
|
+
shuffle (bool, optional): If set to :obj:`True`, will have the data reshuffled at every epoch. (default: :obj:`False`)
|
|
25
|
+
skip_too_big (bool, optional): If set to :obj:`True`, skip samples which cannot fit in a batch by itself. (default: :obj:`False`)
|
|
26
|
+
num_steps (int, optional): The number of mini-batches to draw for a single epoch. (default: :obj:`None`)
|
|
27
|
+
"""
|
|
28
|
+
def __init__(
|
|
29
|
+
self,
|
|
30
|
+
dataset: Dataset,
|
|
31
|
+
max_num: int,
|
|
32
|
+
mode: str = 'node',
|
|
33
|
+
shuffle: bool = False,
|
|
34
|
+
skip_too_big: bool = False,
|
|
35
|
+
num_steps: Optional[int] = None,
|
|
36
|
+
):
|
|
37
|
+
if max_num <= 0:
|
|
38
|
+
raise ValueError(f"`max_num` should be a positive integer value (got {max_num})")
|
|
39
|
+
if mode not in ['node', 'edge']:
|
|
40
|
+
raise ValueError(f"`mode` choice should be either 'node' or 'edge' (got '{mode}')")
|
|
41
|
+
|
|
42
|
+
self.dataset = dataset
|
|
43
|
+
self.max_num = max_num
|
|
44
|
+
self.mode = mode
|
|
45
|
+
self.shuffle = shuffle
|
|
46
|
+
self.skip_too_big = skip_too_big
|
|
47
|
+
self.num_steps = num_steps
|
|
48
|
+
self.max_steps = num_steps or len(dataset)
|
|
49
|
+
|
|
50
|
+
def __iter__(self) -> Iterator[List[int]]:
|
|
51
|
+
if self.shuffle:
|
|
52
|
+
if torch is not None:
|
|
53
|
+
indices = torch.randperm(len(self.dataset)).tolist()
|
|
54
|
+
else:
|
|
55
|
+
indices = np.random.permutation(len(self.dataset)).tolist()
|
|
56
|
+
else:
|
|
57
|
+
indices = list(range(len(self.dataset)))
|
|
58
|
+
|
|
59
|
+
samples: List[int] = []
|
|
60
|
+
current_num: int = 0
|
|
61
|
+
num_steps: int = 0
|
|
62
|
+
num_processed: int = 0
|
|
63
|
+
|
|
64
|
+
while num_processed < len(self.dataset) and num_steps < self.max_steps:
|
|
65
|
+
for i in indices[num_processed:]:
|
|
66
|
+
data = self.dataset[i]
|
|
67
|
+
num = data.num_nodes if self.mode == 'node' else data.num_edges
|
|
68
|
+
|
|
69
|
+
if current_num + num > self.max_num:
|
|
70
|
+
if current_num == 0:
|
|
71
|
+
if self.skip_too_big:
|
|
72
|
+
num_processed += 1
|
|
73
|
+
continue
|
|
74
|
+
else:
|
|
75
|
+
break
|
|
76
|
+
|
|
77
|
+
samples.append(i)
|
|
78
|
+
num_processed += 1
|
|
79
|
+
current_num += num
|
|
80
|
+
|
|
81
|
+
yield samples
|
|
82
|
+
samples = []
|
|
83
|
+
current_num = 0
|
|
84
|
+
num_steps += 1
|
|
85
|
+
|
|
86
|
+
def __len__(self) -> int:
|
|
87
|
+
if self.num_steps is None:
|
|
88
|
+
raise ValueError(
|
|
89
|
+
f"The length of '{self.__class__.__name__}' is undefined since the number of steps per epoch "
|
|
90
|
+
f"is ambiguous. Either specify `num_steps` or use a static batch sampler."
|
|
91
|
+
)
|
|
92
|
+
return self.num_steps
|
|
93
|
+
|