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,424 @@
|
|
|
1
|
+
r"""k3-node ports of `torch_geometric.nn.models`."""
|
|
2
|
+
|
|
3
|
+
from .mlp import MLP
|
|
4
|
+
from .attract_repel import ARLinkPredictor
|
|
5
|
+
from .autoencoder import InnerProductDecoder, GAE, VGAE, ARGA, ARGVA
|
|
6
|
+
from .deep_graph_infomax import DeepGraphInfomax
|
|
7
|
+
from .deepgcn import DeepGCNLayer
|
|
8
|
+
from .attentive_fp import AttentiveFP
|
|
9
|
+
from .jumping_knowledge import JumpingKnowledge, HeteroJumpingKnowledge
|
|
10
|
+
from .mask_label import MaskLabel
|
|
11
|
+
from .meta import MetaLayer
|
|
12
|
+
from .pmlp import PMLP
|
|
13
|
+
from .polynormer import Polynormer
|
|
14
|
+
from .basic_gnn import BasicGNN, GCN, GraphSAGE, GIN, GAT, PNA, EdgeCNN
|
|
15
|
+
from .label_prop import LabelPropagation
|
|
16
|
+
from .correct_and_smooth import CorrectAndSmooth
|
|
17
|
+
from .lightgcn import LightGCN, BPRLoss
|
|
18
|
+
from .linkx import LINKX, SparseLinear
|
|
19
|
+
from .rect import RECT_L
|
|
20
|
+
from .signed_gcn import SignedGCN
|
|
21
|
+
from .neural_fingerprint import NeuralFingerprint
|
|
22
|
+
from .graph_unet import GraphUNet
|
|
23
|
+
from .rev_gnn import GroupAddRev
|
|
24
|
+
from .sgformer import SGFormer
|
|
25
|
+
from .node2vec import Node2Vec
|
|
26
|
+
from .metapath2vec import MetaPath2Vec
|
|
27
|
+
from .renet import RENet
|
|
28
|
+
from .tgn import (
|
|
29
|
+
TGNMemory,
|
|
30
|
+
IdentityMessage,
|
|
31
|
+
LastAggregator,
|
|
32
|
+
MeanAggregator,
|
|
33
|
+
TimeEncoder,
|
|
34
|
+
LastNeighborLoader,
|
|
35
|
+
)
|
|
36
|
+
from .schnet import (
|
|
37
|
+
SchNet,
|
|
38
|
+
CFConv,
|
|
39
|
+
InteractionBlock as SchNetInteractionBlock,
|
|
40
|
+
GaussianSmearing,
|
|
41
|
+
ShiftedSoftplus,
|
|
42
|
+
RadiusInteractionGraph,
|
|
43
|
+
)
|
|
44
|
+
from .dimenet import (
|
|
45
|
+
DimeNet,
|
|
46
|
+
DimeNetPlusPlus,
|
|
47
|
+
BesselBasisLayer,
|
|
48
|
+
SphericalBasisLayer,
|
|
49
|
+
triplets,
|
|
50
|
+
)
|
|
51
|
+
from .gnnff import (
|
|
52
|
+
GNNFF,
|
|
53
|
+
NodeBlock,
|
|
54
|
+
EdgeBlock,
|
|
55
|
+
GaussianFilter,
|
|
56
|
+
)
|
|
57
|
+
from .gpse import (
|
|
58
|
+
GPSE,
|
|
59
|
+
GPSENodeEncoder,
|
|
60
|
+
GeneralLayer,
|
|
61
|
+
GeneralMultiLayer,
|
|
62
|
+
GNNStackStage,
|
|
63
|
+
GNNInductiveHybridMultiHead,
|
|
64
|
+
)
|
|
65
|
+
from .visnet import ViSNet
|
|
66
|
+
from .lpformer import LPFormer, LPAttLayer
|
|
67
|
+
from .graphmae2 import (
|
|
68
|
+
GraphMAE2,
|
|
69
|
+
sce_loss,
|
|
70
|
+
load_graphmae2_weights,
|
|
71
|
+
download_graphmae2_checkpoint,
|
|
72
|
+
)
|
|
73
|
+
from .graphormer import (
|
|
74
|
+
Graphormer,
|
|
75
|
+
GraphNodeFeature,
|
|
76
|
+
GraphAttnBias,
|
|
77
|
+
GraphormerMultiheadAttention,
|
|
78
|
+
GraphormerGraphEncoderLayer,
|
|
79
|
+
GraphormerGraphEncoder,
|
|
80
|
+
load_graphormer_weights,
|
|
81
|
+
download_graphormer_checkpoint,
|
|
82
|
+
)
|
|
83
|
+
from .graphormer_3d import (
|
|
84
|
+
Graphormer3D,
|
|
85
|
+
GaussianLayer,
|
|
86
|
+
RBF,
|
|
87
|
+
Graphormer3DEncoderLayer,
|
|
88
|
+
NodeTaskHead,
|
|
89
|
+
load_graphormer3d_weights,
|
|
90
|
+
download_graphormer3d_checkpoint,
|
|
91
|
+
)
|
|
92
|
+
from .captum import to_captum_model, to_captum_input, captum_output_to_dicts
|
|
93
|
+
from .gps_model import (
|
|
94
|
+
GPSModel,
|
|
95
|
+
GPSLayer,
|
|
96
|
+
CustomGatedGCN,
|
|
97
|
+
AtomEncoder,
|
|
98
|
+
BondEncoder,
|
|
99
|
+
RWSEEncoder,
|
|
100
|
+
SANGraphHead,
|
|
101
|
+
load_gps_weights,
|
|
102
|
+
download_gps_checkpoint,
|
|
103
|
+
)
|
|
104
|
+
from .grover import (
|
|
105
|
+
GROVER,
|
|
106
|
+
GTransEncoder,
|
|
107
|
+
Readout,
|
|
108
|
+
load_grover_weights,
|
|
109
|
+
download_grover_checkpoint,
|
|
110
|
+
)
|
|
111
|
+
from .mole_bert import (
|
|
112
|
+
MoleBERT,
|
|
113
|
+
MoleBERTGNN,
|
|
114
|
+
MoleBERTGINConv,
|
|
115
|
+
load_mole_bert_weights,
|
|
116
|
+
download_mole_bert_checkpoint,
|
|
117
|
+
)
|
|
118
|
+
from .unimol import (
|
|
119
|
+
UniMolModel,
|
|
120
|
+
UniMolConfGenModel,
|
|
121
|
+
UniMolDockingModel,
|
|
122
|
+
GaussianLayer as UniMolGaussianLayer,
|
|
123
|
+
NumericalEmbed as UniMolNumericalEmbed,
|
|
124
|
+
NonLinearHead as UniMolNonLinearHead,
|
|
125
|
+
DistanceHead as UniMolDistanceHead,
|
|
126
|
+
ClassificationHead as UniMolClassificationHead,
|
|
127
|
+
LinearHead as UniMolLinearHead,
|
|
128
|
+
MaskLMHead as UniMolMaskLMHead,
|
|
129
|
+
download_unimol_checkpoint,
|
|
130
|
+
load_unimol_weights,
|
|
131
|
+
)
|
|
132
|
+
from .unimol2 import (
|
|
133
|
+
UniMol2Model,
|
|
134
|
+
AtomFeature as UniMol2AtomFeature,
|
|
135
|
+
EdgeFeature as UniMol2EdgeFeature,
|
|
136
|
+
SE3InvariantKernel as UniMol2SE3Kernel,
|
|
137
|
+
MovementPredictionHead as UniMol2MovementHead,
|
|
138
|
+
download_unimol2_checkpoint,
|
|
139
|
+
load_unimol2_weights,
|
|
140
|
+
)
|
|
141
|
+
from .unimol_plus import (
|
|
142
|
+
UniMolPlusPCQModel,
|
|
143
|
+
UniMolPlusOC20Model,
|
|
144
|
+
EnergyHead as UniMolPlusEnergyHead,
|
|
145
|
+
download_unimol_plus_checkpoint,
|
|
146
|
+
load_unimol_plus_weights,
|
|
147
|
+
)
|
|
148
|
+
from .unimol_docking_v2 import (
|
|
149
|
+
DockingPoseModelV2,
|
|
150
|
+
download_unimol_docking_checkpoint,
|
|
151
|
+
load_unimol_docking_weights,
|
|
152
|
+
)
|
|
153
|
+
from . import materials
|
|
154
|
+
from . import bio
|
|
155
|
+
from . import chemistry
|
|
156
|
+
from .materials import (
|
|
157
|
+
MEGNet,
|
|
158
|
+
MEGNetBlock,
|
|
159
|
+
MEGNetGraphConv,
|
|
160
|
+
M3GNet,
|
|
161
|
+
M3GNetBlock,
|
|
162
|
+
M3GNetGraphConv,
|
|
163
|
+
ThreeBodyInteractions,
|
|
164
|
+
TensorNet,
|
|
165
|
+
TensorEmbedding,
|
|
166
|
+
TensorNetInteraction,
|
|
167
|
+
CHGNet,
|
|
168
|
+
CHGNetAtomGraphBlock,
|
|
169
|
+
CHGNetBondGraphBlock,
|
|
170
|
+
SO3Net,
|
|
171
|
+
SO3Convolution,
|
|
172
|
+
RealSphericalHarmonics,
|
|
173
|
+
GRACE,
|
|
174
|
+
GraceSPBasis,
|
|
175
|
+
GraceACEStack,
|
|
176
|
+
QET,
|
|
177
|
+
LinearQeq,
|
|
178
|
+
ElectrostaticPotential,
|
|
179
|
+
TransformedTargetModel,
|
|
180
|
+
Potential,
|
|
181
|
+
BondExpansion as MatGLBondExpansion,
|
|
182
|
+
GaussianExpansion as MatGLGaussianExpansion,
|
|
183
|
+
RadialBesselFunction as MatGLRadialBesselFunction,
|
|
184
|
+
FourierExpansion as MatGLFourierExpansion,
|
|
185
|
+
ChebyshevRadialBasis as MatGLChebyshevRadialBasis,
|
|
186
|
+
SphericalBesselFunction as MatGLSphericalBesselFunction,
|
|
187
|
+
SphericalBesselWithHarmonics as MatGLSphericalBesselWithHarmonics,
|
|
188
|
+
ReduceReadOut as MatGLReduceReadOut,
|
|
189
|
+
WeightedReadOut as MatGLWeightedReadOut,
|
|
190
|
+
WeightedAtomReadOut as MatGLWeightedAtomReadOut,
|
|
191
|
+
Set2SetReadOut as MatGLSet2SetReadOut,
|
|
192
|
+
EdgeSet2Set as MatGLEdgeSet2Set,
|
|
193
|
+
download_matgl_checkpoint,
|
|
194
|
+
load_matgl_weights,
|
|
195
|
+
load_model as load_matgl_model,
|
|
196
|
+
get_available_pretrained_models as get_available_matgl_models,
|
|
197
|
+
)
|
|
198
|
+
|
|
199
|
+
__all__ = [
|
|
200
|
+
"materials",
|
|
201
|
+
"bio",
|
|
202
|
+
"chemistry",
|
|
203
|
+
"MEGNet",
|
|
204
|
+
"MEGNetBlock",
|
|
205
|
+
"MEGNetGraphConv",
|
|
206
|
+
"M3GNet",
|
|
207
|
+
"M3GNetBlock",
|
|
208
|
+
"M3GNetGraphConv",
|
|
209
|
+
"ThreeBodyInteractions",
|
|
210
|
+
"TensorNet",
|
|
211
|
+
"TensorEmbedding",
|
|
212
|
+
"TensorNetInteraction",
|
|
213
|
+
"CHGNet",
|
|
214
|
+
"CHGNetAtomGraphBlock",
|
|
215
|
+
"CHGNetBondGraphBlock",
|
|
216
|
+
"SO3Net",
|
|
217
|
+
"SO3Convolution",
|
|
218
|
+
"RealSphericalHarmonics",
|
|
219
|
+
"GRACE",
|
|
220
|
+
"GraceSPBasis",
|
|
221
|
+
"GraceACEStack",
|
|
222
|
+
"QET",
|
|
223
|
+
"LinearQeq",
|
|
224
|
+
"ElectrostaticPotential",
|
|
225
|
+
"TransformedTargetModel",
|
|
226
|
+
"Potential",
|
|
227
|
+
"MatGLBondExpansion",
|
|
228
|
+
"MatGLGaussianExpansion",
|
|
229
|
+
"MatGLRadialBesselFunction",
|
|
230
|
+
"MatGLFourierExpansion",
|
|
231
|
+
"MatGLChebyshevRadialBasis",
|
|
232
|
+
"MatGLSphericalBesselFunction",
|
|
233
|
+
"MatGLSphericalBesselWithHarmonics",
|
|
234
|
+
"MatGLReduceReadOut",
|
|
235
|
+
"MatGLWeightedReadOut",
|
|
236
|
+
"MatGLWeightedAtomReadOut",
|
|
237
|
+
"MatGLSet2SetReadOut",
|
|
238
|
+
"MatGLEdgeSet2Set",
|
|
239
|
+
"download_matgl_checkpoint",
|
|
240
|
+
"load_matgl_weights",
|
|
241
|
+
"load_matgl_model",
|
|
242
|
+
"get_available_matgl_models",
|
|
243
|
+
"MLP",
|
|
244
|
+
"ARLinkPredictor",
|
|
245
|
+
"InnerProductDecoder",
|
|
246
|
+
"GAE",
|
|
247
|
+
"VGAE",
|
|
248
|
+
"ARGA",
|
|
249
|
+
"ARGVA",
|
|
250
|
+
"DeepGraphInfomax",
|
|
251
|
+
"DeepGCNLayer",
|
|
252
|
+
"AttentiveFP",
|
|
253
|
+
"JumpingKnowledge",
|
|
254
|
+
"HeteroJumpingKnowledge",
|
|
255
|
+
"MaskLabel",
|
|
256
|
+
"MetaLayer",
|
|
257
|
+
"PMLP",
|
|
258
|
+
"Polynormer",
|
|
259
|
+
"BasicGNN",
|
|
260
|
+
"GCN",
|
|
261
|
+
"GraphSAGE",
|
|
262
|
+
"GIN",
|
|
263
|
+
"GAT",
|
|
264
|
+
"PNA",
|
|
265
|
+
"EdgeCNN",
|
|
266
|
+
"LabelPropagation",
|
|
267
|
+
"CorrectAndSmooth",
|
|
268
|
+
"LightGCN",
|
|
269
|
+
"BPRLoss",
|
|
270
|
+
"LINKX",
|
|
271
|
+
"SparseLinear",
|
|
272
|
+
"RECT_L",
|
|
273
|
+
"SignedGCN",
|
|
274
|
+
"NeuralFingerprint",
|
|
275
|
+
"GraphUNet",
|
|
276
|
+
"GroupAddRev",
|
|
277
|
+
"SGFormer",
|
|
278
|
+
"Node2Vec",
|
|
279
|
+
"MetaPath2Vec",
|
|
280
|
+
"RENet",
|
|
281
|
+
"TGNMemory",
|
|
282
|
+
"IdentityMessage",
|
|
283
|
+
"LastAggregator",
|
|
284
|
+
"MeanAggregator",
|
|
285
|
+
"TimeEncoder",
|
|
286
|
+
"LastNeighborLoader",
|
|
287
|
+
"SchNet",
|
|
288
|
+
"CFConv",
|
|
289
|
+
"SchNetInteractionBlock",
|
|
290
|
+
"GaussianSmearing",
|
|
291
|
+
"ShiftedSoftplus",
|
|
292
|
+
"RadiusInteractionGraph",
|
|
293
|
+
"DimeNet",
|
|
294
|
+
"DimeNetPlusPlus",
|
|
295
|
+
"BesselBasisLayer",
|
|
296
|
+
"SphericalBasisLayer",
|
|
297
|
+
"triplets",
|
|
298
|
+
"GNNFF",
|
|
299
|
+
"NodeBlock",
|
|
300
|
+
"EdgeBlock",
|
|
301
|
+
"GaussianFilter",
|
|
302
|
+
"GPSE",
|
|
303
|
+
"GPSENodeEncoder",
|
|
304
|
+
"GeneralLayer",
|
|
305
|
+
"GeneralMultiLayer",
|
|
306
|
+
"GNNStackStage",
|
|
307
|
+
"GNNInductiveHybridMultiHead",
|
|
308
|
+
"ViSNet",
|
|
309
|
+
"LPFormer",
|
|
310
|
+
"LPAttLayer",
|
|
311
|
+
"GraphMAE2",
|
|
312
|
+
"sce_loss",
|
|
313
|
+
"load_graphmae2_weights",
|
|
314
|
+
"download_graphmae2_checkpoint",
|
|
315
|
+
"Graphormer",
|
|
316
|
+
"GraphNodeFeature",
|
|
317
|
+
"GraphAttnBias",
|
|
318
|
+
"GraphormerMultiheadAttention",
|
|
319
|
+
"GraphormerGraphEncoderLayer",
|
|
320
|
+
"GraphormerGraphEncoder",
|
|
321
|
+
"load_graphormer_weights",
|
|
322
|
+
"download_graphormer_checkpoint",
|
|
323
|
+
"Graphormer3D",
|
|
324
|
+
"GaussianLayer",
|
|
325
|
+
"RBF",
|
|
326
|
+
"Graphormer3DEncoderLayer",
|
|
327
|
+
"NodeTaskHead",
|
|
328
|
+
"load_graphormer3d_weights",
|
|
329
|
+
"download_graphormer3d_checkpoint",
|
|
330
|
+
"to_captum_model",
|
|
331
|
+
"to_captum_input",
|
|
332
|
+
"captum_output_to_dicts",
|
|
333
|
+
"GPSModel",
|
|
334
|
+
"GPSLayer",
|
|
335
|
+
"CustomGatedGCN",
|
|
336
|
+
"AtomEncoder",
|
|
337
|
+
"BondEncoder",
|
|
338
|
+
"RWSEEncoder",
|
|
339
|
+
"SANGraphHead",
|
|
340
|
+
"load_gps_weights",
|
|
341
|
+
"download_gps_checkpoint",
|
|
342
|
+
"GROVER",
|
|
343
|
+
"GTransEncoder",
|
|
344
|
+
"Readout",
|
|
345
|
+
"load_grover_weights",
|
|
346
|
+
"download_grover_checkpoint",
|
|
347
|
+
"MoleBERT",
|
|
348
|
+
"MoleBERTGNN",
|
|
349
|
+
"MoleBERTGINConv",
|
|
350
|
+
"load_mole_bert_weights",
|
|
351
|
+
"download_mole_bert_checkpoint",
|
|
352
|
+
"UniMolModel",
|
|
353
|
+
"UniMolConfGenModel",
|
|
354
|
+
"UniMolDockingModel",
|
|
355
|
+
"UniMolGaussianLayer",
|
|
356
|
+
"UniMolNumericalEmbed",
|
|
357
|
+
"UniMolNonLinearHead",
|
|
358
|
+
"UniMolDistanceHead",
|
|
359
|
+
"UniMolClassificationHead",
|
|
360
|
+
"UniMolLinearHead",
|
|
361
|
+
"UniMolMaskLMHead",
|
|
362
|
+
"download_unimol_checkpoint",
|
|
363
|
+
"load_unimol_weights",
|
|
364
|
+
"UniMol2Model",
|
|
365
|
+
"UniMol2AtomFeature",
|
|
366
|
+
"UniMol2EdgeFeature",
|
|
367
|
+
"UniMol2SE3Kernel",
|
|
368
|
+
"UniMol2MovementHead",
|
|
369
|
+
"download_unimol2_checkpoint",
|
|
370
|
+
"load_unimol2_weights",
|
|
371
|
+
"UniMolPlusPCQModel",
|
|
372
|
+
"UniMolPlusOC20Model",
|
|
373
|
+
"UniMolPlusEnergyHead",
|
|
374
|
+
"download_unimol_plus_checkpoint",
|
|
375
|
+
"load_unimol_plus_weights",
|
|
376
|
+
"DockingPoseModelV2",
|
|
377
|
+
"download_unimol_docking_checkpoint",
|
|
378
|
+
"load_unimol_docking_weights",
|
|
379
|
+
]
|
|
380
|
+
|
|
381
|
+
# Inject Hugging Face Hub capabilities (from_pretrained, save_pretrained, push_to_hub, predict)
|
|
382
|
+
# to all models in k3_node.models
|
|
383
|
+
import keras
|
|
384
|
+
from k3_node.hub.hub_mixin import K3NodeHubMixin
|
|
385
|
+
|
|
386
|
+
def _defines_own(cls, attr):
|
|
387
|
+
"""True if a k3_node class in ``cls``'s MRO defines ``attr`` itself (e.g. a model-specific
|
|
388
|
+
``from_pretrained`` that loads original checkpoints), which must not be overwritten."""
|
|
389
|
+
return any(attr in vars(klass) for klass in cls.__mro__ if klass.__module__.startswith("k3_node"))
|
|
390
|
+
|
|
391
|
+
|
|
392
|
+
def _is_graph_input(data):
|
|
393
|
+
if hasattr(data, "edge_index") or (hasattr(data, "z") and hasattr(data, "pos")):
|
|
394
|
+
return True
|
|
395
|
+
return isinstance(data, dict) and any(k in data for k in ("edge_index", "pos", "z"))
|
|
396
|
+
|
|
397
|
+
|
|
398
|
+
def _graph_aware_predict(self, data=None, *args, **kwargs):
|
|
399
|
+
"""Graph inputs (``Data``, ``Batch``, graph dicts) use the hub-style ``predict``; everything
|
|
400
|
+
else keeps Keras' batched ``Model.predict`` (arrays, ``tf.data``, ``PyDataset``, ...)."""
|
|
401
|
+
if _is_graph_input(data):
|
|
402
|
+
return K3NodeHubMixin.predict(self, data, *args, **kwargs)
|
|
403
|
+
return keras.Model.predict(self, data, *args, **kwargs)
|
|
404
|
+
|
|
405
|
+
|
|
406
|
+
for _name in list(__all__):
|
|
407
|
+
_obj = globals().get(_name)
|
|
408
|
+
if isinstance(_obj, type) and issubclass(_obj, (keras.Model, keras.layers.Layer)):
|
|
409
|
+
if not issubclass(_obj, K3NodeHubMixin):
|
|
410
|
+
if not _defines_own(_obj, "from_pretrained"):
|
|
411
|
+
_obj.from_pretrained = classmethod(K3NodeHubMixin.from_pretrained.__func__)
|
|
412
|
+
if not _defines_own(_obj, "save_pretrained"):
|
|
413
|
+
_obj.save_pretrained = K3NodeHubMixin.save_pretrained
|
|
414
|
+
if not _defines_own(_obj, "push_to_hub"):
|
|
415
|
+
_obj.push_to_hub = K3NodeHubMixin.push_to_hub
|
|
416
|
+
if not _defines_own(_obj, "predict"):
|
|
417
|
+
if issubclass(_obj, keras.Model):
|
|
418
|
+
_obj.predict = _graph_aware_predict
|
|
419
|
+
else:
|
|
420
|
+
_obj.predict = K3NodeHubMixin.predict
|
|
421
|
+
if not hasattr(_obj, "_get_config"):
|
|
422
|
+
_obj._get_config = K3NodeHubMixin._get_config
|
|
423
|
+
|
|
424
|
+
|
|
@@ -0,0 +1,232 @@
|
|
|
1
|
+
from typing import Optional
|
|
2
|
+
import keras
|
|
3
|
+
from keras import ops
|
|
4
|
+
|
|
5
|
+
from k3_node.layers.conv import GATConv
|
|
6
|
+
from k3_node.layers.conv.message_passing import MessagePassing
|
|
7
|
+
from k3_node.layers.conv.utils import softmax
|
|
8
|
+
from k3_node.layers.pool import global_add_pool
|
|
9
|
+
from k3_node.layers.pool.glob import _infer_size
|
|
10
|
+
from k3_node.ops.segment import segment_sum
|
|
11
|
+
|
|
12
|
+
try:
|
|
13
|
+
from keras.src.backend.common.symbolic_scope import in_symbolic_scope
|
|
14
|
+
except ImportError:
|
|
15
|
+
def in_symbolic_scope():
|
|
16
|
+
return False
|
|
17
|
+
|
|
18
|
+
|
|
19
|
+
def _gru_step(cell, x, h):
|
|
20
|
+
out, _ = cell(x, [h])
|
|
21
|
+
return out
|
|
22
|
+
|
|
23
|
+
|
|
24
|
+
class GATEConv(MessagePassing):
|
|
25
|
+
r"""The edge-conditioned attention layer used as the first message
|
|
26
|
+
passing step of `AttentiveFP`."""
|
|
27
|
+
def __init__(self, in_channels: int, out_channels: int, edge_dim: int,
|
|
28
|
+
dropout: float = 0.0, **kwargs):
|
|
29
|
+
super().__init__(aggr="add", node_dim=0, **kwargs)
|
|
30
|
+
|
|
31
|
+
self.in_channels = in_channels
|
|
32
|
+
self.out_channels = out_channels
|
|
33
|
+
self.edge_dim = edge_dim
|
|
34
|
+
self.dropout_rate = dropout
|
|
35
|
+
|
|
36
|
+
self.lin1 = keras.layers.Dense(out_channels, use_bias=False)
|
|
37
|
+
self.lin2 = keras.layers.Dense(out_channels, use_bias=False)
|
|
38
|
+
self.dropout = keras.layers.Dropout(dropout) if dropout > 0.0 else None
|
|
39
|
+
|
|
40
|
+
self.lin1.build((None, in_channels + edge_dim))
|
|
41
|
+
self.lin2.build((None, out_channels))
|
|
42
|
+
|
|
43
|
+
self.att_l = self.add_weight(shape=(1, out_channels), initializer="glorot_uniform", name="att_l")
|
|
44
|
+
self.att_r = self.add_weight(shape=(1, in_channels), initializer="glorot_uniform", name="att_r")
|
|
45
|
+
self.bias = self.add_weight(shape=(out_channels,), initializer="zeros", name="bias")
|
|
46
|
+
|
|
47
|
+
def build(self, input_shape=None):
|
|
48
|
+
self.built = True
|
|
49
|
+
|
|
50
|
+
def call(self, x, edge_index, edge_attr, training=None):
|
|
51
|
+
row, col = ops.cast(edge_index[0], "int32"), ops.cast(edge_index[1], "int32")
|
|
52
|
+
|
|
53
|
+
x_j = ops.take(x, row, axis=0)
|
|
54
|
+
x_i = ops.take(x, col, axis=0)
|
|
55
|
+
|
|
56
|
+
edge_attr = ops.cast(edge_attr, x.dtype)
|
|
57
|
+
h = ops.leaky_relu(self.lin1(ops.concatenate([x_j, edge_attr], axis=-1)), negative_slope=0.01)
|
|
58
|
+
alpha_j = ops.sum(h * self.att_l, axis=-1)
|
|
59
|
+
alpha_i = ops.sum(x_i * self.att_r, axis=-1)
|
|
60
|
+
alpha = ops.leaky_relu(alpha_j + alpha_i, negative_slope=0.01)
|
|
61
|
+
|
|
62
|
+
num_nodes = ops.shape(x)[0]
|
|
63
|
+
alpha = softmax(alpha, col, num_nodes=num_nodes, dim=0)
|
|
64
|
+
if self.dropout is not None:
|
|
65
|
+
alpha = self.dropout(alpha, training=training)
|
|
66
|
+
|
|
67
|
+
message = self.lin2(x_j) * ops.expand_dims(alpha, -1)
|
|
68
|
+
out = segment_sum(message, col, num_segments=num_nodes)
|
|
69
|
+
return out + self.bias
|
|
70
|
+
|
|
71
|
+
|
|
72
|
+
class AttentiveFP(keras.Model):
|
|
73
|
+
r"""The Attentive FP model for molecular representation learning from the
|
|
74
|
+
`"Pushing the Boundaries of Molecular Representation for Drug Discovery
|
|
75
|
+
with the Graph Attention Mechanism"
|
|
76
|
+
<https://pubs.acs.org/doi/10.1021/acs.jmedchem.9b00959>`_ paper, based on
|
|
77
|
+
graph attention mechanisms.
|
|
78
|
+
|
|
79
|
+
Args:
|
|
80
|
+
in_channels (int): Size of each input sample.
|
|
81
|
+
hidden_channels (int): Hidden node feature dimensionality.
|
|
82
|
+
out_channels (int): Size of each output sample.
|
|
83
|
+
edge_dim (int): Edge feature dimensionality.
|
|
84
|
+
num_layers (int): Number of GNN layers.
|
|
85
|
+
num_timesteps (int): Number of iterative refinement steps for global
|
|
86
|
+
readout.
|
|
87
|
+
dropout (float, optional): Dropout probability. (default: `0.0`)
|
|
88
|
+
batch_size (int, optional): Fixed batch size (number of graphs) for JAX/XLA
|
|
89
|
+
static shape compatibility. (default: `None`)
|
|
90
|
+
|
|
91
|
+
Example:
|
|
92
|
+
```python
|
|
93
|
+
import numpy as np
|
|
94
|
+
from k3_node.models import AttentiveFP
|
|
95
|
+
|
|
96
|
+
x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
|
|
97
|
+
edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
|
|
98
|
+
edge_attr = np.random.rand(30, 3).astype("float32") # bond features
|
|
99
|
+
|
|
100
|
+
batch = np.repeat([0, 1], 5) # two molecules with 5 atoms each
|
|
101
|
+
model = AttentiveFP(in_channels=8, hidden_channels=16, out_channels=1, edge_dim=3,
|
|
102
|
+
num_layers=2, num_timesteps=2)
|
|
103
|
+
out = model(x, edge_index, edge_attr, batch) # one prediction per molecule
|
|
104
|
+
print(tuple(out.shape)) # (2, 1)
|
|
105
|
+
```
|
|
106
|
+
"""
|
|
107
|
+
def __init__(
|
|
108
|
+
self,
|
|
109
|
+
in_channels: int,
|
|
110
|
+
hidden_channels: int,
|
|
111
|
+
out_channels: int,
|
|
112
|
+
edge_dim: int,
|
|
113
|
+
num_layers: int,
|
|
114
|
+
num_timesteps: int,
|
|
115
|
+
dropout: float = 0.0,
|
|
116
|
+
batch_size: Optional[int] = None,
|
|
117
|
+
**kwargs,
|
|
118
|
+
):
|
|
119
|
+
super().__init__(**kwargs)
|
|
120
|
+
|
|
121
|
+
self.in_channels = in_channels
|
|
122
|
+
self.hidden_channels = hidden_channels
|
|
123
|
+
self.out_channels = out_channels
|
|
124
|
+
self.edge_dim = edge_dim
|
|
125
|
+
self.num_layers = num_layers
|
|
126
|
+
self.num_timesteps = num_timesteps
|
|
127
|
+
self.dropout_rate = dropout
|
|
128
|
+
self.batch_size = batch_size
|
|
129
|
+
|
|
130
|
+
self.lin1 = keras.layers.Dense(hidden_channels)
|
|
131
|
+
self.lin1.build((None, in_channels))
|
|
132
|
+
|
|
133
|
+
self.gate_conv = GATEConv(hidden_channels, hidden_channels, edge_dim, dropout)
|
|
134
|
+
self.gru = keras.layers.GRUCell(hidden_channels)
|
|
135
|
+
self.gru.build((None, hidden_channels))
|
|
136
|
+
|
|
137
|
+
self.atom_convs = []
|
|
138
|
+
self.atom_grus = []
|
|
139
|
+
for _ in range(num_layers - 1):
|
|
140
|
+
conv = GATConv(hidden_channels, hidden_channels, dropout=dropout,
|
|
141
|
+
add_self_loops=False, negative_slope=0.01)
|
|
142
|
+
conv.build((None, hidden_channels))
|
|
143
|
+
self.atom_convs.append(conv)
|
|
144
|
+
gru = keras.layers.GRUCell(hidden_channels)
|
|
145
|
+
gru.build((None, hidden_channels))
|
|
146
|
+
self.atom_grus.append(gru)
|
|
147
|
+
|
|
148
|
+
self.mol_conv = GATConv(hidden_channels, hidden_channels, dropout=dropout,
|
|
149
|
+
add_self_loops=False, negative_slope=0.01)
|
|
150
|
+
self.mol_conv.build([(None, hidden_channels), (None, hidden_channels)])
|
|
151
|
+
self.mol_gru = keras.layers.GRUCell(hidden_channels)
|
|
152
|
+
self.mol_gru.build((None, hidden_channels))
|
|
153
|
+
|
|
154
|
+
self.lin2 = keras.layers.Dense(out_channels)
|
|
155
|
+
self.lin2.build((None, hidden_channels))
|
|
156
|
+
|
|
157
|
+
self.dropout = keras.layers.Dropout(dropout) if dropout > 0.0 else None
|
|
158
|
+
self.built = True
|
|
159
|
+
|
|
160
|
+
def call(self, x, edge_index=None, edge_attr=None, batch=None, batch_size=None, training=None):
|
|
161
|
+
if isinstance(x, dict):
|
|
162
|
+
edge_index = x.get("edge_index")
|
|
163
|
+
edge_attr = x.get("edge_attr")
|
|
164
|
+
batch = x.get("batch")
|
|
165
|
+
batch_size = x.get("batch_size", batch_size)
|
|
166
|
+
x = x.get("x")
|
|
167
|
+
elif isinstance(x, (tuple, list)) and edge_index is None:
|
|
168
|
+
if len(x) >= 4:
|
|
169
|
+
x, edge_index, edge_attr, batch = x[0], x[1], x[2], x[3]
|
|
170
|
+
elif len(x) == 3:
|
|
171
|
+
x, edge_index, edge_attr = x[0], x[1], x[2]
|
|
172
|
+
|
|
173
|
+
bs = batch_size if batch_size is not None else self.batch_size
|
|
174
|
+
x = ops.cast(x, "float32")
|
|
175
|
+
if edge_attr is not None:
|
|
176
|
+
edge_attr = ops.cast(edge_attr, "float32")
|
|
177
|
+
# Atom Embedding:
|
|
178
|
+
x = ops.leaky_relu(self.lin1(x), negative_slope=0.01)
|
|
179
|
+
|
|
180
|
+
h = ops.elu(self.gate_conv(x, edge_index, edge_attr, training=training))
|
|
181
|
+
if self.dropout is not None:
|
|
182
|
+
h = self.dropout(h, training=training)
|
|
183
|
+
x = ops.relu(_gru_step(self.gru, h, x))
|
|
184
|
+
|
|
185
|
+
for conv, gru in zip(self.atom_convs, self.atom_grus):
|
|
186
|
+
h = conv(x, edge_index, training=training)
|
|
187
|
+
h = ops.elu(h)
|
|
188
|
+
if self.dropout is not None:
|
|
189
|
+
h = self.dropout(h, training=training)
|
|
190
|
+
x = ops.relu(_gru_step(gru, h, x))
|
|
191
|
+
|
|
192
|
+
# Molecule Embedding:
|
|
193
|
+
if batch is None:
|
|
194
|
+
batch = ops.zeros((ops.shape(x)[0],), dtype="int32")
|
|
195
|
+
else:
|
|
196
|
+
batch = ops.cast(batch, "int32")
|
|
197
|
+
num_nodes = ops.shape(batch)[0]
|
|
198
|
+
row = ops.arange(num_nodes, dtype="int32")
|
|
199
|
+
mol_edge_index = ops.stack([row, batch], axis=0)
|
|
200
|
+
|
|
201
|
+
if keras.config.backend() == "jax":
|
|
202
|
+
size = bs
|
|
203
|
+
elif in_symbolic_scope():
|
|
204
|
+
size = bs
|
|
205
|
+
else:
|
|
206
|
+
size = _infer_size(batch)
|
|
207
|
+
if size is None:
|
|
208
|
+
size = bs
|
|
209
|
+
|
|
210
|
+
out = ops.relu(global_add_pool(x, batch, size=size))
|
|
211
|
+
for _ in range(self.num_timesteps):
|
|
212
|
+
h = ops.elu(self.mol_conv((x, out), mol_edge_index, training=training))
|
|
213
|
+
if self.dropout is not None:
|
|
214
|
+
h = self.dropout(h, training=training)
|
|
215
|
+
out = ops.relu(_gru_step(self.mol_gru, h, out))
|
|
216
|
+
|
|
217
|
+
# Predictor:
|
|
218
|
+
if self.dropout is not None:
|
|
219
|
+
out = self.dropout(out, training=training)
|
|
220
|
+
return self.lin2(out)
|
|
221
|
+
|
|
222
|
+
def __repr__(self) -> str:
|
|
223
|
+
return (
|
|
224
|
+
f"{self.__class__.__name__}("
|
|
225
|
+
f"in_channels={self.in_channels}, "
|
|
226
|
+
f"hidden_channels={self.hidden_channels}, "
|
|
227
|
+
f"out_channels={self.out_channels}, "
|
|
228
|
+
f"edge_dim={self.edge_dim}, "
|
|
229
|
+
f"num_layers={self.num_layers}, "
|
|
230
|
+
f"num_timesteps={self.num_timesteps}"
|
|
231
|
+
f")"
|
|
232
|
+
)
|