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/datasets/fake.py
ADDED
|
@@ -0,0 +1,256 @@
|
|
|
1
|
+
import random
|
|
2
|
+
from collections import defaultdict
|
|
3
|
+
from itertools import product
|
|
4
|
+
from typing import Callable, Dict, List, Optional, Tuple, Union
|
|
5
|
+
import numpy as np
|
|
6
|
+
from keras import ops
|
|
7
|
+
|
|
8
|
+
from k3_node.data import Data, HeteroData, InMemoryDataset
|
|
9
|
+
from k3_node.layers.conv.utils import remove_self_loops
|
|
10
|
+
from k3_node.transforms.utils import to_undirected
|
|
11
|
+
from k3_node.utils.graph import coalesce
|
|
12
|
+
|
|
13
|
+
|
|
14
|
+
def get_num_nodes(avg_num_nodes: int, avg_degree: float) -> int:
|
|
15
|
+
min_num_nodes = max(3 * avg_num_nodes // 4, int(avg_degree))
|
|
16
|
+
max_num_nodes = 5 * avg_num_nodes // 4
|
|
17
|
+
return random.randint(min_num_nodes, max_num_nodes)
|
|
18
|
+
|
|
19
|
+
|
|
20
|
+
def get_num_channels(num_channels: int) -> int:
|
|
21
|
+
min_num_channels = 3 * num_channels // 4
|
|
22
|
+
max_num_channels = 5 * num_channels // 4
|
|
23
|
+
return random.randint(min_num_channels, max_num_channels)
|
|
24
|
+
|
|
25
|
+
|
|
26
|
+
def get_edge_index(
|
|
27
|
+
num_src_nodes: int,
|
|
28
|
+
num_dst_nodes: int,
|
|
29
|
+
avg_degree: float,
|
|
30
|
+
is_undirected: bool = False,
|
|
31
|
+
remove_loops: bool = False,
|
|
32
|
+
):
|
|
33
|
+
num_edges = int(num_src_nodes * avg_degree)
|
|
34
|
+
row = np.random.randint(0, num_src_nodes, size=(num_edges,), dtype=np.int64)
|
|
35
|
+
col = np.random.randint(0, num_dst_nodes, size=(num_edges,), dtype=np.int64)
|
|
36
|
+
edge_index = np.stack([row, col], axis=0)
|
|
37
|
+
|
|
38
|
+
if remove_loops:
|
|
39
|
+
edge_index, _ = remove_self_loops(edge_index)
|
|
40
|
+
|
|
41
|
+
num_nodes = max(num_src_nodes, num_dst_nodes)
|
|
42
|
+
if is_undirected:
|
|
43
|
+
edge_index = to_undirected(edge_index, num_nodes=num_nodes)
|
|
44
|
+
else:
|
|
45
|
+
edge_index, _ = coalesce(edge_index, num_nodes=num_nodes)
|
|
46
|
+
|
|
47
|
+
return ops.convert_to_tensor(edge_index, dtype="int64")
|
|
48
|
+
|
|
49
|
+
|
|
50
|
+
class FakeDataset(InMemoryDataset):
|
|
51
|
+
r"""A fake dataset that returns randomly generated `k3_node.data.Data` objects."""
|
|
52
|
+
|
|
53
|
+
def __init__(
|
|
54
|
+
self,
|
|
55
|
+
num_graphs: int = 1,
|
|
56
|
+
avg_num_nodes: int = 1000,
|
|
57
|
+
avg_degree: float = 10.0,
|
|
58
|
+
num_channels: int = 64,
|
|
59
|
+
edge_dim: int = 0,
|
|
60
|
+
num_classes: int = 10,
|
|
61
|
+
task: str = "auto",
|
|
62
|
+
is_undirected: bool = True,
|
|
63
|
+
transform: Optional[Callable] = None,
|
|
64
|
+
pre_transform: Optional[Callable] = None,
|
|
65
|
+
**kwargs: Union[int, Tuple[int, ...]],
|
|
66
|
+
):
|
|
67
|
+
super().__init__(None, transform)
|
|
68
|
+
|
|
69
|
+
if task == "auto":
|
|
70
|
+
task = "graph" if num_graphs > 1 else "node"
|
|
71
|
+
assert task in ["node", "graph"]
|
|
72
|
+
|
|
73
|
+
self.num_graphs_val = num_graphs
|
|
74
|
+
self.avg_num_nodes = max(avg_num_nodes, int(avg_degree))
|
|
75
|
+
self.avg_degree = max(avg_degree, 1)
|
|
76
|
+
self.num_channels = num_channels
|
|
77
|
+
self.edge_dim = edge_dim
|
|
78
|
+
self._num_classes = num_classes
|
|
79
|
+
self.task = task
|
|
80
|
+
self.is_undirected = is_undirected
|
|
81
|
+
self.kwargs = kwargs
|
|
82
|
+
|
|
83
|
+
data_list = [self.generate_data() for _ in range(max(num_graphs, 1))]
|
|
84
|
+
self.data, self.slices = self.collate(data_list)
|
|
85
|
+
|
|
86
|
+
def __repr__(self) -> str:
|
|
87
|
+
return f"FakeDataset({self.num_graphs_val})" if self.num_graphs_val > 1 else "FakeDataset()"
|
|
88
|
+
|
|
89
|
+
def generate_data(self) -> Data:
|
|
90
|
+
num_nodes = get_num_nodes(self.avg_num_nodes, self.avg_degree)
|
|
91
|
+
data = Data()
|
|
92
|
+
|
|
93
|
+
if self._num_classes > 0 and self.task == "node":
|
|
94
|
+
data.y = ops.convert_to_tensor(
|
|
95
|
+
np.random.randint(0, self._num_classes, size=(num_nodes,), dtype=np.int64),
|
|
96
|
+
dtype="int64",
|
|
97
|
+
)
|
|
98
|
+
elif self._num_classes > 0 and self.task == "graph":
|
|
99
|
+
data.y = ops.convert_to_tensor(
|
|
100
|
+
np.array([random.randint(0, self._num_classes - 1)], dtype=np.int64),
|
|
101
|
+
dtype="int64",
|
|
102
|
+
)
|
|
103
|
+
|
|
104
|
+
data.edge_index = get_edge_index(
|
|
105
|
+
num_nodes, num_nodes, self.avg_degree, self.is_undirected, remove_loops=True
|
|
106
|
+
)
|
|
107
|
+
|
|
108
|
+
if self.num_channels > 0:
|
|
109
|
+
x = np.random.randn(num_nodes, self.num_channels).astype(np.float32)
|
|
110
|
+
if self._num_classes > 0 and self.task == "node":
|
|
111
|
+
y_np = ops.convert_to_numpy(data.y)
|
|
112
|
+
x = x + y_np[:, None]
|
|
113
|
+
elif self._num_classes > 0 and self.task == "graph":
|
|
114
|
+
y_np = ops.convert_to_numpy(data.y)
|
|
115
|
+
x = x + y_np
|
|
116
|
+
data.x = ops.convert_to_tensor(x, dtype="float32")
|
|
117
|
+
else:
|
|
118
|
+
data.num_nodes = num_nodes
|
|
119
|
+
|
|
120
|
+
num_edges = int(ops.shape(data.edge_index)[1])
|
|
121
|
+
if self.edge_dim > 1:
|
|
122
|
+
data.edge_attr = ops.convert_to_tensor(
|
|
123
|
+
np.random.rand(num_edges, self.edge_dim).astype(np.float32),
|
|
124
|
+
dtype="float32",
|
|
125
|
+
)
|
|
126
|
+
elif self.edge_dim == 1:
|
|
127
|
+
data.edge_weight = ops.convert_to_tensor(
|
|
128
|
+
np.random.rand(num_edges).astype(np.float32),
|
|
129
|
+
dtype="float32",
|
|
130
|
+
)
|
|
131
|
+
|
|
132
|
+
for feature_name, feature_shape in self.kwargs.items():
|
|
133
|
+
shape = (feature_shape,) if isinstance(feature_shape, int) else feature_shape
|
|
134
|
+
setattr(
|
|
135
|
+
data,
|
|
136
|
+
feature_name,
|
|
137
|
+
ops.convert_to_tensor(np.random.randn(*shape).astype(np.float32), dtype="float32"),
|
|
138
|
+
)
|
|
139
|
+
|
|
140
|
+
return data
|
|
141
|
+
|
|
142
|
+
|
|
143
|
+
class FakeHeteroDataset(InMemoryDataset):
|
|
144
|
+
r"""A fake dataset that returns randomly generated `k3_node.data.HeteroData` objects."""
|
|
145
|
+
|
|
146
|
+
def __init__(
|
|
147
|
+
self,
|
|
148
|
+
num_graphs: int = 1,
|
|
149
|
+
num_node_types: int = 3,
|
|
150
|
+
num_edge_types: int = 6,
|
|
151
|
+
avg_num_nodes: int = 1000,
|
|
152
|
+
avg_degree: float = 10.0,
|
|
153
|
+
avg_num_channels: int = 64,
|
|
154
|
+
edge_dim: int = 0,
|
|
155
|
+
num_classes: int = 10,
|
|
156
|
+
task: str = "auto",
|
|
157
|
+
transform: Optional[Callable] = None,
|
|
158
|
+
pre_transform: Optional[Callable] = None,
|
|
159
|
+
**kwargs: Union[int, Tuple[int, ...]],
|
|
160
|
+
):
|
|
161
|
+
super().__init__(None, transform)
|
|
162
|
+
|
|
163
|
+
if task == "auto":
|
|
164
|
+
task = "graph" if num_graphs > 1 else "node"
|
|
165
|
+
assert task in ["node", "graph"]
|
|
166
|
+
|
|
167
|
+
self.num_graphs_val = num_graphs
|
|
168
|
+
self.node_types = [f"v{i}" for i in range(max(num_node_types, 1))]
|
|
169
|
+
|
|
170
|
+
edge_types: List[Tuple[str, str]] = []
|
|
171
|
+
edge_type_product = list(product(self.node_types, self.node_types))
|
|
172
|
+
while len(edge_types) < max(num_edge_types, 1):
|
|
173
|
+
edge_types.extend(edge_type_product)
|
|
174
|
+
random.shuffle(edge_types)
|
|
175
|
+
|
|
176
|
+
self.edge_types: List[Tuple[str, str, str]] = []
|
|
177
|
+
count: Dict[Tuple[str, str], int] = defaultdict(int)
|
|
178
|
+
for edge_type in edge_types[: max(num_edge_types, 1)]:
|
|
179
|
+
rel = f"e{count[edge_type]}"
|
|
180
|
+
count[edge_type] += 1
|
|
181
|
+
self.edge_types.append((edge_type[0], rel, edge_type[1]))
|
|
182
|
+
|
|
183
|
+
self.avg_num_nodes = max(avg_num_nodes, int(avg_degree))
|
|
184
|
+
self.avg_degree = max(avg_degree, 1)
|
|
185
|
+
self.avg_num_channels = avg_num_channels
|
|
186
|
+
self.edge_dim = edge_dim
|
|
187
|
+
self._num_classes = num_classes
|
|
188
|
+
self.task = task
|
|
189
|
+
self.kwargs = kwargs
|
|
190
|
+
|
|
191
|
+
data_list = [self.generate_data() for _ in range(max(num_graphs, 1))]
|
|
192
|
+
self.data, self.slices = self.collate(data_list)
|
|
193
|
+
|
|
194
|
+
def __repr__(self) -> str:
|
|
195
|
+
return f"FakeHeteroDataset({self.num_graphs_val})" if self.num_graphs_val > 1 else "FakeHeteroDataset()"
|
|
196
|
+
|
|
197
|
+
def generate_data(self) -> HeteroData:
|
|
198
|
+
data = HeteroData()
|
|
199
|
+
|
|
200
|
+
for node_type in self.node_types:
|
|
201
|
+
num_nodes = get_num_nodes(self.avg_num_nodes, self.avg_degree)
|
|
202
|
+
num_channels = get_num_channels(self.avg_num_channels)
|
|
203
|
+
store = data[node_type]
|
|
204
|
+
|
|
205
|
+
if self.avg_num_channels > 0:
|
|
206
|
+
store.x = ops.convert_to_tensor(
|
|
207
|
+
np.random.randn(num_nodes, num_channels).astype(np.float32),
|
|
208
|
+
dtype="float32",
|
|
209
|
+
)
|
|
210
|
+
else:
|
|
211
|
+
store.num_nodes = num_nodes
|
|
212
|
+
|
|
213
|
+
if self._num_classes > 0 and self.task == "node":
|
|
214
|
+
store.y = ops.convert_to_tensor(
|
|
215
|
+
np.random.randint(0, self._num_classes, size=(num_nodes,), dtype=np.int64),
|
|
216
|
+
dtype="int64",
|
|
217
|
+
)
|
|
218
|
+
|
|
219
|
+
for edge_type in self.edge_types:
|
|
220
|
+
src, rel, dst = edge_type
|
|
221
|
+
store = data[edge_type]
|
|
222
|
+
store.edge_index = get_edge_index(
|
|
223
|
+
data[src].num_nodes,
|
|
224
|
+
data[dst].num_nodes,
|
|
225
|
+
self.avg_degree,
|
|
226
|
+
is_undirected=False,
|
|
227
|
+
remove_loops=False,
|
|
228
|
+
)
|
|
229
|
+
|
|
230
|
+
num_edges = int(ops.shape(store.edge_index)[1])
|
|
231
|
+
if self.edge_dim > 1:
|
|
232
|
+
store.edge_attr = ops.convert_to_tensor(
|
|
233
|
+
np.random.rand(num_edges, self.edge_dim).astype(np.float32),
|
|
234
|
+
dtype="float32",
|
|
235
|
+
)
|
|
236
|
+
elif self.edge_dim == 1:
|
|
237
|
+
store.edge_weight = ops.convert_to_tensor(
|
|
238
|
+
np.random.rand(num_edges).astype(np.float32),
|
|
239
|
+
dtype="float32",
|
|
240
|
+
)
|
|
241
|
+
|
|
242
|
+
if self._num_classes > 0 and self.task == "graph":
|
|
243
|
+
data.y = ops.convert_to_tensor(
|
|
244
|
+
np.array([random.randint(0, self._num_classes - 1)], dtype=np.int64),
|
|
245
|
+
dtype="int64",
|
|
246
|
+
)
|
|
247
|
+
|
|
248
|
+
for feature_name, feature_shape in self.kwargs.items():
|
|
249
|
+
shape = (feature_shape,) if isinstance(feature_shape, int) else feature_shape
|
|
250
|
+
setattr(
|
|
251
|
+
data,
|
|
252
|
+
feature_name,
|
|
253
|
+
ops.convert_to_tensor(np.random.randn(*shape).astype(np.float32), dtype="float32"),
|
|
254
|
+
)
|
|
255
|
+
|
|
256
|
+
return data
|
|
@@ -0,0 +1,90 @@
|
|
|
1
|
+
from typing import Callable, List, Optional
|
|
2
|
+
import numpy as np
|
|
3
|
+
from keras import ops
|
|
4
|
+
|
|
5
|
+
from k3_node.data import Data, InMemoryDataset
|
|
6
|
+
from k3_node.io import fs
|
|
7
|
+
|
|
8
|
+
|
|
9
|
+
class FB15k_237(InMemoryDataset):
|
|
10
|
+
r"""The FB15K237 dataset containing 14,541 entities, 237 relations and 310,116 fact triples.
|
|
11
|
+
|
|
12
|
+
Args:
|
|
13
|
+
root (str): Root directory where the dataset should be saved.
|
|
14
|
+
split (str, optional): "train", "val", or "test". (default: "train")
|
|
15
|
+
transform (callable, optional): Transform function.
|
|
16
|
+
pre_transform (callable, optional): Pre-transform function.
|
|
17
|
+
force_reload (bool, optional): Whether to re-process the dataset.
|
|
18
|
+
"""
|
|
19
|
+
|
|
20
|
+
url = "https://raw.githubusercontent.com/villmow/datasets_knowledge_embedding/master/FB15k-237"
|
|
21
|
+
|
|
22
|
+
def __init__(
|
|
23
|
+
self,
|
|
24
|
+
root: str,
|
|
25
|
+
split: str = "train",
|
|
26
|
+
transform: Optional[Callable] = None,
|
|
27
|
+
pre_transform: Optional[Callable] = None,
|
|
28
|
+
force_reload: bool = False,
|
|
29
|
+
):
|
|
30
|
+
if split not in {"train", "val", "test"}:
|
|
31
|
+
raise ValueError(f"Invalid split argument (got {split})")
|
|
32
|
+
self.split = split
|
|
33
|
+
super().__init__(root, transform, pre_transform, force_reload=force_reload)
|
|
34
|
+
idx = ["train", "val", "test"].index(split)
|
|
35
|
+
self.load(self.processed_paths[idx])
|
|
36
|
+
|
|
37
|
+
@property
|
|
38
|
+
def raw_file_names(self) -> List[str]:
|
|
39
|
+
return ["train.txt", "valid.txt", "test.txt"]
|
|
40
|
+
|
|
41
|
+
@property
|
|
42
|
+
def processed_file_names(self) -> List[str]:
|
|
43
|
+
return ["train_data.pt", "val_data.pt", "test_data.pt"]
|
|
44
|
+
|
|
45
|
+
def download(self):
|
|
46
|
+
for filename in self.raw_file_names:
|
|
47
|
+
fs.cp(f"{self.url}/{filename}", self.raw_dir)
|
|
48
|
+
|
|
49
|
+
def process(self):
|
|
50
|
+
# Map entities and relations to integer IDs
|
|
51
|
+
entities, relations = {}, {}
|
|
52
|
+
for path in self.raw_paths:
|
|
53
|
+
with open(path) as f:
|
|
54
|
+
lines = f.read().split("\n")[:-1]
|
|
55
|
+
for line in lines:
|
|
56
|
+
parts = line.split()
|
|
57
|
+
if len(parts) >= 3:
|
|
58
|
+
s, r, d = parts[0], parts[1], parts[2]
|
|
59
|
+
if s not in entities:
|
|
60
|
+
entities[s] = len(entities)
|
|
61
|
+
if d not in entities:
|
|
62
|
+
entities[d] = len(entities)
|
|
63
|
+
if r not in relations:
|
|
64
|
+
relations[r] = len(relations)
|
|
65
|
+
|
|
66
|
+
for in_path, out_path in zip(self.raw_paths, self.processed_paths):
|
|
67
|
+
srcs, dsts, rels = [], [], []
|
|
68
|
+
with open(in_path) as f:
|
|
69
|
+
lines = f.read().split("\n")[:-1]
|
|
70
|
+
for line in lines:
|
|
71
|
+
parts = line.split()
|
|
72
|
+
if len(parts) >= 3:
|
|
73
|
+
srcs.append(entities[parts[0]])
|
|
74
|
+
rels.append(relations[parts[1]])
|
|
75
|
+
dsts.append(entities[parts[2]])
|
|
76
|
+
|
|
77
|
+
edge_index = np.array([srcs, dsts], dtype=np.int64)
|
|
78
|
+
edge_type = np.array(rels, dtype=np.int64)
|
|
79
|
+
|
|
80
|
+
data = Data(
|
|
81
|
+
edge_index=ops.convert_to_tensor(edge_index, dtype="int64"),
|
|
82
|
+
edge_type=ops.convert_to_tensor(edge_type, dtype="int64"),
|
|
83
|
+
num_nodes=len(entities),
|
|
84
|
+
)
|
|
85
|
+
|
|
86
|
+
if self.pre_transform is not None:
|
|
87
|
+
data = self.pre_transform(data)
|
|
88
|
+
|
|
89
|
+
self.save([data], out_path)
|
|
90
|
+
|
|
@@ -0,0 +1,69 @@
|
|
|
1
|
+
import glob
|
|
2
|
+
import os
|
|
3
|
+
import os.path as osp
|
|
4
|
+
from typing import Callable, List, Optional
|
|
5
|
+
|
|
6
|
+
import numpy as np
|
|
7
|
+
|
|
8
|
+
from k3_node.data import InMemoryDataset
|
|
9
|
+
from k3_node.data.download import download_url
|
|
10
|
+
from k3_node.data.extract import extract_zip
|
|
11
|
+
from k3_node.io.off import read_off
|
|
12
|
+
|
|
13
|
+
|
|
14
|
+
class GeometricShapes(InMemoryDataset):
|
|
15
|
+
r"""Synthetic meshes of 40 geometric shapes such as cubes, spheres or pyramids (one training and
|
|
16
|
+
one test mesh per shape), as in PyG.
|
|
17
|
+
|
|
18
|
+
The graphs hold mesh faces (``face``) and vertex positions (``pos``) but no edges. Use
|
|
19
|
+
:class:`~k3_node.transforms.FaceToEdge` to turn a mesh into a graph, or
|
|
20
|
+
:class:`~k3_node.transforms.SamplePoints` to sample a point cloud from its surface.
|
|
21
|
+
|
|
22
|
+
Args:
|
|
23
|
+
root (str): Root directory where the dataset should be saved.
|
|
24
|
+
train (bool, optional): Loads the training meshes if ``True``, else the test meshes.
|
|
25
|
+
transform (callable, optional): A function applied to each graph when it is accessed.
|
|
26
|
+
pre_transform (callable, optional): A function applied to each graph before saving.
|
|
27
|
+
pre_filter (callable, optional): A function deciding which graphs to keep.
|
|
28
|
+
force_reload (bool, optional): Whether to re-process the dataset. (default: ``False``)
|
|
29
|
+
"""
|
|
30
|
+
|
|
31
|
+
url = 'https://github.com/Yannick-S/geometric_shapes/raw/master/raw.zip'
|
|
32
|
+
|
|
33
|
+
def __init__(self, root: str, train: bool = True, transform: Optional[Callable] = None,
|
|
34
|
+
pre_transform: Optional[Callable] = None, pre_filter: Optional[Callable] = None,
|
|
35
|
+
force_reload: bool = False):
|
|
36
|
+
super().__init__(root, transform, pre_transform, pre_filter, force_reload=force_reload)
|
|
37
|
+
self.load(self.processed_paths[0] if train else self.processed_paths[1])
|
|
38
|
+
|
|
39
|
+
@property
|
|
40
|
+
def raw_file_names(self) -> str:
|
|
41
|
+
return '2d_circle'
|
|
42
|
+
|
|
43
|
+
@property
|
|
44
|
+
def processed_file_names(self) -> List[str]:
|
|
45
|
+
return ['training.pt', 'test.pt']
|
|
46
|
+
|
|
47
|
+
def download(self):
|
|
48
|
+
path = download_url(self.url, self.root)
|
|
49
|
+
extract_zip(path, self.root)
|
|
50
|
+
os.unlink(path)
|
|
51
|
+
|
|
52
|
+
def process(self):
|
|
53
|
+
self.save(self._process_set('train'), self.processed_paths[0])
|
|
54
|
+
self.save(self._process_set('test'), self.processed_paths[1])
|
|
55
|
+
|
|
56
|
+
def _process_set(self, split: str):
|
|
57
|
+
categories = sorted(x.split(os.sep)[-2] for x in glob.glob(osp.join(self.raw_dir, '*', '')))
|
|
58
|
+
data_list = []
|
|
59
|
+
for target, category in enumerate(categories):
|
|
60
|
+
for path in sorted(glob.glob(osp.join(self.raw_dir, category, split, '*.off'))):
|
|
61
|
+
data = read_off(path)
|
|
62
|
+
data.pos = data.pos - data.pos.mean(axis=0, keepdims=True)
|
|
63
|
+
data.y = np.array([target], dtype=np.int64)
|
|
64
|
+
data_list.append(data)
|
|
65
|
+
if self.pre_filter is not None:
|
|
66
|
+
data_list = [d for d in data_list if self.pre_filter(d)]
|
|
67
|
+
if self.pre_transform is not None:
|
|
68
|
+
data_list = [self.pre_transform(d) for d in data_list]
|
|
69
|
+
return data_list
|
|
@@ -0,0 +1,51 @@
|
|
|
1
|
+
from typing import Callable, Optional
|
|
2
|
+
import numpy as np
|
|
3
|
+
from keras import ops
|
|
4
|
+
|
|
5
|
+
from k3_node.data import Data, InMemoryDataset
|
|
6
|
+
from k3_node.io import fs
|
|
7
|
+
|
|
8
|
+
|
|
9
|
+
class GitHub(InMemoryDataset):
|
|
10
|
+
r"""The GitHub Web and ML Developers dataset."""
|
|
11
|
+
|
|
12
|
+
url = "https://graphmining.ai/datasets/ptg/github.npz"
|
|
13
|
+
|
|
14
|
+
def __init__(
|
|
15
|
+
self,
|
|
16
|
+
root: str,
|
|
17
|
+
transform: Optional[Callable] = None,
|
|
18
|
+
pre_transform: Optional[Callable] = None,
|
|
19
|
+
force_reload: bool = False,
|
|
20
|
+
):
|
|
21
|
+
super().__init__(root, transform, pre_transform, force_reload=force_reload)
|
|
22
|
+
self.load(self.processed_paths[0])
|
|
23
|
+
|
|
24
|
+
@property
|
|
25
|
+
def raw_file_names(self) -> str:
|
|
26
|
+
return "github.npz"
|
|
27
|
+
|
|
28
|
+
@property
|
|
29
|
+
def processed_file_names(self) -> str:
|
|
30
|
+
return "data.pt"
|
|
31
|
+
|
|
32
|
+
def download(self):
|
|
33
|
+
fs.cp(self.url, self.raw_dir)
|
|
34
|
+
|
|
35
|
+
def process(self):
|
|
36
|
+
data = np.load(self.raw_paths[0], allow_pickle=True)
|
|
37
|
+
x = data["features"].astype(np.float32)
|
|
38
|
+
y = data["target"].astype(np.int64)
|
|
39
|
+
edge_index = data["edges"].astype(np.int64).T
|
|
40
|
+
|
|
41
|
+
data_obj = Data(
|
|
42
|
+
x=ops.convert_to_tensor(x, dtype="float32"),
|
|
43
|
+
y=ops.convert_to_tensor(y, dtype="int64"),
|
|
44
|
+
edge_index=ops.convert_to_tensor(edge_index, dtype="int64"),
|
|
45
|
+
)
|
|
46
|
+
|
|
47
|
+
if self.pre_transform is not None:
|
|
48
|
+
data_obj = self.pre_transform(data_obj)
|
|
49
|
+
|
|
50
|
+
self.save([data_obj], self.processed_paths[0])
|
|
51
|
+
|
|
@@ -0,0 +1,20 @@
|
|
|
1
|
+
from k3_node.data import Data
|
|
2
|
+
from k3_node.datasets.graph_generator.base import GraphGenerator
|
|
3
|
+
from k3_node.utils.random import barabasi_albert_graph
|
|
4
|
+
|
|
5
|
+
|
|
6
|
+
class BAGraph(GraphGenerator):
|
|
7
|
+
r"""Generates random Barabasi-Albert (BA) graphs."""
|
|
8
|
+
|
|
9
|
+
def __init__(self, num_nodes: int, num_edges: int):
|
|
10
|
+
super().__init__()
|
|
11
|
+
self.num_nodes = num_nodes
|
|
12
|
+
self.num_edges = num_edges
|
|
13
|
+
|
|
14
|
+
def __call__(self) -> Data:
|
|
15
|
+
edge_index = barabasi_albert_graph(self.num_nodes, self.num_edges)
|
|
16
|
+
return Data(num_nodes=self.num_nodes, edge_index=edge_index)
|
|
17
|
+
|
|
18
|
+
def __repr__(self) -> str:
|
|
19
|
+
return f"{self.__class__.__name__}(num_nodes={self.num_nodes}, num_edges={self.num_edges})"
|
|
20
|
+
|
|
@@ -0,0 +1,29 @@
|
|
|
1
|
+
from abc import ABC, abstractmethod
|
|
2
|
+
from typing import Any
|
|
3
|
+
from k3_node.data import Data
|
|
4
|
+
|
|
5
|
+
|
|
6
|
+
class GraphGenerator(ABC):
|
|
7
|
+
r"""An abstract base class for generating synthetic graphs."""
|
|
8
|
+
|
|
9
|
+
@abstractmethod
|
|
10
|
+
def __call__(self) -> Data:
|
|
11
|
+
raise NotImplementedError
|
|
12
|
+
|
|
13
|
+
@staticmethod
|
|
14
|
+
def resolve(query: Any, *args: Any, **kwargs: Any) -> "GraphGenerator":
|
|
15
|
+
if isinstance(query, GraphGenerator):
|
|
16
|
+
return query
|
|
17
|
+
if isinstance(query, str):
|
|
18
|
+
query = query.lower()
|
|
19
|
+
if query in ["ba", "bagraph", "barabasi_albert"]:
|
|
20
|
+
from k3_node.datasets.graph_generator.ba_graph import BAGraph
|
|
21
|
+
return BAGraph(*args, **kwargs)
|
|
22
|
+
elif query in ["er", "ergraph", "erdos_renyi"]:
|
|
23
|
+
from k3_node.datasets.graph_generator.er_graph import ERGraph
|
|
24
|
+
return ERGraph(*args, **kwargs)
|
|
25
|
+
raise ValueError(f"Could not resolve graph generator: {query}")
|
|
26
|
+
|
|
27
|
+
def __repr__(self) -> str:
|
|
28
|
+
return f"{self.__class__.__name__}()"
|
|
29
|
+
|
|
@@ -0,0 +1,21 @@
|
|
|
1
|
+
from k3_node.data import Data
|
|
2
|
+
from k3_node.datasets.graph_generator.base import GraphGenerator
|
|
3
|
+
from k3_node.utils.random import erdos_renyi_graph
|
|
4
|
+
|
|
5
|
+
|
|
6
|
+
class ERGraph(GraphGenerator):
|
|
7
|
+
r"""Generates random Erdos-Renyi (ER) graphs."""
|
|
8
|
+
|
|
9
|
+
def __init__(self, num_nodes: int, edge_prob: float, directed: bool = False):
|
|
10
|
+
super().__init__()
|
|
11
|
+
self.num_nodes = num_nodes
|
|
12
|
+
self.edge_prob = edge_prob
|
|
13
|
+
self.directed = directed
|
|
14
|
+
|
|
15
|
+
def __call__(self) -> Data:
|
|
16
|
+
edge_index = erdos_renyi_graph(self.num_nodes, self.edge_prob, directed=self.directed)
|
|
17
|
+
return Data(num_nodes=self.num_nodes, edge_index=edge_index)
|
|
18
|
+
|
|
19
|
+
def __repr__(self) -> str:
|
|
20
|
+
return f"{self.__class__.__name__}(num_nodes={self.num_nodes}, edge_prob={self.edge_prob})"
|
|
21
|
+
|
|
@@ -0,0 +1,58 @@
|
|
|
1
|
+
from typing import Callable, List, Optional
|
|
2
|
+
|
|
3
|
+
import numpy as np
|
|
4
|
+
|
|
5
|
+
from k3_node.data import Data, InMemoryDataset
|
|
6
|
+
from k3_node.data.download import download_url
|
|
7
|
+
|
|
8
|
+
|
|
9
|
+
class ICEWS18(InMemoryDataset):
|
|
10
|
+
r"""The ICEWS18 temporal knowledge graph (Integrated Crisis Early Warning System, events from
|
|
11
|
+
1/1/2018 to 10/31/2018 at a daily resolution), used by RE-Net: every graph is one event with
|
|
12
|
+
subject ``sub``, relation ``rel``, object ``obj`` and day ``t``.
|
|
13
|
+
|
|
14
|
+
Args:
|
|
15
|
+
root (str): Root directory where the dataset should be saved.
|
|
16
|
+
split (str, optional): ``"train"``, ``"val"`` or ``"test"``. (default: ``"train"``)
|
|
17
|
+
transform (callable, optional): A function applied to each event when it is accessed.
|
|
18
|
+
pre_transform (callable, optional): A function applied to the events in time order before
|
|
19
|
+
saving, e.g. :meth:`~k3_node.models.RENet.pre_transform`.
|
|
20
|
+
force_reload (bool, optional): Whether to re-process the dataset. (default: ``False``)
|
|
21
|
+
"""
|
|
22
|
+
|
|
23
|
+
url = 'https://github.com/INK-USC/RE-Net/raw/master/data/ICEWS18'
|
|
24
|
+
splits = [0, 373018, 419013, 468558]
|
|
25
|
+
num_nodes, num_rels = 23033, 256
|
|
26
|
+
|
|
27
|
+
def __init__(self, root: str, split: str = 'train', transform: Optional[Callable] = None,
|
|
28
|
+
pre_transform: Optional[Callable] = None, force_reload: bool = False):
|
|
29
|
+
assert split in ['train', 'val', 'test']
|
|
30
|
+
super().__init__(root, transform, pre_transform, force_reload=force_reload)
|
|
31
|
+
self.load(self.processed_paths[['train', 'val', 'test'].index(split)])
|
|
32
|
+
|
|
33
|
+
@property
|
|
34
|
+
def raw_file_names(self) -> List[str]:
|
|
35
|
+
return [f'{name}.txt' for name in ['train', 'valid', 'test']]
|
|
36
|
+
|
|
37
|
+
@property
|
|
38
|
+
def processed_file_names(self) -> List[str]:
|
|
39
|
+
return ['train.pt', 'val.pt', 'test.pt']
|
|
40
|
+
|
|
41
|
+
def download(self):
|
|
42
|
+
for filename in self.raw_file_names:
|
|
43
|
+
download_url(f'{self.url}/{filename}', self.raw_dir)
|
|
44
|
+
|
|
45
|
+
def process(self):
|
|
46
|
+
events = np.concatenate([np.loadtxt(path, delimiter='\t', usecols=range(4), dtype=np.int64)
|
|
47
|
+
for path in self.raw_paths])
|
|
48
|
+
events[:, 3] //= 24 # hours -> days
|
|
49
|
+
events -= events.min(axis=0, keepdims=True)
|
|
50
|
+
data_list = []
|
|
51
|
+
for sub, rel, obj, t in events.tolist():
|
|
52
|
+
data = Data(sub=sub, rel=rel, obj=obj, t=t)
|
|
53
|
+
if self.pre_transform is not None:
|
|
54
|
+
data = self.pre_transform(data)
|
|
55
|
+
data_list.append(data)
|
|
56
|
+
s = self.splits
|
|
57
|
+
for i in range(3):
|
|
58
|
+
self.save(data_list[s[i]:s[i + 1]], self.processed_paths[i])
|