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/imdb.py
ADDED
|
@@ -0,0 +1,96 @@
|
|
|
1
|
+
import os
|
|
2
|
+
import os.path as osp
|
|
3
|
+
from itertools import product
|
|
4
|
+
from typing import Callable, List, Optional
|
|
5
|
+
import numpy as np
|
|
6
|
+
from keras import ops
|
|
7
|
+
|
|
8
|
+
from k3_node.data import HeteroData, InMemoryDataset
|
|
9
|
+
from k3_node.io import fs
|
|
10
|
+
|
|
11
|
+
|
|
12
|
+
class IMDB(InMemoryDataset):
|
|
13
|
+
r"""A subset of the Internet Movie Database (IMDB) containing three types of entities:
|
|
14
|
+
movies, actors, and directors.
|
|
15
|
+
"""
|
|
16
|
+
|
|
17
|
+
url = "https://www.dropbox.com/s/g0btk9ctr1es39x/IMDB_processed.zip?dl=1"
|
|
18
|
+
|
|
19
|
+
def __init__(
|
|
20
|
+
self,
|
|
21
|
+
root: str,
|
|
22
|
+
transform: Optional[Callable] = None,
|
|
23
|
+
pre_transform: Optional[Callable] = None,
|
|
24
|
+
force_reload: bool = False,
|
|
25
|
+
):
|
|
26
|
+
super().__init__(root, transform, pre_transform, force_reload=force_reload)
|
|
27
|
+
self.load(self.processed_paths[0])
|
|
28
|
+
|
|
29
|
+
@property
|
|
30
|
+
def raw_file_names(self) -> List[str]:
|
|
31
|
+
return [
|
|
32
|
+
"adjM.npz",
|
|
33
|
+
"features_0.npz",
|
|
34
|
+
"features_1.npz",
|
|
35
|
+
"features_2.npz",
|
|
36
|
+
"labels.npy",
|
|
37
|
+
"train_val_test_idx.npz",
|
|
38
|
+
]
|
|
39
|
+
|
|
40
|
+
@property
|
|
41
|
+
def processed_file_names(self) -> str:
|
|
42
|
+
return "data.pt"
|
|
43
|
+
|
|
44
|
+
def download(self):
|
|
45
|
+
zip_path = osp.join(self.raw_dir, "IMDB_processed.zip")
|
|
46
|
+
fs.cp(self.url, zip_path, extract=True)
|
|
47
|
+
if osp.exists(zip_path):
|
|
48
|
+
fs.rm(zip_path)
|
|
49
|
+
|
|
50
|
+
def process(self):
|
|
51
|
+
import scipy.sparse as sp
|
|
52
|
+
|
|
53
|
+
data = HeteroData()
|
|
54
|
+
node_types = ["movie", "director", "actor"]
|
|
55
|
+
|
|
56
|
+
for i, node_type in enumerate(node_types):
|
|
57
|
+
x = sp.load_npz(osp.join(self.raw_dir, f"features_{i}.npz"))
|
|
58
|
+
data[node_type].x = ops.convert_to_tensor(
|
|
59
|
+
np.array(x.todense(), dtype=np.float32), dtype="float32"
|
|
60
|
+
)
|
|
61
|
+
|
|
62
|
+
y = np.load(osp.join(self.raw_dir, "labels.npy"))
|
|
63
|
+
data["movie"].y = ops.convert_to_tensor(y.astype(np.int64), dtype="int64")
|
|
64
|
+
|
|
65
|
+
split = np.load(osp.join(self.raw_dir, "train_val_test_idx.npz"))
|
|
66
|
+
for name in ["train", "val", "test"]:
|
|
67
|
+
idx = split[f"{name}_idx"]
|
|
68
|
+
mask = np.zeros(data["movie"].num_nodes, dtype=bool)
|
|
69
|
+
mask[idx] = True
|
|
70
|
+
data["movie"][f"{name}_mask"] = ops.convert_to_tensor(mask, dtype="bool")
|
|
71
|
+
|
|
72
|
+
s = {}
|
|
73
|
+
N_m = data["movie"].num_nodes
|
|
74
|
+
N_d = data["director"].num_nodes
|
|
75
|
+
N_a = data["actor"].num_nodes
|
|
76
|
+
s["movie"] = (0, N_m)
|
|
77
|
+
s["director"] = (N_m, N_m + N_d)
|
|
78
|
+
s["actor"] = (N_m + N_d, N_m + N_d + N_a)
|
|
79
|
+
|
|
80
|
+
A = sp.load_npz(osp.join(self.raw_dir, "adjM.npz"))
|
|
81
|
+
for src, dst in product(node_types, node_types):
|
|
82
|
+
A_sub = A[s[src][0] : s[src][1], s[dst][0] : s[dst][1]].tocoo()
|
|
83
|
+
if A_sub.nnz > 0:
|
|
84
|
+
row = np.array(A_sub.row, dtype=np.int64)
|
|
85
|
+
col = np.array(A_sub.col, dtype=np.int64)
|
|
86
|
+
edge_index = np.stack([row, col], axis=0)
|
|
87
|
+
data[src, dst].edge_index = ops.convert_to_tensor(edge_index, dtype="int64")
|
|
88
|
+
|
|
89
|
+
if self.pre_transform is not None:
|
|
90
|
+
data = self.pre_transform(data)
|
|
91
|
+
|
|
92
|
+
self.save([data], self.processed_paths[0])
|
|
93
|
+
|
|
94
|
+
def __repr__(self) -> str:
|
|
95
|
+
return f"{self.__class__.__name__}()"
|
|
96
|
+
|
|
@@ -0,0 +1,56 @@
|
|
|
1
|
+
import os
|
|
2
|
+
import os.path as osp
|
|
3
|
+
|
|
4
|
+
import numpy as np
|
|
5
|
+
|
|
6
|
+
|
|
7
|
+
class JODIEDataset:
|
|
8
|
+
r"""The temporal interaction datasets of `"JODIE: Predicting Dynamic Embedding Trajectory in
|
|
9
|
+
Temporal Interaction Networks" <https://cs.stanford.edu/~srijan/pubs/jodie-kdd2019.pdf>`_:
|
|
10
|
+
``"wikipedia"``, ``"reddit"``, ``"mooc"`` and ``"lastfm"``.
|
|
11
|
+
|
|
12
|
+
``dataset[0]`` is a :class:`~k3_node.data.TemporalData` event stream: user ``src`` interacts
|
|
13
|
+
with item ``dst`` (numbered after the users) at time ``t``, with features ``msg`` and label
|
|
14
|
+
``y``. MOOC (7,144 nodes, 411,749 events, 4 features) is the smallest download (40 MB).
|
|
15
|
+
|
|
16
|
+
Args:
|
|
17
|
+
root (str): Root directory where the dataset should be saved.
|
|
18
|
+
name (str): The name of the dataset.
|
|
19
|
+
"""
|
|
20
|
+
|
|
21
|
+
url = "https://snap.stanford.edu/jodie/{}.csv"
|
|
22
|
+
names = ["wikipedia", "reddit", "mooc", "lastfm"]
|
|
23
|
+
|
|
24
|
+
def __init__(self, root: str, name: str):
|
|
25
|
+
self.name = name.lower()
|
|
26
|
+
assert self.name in self.names
|
|
27
|
+
folder = osp.join(root, self.name)
|
|
28
|
+
cache = osp.join(folder, "processed.npz")
|
|
29
|
+
if not osp.exists(cache):
|
|
30
|
+
os.makedirs(folder, exist_ok=True)
|
|
31
|
+
csv = osp.join(folder, f"{self.name}.csv")
|
|
32
|
+
if not osp.exists(csv):
|
|
33
|
+
from k3_node.data.download import download_url
|
|
34
|
+
|
|
35
|
+
download_url(self.url.format(self.name), folder)
|
|
36
|
+
import pandas as pd
|
|
37
|
+
|
|
38
|
+
df = pd.read_csv(csv, skiprows=1, header=None)
|
|
39
|
+
src = df.iloc[:, 0].values.astype(np.int64)
|
|
40
|
+
dst = df.iloc[:, 1].values.astype(np.int64) + int(src.max()) + 1
|
|
41
|
+
np.savez(cache, src=src, dst=dst, t=df.iloc[:, 2].values.astype(np.int64),
|
|
42
|
+
y=df.iloc[:, 3].values.astype(np.int64), msg=df.iloc[:, 4:].values.astype(np.float32))
|
|
43
|
+
self._arrays = dict(np.load(cache))
|
|
44
|
+
|
|
45
|
+
def __len__(self):
|
|
46
|
+
return 1
|
|
47
|
+
|
|
48
|
+
def __getitem__(self, idx):
|
|
49
|
+
from k3_node.data import TemporalData
|
|
50
|
+
|
|
51
|
+
if idx != 0:
|
|
52
|
+
raise IndexError(idx)
|
|
53
|
+
return TemporalData(**self._arrays)
|
|
54
|
+
|
|
55
|
+
def __repr__(self):
|
|
56
|
+
return f"JODIEDataset({self.name})"
|
|
@@ -0,0 +1,56 @@
|
|
|
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
|
+
|
|
7
|
+
|
|
8
|
+
class KarateClub(InMemoryDataset):
|
|
9
|
+
r"""Zachary's karate club network from the `"An Information Flow Model for
|
|
10
|
+
Conflict and Fission in Small Groups" paper, containing 34 nodes and 156
|
|
11
|
+
undirected edges labeled into 4 community classes.
|
|
12
|
+
"""
|
|
13
|
+
|
|
14
|
+
def __init__(self, transform: Optional[Callable] = None):
|
|
15
|
+
super().__init__(None, transform)
|
|
16
|
+
|
|
17
|
+
row = [
|
|
18
|
+
0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 1, 1, 1, 1, 1, 1,
|
|
19
|
+
1, 1, 1, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 3, 3, 3, 3, 3, 3, 4, 4, 4,
|
|
20
|
+
5, 5, 5, 5, 6, 6, 6, 6, 7, 7, 7, 7, 8, 8, 8, 8, 8, 9, 9, 10, 10,
|
|
21
|
+
10, 11, 12, 12, 13, 13, 13, 13, 13, 14, 14, 15, 15, 16, 16, 17, 17,
|
|
22
|
+
18, 18, 19, 19, 19, 20, 20, 21, 21, 22, 22, 23, 23, 23, 23, 23, 24,
|
|
23
|
+
24, 24, 25, 25, 25, 26, 26, 27, 27, 27, 27, 28, 28, 28, 29, 29, 29,
|
|
24
|
+
29, 30, 30, 30, 30, 31, 31, 31, 31, 31, 31, 32, 32, 32, 32, 32, 32,
|
|
25
|
+
32, 32, 32, 32, 32, 32, 33, 33, 33, 33, 33, 33, 33, 33, 33, 33, 33,
|
|
26
|
+
33, 33, 33, 33, 33, 33
|
|
27
|
+
]
|
|
28
|
+
col = [
|
|
29
|
+
1, 2, 3, 4, 5, 6, 7, 8, 10, 11, 12, 13, 17, 19, 21, 31, 0, 2, 3, 7,
|
|
30
|
+
13, 17, 19, 21, 30, 0, 1, 3, 7, 8, 9, 13, 27, 28, 32, 0, 1, 2, 7,
|
|
31
|
+
12, 13, 0, 6, 10, 0, 6, 10, 16, 0, 4, 5, 16, 0, 1, 2, 3, 0, 2, 30,
|
|
32
|
+
32, 33, 2, 33, 0, 4, 5, 0, 0, 3, 0, 1, 2, 3, 33, 32, 33, 32, 33, 5,
|
|
33
|
+
6, 0, 1, 32, 33, 0, 1, 33, 32, 33, 0, 1, 32, 33, 25, 27, 29, 32,
|
|
34
|
+
33, 25, 27, 31, 23, 24, 31, 29, 33, 2, 23, 24, 33, 2, 31, 33, 23,
|
|
35
|
+
26, 32, 33, 1, 8, 32, 33, 0, 24, 25, 28, 32, 33, 2, 8, 14, 15, 18,
|
|
36
|
+
20, 22, 23, 29, 30, 31, 33, 8, 9, 13, 14, 15, 18, 19, 20, 22, 23,
|
|
37
|
+
26, 27, 28, 29, 30, 31, 32
|
|
38
|
+
]
|
|
39
|
+
edge_index = ops.convert_to_tensor(np.array([row, col], dtype=np.int64), dtype="int64")
|
|
40
|
+
|
|
41
|
+
y_np = np.array([
|
|
42
|
+
1, 1, 1, 1, 3, 3, 3, 1, 0, 1, 3, 1, 1, 1, 0, 0, 3, 1, 0, 1, 0, 1,
|
|
43
|
+
0, 0, 2, 2, 0, 0, 2, 0, 0, 2, 0, 0
|
|
44
|
+
], dtype=np.int64)
|
|
45
|
+
y = ops.convert_to_tensor(y_np, dtype="int64")
|
|
46
|
+
|
|
47
|
+
x = ops.convert_to_tensor(np.eye(34, dtype=np.float32), dtype="float32")
|
|
48
|
+
|
|
49
|
+
train_mask_np = np.zeros(34, dtype=bool)
|
|
50
|
+
for i in range(int(np.max(y_np)) + 1):
|
|
51
|
+
train_mask_np[np.where(y_np == i)[0][0]] = True
|
|
52
|
+
train_mask = ops.convert_to_tensor(train_mask_np, dtype="bool")
|
|
53
|
+
|
|
54
|
+
data = Data(x=x, edge_index=edge_index, y=y, train_mask=train_mask)
|
|
55
|
+
self.data, self.slices = self.collate([data])
|
|
56
|
+
|
|
@@ -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 LastFMAsia(InMemoryDataset):
|
|
10
|
+
r"""The LastFM Asia Network dataset."""
|
|
11
|
+
|
|
12
|
+
url = "https://graphmining.ai/datasets/ptg/lastfm_asia.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 "lastfm_asia.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,50 @@
|
|
|
1
|
+
import os
|
|
2
|
+
from typing import Callable, Optional
|
|
3
|
+
|
|
4
|
+
import numpy as np
|
|
5
|
+
|
|
6
|
+
from k3_node.data import Data, InMemoryDataset
|
|
7
|
+
|
|
8
|
+
|
|
9
|
+
class MeshCorrespondence(InMemoryDataset):
|
|
10
|
+
r"""A small shape correspondence dataset in the spirit of FAUST: randomly deformed copies of one
|
|
11
|
+
:class:`GeometricShapes` mesh (the monkey head by default). Every copy is smoothly bent and
|
|
12
|
+
slightly rotated; the task is to recognize every vertex, so ``y`` holds the vertex indices.
|
|
13
|
+
|
|
14
|
+
The graphs hold ``pos`` and mesh faces ``face``; use :class:`~k3_node.transforms.FaceToEdge`
|
|
15
|
+
to connect the vertices. As in FAUST, there are 80 training and 20 test meshes by default.
|
|
16
|
+
|
|
17
|
+
Args:
|
|
18
|
+
root (str): Directory where GeometricShapes is (or will be) downloaded.
|
|
19
|
+
train (bool, optional): Build the training meshes if ``True``, else the test meshes.
|
|
20
|
+
num_meshes (int, optional): Number of meshes. (default: ``80`` / ``20``)
|
|
21
|
+
shape (str, optional): The GeometricShapes mesh to deform. (default: ``"3d_monkey"``)
|
|
22
|
+
transform (callable, optional): A function applied to each mesh when it is accessed.
|
|
23
|
+
pre_transform (callable, optional): A function applied to each mesh once, when built.
|
|
24
|
+
seed (int, optional): Random seed. (default: ``0``)
|
|
25
|
+
"""
|
|
26
|
+
|
|
27
|
+
def __init__(self, root: str, train: bool = True, num_meshes: Optional[int] = None, shape: str = "3d_monkey",
|
|
28
|
+
transform: Optional[Callable] = None, pre_transform: Optional[Callable] = None, seed: int = 0):
|
|
29
|
+
super().__init__(None, transform)
|
|
30
|
+
from k3_node.datasets.geometric_shapes import GeometricShapes
|
|
31
|
+
from k3_node.transforms.spatial import _np
|
|
32
|
+
|
|
33
|
+
shapes = GeometricShapes(root, train=True)
|
|
34
|
+
names = sorted(os.listdir(shapes.raw_dir))
|
|
35
|
+
base = next(shapes[i] for i in range(len(shapes)) if names[int(_np(shapes[i].y)[0])] == shape)
|
|
36
|
+
pos, face = _np(base.pos).astype(np.float64), _np(base.face)
|
|
37
|
+
pos = pos / np.abs(pos).max()
|
|
38
|
+
|
|
39
|
+
rng = np.random.default_rng(seed + (0 if train else 1))
|
|
40
|
+
graphs = []
|
|
41
|
+
for _ in range(num_meshes or (80 if train else 20)):
|
|
42
|
+
w = rng.normal(size=(3, 3)) * 1.5
|
|
43
|
+
bent = pos + 0.12 * np.sin(pos @ w + rng.uniform(0, 2 * np.pi, size=3)) # smooth deformation
|
|
44
|
+
angle = np.deg2rad(rng.uniform(-15, 15))
|
|
45
|
+
c, s = np.cos(angle), np.sin(angle)
|
|
46
|
+
rot = np.array([[c, 0, s], [0, 1, 0], [-s, 0, c]])
|
|
47
|
+
data = Data(pos=(bent @ rot.T).astype(np.float32), face=face.copy(),
|
|
48
|
+
y=np.arange(len(pos), dtype=np.int64))
|
|
49
|
+
graphs.append(pre_transform(data) if pre_transform is not None else data)
|
|
50
|
+
self.data, self.slices = self.collate(graphs)
|
|
@@ -0,0 +1,148 @@
|
|
|
1
|
+
import os
|
|
2
|
+
import os.path as osp
|
|
3
|
+
import re
|
|
4
|
+
import warnings
|
|
5
|
+
from typing import Callable, Dict, Optional, Tuple, Union
|
|
6
|
+
|
|
7
|
+
import numpy as np
|
|
8
|
+
from keras import ops
|
|
9
|
+
|
|
10
|
+
from k3_node.data import InMemoryDataset, download_url, extract_gz
|
|
11
|
+
from k3_node.utils.smiles import from_smiles as default_from_smiles
|
|
12
|
+
|
|
13
|
+
|
|
14
|
+
class MoleculeNet(InMemoryDataset):
|
|
15
|
+
r"""The `MoleculeNet <http://moleculenet.org/datasets-1>`_ benchmark collection
|
|
16
|
+
from the `"MoleculeNet: A Benchmark for Molecular Machine Learning"
|
|
17
|
+
<https://arxiv.org/abs/1703.00564>`_ paper, containing datasets from physical
|
|
18
|
+
chemistry, biophysics and physiology.
|
|
19
|
+
|
|
20
|
+
Args:
|
|
21
|
+
root (str): Root directory where the dataset should be saved.
|
|
22
|
+
name (str): The name of the dataset (:obj:`"ESOL"`, :obj:`"FreeSolv"`,
|
|
23
|
+
:obj:`"Lipo"`, :obj:`"PCBA"`, :obj:`"MUV"`, :obj:`"HIV"`,
|
|
24
|
+
:obj:`"BACE"`, :obj:`"BBBP"`, :obj:`"Tox21"`, :obj:`"ToxCast"`,
|
|
25
|
+
:obj:`"SIDER"`, :obj:`"ClinTox"`).
|
|
26
|
+
transform (callable, optional): A function/transform that takes in a
|
|
27
|
+
:obj:`k3_node.data.Data` object and returns a transformed version.
|
|
28
|
+
(default: :obj:`None`)
|
|
29
|
+
pre_transform (callable, optional): A function/transform that takes in a
|
|
30
|
+
:obj:`k3_node.data.Data` object and returns a transformed version.
|
|
31
|
+
(default: :obj:`None`)
|
|
32
|
+
pre_filter (callable, optional): A function that takes in a
|
|
33
|
+
:obj:`k3_node.data.Data` object and returns a boolean value,
|
|
34
|
+
indicating whether the data object should be included in the final
|
|
35
|
+
dataset. (default: :obj:`None`)
|
|
36
|
+
force_reload (bool, optional): Whether to re-process the dataset.
|
|
37
|
+
(default: :obj:`False`)
|
|
38
|
+
from_smiles (callable, optional): A custom function that takes a SMILES
|
|
39
|
+
string and outputs a :obj:`k3_node.data.Data` object.
|
|
40
|
+
(default: :obj:`None`)
|
|
41
|
+
"""
|
|
42
|
+
|
|
43
|
+
url = "https://deepchemdata.s3-us-west-1.amazonaws.com/datasets/{}"
|
|
44
|
+
|
|
45
|
+
# Format: name: (display_name, url_name, csv_name, smiles_idx, y_idx)
|
|
46
|
+
names: Dict[str, Tuple[str, str, str, int, Union[int, slice]]] = {
|
|
47
|
+
"esol": ("ESOL", "delaney-processed.csv", "delaney-processed", -1, -2),
|
|
48
|
+
"freesolv": ("FreeSolv", "SAMPL.csv", "SAMPL", 1, 2),
|
|
49
|
+
"lipo": ("Lipophilicity", "Lipophilicity.csv", "Lipophilicity", 2, 1),
|
|
50
|
+
"pcba": ("PCBA", "pcba.csv.gz", "pcba", -1, slice(0, 128)),
|
|
51
|
+
"muv": ("MUV", "muv.csv.gz", "muv", -1, slice(0, 17)),
|
|
52
|
+
"hiv": ("HIV", "HIV.csv", "HIV", 0, -1),
|
|
53
|
+
"bace": ("BACE", "bace.csv", "bace", 0, 2),
|
|
54
|
+
"bbbp": ("BBBP", "BBBP.csv", "BBBP", -1, -2),
|
|
55
|
+
"tox21": ("Tox21", "tox21.csv.gz", "tox21", -1, slice(0, 12)),
|
|
56
|
+
"toxcast": ("ToxCast", "toxcast_data.csv.gz", "toxcast_data", 0, slice(1, 618)),
|
|
57
|
+
"sider": ("SIDER", "sider.csv.gz", "sider", 0, slice(1, 28)),
|
|
58
|
+
"clintox": ("ClinTox", "clintox.csv.gz", "clintox", 0, slice(1, 3)),
|
|
59
|
+
}
|
|
60
|
+
|
|
61
|
+
def __init__(
|
|
62
|
+
self,
|
|
63
|
+
root: str,
|
|
64
|
+
name: str,
|
|
65
|
+
transform: Optional[Callable] = None,
|
|
66
|
+
pre_transform: Optional[Callable] = None,
|
|
67
|
+
pre_filter: Optional[Callable] = None,
|
|
68
|
+
force_reload: bool = False,
|
|
69
|
+
from_smiles: Optional[Callable] = None,
|
|
70
|
+
) -> None:
|
|
71
|
+
self.name = name.lower()
|
|
72
|
+
if self.name not in self.names:
|
|
73
|
+
raise ValueError(
|
|
74
|
+
f"Unknown dataset name '{name}'. Available names: {list(self.names.keys())}"
|
|
75
|
+
)
|
|
76
|
+
self.from_smiles = from_smiles or default_from_smiles
|
|
77
|
+
super().__init__(
|
|
78
|
+
root,
|
|
79
|
+
transform,
|
|
80
|
+
pre_transform,
|
|
81
|
+
pre_filter,
|
|
82
|
+
force_reload=force_reload,
|
|
83
|
+
)
|
|
84
|
+
self.load(self.processed_paths[0])
|
|
85
|
+
|
|
86
|
+
@property
|
|
87
|
+
def raw_dir(self) -> str:
|
|
88
|
+
return osp.join(self.root, self.name, "raw")
|
|
89
|
+
|
|
90
|
+
@property
|
|
91
|
+
def processed_dir(self) -> str:
|
|
92
|
+
return osp.join(self.root, self.name, "processed")
|
|
93
|
+
|
|
94
|
+
@property
|
|
95
|
+
def raw_file_names(self) -> str:
|
|
96
|
+
return f"{self.names[self.name][2]}.csv"
|
|
97
|
+
|
|
98
|
+
@property
|
|
99
|
+
def processed_file_names(self) -> str:
|
|
100
|
+
return "data.pt"
|
|
101
|
+
|
|
102
|
+
def download(self) -> None:
|
|
103
|
+
url = self.url.format(self.names[self.name][1])
|
|
104
|
+
path = download_url(url, self.raw_dir)
|
|
105
|
+
if self.names[self.name][1].endswith("gz"):
|
|
106
|
+
extract_gz(path, self.raw_dir)
|
|
107
|
+
os.unlink(path)
|
|
108
|
+
|
|
109
|
+
def process(self) -> None:
|
|
110
|
+
with open(self.raw_paths[0], "r", encoding="utf-8") as f:
|
|
111
|
+
dataset = f.read().split("\n")[1:-1]
|
|
112
|
+
dataset = [x for x in dataset if len(x) > 0]
|
|
113
|
+
|
|
114
|
+
data_list = []
|
|
115
|
+
for line in dataset:
|
|
116
|
+
line = re.sub(r'".*"', "", line) # Replace quoted substrings
|
|
117
|
+
values = line.split(",")
|
|
118
|
+
|
|
119
|
+
smiles = values[self.names[self.name][3]]
|
|
120
|
+
labels = values[self.names[self.name][4]]
|
|
121
|
+
labels = labels if isinstance(labels, list) else [labels]
|
|
122
|
+
|
|
123
|
+
ys = [float(y) if len(y) > 0 else float("nan") for y in labels]
|
|
124
|
+
y = ops.convert_to_tensor(np.array(ys, dtype=np.float32).reshape(1, -1), dtype="float32")
|
|
125
|
+
|
|
126
|
+
data = self.from_smiles(smiles)
|
|
127
|
+
data.y = y
|
|
128
|
+
|
|
129
|
+
if data.num_nodes == 0:
|
|
130
|
+
warnings.warn(
|
|
131
|
+
f"Skipping molecule '{smiles}' since it resulted in zero atoms",
|
|
132
|
+
stacklevel=2,
|
|
133
|
+
)
|
|
134
|
+
continue
|
|
135
|
+
|
|
136
|
+
if self.pre_filter is not None and not self.pre_filter(data):
|
|
137
|
+
continue
|
|
138
|
+
|
|
139
|
+
if self.pre_transform is not None:
|
|
140
|
+
data = self.pre_transform(data)
|
|
141
|
+
|
|
142
|
+
data_list.append(data)
|
|
143
|
+
|
|
144
|
+
self.save(data_list, self.processed_paths[0])
|
|
145
|
+
|
|
146
|
+
def __repr__(self) -> str:
|
|
147
|
+
return f"{self.names[self.name][0]}({len(self)})"
|
|
148
|
+
|
|
@@ -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 MotifGenerator(ABC):
|
|
7
|
+
r"""An abstract base class for generating motifs."""
|
|
8
|
+
|
|
9
|
+
@abstractmethod
|
|
10
|
+
def __call__(self) -> Data:
|
|
11
|
+
raise NotImplementedError
|
|
12
|
+
|
|
13
|
+
@staticmethod
|
|
14
|
+
def resolve(query: Any, *args: Any, **kwargs: Any) -> "MotifGenerator":
|
|
15
|
+
if isinstance(query, MotifGenerator):
|
|
16
|
+
return query
|
|
17
|
+
if isinstance(query, str):
|
|
18
|
+
query = query.lower()
|
|
19
|
+
if query in ["house", "housemotif"]:
|
|
20
|
+
from k3_node.datasets.motif_generator.house import HouseMotif
|
|
21
|
+
return HouseMotif(*args, **kwargs)
|
|
22
|
+
elif query in ["cycle", "cyclemotif"]:
|
|
23
|
+
from k3_node.datasets.motif_generator.cycle import CycleMotif
|
|
24
|
+
return CycleMotif(*args, **kwargs)
|
|
25
|
+
raise ValueError(f"Could not resolve motif generator: {query}")
|
|
26
|
+
|
|
27
|
+
def __repr__(self) -> str:
|
|
28
|
+
return f"{self.__class__.__name__}()"
|
|
29
|
+
|
|
@@ -0,0 +1,17 @@
|
|
|
1
|
+
from typing import Any, Optional
|
|
2
|
+
from k3_node.data import Data
|
|
3
|
+
from k3_node.datasets.motif_generator.base import MotifGenerator
|
|
4
|
+
|
|
5
|
+
|
|
6
|
+
class CustomMotif(MotifGenerator):
|
|
7
|
+
r"""Generates a motif based on a custom structure coming from a Data object."""
|
|
8
|
+
|
|
9
|
+
def __init__(self, structure: Any):
|
|
10
|
+
super().__init__()
|
|
11
|
+
if not isinstance(structure, Data):
|
|
12
|
+
raise ValueError(f"Expected structure of type Data, got {type(structure)}")
|
|
13
|
+
self.structure = structure
|
|
14
|
+
|
|
15
|
+
def __call__(self) -> Data:
|
|
16
|
+
return self.structure
|
|
17
|
+
|
|
@@ -0,0 +1,25 @@
|
|
|
1
|
+
import numpy as np
|
|
2
|
+
from keras import ops
|
|
3
|
+
|
|
4
|
+
from k3_node.data import Data
|
|
5
|
+
from k3_node.datasets.motif_generator.custom import CustomMotif
|
|
6
|
+
|
|
7
|
+
|
|
8
|
+
class CycleMotif(CustomMotif):
|
|
9
|
+
r"""Generates the cycle motif from the "GNNExplainer" paper."""
|
|
10
|
+
|
|
11
|
+
def __init__(self, num_nodes: int):
|
|
12
|
+
self.num_nodes = num_nodes
|
|
13
|
+
|
|
14
|
+
row = np.repeat(np.arange(num_nodes), 2)
|
|
15
|
+
col1 = np.mod(np.arange(-1, num_nodes - 1), num_nodes)
|
|
16
|
+
col2 = np.mod(np.arange(1, num_nodes + 1), num_nodes)
|
|
17
|
+
col = np.sort(np.stack([col1, col2], axis=1), axis=-1).flatten()
|
|
18
|
+
|
|
19
|
+
edge_index = ops.convert_to_tensor(np.stack([row, col], axis=0), dtype="int64")
|
|
20
|
+
structure = Data(num_nodes=num_nodes, edge_index=edge_index)
|
|
21
|
+
super().__init__(structure)
|
|
22
|
+
|
|
23
|
+
def __repr__(self) -> str:
|
|
24
|
+
return f"{self.__class__.__name__}({self.num_nodes})"
|
|
25
|
+
|
|
@@ -0,0 +1,27 @@
|
|
|
1
|
+
import numpy as np
|
|
2
|
+
from keras import ops
|
|
3
|
+
|
|
4
|
+
from k3_node.data import Data
|
|
5
|
+
from k3_node.datasets.motif_generator.custom import CustomMotif
|
|
6
|
+
|
|
7
|
+
|
|
8
|
+
class HouseMotif(CustomMotif):
|
|
9
|
+
r"""Generates the house-structured motif from the "GNNExplainer" paper,
|
|
10
|
+
containing 5 nodes and 6 undirected edges.
|
|
11
|
+
"""
|
|
12
|
+
|
|
13
|
+
def __init__(self):
|
|
14
|
+
edge_index = ops.convert_to_tensor(
|
|
15
|
+
np.array(
|
|
16
|
+
[
|
|
17
|
+
[0, 0, 0, 1, 1, 1, 2, 2, 3, 3, 4, 4],
|
|
18
|
+
[1, 3, 4, 4, 2, 0, 1, 3, 2, 0, 0, 1],
|
|
19
|
+
],
|
|
20
|
+
dtype=np.int64,
|
|
21
|
+
),
|
|
22
|
+
dtype="int64",
|
|
23
|
+
)
|
|
24
|
+
y = ops.convert_to_tensor(np.array([0, 0, 1, 1, 2], dtype=np.int64), dtype="int64")
|
|
25
|
+
structure = Data(num_nodes=5, edge_index=edge_index, y=y)
|
|
26
|
+
super().__init__(structure)
|
|
27
|
+
|
|
@@ -0,0 +1,55 @@
|
|
|
1
|
+
import os
|
|
2
|
+
import os.path as osp
|
|
3
|
+
from typing import Callable, Optional
|
|
4
|
+
|
|
5
|
+
import numpy as np
|
|
6
|
+
|
|
7
|
+
from k3_node.data import Data, InMemoryDataset
|
|
8
|
+
from k3_node.data.download import download_url
|
|
9
|
+
from k3_node.data.extract import extract_zip
|
|
10
|
+
|
|
11
|
+
|
|
12
|
+
class MovieLens100K(InMemoryDataset):
|
|
13
|
+
r"""The MovieLens 100K ratings (943 users, 1,682 movies, 100,000 ratings) as a user-movie graph
|
|
14
|
+
for recommendation, a small stand-in for PyG's ``AmazonBook``.
|
|
15
|
+
|
|
16
|
+
Every rating counts as an interaction. The graph is homogeneous: users are nodes
|
|
17
|
+
``0 .. num_users - 1`` and movies ``num_users .. num_nodes - 1``. ``edge_index`` holds the
|
|
18
|
+
training interactions in both directions, ``edge_label_index`` the test interactions
|
|
19
|
+
(user, movie), from the official 80/20 split ``u1.base`` / ``u1.test``.
|
|
20
|
+
|
|
21
|
+
Args:
|
|
22
|
+
root (str): Root directory where the dataset should be saved.
|
|
23
|
+
transform (callable, optional): A function applied to the graph when it is accessed.
|
|
24
|
+
force_reload (bool, optional): Whether to re-process the dataset. (default: ``False``)
|
|
25
|
+
"""
|
|
26
|
+
|
|
27
|
+
url = "https://files.grouplens.org/datasets/movielens/ml-100k.zip"
|
|
28
|
+
num_users, num_items = 943, 1682
|
|
29
|
+
|
|
30
|
+
def __init__(self, root: str, transform: Optional[Callable] = None, force_reload: bool = False):
|
|
31
|
+
super().__init__(root, transform, force_reload=force_reload)
|
|
32
|
+
self.load(self.processed_paths[0])
|
|
33
|
+
|
|
34
|
+
@property
|
|
35
|
+
def raw_file_names(self):
|
|
36
|
+
return [osp.join("ml-100k", "u1.base"), osp.join("ml-100k", "u1.test")]
|
|
37
|
+
|
|
38
|
+
@property
|
|
39
|
+
def processed_file_names(self) -> str:
|
|
40
|
+
return "data.pt"
|
|
41
|
+
|
|
42
|
+
def download(self):
|
|
43
|
+
path = download_url(self.url, self.raw_dir)
|
|
44
|
+
extract_zip(path, self.raw_dir)
|
|
45
|
+
os.unlink(path)
|
|
46
|
+
|
|
47
|
+
def _read(self, path):
|
|
48
|
+
ratings = np.loadtxt(path, dtype=np.int64)[:, :2] - 1 # user id, movie id (1-based)
|
|
49
|
+
return np.stack([ratings[:, 0], ratings[:, 1] + self.num_users])
|
|
50
|
+
|
|
51
|
+
def process(self):
|
|
52
|
+
train, test = self._read(self.raw_paths[0]), self._read(self.raw_paths[1])
|
|
53
|
+
data = Data(edge_index=np.concatenate([train, train[::-1]], axis=1), edge_label_index=test,
|
|
54
|
+
num_nodes=self.num_users + self.num_items)
|
|
55
|
+
self.save([data], self.processed_paths[0])
|