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,1258 @@
|
|
|
1
|
+
import math
|
|
2
|
+
import os
|
|
3
|
+
import urllib.request
|
|
4
|
+
from typing import Optional, Union, Tuple, List, Dict, Any
|
|
5
|
+
|
|
6
|
+
import keras
|
|
7
|
+
from keras import layers, ops
|
|
8
|
+
from k3_node.ops.creation import repeat
|
|
9
|
+
|
|
10
|
+
|
|
11
|
+
class GraphNodeFeature(layers.Layer):
|
|
12
|
+
r"""Computes initial node representations by summing atom feature embeddings,
|
|
13
|
+
in-degree embeddings, out-degree embeddings, and prepending a learnable graph token.
|
|
14
|
+
|
|
15
|
+
Args:
|
|
16
|
+
num_atoms (int): Maximum number of atom types.
|
|
17
|
+
num_in_degree (int): Maximum in-degree value.
|
|
18
|
+
num_out_degree (int): Maximum out-degree value.
|
|
19
|
+
hidden_dim (int): Embedding dimension.
|
|
20
|
+
**kwargs: Additional layer arguments.
|
|
21
|
+
|
|
22
|
+
Example:
|
|
23
|
+
```python
|
|
24
|
+
import numpy as np
|
|
25
|
+
from k3_node.models import GraphNodeFeature
|
|
26
|
+
|
|
27
|
+
x = np.random.randint(1, 16, size=(2, 5)) # atom types of 2 graphs with 5 nodes
|
|
28
|
+
in_degree = np.random.randint(0, 10, size=(2, 5))
|
|
29
|
+
out_degree = np.random.randint(0, 10, size=(2, 5))
|
|
30
|
+
layer = GraphNodeFeature(num_atoms=16, num_in_degree=10, num_out_degree=10, hidden_dim=32)
|
|
31
|
+
print(tuple(layer(x, in_degree, out_degree).shape)) # (2, 6, 32): nodes + a virtual graph token
|
|
32
|
+
```
|
|
33
|
+
"""
|
|
34
|
+
|
|
35
|
+
def __init__(
|
|
36
|
+
self,
|
|
37
|
+
num_atoms: int,
|
|
38
|
+
num_in_degree: int,
|
|
39
|
+
num_out_degree: int,
|
|
40
|
+
hidden_dim: int,
|
|
41
|
+
**kwargs,
|
|
42
|
+
):
|
|
43
|
+
super().__init__(**kwargs)
|
|
44
|
+
self.num_atoms = num_atoms
|
|
45
|
+
self.num_in_degree = num_in_degree
|
|
46
|
+
self.num_out_degree = num_out_degree
|
|
47
|
+
self.hidden_dim = hidden_dim
|
|
48
|
+
|
|
49
|
+
# 1-indexed padding_idx=0 in fairseq / Graphormer
|
|
50
|
+
self.atom_encoder = layers.Embedding(
|
|
51
|
+
input_dim=num_atoms + 1,
|
|
52
|
+
output_dim=hidden_dim,
|
|
53
|
+
mask_zero=False,
|
|
54
|
+
name="atom_encoder",
|
|
55
|
+
)
|
|
56
|
+
self.in_degree_encoder = layers.Embedding(
|
|
57
|
+
input_dim=num_in_degree,
|
|
58
|
+
output_dim=hidden_dim,
|
|
59
|
+
mask_zero=False,
|
|
60
|
+
name="in_degree_encoder",
|
|
61
|
+
)
|
|
62
|
+
self.out_degree_encoder = layers.Embedding(
|
|
63
|
+
input_dim=num_out_degree,
|
|
64
|
+
output_dim=hidden_dim,
|
|
65
|
+
mask_zero=False,
|
|
66
|
+
name="out_degree_encoder",
|
|
67
|
+
)
|
|
68
|
+
self.graph_token = layers.Embedding(
|
|
69
|
+
input_dim=1,
|
|
70
|
+
output_dim=hidden_dim,
|
|
71
|
+
name="graph_token",
|
|
72
|
+
)
|
|
73
|
+
|
|
74
|
+
def build(self, input_shape=None):
|
|
75
|
+
if not self.built:
|
|
76
|
+
self.atom_encoder.build(None)
|
|
77
|
+
self.in_degree_encoder.build(None)
|
|
78
|
+
self.out_degree_encoder.build(None)
|
|
79
|
+
self.graph_token.build(None)
|
|
80
|
+
super().build(input_shape)
|
|
81
|
+
|
|
82
|
+
def call(self, x, in_degree, out_degree):
|
|
83
|
+
r"""
|
|
84
|
+
Args:
|
|
85
|
+
x (Tensor): Atom features of shape ``[batch_size, num_nodes, num_features]``
|
|
86
|
+
or ``[batch_size, num_nodes]``.
|
|
87
|
+
in_degree (Tensor): In-degrees of shape ``[batch_size, num_nodes]``.
|
|
88
|
+
out_degree (Tensor): Out-degrees of shape ``[batch_size, num_nodes]``.
|
|
89
|
+
|
|
90
|
+
Returns:
|
|
91
|
+
Tensor: Node features with prepended graph token of shape
|
|
92
|
+
``[batch_size, num_nodes + 1, hidden_dim]``.
|
|
93
|
+
"""
|
|
94
|
+
x_shape = ops.shape(x)
|
|
95
|
+
batch_size, num_nodes = x_shape[0], x_shape[1]
|
|
96
|
+
|
|
97
|
+
if len(ops.shape(x)) == 2:
|
|
98
|
+
x_emb = self.atom_encoder(x)
|
|
99
|
+
else:
|
|
100
|
+
# Multi-dimensional atom features: sum over feature dim
|
|
101
|
+
x_emb = ops.sum(self.atom_encoder(x), axis=-2)
|
|
102
|
+
|
|
103
|
+
node_feature = (
|
|
104
|
+
x_emb
|
|
105
|
+
+ self.in_degree_encoder(in_degree)
|
|
106
|
+
+ self.out_degree_encoder(out_degree)
|
|
107
|
+
)
|
|
108
|
+
|
|
109
|
+
# Graph token feature: shape [batch_size, 1, hidden_dim]
|
|
110
|
+
token_id = ops.zeros((batch_size, 1), dtype="int32")
|
|
111
|
+
graph_token_feature = self.graph_token(token_id)
|
|
112
|
+
|
|
113
|
+
graph_node_feature = ops.concatenate([graph_token_feature, node_feature], axis=1)
|
|
114
|
+
return graph_node_feature
|
|
115
|
+
|
|
116
|
+
|
|
117
|
+
class GraphAttnBias(layers.Layer):
|
|
118
|
+
r"""Computes the structural attention bias for each attention head from shortest path
|
|
119
|
+
distances (spatial encoding) and edge features (edge encoding).
|
|
120
|
+
|
|
121
|
+
Args:
|
|
122
|
+
num_heads (int): Number of attention heads.
|
|
123
|
+
num_atoms (int): Maximum number of atom types.
|
|
124
|
+
num_edges (int): Maximum number of edge types.
|
|
125
|
+
num_spatial (int): Maximum spatial distance.
|
|
126
|
+
num_edge_dis (int): Maximum edge distance multiplier for multi-hop encoding.
|
|
127
|
+
edge_type (str, optional): Type of edge encoding (``"multi_hop"`` or ``"single_hop"``).
|
|
128
|
+
(default: ``"multi_hop"``)
|
|
129
|
+
multi_hop_max_dist (int, optional): Maximum distance for multi-hop paths. (default: ``20``)
|
|
130
|
+
**kwargs: Additional layer arguments.
|
|
131
|
+
|
|
132
|
+
Example:
|
|
133
|
+
```python
|
|
134
|
+
import numpy as np
|
|
135
|
+
from k3_node.models import GraphAttnBias
|
|
136
|
+
|
|
137
|
+
attn_bias = np.zeros((2, 5, 5), dtype="float32") # 2 graphs, 4 nodes + graph token
|
|
138
|
+
spatial_pos = np.random.randint(0, 10, size=(2, 4, 4)) # shortest-path distances
|
|
139
|
+
x = np.zeros((2, 4, 1), dtype="int32")
|
|
140
|
+
edge_input = np.random.randint(0, 8, size=(2, 4, 4, 3, 2)) # edge types along each shortest path
|
|
141
|
+
layer = GraphAttnBias(num_heads=4, num_atoms=16, num_edges=8, num_spatial=10, num_edge_dis=5)
|
|
142
|
+
bias = layer(attn_bias=attn_bias, spatial_pos=spatial_pos, x=x, edge_input=edge_input)
|
|
143
|
+
print(tuple(bias.shape)) # (2, 4, 5, 5): one attention bias per head
|
|
144
|
+
```
|
|
145
|
+
"""
|
|
146
|
+
|
|
147
|
+
def __init__(
|
|
148
|
+
self,
|
|
149
|
+
num_heads: int,
|
|
150
|
+
num_atoms: int,
|
|
151
|
+
num_edges: int,
|
|
152
|
+
num_spatial: int,
|
|
153
|
+
num_edge_dis: int,
|
|
154
|
+
edge_type: str = "multi_hop",
|
|
155
|
+
multi_hop_max_dist: int = 20,
|
|
156
|
+
**kwargs,
|
|
157
|
+
):
|
|
158
|
+
super().__init__(**kwargs)
|
|
159
|
+
self.num_heads = num_heads
|
|
160
|
+
self.num_atoms = num_atoms
|
|
161
|
+
self.num_edges = num_edges
|
|
162
|
+
self.num_spatial = num_spatial
|
|
163
|
+
self.num_edge_dis = num_edge_dis
|
|
164
|
+
self.edge_type = edge_type
|
|
165
|
+
self.multi_hop_max_dist = multi_hop_max_dist
|
|
166
|
+
|
|
167
|
+
self.edge_encoder = layers.Embedding(
|
|
168
|
+
input_dim=num_edges + 1,
|
|
169
|
+
output_dim=num_heads,
|
|
170
|
+
name="edge_encoder",
|
|
171
|
+
)
|
|
172
|
+
if self.edge_type == "multi_hop":
|
|
173
|
+
self.edge_dis_encoder = layers.Embedding(
|
|
174
|
+
input_dim=num_edge_dis * num_heads * num_heads,
|
|
175
|
+
output_dim=1,
|
|
176
|
+
name="edge_dis_encoder",
|
|
177
|
+
)
|
|
178
|
+
self.spatial_pos_encoder = layers.Embedding(
|
|
179
|
+
input_dim=num_spatial,
|
|
180
|
+
output_dim=num_heads,
|
|
181
|
+
name="spatial_pos_encoder",
|
|
182
|
+
)
|
|
183
|
+
self.graph_token_virtual_distance = layers.Embedding(
|
|
184
|
+
input_dim=1,
|
|
185
|
+
output_dim=num_heads,
|
|
186
|
+
name="graph_token_virtual_distance",
|
|
187
|
+
)
|
|
188
|
+
|
|
189
|
+
def build(self, input_shape=None):
|
|
190
|
+
if not self.built:
|
|
191
|
+
self.edge_encoder.build(None)
|
|
192
|
+
self.spatial_pos_encoder.build(None)
|
|
193
|
+
self.graph_token_virtual_distance.build(None)
|
|
194
|
+
if hasattr(self, "edge_dis_encoder"):
|
|
195
|
+
self.edge_dis_encoder.build(None)
|
|
196
|
+
super().build(input_shape)
|
|
197
|
+
|
|
198
|
+
def call(
|
|
199
|
+
self,
|
|
200
|
+
attn_bias,
|
|
201
|
+
spatial_pos,
|
|
202
|
+
x,
|
|
203
|
+
edge_input=None,
|
|
204
|
+
attn_edge_type=None,
|
|
205
|
+
):
|
|
206
|
+
r"""
|
|
207
|
+
Args:
|
|
208
|
+
attn_bias (Tensor): Base attention mask of shape ``[batch_size, num_nodes + 1, num_nodes + 1]``.
|
|
209
|
+
spatial_pos (Tensor): Shortest path matrix of shape ``[batch_size, num_nodes, num_nodes]``.
|
|
210
|
+
x (Tensor): Atom features of shape ``[batch_size, num_nodes, ...]``.
|
|
211
|
+
edge_input (Tensor, optional): Multi-hop edge features along shortest paths of shape
|
|
212
|
+
``[batch_size, num_nodes, num_nodes, max_dist, edge_feat_dim]``.
|
|
213
|
+
attn_edge_type (Tensor, optional): Single-hop edge types of shape
|
|
214
|
+
``[batch_size, num_nodes, num_nodes, edge_feat_dim]``.
|
|
215
|
+
|
|
216
|
+
Returns:
|
|
217
|
+
Tensor: Attention bias tensor of shape ``[batch_size, num_heads, num_nodes + 1, num_nodes + 1]``.
|
|
218
|
+
"""
|
|
219
|
+
x_shape = ops.shape(x)
|
|
220
|
+
batch_size, num_nodes = x_shape[0], x_shape[1]
|
|
221
|
+
|
|
222
|
+
# [batch_size, num_heads, num_nodes + 1, num_nodes + 1]
|
|
223
|
+
graph_attn_bias = repeat(
|
|
224
|
+
ops.expand_dims(attn_bias, axis=1), repeats=self.num_heads, axis=1
|
|
225
|
+
)
|
|
226
|
+
|
|
227
|
+
# Spatial position bias: [batch_size, num_nodes, num_nodes, num_heads] -> [batch_size, num_heads, num_nodes, num_nodes]
|
|
228
|
+
spatial_pos_bias = ops.transpose(self.spatial_pos_encoder(spatial_pos), (0, 3, 1, 2))
|
|
229
|
+
|
|
230
|
+
# Update node-to-node submatrix [1:, 1:]
|
|
231
|
+
sub_bias = graph_attn_bias[:, :, 1:, 1:] + spatial_pos_bias
|
|
232
|
+
|
|
233
|
+
# Virtual distance bias for graph token: shape [1, num_heads, 1]
|
|
234
|
+
t = ops.reshape(self.graph_token_virtual_distance(ops.zeros((1,), dtype="int32")), (1, self.num_heads, 1, 1))
|
|
235
|
+
# Add to row 0 (graph token to all nodes) and column 0 (all nodes to graph token)
|
|
236
|
+
row0 = graph_attn_bias[:, :, 0:1, 1:] + t[:, :, :, 0:]
|
|
237
|
+
col0 = graph_attn_bias[:, :, 1:, 0:1] + t[:, :, 0:, :]
|
|
238
|
+
corner = graph_attn_bias[:, :, 0:1, 0:1] + t
|
|
239
|
+
|
|
240
|
+
# Edge feature bias
|
|
241
|
+
if self.edge_type == "multi_hop" and edge_input is not None:
|
|
242
|
+
spatial_pos_ = ops.copy(spatial_pos)
|
|
243
|
+
# Replace 0 with 1 for padding
|
|
244
|
+
spatial_pos_ = ops.where(ops.equal(spatial_pos_, 0), 1, spatial_pos_)
|
|
245
|
+
spatial_pos_ = ops.where(spatial_pos_ > 1, spatial_pos_ - 1, spatial_pos_)
|
|
246
|
+
if self.multi_hop_max_dist > 0:
|
|
247
|
+
spatial_pos_ = ops.clip(spatial_pos_, 0, self.multi_hop_max_dist)
|
|
248
|
+
edge_input = edge_input[:, :, :, : self.multi_hop_max_dist, :]
|
|
249
|
+
|
|
250
|
+
# edge_input: [batch_size, num_nodes, num_nodes, max_dist, edge_feat_dim]
|
|
251
|
+
# edge_encoder -> [batch_size, num_nodes, num_nodes, max_dist, edge_feat_dim, num_heads]
|
|
252
|
+
edge_enc = ops.mean(self.edge_encoder(edge_input), axis=-2)
|
|
253
|
+
max_dist = ops.shape(edge_enc)[3]
|
|
254
|
+
|
|
255
|
+
# edge_enc: [batch_size, num_nodes, num_nodes, max_dist, num_heads]
|
|
256
|
+
# permute to [max_dist, batch_size * num_nodes * num_nodes, num_heads]
|
|
257
|
+
edge_enc_perm = ops.transpose(edge_enc, (3, 0, 1, 2, 4))
|
|
258
|
+
edge_input_flat = ops.reshape(edge_enc_perm, (max_dist, -1, self.num_heads))
|
|
259
|
+
|
|
260
|
+
# Weight shape for edge_dis_encoder: [num_edge_dis * num_heads * num_heads, 1]
|
|
261
|
+
if not self.edge_dis_encoder.built:
|
|
262
|
+
self.edge_dis_encoder.build(None)
|
|
263
|
+
dis_weights = ops.reshape(
|
|
264
|
+
self.edge_dis_encoder.weights[0], (-1, self.num_heads, self.num_heads)
|
|
265
|
+
)[:max_dist, :, :]
|
|
266
|
+
|
|
267
|
+
edge_input_flat = ops.matmul(edge_input_flat, dis_weights)
|
|
268
|
+
edge_enc_back = ops.reshape(
|
|
269
|
+
edge_input_flat, (max_dist, batch_size, num_nodes, num_nodes, self.num_heads)
|
|
270
|
+
)
|
|
271
|
+
# Permute to [batch_size, num_nodes, num_nodes, max_dist, num_heads]
|
|
272
|
+
edge_enc_back = ops.transpose(edge_enc_back, (1, 2, 3, 0, 4))
|
|
273
|
+
|
|
274
|
+
# Sum over distance dim and divide by path length:
|
|
275
|
+
sp_float = ops.expand_dims(ops.cast(spatial_pos_, "float32"), axis=-1)
|
|
276
|
+
edge_bias = ops.sum(edge_enc_back, axis=3) / sp_float
|
|
277
|
+
edge_bias = ops.transpose(edge_bias, (0, 3, 1, 2))
|
|
278
|
+
sub_bias = sub_bias + edge_bias
|
|
279
|
+
elif attn_edge_type is not None:
|
|
280
|
+
edge_bias = ops.mean(self.edge_encoder(attn_edge_type), axis=-2)
|
|
281
|
+
edge_bias = ops.transpose(edge_bias, (0, 3, 1, 2))
|
|
282
|
+
sub_bias = sub_bias + edge_bias
|
|
283
|
+
|
|
284
|
+
# Reconstruct graph_attn_bias:
|
|
285
|
+
top_row = ops.concatenate([corner, row0], axis=3)
|
|
286
|
+
bottom_rows = ops.concatenate([col0, sub_bias], axis=3)
|
|
287
|
+
graph_attn_bias = ops.concatenate([top_row, bottom_rows], axis=2)
|
|
288
|
+
|
|
289
|
+
# Reset padding elements with -inf mask
|
|
290
|
+
graph_attn_bias = graph_attn_bias + ops.expand_dims(attn_bias, axis=1)
|
|
291
|
+
return graph_attn_bias
|
|
292
|
+
|
|
293
|
+
|
|
294
|
+
class GraphormerMultiheadAttention(layers.Layer):
|
|
295
|
+
r"""Multi-head self-attention layer with support for additive graph structural attention bias
|
|
296
|
+
and key padding masks.
|
|
297
|
+
|
|
298
|
+
Args:
|
|
299
|
+
embed_dim (int): Total embedding dimension.
|
|
300
|
+
num_heads (int): Number of attention heads.
|
|
301
|
+
dropout (float, optional): Attention dropout probability. (default: ``0.0``)
|
|
302
|
+
bias (bool, optional): Whether to use bias in linear projections. (default: ``True``)
|
|
303
|
+
**kwargs: Additional layer arguments.
|
|
304
|
+
|
|
305
|
+
Example:
|
|
306
|
+
```python
|
|
307
|
+
import numpy as np
|
|
308
|
+
from k3_node.models import GraphormerMultiheadAttention
|
|
309
|
+
|
|
310
|
+
x = np.random.rand(2, 6, 32).astype("float32") # [batch, tokens, embed_dim]
|
|
311
|
+
attn_bias = np.zeros((2, 4, 6, 6), dtype="float32") # structural bias per head
|
|
312
|
+
attn = GraphormerMultiheadAttention(embed_dim=32, num_heads=4, dropout=0.0)
|
|
313
|
+
out, weights = attn(x, attn_bias=attn_bias)
|
|
314
|
+
print(tuple(out.shape), tuple(weights.shape)) # (2, 6, 32) (2, 4, 6, 6)
|
|
315
|
+
```
|
|
316
|
+
"""
|
|
317
|
+
|
|
318
|
+
def __init__(
|
|
319
|
+
self,
|
|
320
|
+
embed_dim: int,
|
|
321
|
+
num_heads: int,
|
|
322
|
+
dropout: float = 0.0,
|
|
323
|
+
bias: bool = True,
|
|
324
|
+
**kwargs,
|
|
325
|
+
):
|
|
326
|
+
super().__init__(**kwargs)
|
|
327
|
+
self.embed_dim = embed_dim
|
|
328
|
+
self.num_heads = num_heads
|
|
329
|
+
self.dropout_rate = dropout
|
|
330
|
+
self.head_dim = embed_dim // num_heads
|
|
331
|
+
if self.head_dim * num_heads != embed_dim:
|
|
332
|
+
raise ValueError(f"embed_dim ({embed_dim}) must be divisible by num_heads ({num_heads}).")
|
|
333
|
+
|
|
334
|
+
self.scaling = self.head_dim ** -0.5
|
|
335
|
+
|
|
336
|
+
self.q_proj = layers.Dense(embed_dim, use_bias=bias, name="q_proj")
|
|
337
|
+
self.k_proj = layers.Dense(embed_dim, use_bias=bias, name="k_proj")
|
|
338
|
+
self.v_proj = layers.Dense(embed_dim, use_bias=bias, name="v_proj")
|
|
339
|
+
self.out_proj = layers.Dense(embed_dim, use_bias=bias, name="out_proj")
|
|
340
|
+
self.dropout = layers.Dropout(dropout)
|
|
341
|
+
|
|
342
|
+
def build(self, input_shape=None):
|
|
343
|
+
if not self.built:
|
|
344
|
+
self.q_proj.build((None, None, self.embed_dim))
|
|
345
|
+
self.k_proj.build((None, None, self.embed_dim))
|
|
346
|
+
self.v_proj.build((None, None, self.embed_dim))
|
|
347
|
+
self.out_proj.build((None, None, self.embed_dim))
|
|
348
|
+
super().build(input_shape)
|
|
349
|
+
|
|
350
|
+
def call(
|
|
351
|
+
self,
|
|
352
|
+
x,
|
|
353
|
+
attn_bias=None,
|
|
354
|
+
key_padding_mask=None,
|
|
355
|
+
training: bool = False,
|
|
356
|
+
):
|
|
357
|
+
r"""
|
|
358
|
+
Args:
|
|
359
|
+
x (Tensor): Sequence representation of shape ``[batch_size, seq_len, embed_dim]``.
|
|
360
|
+
attn_bias (Tensor, optional): Additive bias of shape
|
|
361
|
+
``[batch_size, num_heads, seq_len, seq_len]``.
|
|
362
|
+
key_padding_mask (Tensor, optional): Boolean padding mask of shape ``[batch_size, seq_len]``,
|
|
363
|
+
where True indicates padding tokens to ignore.
|
|
364
|
+
training (bool, optional): Whether in training mode. (default: ``False``)
|
|
365
|
+
|
|
366
|
+
Returns:
|
|
367
|
+
Tuple[Tensor, Tensor]: Output tensor of shape ``[batch_size, seq_len, embed_dim]``
|
|
368
|
+
and attention weights of shape ``[batch_size, num_heads, seq_len, seq_len]``.
|
|
369
|
+
"""
|
|
370
|
+
shape = ops.shape(x)
|
|
371
|
+
batch_size, seq_len = shape[0], shape[1]
|
|
372
|
+
|
|
373
|
+
q = self.q_proj(x) * self.scaling
|
|
374
|
+
k = self.k_proj(x)
|
|
375
|
+
v = self.v_proj(x)
|
|
376
|
+
|
|
377
|
+
# Reshape to [batch_size, num_heads, seq_len, head_dim]
|
|
378
|
+
q = ops.transpose(ops.reshape(q, (batch_size, seq_len, self.num_heads, self.head_dim)), (0, 2, 1, 3))
|
|
379
|
+
k = ops.transpose(ops.reshape(k, (batch_size, seq_len, self.num_heads, self.head_dim)), (0, 2, 1, 3))
|
|
380
|
+
v = ops.transpose(ops.reshape(v, (batch_size, seq_len, self.num_heads, self.head_dim)), (0, 2, 1, 3))
|
|
381
|
+
|
|
382
|
+
# [batch_size, num_heads, seq_len, seq_len]
|
|
383
|
+
attn_weights = ops.matmul(q, ops.transpose(k, (0, 1, 3, 2)))
|
|
384
|
+
|
|
385
|
+
if attn_bias is not None:
|
|
386
|
+
attn_weights = attn_weights + attn_bias
|
|
387
|
+
|
|
388
|
+
if key_padding_mask is not None:
|
|
389
|
+
# key_padding_mask: [batch_size, seq_len] -> [batch_size, 1, 1, seq_len]
|
|
390
|
+
mask = ops.expand_dims(ops.expand_dims(key_padding_mask, axis=1), axis=2)
|
|
391
|
+
attn_weights = ops.where(mask, -1e9, attn_weights)
|
|
392
|
+
|
|
393
|
+
attn_probs = ops.softmax(attn_weights, axis=-1)
|
|
394
|
+
attn_probs = self.dropout(attn_probs, training=training)
|
|
395
|
+
|
|
396
|
+
# [batch_size, num_heads, seq_len, head_dim]
|
|
397
|
+
attn = ops.matmul(attn_probs, v)
|
|
398
|
+
|
|
399
|
+
# Reshape to [batch_size, seq_len, embed_dim]
|
|
400
|
+
attn = ops.reshape(ops.transpose(attn, (0, 2, 1, 3)), (batch_size, seq_len, self.embed_dim))
|
|
401
|
+
out = self.out_proj(attn)
|
|
402
|
+
return out, attn_probs
|
|
403
|
+
|
|
404
|
+
|
|
405
|
+
class GraphormerGraphEncoderLayer(layers.Layer):
|
|
406
|
+
r"""A single Graphormer Transformer Encoder Layer, supporting Pre-LN or Post-LN,
|
|
407
|
+
multi-head attention with structural attention bias, and a 2-layer FFN.
|
|
408
|
+
|
|
409
|
+
Args:
|
|
410
|
+
embedding_dim (int, optional): Embedding dimension. (default: ``768``)
|
|
411
|
+
ffn_embedding_dim (int, optional): FFN intermediate dimension. (default: ``768``)
|
|
412
|
+
num_attention_heads (int, optional): Number of attention heads. (default: ``32``)
|
|
413
|
+
dropout (float, optional): Dropout probability. (default: ``0.1``)
|
|
414
|
+
attention_dropout (float, optional): Attention dropout. (default: ``0.1``)
|
|
415
|
+
activation_dropout (float, optional): FFN activation dropout. (default: ``0.1``)
|
|
416
|
+
activation_fn (str, optional): Activation function (``"gelu"`` or ``"relu"``). (default: ``"gelu"``)
|
|
417
|
+
pre_layernorm (bool, optional): Whether to use Pre-LN. (default: ``False``)
|
|
418
|
+
**kwargs: Additional layer arguments.
|
|
419
|
+
|
|
420
|
+
Example:
|
|
421
|
+
```python
|
|
422
|
+
import numpy as np
|
|
423
|
+
from k3_node.models import GraphormerGraphEncoderLayer
|
|
424
|
+
|
|
425
|
+
x = np.random.rand(2, 6, 32).astype("float32")
|
|
426
|
+
attn_bias = np.zeros((2, 4, 6, 6), dtype="float32")
|
|
427
|
+
layer = GraphormerGraphEncoderLayer(embedding_dim=32, ffn_embedding_dim=64, num_attention_heads=4)
|
|
428
|
+
print(tuple(layer(x, attn_bias=attn_bias).shape)) # (2, 6, 32)
|
|
429
|
+
```
|
|
430
|
+
"""
|
|
431
|
+
|
|
432
|
+
def __init__(
|
|
433
|
+
self,
|
|
434
|
+
embedding_dim: int = 768,
|
|
435
|
+
ffn_embedding_dim: int = 768,
|
|
436
|
+
num_attention_heads: int = 32,
|
|
437
|
+
dropout: float = 0.1,
|
|
438
|
+
attention_dropout: float = 0.1,
|
|
439
|
+
activation_dropout: float = 0.1,
|
|
440
|
+
activation_fn: str = "gelu",
|
|
441
|
+
pre_layernorm: bool = False,
|
|
442
|
+
**kwargs,
|
|
443
|
+
):
|
|
444
|
+
super().__init__(**kwargs)
|
|
445
|
+
self.embedding_dim = embedding_dim
|
|
446
|
+
self.ffn_embedding_dim = ffn_embedding_dim
|
|
447
|
+
self.num_attention_heads = num_attention_heads
|
|
448
|
+
self.dropout_rate = dropout
|
|
449
|
+
self.pre_layernorm = pre_layernorm
|
|
450
|
+
|
|
451
|
+
self.self_attn = GraphormerMultiheadAttention(
|
|
452
|
+
embed_dim=embedding_dim,
|
|
453
|
+
num_heads=num_attention_heads,
|
|
454
|
+
dropout=attention_dropout,
|
|
455
|
+
name="self_attn",
|
|
456
|
+
)
|
|
457
|
+
self.self_attn_layer_norm = layers.LayerNormalization(
|
|
458
|
+
epsilon=1e-5, name="self_attn_layer_norm"
|
|
459
|
+
)
|
|
460
|
+
self.dropout = layers.Dropout(dropout)
|
|
461
|
+
|
|
462
|
+
self.fc1 = layers.Dense(ffn_embedding_dim, name="fc1")
|
|
463
|
+
self.fc2 = layers.Dense(embedding_dim, name="fc2")
|
|
464
|
+
self.act_dropout = layers.Dropout(activation_dropout)
|
|
465
|
+
self.final_layer_norm = layers.LayerNormalization(
|
|
466
|
+
epsilon=1e-5, name="final_layer_norm"
|
|
467
|
+
)
|
|
468
|
+
|
|
469
|
+
if activation_fn == "gelu":
|
|
470
|
+
self.activation = ops.gelu
|
|
471
|
+
elif activation_fn == "relu":
|
|
472
|
+
self.activation = ops.relu
|
|
473
|
+
else:
|
|
474
|
+
self.activation = keras.activations.get(activation_fn)
|
|
475
|
+
|
|
476
|
+
def build(self, input_shape=None):
|
|
477
|
+
if not self.built:
|
|
478
|
+
self.self_attn.build((None, None, self.embedding_dim))
|
|
479
|
+
self.self_attn_layer_norm.build((None, None, self.embedding_dim))
|
|
480
|
+
self.fc1.build((None, None, self.embedding_dim))
|
|
481
|
+
self.fc2.build((None, None, self.ffn_embedding_dim))
|
|
482
|
+
self.final_layer_norm.build((None, None, self.embedding_dim))
|
|
483
|
+
super().build(input_shape)
|
|
484
|
+
|
|
485
|
+
def call(
|
|
486
|
+
self,
|
|
487
|
+
x,
|
|
488
|
+
attn_bias=None,
|
|
489
|
+
key_padding_mask=None,
|
|
490
|
+
training: bool = False,
|
|
491
|
+
):
|
|
492
|
+
residual = x
|
|
493
|
+
if self.pre_layernorm:
|
|
494
|
+
x = self.self_attn_layer_norm(x)
|
|
495
|
+
|
|
496
|
+
x, _ = self.self_attn(
|
|
497
|
+
x,
|
|
498
|
+
attn_bias=attn_bias,
|
|
499
|
+
key_padding_mask=key_padding_mask,
|
|
500
|
+
training=training,
|
|
501
|
+
)
|
|
502
|
+
x = self.dropout(x, training=training)
|
|
503
|
+
x = residual + x
|
|
504
|
+
|
|
505
|
+
if not self.pre_layernorm:
|
|
506
|
+
x = self.self_attn_layer_norm(x)
|
|
507
|
+
|
|
508
|
+
residual = x
|
|
509
|
+
if self.pre_layernorm:
|
|
510
|
+
x = self.final_layer_norm(x)
|
|
511
|
+
|
|
512
|
+
x = self.fc2(self.act_dropout(self.activation(self.fc1(x)), training=training))
|
|
513
|
+
x = self.dropout(x, training=training)
|
|
514
|
+
x = residual + x
|
|
515
|
+
|
|
516
|
+
if not self.pre_layernorm:
|
|
517
|
+
x = self.final_layer_norm(x)
|
|
518
|
+
|
|
519
|
+
return x
|
|
520
|
+
|
|
521
|
+
|
|
522
|
+
class GraphormerGraphEncoder(layers.Layer):
|
|
523
|
+
r"""Graphormer Graph Encoder stack consisting of graph node features, structural attention bias,
|
|
524
|
+
and multiple stacked encoder layers.
|
|
525
|
+
|
|
526
|
+
Args:
|
|
527
|
+
num_atoms (int): Maximum number of atom types.
|
|
528
|
+
num_in_degree (int): Maximum in-degree value.
|
|
529
|
+
num_out_degree (int): Maximum out-degree value.
|
|
530
|
+
num_edges (int): Maximum number of edge types.
|
|
531
|
+
num_spatial (int): Maximum spatial distance.
|
|
532
|
+
num_edge_dis (int): Maximum edge distance multiplier.
|
|
533
|
+
edge_type (str, optional): Edge encoding mode (``"multi_hop"`` or ``"single_hop"``). (default: ``"multi_hop"``)
|
|
534
|
+
multi_hop_max_dist (int, optional): Max distance for multi-hop. (default: ``20``)
|
|
535
|
+
num_encoder_layers (int, optional): Number of encoder layers. (default: ``12``)
|
|
536
|
+
embedding_dim (int, optional): Embedding dimension. (default: ``768``)
|
|
537
|
+
ffn_embedding_dim (int, optional): FFN intermediate dimension. (default: ``768``)
|
|
538
|
+
num_attention_heads (int, optional): Number of attention heads. (default: ``32``)
|
|
539
|
+
dropout (float, optional): Dropout probability. (default: ``0.1``)
|
|
540
|
+
attention_dropout (float, optional): Attention dropout. (default: ``0.1``)
|
|
541
|
+
activation_dropout (float, optional): FFN activation dropout. (default: ``0.1``)
|
|
542
|
+
pre_layernorm (bool, optional): Whether to use Pre-LN. (default: ``False``)
|
|
543
|
+
encoder_normalize_before (bool, optional): Whether to normalize before encoder blocks. (default: ``False``)
|
|
544
|
+
activation_fn (str, optional): Activation function name. (default: ``"gelu"``)
|
|
545
|
+
**kwargs: Additional layer arguments.
|
|
546
|
+
|
|
547
|
+
Example:
|
|
548
|
+
```python
|
|
549
|
+
import numpy as np
|
|
550
|
+
from k3_node.models import GraphormerGraphEncoder
|
|
551
|
+
|
|
552
|
+
# A batch of 2 graphs padded to 4 nodes, preprocessed into Graphormer's dense inputs
|
|
553
|
+
data = {
|
|
554
|
+
"x": np.random.randint(1, 15, size=(2, 4, 2)), # 2 categorical features per node
|
|
555
|
+
"in_degree": np.random.randint(0, 7, size=(2, 4)),
|
|
556
|
+
"out_degree": np.random.randint(0, 7, size=(2, 4)),
|
|
557
|
+
"attn_bias": np.zeros((2, 5, 5), dtype="float32"), # +1 for the virtual graph token
|
|
558
|
+
"spatial_pos": np.random.randint(0, 7, size=(2, 4, 4)), # shortest-path distances
|
|
559
|
+
"edge_input": np.random.randint(0, 7, size=(2, 4, 4, 2, 2)), # edge types along shortest paths
|
|
560
|
+
}
|
|
561
|
+
|
|
562
|
+
encoder = GraphormerGraphEncoder(num_atoms=16, num_in_degree=8, num_out_degree=8, num_edges=8,
|
|
563
|
+
num_spatial=8, num_edge_dis=4, num_encoder_layers=2,
|
|
564
|
+
embedding_dim=32, ffn_embedding_dim=64, num_attention_heads=4)
|
|
565
|
+
node_states, graph_rep = encoder(**data)
|
|
566
|
+
print(tuple(node_states.shape), tuple(graph_rep.shape)) # (2, 5, 32) (2, 32): per-token states, graph-token embedding
|
|
567
|
+
```
|
|
568
|
+
"""
|
|
569
|
+
|
|
570
|
+
def __init__(
|
|
571
|
+
self,
|
|
572
|
+
num_atoms: int = 512,
|
|
573
|
+
num_in_degree: int = 512,
|
|
574
|
+
num_out_degree: int = 512,
|
|
575
|
+
num_edges: int = 512,
|
|
576
|
+
num_spatial: int = 512,
|
|
577
|
+
num_edge_dis: int = 128,
|
|
578
|
+
edge_type: str = "multi_hop",
|
|
579
|
+
multi_hop_max_dist: int = 20,
|
|
580
|
+
num_encoder_layers: int = 12,
|
|
581
|
+
embedding_dim: int = 768,
|
|
582
|
+
ffn_embedding_dim: int = 768,
|
|
583
|
+
num_attention_heads: int = 32,
|
|
584
|
+
dropout: float = 0.1,
|
|
585
|
+
attention_dropout: float = 0.1,
|
|
586
|
+
activation_dropout: float = 0.1,
|
|
587
|
+
pre_layernorm: bool = False,
|
|
588
|
+
encoder_normalize_before: bool = False,
|
|
589
|
+
activation_fn: str = "gelu",
|
|
590
|
+
**kwargs,
|
|
591
|
+
):
|
|
592
|
+
super().__init__(**kwargs)
|
|
593
|
+
self.embedding_dim = embedding_dim
|
|
594
|
+
self.num_encoder_layers = num_encoder_layers
|
|
595
|
+
self.pre_layernorm = pre_layernorm
|
|
596
|
+
|
|
597
|
+
self.graph_node_feature = GraphNodeFeature(
|
|
598
|
+
num_atoms=num_atoms,
|
|
599
|
+
num_in_degree=num_in_degree,
|
|
600
|
+
num_out_degree=num_out_degree,
|
|
601
|
+
hidden_dim=embedding_dim,
|
|
602
|
+
name="graph_node_feature",
|
|
603
|
+
)
|
|
604
|
+
self.graph_attn_bias = GraphAttnBias(
|
|
605
|
+
num_heads=num_attention_heads,
|
|
606
|
+
num_atoms=num_atoms,
|
|
607
|
+
num_edges=num_edges,
|
|
608
|
+
num_spatial=num_spatial,
|
|
609
|
+
num_edge_dis=num_edge_dis,
|
|
610
|
+
edge_type=edge_type,
|
|
611
|
+
multi_hop_max_dist=multi_hop_max_dist,
|
|
612
|
+
name="graph_attn_bias",
|
|
613
|
+
)
|
|
614
|
+
if encoder_normalize_before:
|
|
615
|
+
self.emb_layer_norm = layers.LayerNormalization(epsilon=1e-5, name="emb_layer_norm")
|
|
616
|
+
else:
|
|
617
|
+
self.emb_layer_norm = None
|
|
618
|
+
|
|
619
|
+
self.dropout = layers.Dropout(dropout)
|
|
620
|
+
|
|
621
|
+
self.encoder_layers = [
|
|
622
|
+
GraphormerGraphEncoderLayer(
|
|
623
|
+
embedding_dim=embedding_dim,
|
|
624
|
+
ffn_embedding_dim=ffn_embedding_dim,
|
|
625
|
+
num_attention_heads=num_attention_heads,
|
|
626
|
+
dropout=dropout,
|
|
627
|
+
attention_dropout=attention_dropout,
|
|
628
|
+
activation_dropout=activation_dropout,
|
|
629
|
+
activation_fn=activation_fn,
|
|
630
|
+
pre_layernorm=pre_layernorm,
|
|
631
|
+
name=f"layers_{i}",
|
|
632
|
+
)
|
|
633
|
+
for i in range(num_encoder_layers)
|
|
634
|
+
]
|
|
635
|
+
|
|
636
|
+
if pre_layernorm:
|
|
637
|
+
self.final_layer_norm = layers.LayerNormalization(epsilon=1e-5, name="final_layer_norm")
|
|
638
|
+
else:
|
|
639
|
+
self.final_layer_norm = None
|
|
640
|
+
|
|
641
|
+
def build(self, input_shape=None):
|
|
642
|
+
if not self.built:
|
|
643
|
+
self.graph_node_feature.build(None)
|
|
644
|
+
self.graph_attn_bias.build(None)
|
|
645
|
+
if self.emb_layer_norm is not None:
|
|
646
|
+
self.emb_layer_norm.build((None, None, self.embedding_dim))
|
|
647
|
+
for layer in self.encoder_layers:
|
|
648
|
+
layer.build((None, None, self.embedding_dim))
|
|
649
|
+
if self.final_layer_norm is not None:
|
|
650
|
+
self.final_layer_norm.build((None, None, self.embedding_dim))
|
|
651
|
+
super().build(input_shape)
|
|
652
|
+
|
|
653
|
+
def call(
|
|
654
|
+
self,
|
|
655
|
+
x,
|
|
656
|
+
in_degree,
|
|
657
|
+
out_degree,
|
|
658
|
+
attn_bias,
|
|
659
|
+
spatial_pos,
|
|
660
|
+
edge_input=None,
|
|
661
|
+
attn_edge_type=None,
|
|
662
|
+
perturb=None,
|
|
663
|
+
training: bool = False,
|
|
664
|
+
):
|
|
665
|
+
batch_size = ops.shape(x)[0]
|
|
666
|
+
# Compute padding mask: [batch_size, num_nodes]
|
|
667
|
+
if len(ops.shape(x)) == 3:
|
|
668
|
+
raw_mask = ops.equal(x[:, :, 0], 0)
|
|
669
|
+
else:
|
|
670
|
+
raw_mask = ops.equal(x, 0)
|
|
671
|
+
|
|
672
|
+
# Prepend False for graph token: [batch_size, num_nodes + 1]
|
|
673
|
+
cls_mask = ops.zeros((batch_size, 1), dtype="bool")
|
|
674
|
+
padding_mask = ops.concatenate([cls_mask, raw_mask], axis=1)
|
|
675
|
+
|
|
676
|
+
# Node features: [batch_size, num_nodes + 1, embedding_dim]
|
|
677
|
+
h = self.graph_node_feature(x, in_degree, out_degree)
|
|
678
|
+
if perturb is not None:
|
|
679
|
+
# perturb is added to non-token nodes
|
|
680
|
+
h_token = h[:, 0:1, :]
|
|
681
|
+
h_nodes = h[:, 1:, :] + perturb
|
|
682
|
+
h = ops.concatenate([h_token, h_nodes], axis=1)
|
|
683
|
+
|
|
684
|
+
bias = self.graph_attn_bias(
|
|
685
|
+
attn_bias=attn_bias,
|
|
686
|
+
spatial_pos=spatial_pos,
|
|
687
|
+
x=x,
|
|
688
|
+
edge_input=edge_input,
|
|
689
|
+
attn_edge_type=attn_edge_type,
|
|
690
|
+
)
|
|
691
|
+
|
|
692
|
+
if self.emb_layer_norm is not None:
|
|
693
|
+
h = self.emb_layer_norm(h)
|
|
694
|
+
|
|
695
|
+
h = self.dropout(h, training=training)
|
|
696
|
+
|
|
697
|
+
for layer in self.encoder_layers:
|
|
698
|
+
h = layer(
|
|
699
|
+
h,
|
|
700
|
+
attn_bias=bias,
|
|
701
|
+
key_padding_mask=padding_mask,
|
|
702
|
+
training=training,
|
|
703
|
+
)
|
|
704
|
+
|
|
705
|
+
if self.final_layer_norm is not None:
|
|
706
|
+
h = self.final_layer_norm(h)
|
|
707
|
+
|
|
708
|
+
graph_rep = h[:, 0, :]
|
|
709
|
+
return h, graph_rep
|
|
710
|
+
|
|
711
|
+
|
|
712
|
+
class Graphormer(keras.Model):
|
|
713
|
+
r"""Graphormer model for molecular graph representation and property prediction
|
|
714
|
+
from `"Do Transformers Really Perform Badly for Graph Representation?" <https://arxiv.org/abs/2106.05234>`_.
|
|
715
|
+
|
|
716
|
+
Args:
|
|
717
|
+
num_atoms (int, optional): Maximum atom vocabulary size. (default: ``512``)
|
|
718
|
+
num_in_degree (int, optional): Maximum in-degree. (default: ``512``)
|
|
719
|
+
num_out_degree (int, optional): Maximum out-degree. (default: ``512``)
|
|
720
|
+
num_edges (int, optional): Maximum edge vocabulary size. (default: ``512``)
|
|
721
|
+
num_spatial (int, optional): Maximum spatial distance. (default: ``512``)
|
|
722
|
+
num_edge_dis (int, optional): Maximum edge distance multiplier. (default: ``128``)
|
|
723
|
+
edge_type (str, optional): Type of edge encoding (``"multi_hop"`` or ``"single_hop"``). (default: ``"multi_hop"``)
|
|
724
|
+
multi_hop_max_dist (int, optional): Max distance for multi-hop paths. (default: ``20``)
|
|
725
|
+
num_encoder_layers (int, optional): Number of Transformer layers. (default: ``12``)
|
|
726
|
+
embedding_dim (int, optional): Hidden embedding dimension. (default: ``768``)
|
|
727
|
+
ffn_embedding_dim (int, optional): FFN intermediate dimension. (default: ``768``)
|
|
728
|
+
num_attention_heads (int, optional): Number of attention heads. (default: ``32``)
|
|
729
|
+
dropout (float, optional): Dropout rate. (default: ``0.0``)
|
|
730
|
+
attention_dropout (float, optional): Attention dropout rate. (default: ``0.1``)
|
|
731
|
+
activation_dropout (float, optional): Activation dropout rate. (default: ``0.1``)
|
|
732
|
+
encoder_normalize_before (bool, optional): Whether to apply LayerNorm before encoder blocks. (default: ``True``)
|
|
733
|
+
pre_layernorm (bool, optional): Whether to use Pre-LN. (default: ``False``)
|
|
734
|
+
num_classes (int, optional): Output dimension for prediction head. (default: ``1``)
|
|
735
|
+
activation_fn (str, optional): Activation function. (default: ``"gelu"``)
|
|
736
|
+
**kwargs: Additional model arguments.
|
|
737
|
+
|
|
738
|
+
Example:
|
|
739
|
+
```python
|
|
740
|
+
import numpy as np
|
|
741
|
+
from k3_node.models import Graphormer
|
|
742
|
+
|
|
743
|
+
# A batch of 2 graphs padded to 4 nodes, preprocessed into Graphormer's dense inputs
|
|
744
|
+
data = {
|
|
745
|
+
"x": np.random.randint(1, 15, size=(2, 4, 2)), # 2 categorical features per node
|
|
746
|
+
"in_degree": np.random.randint(0, 7, size=(2, 4)),
|
|
747
|
+
"out_degree": np.random.randint(0, 7, size=(2, 4)),
|
|
748
|
+
"attn_bias": np.zeros((2, 5, 5), dtype="float32"), # +1 for the virtual graph token
|
|
749
|
+
"spatial_pos": np.random.randint(0, 7, size=(2, 4, 4)), # shortest-path distances
|
|
750
|
+
"edge_input": np.random.randint(0, 7, size=(2, 4, 4, 2, 2)), # edge types along shortest paths
|
|
751
|
+
}
|
|
752
|
+
|
|
753
|
+
model = Graphormer(num_atoms=16, num_in_degree=8, num_out_degree=8, num_edges=8, num_spatial=8,
|
|
754
|
+
num_edge_dis=4, num_encoder_layers=2, embedding_dim=32, ffn_embedding_dim=64,
|
|
755
|
+
num_attention_heads=4, num_classes=1)
|
|
756
|
+
out = model(data) # one prediction per graph
|
|
757
|
+
print(tuple(out.shape)) # (2, 1)
|
|
758
|
+
```
|
|
759
|
+
"""
|
|
760
|
+
|
|
761
|
+
def __init__(
|
|
762
|
+
self,
|
|
763
|
+
num_atoms: int = 512,
|
|
764
|
+
num_in_degree: int = 512,
|
|
765
|
+
num_out_degree: int = 512,
|
|
766
|
+
num_edges: int = 512,
|
|
767
|
+
num_spatial: int = 512,
|
|
768
|
+
num_edge_dis: int = 128,
|
|
769
|
+
edge_type: str = "multi_hop",
|
|
770
|
+
multi_hop_max_dist: int = 20,
|
|
771
|
+
num_encoder_layers: int = 12,
|
|
772
|
+
embedding_dim: int = 768,
|
|
773
|
+
ffn_embedding_dim: int = 768,
|
|
774
|
+
num_attention_heads: int = 32,
|
|
775
|
+
dropout: float = 0.0,
|
|
776
|
+
attention_dropout: float = 0.1,
|
|
777
|
+
activation_dropout: float = 0.1,
|
|
778
|
+
encoder_normalize_before: bool = True,
|
|
779
|
+
pre_layernorm: bool = False,
|
|
780
|
+
num_classes: int = 1,
|
|
781
|
+
activation_fn: str = "gelu",
|
|
782
|
+
**kwargs,
|
|
783
|
+
):
|
|
784
|
+
super().__init__(**kwargs)
|
|
785
|
+
self.num_atoms = num_atoms
|
|
786
|
+
self.num_in_degree = num_in_degree
|
|
787
|
+
self.num_out_degree = num_out_degree
|
|
788
|
+
self.num_edges = num_edges
|
|
789
|
+
self.num_spatial = num_spatial
|
|
790
|
+
self.num_edge_dis = num_edge_dis
|
|
791
|
+
self.edge_type = edge_type
|
|
792
|
+
self.multi_hop_max_dist = multi_hop_max_dist
|
|
793
|
+
self.num_encoder_layers = num_encoder_layers
|
|
794
|
+
self.embedding_dim = embedding_dim
|
|
795
|
+
self.ffn_embedding_dim = ffn_embedding_dim
|
|
796
|
+
self.num_attention_heads = num_attention_heads
|
|
797
|
+
self.num_classes = num_classes
|
|
798
|
+
self.pre_layernorm = pre_layernorm
|
|
799
|
+
|
|
800
|
+
self.graph_encoder = GraphormerGraphEncoder(
|
|
801
|
+
num_atoms=num_atoms,
|
|
802
|
+
num_in_degree=num_in_degree,
|
|
803
|
+
num_out_degree=num_out_degree,
|
|
804
|
+
num_edges=num_edges,
|
|
805
|
+
num_spatial=num_spatial,
|
|
806
|
+
num_edge_dis=num_edge_dis,
|
|
807
|
+
edge_type=edge_type,
|
|
808
|
+
multi_hop_max_dist=multi_hop_max_dist,
|
|
809
|
+
num_encoder_layers=num_encoder_layers,
|
|
810
|
+
embedding_dim=embedding_dim,
|
|
811
|
+
ffn_embedding_dim=ffn_embedding_dim,
|
|
812
|
+
num_attention_heads=num_attention_heads,
|
|
813
|
+
dropout=dropout,
|
|
814
|
+
attention_dropout=attention_dropout,
|
|
815
|
+
activation_dropout=activation_dropout,
|
|
816
|
+
encoder_normalize_before=encoder_normalize_before,
|
|
817
|
+
pre_layernorm=pre_layernorm,
|
|
818
|
+
activation_fn=activation_fn,
|
|
819
|
+
name="graph_encoder",
|
|
820
|
+
)
|
|
821
|
+
|
|
822
|
+
# Output prediction head (matching GraphormerEncoder in Graphormer-main)
|
|
823
|
+
self.lm_head_transform_weight = layers.Dense(
|
|
824
|
+
embedding_dim, name="lm_head_transform_weight"
|
|
825
|
+
)
|
|
826
|
+
self.layer_norm = layers.LayerNormalization(epsilon=1e-5, name="layer_norm")
|
|
827
|
+
self.embed_out = layers.Dense(num_classes, use_bias=False, name="embed_out")
|
|
828
|
+
|
|
829
|
+
if activation_fn == "gelu":
|
|
830
|
+
self.head_act = ops.gelu
|
|
831
|
+
elif activation_fn == "relu":
|
|
832
|
+
self.head_act = ops.relu
|
|
833
|
+
else:
|
|
834
|
+
self.head_act = keras.activations.get(activation_fn)
|
|
835
|
+
|
|
836
|
+
def build(self, input_shape=None):
|
|
837
|
+
if not self.built:
|
|
838
|
+
self.graph_encoder.build(None)
|
|
839
|
+
self.lm_head_transform_weight.build((None, self.embedding_dim))
|
|
840
|
+
self.layer_norm.build((None, self.embedding_dim))
|
|
841
|
+
self.embed_out.build((None, self.embedding_dim))
|
|
842
|
+
self.lm_output_learned_bias = self.add_weight(
|
|
843
|
+
shape=(1,),
|
|
844
|
+
initializer="zeros",
|
|
845
|
+
trainable=True,
|
|
846
|
+
name="lm_output_learned_bias",
|
|
847
|
+
)
|
|
848
|
+
super().build(input_shape)
|
|
849
|
+
|
|
850
|
+
def call(
|
|
851
|
+
self,
|
|
852
|
+
batched_data: Optional[Dict[str, Any]] = None,
|
|
853
|
+
x=None,
|
|
854
|
+
in_degree=None,
|
|
855
|
+
out_degree=None,
|
|
856
|
+
attn_bias=None,
|
|
857
|
+
spatial_pos=None,
|
|
858
|
+
edge_input=None,
|
|
859
|
+
attn_edge_type=None,
|
|
860
|
+
perturb=None,
|
|
861
|
+
return_all: bool = False,
|
|
862
|
+
training: bool = False,
|
|
863
|
+
):
|
|
864
|
+
r"""Forward pass for Graphormer.
|
|
865
|
+
|
|
866
|
+
Accepts either a single dictionary ``batched_data`` containing graph tensors,
|
|
867
|
+
or individual tensor arguments.
|
|
868
|
+
|
|
869
|
+
Returns:
|
|
870
|
+
Tensor or Tuple[Tensor, Tensor]: Graph prediction tensor of shape ``[batch_size, num_classes]``,
|
|
871
|
+
or if return_all=True, a tuple of ``(graph_pred, all_node_features)``.
|
|
872
|
+
"""
|
|
873
|
+
if batched_data is not None:
|
|
874
|
+
x = batched_data["x"]
|
|
875
|
+
in_degree = batched_data["in_degree"]
|
|
876
|
+
out_degree = batched_data["out_degree"]
|
|
877
|
+
attn_bias = batched_data["attn_bias"]
|
|
878
|
+
spatial_pos = batched_data["spatial_pos"]
|
|
879
|
+
edge_input = batched_data.get("edge_input", None)
|
|
880
|
+
attn_edge_type = batched_data.get("attn_edge_type", None)
|
|
881
|
+
|
|
882
|
+
h, graph_rep = self.graph_encoder(
|
|
883
|
+
x=x,
|
|
884
|
+
in_degree=in_degree,
|
|
885
|
+
out_degree=out_degree,
|
|
886
|
+
attn_bias=attn_bias,
|
|
887
|
+
spatial_pos=spatial_pos,
|
|
888
|
+
edge_input=edge_input,
|
|
889
|
+
attn_edge_type=attn_edge_type,
|
|
890
|
+
perturb=perturb,
|
|
891
|
+
training=training,
|
|
892
|
+
)
|
|
893
|
+
|
|
894
|
+
# Output projection on graph token: shape [batch_size, embedding_dim]
|
|
895
|
+
token_h = self.layer_norm(self.head_act(self.lm_head_transform_weight(graph_rep)))
|
|
896
|
+
out = self.embed_out(token_h)
|
|
897
|
+
if hasattr(self, "lm_output_learned_bias") and self.lm_output_learned_bias is not None:
|
|
898
|
+
out = out + self.lm_output_learned_bias
|
|
899
|
+
|
|
900
|
+
if return_all:
|
|
901
|
+
return out, h
|
|
902
|
+
return out
|
|
903
|
+
|
|
904
|
+
def embed(
|
|
905
|
+
self,
|
|
906
|
+
batched_data: Optional[Dict[str, Any]] = None,
|
|
907
|
+
x=None,
|
|
908
|
+
in_degree=None,
|
|
909
|
+
out_degree=None,
|
|
910
|
+
attn_bias=None,
|
|
911
|
+
spatial_pos=None,
|
|
912
|
+
edge_input=None,
|
|
913
|
+
attn_edge_type=None,
|
|
914
|
+
):
|
|
915
|
+
r"""Computes graph and node embeddings without applying the prediction head."""
|
|
916
|
+
if batched_data is not None:
|
|
917
|
+
x = batched_data["x"]
|
|
918
|
+
in_degree = batched_data["in_degree"]
|
|
919
|
+
out_degree = batched_data["out_degree"]
|
|
920
|
+
attn_bias = batched_data["attn_bias"]
|
|
921
|
+
spatial_pos = batched_data["spatial_pos"]
|
|
922
|
+
edge_input = batched_data.get("edge_input", None)
|
|
923
|
+
attn_edge_type = batched_data.get("attn_edge_type", None)
|
|
924
|
+
|
|
925
|
+
h, graph_rep = self.graph_encoder(
|
|
926
|
+
x=x,
|
|
927
|
+
in_degree=in_degree,
|
|
928
|
+
out_degree=out_degree,
|
|
929
|
+
attn_bias=attn_bias,
|
|
930
|
+
spatial_pos=spatial_pos,
|
|
931
|
+
edge_input=edge_input,
|
|
932
|
+
attn_edge_type=attn_edge_type,
|
|
933
|
+
training=False,
|
|
934
|
+
)
|
|
935
|
+
return graph_rep
|
|
936
|
+
|
|
937
|
+
@classmethod
|
|
938
|
+
def from_pretrained(
|
|
939
|
+
cls,
|
|
940
|
+
pretrained_name: str = "pcqm4mv1_graphormer_base",
|
|
941
|
+
folder: str = "checkpoints",
|
|
942
|
+
download: bool = True,
|
|
943
|
+
**kwargs,
|
|
944
|
+
) -> "Graphormer":
|
|
945
|
+
r"""Instantiates a Graphormer model with pre-trained weights."""
|
|
946
|
+
cfg = get_graphormer_config(pretrained_name)
|
|
947
|
+
cfg.update(kwargs)
|
|
948
|
+
model = cls(**cfg)
|
|
949
|
+
load_graphormer_weights(model, pretrained_name=pretrained_name, folder=folder, download=download)
|
|
950
|
+
return model
|
|
951
|
+
|
|
952
|
+
def __repr__(self) -> str:
|
|
953
|
+
return (
|
|
954
|
+
f"{self.__class__.__name__}("
|
|
955
|
+
f"embedding_dim={self.embedding_dim}, "
|
|
956
|
+
f"num_encoder_layers={self.num_encoder_layers}, "
|
|
957
|
+
f"num_attention_heads={self.num_attention_heads}, "
|
|
958
|
+
f"num_classes={self.num_classes})"
|
|
959
|
+
)
|
|
960
|
+
|
|
961
|
+
|
|
962
|
+
# =========================================================================
|
|
963
|
+
# Configuration presets & Pre-trained URLs
|
|
964
|
+
# =========================================================================
|
|
965
|
+
|
|
966
|
+
PRETRAINED_MODEL_URLS = {
|
|
967
|
+
"pcqm4mv1_graphormer_base": "https://huggingface.co/clefourrier/graphormer-base-pcqm4mv1/resolve/main/pytorch_model.bin",
|
|
968
|
+
"pcqm4mv2_graphormer_base": "https://huggingface.co/clefourrier/graphormer-base-pcqm4mv2/resolve/main/pytorch_model.bin",
|
|
969
|
+
"pcqm4mv1_graphormer_base_for_molhiv": "https://ml2md.blob.core.windows.net/graphormer-ckpts/checkpoint_base_preln_pcqm4mv1_for_hiv.pt",
|
|
970
|
+
}
|
|
971
|
+
|
|
972
|
+
LEGACY_URLS = {
|
|
973
|
+
"pcqm4mv1_graphormer_base": "https://ml2md.blob.core.windows.net/graphormer-ckpts/checkpoint_best_pcqm4mv1.pt",
|
|
974
|
+
"pcqm4mv2_graphormer_base": "https://ml2md.blob.core.windows.net/graphormer-ckpts/checkpoint_best_pcqm4mv2.pt",
|
|
975
|
+
}
|
|
976
|
+
|
|
977
|
+
|
|
978
|
+
def get_graphormer_config(name_or_variant: str) -> Dict[str, Any]:
|
|
979
|
+
r"""Returns configuration dictionary for known Graphormer architectures."""
|
|
980
|
+
key = name_or_variant.lower().replace("-", "_")
|
|
981
|
+
if "slim" in key:
|
|
982
|
+
return {
|
|
983
|
+
"num_encoder_layers": 12,
|
|
984
|
+
"num_attention_heads": 8,
|
|
985
|
+
"embedding_dim": 80,
|
|
986
|
+
"ffn_embedding_dim": 80,
|
|
987
|
+
"pre_layernorm": False,
|
|
988
|
+
"encoder_normalize_before": True,
|
|
989
|
+
"dropout": 0.0,
|
|
990
|
+
"attention_dropout": 0.1,
|
|
991
|
+
"activation_dropout": 0.1,
|
|
992
|
+
}
|
|
993
|
+
elif "large" in key:
|
|
994
|
+
return {
|
|
995
|
+
"num_encoder_layers": 24,
|
|
996
|
+
"num_attention_heads": 32,
|
|
997
|
+
"embedding_dim": 1024,
|
|
998
|
+
"ffn_embedding_dim": 1024,
|
|
999
|
+
"pre_layernorm": False,
|
|
1000
|
+
"encoder_normalize_before": True,
|
|
1001
|
+
"dropout": 0.0,
|
|
1002
|
+
"attention_dropout": 0.1,
|
|
1003
|
+
"activation_dropout": 0.1,
|
|
1004
|
+
}
|
|
1005
|
+
elif "base" in key or "pcqm4m" in key:
|
|
1006
|
+
pre_ln = "molhiv" in key or "preln" in key
|
|
1007
|
+
num_classes = 1
|
|
1008
|
+
if "molhiv" in key:
|
|
1009
|
+
num_classes = 2
|
|
1010
|
+
return {
|
|
1011
|
+
"num_encoder_layers": 12,
|
|
1012
|
+
"num_attention_heads": 32,
|
|
1013
|
+
"embedding_dim": 768,
|
|
1014
|
+
"ffn_embedding_dim": 768,
|
|
1015
|
+
"pre_layernorm": pre_ln,
|
|
1016
|
+
"encoder_normalize_before": True,
|
|
1017
|
+
"dropout": 0.0,
|
|
1018
|
+
"attention_dropout": 0.1,
|
|
1019
|
+
"activation_dropout": 0.1,
|
|
1020
|
+
"num_classes": num_classes,
|
|
1021
|
+
}
|
|
1022
|
+
else:
|
|
1023
|
+
# Standard default architecture
|
|
1024
|
+
return {
|
|
1025
|
+
"num_encoder_layers": 6,
|
|
1026
|
+
"num_attention_heads": 8,
|
|
1027
|
+
"embedding_dim": 1024,
|
|
1028
|
+
"ffn_embedding_dim": 4096,
|
|
1029
|
+
"pre_layernorm": False,
|
|
1030
|
+
"encoder_normalize_before": True,
|
|
1031
|
+
"dropout": 0.1,
|
|
1032
|
+
"attention_dropout": 0.1,
|
|
1033
|
+
"activation_dropout": 0.0,
|
|
1034
|
+
"num_classes": 1,
|
|
1035
|
+
}
|
|
1036
|
+
|
|
1037
|
+
|
|
1038
|
+
def download_graphormer_checkpoint(
|
|
1039
|
+
name: str,
|
|
1040
|
+
folder: str = "checkpoints",
|
|
1041
|
+
log: bool = True,
|
|
1042
|
+
) -> str:
|
|
1043
|
+
r"""Downloads a pre-trained Graphormer checkpoint to the specified folder.
|
|
1044
|
+
|
|
1045
|
+
Args:
|
|
1046
|
+
name (str): Pretrained checkpoint name (e.g., ``"pcqm4mv1_graphormer_base"``,
|
|
1047
|
+
``"pcqm4mv2_graphormer_base"``).
|
|
1048
|
+
folder (str, optional): Destination folder. (default: ``"checkpoints"``)
|
|
1049
|
+
log (bool, optional): Whether to print download progress. (default: ``True``)
|
|
1050
|
+
|
|
1051
|
+
Returns:
|
|
1052
|
+
str: Absolute path to the downloaded file.
|
|
1053
|
+
"""
|
|
1054
|
+
from k3_node.data.download import download_url
|
|
1055
|
+
|
|
1056
|
+
clean_name = name.lower().replace("-", "_")
|
|
1057
|
+
if clean_name not in PRETRAINED_MODEL_URLS and name not in PRETRAINED_MODEL_URLS:
|
|
1058
|
+
raise ValueError(
|
|
1059
|
+
f"Unknown pretrained model name '{name}'. Available: {list(PRETRAINED_MODEL_URLS.keys())}"
|
|
1060
|
+
)
|
|
1061
|
+
|
|
1062
|
+
matched_key = clean_name if clean_name in PRETRAINED_MODEL_URLS else name
|
|
1063
|
+
url = PRETRAINED_MODEL_URLS[matched_key]
|
|
1064
|
+
filename = f"{matched_key}.pt"
|
|
1065
|
+
local_path = os.path.join(folder, filename)
|
|
1066
|
+
|
|
1067
|
+
if os.path.exists(local_path):
|
|
1068
|
+
return local_path
|
|
1069
|
+
|
|
1070
|
+
# Check local repo path if present
|
|
1071
|
+
repo_alt = os.path.join("Graphormer-main", "checkpoints", filename)
|
|
1072
|
+
if os.path.exists(repo_alt):
|
|
1073
|
+
return repo_alt
|
|
1074
|
+
|
|
1075
|
+
try:
|
|
1076
|
+
return download_url(url, folder=folder, filename=filename, log=log)
|
|
1077
|
+
except Exception as e:
|
|
1078
|
+
if matched_key in LEGACY_URLS:
|
|
1079
|
+
fallback_url = LEGACY_URLS[matched_key]
|
|
1080
|
+
try:
|
|
1081
|
+
return download_url(fallback_url, folder=folder, filename=filename, log=log)
|
|
1082
|
+
except Exception:
|
|
1083
|
+
pass
|
|
1084
|
+
raise RuntimeError(f"Failed to download checkpoint for '{name}' from {url}: {e}")
|
|
1085
|
+
|
|
1086
|
+
|
|
1087
|
+
def load_graphormer_weights(
|
|
1088
|
+
model: Graphormer,
|
|
1089
|
+
checkpoint_path: Optional[str] = None,
|
|
1090
|
+
pretrained_name: Optional[str] = None,
|
|
1091
|
+
folder: str = "checkpoints",
|
|
1092
|
+
download: bool = True,
|
|
1093
|
+
) -> Graphormer:
|
|
1094
|
+
r"""Loads weights from a PyTorch (.pt or .bin) checkpoint into a Keras Graphormer model.
|
|
1095
|
+
|
|
1096
|
+
Args:
|
|
1097
|
+
model (Graphormer): Target model instance.
|
|
1098
|
+
checkpoint_path (str, optional): Path to .pt or .bin file.
|
|
1099
|
+
pretrained_name (str, optional): Name of pre-trained model to load or download.
|
|
1100
|
+
folder (str, optional): Checkpoints directory. (default: ``"checkpoints"``)
|
|
1101
|
+
download (bool, optional): Whether to download if missing. (default: ``True``)
|
|
1102
|
+
|
|
1103
|
+
Returns:
|
|
1104
|
+
Graphormer: The model with loaded weights.
|
|
1105
|
+
"""
|
|
1106
|
+
path_to_load = checkpoint_path
|
|
1107
|
+
|
|
1108
|
+
if path_to_load is None:
|
|
1109
|
+
if pretrained_name is None:
|
|
1110
|
+
raise ValueError("Either checkpoint_path or pretrained_name must be specified.")
|
|
1111
|
+
candidate = os.path.join(folder, f"{pretrained_name}.pt")
|
|
1112
|
+
if os.path.isfile(candidate):
|
|
1113
|
+
path_to_load = candidate
|
|
1114
|
+
elif download:
|
|
1115
|
+
path_to_load = download_graphormer_checkpoint(pretrained_name, folder=folder)
|
|
1116
|
+
else:
|
|
1117
|
+
raise FileNotFoundError(f"Checkpoint for '{pretrained_name}' not found at '{candidate}'.")
|
|
1118
|
+
elif not os.path.isfile(path_to_load) and download and pretrained_name:
|
|
1119
|
+
path_to_load = download_graphormer_checkpoint(pretrained_name, folder=folder)
|
|
1120
|
+
|
|
1121
|
+
import torch
|
|
1122
|
+
import numpy as np
|
|
1123
|
+
|
|
1124
|
+
state = torch.load(path_to_load, map_location="cpu")
|
|
1125
|
+
if isinstance(state, dict) and "model" in state:
|
|
1126
|
+
state_dict = state["model"]
|
|
1127
|
+
elif isinstance(state, dict):
|
|
1128
|
+
state_dict = state
|
|
1129
|
+
else:
|
|
1130
|
+
raise ValueError(f"Unexpected checkpoint state format: {type(state)}")
|
|
1131
|
+
|
|
1132
|
+
if not model.built:
|
|
1133
|
+
model.build(None)
|
|
1134
|
+
|
|
1135
|
+
def _to_tensor(t):
|
|
1136
|
+
if hasattr(t, "detach"):
|
|
1137
|
+
t = t.detach()
|
|
1138
|
+
if hasattr(t, "numpy"):
|
|
1139
|
+
t = t.numpy()
|
|
1140
|
+
return ops.convert_to_tensor(np.array(t, dtype=np.float32), dtype="float32")
|
|
1141
|
+
|
|
1142
|
+
# Prefix stripping (fairseq checkpoints often have 'encoder.' or 'graph_encoder.')
|
|
1143
|
+
clean_dict = {}
|
|
1144
|
+
for k, v in state_dict.items():
|
|
1145
|
+
ck = k
|
|
1146
|
+
if ck.startswith("encoder."):
|
|
1147
|
+
ck = ck[len("encoder.") :]
|
|
1148
|
+
clean_dict[ck] = v
|
|
1149
|
+
|
|
1150
|
+
# 1. GraphNodeFeature
|
|
1151
|
+
gnf = model.graph_encoder.graph_node_feature
|
|
1152
|
+
for param_name, layer in [
|
|
1153
|
+
("atom_encoder.weight", gnf.atom_encoder),
|
|
1154
|
+
("in_degree_encoder.weight", gnf.in_degree_encoder),
|
|
1155
|
+
("out_degree_encoder.weight", gnf.out_degree_encoder),
|
|
1156
|
+
("graph_token.weight", gnf.graph_token),
|
|
1157
|
+
]:
|
|
1158
|
+
key = f"graph_encoder.graph_node_feature.{param_name}"
|
|
1159
|
+
if key in clean_dict:
|
|
1160
|
+
layer.weights[0].assign(_to_tensor(clean_dict[key]))
|
|
1161
|
+
|
|
1162
|
+
# 2. GraphAttnBias
|
|
1163
|
+
gab = model.graph_encoder.graph_attn_bias
|
|
1164
|
+
for param_name, layer in [
|
|
1165
|
+
("edge_encoder.weight", gab.edge_encoder),
|
|
1166
|
+
("spatial_pos_encoder.weight", gab.spatial_pos_encoder),
|
|
1167
|
+
("graph_token_virtual_distance.weight", gab.graph_token_virtual_distance),
|
|
1168
|
+
]:
|
|
1169
|
+
key = f"graph_encoder.graph_attn_bias.{param_name}"
|
|
1170
|
+
if key in clean_dict:
|
|
1171
|
+
layer.weights[0].assign(_to_tensor(clean_dict[key]))
|
|
1172
|
+
|
|
1173
|
+
if hasattr(gab, "edge_dis_encoder"):
|
|
1174
|
+
key = "graph_encoder.graph_attn_bias.edge_dis_encoder.weight"
|
|
1175
|
+
if key in clean_dict:
|
|
1176
|
+
gab.edge_dis_encoder.weights[0].assign(_to_tensor(clean_dict[key]))
|
|
1177
|
+
|
|
1178
|
+
# 3. emb_layer_norm
|
|
1179
|
+
if model.graph_encoder.emb_layer_norm is not None:
|
|
1180
|
+
eln = model.graph_encoder.emb_layer_norm
|
|
1181
|
+
if "graph_encoder.emb_layer_norm.weight" in clean_dict:
|
|
1182
|
+
eln.gamma.assign(_to_tensor(clean_dict["graph_encoder.emb_layer_norm.weight"]))
|
|
1183
|
+
if "graph_encoder.emb_layer_norm.bias" in clean_dict:
|
|
1184
|
+
eln.beta.assign(_to_tensor(clean_dict["graph_encoder.emb_layer_norm.bias"]))
|
|
1185
|
+
|
|
1186
|
+
# 4. Encoder layers
|
|
1187
|
+
for i, enc_layer in enumerate(model.graph_encoder.encoder_layers):
|
|
1188
|
+
p = f"graph_encoder.layers.{i}"
|
|
1189
|
+
# Attention projections
|
|
1190
|
+
for proj_name in ["q_proj", "k_proj", "v_proj", "out_proj"]:
|
|
1191
|
+
proj = getattr(enc_layer.self_attn, proj_name)
|
|
1192
|
+
w_key = f"{p}.self_attn.{proj_name}.weight"
|
|
1193
|
+
b_key = f"{p}.self_attn.{proj_name}.bias"
|
|
1194
|
+
if w_key in clean_dict:
|
|
1195
|
+
proj.kernel.assign(_to_tensor(clean_dict[w_key].t()))
|
|
1196
|
+
if b_key in clean_dict and proj.bias is not None:
|
|
1197
|
+
proj.bias.assign(_to_tensor(clean_dict[b_key]))
|
|
1198
|
+
|
|
1199
|
+
# self_attn_layer_norm
|
|
1200
|
+
if f"{p}.self_attn_layer_norm.weight" in clean_dict:
|
|
1201
|
+
enc_layer.self_attn_layer_norm.gamma.assign(
|
|
1202
|
+
_to_tensor(clean_dict[f"{p}.self_attn_layer_norm.weight"])
|
|
1203
|
+
)
|
|
1204
|
+
if f"{p}.self_attn_layer_norm.bias" in clean_dict:
|
|
1205
|
+
enc_layer.self_attn_layer_norm.beta.assign(
|
|
1206
|
+
_to_tensor(clean_dict[f"{p}.self_attn_layer_norm.bias"])
|
|
1207
|
+
)
|
|
1208
|
+
|
|
1209
|
+
# fc1 & fc2
|
|
1210
|
+
if f"{p}.fc1.weight" in clean_dict:
|
|
1211
|
+
enc_layer.fc1.kernel.assign(_to_tensor(clean_dict[f"{p}.fc1.weight"].t()))
|
|
1212
|
+
if f"{p}.fc1.bias" in clean_dict:
|
|
1213
|
+
enc_layer.fc1.bias.assign(_to_tensor(clean_dict[f"{p}.fc1.bias"]))
|
|
1214
|
+
if f"{p}.fc2.weight" in clean_dict:
|
|
1215
|
+
enc_layer.fc2.kernel.assign(_to_tensor(clean_dict[f"{p}.fc2.weight"].t()))
|
|
1216
|
+
if f"{p}.fc2.bias" in clean_dict:
|
|
1217
|
+
enc_layer.fc2.bias.assign(_to_tensor(clean_dict[f"{p}.fc2.bias"]))
|
|
1218
|
+
|
|
1219
|
+
# final_layer_norm
|
|
1220
|
+
if f"{p}.final_layer_norm.weight" in clean_dict:
|
|
1221
|
+
enc_layer.final_layer_norm.gamma.assign(
|
|
1222
|
+
_to_tensor(clean_dict[f"{p}.final_layer_norm.weight"])
|
|
1223
|
+
)
|
|
1224
|
+
if f"{p}.final_layer_norm.bias" in clean_dict:
|
|
1225
|
+
enc_layer.final_layer_norm.beta.assign(
|
|
1226
|
+
_to_tensor(clean_dict[f"{p}.final_layer_norm.bias"])
|
|
1227
|
+
)
|
|
1228
|
+
|
|
1229
|
+
# 5. final_layer_norm of encoder stack (if pre_layernorm)
|
|
1230
|
+
if model.graph_encoder.final_layer_norm is not None:
|
|
1231
|
+
fln = model.graph_encoder.final_layer_norm
|
|
1232
|
+
if "graph_encoder.final_layer_norm.weight" in clean_dict:
|
|
1233
|
+
fln.gamma.assign(_to_tensor(clean_dict["graph_encoder.final_layer_norm.weight"]))
|
|
1234
|
+
if "graph_encoder.final_layer_norm.bias" in clean_dict:
|
|
1235
|
+
fln.beta.assign(_to_tensor(clean_dict["graph_encoder.final_layer_norm.bias"]))
|
|
1236
|
+
|
|
1237
|
+
# 6. Prediction Head
|
|
1238
|
+
if "lm_head_transform_weight.weight" in clean_dict:
|
|
1239
|
+
model.lm_head_transform_weight.kernel.assign(
|
|
1240
|
+
_to_tensor(clean_dict["lm_head_transform_weight.weight"].t())
|
|
1241
|
+
)
|
|
1242
|
+
if "lm_head_transform_weight.bias" in clean_dict:
|
|
1243
|
+
model.lm_head_transform_weight.bias.assign(
|
|
1244
|
+
_to_tensor(clean_dict["lm_head_transform_weight.bias"])
|
|
1245
|
+
)
|
|
1246
|
+
|
|
1247
|
+
if "layer_norm.weight" in clean_dict:
|
|
1248
|
+
model.layer_norm.gamma.assign(_to_tensor(clean_dict["layer_norm.weight"]))
|
|
1249
|
+
if "layer_norm.bias" in clean_dict:
|
|
1250
|
+
model.layer_norm.beta.assign(_to_tensor(clean_dict["layer_norm.bias"]))
|
|
1251
|
+
|
|
1252
|
+
if "embed_out.weight" in clean_dict and hasattr(model, "embed_out"):
|
|
1253
|
+
model.embed_out.kernel.assign(_to_tensor(clean_dict["embed_out.weight"].t()))
|
|
1254
|
+
|
|
1255
|
+
if "lm_output_learned_bias" in clean_dict and hasattr(model, "lm_output_learned_bias"):
|
|
1256
|
+
model.lm_output_learned_bias.assign(_to_tensor(clean_dict["lm_output_learned_bias"]))
|
|
1257
|
+
|
|
1258
|
+
return model
|