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,1122 @@
|
|
|
1
|
+
import os
|
|
2
|
+
import math
|
|
3
|
+
import zipfile
|
|
4
|
+
import urllib.request
|
|
5
|
+
from typing import Optional, List, Union, Tuple, Dict, Any
|
|
6
|
+
|
|
7
|
+
import numpy as np
|
|
8
|
+
import keras
|
|
9
|
+
from keras import layers, ops
|
|
10
|
+
|
|
11
|
+
from k3_node.layers.pool import global_add_pool, global_max_pool, global_mean_pool
|
|
12
|
+
from k3_node.ops.segment import segment_sum
|
|
13
|
+
|
|
14
|
+
|
|
15
|
+
def _get_act(act_str: str):
|
|
16
|
+
act_str = act_str.lower()
|
|
17
|
+
if act_str in ["gelu", "quick_gelu"]:
|
|
18
|
+
return ops.gelu
|
|
19
|
+
elif act_str == "relu":
|
|
20
|
+
return ops.relu
|
|
21
|
+
elif act_str == "silu" or act_str == "swish":
|
|
22
|
+
return ops.silu
|
|
23
|
+
elif act_str == "tanh":
|
|
24
|
+
return ops.tanh
|
|
25
|
+
elif act_str == "sigmoid":
|
|
26
|
+
return ops.sigmoid
|
|
27
|
+
return ops.relu
|
|
28
|
+
|
|
29
|
+
|
|
30
|
+
# 2022 OGB molecule atom feature dimensions (used in GraphGPS pretraining on PCQM4Mv2)
|
|
31
|
+
DEFAULT_ATOM_FEATURE_DIMS = [119, 4, 12, 12, 10, 6, 6, 2, 2]
|
|
32
|
+
DEFAULT_BOND_FEATURE_DIMS = [5, 6, 2]
|
|
33
|
+
|
|
34
|
+
|
|
35
|
+
class AtomEncoder(layers.Layer):
|
|
36
|
+
r"""OGB Molecule categorical atom feature encoder.
|
|
37
|
+
|
|
38
|
+
Args:
|
|
39
|
+
emb_dim (int): Output embedding dimension.
|
|
40
|
+
feature_dims (List[int], optional): Categorical feature vocabulary sizes for each
|
|
41
|
+
atom feature column. (default: ``[119, 4, 12, 12, 10, 6, 6, 2, 2]``)
|
|
42
|
+
**kwargs: Additional layer arguments.
|
|
43
|
+
|
|
44
|
+
Example:
|
|
45
|
+
```python
|
|
46
|
+
import numpy as np
|
|
47
|
+
from k3_node.models import AtomEncoder
|
|
48
|
+
|
|
49
|
+
# OGB-style integer atom and bond features (9 per atom, 3 per bond)
|
|
50
|
+
x = np.random.randint(0, 2, size=(5, 9))
|
|
51
|
+
edge_index = np.array([[0, 1, 1, 2, 3, 4], [1, 0, 2, 1, 4, 3]])
|
|
52
|
+
edge_attr = np.random.randint(0, 2, size=(6, 3))
|
|
53
|
+
batch = np.array([0, 0, 0, 1, 1]) # two graphs
|
|
54
|
+
|
|
55
|
+
print(tuple(AtomEncoder(emb_dim=32)(x).shape)) # (5, 32): sum of per-feature embeddings
|
|
56
|
+
```
|
|
57
|
+
"""
|
|
58
|
+
|
|
59
|
+
def __init__(
|
|
60
|
+
self,
|
|
61
|
+
emb_dim: int,
|
|
62
|
+
feature_dims: Optional[List[int]] = None,
|
|
63
|
+
**kwargs,
|
|
64
|
+
):
|
|
65
|
+
super().__init__(**kwargs)
|
|
66
|
+
self.emb_dim = emb_dim
|
|
67
|
+
self.feature_dims = feature_dims if feature_dims is not None else DEFAULT_ATOM_FEATURE_DIMS
|
|
68
|
+
|
|
69
|
+
self.atom_embedding_list = [
|
|
70
|
+
layers.Embedding(
|
|
71
|
+
input_dim=dim,
|
|
72
|
+
output_dim=emb_dim,
|
|
73
|
+
embeddings_initializer="glorot_uniform",
|
|
74
|
+
name=f"atom_embedding_{i}",
|
|
75
|
+
)
|
|
76
|
+
for i, dim in enumerate(self.feature_dims)
|
|
77
|
+
]
|
|
78
|
+
|
|
79
|
+
def build(self, input_shape=None):
|
|
80
|
+
for emb in self.atom_embedding_list:
|
|
81
|
+
if not emb.built:
|
|
82
|
+
emb.build(None)
|
|
83
|
+
super().build(input_shape)
|
|
84
|
+
|
|
85
|
+
def call(self, x):
|
|
86
|
+
x = ops.cast(x, "int32")
|
|
87
|
+
out = 0
|
|
88
|
+
for i, emb in enumerate(self.atom_embedding_list):
|
|
89
|
+
feat = x[:, i]
|
|
90
|
+
out = out + emb(feat)
|
|
91
|
+
return out
|
|
92
|
+
|
|
93
|
+
|
|
94
|
+
class BondEncoder(layers.Layer):
|
|
95
|
+
r"""OGB Molecule categorical bond feature encoder.
|
|
96
|
+
|
|
97
|
+
Args:
|
|
98
|
+
emb_dim (int): Output embedding dimension.
|
|
99
|
+
feature_dims (List[int], optional): Categorical feature vocabulary sizes for each
|
|
100
|
+
bond feature column. (default: ``[5, 6, 2]``)
|
|
101
|
+
**kwargs: Additional layer arguments.
|
|
102
|
+
|
|
103
|
+
Example:
|
|
104
|
+
```python
|
|
105
|
+
import numpy as np
|
|
106
|
+
from k3_node.models import BondEncoder
|
|
107
|
+
|
|
108
|
+
# OGB-style integer atom and bond features (9 per atom, 3 per bond)
|
|
109
|
+
x = np.random.randint(0, 2, size=(5, 9))
|
|
110
|
+
edge_index = np.array([[0, 1, 1, 2, 3, 4], [1, 0, 2, 1, 4, 3]])
|
|
111
|
+
edge_attr = np.random.randint(0, 2, size=(6, 3))
|
|
112
|
+
batch = np.array([0, 0, 0, 1, 1]) # two graphs
|
|
113
|
+
|
|
114
|
+
print(tuple(BondEncoder(emb_dim=32)(edge_attr).shape)) # (6, 32)
|
|
115
|
+
```
|
|
116
|
+
"""
|
|
117
|
+
|
|
118
|
+
def __init__(
|
|
119
|
+
self,
|
|
120
|
+
emb_dim: int,
|
|
121
|
+
feature_dims: Optional[List[int]] = None,
|
|
122
|
+
**kwargs,
|
|
123
|
+
):
|
|
124
|
+
super().__init__(**kwargs)
|
|
125
|
+
self.emb_dim = emb_dim
|
|
126
|
+
self.feature_dims = feature_dims if feature_dims is not None else DEFAULT_BOND_FEATURE_DIMS
|
|
127
|
+
|
|
128
|
+
self.bond_embedding_list = [
|
|
129
|
+
layers.Embedding(
|
|
130
|
+
input_dim=dim,
|
|
131
|
+
output_dim=emb_dim,
|
|
132
|
+
embeddings_initializer="glorot_uniform",
|
|
133
|
+
name=f"bond_embedding_{i}",
|
|
134
|
+
)
|
|
135
|
+
for i, dim in enumerate(self.feature_dims)
|
|
136
|
+
]
|
|
137
|
+
|
|
138
|
+
def build(self, input_shape=None):
|
|
139
|
+
for emb in self.bond_embedding_list:
|
|
140
|
+
if not emb.built:
|
|
141
|
+
emb.build(None)
|
|
142
|
+
super().build(input_shape)
|
|
143
|
+
|
|
144
|
+
def call(self, edge_attr):
|
|
145
|
+
edge_attr = ops.cast(edge_attr, "int32")
|
|
146
|
+
out = 0
|
|
147
|
+
for i, emb in enumerate(self.bond_embedding_list):
|
|
148
|
+
feat = edge_attr[:, i]
|
|
149
|
+
out = out + emb(feat)
|
|
150
|
+
return out
|
|
151
|
+
|
|
152
|
+
|
|
153
|
+
class RWSEEncoder(layers.Layer):
|
|
154
|
+
r"""Random Walk Structural Encoding (RWSE) node encoder.
|
|
155
|
+
|
|
156
|
+
Normalizes the precomputed $k$-step diagonal random walk landing probabilities
|
|
157
|
+
using Batch Normalization and projects them into `pe_dim` dimension.
|
|
158
|
+
|
|
159
|
+
Args:
|
|
160
|
+
num_rw_steps (int): Number of random walk steps. (default: ``16``)
|
|
161
|
+
pe_dim (int): Output structural encoding dimension. (default: ``20``)
|
|
162
|
+
**kwargs: Additional layer arguments.
|
|
163
|
+
|
|
164
|
+
Example:
|
|
165
|
+
```python
|
|
166
|
+
import numpy as np
|
|
167
|
+
from k3_node.models import RWSEEncoder
|
|
168
|
+
|
|
169
|
+
rwse = np.random.rand(5, 16).astype("float32") # 16-step random-walk return probabilities
|
|
170
|
+
print(tuple(RWSEEncoder(num_rw_steps=16, pe_dim=8)(rwse).shape)) # (5, 8)
|
|
171
|
+
```
|
|
172
|
+
"""
|
|
173
|
+
|
|
174
|
+
def __init__(
|
|
175
|
+
self,
|
|
176
|
+
num_rw_steps: int = 16,
|
|
177
|
+
pe_dim: int = 20,
|
|
178
|
+
**kwargs,
|
|
179
|
+
):
|
|
180
|
+
super().__init__(**kwargs)
|
|
181
|
+
self.num_rw_steps = num_rw_steps
|
|
182
|
+
self.pe_dim = pe_dim
|
|
183
|
+
|
|
184
|
+
self.raw_norm = layers.BatchNormalization(
|
|
185
|
+
axis=-1,
|
|
186
|
+
epsilon=1e-5,
|
|
187
|
+
momentum=0.9,
|
|
188
|
+
name="raw_norm",
|
|
189
|
+
)
|
|
190
|
+
self.pe_encoder = layers.Dense(pe_dim, use_bias=True, name="pe_encoder")
|
|
191
|
+
|
|
192
|
+
def build(self, input_shape=None):
|
|
193
|
+
if not self.raw_norm.built:
|
|
194
|
+
self.raw_norm.build((None, self.num_rw_steps))
|
|
195
|
+
if not self.pe_encoder.built:
|
|
196
|
+
self.pe_encoder.build((None, self.num_rw_steps))
|
|
197
|
+
super().build(input_shape)
|
|
198
|
+
|
|
199
|
+
def call(self, pestat_RWSE, training=False):
|
|
200
|
+
pe = self.raw_norm(pestat_RWSE, training=training)
|
|
201
|
+
return self.pe_encoder(pe)
|
|
202
|
+
|
|
203
|
+
|
|
204
|
+
class CustomGatedGCN(layers.Layer):
|
|
205
|
+
r"""Residual Gated Graph ConvNet layer with edge feature updates.
|
|
206
|
+
|
|
207
|
+
Reference:
|
|
208
|
+
"Residual Gated Graph ConvNets" (Bresson & Laurent, 2017).
|
|
209
|
+
|
|
210
|
+
Args:
|
|
211
|
+
in_dim (int): Input feature dimension.
|
|
212
|
+
out_dim (int): Output feature dimension.
|
|
213
|
+
dropout (float, optional): Dropout rate. (default: ``0.0``)
|
|
214
|
+
residual (bool, optional): Whether to use residual connections. (default: ``True``)
|
|
215
|
+
act (str, optional): Activation function. (default: ``"gelu"``)
|
|
216
|
+
**kwargs: Additional layer arguments.
|
|
217
|
+
|
|
218
|
+
Example:
|
|
219
|
+
```python
|
|
220
|
+
import numpy as np
|
|
221
|
+
from k3_node.models import CustomGatedGCN
|
|
222
|
+
|
|
223
|
+
x = np.random.rand(5, 32).astype("float32") # node features
|
|
224
|
+
e = np.random.rand(6, 32).astype("float32") # edge features
|
|
225
|
+
edge_index = np.array([[0, 1, 1, 2, 3, 4], [1, 0, 2, 1, 4, 3]])
|
|
226
|
+
batch = np.array([0, 0, 0, 1, 1]) # two graphs
|
|
227
|
+
|
|
228
|
+
layer = CustomGatedGCN(in_dim=32, out_dim=32, dropout=0.0, residual=True)
|
|
229
|
+
x_out, e_out = layer(x, edge_index, e) # updates node and edge features
|
|
230
|
+
print(tuple(x_out.shape), tuple(e_out.shape)) # (5, 32) (6, 32)
|
|
231
|
+
```
|
|
232
|
+
"""
|
|
233
|
+
|
|
234
|
+
def __init__(
|
|
235
|
+
self,
|
|
236
|
+
in_dim: int,
|
|
237
|
+
out_dim: int,
|
|
238
|
+
dropout: float = 0.0,
|
|
239
|
+
residual: bool = True,
|
|
240
|
+
act: str = "gelu",
|
|
241
|
+
**kwargs,
|
|
242
|
+
):
|
|
243
|
+
super().__init__(**kwargs)
|
|
244
|
+
self.in_dim = in_dim
|
|
245
|
+
self.out_dim = out_dim
|
|
246
|
+
self.dropout_rate = dropout
|
|
247
|
+
self.residual = residual
|
|
248
|
+
self.act_name = act
|
|
249
|
+
|
|
250
|
+
self.A = layers.Dense(out_dim, use_bias=True, name="A")
|
|
251
|
+
self.B = layers.Dense(out_dim, use_bias=True, name="B")
|
|
252
|
+
self.C = layers.Dense(out_dim, use_bias=True, name="C")
|
|
253
|
+
self.D = layers.Dense(out_dim, use_bias=True, name="D")
|
|
254
|
+
self.E = layers.Dense(out_dim, use_bias=True, name="E")
|
|
255
|
+
|
|
256
|
+
self.bn_node_x = layers.BatchNormalization(
|
|
257
|
+
axis=-1, epsilon=1e-5, momentum=0.9, name="bn_node_x"
|
|
258
|
+
)
|
|
259
|
+
self.bn_edge_e = layers.BatchNormalization(
|
|
260
|
+
axis=-1, epsilon=1e-5, momentum=0.9, name="bn_edge_e"
|
|
261
|
+
)
|
|
262
|
+
self.dropout_node = layers.Dropout(dropout)
|
|
263
|
+
self.dropout_edge = layers.Dropout(dropout)
|
|
264
|
+
|
|
265
|
+
def build(self, input_shape=None):
|
|
266
|
+
if not self.A.built:
|
|
267
|
+
self.A.build((None, self.in_dim))
|
|
268
|
+
self.B.build((None, self.in_dim))
|
|
269
|
+
self.C.build((None, self.in_dim))
|
|
270
|
+
self.D.build((None, self.in_dim))
|
|
271
|
+
self.E.build((None, self.in_dim))
|
|
272
|
+
self.bn_node_x.build((None, self.out_dim))
|
|
273
|
+
self.bn_edge_e.build((None, self.out_dim))
|
|
274
|
+
super().build(input_shape)
|
|
275
|
+
|
|
276
|
+
def call(self, x, edge_index, edge_attr, training=False):
|
|
277
|
+
x_in = x
|
|
278
|
+
e_in = edge_attr
|
|
279
|
+
|
|
280
|
+
Ax = self.A(x)
|
|
281
|
+
Bx = self.B(x)
|
|
282
|
+
Ce = self.C(edge_attr)
|
|
283
|
+
Dx = self.D(x)
|
|
284
|
+
Ex = self.E(x)
|
|
285
|
+
|
|
286
|
+
src = ops.cast(edge_index[0], "int32")
|
|
287
|
+
dst = ops.cast(edge_index[1], "int32")
|
|
288
|
+
|
|
289
|
+
Dx_i = ops.take(Dx, dst, axis=0)
|
|
290
|
+
Ex_j = ops.take(Ex, src, axis=0)
|
|
291
|
+
Bx_j = ops.take(Bx, src, axis=0)
|
|
292
|
+
|
|
293
|
+
e_ij = Dx_i + Ex_j + Ce
|
|
294
|
+
sigma_ij = ops.sigmoid(e_ij)
|
|
295
|
+
|
|
296
|
+
num_nodes = ops.shape(x)[0]
|
|
297
|
+
sum_sigma_x = segment_sum(sigma_ij * Bx_j, dst, num_segments=num_nodes)
|
|
298
|
+
sum_sigma = segment_sum(sigma_ij, dst, num_segments=num_nodes)
|
|
299
|
+
aggr_out = sum_sigma_x / (sum_sigma + 1e-6)
|
|
300
|
+
|
|
301
|
+
x_out = self.bn_node_x(Ax + aggr_out, training=training)
|
|
302
|
+
e_out = self.bn_edge_e(e_ij, training=training)
|
|
303
|
+
|
|
304
|
+
act_fn = _get_act(self.act_name)
|
|
305
|
+
x_out = act_fn(x_out)
|
|
306
|
+
e_out = act_fn(e_out)
|
|
307
|
+
|
|
308
|
+
x_out = self.dropout_node(x_out, training=training)
|
|
309
|
+
e_out = self.dropout_edge(e_out, training=training)
|
|
310
|
+
|
|
311
|
+
if self.residual:
|
|
312
|
+
x_out = x_in + x_out
|
|
313
|
+
e_out = e_in + e_out
|
|
314
|
+
|
|
315
|
+
return x_out, e_out
|
|
316
|
+
|
|
317
|
+
|
|
318
|
+
class TransformerSelfAttention(layers.Layer):
|
|
319
|
+
r"""Multi-Head Self-Attention layer matching PyTorch `nn.MultiheadAttention`.
|
|
320
|
+
|
|
321
|
+
Args:
|
|
322
|
+
embed_dim (int): Total embedding dimension.
|
|
323
|
+
num_heads (int): Number of attention heads.
|
|
324
|
+
dropout (float, optional): Attention dropout rate. (default: ``0.0``)
|
|
325
|
+
**kwargs: Additional layer arguments.
|
|
326
|
+
"""
|
|
327
|
+
|
|
328
|
+
def __init__(
|
|
329
|
+
self,
|
|
330
|
+
embed_dim: int,
|
|
331
|
+
num_heads: int,
|
|
332
|
+
dropout: float = 0.0,
|
|
333
|
+
**kwargs,
|
|
334
|
+
):
|
|
335
|
+
super().__init__(**kwargs)
|
|
336
|
+
self.supports_masking = True
|
|
337
|
+
self.embed_dim = embed_dim
|
|
338
|
+
self.num_heads = num_heads
|
|
339
|
+
self.head_dim = embed_dim // num_heads
|
|
340
|
+
self.scaling = 1.0 / math.sqrt(self.head_dim)
|
|
341
|
+
self.dropout_rate = dropout
|
|
342
|
+
|
|
343
|
+
self.q_proj = layers.Dense(embed_dim, use_bias=True, name="q_proj")
|
|
344
|
+
self.k_proj = layers.Dense(embed_dim, use_bias=True, name="k_proj")
|
|
345
|
+
self.v_proj = layers.Dense(embed_dim, use_bias=True, name="v_proj")
|
|
346
|
+
self.out_proj = layers.Dense(embed_dim, use_bias=True, name="out_proj")
|
|
347
|
+
self.dropout = layers.Dropout(dropout)
|
|
348
|
+
|
|
349
|
+
def build(self, input_shape=None):
|
|
350
|
+
if not self.q_proj.built:
|
|
351
|
+
self.q_proj.build((None, self.embed_dim))
|
|
352
|
+
self.k_proj.build((None, self.embed_dim))
|
|
353
|
+
self.v_proj.build((None, self.embed_dim))
|
|
354
|
+
self.out_proj.build((None, self.embed_dim))
|
|
355
|
+
super().build(input_shape)
|
|
356
|
+
|
|
357
|
+
def call(self, x, mask=None, training=False):
|
|
358
|
+
shape = ops.shape(x)
|
|
359
|
+
bsz, n_node = shape[0], shape[1]
|
|
360
|
+
|
|
361
|
+
q = self.q_proj(x)
|
|
362
|
+
k = self.k_proj(x)
|
|
363
|
+
v = self.v_proj(x)
|
|
364
|
+
|
|
365
|
+
q = ops.transpose(
|
|
366
|
+
ops.reshape(q, (bsz, n_node, self.num_heads, self.head_dim)),
|
|
367
|
+
(0, 2, 1, 3),
|
|
368
|
+
) * self.scaling
|
|
369
|
+
k = ops.transpose(
|
|
370
|
+
ops.reshape(k, (bsz, n_node, self.num_heads, self.head_dim)),
|
|
371
|
+
(0, 2, 1, 3),
|
|
372
|
+
)
|
|
373
|
+
v = ops.transpose(
|
|
374
|
+
ops.reshape(v, (bsz, n_node, self.num_heads, self.head_dim)),
|
|
375
|
+
(0, 2, 1, 3),
|
|
376
|
+
)
|
|
377
|
+
|
|
378
|
+
scores = ops.matmul(q, ops.transpose(k, (0, 1, 3, 2))) # [B, H, N, N]
|
|
379
|
+
|
|
380
|
+
if mask is not None:
|
|
381
|
+
# mask: [B, N] boolean tensor, True for valid tokens, False for padding
|
|
382
|
+
# PyTorch key_padding_mask: True where padding
|
|
383
|
+
key_padding_mask = ~mask
|
|
384
|
+
scores = ops.where(
|
|
385
|
+
ops.expand_dims(ops.expand_dims(key_padding_mask, axis=1), axis=2),
|
|
386
|
+
float("-inf"),
|
|
387
|
+
scores,
|
|
388
|
+
)
|
|
389
|
+
|
|
390
|
+
attn_weights = ops.softmax(scores, axis=-1)
|
|
391
|
+
# Replace NaNs from all-inf rows
|
|
392
|
+
attn_weights = ops.where(ops.isnan(attn_weights), 0.0, attn_weights)
|
|
393
|
+
attn_weights = self.dropout(attn_weights, training=training)
|
|
394
|
+
|
|
395
|
+
attn = ops.matmul(attn_weights, v) # [B, H, N, head_dim]
|
|
396
|
+
attn = ops.transpose(attn, (0, 2, 1, 3))
|
|
397
|
+
attn = ops.reshape(attn, (bsz, n_node, self.embed_dim))
|
|
398
|
+
|
|
399
|
+
return self.out_proj(attn)
|
|
400
|
+
|
|
401
|
+
|
|
402
|
+
class GPSLayer(layers.Layer):
|
|
403
|
+
r"""GraphGPS hybrid layer combining local MPNN (e.g. `CustomGatedGCN`) and global Multi-Head Attention.
|
|
404
|
+
|
|
405
|
+
Reference:
|
|
406
|
+
"Recipe for a General, Powerful, Scalable Graph Transformer" (NeurIPS 2022).
|
|
407
|
+
|
|
408
|
+
Args:
|
|
409
|
+
dim_h (int): Hidden embedding dimension.
|
|
410
|
+
local_gnn_type (str, optional): Local MPNN type. (default: ``"CustomGatedGCN"``)
|
|
411
|
+
global_model_type (str, optional): Global attention type. (default: ``"Transformer"``)
|
|
412
|
+
num_heads (int, optional): Number of attention heads. (default: ``8``)
|
|
413
|
+
act (str, optional): Activation function. (default: ``"gelu"``)
|
|
414
|
+
dropout (float, optional): Dropout rate. (default: ``0.0``)
|
|
415
|
+
attn_dropout (float, optional): Attention dropout rate. (default: ``0.0``)
|
|
416
|
+
layer_norm (bool, optional): Whether to use LayerNorm. (default: ``False``)
|
|
417
|
+
batch_norm (bool, optional): Whether to use BatchNorm. (default: ``True``)
|
|
418
|
+
**kwargs: Additional layer arguments.
|
|
419
|
+
|
|
420
|
+
Example:
|
|
421
|
+
```python
|
|
422
|
+
import numpy as np
|
|
423
|
+
from k3_node.models import GPSLayer
|
|
424
|
+
|
|
425
|
+
x = np.random.rand(5, 32).astype("float32") # node features
|
|
426
|
+
e = np.random.rand(6, 32).astype("float32") # edge features
|
|
427
|
+
edge_index = np.array([[0, 1, 1, 2, 3, 4], [1, 0, 2, 1, 4, 3]])
|
|
428
|
+
batch = np.array([0, 0, 0, 1, 1]) # two graphs
|
|
429
|
+
|
|
430
|
+
# Local message passing plus global attention over each graph's nodes
|
|
431
|
+
layer = GPSLayer(dim_h=32, local_gnn_type="CustomGatedGCN", global_model_type="Transformer", num_heads=4)
|
|
432
|
+
x_out, e_out = layer(x, edge_index, e, batch=batch)
|
|
433
|
+
print(tuple(x_out.shape), tuple(e_out.shape)) # (5, 32) (6, 32)
|
|
434
|
+
```
|
|
435
|
+
"""
|
|
436
|
+
|
|
437
|
+
def __init__(
|
|
438
|
+
self,
|
|
439
|
+
dim_h: int,
|
|
440
|
+
local_gnn_type: str = "CustomGatedGCN",
|
|
441
|
+
global_model_type: str = "Transformer",
|
|
442
|
+
num_heads: int = 8,
|
|
443
|
+
act: str = "gelu",
|
|
444
|
+
dropout: float = 0.0,
|
|
445
|
+
attn_dropout: float = 0.0,
|
|
446
|
+
layer_norm: bool = False,
|
|
447
|
+
batch_norm: bool = True,
|
|
448
|
+
**kwargs,
|
|
449
|
+
):
|
|
450
|
+
super().__init__(**kwargs)
|
|
451
|
+
self.dim_h = dim_h
|
|
452
|
+
self.local_gnn_type = local_gnn_type
|
|
453
|
+
self.global_model_type = global_model_type
|
|
454
|
+
self.num_heads = num_heads
|
|
455
|
+
self.act_name = act
|
|
456
|
+
self.dropout_rate = dropout
|
|
457
|
+
self.attn_dropout_rate = attn_dropout
|
|
458
|
+
self.layer_norm = layer_norm
|
|
459
|
+
self.batch_norm = batch_norm
|
|
460
|
+
|
|
461
|
+
if local_gnn_type == "CustomGatedGCN":
|
|
462
|
+
self.local_model = CustomGatedGCN(
|
|
463
|
+
in_dim=dim_h,
|
|
464
|
+
out_dim=dim_h,
|
|
465
|
+
dropout=dropout,
|
|
466
|
+
residual=True,
|
|
467
|
+
act=act,
|
|
468
|
+
name="local_model",
|
|
469
|
+
)
|
|
470
|
+
else:
|
|
471
|
+
self.local_model = None
|
|
472
|
+
|
|
473
|
+
if global_model_type in ["Transformer", "BiasedTransformer"]:
|
|
474
|
+
self.self_attn = TransformerSelfAttention(
|
|
475
|
+
embed_dim=dim_h,
|
|
476
|
+
num_heads=num_heads,
|
|
477
|
+
dropout=attn_dropout,
|
|
478
|
+
name="self_attn",
|
|
479
|
+
)
|
|
480
|
+
else:
|
|
481
|
+
self.self_attn = None
|
|
482
|
+
|
|
483
|
+
if layer_norm:
|
|
484
|
+
self.norm1_local = layers.LayerNormalization(epsilon=1e-5, name="norm1_local")
|
|
485
|
+
self.norm1_attn = layers.LayerNormalization(epsilon=1e-5, name="norm1_attn")
|
|
486
|
+
self.norm2 = layers.LayerNormalization(epsilon=1e-5, name="norm2")
|
|
487
|
+
elif batch_norm:
|
|
488
|
+
self.norm1_local = layers.BatchNormalization(
|
|
489
|
+
axis=-1, epsilon=1e-5, momentum=0.9, name="norm1_local"
|
|
490
|
+
)
|
|
491
|
+
self.norm1_attn = layers.BatchNormalization(
|
|
492
|
+
axis=-1, epsilon=1e-5, momentum=0.9, name="norm1_attn"
|
|
493
|
+
)
|
|
494
|
+
self.norm2 = layers.BatchNormalization(
|
|
495
|
+
axis=-1, epsilon=1e-5, momentum=0.9, name="norm2"
|
|
496
|
+
)
|
|
497
|
+
else:
|
|
498
|
+
self.norm1_local = None
|
|
499
|
+
self.norm1_attn = None
|
|
500
|
+
self.norm2 = None
|
|
501
|
+
|
|
502
|
+
self.dropout_local = layers.Dropout(dropout)
|
|
503
|
+
self.dropout_attn = layers.Dropout(dropout)
|
|
504
|
+
|
|
505
|
+
self.ff_linear1 = layers.Dense(dim_h * 2, use_bias=True, name="ff_linear1")
|
|
506
|
+
self.ff_linear2 = layers.Dense(dim_h, use_bias=True, name="ff_linear2")
|
|
507
|
+
self.ff_dropout1 = layers.Dropout(dropout)
|
|
508
|
+
self.ff_dropout2 = layers.Dropout(dropout)
|
|
509
|
+
|
|
510
|
+
def build(self, input_shape=None):
|
|
511
|
+
if self.local_model is not None and not self.local_model.built:
|
|
512
|
+
self.local_model.build(None)
|
|
513
|
+
if self.self_attn is not None and not self.self_attn.built:
|
|
514
|
+
self.self_attn.build(None)
|
|
515
|
+
if self.norm1_local is not None and not self.norm1_local.built:
|
|
516
|
+
self.norm1_local.build((None, self.dim_h))
|
|
517
|
+
if self.norm1_attn is not None and not self.norm1_attn.built:
|
|
518
|
+
self.norm1_attn.build((None, self.dim_h))
|
|
519
|
+
if not self.ff_linear1.built:
|
|
520
|
+
self.ff_linear1.build((None, self.dim_h))
|
|
521
|
+
self.ff_linear2.build((None, self.dim_h * 2))
|
|
522
|
+
if self.norm2 is not None and not self.norm2.built:
|
|
523
|
+
self.norm2.build((None, self.dim_h))
|
|
524
|
+
super().build(input_shape)
|
|
525
|
+
|
|
526
|
+
def call(self, x, edge_index, edge_attr, batch=None, training=False):
|
|
527
|
+
h_in1 = x
|
|
528
|
+
h_out_list = []
|
|
529
|
+
|
|
530
|
+
# Local MPNN
|
|
531
|
+
if self.local_model is not None:
|
|
532
|
+
h_local, edge_attr = self.local_model(
|
|
533
|
+
x, edge_index, edge_attr, training=training
|
|
534
|
+
)
|
|
535
|
+
# CustomGatedGCN handles residual internally
|
|
536
|
+
if self.norm1_local is not None:
|
|
537
|
+
h_local = self.norm1_local(h_local, training=training)
|
|
538
|
+
h_out_list.append(h_local)
|
|
539
|
+
|
|
540
|
+
# Global Attention
|
|
541
|
+
if self.self_attn is not None:
|
|
542
|
+
if batch is None:
|
|
543
|
+
h_dense = ops.expand_dims(x, axis=0)
|
|
544
|
+
mask = ops.ones((1, ops.shape(x)[0]), dtype="bool")
|
|
545
|
+
h_attn = self.self_attn(h_dense, mask=mask, training=training)[0]
|
|
546
|
+
else:
|
|
547
|
+
batch_np = ops.convert_to_numpy(batch)
|
|
548
|
+
B_int = int(batch_np.max()) + 1 if len(batch_np) > 0 else 1
|
|
549
|
+
counts = np.bincount(batch_np, minlength=B_int)
|
|
550
|
+
max_nodes = int(counts.max()) if len(counts) > 0 else 0
|
|
551
|
+
|
|
552
|
+
offsets = np.zeros(len(batch_np), dtype=np.int32)
|
|
553
|
+
curr = np.zeros(B_int, dtype=np.int32)
|
|
554
|
+
for i, b in enumerate(batch_np):
|
|
555
|
+
offsets[i] = curr[b]
|
|
556
|
+
curr[b] += 1
|
|
557
|
+
|
|
558
|
+
offsets_t = ops.convert_to_tensor(offsets, dtype="int32")
|
|
559
|
+
batch_cast = ops.cast(batch, "int32")
|
|
560
|
+
indices = ops.stack([batch_cast, offsets_t], axis=1)
|
|
561
|
+
|
|
562
|
+
dense_x = ops.scatter_update(
|
|
563
|
+
ops.zeros((B_int, max_nodes, self.dim_h), dtype=x.dtype),
|
|
564
|
+
indices,
|
|
565
|
+
x,
|
|
566
|
+
)
|
|
567
|
+
|
|
568
|
+
mask_np = np.zeros((B_int, max_nodes), dtype=bool)
|
|
569
|
+
for b, count in enumerate(counts):
|
|
570
|
+
mask_np[b, :count] = True
|
|
571
|
+
mask = ops.convert_to_tensor(mask_np)
|
|
572
|
+
|
|
573
|
+
h_attn_dense = self.self_attn(dense_x, mask=mask, training=training)
|
|
574
|
+
|
|
575
|
+
flat_attn = ops.reshape(h_attn_dense, (B_int * max_nodes, self.dim_h))
|
|
576
|
+
flat_indices = batch_cast * max_nodes + offsets_t
|
|
577
|
+
h_attn = ops.take(flat_attn, flat_indices, axis=0)
|
|
578
|
+
|
|
579
|
+
h_attn = self.dropout_attn(h_attn, training=training)
|
|
580
|
+
h_attn = h_in1 + h_attn
|
|
581
|
+
if self.norm1_attn is not None:
|
|
582
|
+
h_attn = self.norm1_attn(h_attn, training=training)
|
|
583
|
+
h_out_list.append(h_attn)
|
|
584
|
+
|
|
585
|
+
# Sum local and global representations
|
|
586
|
+
h = sum(h_out_list)
|
|
587
|
+
|
|
588
|
+
# Feed Forward block
|
|
589
|
+
act_fn = _get_act(self.act_name)
|
|
590
|
+
ff_out = self.ff_dropout1(act_fn(self.ff_linear1(h)), training=training)
|
|
591
|
+
ff_out = self.ff_dropout2(self.ff_linear2(ff_out), training=training)
|
|
592
|
+
h = h + ff_out
|
|
593
|
+
|
|
594
|
+
if self.norm2 is not None:
|
|
595
|
+
h = self.norm2(h, training=training)
|
|
596
|
+
|
|
597
|
+
return h, edge_attr
|
|
598
|
+
|
|
599
|
+
|
|
600
|
+
class SANGraphHead(layers.Layer):
|
|
601
|
+
r"""Prediction head for graph-level tasks from the Spectral Attention Network (SAN).
|
|
602
|
+
|
|
603
|
+
Args:
|
|
604
|
+
dim_in (int): Input feature dimension.
|
|
605
|
+
dim_out (int): Output feature dimension. (default: ``1``)
|
|
606
|
+
L (int, optional): Number of hidden layers. (default: ``2``)
|
|
607
|
+
act (str, optional): Activation function. (default: ``"gelu"``)
|
|
608
|
+
pooling (str, optional): Graph pooling method ('mean', 'add', 'max'). (default: ``"mean"``)
|
|
609
|
+
**kwargs: Additional layer arguments.
|
|
610
|
+
|
|
611
|
+
Example:
|
|
612
|
+
```python
|
|
613
|
+
import numpy as np
|
|
614
|
+
from k3_node.models import SANGraphHead
|
|
615
|
+
|
|
616
|
+
x = np.random.rand(5, 32).astype("float32") # node features
|
|
617
|
+
e = np.random.rand(6, 32).astype("float32") # edge features
|
|
618
|
+
edge_index = np.array([[0, 1, 1, 2, 3, 4], [1, 0, 2, 1, 4, 3]])
|
|
619
|
+
batch = np.array([0, 0, 0, 1, 1]) # two graphs
|
|
620
|
+
|
|
621
|
+
head = SANGraphHead(dim_in=32, dim_out=1, L=2, pooling="mean")
|
|
622
|
+
print(tuple(head(x, batch=batch).shape)) # (2, 1): one prediction per graph
|
|
623
|
+
```
|
|
624
|
+
"""
|
|
625
|
+
|
|
626
|
+
def __init__(
|
|
627
|
+
self,
|
|
628
|
+
dim_in: int,
|
|
629
|
+
dim_out: int = 1,
|
|
630
|
+
L: int = 2,
|
|
631
|
+
act: str = "gelu",
|
|
632
|
+
pooling: str = "mean",
|
|
633
|
+
**kwargs,
|
|
634
|
+
):
|
|
635
|
+
super().__init__(**kwargs)
|
|
636
|
+
self.dim_in = dim_in
|
|
637
|
+
self.dim_out = dim_out
|
|
638
|
+
self.L = L
|
|
639
|
+
self.act_name = act
|
|
640
|
+
self.pooling = pooling.lower()
|
|
641
|
+
|
|
642
|
+
self.FC_layers = []
|
|
643
|
+
for l in range(L):
|
|
644
|
+
in_dim_l = dim_in // (2**l)
|
|
645
|
+
out_dim_l = dim_in // (2 ** (l + 1))
|
|
646
|
+
self.FC_layers.append(
|
|
647
|
+
layers.Dense(out_dim_l, use_bias=True, name=f"FC_layers_{l}")
|
|
648
|
+
)
|
|
649
|
+
self.FC_layers.append(
|
|
650
|
+
layers.Dense(dim_out, use_bias=True, name=f"FC_layers_{L}")
|
|
651
|
+
)
|
|
652
|
+
|
|
653
|
+
def build(self, input_shape=None):
|
|
654
|
+
curr_dim = self.dim_in
|
|
655
|
+
for layer in self.FC_layers:
|
|
656
|
+
if not layer.built:
|
|
657
|
+
layer.build((None, curr_dim))
|
|
658
|
+
curr_dim = layer.units
|
|
659
|
+
super().build(input_shape)
|
|
660
|
+
|
|
661
|
+
def call(self, x, batch=None, training=False, batch_size=None):
|
|
662
|
+
if self.pooling in ["mean", "avg"]:
|
|
663
|
+
graph_emb = global_mean_pool(x, batch, size=batch_size)
|
|
664
|
+
elif self.pooling in ["add", "sum"]:
|
|
665
|
+
graph_emb = global_add_pool(x, batch, size=batch_size)
|
|
666
|
+
elif self.pooling == "max":
|
|
667
|
+
graph_emb = global_max_pool(x, batch, size=batch_size)
|
|
668
|
+
else:
|
|
669
|
+
raise ValueError(f"Unknown pooling method '{self.pooling}'")
|
|
670
|
+
|
|
671
|
+
act_fn = _get_act(self.act_name)
|
|
672
|
+
for l in range(self.L):
|
|
673
|
+
graph_emb = self.FC_layers[l](graph_emb)
|
|
674
|
+
graph_emb = act_fn(graph_emb)
|
|
675
|
+
|
|
676
|
+
graph_emb = self.FC_layers[self.L](graph_emb)
|
|
677
|
+
return graph_emb
|
|
678
|
+
|
|
679
|
+
|
|
680
|
+
class GPSModel(keras.Model):
|
|
681
|
+
r"""GraphGPS: General Powerful Scalable Graph Transformer from the
|
|
682
|
+
`"Recipe for a General, Powerful, Scalable Graph Transformer"
|
|
683
|
+
<https://arxiv.org/abs/2205.12454>`_ paper (NeurIPS 2022).
|
|
684
|
+
|
|
685
|
+
Args:
|
|
686
|
+
dim_in (int, optional): Initial input feature dimension. (default: ``256``)
|
|
687
|
+
dim_out (int, optional): Target output dimension. (default: ``1``)
|
|
688
|
+
num_layers (int, optional): Number of GPS layers. (default: ``16``)
|
|
689
|
+
dim_hidden (int, optional): Hidden embedding dimension. (default: ``256``)
|
|
690
|
+
num_heads (int, optional): Number of attention heads. (default: ``8``)
|
|
691
|
+
local_gnn_type (str, optional): Local MPNN layer type. (default: ``"CustomGatedGCN"``)
|
|
692
|
+
act (str, optional): Activation function. (default: ``"gelu"``)
|
|
693
|
+
dropout (float, optional): Dropout probability. (default: ``0.1``)
|
|
694
|
+
attn_dropout (float, optional): Attention dropout probability. (default: ``0.1``)
|
|
695
|
+
batch_norm (bool, optional): Whether to use batch normalization. (default: ``True``)
|
|
696
|
+
layer_norm (bool, optional): Whether to use layer normalization. (default: ``False``)
|
|
697
|
+
node_encoder_type (str, optional): Node encoder type ("Atom+RWSE", "Atom", "Linear", or None). (default: ``"Atom+RWSE"``)
|
|
698
|
+
edge_encoder_type (str, optional): Edge encoder type ("Bond", "Linear", or None). (default: ``"Bond"``)
|
|
699
|
+
atom_feature_dims (List[int], optional): Categorical feature vocabulary sizes for atom features.
|
|
700
|
+
bond_feature_dims (List[int], optional): Categorical feature vocabulary sizes for bond features.
|
|
701
|
+
rwse_num_steps (int, optional): Number of RWSE steps. (default: ``16``)
|
|
702
|
+
rwse_dim_pe (int, optional): RWSE embedding dimension. (default: ``20``)
|
|
703
|
+
graph_pooling (str, optional): Graph pooling type ('mean', 'add', 'max'). (default: ``"mean"``)
|
|
704
|
+
head_layers (int, optional): Number of hidden layers in prediction head. (default: ``2``)
|
|
705
|
+
**kwargs: Additional model arguments.
|
|
706
|
+
|
|
707
|
+
Example:
|
|
708
|
+
```python
|
|
709
|
+
import numpy as np
|
|
710
|
+
from k3_node.models import GPSModel
|
|
711
|
+
|
|
712
|
+
# OGB-style integer atom and bond features (9 per atom, 3 per bond)
|
|
713
|
+
x = np.random.randint(0, 2, size=(5, 9))
|
|
714
|
+
edge_index = np.array([[0, 1, 1, 2, 3, 4], [1, 0, 2, 1, 4, 3]])
|
|
715
|
+
edge_attr = np.random.randint(0, 2, size=(6, 3))
|
|
716
|
+
batch = np.array([0, 0, 0, 1, 1]) # two graphs
|
|
717
|
+
rwse = np.random.rand(5, 16).astype("float32") # random-walk structural encodings
|
|
718
|
+
|
|
719
|
+
model = GPSModel(dim_in=32, dim_out=1, num_layers=2, dim_hidden=32, num_heads=4,
|
|
720
|
+
node_encoder_type="Atom+RWSE", edge_encoder_type="Bond",
|
|
721
|
+
rwse_num_steps=16, rwse_dim_pe=8)
|
|
722
|
+
pred = model(x, edge_index, edge_attr=edge_attr, pestat_RWSE=rwse, batch=batch, batch_size=2)
|
|
723
|
+
print(tuple(pred.shape)) # (2, 1): one prediction per graph
|
|
724
|
+
```
|
|
725
|
+
"""
|
|
726
|
+
|
|
727
|
+
def __init__(
|
|
728
|
+
self,
|
|
729
|
+
dim_in: int = 256,
|
|
730
|
+
dim_out: int = 1,
|
|
731
|
+
num_layers: int = 16,
|
|
732
|
+
dim_hidden: int = 256,
|
|
733
|
+
num_heads: int = 8,
|
|
734
|
+
local_gnn_type: str = "CustomGatedGCN",
|
|
735
|
+
act: str = "gelu",
|
|
736
|
+
dropout: float = 0.1,
|
|
737
|
+
attn_dropout: float = 0.1,
|
|
738
|
+
batch_norm: bool = True,
|
|
739
|
+
layer_norm: bool = False,
|
|
740
|
+
node_encoder_type: Optional[str] = "Atom+RWSE",
|
|
741
|
+
edge_encoder_type: Optional[str] = "Bond",
|
|
742
|
+
atom_feature_dims: Optional[List[int]] = None,
|
|
743
|
+
bond_feature_dims: Optional[List[int]] = None,
|
|
744
|
+
rwse_num_steps: int = 16,
|
|
745
|
+
rwse_dim_pe: int = 20,
|
|
746
|
+
graph_pooling: str = "mean",
|
|
747
|
+
head_layers: int = 2,
|
|
748
|
+
**kwargs,
|
|
749
|
+
):
|
|
750
|
+
super().__init__(**kwargs)
|
|
751
|
+
self.dim_in = dim_in
|
|
752
|
+
self.dim_out = dim_out
|
|
753
|
+
self.num_layers = num_layers
|
|
754
|
+
self.dim_hidden = dim_hidden
|
|
755
|
+
self.num_heads = num_heads
|
|
756
|
+
self.local_gnn_type = local_gnn_type
|
|
757
|
+
self.act_name = act
|
|
758
|
+
self.dropout_rate = dropout
|
|
759
|
+
self.attn_dropout_rate = attn_dropout
|
|
760
|
+
self.batch_norm = batch_norm
|
|
761
|
+
self.layer_norm = layer_norm
|
|
762
|
+
self.node_encoder_type = node_encoder_type
|
|
763
|
+
self.edge_encoder_type = edge_encoder_type
|
|
764
|
+
self.rwse_num_steps = rwse_num_steps
|
|
765
|
+
self.rwse_dim_pe = rwse_dim_pe
|
|
766
|
+
self.graph_pooling = graph_pooling
|
|
767
|
+
self.head_layers = head_layers
|
|
768
|
+
|
|
769
|
+
# Node Encoder
|
|
770
|
+
if node_encoder_type == "Atom+RWSE":
|
|
771
|
+
self.atom_encoder = AtomEncoder(
|
|
772
|
+
emb_dim=dim_hidden - rwse_dim_pe,
|
|
773
|
+
feature_dims=atom_feature_dims,
|
|
774
|
+
name="atom_encoder",
|
|
775
|
+
)
|
|
776
|
+
self.rwse_encoder = RWSEEncoder(
|
|
777
|
+
num_rw_steps=rwse_num_steps,
|
|
778
|
+
pe_dim=rwse_dim_pe,
|
|
779
|
+
name="rwse_encoder",
|
|
780
|
+
)
|
|
781
|
+
elif node_encoder_type == "Atom":
|
|
782
|
+
self.atom_encoder = AtomEncoder(
|
|
783
|
+
emb_dim=dim_hidden,
|
|
784
|
+
feature_dims=atom_feature_dims,
|
|
785
|
+
name="atom_encoder",
|
|
786
|
+
)
|
|
787
|
+
self.rwse_encoder = None
|
|
788
|
+
elif node_encoder_type == "Linear":
|
|
789
|
+
self.linear_node_encoder = layers.Dense(dim_hidden, name="linear_node_encoder")
|
|
790
|
+
self.atom_encoder = None
|
|
791
|
+
self.rwse_encoder = None
|
|
792
|
+
else:
|
|
793
|
+
self.atom_encoder = None
|
|
794
|
+
self.rwse_encoder = None
|
|
795
|
+
self.linear_node_encoder = None
|
|
796
|
+
|
|
797
|
+
# Edge Encoder
|
|
798
|
+
if edge_encoder_type == "Bond":
|
|
799
|
+
self.bond_encoder = BondEncoder(
|
|
800
|
+
emb_dim=dim_hidden,
|
|
801
|
+
feature_dims=bond_feature_dims,
|
|
802
|
+
name="bond_encoder",
|
|
803
|
+
)
|
|
804
|
+
elif edge_encoder_type == "Linear":
|
|
805
|
+
self.linear_edge_encoder = layers.Dense(dim_hidden, name="linear_edge_encoder")
|
|
806
|
+
self.bond_encoder = None
|
|
807
|
+
else:
|
|
808
|
+
self.bond_encoder = None
|
|
809
|
+
self.linear_edge_encoder = None
|
|
810
|
+
|
|
811
|
+
# GPS Layers
|
|
812
|
+
self.gps_layers = [
|
|
813
|
+
GPSLayer(
|
|
814
|
+
dim_h=dim_hidden,
|
|
815
|
+
local_gnn_type=local_gnn_type,
|
|
816
|
+
global_model_type="Transformer",
|
|
817
|
+
num_heads=num_heads,
|
|
818
|
+
act=act,
|
|
819
|
+
dropout=dropout,
|
|
820
|
+
attn_dropout=attn_dropout,
|
|
821
|
+
layer_norm=layer_norm,
|
|
822
|
+
batch_norm=batch_norm,
|
|
823
|
+
name=f"gps_layer_{i}",
|
|
824
|
+
)
|
|
825
|
+
for i in range(num_layers)
|
|
826
|
+
]
|
|
827
|
+
|
|
828
|
+
# SANGraphHead
|
|
829
|
+
self.post_mp = SANGraphHead(
|
|
830
|
+
dim_in=dim_hidden,
|
|
831
|
+
dim_out=dim_out,
|
|
832
|
+
L=head_layers,
|
|
833
|
+
act=act,
|
|
834
|
+
pooling=graph_pooling,
|
|
835
|
+
name="post_mp",
|
|
836
|
+
)
|
|
837
|
+
|
|
838
|
+
def build(self, input_shape=None):
|
|
839
|
+
if self.atom_encoder is not None and not self.atom_encoder.built:
|
|
840
|
+
self.atom_encoder.build(None)
|
|
841
|
+
if self.rwse_encoder is not None and not self.rwse_encoder.built:
|
|
842
|
+
self.rwse_encoder.build(None)
|
|
843
|
+
if self.bond_encoder is not None and not self.bond_encoder.built:
|
|
844
|
+
self.bond_encoder.build(None)
|
|
845
|
+
for layer in self.gps_layers:
|
|
846
|
+
if not layer.built:
|
|
847
|
+
layer.build(None)
|
|
848
|
+
if not self.post_mp.built:
|
|
849
|
+
self.post_mp.build(None)
|
|
850
|
+
super().build(input_shape)
|
|
851
|
+
|
|
852
|
+
def call(
|
|
853
|
+
self,
|
|
854
|
+
x,
|
|
855
|
+
edge_index,
|
|
856
|
+
edge_attr=None,
|
|
857
|
+
pestat_RWSE=None,
|
|
858
|
+
batch=None,
|
|
859
|
+
training=False,
|
|
860
|
+
batch_size=None,
|
|
861
|
+
):
|
|
862
|
+
# Node encoding
|
|
863
|
+
if self.node_encoder_type == "Atom+RWSE":
|
|
864
|
+
x_emb = self.atom_encoder(x)
|
|
865
|
+
if pestat_RWSE is not None and self.rwse_encoder is not None:
|
|
866
|
+
pe_emb = self.rwse_encoder(pestat_RWSE, training=training)
|
|
867
|
+
x = ops.concatenate([x_emb, pe_emb], axis=-1)
|
|
868
|
+
else:
|
|
869
|
+
x = x_emb
|
|
870
|
+
elif self.node_encoder_type == "Atom":
|
|
871
|
+
x = self.atom_encoder(x)
|
|
872
|
+
elif self.node_encoder_type == "Linear" and self.linear_node_encoder is not None:
|
|
873
|
+
x = self.linear_node_encoder(x)
|
|
874
|
+
|
|
875
|
+
# Edge encoding
|
|
876
|
+
if self.edge_encoder_type == "Bond" and edge_attr is not None:
|
|
877
|
+
edge_attr = self.bond_encoder(edge_attr)
|
|
878
|
+
elif self.edge_encoder_type == "Linear" and self.linear_edge_encoder is not None:
|
|
879
|
+
edge_attr = self.linear_edge_encoder(edge_attr)
|
|
880
|
+
|
|
881
|
+
# GPS layers
|
|
882
|
+
for layer in self.gps_layers:
|
|
883
|
+
x, edge_attr = layer(
|
|
884
|
+
x, edge_index, edge_attr, batch=batch, training=training
|
|
885
|
+
)
|
|
886
|
+
|
|
887
|
+
# Head
|
|
888
|
+
pred = self.post_mp(x, batch=batch, training=training, batch_size=batch_size)
|
|
889
|
+
return pred
|
|
890
|
+
|
|
891
|
+
|
|
892
|
+
def load_gps_weights(model: GPSModel, checkpoint_path: str):
|
|
893
|
+
r"""Loads trained PyTorch GraphGPS checkpoint weights into a Keras 3 `GPSModel`.
|
|
894
|
+
|
|
895
|
+
Args:
|
|
896
|
+
model (GPSModel): Target `GPSModel` instance.
|
|
897
|
+
checkpoint_path (str): Path to PyTorch `.ckpt` or `.pt` checkpoint file.
|
|
898
|
+
|
|
899
|
+
Returns:
|
|
900
|
+
GPSModel: The model with loaded weights.
|
|
901
|
+
"""
|
|
902
|
+
import torch
|
|
903
|
+
|
|
904
|
+
ckpt = torch.load(checkpoint_path, map_location="cpu")
|
|
905
|
+
if "model_state" in ckpt:
|
|
906
|
+
state_dict = ckpt["model_state"]
|
|
907
|
+
elif "state_dict" in ckpt:
|
|
908
|
+
state_dict = ckpt["state_dict"]
|
|
909
|
+
else:
|
|
910
|
+
state_dict = ckpt
|
|
911
|
+
|
|
912
|
+
clean_dict = {}
|
|
913
|
+
for k, v in state_dict.items():
|
|
914
|
+
if k.startswith("model."):
|
|
915
|
+
k = k[6:]
|
|
916
|
+
clean_dict[k] = v
|
|
917
|
+
|
|
918
|
+
def _to_tensor(t):
|
|
919
|
+
arr = t.detach().cpu().numpy()
|
|
920
|
+
return ops.convert_to_tensor(arr, dtype="float32")
|
|
921
|
+
|
|
922
|
+
# Build model if needed
|
|
923
|
+
if not model.built:
|
|
924
|
+
model.build(None)
|
|
925
|
+
|
|
926
|
+
# 1. Node Encoder
|
|
927
|
+
if model.node_encoder_type in ["Atom+RWSE", "Atom"] and model.atom_encoder is not None:
|
|
928
|
+
for i, emb in enumerate(model.atom_encoder.atom_embedding_list):
|
|
929
|
+
k = f"encoder.node_encoder.encoder1.atom_embedding_list.{i}.weight"
|
|
930
|
+
if k in clean_dict:
|
|
931
|
+
emb.embeddings.assign(_to_tensor(clean_dict[k]))
|
|
932
|
+
|
|
933
|
+
if model.node_encoder_type == "Atom+RWSE" and model.rwse_encoder is not None:
|
|
934
|
+
p = "encoder.node_encoder.encoder2"
|
|
935
|
+
if f"{p}.raw_norm.weight" in clean_dict:
|
|
936
|
+
model.rwse_encoder.raw_norm.gamma.assign(_to_tensor(clean_dict[f"{p}.raw_norm.weight"]))
|
|
937
|
+
if f"{p}.raw_norm.bias" in clean_dict:
|
|
938
|
+
model.rwse_encoder.raw_norm.beta.assign(_to_tensor(clean_dict[f"{p}.raw_norm.bias"]))
|
|
939
|
+
if f"{p}.raw_norm.running_mean" in clean_dict:
|
|
940
|
+
model.rwse_encoder.raw_norm.moving_mean.assign(_to_tensor(clean_dict[f"{p}.raw_norm.running_mean"]))
|
|
941
|
+
if f"{p}.raw_norm.running_var" in clean_dict:
|
|
942
|
+
model.rwse_encoder.raw_norm.moving_variance.assign(_to_tensor(clean_dict[f"{p}.raw_norm.running_var"]))
|
|
943
|
+
|
|
944
|
+
if f"{p}.pe_encoder.weight" in clean_dict:
|
|
945
|
+
model.rwse_encoder.pe_encoder.kernel.assign(_to_tensor(clean_dict[f"{p}.pe_encoder.weight"].t()))
|
|
946
|
+
if f"{p}.pe_encoder.bias" in clean_dict:
|
|
947
|
+
model.rwse_encoder.pe_encoder.bias.assign(_to_tensor(clean_dict[f"{p}.pe_encoder.bias"]))
|
|
948
|
+
|
|
949
|
+
# 2. Edge Encoder
|
|
950
|
+
if model.edge_encoder_type == "Bond" and model.bond_encoder is not None:
|
|
951
|
+
for i, emb in enumerate(model.bond_encoder.bond_embedding_list):
|
|
952
|
+
k = f"encoder.edge_encoder.bond_embedding_list.{i}.weight"
|
|
953
|
+
if k in clean_dict:
|
|
954
|
+
emb.embeddings.assign(_to_tensor(clean_dict[k]))
|
|
955
|
+
|
|
956
|
+
# 3. GPS Layers
|
|
957
|
+
dim_h = model.dim_hidden
|
|
958
|
+
for i, gps_layer in enumerate(model.gps_layers):
|
|
959
|
+
p = f"layers.{i}"
|
|
960
|
+
|
|
961
|
+
# Local MPNN: CustomGatedGCN
|
|
962
|
+
if gps_layer.local_model is not None:
|
|
963
|
+
lm = gps_layer.local_model
|
|
964
|
+
for proj_name in ["A", "B", "C", "D", "E"]:
|
|
965
|
+
dense_proj = getattr(lm, proj_name)
|
|
966
|
+
if f"{p}.local_model.{proj_name}.weight" in clean_dict:
|
|
967
|
+
dense_proj.kernel.assign(_to_tensor(clean_dict[f"{p}.local_model.{proj_name}.weight"].t()))
|
|
968
|
+
if f"{p}.local_model.{proj_name}.bias" in clean_dict:
|
|
969
|
+
dense_proj.bias.assign(_to_tensor(clean_dict[f"{p}.local_model.{proj_name}.bias"]))
|
|
970
|
+
|
|
971
|
+
for bn_name, bn_layer in [("bn_node_x", lm.bn_node_x), ("bn_edge_e", lm.bn_edge_e)]:
|
|
972
|
+
if f"{p}.local_model.{bn_name}.weight" in clean_dict:
|
|
973
|
+
bn_layer.gamma.assign(_to_tensor(clean_dict[f"{p}.local_model.{bn_name}.weight"]))
|
|
974
|
+
if f"{p}.local_model.{bn_name}.bias" in clean_dict:
|
|
975
|
+
bn_layer.beta.assign(_to_tensor(clean_dict[f"{p}.local_model.{bn_name}.bias"]))
|
|
976
|
+
if f"{p}.local_model.{bn_name}.running_mean" in clean_dict:
|
|
977
|
+
bn_layer.moving_mean.assign(_to_tensor(clean_dict[f"{p}.local_model.{bn_name}.running_mean"]))
|
|
978
|
+
if f"{p}.local_model.{bn_name}.running_var" in clean_dict:
|
|
979
|
+
bn_layer.moving_variance.assign(_to_tensor(clean_dict[f"{p}.local_model.{bn_name}.running_var"]))
|
|
980
|
+
|
|
981
|
+
if gps_layer.norm1_local is not None:
|
|
982
|
+
if f"{p}.norm1_local.weight" in clean_dict:
|
|
983
|
+
gps_layer.norm1_local.gamma.assign(_to_tensor(clean_dict[f"{p}.norm1_local.weight"]))
|
|
984
|
+
if f"{p}.norm1_local.bias" in clean_dict:
|
|
985
|
+
gps_layer.norm1_local.beta.assign(_to_tensor(clean_dict[f"{p}.norm1_local.bias"]))
|
|
986
|
+
if hasattr(gps_layer.norm1_local, "moving_mean") and f"{p}.norm1_local.running_mean" in clean_dict:
|
|
987
|
+
gps_layer.norm1_local.moving_mean.assign(_to_tensor(clean_dict[f"{p}.norm1_local.running_mean"]))
|
|
988
|
+
if hasattr(gps_layer.norm1_local, "moving_variance") and f"{p}.norm1_local.running_var" in clean_dict:
|
|
989
|
+
gps_layer.norm1_local.moving_variance.assign(_to_tensor(clean_dict[f"{p}.norm1_local.running_var"]))
|
|
990
|
+
|
|
991
|
+
# Global Attention: MultiHeadAttention
|
|
992
|
+
if gps_layer.self_attn is not None:
|
|
993
|
+
sa = gps_layer.self_attn
|
|
994
|
+
if f"{p}.self_attn.in_proj_weight" in clean_dict:
|
|
995
|
+
in_w = clean_dict[f"{p}.self_attn.in_proj_weight"]
|
|
996
|
+
sa.q_proj.kernel.assign(_to_tensor(in_w[:dim_h, :].t()))
|
|
997
|
+
sa.k_proj.kernel.assign(_to_tensor(in_w[dim_h:2*dim_h, :].t()))
|
|
998
|
+
sa.v_proj.kernel.assign(_to_tensor(in_w[2*dim_h:, :].t()))
|
|
999
|
+
|
|
1000
|
+
if f"{p}.self_attn.in_proj_bias" in clean_dict:
|
|
1001
|
+
in_b = clean_dict[f"{p}.self_attn.in_proj_bias"]
|
|
1002
|
+
sa.q_proj.bias.assign(_to_tensor(in_b[:dim_h]))
|
|
1003
|
+
sa.k_proj.bias.assign(_to_tensor(in_b[dim_h:2*dim_h]))
|
|
1004
|
+
sa.v_proj.bias.assign(_to_tensor(in_b[2*dim_h:]))
|
|
1005
|
+
|
|
1006
|
+
if f"{p}.self_attn.out_proj.weight" in clean_dict:
|
|
1007
|
+
sa.out_proj.kernel.assign(_to_tensor(clean_dict[f"{p}.self_attn.out_proj.weight"].t()))
|
|
1008
|
+
if f"{p}.self_attn.out_proj.bias" in clean_dict:
|
|
1009
|
+
sa.out_proj.bias.assign(_to_tensor(clean_dict[f"{p}.self_attn.out_proj.bias"]))
|
|
1010
|
+
|
|
1011
|
+
if gps_layer.norm1_attn is not None:
|
|
1012
|
+
if f"{p}.norm1_attn.weight" in clean_dict:
|
|
1013
|
+
gps_layer.norm1_attn.gamma.assign(_to_tensor(clean_dict[f"{p}.norm1_attn.weight"]))
|
|
1014
|
+
if f"{p}.norm1_attn.bias" in clean_dict:
|
|
1015
|
+
gps_layer.norm1_attn.beta.assign(_to_tensor(clean_dict[f"{p}.norm1_attn.bias"]))
|
|
1016
|
+
if hasattr(gps_layer.norm1_attn, "moving_mean") and f"{p}.norm1_attn.running_mean" in clean_dict:
|
|
1017
|
+
gps_layer.norm1_attn.moving_mean.assign(_to_tensor(clean_dict[f"{p}.norm1_attn.running_mean"]))
|
|
1018
|
+
if hasattr(gps_layer.norm1_attn, "moving_variance") and f"{p}.norm1_attn.running_var" in clean_dict:
|
|
1019
|
+
gps_layer.norm1_attn.moving_variance.assign(_to_tensor(clean_dict[f"{p}.norm1_attn.running_var"]))
|
|
1020
|
+
|
|
1021
|
+
# FFN
|
|
1022
|
+
if f"{p}.ff_linear1.weight" in clean_dict:
|
|
1023
|
+
gps_layer.ff_linear1.kernel.assign(_to_tensor(clean_dict[f"{p}.ff_linear1.weight"].t()))
|
|
1024
|
+
if f"{p}.ff_linear1.bias" in clean_dict:
|
|
1025
|
+
gps_layer.ff_linear1.bias.assign(_to_tensor(clean_dict[f"{p}.ff_linear1.bias"]))
|
|
1026
|
+
if f"{p}.ff_linear2.weight" in clean_dict:
|
|
1027
|
+
gps_layer.ff_linear2.kernel.assign(_to_tensor(clean_dict[f"{p}.ff_linear2.weight"].t()))
|
|
1028
|
+
if f"{p}.ff_linear2.bias" in clean_dict:
|
|
1029
|
+
gps_layer.ff_linear2.bias.assign(_to_tensor(clean_dict[f"{p}.ff_linear2.bias"]))
|
|
1030
|
+
|
|
1031
|
+
if gps_layer.norm2 is not None:
|
|
1032
|
+
if f"{p}.norm2.weight" in clean_dict:
|
|
1033
|
+
gps_layer.norm2.gamma.assign(_to_tensor(clean_dict[f"{p}.norm2.weight"]))
|
|
1034
|
+
if f"{p}.norm2.bias" in clean_dict:
|
|
1035
|
+
gps_layer.norm2.beta.assign(_to_tensor(clean_dict[f"{p}.norm2.bias"]))
|
|
1036
|
+
if hasattr(gps_layer.norm2, "moving_mean") and f"{p}.norm2.running_mean" in clean_dict:
|
|
1037
|
+
gps_layer.norm2.moving_mean.assign(_to_tensor(clean_dict[f"{p}.norm2.running_mean"]))
|
|
1038
|
+
if hasattr(gps_layer.norm2, "moving_variance") and f"{p}.norm2.running_var" in clean_dict:
|
|
1039
|
+
gps_layer.norm2.moving_variance.assign(_to_tensor(clean_dict[f"{p}.norm2.running_var"]))
|
|
1040
|
+
if hasattr(gps_layer.norm2, "moving_mean") and f"{p}.norm2.running_mean" in clean_dict:
|
|
1041
|
+
gps_layer.norm2.moving_mean.assign(_to_tensor(clean_dict[f"{p}.norm2.running_mean"]))
|
|
1042
|
+
if hasattr(gps_layer.norm2, "moving_variance") and f"{p}.norm2.running_var" in clean_dict:
|
|
1043
|
+
gps_layer.norm2.moving_variance.assign(_to_tensor(clean_dict[f"{p}.norm2.running_var"]))
|
|
1044
|
+
|
|
1045
|
+
# 4. SANGraphHead (post_mp)
|
|
1046
|
+
for l, fc in enumerate(model.post_mp.FC_layers):
|
|
1047
|
+
if f"post_mp.FC_layers.{l}.weight" in clean_dict:
|
|
1048
|
+
fc.kernel.assign(_to_tensor(clean_dict[f"post_mp.FC_layers.{l}.weight"].t()))
|
|
1049
|
+
if f"post_mp.FC_layers.{l}.bias" in clean_dict:
|
|
1050
|
+
fc.bias.assign(_to_tensor(clean_dict[f"post_mp.FC_layers.{l}.bias"]))
|
|
1051
|
+
|
|
1052
|
+
return model
|
|
1053
|
+
|
|
1054
|
+
|
|
1055
|
+
DROPBOX_CHECKPOINT_URLS = {
|
|
1056
|
+
"pcqm4m-GPS+RWSE.deep": "https://www.dropbox.com/s/aomimvak4gb6et3/pcqm4m-GPS%2BRWSE.deep.zip?dl=1",
|
|
1057
|
+
}
|
|
1058
|
+
|
|
1059
|
+
|
|
1060
|
+
def download_gps_checkpoint(
|
|
1061
|
+
checkpoint_name: str = "pcqm4m-GPS+RWSE.deep",
|
|
1062
|
+
cache_dir: Optional[str] = None,
|
|
1063
|
+
) -> str:
|
|
1064
|
+
r"""Downloads and extracts a pretrained GraphGPS checkpoint.
|
|
1065
|
+
|
|
1066
|
+
Args:
|
|
1067
|
+
checkpoint_name (str, optional): Name of the checkpoint.
|
|
1068
|
+
Currently supported: ``"pcqm4m-GPS+RWSE.deep"``.
|
|
1069
|
+
cache_dir (str, optional): Cache directory to store downloaded checkpoint.
|
|
1070
|
+
|
|
1071
|
+
Returns:
|
|
1072
|
+
str: Absolute path to the extracted `.ckpt` checkpoint file.
|
|
1073
|
+
"""
|
|
1074
|
+
if cache_dir is None:
|
|
1075
|
+
cache_dir = os.path.expanduser("~/.cache/k3_node/graphgps")
|
|
1076
|
+
os.makedirs(cache_dir, exist_ok=True)
|
|
1077
|
+
|
|
1078
|
+
expected_ckpt = os.path.join(cache_dir, checkpoint_name, "0", "ckpt", "148.ckpt")
|
|
1079
|
+
if os.path.exists(expected_ckpt):
|
|
1080
|
+
return expected_ckpt
|
|
1081
|
+
|
|
1082
|
+
# Also check repo pretrained dir if available
|
|
1083
|
+
local_repo_ckpt = os.path.join(
|
|
1084
|
+
os.path.dirname(os.path.dirname(os.path.abspath(__file__))),
|
|
1085
|
+
"..",
|
|
1086
|
+
"GraphGPS",
|
|
1087
|
+
"pretrained",
|
|
1088
|
+
checkpoint_name,
|
|
1089
|
+
"0",
|
|
1090
|
+
"ckpt",
|
|
1091
|
+
"148.ckpt",
|
|
1092
|
+
)
|
|
1093
|
+
if os.path.exists(local_repo_ckpt):
|
|
1094
|
+
return os.path.abspath(local_repo_ckpt)
|
|
1095
|
+
|
|
1096
|
+
url = DROPBOX_CHECKPOINT_URLS.get(checkpoint_name)
|
|
1097
|
+
if url is None:
|
|
1098
|
+
raise ValueError(f"Unknown checkpoint '{checkpoint_name}'")
|
|
1099
|
+
|
|
1100
|
+
zip_path = os.path.join(cache_dir, f"{checkpoint_name}.zip")
|
|
1101
|
+
if not os.path.exists(zip_path):
|
|
1102
|
+
print(f"Downloading {checkpoint_name} from {url}...")
|
|
1103
|
+
req = urllib.request.Request(url, headers={"User-Agent": "Mozilla/5.0"})
|
|
1104
|
+
with urllib.request.urlopen(req) as resp, open(zip_path, "wb") as f:
|
|
1105
|
+
while True:
|
|
1106
|
+
chunk = resp.read(1024 * 1024)
|
|
1107
|
+
if not chunk:
|
|
1108
|
+
break
|
|
1109
|
+
f.write(chunk)
|
|
1110
|
+
|
|
1111
|
+
with zipfile.ZipFile(zip_path, "r") as zip_ref:
|
|
1112
|
+
zip_ref.extractall(cache_dir)
|
|
1113
|
+
|
|
1114
|
+
if not os.path.exists(expected_ckpt):
|
|
1115
|
+
# Search for any ckpt file inside extracted directory
|
|
1116
|
+
for root, _, files in os.walk(cache_dir):
|
|
1117
|
+
for file in files:
|
|
1118
|
+
if file.endswith(".ckpt"):
|
|
1119
|
+
return os.path.join(root, file)
|
|
1120
|
+
raise FileNotFoundError(f"Could not find .ckpt inside extracted {zip_path}")
|
|
1121
|
+
|
|
1122
|
+
return expected_ckpt
|