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,379 @@
|
|
|
1
|
+
import os
|
|
2
|
+
import os.path as osp
|
|
3
|
+
import shutil
|
|
4
|
+
from typing import Optional, Tuple, Union
|
|
5
|
+
|
|
6
|
+
import keras
|
|
7
|
+
from keras import layers, ops
|
|
8
|
+
|
|
9
|
+
from k3_node.layers.conv.message_passing import MessagePassing
|
|
10
|
+
from k3_node.layers.conv.utils import add_self_loops
|
|
11
|
+
from k3_node.layers.pool import global_add_pool, global_max_pool, global_mean_pool
|
|
12
|
+
from k3_node.data.download import download_url
|
|
13
|
+
from k3_node.ops.creation import full
|
|
14
|
+
|
|
15
|
+
|
|
16
|
+
MOLE_BERT_URL = (
|
|
17
|
+
"https://github.com/junxia97/Mole-BERT/raw/refs/heads/main/model_gin/Mole-BERT.pth"
|
|
18
|
+
)
|
|
19
|
+
|
|
20
|
+
|
|
21
|
+
class MoleBERTGINConv(MessagePassing):
|
|
22
|
+
"""Extension of GIN aggregation to incorporate categorical edge information with self-loops.
|
|
23
|
+
|
|
24
|
+
Matches the GINConv variant from Hu et al. used in Mole-BERT.
|
|
25
|
+
|
|
26
|
+
Example:
|
|
27
|
+
```python
|
|
28
|
+
import numpy as np
|
|
29
|
+
from k3_node.models import MoleBERTGINConv
|
|
30
|
+
|
|
31
|
+
x = np.random.rand(4, 32).astype("float32") # atom embeddings
|
|
32
|
+
edge_index = np.array([[0, 1, 1, 2], [1, 0, 2, 1]])
|
|
33
|
+
edge_attr = np.array([[0, 0], [0, 0], [1, 0], [1, 0]]) # bond type, bond direction
|
|
34
|
+
conv = MoleBERTGINConv(emb_dim=32)
|
|
35
|
+
print(tuple(conv(x, edge_index, edge_attr).shape)) # (4, 32)
|
|
36
|
+
```
|
|
37
|
+
"""
|
|
38
|
+
|
|
39
|
+
def __init__(
|
|
40
|
+
self,
|
|
41
|
+
emb_dim: int = 300,
|
|
42
|
+
out_dim: Optional[int] = None,
|
|
43
|
+
num_bond_type: int = 6,
|
|
44
|
+
num_bond_direction: int = 3,
|
|
45
|
+
**kwargs,
|
|
46
|
+
):
|
|
47
|
+
super().__init__(aggr="add", **kwargs)
|
|
48
|
+
self.emb_dim = emb_dim
|
|
49
|
+
self.out_dim = emb_dim if out_dim is None else out_dim
|
|
50
|
+
self.num_bond_type = num_bond_type
|
|
51
|
+
self.num_bond_direction = num_bond_direction
|
|
52
|
+
|
|
53
|
+
self.mlp = keras.Sequential(
|
|
54
|
+
[
|
|
55
|
+
layers.Dense(2 * emb_dim, activation="relu", name="mlp_0"),
|
|
56
|
+
layers.Dense(self.out_dim, name="mlp_2"),
|
|
57
|
+
],
|
|
58
|
+
name="mlp",
|
|
59
|
+
)
|
|
60
|
+
self.edge_embedding1 = layers.Embedding(
|
|
61
|
+
num_bond_type, emb_dim, name="edge_embedding1"
|
|
62
|
+
)
|
|
63
|
+
self.edge_embedding2 = layers.Embedding(
|
|
64
|
+
num_bond_direction, emb_dim, name="edge_embedding2"
|
|
65
|
+
)
|
|
66
|
+
|
|
67
|
+
def build(self, input_shape=None):
|
|
68
|
+
self.mlp.build((None, self.emb_dim))
|
|
69
|
+
self.edge_embedding1.build((None,))
|
|
70
|
+
self.edge_embedding2.build((None,))
|
|
71
|
+
super().build(input_shape)
|
|
72
|
+
|
|
73
|
+
def call(self, x, edge_index, edge_attr):
|
|
74
|
+
num_nodes = ops.shape(x)[0]
|
|
75
|
+
|
|
76
|
+
# Add self-loops to edge space
|
|
77
|
+
edge_index, _ = add_self_loops(edge_index, num_nodes=num_nodes)
|
|
78
|
+
|
|
79
|
+
# Add features corresponding to self-loop edges: [4, 0]
|
|
80
|
+
self_loop_attr = ops.stack(
|
|
81
|
+
[
|
|
82
|
+
full((num_nodes,), 4, dtype=edge_attr.dtype),
|
|
83
|
+
ops.zeros((num_nodes,), dtype=edge_attr.dtype),
|
|
84
|
+
],
|
|
85
|
+
axis=1,
|
|
86
|
+
)
|
|
87
|
+
edge_attr = ops.concatenate([edge_attr, self_loop_attr], axis=0)
|
|
88
|
+
|
|
89
|
+
edge_embeddings = self.edge_embedding1(edge_attr[:, 0]) + self.edge_embedding2(
|
|
90
|
+
edge_attr[:, 1]
|
|
91
|
+
)
|
|
92
|
+
|
|
93
|
+
return self.propagate(
|
|
94
|
+
edge_index, x=x, edge_attr=edge_embeddings, size=(num_nodes, num_nodes)
|
|
95
|
+
)
|
|
96
|
+
|
|
97
|
+
def message(self, x_j, edge_attr):
|
|
98
|
+
return x_j + edge_attr
|
|
99
|
+
|
|
100
|
+
def update(self, aggr_out):
|
|
101
|
+
return self.mlp(aggr_out)
|
|
102
|
+
|
|
103
|
+
|
|
104
|
+
class MoleBERTGNN(layers.Layer):
|
|
105
|
+
"""5-layer GIN encoder backbone of Mole-BERT with Jumping Knowledge.
|
|
106
|
+
|
|
107
|
+
Example:
|
|
108
|
+
```python
|
|
109
|
+
import numpy as np
|
|
110
|
+
from k3_node.models import MoleBERTGNN
|
|
111
|
+
|
|
112
|
+
x = np.stack([np.random.randint(0, 119, size=5), np.random.randint(0, 3, size=5)], axis=1) # atom type, chirality
|
|
113
|
+
edge_index = np.array([[0, 1, 1, 2, 3, 4], [1, 0, 2, 1, 4, 3]])
|
|
114
|
+
edge_attr = np.stack([np.random.randint(0, 5, size=6), np.random.randint(0, 3, size=6)], axis=1) # bond type, direction
|
|
115
|
+
batch = np.array([0, 0, 0, 1, 1]) # two molecules
|
|
116
|
+
|
|
117
|
+
gnn = MoleBERTGNN(num_layer=3, emb_dim=32, JK="last")
|
|
118
|
+
print(tuple(gnn(x, edge_index, edge_attr).shape)) # (5, 32)
|
|
119
|
+
```
|
|
120
|
+
"""
|
|
121
|
+
|
|
122
|
+
def __init__(
|
|
123
|
+
self,
|
|
124
|
+
num_layer: int = 5,
|
|
125
|
+
emb_dim: int = 300,
|
|
126
|
+
num_atom_type: int = 120,
|
|
127
|
+
num_chirality_tag: int = 3,
|
|
128
|
+
JK: str = "last",
|
|
129
|
+
drop_ratio: float = 0.0,
|
|
130
|
+
**kwargs,
|
|
131
|
+
):
|
|
132
|
+
super().__init__(**kwargs)
|
|
133
|
+
if num_layer < 2:
|
|
134
|
+
raise ValueError("Number of GNN layers must be greater than 1.")
|
|
135
|
+
self.num_layer = num_layer
|
|
136
|
+
self.emb_dim = emb_dim
|
|
137
|
+
self.num_atom_type = num_atom_type
|
|
138
|
+
self.num_chirality_tag = num_chirality_tag
|
|
139
|
+
self.JK = JK
|
|
140
|
+
self.drop_ratio = drop_ratio
|
|
141
|
+
|
|
142
|
+
self.x_embedding1 = layers.Embedding(
|
|
143
|
+
num_atom_type, emb_dim, name="x_embedding1"
|
|
144
|
+
)
|
|
145
|
+
self.x_embedding2 = layers.Embedding(
|
|
146
|
+
num_chirality_tag, emb_dim, name="x_embedding2"
|
|
147
|
+
)
|
|
148
|
+
|
|
149
|
+
self.gnns = [
|
|
150
|
+
MoleBERTGINConv(emb_dim=emb_dim, name=f"gnns_{i}")
|
|
151
|
+
for i in range(num_layer)
|
|
152
|
+
]
|
|
153
|
+
self.batch_norms = [
|
|
154
|
+
layers.BatchNormalization(
|
|
155
|
+
axis=-1, epsilon=1e-5, momentum=0.9, name=f"batch_norms_{i}"
|
|
156
|
+
)
|
|
157
|
+
for i in range(num_layer)
|
|
158
|
+
]
|
|
159
|
+
self.dropout_layer = layers.Dropout(drop_ratio)
|
|
160
|
+
|
|
161
|
+
def build(self, input_shape=None):
|
|
162
|
+
self.x_embedding1.build((None,))
|
|
163
|
+
self.x_embedding2.build((None,))
|
|
164
|
+
for i in range(self.num_layer):
|
|
165
|
+
self.gnns[i].build(None)
|
|
166
|
+
self.batch_norms[i].build((None, self.emb_dim))
|
|
167
|
+
super().build(input_shape)
|
|
168
|
+
|
|
169
|
+
def call(self, x, edge_index, edge_attr, training: bool = False):
|
|
170
|
+
h = self.x_embedding1(x[:, 0]) + self.x_embedding2(x[:, 1])
|
|
171
|
+
h_list = [h]
|
|
172
|
+
|
|
173
|
+
for layer in range(self.num_layer):
|
|
174
|
+
h = self.gnns[layer](h_list[layer], edge_index, edge_attr)
|
|
175
|
+
h = self.batch_norms[layer](h, training=training)
|
|
176
|
+
if layer == self.num_layer - 1:
|
|
177
|
+
# Remove ReLU for the last layer
|
|
178
|
+
h = self.dropout_layer(h, training=training)
|
|
179
|
+
else:
|
|
180
|
+
h = self.dropout_layer(ops.relu(h), training=training)
|
|
181
|
+
h_list.append(h)
|
|
182
|
+
|
|
183
|
+
if self.JK == "concat":
|
|
184
|
+
node_representation = ops.concatenate(h_list, axis=1)
|
|
185
|
+
elif self.JK == "last":
|
|
186
|
+
node_representation = h_list[-1]
|
|
187
|
+
elif self.JK == "max":
|
|
188
|
+
node_representation = ops.max(ops.stack(h_list, axis=0), axis=0)
|
|
189
|
+
elif self.JK == "sum":
|
|
190
|
+
node_representation = ops.sum(ops.stack(h_list, axis=0), axis=0)
|
|
191
|
+
else:
|
|
192
|
+
raise ValueError(f"Unknown JK mode: {self.JK}")
|
|
193
|
+
|
|
194
|
+
return node_representation
|
|
195
|
+
|
|
196
|
+
|
|
197
|
+
class MoleBERT(keras.Model):
|
|
198
|
+
"""Complete Mole-BERT Model with graph-level pooling and property prediction head.
|
|
199
|
+
|
|
200
|
+
Example:
|
|
201
|
+
```python
|
|
202
|
+
import numpy as np
|
|
203
|
+
from k3_node.models import MoleBERT
|
|
204
|
+
|
|
205
|
+
x = np.stack([np.random.randint(0, 119, size=5), np.random.randint(0, 3, size=5)], axis=1) # atom type, chirality
|
|
206
|
+
edge_index = np.array([[0, 1, 1, 2, 3, 4], [1, 0, 2, 1, 4, 3]])
|
|
207
|
+
edge_attr = np.stack([np.random.randint(0, 5, size=6), np.random.randint(0, 3, size=6)], axis=1) # bond type, direction
|
|
208
|
+
batch = np.array([0, 0, 0, 1, 1]) # two molecules
|
|
209
|
+
|
|
210
|
+
model = MoleBERT(num_layer=3, emb_dim=32, num_tasks=2, graph_pooling="mean")
|
|
211
|
+
logits, node_rep = model((x, edge_index, edge_attr, batch))
|
|
212
|
+
print(tuple(logits.shape), tuple(node_rep.shape)) # (2, 2) (5, 32): per-molecule predictions, per-atom embeddings
|
|
213
|
+
```
|
|
214
|
+
"""
|
|
215
|
+
|
|
216
|
+
def __init__(
|
|
217
|
+
self,
|
|
218
|
+
num_layer: int = 5,
|
|
219
|
+
emb_dim: int = 300,
|
|
220
|
+
num_tasks: Optional[int] = None,
|
|
221
|
+
JK: str = "last",
|
|
222
|
+
drop_ratio: float = 0.0,
|
|
223
|
+
graph_pooling: str = "mean",
|
|
224
|
+
**kwargs,
|
|
225
|
+
):
|
|
226
|
+
super().__init__(**kwargs)
|
|
227
|
+
self.num_layer = num_layer
|
|
228
|
+
self.emb_dim = emb_dim
|
|
229
|
+
self.num_tasks = num_tasks
|
|
230
|
+
self.JK = JK
|
|
231
|
+
self.drop_ratio = drop_ratio
|
|
232
|
+
self.graph_pooling = graph_pooling
|
|
233
|
+
|
|
234
|
+
self.gnn = MoleBERTGNN(
|
|
235
|
+
num_layer=num_layer,
|
|
236
|
+
emb_dim=emb_dim,
|
|
237
|
+
JK=JK,
|
|
238
|
+
drop_ratio=drop_ratio,
|
|
239
|
+
name="gnn",
|
|
240
|
+
)
|
|
241
|
+
|
|
242
|
+
if graph_pooling in ("sum", "add"):
|
|
243
|
+
self.pool_fn = global_add_pool
|
|
244
|
+
elif graph_pooling == "mean":
|
|
245
|
+
self.pool_fn = global_mean_pool
|
|
246
|
+
elif graph_pooling == "max":
|
|
247
|
+
self.pool_fn = global_max_pool
|
|
248
|
+
else:
|
|
249
|
+
raise ValueError(f"Invalid graph pooling type: '{graph_pooling}'")
|
|
250
|
+
|
|
251
|
+
if num_tasks is not None:
|
|
252
|
+
mult = (num_layer + 1) if JK == "concat" else 1
|
|
253
|
+
self.graph_pred_linear = layers.Dense(
|
|
254
|
+
num_tasks, name="graph_pred_linear"
|
|
255
|
+
)
|
|
256
|
+
else:
|
|
257
|
+
self.graph_pred_linear = None
|
|
258
|
+
|
|
259
|
+
def build(self, input_shape=None):
|
|
260
|
+
self.gnn.build(None)
|
|
261
|
+
if self.graph_pred_linear is not None:
|
|
262
|
+
mult = (self.num_layer + 1) if self.JK == "concat" else 1
|
|
263
|
+
self.graph_pred_linear.build((None, mult * self.emb_dim))
|
|
264
|
+
super().build(input_shape)
|
|
265
|
+
|
|
266
|
+
def call(self, inputs, training: bool = False):
|
|
267
|
+
"""Call MoleBERT model.
|
|
268
|
+
|
|
269
|
+
inputs can be a tuple: `(x, edge_index, edge_attr)` or `(x, edge_index, edge_attr, batch)`.
|
|
270
|
+
"""
|
|
271
|
+
if isinstance(inputs, (list, tuple)):
|
|
272
|
+
if len(inputs) == 3:
|
|
273
|
+
x, edge_index, edge_attr = inputs
|
|
274
|
+
batch = None
|
|
275
|
+
elif len(inputs) == 4:
|
|
276
|
+
x, edge_index, edge_attr, batch = inputs
|
|
277
|
+
else:
|
|
278
|
+
raise ValueError("Expected 3 or 4 input tensors.")
|
|
279
|
+
elif isinstance(inputs, dict):
|
|
280
|
+
x = inputs["x"]
|
|
281
|
+
edge_index = inputs["edge_index"]
|
|
282
|
+
edge_attr = inputs["edge_attr"]
|
|
283
|
+
batch = inputs.get("batch", None)
|
|
284
|
+
else:
|
|
285
|
+
raise ValueError("inputs must be a tuple, list, or dict.")
|
|
286
|
+
|
|
287
|
+
node_rep = self.gnn(x, edge_index, edge_attr, training=training)
|
|
288
|
+
|
|
289
|
+
if batch is not None:
|
|
290
|
+
graph_rep = self.pool_fn(node_rep, batch)
|
|
291
|
+
if self.graph_pred_linear is not None:
|
|
292
|
+
logits = self.graph_pred_linear(graph_rep)
|
|
293
|
+
return logits, node_rep
|
|
294
|
+
return graph_rep, node_rep
|
|
295
|
+
|
|
296
|
+
if self.graph_pred_linear is not None:
|
|
297
|
+
# If no batch given, default all nodes to batch 0
|
|
298
|
+
num_nodes = ops.shape(node_rep)[0]
|
|
299
|
+
batch_zeros = ops.zeros((num_nodes,), dtype="int32")
|
|
300
|
+
graph_rep = self.pool_fn(node_rep, batch_zeros)
|
|
301
|
+
logits = self.graph_pred_linear(graph_rep)
|
|
302
|
+
return logits, node_rep
|
|
303
|
+
|
|
304
|
+
return node_rep
|
|
305
|
+
|
|
306
|
+
|
|
307
|
+
def load_mole_bert_weights(model: Union[MoleBERT, MoleBERTGNN], checkpoint_path: str):
|
|
308
|
+
"""Loads official PyTorch Mole-BERT.pth checkpoint state dict into Keras 3 MoleBERT model."""
|
|
309
|
+
import torch
|
|
310
|
+
|
|
311
|
+
try:
|
|
312
|
+
state_dict = torch.load(checkpoint_path, map_location="cpu", weights_only=False)
|
|
313
|
+
except Exception:
|
|
314
|
+
state_dict = torch.load(checkpoint_path, map_location="cpu")
|
|
315
|
+
|
|
316
|
+
if "state_dict" in state_dict:
|
|
317
|
+
state_dict = state_dict["state_dict"]
|
|
318
|
+
|
|
319
|
+
def get_np(key):
|
|
320
|
+
return state_dict[key].detach().cpu().float().numpy()
|
|
321
|
+
|
|
322
|
+
if not model.built:
|
|
323
|
+
model.build(None)
|
|
324
|
+
|
|
325
|
+
gnn = model.gnn if hasattr(model, "gnn") else model
|
|
326
|
+
|
|
327
|
+
# Embeddings
|
|
328
|
+
gnn.x_embedding1.set_weights([get_np("x_embedding1.weight")])
|
|
329
|
+
gnn.x_embedding2.set_weights([get_np("x_embedding2.weight")])
|
|
330
|
+
|
|
331
|
+
for l in range(gnn.num_layer):
|
|
332
|
+
conv = gnn.gnns[l]
|
|
333
|
+
bn = gnn.batch_norms[l]
|
|
334
|
+
|
|
335
|
+
# GINConv edge embeddings
|
|
336
|
+
conv.edge_embedding1.set_weights([get_np(f"gnns.{l}.edge_embedding1.weight")])
|
|
337
|
+
conv.edge_embedding2.set_weights([get_np(f"gnns.{l}.edge_embedding2.weight")])
|
|
338
|
+
|
|
339
|
+
# GINConv MLP layers
|
|
340
|
+
w0 = get_np(f"gnns.{l}.mlp.0.weight").T
|
|
341
|
+
b0 = get_np(f"gnns.{l}.mlp.0.bias")
|
|
342
|
+
conv.mlp.layers[0].set_weights([w0, b0])
|
|
343
|
+
|
|
344
|
+
w2 = get_np(f"gnns.{l}.mlp.2.weight").T
|
|
345
|
+
b2 = get_np(f"gnns.{l}.mlp.2.bias")
|
|
346
|
+
conv.mlp.layers[1].set_weights([w2, b2])
|
|
347
|
+
|
|
348
|
+
# Batch Normalization
|
|
349
|
+
gamma = get_np(f"batch_norms.{l}.weight")
|
|
350
|
+
beta = get_np(f"batch_norms.{l}.bias")
|
|
351
|
+
mean = get_np(f"batch_norms.{l}.running_mean")
|
|
352
|
+
var = get_np(f"batch_norms.{l}.running_var")
|
|
353
|
+
bn.set_weights([gamma, beta, mean, var])
|
|
354
|
+
|
|
355
|
+
|
|
356
|
+
def download_mole_bert_checkpoint(cache_dir: Optional[str] = None) -> str:
|
|
357
|
+
"""Downloads official Mole-BERT.pth checkpoint from GitHub."""
|
|
358
|
+
if cache_dir is None:
|
|
359
|
+
cache_dir = osp.expanduser("~/.cache/k3_node/mole_bert")
|
|
360
|
+
|
|
361
|
+
os.makedirs(cache_dir, exist_ok=True)
|
|
362
|
+
target_path = osp.join(cache_dir, "Mole-BERT.pth")
|
|
363
|
+
|
|
364
|
+
if osp.exists(target_path) and osp.getsize(target_path) > 1000:
|
|
365
|
+
return target_path
|
|
366
|
+
|
|
367
|
+
# Check local path in Mole-BERT/model_gin/Mole-BERT.pth
|
|
368
|
+
local_path = osp.join(
|
|
369
|
+
osp.dirname(osp.dirname(osp.dirname(osp.abspath(__file__)))),
|
|
370
|
+
"Mole-BERT",
|
|
371
|
+
"model_gin",
|
|
372
|
+
"Mole-BERT.pth",
|
|
373
|
+
)
|
|
374
|
+
if osp.exists(local_path) and osp.getsize(local_path) > 1000:
|
|
375
|
+
shutil.copyfile(local_path, target_path)
|
|
376
|
+
return target_path
|
|
377
|
+
|
|
378
|
+
return download_url(MOLE_BERT_URL, cache_dir, filename="Mole-BERT.pth")
|
|
379
|
+
|
|
@@ -0,0 +1,95 @@
|
|
|
1
|
+
from typing import Optional
|
|
2
|
+
|
|
3
|
+
import keras
|
|
4
|
+
from keras import ops
|
|
5
|
+
|
|
6
|
+
from k3_node.layers.conv import MFConv
|
|
7
|
+
from k3_node.layers.pool import global_add_pool
|
|
8
|
+
|
|
9
|
+
|
|
10
|
+
class NeuralFingerprint(keras.layers.Layer):
|
|
11
|
+
r"""The Neural Fingerprint model from the
|
|
12
|
+
`"Convolutional Networks on Graphs for Learning Molecular Fingerprints"
|
|
13
|
+
<https://arxiv.org/abs/1509.09292>`__ paper to generate fingerprints
|
|
14
|
+
of molecules.
|
|
15
|
+
|
|
16
|
+
Args:
|
|
17
|
+
in_channels (int): Size of each input sample.
|
|
18
|
+
hidden_channels (int): Size of each hidden sample.
|
|
19
|
+
out_channels (int): Size of each output fingerprint.
|
|
20
|
+
num_layers (int): Number of layers.
|
|
21
|
+
**kwargs (optional): Additional arguments of
|
|
22
|
+
:class:`~k3_node.layers.conv.MFConv`.
|
|
23
|
+
|
|
24
|
+
Example:
|
|
25
|
+
```python
|
|
26
|
+
import numpy as np
|
|
27
|
+
from k3_node.models import NeuralFingerprint
|
|
28
|
+
|
|
29
|
+
x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
|
|
30
|
+
edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
|
|
31
|
+
|
|
32
|
+
batch = np.repeat([0, 1], 5) # two molecules with 5 atoms each
|
|
33
|
+
model = NeuralFingerprint(in_channels=8, hidden_channels=32, out_channels=16, num_layers=3)
|
|
34
|
+
fingerprint = model(x, edge_index, batch) # one learned fingerprint per molecule
|
|
35
|
+
print(tuple(fingerprint.shape)) # (2, 16)
|
|
36
|
+
```
|
|
37
|
+
"""
|
|
38
|
+
def __init__(
|
|
39
|
+
self,
|
|
40
|
+
in_channels: int,
|
|
41
|
+
hidden_channels: int,
|
|
42
|
+
out_channels: int,
|
|
43
|
+
num_layers: int,
|
|
44
|
+
**kwargs,
|
|
45
|
+
):
|
|
46
|
+
super().__init__()
|
|
47
|
+
|
|
48
|
+
self.in_channels = in_channels
|
|
49
|
+
self.hidden_channels = hidden_channels
|
|
50
|
+
self.out_channels = out_channels
|
|
51
|
+
self.num_layers = num_layers
|
|
52
|
+
|
|
53
|
+
self.convs = []
|
|
54
|
+
for i in range(self.num_layers):
|
|
55
|
+
in_c = self.in_channels if i == 0 else self.hidden_channels
|
|
56
|
+
self.convs.append(MFConv(in_c, hidden_channels, **kwargs))
|
|
57
|
+
|
|
58
|
+
self.lins = []
|
|
59
|
+
for _ in range(self.num_layers):
|
|
60
|
+
self.lins.append(keras.layers.Dense(out_channels, use_bias=False))
|
|
61
|
+
|
|
62
|
+
def build(self, input_shape=None):
|
|
63
|
+
self.built = True
|
|
64
|
+
|
|
65
|
+
def reset_parameters(self):
|
|
66
|
+
r"""Resets all learnable parameters of the module."""
|
|
67
|
+
for conv in self.convs:
|
|
68
|
+
if hasattr(conv, "reset_parameters"):
|
|
69
|
+
conv.reset_parameters()
|
|
70
|
+
for lin in self.lins:
|
|
71
|
+
if lin.built:
|
|
72
|
+
lin.kernel.assign(keras.initializers.GlorotUniform()(lin.kernel.shape))
|
|
73
|
+
|
|
74
|
+
def call(
|
|
75
|
+
self,
|
|
76
|
+
x,
|
|
77
|
+
edge_index,
|
|
78
|
+
batch: Optional[any] = None,
|
|
79
|
+
batch_size: Optional[int] = None,
|
|
80
|
+
):
|
|
81
|
+
outs = []
|
|
82
|
+
for conv, lin in zip(self.convs, self.lins):
|
|
83
|
+
x = ops.sigmoid(conv(x, edge_index))
|
|
84
|
+
y = ops.softmax(lin(x), axis=-1)
|
|
85
|
+
outs.append(global_add_pool(y, batch, size=batch_size))
|
|
86
|
+
|
|
87
|
+
out = outs[0]
|
|
88
|
+
for item in outs[1:]:
|
|
89
|
+
out = out + item
|
|
90
|
+
return out
|
|
91
|
+
|
|
92
|
+
def __repr__(self) -> str:
|
|
93
|
+
return (f'{self.__class__.__name__}({self.in_channels}, '
|
|
94
|
+
f'{self.out_channels}, num_layers={self.num_layers})')
|
|
95
|
+
|
|
@@ -0,0 +1,213 @@
|
|
|
1
|
+
from typing import Optional
|
|
2
|
+
import numpy as np
|
|
3
|
+
import keras
|
|
4
|
+
from keras import ops
|
|
5
|
+
|
|
6
|
+
|
|
7
|
+
class Node2Vec(keras.layers.Layer):
|
|
8
|
+
r"""The Node2Vec model from the
|
|
9
|
+
`"node2vec: Scalable Feature Learning for Networks"
|
|
10
|
+
<https://arxiv.org/abs/1607.00653>`_ paper where random walks of
|
|
11
|
+
length :obj:`walk_length` are sampled in a given graph, and node embeddings
|
|
12
|
+
are learned via negative sampling optimization.
|
|
13
|
+
|
|
14
|
+
Args:
|
|
15
|
+
edge_index: The edge indices.
|
|
16
|
+
embedding_dim (int): The size of each embedding vector.
|
|
17
|
+
walk_length (int): The walk length.
|
|
18
|
+
context_size (int): The actual context size which is considered for
|
|
19
|
+
positive samples.
|
|
20
|
+
walks_per_node (int, optional): The number of walks to sample for each
|
|
21
|
+
node. (default: :obj:`1`)
|
|
22
|
+
p (float, optional): Likelihood of immediately revisiting a node in the
|
|
23
|
+
walk. (default: :obj:`1.0`)
|
|
24
|
+
q (float, optional): Control parameter to interpolate between
|
|
25
|
+
breadth-first strategy and depth-first strategy. (default: :obj:`1.0`)
|
|
26
|
+
num_negative_samples (int, optional): The number of negative samples to
|
|
27
|
+
use for each positive sample. (default: :obj:`1`)
|
|
28
|
+
num_nodes (int, optional): The number of nodes. (default: :obj:`None`)
|
|
29
|
+
|
|
30
|
+
Example:
|
|
31
|
+
```python
|
|
32
|
+
import numpy as np
|
|
33
|
+
from k3_node.models import Node2Vec
|
|
34
|
+
|
|
35
|
+
edge_index = np.array([[0, 1, 2, 3, 0, 2], [1, 2, 3, 0, 2, 0]])
|
|
36
|
+
model = Node2Vec(edge_index, embedding_dim=16, walk_length=4, context_size=3, walks_per_node=2)
|
|
37
|
+
print(tuple(model().shape)) # (4, 16): embeddings of all nodes
|
|
38
|
+
|
|
39
|
+
batch = np.array([0, 1])
|
|
40
|
+
loss = model.loss(model.pos_sample(batch), model.neg_sample(batch)) # skip-gram loss on random walks
|
|
41
|
+
print(tuple(loss.shape)) # (): a scalar
|
|
42
|
+
```
|
|
43
|
+
"""
|
|
44
|
+
def __init__(
|
|
45
|
+
self,
|
|
46
|
+
edge_index,
|
|
47
|
+
embedding_dim: int,
|
|
48
|
+
walk_length: int,
|
|
49
|
+
context_size: int,
|
|
50
|
+
walks_per_node: int = 1,
|
|
51
|
+
p: float = 1.0,
|
|
52
|
+
q: float = 1.0,
|
|
53
|
+
num_negative_samples: int = 1,
|
|
54
|
+
num_nodes: Optional[int] = None,
|
|
55
|
+
**kwargs,
|
|
56
|
+
):
|
|
57
|
+
super().__init__(**kwargs)
|
|
58
|
+
|
|
59
|
+
edge_index_np = ops.convert_to_numpy(edge_index).astype(np.int64)
|
|
60
|
+
if num_nodes is None:
|
|
61
|
+
num_nodes = int(edge_index_np.max()) + 1 if edge_index_np.size > 0 else 0
|
|
62
|
+
|
|
63
|
+
self.num_nodes = num_nodes
|
|
64
|
+
self.embedding_dim = embedding_dim
|
|
65
|
+
self.walk_length = walk_length - 1
|
|
66
|
+
self.context_size = context_size
|
|
67
|
+
self.walks_per_node = walks_per_node
|
|
68
|
+
self.p = p
|
|
69
|
+
self.q = q
|
|
70
|
+
self.num_negative_samples = num_negative_samples
|
|
71
|
+
self.EPS = 1e-15
|
|
72
|
+
|
|
73
|
+
# Build adjacency list for fast random walks
|
|
74
|
+
self.adj = [[] for _ in range(num_nodes)]
|
|
75
|
+
if edge_index_np.size > 0:
|
|
76
|
+
for src, dst in zip(edge_index_np[0], edge_index_np[1]):
|
|
77
|
+
self.adj[src].append(int(dst))
|
|
78
|
+
self.adj_sets = [set(nbrs) for nbrs in self.adj]
|
|
79
|
+
|
|
80
|
+
self.embedding = keras.layers.Embedding(num_nodes, embedding_dim)
|
|
81
|
+
|
|
82
|
+
def reset_parameters(self):
|
|
83
|
+
if self.embedding.built:
|
|
84
|
+
self.embedding.embeddings.assign(
|
|
85
|
+
keras.initializers.GlorotUniform()(self.embedding.embeddings.shape)
|
|
86
|
+
)
|
|
87
|
+
|
|
88
|
+
def call(self, batch: Optional[any] = None):
|
|
89
|
+
"""Returns the embeddings for the nodes in :obj:`batch`."""
|
|
90
|
+
if batch is None:
|
|
91
|
+
batch = ops.arange(self.num_nodes, dtype="int64")
|
|
92
|
+
return self.embedding(batch)
|
|
93
|
+
|
|
94
|
+
def pos_sample(self, batch):
|
|
95
|
+
batch_np = ops.convert_to_numpy(batch).astype(np.int64)
|
|
96
|
+
repeated = np.repeat(batch_np, self.walks_per_node)
|
|
97
|
+
|
|
98
|
+
# Sample random walks
|
|
99
|
+
all_walks = [self._random_walk(int(node)) for node in repeated]
|
|
100
|
+
|
|
101
|
+
rw = np.array(all_walks, dtype=np.int64)
|
|
102
|
+
walks = []
|
|
103
|
+
num_walks_per_rw = 1 + self.walk_length + 1 - self.context_size
|
|
104
|
+
for j in range(num_walks_per_rw):
|
|
105
|
+
walks.append(rw[:, j : j + self.context_size])
|
|
106
|
+
out = np.concatenate(walks, axis=0) if len(walks) > 0 else rw
|
|
107
|
+
return ops.convert_to_tensor(out, dtype="int64")
|
|
108
|
+
|
|
109
|
+
def _random_walk(self, start: int):
|
|
110
|
+
"""A node2vec walk: from ``cur`` (reached from ``prev``) the next node is weighted by 1/p
|
|
111
|
+
if it returns to ``prev``, 1 if it is also a neighbor of ``prev`` and 1/q otherwise."""
|
|
112
|
+
walk = [start]
|
|
113
|
+
for _ in range(self.walk_length):
|
|
114
|
+
cur = walk[-1]
|
|
115
|
+
nbrs = self.adj[cur]
|
|
116
|
+
if not nbrs:
|
|
117
|
+
walk.append(cur)
|
|
118
|
+
continue
|
|
119
|
+
if (self.p == 1.0 and self.q == 1.0) or len(walk) == 1:
|
|
120
|
+
walk.append(nbrs[np.random.randint(len(nbrs))])
|
|
121
|
+
continue
|
|
122
|
+
prev = walk[-2]
|
|
123
|
+
prev_nbrs = self.adj_sets[prev]
|
|
124
|
+
weights = np.array([1.0 / self.p if n == prev else (1.0 if n in prev_nbrs else 1.0 / self.q)
|
|
125
|
+
for n in nbrs])
|
|
126
|
+
walk.append(nbrs[np.random.choice(len(nbrs), p=weights / weights.sum())])
|
|
127
|
+
return walk
|
|
128
|
+
|
|
129
|
+
def loader(self, batch_size: int = 128, shuffle: bool = True):
|
|
130
|
+
r"""Yields ``(pos_rw, neg_rw)`` random-walk batches, starting from ``batch_size`` nodes
|
|
131
|
+
at a time, as PyG's ``Node2Vec.loader``."""
|
|
132
|
+
nodes = np.random.permutation(self.num_nodes) if shuffle else np.arange(self.num_nodes)
|
|
133
|
+
for start in range(0, self.num_nodes, batch_size):
|
|
134
|
+
batch = nodes[start:start + batch_size]
|
|
135
|
+
yield self.pos_sample(batch), self.neg_sample(batch)
|
|
136
|
+
|
|
137
|
+
def compile(self, optimizer):
|
|
138
|
+
r"""Sets the optimizer used by :meth:`fit`."""
|
|
139
|
+
self.optimizer = optimizer
|
|
140
|
+
|
|
141
|
+
def fit(self, epochs: int = 1, batch_size: int = 128, verbose: int = 1):
|
|
142
|
+
r"""Trains the embeddings on freshly sampled random walks for ``epochs`` passes over
|
|
143
|
+
all nodes; returns the mean loss of every epoch."""
|
|
144
|
+
from k3_node.training import gradient_step
|
|
145
|
+
|
|
146
|
+
if getattr(self, "optimizer", None) is None:
|
|
147
|
+
raise ValueError("Call `compile(optimizer=...)` before `fit`.")
|
|
148
|
+
self(ops.arange(1)) # create the embeddings
|
|
149
|
+
history = {"loss": []}
|
|
150
|
+
for epoch in range(1, epochs + 1):
|
|
151
|
+
losses = [gradient_step(lambda: self.loss(pos_rw, neg_rw), self.trainable_variables, self.optimizer)
|
|
152
|
+
for pos_rw, neg_rw in self.loader(batch_size, shuffle=True)]
|
|
153
|
+
history["loss"].append(float(np.mean(losses)))
|
|
154
|
+
if verbose:
|
|
155
|
+
print(f"Epoch {epoch:03d}: loss: {history['loss'][-1]:.4f}")
|
|
156
|
+
return history
|
|
157
|
+
|
|
158
|
+
def test(self, train_z, train_y, test_z, test_y, solver: str = "lbfgs", *args, **kwargs):
|
|
159
|
+
r"""Evaluates the embeddings with a logistic regression classifier; returns its accuracy."""
|
|
160
|
+
from sklearn.linear_model import LogisticRegression
|
|
161
|
+
|
|
162
|
+
clf = LogisticRegression(*args, solver=solver, **kwargs).fit(
|
|
163
|
+
ops.convert_to_numpy(train_z), ops.convert_to_numpy(train_y))
|
|
164
|
+
return clf.score(ops.convert_to_numpy(test_z), ops.convert_to_numpy(test_y))
|
|
165
|
+
|
|
166
|
+
def neg_sample(self, batch):
|
|
167
|
+
batch_np = ops.convert_to_numpy(batch).astype(np.int64)
|
|
168
|
+
repeated = np.repeat(
|
|
169
|
+
batch_np, self.walks_per_node * self.num_negative_samples
|
|
170
|
+
)
|
|
171
|
+
rand_rest = np.random.randint(
|
|
172
|
+
0, max(self.num_nodes, 1), size=(len(repeated), self.walk_length)
|
|
173
|
+
)
|
|
174
|
+
rw = np.concatenate([repeated[:, None], rand_rest], axis=-1)
|
|
175
|
+
|
|
176
|
+
walks = []
|
|
177
|
+
num_walks_per_rw = 1 + self.walk_length + 1 - self.context_size
|
|
178
|
+
for j in range(num_walks_per_rw):
|
|
179
|
+
walks.append(rw[:, j : j + self.context_size])
|
|
180
|
+
out = np.concatenate(walks, axis=0) if len(walks) > 0 else rw
|
|
181
|
+
return ops.convert_to_tensor(out, dtype="int64")
|
|
182
|
+
|
|
183
|
+
def loss(self, pos_rw, neg_rw):
|
|
184
|
+
r"""Computes the loss given positive and negative random walks."""
|
|
185
|
+
# Positive loss
|
|
186
|
+
start = pos_rw[:, 0]
|
|
187
|
+
rest = pos_rw[:, 1:]
|
|
188
|
+
pos_b = ops.shape(pos_rw)[0]
|
|
189
|
+
|
|
190
|
+
h_start = ops.reshape(self.embedding(start), (pos_b, 1, self.embedding_dim))
|
|
191
|
+
h_rest = ops.reshape(
|
|
192
|
+
self.embedding(ops.reshape(rest, (-1,))),
|
|
193
|
+
(pos_b, -1, self.embedding_dim),
|
|
194
|
+
)
|
|
195
|
+
|
|
196
|
+
out = ops.reshape(ops.sum(h_start * h_rest, axis=-1), (-1,))
|
|
197
|
+
pos_loss = -ops.mean(ops.log(ops.sigmoid(out) + self.EPS))
|
|
198
|
+
|
|
199
|
+
# Negative loss
|
|
200
|
+
start = neg_rw[:, 0]
|
|
201
|
+
rest = neg_rw[:, 1:]
|
|
202
|
+
neg_b = ops.shape(neg_rw)[0]
|
|
203
|
+
|
|
204
|
+
h_start = ops.reshape(self.embedding(start), (neg_b, 1, self.embedding_dim))
|
|
205
|
+
h_rest = ops.reshape(
|
|
206
|
+
self.embedding(ops.reshape(rest, (-1,))),
|
|
207
|
+
(neg_b, -1, self.embedding_dim),
|
|
208
|
+
)
|
|
209
|
+
|
|
210
|
+
out = ops.reshape(ops.sum(h_start * h_rest, axis=-1), (-1,))
|
|
211
|
+
neg_loss = -ops.mean(ops.log(1.0 - ops.sigmoid(out) + self.EPS))
|
|
212
|
+
|
|
213
|
+
return pos_loss + neg_loss
|