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,137 @@
|
|
|
1
|
+
import os.path as osp
|
|
2
|
+
from typing import Callable, List, Optional
|
|
3
|
+
import numpy as np
|
|
4
|
+
from keras import ops
|
|
5
|
+
|
|
6
|
+
from k3_node.data import InMemoryDataset
|
|
7
|
+
from k3_node.io import fs, read_planetoid_data
|
|
8
|
+
|
|
9
|
+
|
|
10
|
+
class Planetoid(InMemoryDataset):
|
|
11
|
+
r"""The citation network datasets "Cora", "CiteSeer" and "PubMed" from the
|
|
12
|
+
"Revisiting Semi-Supervised Learning with Graph Embeddings" paper.
|
|
13
|
+
Nodes represent documents and edges represent citation links.
|
|
14
|
+
Training, validation and test splits are given by binary masks.
|
|
15
|
+
|
|
16
|
+
Args:
|
|
17
|
+
root (str): Root directory where the dataset should be saved.
|
|
18
|
+
name (str): The name of the dataset ("Cora", "CiteSeer", "PubMed").
|
|
19
|
+
split (str, optional): The type of dataset split ("public", "full", "geom-gcn", "random").
|
|
20
|
+
(default: "public")
|
|
21
|
+
num_train_per_class (int, optional): The number of training samples per class for "random" split.
|
|
22
|
+
(default: 20)
|
|
23
|
+
num_val (int, optional): The number of validation samples for "random" split. (default: 500)
|
|
24
|
+
num_test (int, optional): The number of test samples for "random" split. (default: 1000)
|
|
25
|
+
transform (callable, optional): A function/transform that takes in a Data object and returns a transformed version.
|
|
26
|
+
pre_transform (callable, optional): A function/transform that takes in a Data object and returns a transformed version.
|
|
27
|
+
force_reload (bool, optional): Whether to re-process the dataset. (default: False)
|
|
28
|
+
"""
|
|
29
|
+
|
|
30
|
+
url = "https://github.com/kimiyoung/planetoid/raw/master/data"
|
|
31
|
+
geom_gcn_url = "https://raw.githubusercontent.com/graphdml-uiuc-jlu/geom-gcn/master"
|
|
32
|
+
|
|
33
|
+
def __init__(
|
|
34
|
+
self,
|
|
35
|
+
root: str,
|
|
36
|
+
name: str,
|
|
37
|
+
split: str = "public",
|
|
38
|
+
num_train_per_class: int = 20,
|
|
39
|
+
num_val: int = 500,
|
|
40
|
+
num_test: int = 1000,
|
|
41
|
+
transform: Optional[Callable] = None,
|
|
42
|
+
pre_transform: Optional[Callable] = None,
|
|
43
|
+
force_reload: bool = False,
|
|
44
|
+
):
|
|
45
|
+
self.name = name
|
|
46
|
+
self.split = split.lower()
|
|
47
|
+
assert self.split in ["public", "full", "geom-gcn", "random"]
|
|
48
|
+
|
|
49
|
+
super().__init__(root, transform, pre_transform, force_reload=force_reload)
|
|
50
|
+
self.load(self.processed_paths[0])
|
|
51
|
+
|
|
52
|
+
if self.split == "full":
|
|
53
|
+
data = self.get(0)
|
|
54
|
+
val_m = ops.convert_to_numpy(data.val_mask)
|
|
55
|
+
test_m = ops.convert_to_numpy(data.test_mask)
|
|
56
|
+
train_mask = np.ones(data.num_nodes, dtype=bool)
|
|
57
|
+
train_mask[val_m | test_m] = False
|
|
58
|
+
data.train_mask = ops.convert_to_tensor(train_mask, dtype="bool")
|
|
59
|
+
self.data, self.slices = self.collate([data])
|
|
60
|
+
|
|
61
|
+
elif self.split == "random":
|
|
62
|
+
data = self.get(0)
|
|
63
|
+
num_nodes = data.num_nodes
|
|
64
|
+
y_np = ops.convert_to_numpy(data.y)
|
|
65
|
+
train_mask = np.zeros(num_nodes, dtype=bool)
|
|
66
|
+
for c in range(self.num_classes):
|
|
67
|
+
idx = np.where(y_np == c)[0]
|
|
68
|
+
perm = np.random.permutation(len(idx))
|
|
69
|
+
train_mask[idx[perm[:num_train_per_class]]] = True
|
|
70
|
+
|
|
71
|
+
remaining = np.where(~train_mask)[0]
|
|
72
|
+
remaining = remaining[np.random.permutation(len(remaining))]
|
|
73
|
+
|
|
74
|
+
val_mask = np.zeros(num_nodes, dtype=bool)
|
|
75
|
+
val_mask[remaining[:num_val]] = True
|
|
76
|
+
|
|
77
|
+
test_mask = np.zeros(num_nodes, dtype=bool)
|
|
78
|
+
test_mask[remaining[num_val : num_val + num_test]] = True
|
|
79
|
+
|
|
80
|
+
data.train_mask = ops.convert_to_tensor(train_mask, dtype="bool")
|
|
81
|
+
data.val_mask = ops.convert_to_tensor(val_mask, dtype="bool")
|
|
82
|
+
data.test_mask = ops.convert_to_tensor(test_mask, dtype="bool")
|
|
83
|
+
|
|
84
|
+
self.data, self.slices = self.collate([data])
|
|
85
|
+
|
|
86
|
+
@property
|
|
87
|
+
def raw_dir(self) -> str:
|
|
88
|
+
if self.split == "geom-gcn":
|
|
89
|
+
return osp.join(self.root, self.name, "geom-gcn", "raw")
|
|
90
|
+
return osp.join(self.root, self.name, "raw")
|
|
91
|
+
|
|
92
|
+
@property
|
|
93
|
+
def processed_dir(self) -> str:
|
|
94
|
+
if self.split == "geom-gcn":
|
|
95
|
+
return osp.join(self.root, self.name, "geom-gcn", "processed")
|
|
96
|
+
return osp.join(self.root, self.name, "processed")
|
|
97
|
+
|
|
98
|
+
@property
|
|
99
|
+
def raw_file_names(self) -> List[str]:
|
|
100
|
+
names = ["x", "tx", "allx", "y", "ty", "ally", "graph", "test.index"]
|
|
101
|
+
return [f"ind.{self.name.lower()}.{name}" for name in names]
|
|
102
|
+
|
|
103
|
+
@property
|
|
104
|
+
def processed_file_names(self) -> str:
|
|
105
|
+
return "data.pt"
|
|
106
|
+
|
|
107
|
+
def download(self):
|
|
108
|
+
for name in self.raw_file_names:
|
|
109
|
+
fs.cp(f"{self.url}/{name}", self.raw_dir)
|
|
110
|
+
if self.split == "geom-gcn":
|
|
111
|
+
for i in range(10):
|
|
112
|
+
url = f"{self.geom_gcn_url}/splits/{self.name.lower()}"
|
|
113
|
+
fs.cp(f"{url}_split_0.6_0.2_{i}.npz", self.raw_dir)
|
|
114
|
+
|
|
115
|
+
def process(self):
|
|
116
|
+
data = read_planetoid_data(self.raw_dir, self.name)
|
|
117
|
+
|
|
118
|
+
if self.split == "geom-gcn":
|
|
119
|
+
train_masks, val_masks, test_masks = [], [], []
|
|
120
|
+
for i in range(10):
|
|
121
|
+
name = f"{self.name.lower()}_split_0.6_0.2_{i}.npz"
|
|
122
|
+
splits = np.load(osp.join(self.raw_dir, name))
|
|
123
|
+
train_masks.append(splits["train_mask"])
|
|
124
|
+
val_masks.append(splits["val_mask"])
|
|
125
|
+
test_masks.append(splits["test_mask"])
|
|
126
|
+
data.train_mask = ops.convert_to_tensor(np.stack(train_masks, axis=1), dtype="bool")
|
|
127
|
+
data.val_mask = ops.convert_to_tensor(np.stack(val_masks, axis=1), dtype="bool")
|
|
128
|
+
data.test_mask = ops.convert_to_tensor(np.stack(test_masks, axis=1), dtype="bool")
|
|
129
|
+
|
|
130
|
+
if self.pre_transform is not None:
|
|
131
|
+
data = self.pre_transform(data)
|
|
132
|
+
|
|
133
|
+
self.save([data], self.processed_paths[0])
|
|
134
|
+
|
|
135
|
+
def __repr__(self) -> str:
|
|
136
|
+
return f"{self.name}()"
|
|
137
|
+
|
|
@@ -0,0 +1,63 @@
|
|
|
1
|
+
import os
|
|
2
|
+
import os.path as osp
|
|
3
|
+
from typing import Callable, List, Optional
|
|
4
|
+
import numpy as np
|
|
5
|
+
from keras import ops
|
|
6
|
+
|
|
7
|
+
from k3_node.data import Data, InMemoryDataset
|
|
8
|
+
from k3_node.io import fs
|
|
9
|
+
|
|
10
|
+
|
|
11
|
+
class PolBlogs(InMemoryDataset):
|
|
12
|
+
r"""The Political Blogs dataset containing 1,490 vertices and 19,025 edges.
|
|
13
|
+
|
|
14
|
+
Args:
|
|
15
|
+
root (str): Root directory where the dataset should be saved.
|
|
16
|
+
transform (callable, optional): Transform function.
|
|
17
|
+
pre_transform (callable, optional): Pre-transform function.
|
|
18
|
+
force_reload (bool, optional): Whether to re-process the dataset.
|
|
19
|
+
"""
|
|
20
|
+
|
|
21
|
+
url = "https://netset.telecom-paris.fr/datasets/polblogs.tar.gz"
|
|
22
|
+
|
|
23
|
+
def __init__(
|
|
24
|
+
self,
|
|
25
|
+
root: str,
|
|
26
|
+
transform: Optional[Callable] = None,
|
|
27
|
+
pre_transform: Optional[Callable] = None,
|
|
28
|
+
force_reload: bool = False,
|
|
29
|
+
):
|
|
30
|
+
super().__init__(root, transform, pre_transform, force_reload=force_reload)
|
|
31
|
+
self.load(self.processed_paths[0])
|
|
32
|
+
|
|
33
|
+
@property
|
|
34
|
+
def raw_file_names(self) -> List[str]:
|
|
35
|
+
return ["adjacency.tsv", "labels.tsv"]
|
|
36
|
+
|
|
37
|
+
@property
|
|
38
|
+
def processed_file_names(self) -> str:
|
|
39
|
+
return "data.pt"
|
|
40
|
+
|
|
41
|
+
def download(self):
|
|
42
|
+
tar_path = osp.join(self.raw_dir, "polblogs.tar.gz")
|
|
43
|
+
fs.cp(self.url, tar_path, extract=True)
|
|
44
|
+
if osp.exists(tar_path):
|
|
45
|
+
fs.rm(tar_path)
|
|
46
|
+
|
|
47
|
+
def process(self):
|
|
48
|
+
adj = np.loadtxt(self.raw_paths[0], delimiter="\t", usecols=(0, 1), dtype=np.int64)
|
|
49
|
+
edge_index = adj.T
|
|
50
|
+
|
|
51
|
+
y = np.loadtxt(self.raw_paths[1], delimiter="\t", usecols=(1,), dtype=np.int64)
|
|
52
|
+
|
|
53
|
+
data = Data(
|
|
54
|
+
edge_index=ops.convert_to_tensor(edge_index, dtype="int64"),
|
|
55
|
+
y=ops.convert_to_tensor(y, dtype="int64"),
|
|
56
|
+
num_nodes=int(y.shape[0]),
|
|
57
|
+
)
|
|
58
|
+
|
|
59
|
+
if self.pre_transform is not None:
|
|
60
|
+
data = self.pre_transform(data)
|
|
61
|
+
|
|
62
|
+
self.save([data], self.processed_paths[0])
|
|
63
|
+
|
k3_node/datasets/ppi.py
ADDED
|
@@ -0,0 +1,189 @@
|
|
|
1
|
+
import json
|
|
2
|
+
import os
|
|
3
|
+
import os.path as osp
|
|
4
|
+
from itertools import product
|
|
5
|
+
from typing import Callable, List, Optional
|
|
6
|
+
|
|
7
|
+
import numpy as np
|
|
8
|
+
from keras import ops
|
|
9
|
+
|
|
10
|
+
from k3_node.data import (
|
|
11
|
+
Data,
|
|
12
|
+
InMemoryDataset,
|
|
13
|
+
download_url,
|
|
14
|
+
extract_zip,
|
|
15
|
+
)
|
|
16
|
+
from k3_node.layers.conv.utils import remove_self_loops
|
|
17
|
+
|
|
18
|
+
|
|
19
|
+
class PPI(InMemoryDataset):
|
|
20
|
+
r"""The protein-protein interaction networks from the `"Predicting
|
|
21
|
+
Multicellular Function through Multi-layer Tissue Networks"
|
|
22
|
+
<https://arxiv.org/abs/1707.04638>`_ paper, containing positional gene
|
|
23
|
+
sets, motif gene sets and immunological signatures as features (50 in
|
|
24
|
+
total) and gene ontology sets as labels (121 in total).
|
|
25
|
+
|
|
26
|
+
Args:
|
|
27
|
+
root (str): Root directory where the dataset should be saved.
|
|
28
|
+
split (str, optional): If :obj:`"train"`, loads the training dataset.
|
|
29
|
+
If :obj:`"val"`, loads the validation dataset.
|
|
30
|
+
If :obj:`"test"`, loads the test dataset. (default: :obj:`"train"`)
|
|
31
|
+
transform (callable, optional): A function/transform that takes in an
|
|
32
|
+
:obj:`k3_node.data.Data` object and returns a transformed
|
|
33
|
+
version. The data object will be transformed before every access.
|
|
34
|
+
(default: :obj:`None`)
|
|
35
|
+
pre_transform (callable, optional): A function/transform that takes in
|
|
36
|
+
an :obj:`k3_node.data.Data` object and returns a
|
|
37
|
+
transformed version. The data object will be transformed before
|
|
38
|
+
being saved to disk. (default: :obj:`None`)
|
|
39
|
+
pre_filter (callable, optional): A function that takes in an
|
|
40
|
+
:obj:`k3_node.data.Data` object and returns a boolean
|
|
41
|
+
value, indicating whether the data object should be included in the
|
|
42
|
+
final dataset. (default: :obj:`None`)
|
|
43
|
+
force_reload (bool, optional): Whether to re-process the dataset.
|
|
44
|
+
(default: :obj:`False`)
|
|
45
|
+
|
|
46
|
+
**STATS:**
|
|
47
|
+
|
|
48
|
+
.. list-table::
|
|
49
|
+
:widths: 10 10 10 10 10
|
|
50
|
+
:header-rows: 1
|
|
51
|
+
|
|
52
|
+
* - #graphs
|
|
53
|
+
- #nodes
|
|
54
|
+
- #edges
|
|
55
|
+
- #features
|
|
56
|
+
- #tasks
|
|
57
|
+
* - 20
|
|
58
|
+
- ~2,245.3
|
|
59
|
+
- ~61,318.4
|
|
60
|
+
- 50
|
|
61
|
+
- 121
|
|
62
|
+
"""
|
|
63
|
+
|
|
64
|
+
url = "https://data.dgl.ai/dataset/ppi.zip"
|
|
65
|
+
|
|
66
|
+
def __init__(
|
|
67
|
+
self,
|
|
68
|
+
root: str,
|
|
69
|
+
split: str = "train",
|
|
70
|
+
transform: Optional[Callable] = None,
|
|
71
|
+
pre_transform: Optional[Callable] = None,
|
|
72
|
+
pre_filter: Optional[Callable] = None,
|
|
73
|
+
force_reload: bool = False,
|
|
74
|
+
) -> None:
|
|
75
|
+
assert split.lower() in ["train", "val", "valid", "test"], f"Invalid split '{split}'"
|
|
76
|
+
self.split = "val" if split.lower() == "valid" else split.lower()
|
|
77
|
+
|
|
78
|
+
super().__init__(
|
|
79
|
+
root,
|
|
80
|
+
transform,
|
|
81
|
+
pre_transform,
|
|
82
|
+
pre_filter,
|
|
83
|
+
force_reload=force_reload,
|
|
84
|
+
)
|
|
85
|
+
|
|
86
|
+
if self.split == "train":
|
|
87
|
+
self.load(self.processed_paths[0])
|
|
88
|
+
elif self.split == "val":
|
|
89
|
+
self.load(self.processed_paths[1])
|
|
90
|
+
elif self.split == "test":
|
|
91
|
+
self.load(self.processed_paths[2])
|
|
92
|
+
|
|
93
|
+
@property
|
|
94
|
+
def raw_file_names(self) -> List[str]:
|
|
95
|
+
splits = ["train", "valid", "test"]
|
|
96
|
+
files = ["feats.npy", "graph_id.npy", "graph.json", "labels.npy"]
|
|
97
|
+
return [f"{split}_{name}" for split, name in product(splits, files)]
|
|
98
|
+
|
|
99
|
+
@property
|
|
100
|
+
def processed_file_names(self) -> List[str]:
|
|
101
|
+
return ["train.pt", "val.pt", "test.pt"]
|
|
102
|
+
|
|
103
|
+
def download(self) -> None:
|
|
104
|
+
path = download_url(self.url, self.root)
|
|
105
|
+
extract_zip(path, self.raw_dir)
|
|
106
|
+
if osp.exists(path):
|
|
107
|
+
os.unlink(path)
|
|
108
|
+
|
|
109
|
+
def process(self) -> None:
|
|
110
|
+
try:
|
|
111
|
+
import networkx as nx
|
|
112
|
+
from networkx.readwrite import json_graph
|
|
113
|
+
has_networkx = True
|
|
114
|
+
except ImportError:
|
|
115
|
+
has_networkx = False
|
|
116
|
+
|
|
117
|
+
for s, split in enumerate(["train", "valid", "test"]):
|
|
118
|
+
path = osp.join(self.raw_dir, f"{split}_graph.json")
|
|
119
|
+
with open(path, "r", encoding="utf-8") as f:
|
|
120
|
+
graph_json = json.load(f)
|
|
121
|
+
|
|
122
|
+
if has_networkx:
|
|
123
|
+
try:
|
|
124
|
+
G = nx.DiGraph(json_graph.node_link_graph(graph_json, edges="links"))
|
|
125
|
+
except TypeError:
|
|
126
|
+
G = nx.DiGraph(json_graph.node_link_graph(graph_json))
|
|
127
|
+
else:
|
|
128
|
+
links = graph_json.get("links", [])
|
|
129
|
+
src_all = np.array([link["source"] for link in links], dtype=np.int64)
|
|
130
|
+
dst_all = np.array([link["target"] for link in links], dtype=np.int64)
|
|
131
|
+
|
|
132
|
+
x_np = np.load(osp.join(self.raw_dir, f"{split}_feats.npy"))
|
|
133
|
+
y_np = np.load(osp.join(self.raw_dir, f"{split}_labels.npy"))
|
|
134
|
+
|
|
135
|
+
data_list = []
|
|
136
|
+
path = osp.join(self.raw_dir, f"{split}_graph_id.npy")
|
|
137
|
+
idx = np.load(path)
|
|
138
|
+
idx = idx - idx.min()
|
|
139
|
+
|
|
140
|
+
for i in range(int(idx.max()) + 1):
|
|
141
|
+
mask = (idx == i)
|
|
142
|
+
node_indices = np.where(mask)[0]
|
|
143
|
+
|
|
144
|
+
if has_networkx:
|
|
145
|
+
G_s = G.subgraph(node_indices.tolist())
|
|
146
|
+
edges = list(G_s.edges)
|
|
147
|
+
if len(edges) > 0:
|
|
148
|
+
edge_index = np.array(edges, dtype=np.int64).T
|
|
149
|
+
edge_index = edge_index - edge_index.min()
|
|
150
|
+
else:
|
|
151
|
+
edge_index = np.zeros((2, 0), dtype=np.int64)
|
|
152
|
+
else:
|
|
153
|
+
min_node, max_node = node_indices.min(), node_indices.max()
|
|
154
|
+
edge_mask = (
|
|
155
|
+
(src_all >= min_node)
|
|
156
|
+
& (src_all <= max_node)
|
|
157
|
+
& (dst_all >= min_node)
|
|
158
|
+
& (dst_all <= max_node)
|
|
159
|
+
)
|
|
160
|
+
sub_src = src_all[edge_mask] - min_node
|
|
161
|
+
sub_dst = dst_all[edge_mask] - min_node
|
|
162
|
+
edge_index = np.stack([sub_src, sub_dst], axis=0)
|
|
163
|
+
|
|
164
|
+
edge_index = ops.convert_to_tensor(edge_index, dtype="int64")
|
|
165
|
+
edge_index, _ = remove_self_loops(edge_index)
|
|
166
|
+
|
|
167
|
+
x = ops.convert_to_tensor(x_np[mask], dtype="float32")
|
|
168
|
+
y = ops.convert_to_tensor(y_np[mask], dtype="float32")
|
|
169
|
+
|
|
170
|
+
data = Data(edge_index=edge_index, x=x, y=y)
|
|
171
|
+
|
|
172
|
+
if self.pre_filter is not None and not self.pre_filter(data):
|
|
173
|
+
continue
|
|
174
|
+
|
|
175
|
+
if self.pre_transform is not None:
|
|
176
|
+
data = self.pre_transform(data)
|
|
177
|
+
|
|
178
|
+
data_list.append(data)
|
|
179
|
+
|
|
180
|
+
self.save(data_list, self.processed_paths[s])
|
|
181
|
+
|
|
182
|
+
@property
|
|
183
|
+
def num_classes(self) -> int:
|
|
184
|
+
data = self.get(0)
|
|
185
|
+
y = getattr(data, "y", None)
|
|
186
|
+
if y is not None and len(y.shape) > 1:
|
|
187
|
+
return y.shape[-1]
|
|
188
|
+
return super().num_classes
|
|
189
|
+
|
k3_node/datasets/qm7.py
ADDED
|
@@ -0,0 +1,65 @@
|
|
|
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 QM7b(InMemoryDataset):
|
|
10
|
+
r"""The QM7b dataset consisting of 7,211 molecules with 14 regression targets."""
|
|
11
|
+
|
|
12
|
+
url = "https://deepchemdata.s3-us-west-1.amazonaws.com/datasets/qm7b.mat"
|
|
13
|
+
|
|
14
|
+
def __init__(
|
|
15
|
+
self,
|
|
16
|
+
root: str,
|
|
17
|
+
transform: Optional[Callable] = None,
|
|
18
|
+
pre_transform: Optional[Callable] = None,
|
|
19
|
+
pre_filter: Optional[Callable] = None,
|
|
20
|
+
force_reload: bool = False,
|
|
21
|
+
):
|
|
22
|
+
super().__init__(root, transform, pre_transform, pre_filter, force_reload=force_reload)
|
|
23
|
+
self.load(self.processed_paths[0])
|
|
24
|
+
|
|
25
|
+
@property
|
|
26
|
+
def raw_file_names(self) -> str:
|
|
27
|
+
return "qm7b.mat"
|
|
28
|
+
|
|
29
|
+
@property
|
|
30
|
+
def processed_file_names(self) -> str:
|
|
31
|
+
return "data.pt"
|
|
32
|
+
|
|
33
|
+
def download(self):
|
|
34
|
+
fs.cp(self.url, self.raw_dir)
|
|
35
|
+
|
|
36
|
+
def process(self):
|
|
37
|
+
from scipy.io import loadmat
|
|
38
|
+
|
|
39
|
+
data = loadmat(self.raw_paths[0])
|
|
40
|
+
coulomb_matrix = data["X"]
|
|
41
|
+
target = data["T"].astype(np.float32)
|
|
42
|
+
|
|
43
|
+
data_list = []
|
|
44
|
+
for i in range(target.shape[0]):
|
|
45
|
+
nz = np.nonzero(coulomb_matrix[i])
|
|
46
|
+
edge_index = np.stack([nz[0], nz[1]], axis=0).astype(np.int64)
|
|
47
|
+
edge_attr = coulomb_matrix[i, edge_index[0], edge_index[1]].astype(np.float32)
|
|
48
|
+
y = target[i].reshape(1, -1)
|
|
49
|
+
num_nodes = int(np.max(edge_index)) + 1 if edge_index.size > 0 else 0
|
|
50
|
+
|
|
51
|
+
d = Data(
|
|
52
|
+
edge_index=ops.convert_to_tensor(edge_index, dtype="int64"),
|
|
53
|
+
edge_attr=ops.convert_to_tensor(edge_attr, dtype="float32"),
|
|
54
|
+
y=ops.convert_to_tensor(y, dtype="float32"),
|
|
55
|
+
num_nodes=num_nodes,
|
|
56
|
+
)
|
|
57
|
+
data_list.append(d)
|
|
58
|
+
|
|
59
|
+
if self.pre_filter is not None:
|
|
60
|
+
data_list = [d for d in data_list if self.pre_filter(d)]
|
|
61
|
+
if self.pre_transform is not None:
|
|
62
|
+
data_list = [self.pre_transform(d) for d in data_list]
|
|
63
|
+
|
|
64
|
+
self.save(data_list, self.processed_paths[0])
|
|
65
|
+
|
k3_node/datasets/qm9.py
ADDED
|
@@ -0,0 +1,132 @@
|
|
|
1
|
+
import os
|
|
2
|
+
import os.path as osp
|
|
3
|
+
from typing import Callable, List, 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
|
+
HAR2EV = 27.211386246
|
|
12
|
+
KCALMOL2EV = 0.04336414
|
|
13
|
+
CONVERSION = np.array([1., 1., HAR2EV, HAR2EV, HAR2EV, 1., HAR2EV, HAR2EV, HAR2EV, HAR2EV, HAR2EV,
|
|
14
|
+
1., KCALMOL2EV, KCALMOL2EV, KCALMOL2EV, KCALMOL2EV, 1., 1., 1.], dtype=np.float32)
|
|
15
|
+
ATOMREFS = {
|
|
16
|
+
6: [0., 0., 0., 0., 0.],
|
|
17
|
+
7: [-13.61312172, -1029.86312267, -1485.30251237, -2042.61123593, -2713.48485589],
|
|
18
|
+
8: [-13.5745904, -1029.82456413, -1485.26398105, -2042.5727046, -2713.44632457],
|
|
19
|
+
9: [-13.54887564, -1029.79887659, -1485.2382935, -2042.54701705, -2713.42063702],
|
|
20
|
+
10: [-13.90303183, -1030.25891228, -1485.71166277, -2043.01812778, -2713.88796536],
|
|
21
|
+
11: [0., 0., 0., 0., 0.],
|
|
22
|
+
}
|
|
23
|
+
|
|
24
|
+
|
|
25
|
+
class QM9(InMemoryDataset):
|
|
26
|
+
r"""The QM9 dataset: about 130,000 small organic molecules with their 3D structure and 19
|
|
27
|
+
regression targets (dipole moment, HOMO/LUMO energies, internal energy, ...), as in PyG.
|
|
28
|
+
|
|
29
|
+
Every molecule has 11 atom features ``x`` (one-hot H/C/N/O/F, atomic number, aromatic,
|
|
30
|
+
sp/sp2/sp3 hybridization, number of hydrogens), atomic numbers ``z``, positions ``pos``,
|
|
31
|
+
one-hot bond types ``edge_attr`` and the targets ``y`` of shape ``[1, 19]`` (energies in eV).
|
|
32
|
+
Requires RDKit to process the raw files.
|
|
33
|
+
|
|
34
|
+
Args:
|
|
35
|
+
root (str): Root directory where the dataset should be saved.
|
|
36
|
+
transform (callable, optional): A function applied to each graph when it is accessed.
|
|
37
|
+
pre_transform (callable, optional): A function applied to each graph before saving.
|
|
38
|
+
pre_filter (callable, optional): A function deciding which graphs to keep.
|
|
39
|
+
force_reload (bool, optional): Whether to re-process the dataset. (default: ``False``)
|
|
40
|
+
"""
|
|
41
|
+
|
|
42
|
+
raw_url = 'https://deepchemdata.s3-us-west-1.amazonaws.com/datasets/molnet_publish/qm9.zip'
|
|
43
|
+
raw_url2 = 'https://ndownloader.figshare.com/files/3195404'
|
|
44
|
+
|
|
45
|
+
def __init__(self, root: str, transform: Optional[Callable] = None, pre_transform: Optional[Callable] = None,
|
|
46
|
+
pre_filter: Optional[Callable] = None, force_reload: bool = False):
|
|
47
|
+
super().__init__(root, transform, pre_transform, pre_filter, force_reload=force_reload)
|
|
48
|
+
self.load(self.processed_paths[0])
|
|
49
|
+
|
|
50
|
+
def mean(self, target: int) -> float:
|
|
51
|
+
return float(np.asarray(self._data.y)[:, target].mean())
|
|
52
|
+
|
|
53
|
+
def std(self, target: int) -> float:
|
|
54
|
+
return float(np.asarray(self._data.y)[:, target].std(ddof=1))
|
|
55
|
+
|
|
56
|
+
def atomref(self, target: int) -> Optional[np.ndarray]:
|
|
57
|
+
r"""Per-element reference energies (a ``[100, 1]`` array indexed by atomic number), or
|
|
58
|
+
``None`` for targets without them."""
|
|
59
|
+
if target not in ATOMREFS:
|
|
60
|
+
return None
|
|
61
|
+
out = np.zeros((100, 1), dtype=np.float32)
|
|
62
|
+
out[[1, 6, 7, 8, 9], 0] = ATOMREFS[target]
|
|
63
|
+
return out
|
|
64
|
+
|
|
65
|
+
@property
|
|
66
|
+
def raw_file_names(self) -> List[str]:
|
|
67
|
+
return ['gdb9.sdf', 'gdb9.sdf.csv', 'uncharacterized.txt']
|
|
68
|
+
|
|
69
|
+
@property
|
|
70
|
+
def processed_file_names(self) -> str:
|
|
71
|
+
return 'data_v3.pt'
|
|
72
|
+
|
|
73
|
+
def download(self):
|
|
74
|
+
path = download_url(self.raw_url, self.raw_dir)
|
|
75
|
+
extract_zip(path, self.raw_dir)
|
|
76
|
+
os.unlink(path)
|
|
77
|
+
download_url(self.raw_url2, self.raw_dir)
|
|
78
|
+
os.rename(osp.join(self.raw_dir, '3195404'), osp.join(self.raw_dir, 'uncharacterized.txt'))
|
|
79
|
+
|
|
80
|
+
def process(self):
|
|
81
|
+
from rdkit import Chem, RDLogger
|
|
82
|
+
from rdkit.Chem.rdchem import BondType as BT
|
|
83
|
+
from rdkit.Chem.rdchem import HybridizationType as HT
|
|
84
|
+
|
|
85
|
+
RDLogger.DisableLog('rdApp.*')
|
|
86
|
+
types = {'H': 0, 'C': 1, 'N': 2, 'O': 3, 'F': 4}
|
|
87
|
+
bonds = {BT.SINGLE: 0, BT.DOUBLE: 1, BT.TRIPLE: 2, BT.AROMATIC: 3}
|
|
88
|
+
|
|
89
|
+
with open(self.raw_paths[1]) as f:
|
|
90
|
+
target = np.array([[float(x) for x in line.split(',')[1:20]] for line in f.read().split('\n')[1:-1]],
|
|
91
|
+
dtype=np.float32)
|
|
92
|
+
target = np.concatenate([target[:, 3:], target[:, :3]], axis=-1) * CONVERSION
|
|
93
|
+
with open(self.raw_paths[2]) as f:
|
|
94
|
+
skip = {int(x.split()[0]) - 1 for x in f.read().split('\n')[9:-2]}
|
|
95
|
+
|
|
96
|
+
data_list = []
|
|
97
|
+
for i, mol in enumerate(Chem.SDMolSupplier(self.raw_paths[0], removeHs=False, sanitize=False)):
|
|
98
|
+
if i in skip:
|
|
99
|
+
continue
|
|
100
|
+
N = mol.GetNumAtoms()
|
|
101
|
+
pos = mol.GetConformer().GetPositions().astype(np.float32)
|
|
102
|
+
atoms = list(mol.GetAtoms())
|
|
103
|
+
z = np.array([a.GetAtomicNum() for a in atoms], dtype=np.int64)
|
|
104
|
+
rows, cols, edge_types = [], [], []
|
|
105
|
+
for bond in mol.GetBonds():
|
|
106
|
+
s, e = bond.GetBeginAtomIdx(), bond.GetEndAtomIdx()
|
|
107
|
+
rows += [s, e]
|
|
108
|
+
cols += [e, s]
|
|
109
|
+
edge_types += 2 * [bonds[bond.GetBondType()]]
|
|
110
|
+
edge_index = np.array([rows, cols], dtype=np.int64).reshape(2, -1)
|
|
111
|
+
edge_type = np.array(edge_types, dtype=np.int64)
|
|
112
|
+
perm = np.argsort(edge_index[0] * N + edge_index[1], kind="stable")
|
|
113
|
+
edge_index, edge_type = edge_index[:, perm], edge_type[perm]
|
|
114
|
+
|
|
115
|
+
num_hs = np.zeros(N, dtype=np.float32)
|
|
116
|
+
np.add.at(num_hs, edge_index[1], (z == 1).astype(np.float32)[edge_index[0]])
|
|
117
|
+
hybrid = [a.GetHybridization() for a in atoms]
|
|
118
|
+
x = np.concatenate([
|
|
119
|
+
np.eye(len(types), dtype=np.float32)[[types[a.GetSymbol()] for a in atoms]],
|
|
120
|
+
np.stack([z, [a.GetIsAromatic() for a in atoms], [h == HT.SP for h in hybrid],
|
|
121
|
+
[h == HT.SP2 for h in hybrid], [h == HT.SP3 for h in hybrid], num_hs], axis=1).astype(np.float32),
|
|
122
|
+
], axis=1)
|
|
123
|
+
data = Data(x=x, z=z, pos=pos, edge_index=edge_index,
|
|
124
|
+
edge_attr=np.eye(len(bonds), dtype=np.float32)[edge_type].reshape(-1, len(bonds)),
|
|
125
|
+
y=target[i][None], smiles=Chem.MolToSmiles(mol, isomericSmiles=True),
|
|
126
|
+
name=mol.GetProp('_Name'), idx=i)
|
|
127
|
+
if self.pre_filter is not None and not self.pre_filter(data):
|
|
128
|
+
continue
|
|
129
|
+
if self.pre_transform is not None:
|
|
130
|
+
data = self.pre_transform(data)
|
|
131
|
+
data_list.append(data)
|
|
132
|
+
self.save(data_list, self.processed_paths[0])
|