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
k3_node/models/unimol.py
ADDED
|
@@ -0,0 +1,1156 @@
|
|
|
1
|
+
import math
|
|
2
|
+
import os
|
|
3
|
+
from typing import Optional, Union, Tuple, List, Dict, Any, Callable
|
|
4
|
+
|
|
5
|
+
import numpy as np
|
|
6
|
+
import keras
|
|
7
|
+
from keras import layers, ops
|
|
8
|
+
|
|
9
|
+
from k3_node.layers.attention.pair_attention import (
|
|
10
|
+
SelfMultiheadAttentionWithPair,
|
|
11
|
+
TransformerEncoderLayerWithPair,
|
|
12
|
+
_get_activation,
|
|
13
|
+
)
|
|
14
|
+
from k3_node.data.download import download_url
|
|
15
|
+
|
|
16
|
+
|
|
17
|
+
# ==============================================================================
|
|
18
|
+
# Pretrained Weights Registry
|
|
19
|
+
# ==============================================================================
|
|
20
|
+
|
|
21
|
+
UNIMOL_PRETRAINED_URLS: Dict[str, Dict[str, str]] = {
|
|
22
|
+
"mol_pre_no_h": {
|
|
23
|
+
"url": "https://github.com/deepmodeling/Uni-Mol/releases/download/v0.1/mol_pre_no_h_220816.pt",
|
|
24
|
+
"filename": "mol_pre_no_h_220816.pt",
|
|
25
|
+
"dict": "mol.dict.txt",
|
|
26
|
+
"dict_url": "https://huggingface.co/dptech/Uni-Mol-Models/resolve/main/mol.dict.txt",
|
|
27
|
+
"description": "Uni-Mol molecular pretraining model (no hydrogen)",
|
|
28
|
+
},
|
|
29
|
+
"mol_pre_all_h": {
|
|
30
|
+
"url": "https://github.com/deepmodeling/Uni-Mol/releases/download/v0.1/mol_pre_all_h_220816.pt",
|
|
31
|
+
"filename": "mol_pre_all_h_220816.pt",
|
|
32
|
+
"dict": "mol.dict.txt",
|
|
33
|
+
"dict_url": "https://huggingface.co/dptech/Uni-Mol-Models/resolve/main/mol.dict.txt",
|
|
34
|
+
"description": "Uni-Mol molecular pretraining model (all hydrogen)",
|
|
35
|
+
},
|
|
36
|
+
"pocket_pre": {
|
|
37
|
+
"url": "https://github.com/deepmodeling/Uni-Mol/releases/download/v0.1/pocket_pre_220816.pt",
|
|
38
|
+
"filename": "pocket_pre_220816.pt",
|
|
39
|
+
"dict": "poc.dict.txt",
|
|
40
|
+
"dict_url": "https://huggingface.co/dptech/Uni-Mol-Models/resolve/main/poc.dict.txt",
|
|
41
|
+
"description": "Uni-Mol candidate protein pocket pretraining model",
|
|
42
|
+
},
|
|
43
|
+
"mp_all_h": {
|
|
44
|
+
"url": "https://huggingface.co/dptech/Uni-Mol-Models/resolve/main/mp_all_h_230313.pt",
|
|
45
|
+
"filename": "mp_all_h_230313.pt",
|
|
46
|
+
"dict": "mp.dict.txt",
|
|
47
|
+
"dict_url": "https://huggingface.co/dptech/Uni-Mol-Models/resolve/main/mp.dict.txt",
|
|
48
|
+
"description": "Uni-Mol crystal Materials Project pretraining model",
|
|
49
|
+
},
|
|
50
|
+
"oled_pre_no_h": {
|
|
51
|
+
"url": "https://huggingface.co/dptech/Uni-Mol-Models/resolve/main/oled_pre_no_h_230101.pt",
|
|
52
|
+
"filename": "oled_pre_no_h_230101.pt",
|
|
53
|
+
"dict": "oled.dict.txt",
|
|
54
|
+
"dict_url": "https://huggingface.co/dptech/Uni-Mol-Models/resolve/main/oled.dict.txt",
|
|
55
|
+
"description": "Uni-Mol OLED molecule pretraining model",
|
|
56
|
+
},
|
|
57
|
+
"qm9": {
|
|
58
|
+
"url": "https://github.com/deepmodeling/Uni-Mol/releases/download/v0.1/qm9_220908.pt",
|
|
59
|
+
"filename": "qm9_220908.pt",
|
|
60
|
+
"dict": "mol.dict.txt",
|
|
61
|
+
"dict_url": "https://huggingface.co/dptech/Uni-Mol-Models/resolve/main/mol.dict.txt",
|
|
62
|
+
"description": "Uni-Mol conformation generation fine-tuned on QM9",
|
|
63
|
+
},
|
|
64
|
+
"drugs": {
|
|
65
|
+
"url": "https://github.com/deepmodeling/Uni-Mol/releases/download/v0.1/drugs_220908.pt",
|
|
66
|
+
"filename": "drugs_220908.pt",
|
|
67
|
+
"dict": "mol.dict.txt",
|
|
68
|
+
"dict_url": "https://huggingface.co/dptech/Uni-Mol-Models/resolve/main/mol.dict.txt",
|
|
69
|
+
"description": "Uni-Mol conformation generation fine-tuned on GEOM-Drugs",
|
|
70
|
+
},
|
|
71
|
+
"binding_pose": {
|
|
72
|
+
"url": "https://github.com/deepmodeling/Uni-Mol/releases/download/v0.1/binding_pose_220908.pt",
|
|
73
|
+
"filename": "binding_pose_220908.pt",
|
|
74
|
+
"dict": "mol.dict.txt",
|
|
75
|
+
"dict_url": "https://huggingface.co/dptech/Uni-Mol-Models/resolve/main/mol.dict.txt",
|
|
76
|
+
"description": "Uni-Mol protein-ligand binding pose prediction",
|
|
77
|
+
},
|
|
78
|
+
}
|
|
79
|
+
|
|
80
|
+
UNIMOL_ALIASES = {
|
|
81
|
+
"molecule": "mol_pre_no_h",
|
|
82
|
+
"molecule_no_h": "mol_pre_no_h",
|
|
83
|
+
"molecule_all_h": "mol_pre_all_h",
|
|
84
|
+
"protein": "pocket_pre",
|
|
85
|
+
"pocket": "pocket_pre",
|
|
86
|
+
"poc_pre": "pocket_pre",
|
|
87
|
+
"crystal": "mp_all_h",
|
|
88
|
+
"mp": "mp_all_h",
|
|
89
|
+
"oled": "oled_pre_no_h",
|
|
90
|
+
}
|
|
91
|
+
|
|
92
|
+
|
|
93
|
+
# ==============================================================================
|
|
94
|
+
# Layers
|
|
95
|
+
# ==============================================================================
|
|
96
|
+
|
|
97
|
+
class GaussianLayer(layers.Layer):
|
|
98
|
+
r"""Gaussian basis function (GBF) expansion over pairwise distances modulated by edge types.
|
|
99
|
+
|
|
100
|
+
Args:
|
|
101
|
+
num_kernel (int, optional): Number of Gaussian kernels. (default: ``128``)
|
|
102
|
+
edge_types (int, optional): Number of distinct pairwise edge types. (default: ``1024``)
|
|
103
|
+
**kwargs: Additional layer arguments.
|
|
104
|
+
|
|
105
|
+
Example:
|
|
106
|
+
```python
|
|
107
|
+
import numpy as np
|
|
108
|
+
from k3_node.models import UniMolGaussianLayer
|
|
109
|
+
|
|
110
|
+
dist = np.random.rand(2, 5, 5).astype("float32") * 5.0 # pairwise atom distances
|
|
111
|
+
edge_type = np.random.randint(0, 128, size=(2, 5, 5)) # atom-pair type ids
|
|
112
|
+
|
|
113
|
+
layer = UniMolGaussianLayer(num_kernel=32, edge_types=128)
|
|
114
|
+
print(tuple(layer(dist, edge_type).shape)) # (2, 5, 5, 32): Gaussian distance features per atom pair
|
|
115
|
+
```
|
|
116
|
+
"""
|
|
117
|
+
|
|
118
|
+
def __init__(self, num_kernel: int = 128, edge_types: int = 1024, **kwargs):
|
|
119
|
+
super().__init__(**kwargs)
|
|
120
|
+
self.num_kernel = num_kernel
|
|
121
|
+
self.edge_types = edge_types
|
|
122
|
+
|
|
123
|
+
self.means = layers.Embedding(1, num_kernel, embeddings_initializer="uniform", name="means")
|
|
124
|
+
self.stds = layers.Embedding(1, num_kernel, embeddings_initializer="uniform", name="stds")
|
|
125
|
+
self.mul = layers.Embedding(edge_types, 1, embeddings_initializer="ones", name="mul")
|
|
126
|
+
self.bias = layers.Embedding(edge_types, 1, embeddings_initializer="zeros", name="bias")
|
|
127
|
+
|
|
128
|
+
def build(self, input_shape=None):
|
|
129
|
+
if not self.built:
|
|
130
|
+
self.means.build(None)
|
|
131
|
+
self.stds.build(None)
|
|
132
|
+
self.mul.build(None)
|
|
133
|
+
self.bias.build(None)
|
|
134
|
+
super().build(input_shape)
|
|
135
|
+
|
|
136
|
+
def call(self, dist, edge_type):
|
|
137
|
+
r"""
|
|
138
|
+
Args:
|
|
139
|
+
dist (Tensor): Pairwise distance matrix of shape ``[batch_size, seq_len, seq_len]``.
|
|
140
|
+
edge_type (Tensor): Pairwise edge type indices of shape ``[batch_size, seq_len, seq_len]``.
|
|
141
|
+
|
|
142
|
+
Returns:
|
|
143
|
+
Tensor: Gaussian basis expansion of shape ``[batch_size, seq_len, seq_len, num_kernel]``.
|
|
144
|
+
"""
|
|
145
|
+
mul = ops.cast(self.mul(edge_type), dist.dtype)
|
|
146
|
+
bias = ops.cast(self.bias(edge_type), dist.dtype)
|
|
147
|
+
|
|
148
|
+
# x: [B, N, N, 1]
|
|
149
|
+
x = mul * ops.expand_dims(dist, axis=-1) + bias
|
|
150
|
+
|
|
151
|
+
# Means and stds: [num_kernel]
|
|
152
|
+
zero_idx = ops.zeros((1,), dtype="int32")
|
|
153
|
+
mean = ops.reshape(self.means(zero_idx), (-1,))
|
|
154
|
+
std = ops.abs(ops.reshape(self.stds(zero_idx), (-1,))) + 1e-5
|
|
155
|
+
|
|
156
|
+
a = math.sqrt(2.0 * math.pi)
|
|
157
|
+
diff = (x - mean) / std
|
|
158
|
+
return ops.exp(-0.5 * ops.power(diff, 2)) / (a * std)
|
|
159
|
+
|
|
160
|
+
|
|
161
|
+
class NumericalEmbed(layers.Layer):
|
|
162
|
+
r"""Numerical embedding layer for continuous edge features.
|
|
163
|
+
|
|
164
|
+
Example:
|
|
165
|
+
```python
|
|
166
|
+
import numpy as np
|
|
167
|
+
from k3_node.models import UniMolNumericalEmbed
|
|
168
|
+
|
|
169
|
+
dist = np.random.rand(2, 5, 5).astype("float32") * 5.0 # pairwise atom distances
|
|
170
|
+
edge_type = np.random.randint(0, 128, size=(2, 5, 5)) # atom-pair type ids
|
|
171
|
+
|
|
172
|
+
layer = UniMolNumericalEmbed(num_kernel=32, edge_types=128)
|
|
173
|
+
print(tuple(layer(dist, edge_type).shape)) # (2, 5, 5, 32)
|
|
174
|
+
```
|
|
175
|
+
"""
|
|
176
|
+
|
|
177
|
+
def __init__(self, num_kernel: int = 128, edge_types: int = 1024, activation_fn: str = "gelu", **kwargs):
|
|
178
|
+
super().__init__(**kwargs)
|
|
179
|
+
self.num_kernel = num_kernel
|
|
180
|
+
self.edge_types = edge_types
|
|
181
|
+
self.mul = layers.Embedding(edge_types, 1, name="mul")
|
|
182
|
+
self.bias = layers.Embedding(edge_types, 1, name="bias")
|
|
183
|
+
self.w_edge = layers.Embedding(edge_types, num_kernel, name="w_edge")
|
|
184
|
+
self.proj = NonLinearHead(1, num_kernel, activation_fn=activation_fn, hidden=2 * num_kernel, name="proj")
|
|
185
|
+
self.ln = layers.LayerNormalization(axis=-1, epsilon=1e-5, name="ln")
|
|
186
|
+
|
|
187
|
+
def build(self, input_shape=None):
|
|
188
|
+
if not self.built:
|
|
189
|
+
self.mul.build(None)
|
|
190
|
+
self.bias.build(None)
|
|
191
|
+
self.w_edge.build(None)
|
|
192
|
+
self.proj.build((None, None, None, 1))
|
|
193
|
+
self.ln.build((None, None, None, self.num_kernel))
|
|
194
|
+
super().build(input_shape)
|
|
195
|
+
|
|
196
|
+
def call(self, dist, edge_type):
|
|
197
|
+
mul = ops.cast(self.mul(edge_type), dist.dtype)
|
|
198
|
+
bias = ops.cast(self.bias(edge_type), dist.dtype)
|
|
199
|
+
w_edge = ops.cast(self.w_edge(edge_type), dist.dtype)
|
|
200
|
+
|
|
201
|
+
edge_feat = mul * ops.expand_dims(dist, axis=-1) + bias
|
|
202
|
+
edge_feat = self.proj(edge_feat)
|
|
203
|
+
edge_feat = edge_feat + w_edge
|
|
204
|
+
return self.ln(edge_feat)
|
|
205
|
+
|
|
206
|
+
|
|
207
|
+
class NonLinearHead(layers.Layer):
|
|
208
|
+
r"""Two-layer feed-forward network with activation for feature projection.
|
|
209
|
+
|
|
210
|
+
Example:
|
|
211
|
+
```python
|
|
212
|
+
import numpy as np
|
|
213
|
+
from k3_node.models import UniMolNonLinearHead
|
|
214
|
+
|
|
215
|
+
x = np.random.rand(2, 5, 32).astype("float32") # [batch, atoms, embed_dim]
|
|
216
|
+
|
|
217
|
+
head = UniMolNonLinearHead(input_dim=32, out_dim=16, activation_fn="gelu")
|
|
218
|
+
print(tuple(head(x).shape)) # (2, 5, 16)
|
|
219
|
+
```
|
|
220
|
+
"""
|
|
221
|
+
|
|
222
|
+
def __init__(
|
|
223
|
+
self,
|
|
224
|
+
input_dim: int,
|
|
225
|
+
out_dim: int,
|
|
226
|
+
activation_fn: Union[str, Callable] = "gelu",
|
|
227
|
+
hidden: Optional[int] = None,
|
|
228
|
+
**kwargs,
|
|
229
|
+
):
|
|
230
|
+
super().__init__(**kwargs)
|
|
231
|
+
self.input_dim = input_dim
|
|
232
|
+
self.out_dim = out_dim
|
|
233
|
+
self.hidden = hidden or input_dim
|
|
234
|
+
self.activation_fn_name = activation_fn
|
|
235
|
+
|
|
236
|
+
self.linear1 = layers.Dense(self.hidden, name="linear1")
|
|
237
|
+
self.act = _get_activation(activation_fn)
|
|
238
|
+
self.linear2 = layers.Dense(out_dim, name="linear2")
|
|
239
|
+
|
|
240
|
+
def build(self, input_shape=None):
|
|
241
|
+
if not self.built:
|
|
242
|
+
self.linear1.build((None, None, None, self.input_dim) if len(input_shape or ()) == 4 else (None, None, self.input_dim))
|
|
243
|
+
self.linear2.build((None, None, None, self.hidden) if len(input_shape or ()) == 4 else (None, None, self.hidden))
|
|
244
|
+
super().build(input_shape)
|
|
245
|
+
|
|
246
|
+
def call(self, x):
|
|
247
|
+
x = self.linear1(x)
|
|
248
|
+
if self.act is not None:
|
|
249
|
+
x = self.act(x)
|
|
250
|
+
x = self.linear2(x)
|
|
251
|
+
return x
|
|
252
|
+
|
|
253
|
+
|
|
254
|
+
class DistanceHead(layers.Layer):
|
|
255
|
+
r"""Symmetrized distance prediction head from pair representations.
|
|
256
|
+
|
|
257
|
+
Example:
|
|
258
|
+
```python
|
|
259
|
+
import numpy as np
|
|
260
|
+
from k3_node.models import UniMolDistanceHead
|
|
261
|
+
|
|
262
|
+
pair = np.random.rand(2, 6, 6, 8).astype("float32") # pair representation with 8 heads
|
|
263
|
+
print(tuple(UniMolDistanceHead(heads=8)(pair).shape)) # (2, 6, 6): predicted distance matrix
|
|
264
|
+
```
|
|
265
|
+
"""
|
|
266
|
+
|
|
267
|
+
def __init__(self, heads: int, activation_fn: Union[str, Callable] = "gelu", **kwargs):
|
|
268
|
+
super().__init__(**kwargs)
|
|
269
|
+
self.heads = heads
|
|
270
|
+
self.dense = layers.Dense(heads, name="dense")
|
|
271
|
+
self.act = _get_activation(activation_fn)
|
|
272
|
+
self.layer_norm = layers.LayerNormalization(axis=-1, epsilon=1e-5, name="layer_norm")
|
|
273
|
+
self.out_proj = layers.Dense(1, name="out_proj")
|
|
274
|
+
|
|
275
|
+
def build(self, input_shape=None):
|
|
276
|
+
if not self.built:
|
|
277
|
+
self.dense.build((None, None, None, self.heads))
|
|
278
|
+
self.layer_norm.build((None, None, None, self.heads))
|
|
279
|
+
self.out_proj.build((None, None, None, self.heads))
|
|
280
|
+
super().build(input_shape)
|
|
281
|
+
|
|
282
|
+
def call(self, x):
|
|
283
|
+
r"""
|
|
284
|
+
Args:
|
|
285
|
+
x (Tensor): Pair tensor of shape ``[batch_size, seq_len, seq_len, heads]``.
|
|
286
|
+
|
|
287
|
+
Returns:
|
|
288
|
+
Tensor: Symmetrized predicted distances ``[batch_size, seq_len, seq_len]``.
|
|
289
|
+
"""
|
|
290
|
+
x = self.dense(x)
|
|
291
|
+
if self.act is not None:
|
|
292
|
+
x = self.act(x)
|
|
293
|
+
x = self.layer_norm(x)
|
|
294
|
+
x = ops.squeeze(self.out_proj(x), axis=-1) # [B, N, N]
|
|
295
|
+
return 0.5 * (x + ops.transpose(x, (0, 2, 1)))
|
|
296
|
+
|
|
297
|
+
|
|
298
|
+
class LinearHead(layers.Layer):
|
|
299
|
+
r"""Linear classification/regression head.
|
|
300
|
+
|
|
301
|
+
Example:
|
|
302
|
+
```python
|
|
303
|
+
import numpy as np
|
|
304
|
+
from k3_node.models import UniMolLinearHead
|
|
305
|
+
|
|
306
|
+
x = np.random.rand(2, 5, 32).astype("float32") # [batch, atoms, embed_dim]
|
|
307
|
+
|
|
308
|
+
head = UniMolLinearHead(input_dim=32, num_classes=3)
|
|
309
|
+
print(tuple(head(x).shape)) # (2, 5, 3)
|
|
310
|
+
```
|
|
311
|
+
"""
|
|
312
|
+
|
|
313
|
+
def __init__(self, input_dim: int, num_classes: int, pooler_dropout: float = 0.0, **kwargs):
|
|
314
|
+
super().__init__(**kwargs)
|
|
315
|
+
self.input_dim = input_dim
|
|
316
|
+
self.num_classes = num_classes
|
|
317
|
+
self.pooler_dropout = pooler_dropout
|
|
318
|
+
|
|
319
|
+
self.dropout = layers.Dropout(pooler_dropout) if pooler_dropout > 0.0 else None
|
|
320
|
+
self.out_proj = layers.Dense(num_classes, name="out_proj")
|
|
321
|
+
|
|
322
|
+
def build(self, input_shape=None):
|
|
323
|
+
if not self.built:
|
|
324
|
+
self.out_proj.build((None, self.input_dim))
|
|
325
|
+
super().build(input_shape)
|
|
326
|
+
|
|
327
|
+
def call(self, features, training: bool = False):
|
|
328
|
+
x = features
|
|
329
|
+
if self.dropout is not None:
|
|
330
|
+
x = self.dropout(x, training=training)
|
|
331
|
+
return self.out_proj(x)
|
|
332
|
+
|
|
333
|
+
|
|
334
|
+
class ClassificationHead(layers.Layer):
|
|
335
|
+
r"""Two-layer sentence/graph-level classification head.
|
|
336
|
+
|
|
337
|
+
Example:
|
|
338
|
+
```python
|
|
339
|
+
import numpy as np
|
|
340
|
+
from k3_node.models import UniMolClassificationHead
|
|
341
|
+
|
|
342
|
+
x = np.random.rand(2, 5, 32).astype("float32") # [batch, atoms, embed_dim]
|
|
343
|
+
|
|
344
|
+
head = UniMolClassificationHead(input_dim=32, inner_dim=32, num_classes=3)
|
|
345
|
+
print(tuple(head(x).shape)) # (2, 3): classifies from the first ([CLS]) token
|
|
346
|
+
```
|
|
347
|
+
"""
|
|
348
|
+
|
|
349
|
+
def __init__(
|
|
350
|
+
self,
|
|
351
|
+
input_dim: int,
|
|
352
|
+
inner_dim: int,
|
|
353
|
+
num_classes: int,
|
|
354
|
+
activation_fn: Union[str, Callable] = "gelu",
|
|
355
|
+
pooler_dropout: float = 0.0,
|
|
356
|
+
**kwargs,
|
|
357
|
+
):
|
|
358
|
+
super().__init__(**kwargs)
|
|
359
|
+
self.input_dim = input_dim
|
|
360
|
+
self.inner_dim = inner_dim
|
|
361
|
+
self.num_classes = num_classes
|
|
362
|
+
self.pooler_dropout = pooler_dropout
|
|
363
|
+
|
|
364
|
+
self.dense = layers.Dense(inner_dim, name="dense")
|
|
365
|
+
self.act = _get_activation(activation_fn)
|
|
366
|
+
self.dropout = layers.Dropout(pooler_dropout) if pooler_dropout > 0.0 else None
|
|
367
|
+
self.out_proj = layers.Dense(num_classes, name="out_proj")
|
|
368
|
+
|
|
369
|
+
def build(self, input_shape=None):
|
|
370
|
+
if not self.built:
|
|
371
|
+
self.dense.build((None, self.input_dim))
|
|
372
|
+
self.out_proj.build((None, self.inner_dim))
|
|
373
|
+
super().build(input_shape)
|
|
374
|
+
|
|
375
|
+
def call(self, features, training: bool = False):
|
|
376
|
+
x = features
|
|
377
|
+
if len(ops.shape(features)) == 3:
|
|
378
|
+
x = features[:, 0, :] # CLS token
|
|
379
|
+
if self.dropout is not None:
|
|
380
|
+
x = self.dropout(x, training=training)
|
|
381
|
+
x = self.dense(x)
|
|
382
|
+
if self.act is not None:
|
|
383
|
+
x = self.act(x)
|
|
384
|
+
if self.dropout is not None:
|
|
385
|
+
x = self.dropout(x, training=training)
|
|
386
|
+
return self.out_proj(x)
|
|
387
|
+
|
|
388
|
+
|
|
389
|
+
class MaskLMHead(layers.Layer):
|
|
390
|
+
r"""Masked language modeling head for predicting masked atom tokens.
|
|
391
|
+
|
|
392
|
+
Example:
|
|
393
|
+
```python
|
|
394
|
+
import numpy as np
|
|
395
|
+
from k3_node.models import UniMolMaskLMHead
|
|
396
|
+
|
|
397
|
+
x = np.random.rand(2, 5, 32).astype("float32") # [batch, atoms, embed_dim]
|
|
398
|
+
|
|
399
|
+
head = UniMolMaskLMHead(embed_dim=32, output_dim=64) # logits over a 64-token vocabulary
|
|
400
|
+
print(tuple(head(x).shape)) # (2, 5, 64)
|
|
401
|
+
```
|
|
402
|
+
"""
|
|
403
|
+
|
|
404
|
+
def __init__(
|
|
405
|
+
self,
|
|
406
|
+
embed_dim: int,
|
|
407
|
+
output_dim: int,
|
|
408
|
+
activation_fn: Union[str, Callable] = "gelu",
|
|
409
|
+
**kwargs,
|
|
410
|
+
):
|
|
411
|
+
super().__init__(**kwargs)
|
|
412
|
+
self.embed_dim = embed_dim
|
|
413
|
+
self.output_dim = output_dim
|
|
414
|
+
self.dense = layers.Dense(embed_dim, name="dense")
|
|
415
|
+
self.act = _get_activation(activation_fn)
|
|
416
|
+
self.layer_norm = layers.LayerNormalization(axis=-1, epsilon=1e-5, name="layer_norm")
|
|
417
|
+
self.out_proj = layers.Dense(output_dim, name="out_proj")
|
|
418
|
+
|
|
419
|
+
def build(self, input_shape=None):
|
|
420
|
+
if not self.built:
|
|
421
|
+
self.dense.build((None, None, self.embed_dim))
|
|
422
|
+
self.layer_norm.build((None, None, self.embed_dim))
|
|
423
|
+
self.out_proj.build((None, None, self.embed_dim))
|
|
424
|
+
super().build(input_shape)
|
|
425
|
+
|
|
426
|
+
def call(self, features, masked_tokens=None):
|
|
427
|
+
x = self.dense(features)
|
|
428
|
+
if self.act is not None:
|
|
429
|
+
x = self.act(x)
|
|
430
|
+
x = self.layer_norm(x)
|
|
431
|
+
x = self.out_proj(x)
|
|
432
|
+
return x
|
|
433
|
+
|
|
434
|
+
|
|
435
|
+
# ==============================================================================
|
|
436
|
+
# Backbone Transformer Encoder
|
|
437
|
+
# ==============================================================================
|
|
438
|
+
|
|
439
|
+
class UniMolTransformerEncoder(layers.Layer):
|
|
440
|
+
r"""Transformer Encoder backbone for Uni-Mol with pair attention bias propagation."""
|
|
441
|
+
|
|
442
|
+
def __init__(
|
|
443
|
+
self,
|
|
444
|
+
encoder_layers: int = 15,
|
|
445
|
+
embed_dim: int = 512,
|
|
446
|
+
ffn_embed_dim: int = 2048,
|
|
447
|
+
attention_heads: int = 64,
|
|
448
|
+
emb_dropout: float = 0.1,
|
|
449
|
+
dropout: float = 0.1,
|
|
450
|
+
attention_dropout: float = 0.1,
|
|
451
|
+
activation_dropout: float = 0.0,
|
|
452
|
+
max_seq_len: int = 512,
|
|
453
|
+
activation_fn: Union[str, Callable] = "gelu",
|
|
454
|
+
post_ln: bool = False,
|
|
455
|
+
no_final_head_layer_norm: bool = False,
|
|
456
|
+
**kwargs,
|
|
457
|
+
):
|
|
458
|
+
super().__init__(**kwargs)
|
|
459
|
+
self.num_layers = encoder_layers
|
|
460
|
+
self.embed_dim = embed_dim
|
|
461
|
+
self.ffn_embed_dim = ffn_embed_dim
|
|
462
|
+
self.attention_heads = attention_heads
|
|
463
|
+
self.emb_dropout_rate = emb_dropout
|
|
464
|
+
self.post_ln = post_ln
|
|
465
|
+
|
|
466
|
+
self.emb_layer_norm = layers.LayerNormalization(axis=-1, epsilon=1e-5, name="emb_layer_norm")
|
|
467
|
+
self.emb_dropout = layers.Dropout(emb_dropout) if emb_dropout > 0.0 else None
|
|
468
|
+
|
|
469
|
+
self.final_layer_norm = None if post_ln else layers.LayerNormalization(axis=-1, epsilon=1e-5, name="final_layer_norm")
|
|
470
|
+
self.final_head_layer_norm = None if no_final_head_layer_norm else layers.LayerNormalization(axis=-1, epsilon=1e-5, name="final_head_layer_norm")
|
|
471
|
+
|
|
472
|
+
self.layers_list = [
|
|
473
|
+
TransformerEncoderLayerWithPair(
|
|
474
|
+
embed_dim=embed_dim,
|
|
475
|
+
ffn_embed_dim=ffn_embed_dim,
|
|
476
|
+
attention_heads=attention_heads,
|
|
477
|
+
dropout=dropout,
|
|
478
|
+
attention_dropout=attention_dropout,
|
|
479
|
+
activation_dropout=activation_dropout,
|
|
480
|
+
activation_fn=activation_fn,
|
|
481
|
+
post_ln=post_ln,
|
|
482
|
+
name=f"layer_{i}",
|
|
483
|
+
)
|
|
484
|
+
for i in range(encoder_layers)
|
|
485
|
+
]
|
|
486
|
+
|
|
487
|
+
def build(self, input_shape=None):
|
|
488
|
+
if not self.built:
|
|
489
|
+
self.emb_layer_norm.build((None, None, self.embed_dim))
|
|
490
|
+
if self.final_layer_norm is not None:
|
|
491
|
+
self.final_layer_norm.build((None, None, self.embed_dim))
|
|
492
|
+
if self.final_head_layer_norm is not None:
|
|
493
|
+
self.final_head_layer_norm.build((None, None, None, self.attention_heads))
|
|
494
|
+
for layer in self.layers_list:
|
|
495
|
+
layer.build((None, None, self.embed_dim))
|
|
496
|
+
super().build(input_shape)
|
|
497
|
+
|
|
498
|
+
def call(self, emb, attn_mask=None, padding_mask=None, training: bool = False):
|
|
499
|
+
shape = ops.shape(emb)
|
|
500
|
+
bsz = shape[0]
|
|
501
|
+
seq_len = shape[1]
|
|
502
|
+
|
|
503
|
+
x = self.emb_layer_norm(emb)
|
|
504
|
+
if self.emb_dropout is not None:
|
|
505
|
+
x = self.emb_dropout(x, training=training)
|
|
506
|
+
|
|
507
|
+
if padding_mask is not None:
|
|
508
|
+
mask_expanded = ops.expand_dims(ops.cast(padding_mask, x.dtype), axis=-1)
|
|
509
|
+
x = x * (1.0 - mask_expanded)
|
|
510
|
+
|
|
511
|
+
if attn_mask is not None:
|
|
512
|
+
mask_shape = ops.shape(attn_mask)
|
|
513
|
+
if len(mask_shape) == 4 and mask_shape[-1] == self.attention_heads:
|
|
514
|
+
attn_mask = ops.transpose(attn_mask, (0, 3, 1, 2))
|
|
515
|
+
elif len(mask_shape) == 3:
|
|
516
|
+
attn_mask = ops.reshape(attn_mask, (bsz, self.attention_heads, seq_len, seq_len))
|
|
517
|
+
|
|
518
|
+
input_attn_mask = attn_mask
|
|
519
|
+
curr_attn_mask = attn_mask
|
|
520
|
+
|
|
521
|
+
for enc_layer in self.layers_list:
|
|
522
|
+
x, curr_attn_mask, _ = enc_layer(
|
|
523
|
+
x,
|
|
524
|
+
attn_bias=curr_attn_mask,
|
|
525
|
+
padding_mask=padding_mask,
|
|
526
|
+
return_attn=True,
|
|
527
|
+
training=training,
|
|
528
|
+
)
|
|
529
|
+
|
|
530
|
+
if self.final_layer_norm is not None:
|
|
531
|
+
x = self.final_layer_norm(x)
|
|
532
|
+
|
|
533
|
+
# Delta pair representation
|
|
534
|
+
delta_pair_repr = curr_attn_mask - input_attn_mask
|
|
535
|
+
|
|
536
|
+
# Reshape to [bsz, seq_len, seq_len, attention_heads]
|
|
537
|
+
pair_shape = ops.shape(curr_attn_mask)
|
|
538
|
+
if len(pair_shape) == 3:
|
|
539
|
+
curr_attn_mask = ops.reshape(curr_attn_mask, (bsz, self.attention_heads, seq_len, seq_len))
|
|
540
|
+
delta_pair_repr = ops.reshape(delta_pair_repr, (bsz, self.attention_heads, seq_len, seq_len))
|
|
541
|
+
|
|
542
|
+
curr_attn_mask = ops.transpose(curr_attn_mask, (0, 2, 3, 1))
|
|
543
|
+
delta_pair_repr = ops.transpose(delta_pair_repr, (0, 2, 3, 1))
|
|
544
|
+
|
|
545
|
+
if self.final_head_layer_norm is not None:
|
|
546
|
+
delta_pair_repr = self.final_head_layer_norm(delta_pair_repr)
|
|
547
|
+
|
|
548
|
+
return x, curr_attn_mask, delta_pair_repr
|
|
549
|
+
|
|
550
|
+
|
|
551
|
+
# ==============================================================================
|
|
552
|
+
# Uni-Mol Model
|
|
553
|
+
# ==============================================================================
|
|
554
|
+
|
|
555
|
+
class UniMolModel(keras.Model):
|
|
556
|
+
r"""Multi-backend Uni-Mol model for 3D molecular representation learning and property prediction.
|
|
557
|
+
|
|
558
|
+
Supports molecular pretraining, candidate pocket pretraining, crystal, and OLED configurations.
|
|
559
|
+
|
|
560
|
+
Args:
|
|
561
|
+
output_dim (int, optional): Number of task output dimensions / classes. (default: ``2``)
|
|
562
|
+
data_type (str, optional): Data domain (``"molecule"``, ``"protein"``, ``"crystal"``, ``"oled"``). (default: ``"molecule"``)
|
|
563
|
+
vocab_size (int, optional): Vocabulary size for token dictionary. (default: ``512``)
|
|
564
|
+
encoder_layers (int, optional): Number of transformer encoder layers. (default: ``15``)
|
|
565
|
+
encoder_embed_dim (int, optional): Node embedding dimension. (default: ``512``)
|
|
566
|
+
encoder_ffn_embed_dim (int, optional): FFN hidden dimension. (default: ``2048``)
|
|
567
|
+
encoder_attention_heads (int, optional): Number of attention heads. (default: ``64``)
|
|
568
|
+
kernel (str, optional): GBF kernel type (``"gaussian"`` or ``"numerical"``). (default: ``"gaussian"``)
|
|
569
|
+
num_kernel (int, optional): Number of radial kernels. (default: ``128``)
|
|
570
|
+
pooler_dropout (float, optional): Dropout for classification head. (default: ``0.0``)
|
|
571
|
+
activation_fn (str, optional): Activation function name. (default: ``"gelu"``)
|
|
572
|
+
post_ln (bool, optional): Post-LN flag. (default: ``False``)
|
|
573
|
+
**kwargs: Additional model arguments.
|
|
574
|
+
|
|
575
|
+
Example:
|
|
576
|
+
```python
|
|
577
|
+
import numpy as np
|
|
578
|
+
from k3_node.models import UniMolModel
|
|
579
|
+
|
|
580
|
+
tokens = np.random.randint(1, 64, size=(2, 6)) # atom tokens of 2 molecules with 6 atoms
|
|
581
|
+
coords = np.random.rand(2, 6, 3).astype("float32") * 3.0 # 3D conformations
|
|
582
|
+
|
|
583
|
+
model = UniMolModel(output_dim=2, vocab_size=64, encoder_layers=2, encoder_embed_dim=32, encoder_ffn_embed_dim=64,
|
|
584
|
+
encoder_attention_heads=4, num_kernel=16)
|
|
585
|
+
logits = model(tokens, src_coord=coords) # molecule-level predictions
|
|
586
|
+
print(tuple(logits.shape)) # (2, 2)
|
|
587
|
+
reprs = model(tokens, src_coord=coords, return_repr=True)
|
|
588
|
+
print(tuple(reprs["cls_repr"].shape), tuple(reprs["encoder_rep"].shape)) # (2, 32) (2, 6, 32): molecule and atom embeddings
|
|
589
|
+
```
|
|
590
|
+
"""
|
|
591
|
+
|
|
592
|
+
def __init__(
|
|
593
|
+
self,
|
|
594
|
+
output_dim: int = 2,
|
|
595
|
+
data_type: str = "molecule",
|
|
596
|
+
vocab_size: int = 512,
|
|
597
|
+
encoder_layers: int = 15,
|
|
598
|
+
encoder_embed_dim: int = 512,
|
|
599
|
+
encoder_ffn_embed_dim: int = 2048,
|
|
600
|
+
encoder_attention_heads: int = 64,
|
|
601
|
+
kernel: str = "gaussian",
|
|
602
|
+
num_kernel: int = 128,
|
|
603
|
+
pooler_dropout: float = 0.0,
|
|
604
|
+
activation_fn: str = "gelu",
|
|
605
|
+
post_ln: bool = False,
|
|
606
|
+
**kwargs,
|
|
607
|
+
):
|
|
608
|
+
super().__init__(**kwargs)
|
|
609
|
+
self.output_dim = output_dim
|
|
610
|
+
self.data_type = data_type
|
|
611
|
+
self.vocab_size = vocab_size
|
|
612
|
+
self.encoder_layers = encoder_layers
|
|
613
|
+
self.encoder_embed_dim = encoder_embed_dim
|
|
614
|
+
self.encoder_ffn_embed_dim = encoder_ffn_embed_dim
|
|
615
|
+
self.encoder_attention_heads = encoder_attention_heads
|
|
616
|
+
self.kernel_type = kernel
|
|
617
|
+
self.num_kernel = num_kernel
|
|
618
|
+
self.pooler_dropout = pooler_dropout
|
|
619
|
+
self.activation_fn_name = activation_fn
|
|
620
|
+
self.post_ln = post_ln
|
|
621
|
+
|
|
622
|
+
self.padding_idx = 0
|
|
623
|
+
|
|
624
|
+
self.embed_tokens = layers.Embedding(vocab_size, encoder_embed_dim, name="embed_tokens")
|
|
625
|
+
|
|
626
|
+
n_edge_type = 1024
|
|
627
|
+
if kernel == "gaussian":
|
|
628
|
+
self.gbf = GaussianLayer(num_kernel, n_edge_type, name="gbf")
|
|
629
|
+
else:
|
|
630
|
+
self.gbf = NumericalEmbed(num_kernel, n_edge_type, activation_fn=activation_fn, name="gbf")
|
|
631
|
+
|
|
632
|
+
self.gbf_proj = NonLinearHead(
|
|
633
|
+
num_kernel, encoder_attention_heads, activation_fn=activation_fn, name="gbf_proj"
|
|
634
|
+
)
|
|
635
|
+
|
|
636
|
+
self.encoder = UniMolTransformerEncoder(
|
|
637
|
+
encoder_layers=encoder_layers,
|
|
638
|
+
embed_dim=encoder_embed_dim,
|
|
639
|
+
ffn_embed_dim=encoder_ffn_embed_dim,
|
|
640
|
+
attention_heads=encoder_attention_heads,
|
|
641
|
+
activation_fn=activation_fn,
|
|
642
|
+
post_ln=post_ln,
|
|
643
|
+
name="encoder",
|
|
644
|
+
)
|
|
645
|
+
|
|
646
|
+
self.classification_head = LinearHead(
|
|
647
|
+
input_dim=encoder_embed_dim,
|
|
648
|
+
num_classes=output_dim,
|
|
649
|
+
pooler_dropout=pooler_dropout,
|
|
650
|
+
name="classification_head",
|
|
651
|
+
)
|
|
652
|
+
|
|
653
|
+
self.pair2coord_proj = NonLinearHead(
|
|
654
|
+
encoder_attention_heads, 1, activation_fn=activation_fn, name="pair2coord_proj"
|
|
655
|
+
)
|
|
656
|
+
self.dist_head = DistanceHead(
|
|
657
|
+
encoder_attention_heads, activation_fn=activation_fn, name="dist_head"
|
|
658
|
+
)
|
|
659
|
+
self.lm_head = MaskLMHead(
|
|
660
|
+
encoder_embed_dim, vocab_size, activation_fn=activation_fn, name="lm_head"
|
|
661
|
+
)
|
|
662
|
+
|
|
663
|
+
def build(self, input_shape=None):
|
|
664
|
+
if not self.built:
|
|
665
|
+
self.embed_tokens.build(None)
|
|
666
|
+
self.gbf.build(None)
|
|
667
|
+
self.gbf_proj.build((None, None, None, self.num_kernel))
|
|
668
|
+
self.encoder.build(None)
|
|
669
|
+
self.classification_head.build((None, self.encoder_embed_dim))
|
|
670
|
+
self.pair2coord_proj.build((None, None, None, self.encoder_attention_heads))
|
|
671
|
+
self.dist_head.build((None, None, None, self.encoder_attention_heads))
|
|
672
|
+
self.lm_head.build((None, None, self.encoder_embed_dim))
|
|
673
|
+
super().build(input_shape)
|
|
674
|
+
|
|
675
|
+
def call(
|
|
676
|
+
self,
|
|
677
|
+
src_tokens,
|
|
678
|
+
src_distance=None,
|
|
679
|
+
src_coord=None,
|
|
680
|
+
src_edge_type=None,
|
|
681
|
+
padding_mask=None,
|
|
682
|
+
return_repr: bool = False,
|
|
683
|
+
return_atomic_reprs: bool = False,
|
|
684
|
+
features_only: bool = False,
|
|
685
|
+
training: bool = False,
|
|
686
|
+
):
|
|
687
|
+
if isinstance(src_tokens, dict):
|
|
688
|
+
src_distance = src_tokens.get("src_distance", src_distance)
|
|
689
|
+
src_coord = src_tokens.get("src_coord", src_coord)
|
|
690
|
+
src_edge_type = src_tokens.get("src_edge_type", src_edge_type)
|
|
691
|
+
padding_mask = src_tokens.get("padding_mask", padding_mask)
|
|
692
|
+
src_tokens = src_tokens.get("src_tokens", src_tokens.get("tokens"))
|
|
693
|
+
elif isinstance(src_tokens, (tuple, list)):
|
|
694
|
+
if len(src_tokens) == 2:
|
|
695
|
+
src_tokens, src_coord = src_tokens
|
|
696
|
+
elif len(src_tokens) >= 3:
|
|
697
|
+
src_tokens, src_distance, src_coord = src_tokens[:3]
|
|
698
|
+
r"""Forward pass for Uni-Mol.
|
|
699
|
+
|
|
700
|
+
Args:
|
|
701
|
+
src_tokens (Tensor): Token indices of shape ``[batch_size, seq_len]``.
|
|
702
|
+
src_distance (Tensor, optional): Pairwise distances ``[batch_size, seq_len, seq_len]``.
|
|
703
|
+
src_coord (Tensor, optional): 3D coordinates ``[batch_size, seq_len, 3]``.
|
|
704
|
+
src_edge_type (Tensor, optional): Pairwise edge types ``[batch_size, seq_len, seq_len]``.
|
|
705
|
+
padding_mask (Tensor, optional): Padding boolean mask ``[batch_size, seq_len]``.
|
|
706
|
+
return_repr (bool, optional): Return CLS/atomic representations.
|
|
707
|
+
features_only (bool, optional): Only return representations, skipping heads.
|
|
708
|
+
training (bool, optional): Training mode flag.
|
|
709
|
+
"""
|
|
710
|
+
shape = ops.shape(src_tokens)
|
|
711
|
+
bsz = shape[0]
|
|
712
|
+
seq_len = shape[1]
|
|
713
|
+
|
|
714
|
+
# Auto-compute padding_mask if not provided
|
|
715
|
+
if padding_mask is None:
|
|
716
|
+
padding_mask = ops.equal(src_tokens, self.padding_idx)
|
|
717
|
+
|
|
718
|
+
# Auto-compute src_distance from src_coord if distance is missing
|
|
719
|
+
if src_distance is None and src_coord is not None:
|
|
720
|
+
diff = ops.expand_dims(src_coord, axis=2) - ops.expand_dims(src_coord, axis=1)
|
|
721
|
+
src_distance = ops.sqrt(ops.sum(ops.power(diff, 2), axis=-1) + 1e-10)
|
|
722
|
+
elif src_distance is None:
|
|
723
|
+
src_distance = ops.zeros((bsz, seq_len, seq_len), dtype="float32")
|
|
724
|
+
|
|
725
|
+
# Auto-compute src_edge_type from tokens if missing
|
|
726
|
+
if src_edge_type is None:
|
|
727
|
+
n_types = 32
|
|
728
|
+
src_edge_type = ops.expand_dims(src_tokens, axis=-1) * n_types + ops.expand_dims(src_tokens, axis=1)
|
|
729
|
+
src_edge_type = ops.cast(ops.mod(src_edge_type, 1024), "int32")
|
|
730
|
+
|
|
731
|
+
# 1. Embeddings & GBF
|
|
732
|
+
x = self.embed_tokens(src_tokens)
|
|
733
|
+
gbf_feature = self.gbf(src_distance, src_edge_type)
|
|
734
|
+
graph_attn_bias = self.gbf_proj(gbf_feature) # [B, N, N, H]
|
|
735
|
+
|
|
736
|
+
# 2. Encoder
|
|
737
|
+
encoder_rep, encoder_pair_rep, delta_pair_rep = self.encoder(
|
|
738
|
+
x,
|
|
739
|
+
attn_mask=graph_attn_bias,
|
|
740
|
+
padding_mask=padding_mask,
|
|
741
|
+
training=training,
|
|
742
|
+
)
|
|
743
|
+
|
|
744
|
+
cls_repr = encoder_rep[:, 0, :]
|
|
745
|
+
|
|
746
|
+
if return_repr:
|
|
747
|
+
res = {"cls_repr": cls_repr, "encoder_rep": encoder_rep, "pair_rep": encoder_pair_rep}
|
|
748
|
+
return res
|
|
749
|
+
|
|
750
|
+
if features_only:
|
|
751
|
+
return encoder_rep, encoder_pair_rep
|
|
752
|
+
|
|
753
|
+
# Classification / Regression logits
|
|
754
|
+
logits = self.classification_head(cls_repr, training=training)
|
|
755
|
+
return logits
|
|
756
|
+
|
|
757
|
+
|
|
758
|
+
class UniMolConfGenModel(UniMolModel):
|
|
759
|
+
r"""Uni-Mol Conformation Generation Model for iterative 3D geometry prediction.
|
|
760
|
+
|
|
761
|
+
Example:
|
|
762
|
+
```python
|
|
763
|
+
import numpy as np
|
|
764
|
+
from k3_node.models import UniMolConfGenModel
|
|
765
|
+
|
|
766
|
+
tokens = np.random.randint(1, 64, size=(2, 6)) # atom tokens of 2 molecules with 6 atoms
|
|
767
|
+
coords = np.random.rand(2, 6, 3).astype("float32") * 3.0 # 3D conformations
|
|
768
|
+
|
|
769
|
+
model = UniMolConfGenModel(vocab_size=64, encoder_layers=2, encoder_embed_dim=32, encoder_ffn_embed_dim=64,
|
|
770
|
+
encoder_attention_heads=4, num_kernel=16)
|
|
771
|
+
new_coords, pred_dist = model(tokens, src_coord=coords) # refined conformation, predicted distances
|
|
772
|
+
print(tuple(new_coords.shape), tuple(pred_dist.shape)) # (2, 6, 3) (2, 6, 6)
|
|
773
|
+
```
|
|
774
|
+
"""
|
|
775
|
+
|
|
776
|
+
def call(
|
|
777
|
+
self,
|
|
778
|
+
src_tokens,
|
|
779
|
+
src_distance=None,
|
|
780
|
+
src_coord=None,
|
|
781
|
+
src_edge_type=None,
|
|
782
|
+
padding_mask=None,
|
|
783
|
+
training: bool = False,
|
|
784
|
+
):
|
|
785
|
+
if isinstance(src_tokens, dict):
|
|
786
|
+
src_distance = src_tokens.get("src_distance", src_distance)
|
|
787
|
+
src_coord = src_tokens.get("src_coord", src_coord)
|
|
788
|
+
src_edge_type = src_tokens.get("src_edge_type", src_edge_type)
|
|
789
|
+
padding_mask = src_tokens.get("padding_mask", padding_mask)
|
|
790
|
+
src_tokens = src_tokens.get("src_tokens", src_tokens.get("tokens"))
|
|
791
|
+
elif isinstance(src_tokens, (tuple, list)):
|
|
792
|
+
if len(src_tokens) == 2:
|
|
793
|
+
src_tokens, src_coord = src_tokens
|
|
794
|
+
elif len(src_tokens) >= 3:
|
|
795
|
+
src_tokens, src_distance, src_coord = src_tokens[:3]
|
|
796
|
+
|
|
797
|
+
if padding_mask is None:
|
|
798
|
+
padding_mask = ops.equal(src_tokens, self.padding_idx)
|
|
799
|
+
|
|
800
|
+
shape = ops.shape(src_tokens)
|
|
801
|
+
bsz = shape[0]
|
|
802
|
+
seq_len = shape[1]
|
|
803
|
+
|
|
804
|
+
if src_coord is None:
|
|
805
|
+
src_coord = ops.zeros((bsz, seq_len, 3), dtype="float32")
|
|
806
|
+
|
|
807
|
+
if src_distance is None:
|
|
808
|
+
diff = ops.expand_dims(src_coord, axis=2) - ops.expand_dims(src_coord, axis=1)
|
|
809
|
+
src_distance = ops.sqrt(ops.sum(ops.power(diff, 2), axis=-1) + 1e-10)
|
|
810
|
+
|
|
811
|
+
if src_edge_type is None:
|
|
812
|
+
src_edge_type = ops.cast(ops.mod(ops.expand_dims(src_tokens, axis=-1) * 32 + ops.expand_dims(src_tokens, axis=1), 1024), "int32")
|
|
813
|
+
|
|
814
|
+
x = self.embed_tokens(src_tokens)
|
|
815
|
+
gbf_feature = self.gbf(src_distance, src_edge_type)
|
|
816
|
+
graph_attn_bias = self.gbf_proj(gbf_feature)
|
|
817
|
+
|
|
818
|
+
encoder_rep, encoder_pair_rep, delta_pair_rep = self.encoder(
|
|
819
|
+
x,
|
|
820
|
+
attn_mask=graph_attn_bias,
|
|
821
|
+
padding_mask=padding_mask,
|
|
822
|
+
training=training,
|
|
823
|
+
)
|
|
824
|
+
|
|
825
|
+
# Coordinate update from delta_pair_rep
|
|
826
|
+
attn_probs = self.pair2coord_proj(delta_pair_rep) # [B, N, N, 1]
|
|
827
|
+
delta_pos = ops.expand_dims(src_coord, axis=1) - ops.expand_dims(src_coord, axis=2)
|
|
828
|
+
coord_update = delta_pos * attn_probs
|
|
829
|
+
updated_coord = src_coord + ops.sum(coord_update, axis=2)
|
|
830
|
+
|
|
831
|
+
pred_dist = self.dist_head(encoder_pair_rep)
|
|
832
|
+
return updated_coord, pred_dist
|
|
833
|
+
|
|
834
|
+
|
|
835
|
+
class UniMolDockingModel(UniMolModel):
|
|
836
|
+
r"""Uni-Mol Protein-Ligand Binding Pose Prediction Model.
|
|
837
|
+
|
|
838
|
+
Example:
|
|
839
|
+
```python
|
|
840
|
+
import numpy as np
|
|
841
|
+
from k3_node.models import UniMolDockingModel
|
|
842
|
+
|
|
843
|
+
tokens = np.random.randint(1, 64, size=(2, 6)) # atom tokens of 2 molecules with 6 atoms
|
|
844
|
+
coords = np.random.rand(2, 6, 3).astype("float32") * 3.0 # 3D conformations
|
|
845
|
+
|
|
846
|
+
model = UniMolDockingModel(vocab_size=64, encoder_layers=2, encoder_embed_dim=32, encoder_ffn_embed_dim=64,
|
|
847
|
+
encoder_attention_heads=4, num_kernel=16)
|
|
848
|
+
pose, pred_dist = model(tokens, src_coord=coords)
|
|
849
|
+
print(tuple(pose.shape), tuple(pred_dist.shape)) # (2, 6, 3) (2, 6, 6)
|
|
850
|
+
```
|
|
851
|
+
"""
|
|
852
|
+
|
|
853
|
+
def __init__(self, **kwargs):
|
|
854
|
+
super().__init__(**kwargs)
|
|
855
|
+
# Cross-layer coordinate prediction head
|
|
856
|
+
self.docking_coord_head = NonLinearHead(
|
|
857
|
+
self.encoder_attention_heads, 1, activation_fn="gelu", name="docking_coord_head"
|
|
858
|
+
)
|
|
859
|
+
|
|
860
|
+
def build(self, input_shape=None):
|
|
861
|
+
if not self.built:
|
|
862
|
+
self.docking_coord_head.build((None, None, None, self.encoder_attention_heads))
|
|
863
|
+
super().build(input_shape)
|
|
864
|
+
|
|
865
|
+
def call(
|
|
866
|
+
self,
|
|
867
|
+
src_tokens,
|
|
868
|
+
src_distance=None,
|
|
869
|
+
src_coord=None,
|
|
870
|
+
src_edge_type=None,
|
|
871
|
+
padding_mask=None,
|
|
872
|
+
training: bool = False,
|
|
873
|
+
):
|
|
874
|
+
if isinstance(src_tokens, dict):
|
|
875
|
+
src_distance = src_tokens.get("src_distance", src_distance)
|
|
876
|
+
src_coord = src_tokens.get("src_coord", src_coord)
|
|
877
|
+
src_edge_type = src_tokens.get("src_edge_type", src_edge_type)
|
|
878
|
+
padding_mask = src_tokens.get("padding_mask", padding_mask)
|
|
879
|
+
src_tokens = src_tokens.get("src_tokens", src_tokens.get("tokens"))
|
|
880
|
+
elif isinstance(src_tokens, (tuple, list)):
|
|
881
|
+
if len(src_tokens) == 2:
|
|
882
|
+
src_tokens, src_coord = src_tokens
|
|
883
|
+
elif len(src_tokens) >= 3:
|
|
884
|
+
src_tokens, src_distance, src_coord = src_tokens[:3]
|
|
885
|
+
|
|
886
|
+
if padding_mask is None:
|
|
887
|
+
padding_mask = ops.equal(src_tokens, self.padding_idx)
|
|
888
|
+
|
|
889
|
+
shape = ops.shape(src_tokens)
|
|
890
|
+
bsz = shape[0]
|
|
891
|
+
seq_len = shape[1]
|
|
892
|
+
|
|
893
|
+
if src_coord is None:
|
|
894
|
+
src_coord = ops.zeros((bsz, seq_len, 3), dtype="float32")
|
|
895
|
+
|
|
896
|
+
if src_distance is None:
|
|
897
|
+
diff = ops.expand_dims(src_coord, axis=2) - ops.expand_dims(src_coord, axis=1)
|
|
898
|
+
src_distance = ops.sqrt(ops.sum(ops.power(diff, 2), axis=-1) + 1e-10)
|
|
899
|
+
|
|
900
|
+
if src_edge_type is None:
|
|
901
|
+
src_edge_type = ops.cast(ops.mod(ops.expand_dims(src_tokens, axis=-1) * 32 + ops.expand_dims(src_tokens, axis=1), 1024), "int32")
|
|
902
|
+
|
|
903
|
+
x = self.embed_tokens(src_tokens)
|
|
904
|
+
gbf_feature = self.gbf(src_distance, src_edge_type)
|
|
905
|
+
graph_attn_bias = self.gbf_proj(gbf_feature)
|
|
906
|
+
|
|
907
|
+
encoder_rep, encoder_pair_rep, delta_pair_rep = self.encoder(
|
|
908
|
+
x,
|
|
909
|
+
attn_mask=graph_attn_bias,
|
|
910
|
+
padding_mask=padding_mask,
|
|
911
|
+
training=training,
|
|
912
|
+
)
|
|
913
|
+
|
|
914
|
+
probs = self.docking_coord_head(delta_pair_rep)
|
|
915
|
+
diff_coord = ops.expand_dims(src_coord, axis=1) - ops.expand_dims(src_coord, axis=2)
|
|
916
|
+
pose_update = src_coord + ops.sum(diff_coord * probs, axis=2)
|
|
917
|
+
pred_dist = self.dist_head(encoder_pair_rep)
|
|
918
|
+
return pose_update, pred_dist
|
|
919
|
+
|
|
920
|
+
|
|
921
|
+
# ==============================================================================
|
|
922
|
+
# Helper Functions: Download & Load Weights
|
|
923
|
+
# ==============================================================================
|
|
924
|
+
|
|
925
|
+
def download_unimol_checkpoint(
|
|
926
|
+
name: str = "mol_pre_no_h",
|
|
927
|
+
folder: str = "checkpoints",
|
|
928
|
+
log: bool = True,
|
|
929
|
+
) -> str:
|
|
930
|
+
r"""Downloads a pre-trained Uni-Mol checkpoint (.pt).
|
|
931
|
+
|
|
932
|
+
Args:
|
|
933
|
+
name (str): Pretrained checkpoint name or alias (e.g. ``"mol_pre_no_h"``,
|
|
934
|
+
``"mol_pre_all_h"``, ``"pocket_pre"``, ``"mp_all_h"``, ``"oled_pre_no_h"``,
|
|
935
|
+
``"qm9"``, ``"drugs"``, ``"binding_pose"``).
|
|
936
|
+
folder (str, optional): Target directory to save the checkpoint. (default: ``"checkpoints"``)
|
|
937
|
+
log (bool, optional): Whether to print download progress. (default: ``True``)
|
|
938
|
+
|
|
939
|
+
Returns:
|
|
940
|
+
str: Absolute path to the downloaded checkpoint file.
|
|
941
|
+
"""
|
|
942
|
+
clean_name = name.strip().lower()
|
|
943
|
+
if clean_name in UNIMOL_ALIASES:
|
|
944
|
+
clean_name = UNIMOL_ALIASES[clean_name]
|
|
945
|
+
|
|
946
|
+
if clean_name not in UNIMOL_PRETRAINED_URLS:
|
|
947
|
+
for k, v in UNIMOL_PRETRAINED_URLS.items():
|
|
948
|
+
if v["filename"] == name or k.lower() == clean_name:
|
|
949
|
+
clean_name = k
|
|
950
|
+
break
|
|
951
|
+
|
|
952
|
+
if clean_name not in UNIMOL_PRETRAINED_URLS:
|
|
953
|
+
raise ValueError(
|
|
954
|
+
f"Unknown Uni-Mol checkpoint '{name}'. Available: {list(UNIMOL_PRETRAINED_URLS.keys())}"
|
|
955
|
+
)
|
|
956
|
+
|
|
957
|
+
info = UNIMOL_PRETRAINED_URLS[clean_name]
|
|
958
|
+
filename = info["filename"]
|
|
959
|
+
local_path = os.path.join(folder, filename)
|
|
960
|
+
|
|
961
|
+
if os.path.exists(local_path):
|
|
962
|
+
return local_path
|
|
963
|
+
|
|
964
|
+
# Check alternative local paths
|
|
965
|
+
alt_paths = [
|
|
966
|
+
os.path.join("Uni-Mol", "unimol", filename),
|
|
967
|
+
os.path.join("Uni-Mol", "unimol_tools", "unimol_tools", "weights", filename),
|
|
968
|
+
]
|
|
969
|
+
for alt in alt_paths:
|
|
970
|
+
if os.path.exists(alt):
|
|
971
|
+
return alt
|
|
972
|
+
|
|
973
|
+
return download_url(info["url"], folder=folder, filename=filename, log=log)
|
|
974
|
+
|
|
975
|
+
|
|
976
|
+
def load_unimol_weights(
|
|
977
|
+
model: UniMolModel,
|
|
978
|
+
checkpoint_path: Optional[str] = None,
|
|
979
|
+
pretrained_name: Optional[str] = None,
|
|
980
|
+
folder: str = "checkpoints",
|
|
981
|
+
download: bool = True,
|
|
982
|
+
) -> UniMolModel:
|
|
983
|
+
r"""Loads pre-trained weights from a PyTorch checkpoint (.pt) into a Keras 3 UniMolModel.
|
|
984
|
+
|
|
985
|
+
Args:
|
|
986
|
+
model (UniMolModel): The target UniMolModel instance.
|
|
987
|
+
checkpoint_path (str, optional): Local path to .pt file or checkpoint name.
|
|
988
|
+
pretrained_name (str, optional): Pretrained model identifier.
|
|
989
|
+
folder (str, optional): Directory to store downloaded checkpoints. (default: ``"checkpoints"``)
|
|
990
|
+
download (bool, optional): Whether to download checkpoint if missing locally. (default: ``True``)
|
|
991
|
+
|
|
992
|
+
Returns:
|
|
993
|
+
UniMolModel: The model with loaded weights.
|
|
994
|
+
"""
|
|
995
|
+
path_to_load = checkpoint_path
|
|
996
|
+
|
|
997
|
+
candidate = pretrained_name or checkpoint_path
|
|
998
|
+
if candidate:
|
|
999
|
+
clean = candidate.strip().lower()
|
|
1000
|
+
if clean in UNIMOL_ALIASES:
|
|
1001
|
+
candidate = UNIMOL_ALIASES[clean]
|
|
1002
|
+
|
|
1003
|
+
if path_to_load is None:
|
|
1004
|
+
if candidate is None:
|
|
1005
|
+
raise ValueError("Either checkpoint_path or pretrained_name must be specified.")
|
|
1006
|
+
if candidate in UNIMOL_PRETRAINED_URLS:
|
|
1007
|
+
fname = UNIMOL_PRETRAINED_URLS[candidate]["filename"]
|
|
1008
|
+
local_target = os.path.join(folder, fname)
|
|
1009
|
+
if os.path.isfile(local_target):
|
|
1010
|
+
path_to_load = local_target
|
|
1011
|
+
elif download:
|
|
1012
|
+
path_to_load = download_unimol_checkpoint(candidate, folder=folder)
|
|
1013
|
+
else:
|
|
1014
|
+
raise FileNotFoundError(f"Checkpoint for '{candidate}' not found at '{local_target}'.")
|
|
1015
|
+
else:
|
|
1016
|
+
raise ValueError(f"Unknown checkpoint '{candidate}'.")
|
|
1017
|
+
elif not os.path.isfile(path_to_load) and download and candidate in UNIMOL_PRETRAINED_URLS:
|
|
1018
|
+
path_to_load = download_unimol_checkpoint(candidate, folder=folder)
|
|
1019
|
+
|
|
1020
|
+
import torch
|
|
1021
|
+
|
|
1022
|
+
state = torch.load(path_to_load, map_location="cpu")
|
|
1023
|
+
if isinstance(state, dict):
|
|
1024
|
+
if "model" in state:
|
|
1025
|
+
state_dict = state["model"]
|
|
1026
|
+
elif "model_state_dict" in state:
|
|
1027
|
+
state_dict = state["model_state_dict"]
|
|
1028
|
+
else:
|
|
1029
|
+
state_dict = state
|
|
1030
|
+
else:
|
|
1031
|
+
raise ValueError(f"Expected dict/state_dict in checkpoint, got {type(state)}")
|
|
1032
|
+
|
|
1033
|
+
if not model.built:
|
|
1034
|
+
model.build(None)
|
|
1035
|
+
|
|
1036
|
+
def _to_tensor(t):
|
|
1037
|
+
if hasattr(t, "detach"):
|
|
1038
|
+
t = t.detach()
|
|
1039
|
+
if hasattr(t, "cpu"):
|
|
1040
|
+
t = t.cpu()
|
|
1041
|
+
if hasattr(t, "numpy"):
|
|
1042
|
+
t = t.numpy()
|
|
1043
|
+
return ops.convert_to_tensor(np.array(t, dtype=np.float32), dtype="float32")
|
|
1044
|
+
|
|
1045
|
+
# 1. Embeddings
|
|
1046
|
+
if "embed_tokens.weight" in state_dict:
|
|
1047
|
+
w = state_dict["embed_tokens.weight"]
|
|
1048
|
+
model.embed_tokens.embeddings.assign(_to_tensor(w))
|
|
1049
|
+
|
|
1050
|
+
# 2. GBF
|
|
1051
|
+
if "gbf.means.weight" in state_dict:
|
|
1052
|
+
model.gbf.means.embeddings.assign(_to_tensor(state_dict["gbf.means.weight"]))
|
|
1053
|
+
if "gbf.stds.weight" in state_dict:
|
|
1054
|
+
model.gbf.stds.embeddings.assign(_to_tensor(state_dict["gbf.stds.weight"]))
|
|
1055
|
+
if "gbf.mul.weight" in state_dict:
|
|
1056
|
+
model.gbf.mul.embeddings.assign(_to_tensor(state_dict["gbf.mul.weight"]))
|
|
1057
|
+
if "gbf.bias.weight" in state_dict:
|
|
1058
|
+
model.gbf.bias.embeddings.assign(_to_tensor(state_dict["gbf.bias.weight"]))
|
|
1059
|
+
|
|
1060
|
+
# 3. GBF Proj
|
|
1061
|
+
if "gbf_proj.linear1.weight" in state_dict:
|
|
1062
|
+
model.gbf_proj.linear1.kernel.assign(_to_tensor(state_dict["gbf_proj.linear1.weight"].t()))
|
|
1063
|
+
if "gbf_proj.linear1.bias" in state_dict:
|
|
1064
|
+
model.gbf_proj.linear1.bias.assign(_to_tensor(state_dict["gbf_proj.linear1.bias"]))
|
|
1065
|
+
if "gbf_proj.linear2.weight" in state_dict:
|
|
1066
|
+
model.gbf_proj.linear2.kernel.assign(_to_tensor(state_dict["gbf_proj.linear2.weight"].t()))
|
|
1067
|
+
if "gbf_proj.linear2.bias" in state_dict:
|
|
1068
|
+
model.gbf_proj.linear2.bias.assign(_to_tensor(state_dict["gbf_proj.linear2.bias"]))
|
|
1069
|
+
|
|
1070
|
+
# 4. Encoder Layers
|
|
1071
|
+
for i, enc_layer in enumerate(model.encoder.layers_list):
|
|
1072
|
+
prefix = f"encoder.layers.{i}"
|
|
1073
|
+
alt_prefix = f"layers.{i}"
|
|
1074
|
+
|
|
1075
|
+
def _get(name):
|
|
1076
|
+
if f"{prefix}.{name}" in state_dict:
|
|
1077
|
+
return state_dict[f"{prefix}.{name}"]
|
|
1078
|
+
elif f"{alt_prefix}.{name}" in state_dict:
|
|
1079
|
+
return state_dict[f"{alt_prefix}.{name}"]
|
|
1080
|
+
return None
|
|
1081
|
+
|
|
1082
|
+
# Self-Attention
|
|
1083
|
+
in_w = _get("self_attn.in_proj.weight")
|
|
1084
|
+
if in_w is not None:
|
|
1085
|
+
enc_layer.self_attn.in_proj.kernel.assign(_to_tensor(in_w.t()))
|
|
1086
|
+
in_b = _get("self_attn.in_proj.bias")
|
|
1087
|
+
if in_b is not None:
|
|
1088
|
+
enc_layer.self_attn.in_proj.bias.assign(_to_tensor(in_b))
|
|
1089
|
+
|
|
1090
|
+
out_w = _get("self_attn.out_proj.weight")
|
|
1091
|
+
if out_w is not None:
|
|
1092
|
+
enc_layer.self_attn.out_proj.kernel.assign(_to_tensor(out_w.t()))
|
|
1093
|
+
out_b = _get("self_attn.out_proj.bias")
|
|
1094
|
+
if out_b is not None:
|
|
1095
|
+
enc_layer.self_attn.out_proj.bias.assign(_to_tensor(out_b))
|
|
1096
|
+
|
|
1097
|
+
# Layer norms
|
|
1098
|
+
attn_ln_w = _get("self_attn_layer_norm.weight")
|
|
1099
|
+
if attn_ln_w is not None and enc_layer.self_attn_layer_norm.gamma is not None:
|
|
1100
|
+
enc_layer.self_attn_layer_norm.gamma.assign(_to_tensor(attn_ln_w))
|
|
1101
|
+
attn_ln_b = _get("self_attn_layer_norm.bias")
|
|
1102
|
+
if attn_ln_b is not None and enc_layer.self_attn_layer_norm.beta is not None:
|
|
1103
|
+
enc_layer.self_attn_layer_norm.beta.assign(_to_tensor(attn_ln_b))
|
|
1104
|
+
|
|
1105
|
+
# FFN
|
|
1106
|
+
fc1_w = _get("fc1.weight")
|
|
1107
|
+
if fc1_w is not None:
|
|
1108
|
+
enc_layer.fc1.kernel.assign(_to_tensor(fc1_w.t()))
|
|
1109
|
+
fc1_b = _get("fc1.bias")
|
|
1110
|
+
if fc1_b is not None:
|
|
1111
|
+
enc_layer.fc1.bias.assign(_to_tensor(fc1_b))
|
|
1112
|
+
|
|
1113
|
+
fc2_w = _get("fc2.weight")
|
|
1114
|
+
if fc2_w is not None:
|
|
1115
|
+
enc_layer.fc2.kernel.assign(_to_tensor(fc2_w.t()))
|
|
1116
|
+
fc2_b = _get("fc2.bias")
|
|
1117
|
+
if fc2_b is not None:
|
|
1118
|
+
enc_layer.fc2.bias.assign(_to_tensor(fc2_b))
|
|
1119
|
+
|
|
1120
|
+
final_ln_w = _get("final_layer_norm.weight")
|
|
1121
|
+
if final_ln_w is not None and enc_layer.final_layer_norm.gamma is not None:
|
|
1122
|
+
enc_layer.final_layer_norm.gamma.assign(_to_tensor(final_ln_w))
|
|
1123
|
+
final_ln_b = _get("final_layer_norm.bias")
|
|
1124
|
+
if final_ln_b is not None and enc_layer.final_layer_norm.beta is not None:
|
|
1125
|
+
enc_layer.final_layer_norm.beta.assign(_to_tensor(final_ln_b))
|
|
1126
|
+
|
|
1127
|
+
# Encoder global norms
|
|
1128
|
+
for key, attr in [
|
|
1129
|
+
("encoder.emb_layer_norm.weight", model.encoder.emb_layer_norm.gamma),
|
|
1130
|
+
("encoder.emb_layer_norm.bias", model.encoder.emb_layer_norm.beta),
|
|
1131
|
+
]:
|
|
1132
|
+
if key in state_dict and attr is not None:
|
|
1133
|
+
attr.assign(_to_tensor(state_dict[key]))
|
|
1134
|
+
|
|
1135
|
+
if model.encoder.final_layer_norm is not None:
|
|
1136
|
+
if "encoder.final_layer_norm.weight" in state_dict:
|
|
1137
|
+
model.encoder.final_layer_norm.gamma.assign(_to_tensor(state_dict["encoder.final_layer_norm.weight"]))
|
|
1138
|
+
if "encoder.final_layer_norm.bias" in state_dict:
|
|
1139
|
+
model.encoder.final_layer_norm.beta.assign(_to_tensor(state_dict["encoder.final_layer_norm.bias"]))
|
|
1140
|
+
|
|
1141
|
+
if model.encoder.final_head_layer_norm is not None:
|
|
1142
|
+
if "encoder.final_head_layer_norm.weight" in state_dict:
|
|
1143
|
+
model.encoder.final_head_layer_norm.gamma.assign(_to_tensor(state_dict["encoder.final_head_layer_norm.weight"]))
|
|
1144
|
+
if "encoder.final_head_layer_norm.bias" in state_dict:
|
|
1145
|
+
model.encoder.final_head_layer_norm.beta.assign(_to_tensor(state_dict["encoder.final_head_layer_norm.bias"]))
|
|
1146
|
+
|
|
1147
|
+
# Classification / linear head
|
|
1148
|
+
if hasattr(model, "classification_head"):
|
|
1149
|
+
for k in ["classification_head.out_proj.weight", "classification_heads.target.out_proj.weight"]:
|
|
1150
|
+
if k in state_dict and hasattr(model.classification_head, "out_proj"):
|
|
1151
|
+
model.classification_head.out_proj.kernel.assign(_to_tensor(state_dict[k].t()))
|
|
1152
|
+
for k in ["classification_head.out_proj.bias", "classification_heads.target.out_proj.bias"]:
|
|
1153
|
+
if k in state_dict and hasattr(model.classification_head, "out_proj"):
|
|
1154
|
+
model.classification_head.out_proj.bias.assign(_to_tensor(state_dict[k]))
|
|
1155
|
+
|
|
1156
|
+
return model
|