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,954 @@
|
|
|
1
|
+
import os
|
|
2
|
+
from typing import Optional, Union, Tuple, List, Callable
|
|
3
|
+
import numpy as np
|
|
4
|
+
|
|
5
|
+
import keras
|
|
6
|
+
from keras import layers, ops
|
|
7
|
+
|
|
8
|
+
from k3_node.layers.conv.utils import softmax
|
|
9
|
+
from k3_node.data.download import download_google_url
|
|
10
|
+
from k3_node.ops.segment import segment_sum
|
|
11
|
+
|
|
12
|
+
|
|
13
|
+
def sce_loss(x, y, alpha: float = 3.0):
|
|
14
|
+
r"""Scaled Cosine Error (SCE) loss from `"GraphMAE: Masked Autoencoding for Graph
|
|
15
|
+
Self-Supervised Learning" <https://arxiv.org/abs/2205.10803>`_ and GraphMAE2.
|
|
16
|
+
|
|
17
|
+
Args:
|
|
18
|
+
x (Tensor): Predicted node representations.
|
|
19
|
+
y (Tensor): Target node representations.
|
|
20
|
+
alpha (float, optional): Scaling exponent. (default: ``3.0``)
|
|
21
|
+
"""
|
|
22
|
+
x_norm = ops.sqrt(ops.sum(ops.power(x, 2), axis=-1, keepdims=True) + 1e-12)
|
|
23
|
+
x = x / x_norm
|
|
24
|
+
y_norm = ops.sqrt(ops.sum(ops.power(y, 2), axis=-1, keepdims=True) + 1e-12)
|
|
25
|
+
y = y / y_norm
|
|
26
|
+
|
|
27
|
+
cos_sim = ops.sum(x * y, axis=-1)
|
|
28
|
+
diff = ops.clip(1.0 - cos_sim, 0.0, 2.0)
|
|
29
|
+
loss = ops.power(diff, alpha)
|
|
30
|
+
return ops.mean(loss)
|
|
31
|
+
|
|
32
|
+
|
|
33
|
+
def _get_activation(name: Optional[Union[str, Callable]]):
|
|
34
|
+
if name is None:
|
|
35
|
+
return None
|
|
36
|
+
if isinstance(name, str):
|
|
37
|
+
name_lower = name.lower()
|
|
38
|
+
if name_lower == "prelu":
|
|
39
|
+
return layers.PReLU(shared_axes=[1])
|
|
40
|
+
elif name_lower == "relu":
|
|
41
|
+
return layers.ReLU()
|
|
42
|
+
elif name_lower == "gelu":
|
|
43
|
+
return layers.Activation("gelu")
|
|
44
|
+
elif name_lower == "silu":
|
|
45
|
+
return layers.Activation("silu")
|
|
46
|
+
elif name_lower == "elu":
|
|
47
|
+
return layers.ELU()
|
|
48
|
+
else:
|
|
49
|
+
return layers.Activation(name)
|
|
50
|
+
elif isinstance(name, layers.Layer):
|
|
51
|
+
return name
|
|
52
|
+
elif callable(name):
|
|
53
|
+
return layers.Activation(name)
|
|
54
|
+
return None
|
|
55
|
+
|
|
56
|
+
|
|
57
|
+
def _get_norm(name: Optional[str], dim: int):
|
|
58
|
+
if name is None:
|
|
59
|
+
return None
|
|
60
|
+
name_lower = name.lower()
|
|
61
|
+
if name_lower in ("layernorm", "layer_norm"):
|
|
62
|
+
return layers.LayerNormalization(axis=-1, epsilon=1e-5)
|
|
63
|
+
elif name_lower in ("batchnorm", "batch_norm"):
|
|
64
|
+
return layers.BatchNormalization(axis=-1, momentum=0.9, epsilon=1e-5)
|
|
65
|
+
return None
|
|
66
|
+
|
|
67
|
+
|
|
68
|
+
class GraphMAE2GATConv(layers.Layer):
|
|
69
|
+
r"""GAT convolution layer matching GraphMAE2's architecture."""
|
|
70
|
+
|
|
71
|
+
def __init__(
|
|
72
|
+
self,
|
|
73
|
+
in_feats: int,
|
|
74
|
+
out_feats: int,
|
|
75
|
+
num_heads: int,
|
|
76
|
+
feat_drop: float = 0.0,
|
|
77
|
+
attn_drop: float = 0.0,
|
|
78
|
+
negative_slope: float = 0.2,
|
|
79
|
+
residual: bool = False,
|
|
80
|
+
activation: Optional[Union[str, Callable]] = None,
|
|
81
|
+
bias: bool = True,
|
|
82
|
+
norm: Optional[str] = None,
|
|
83
|
+
concat_out: bool = True,
|
|
84
|
+
**kwargs,
|
|
85
|
+
):
|
|
86
|
+
super().__init__(**kwargs)
|
|
87
|
+
self.in_feats = in_feats
|
|
88
|
+
self.out_feats = out_feats
|
|
89
|
+
self.num_heads = num_heads
|
|
90
|
+
self.feat_drop_rate = feat_drop
|
|
91
|
+
self.attn_drop_rate = attn_drop
|
|
92
|
+
self.negative_slope = negative_slope
|
|
93
|
+
self.use_residual = residual
|
|
94
|
+
self.concat_out = concat_out
|
|
95
|
+
self.use_bias = bias
|
|
96
|
+
self.norm_name = norm
|
|
97
|
+
self.act_name = activation
|
|
98
|
+
|
|
99
|
+
self.fc = layers.Dense(num_heads * out_feats, use_bias=False)
|
|
100
|
+
self.feat_drop = layers.Dropout(feat_drop) if feat_drop > 0.0 else None
|
|
101
|
+
self.attn_drop = layers.Dropout(attn_drop) if attn_drop > 0.0 else None
|
|
102
|
+
|
|
103
|
+
if residual and in_feats != num_heads * out_feats:
|
|
104
|
+
self.res_fc = layers.Dense(num_heads * out_feats, use_bias=False)
|
|
105
|
+
else:
|
|
106
|
+
self.res_fc = None
|
|
107
|
+
|
|
108
|
+
total_dim = num_heads * out_feats if concat_out else out_feats
|
|
109
|
+
self.norm = _get_norm(norm, total_dim)
|
|
110
|
+
self.activation = _get_activation(activation)
|
|
111
|
+
|
|
112
|
+
def build(self, input_shape=None):
|
|
113
|
+
shape = input_shape or (None, self.in_feats)
|
|
114
|
+
in_dim = shape[-1] if shape is not None and shape[-1] is not None else self.in_feats
|
|
115
|
+
|
|
116
|
+
self.fc.build((None, in_dim))
|
|
117
|
+
if self.res_fc is not None:
|
|
118
|
+
self.res_fc.build((None, in_dim))
|
|
119
|
+
|
|
120
|
+
self.attn_l = self.add_weight(
|
|
121
|
+
shape=(1, self.num_heads, self.out_feats),
|
|
122
|
+
initializer="glorot_uniform",
|
|
123
|
+
trainable=True,
|
|
124
|
+
name="attn_l",
|
|
125
|
+
)
|
|
126
|
+
self.attn_r = self.add_weight(
|
|
127
|
+
shape=(1, self.num_heads, self.out_feats),
|
|
128
|
+
initializer="glorot_uniform",
|
|
129
|
+
trainable=True,
|
|
130
|
+
name="attn_r",
|
|
131
|
+
)
|
|
132
|
+
|
|
133
|
+
if self.use_bias:
|
|
134
|
+
self.bias = self.add_weight(
|
|
135
|
+
shape=(self.num_heads * self.out_feats,),
|
|
136
|
+
initializer="zeros",
|
|
137
|
+
trainable=True,
|
|
138
|
+
name="bias",
|
|
139
|
+
)
|
|
140
|
+
else:
|
|
141
|
+
self.bias = None
|
|
142
|
+
|
|
143
|
+
total_dim = self.num_heads * self.out_feats if self.concat_out else self.out_feats
|
|
144
|
+
if self.norm is not None:
|
|
145
|
+
self.norm.build((None, total_dim))
|
|
146
|
+
if self.activation is not None and hasattr(self.activation, "build"):
|
|
147
|
+
self.activation.build((None, total_dim))
|
|
148
|
+
|
|
149
|
+
self.built = True
|
|
150
|
+
|
|
151
|
+
def call(self, x, edge_index, training=False):
|
|
152
|
+
h = self.feat_drop(x, training=training) if self.feat_drop is not None else x
|
|
153
|
+
feat_src = ops.reshape(self.fc(h), (-1, self.num_heads, self.out_feats))
|
|
154
|
+
feat_dst = feat_src
|
|
155
|
+
|
|
156
|
+
el = ops.sum(feat_src * self.attn_l, axis=-1, keepdims=True)
|
|
157
|
+
er = ops.sum(feat_dst * self.attn_r, axis=-1, keepdims=True)
|
|
158
|
+
|
|
159
|
+
row = ops.cast(edge_index[0], "int32")
|
|
160
|
+
col = ops.cast(edge_index[1], "int32")
|
|
161
|
+
|
|
162
|
+
el_src = ops.take(el, row, axis=0)
|
|
163
|
+
er_dst = ops.take(er, col, axis=0)
|
|
164
|
+
e = ops.leaky_relu(el_src + er_dst, negative_slope=self.negative_slope)
|
|
165
|
+
|
|
166
|
+
num_nodes = ops.shape(x)[0]
|
|
167
|
+
a = softmax(e, col, num_nodes=num_nodes, dim=0)
|
|
168
|
+
if self.attn_drop is not None:
|
|
169
|
+
a = self.attn_drop(a, training=training)
|
|
170
|
+
|
|
171
|
+
msg = a * ops.take(feat_src, row, axis=0)
|
|
172
|
+
rst = segment_sum(msg, col, num_segments=num_nodes)
|
|
173
|
+
|
|
174
|
+
if self.bias is not None:
|
|
175
|
+
rst = rst + ops.reshape(self.bias, (1, self.num_heads, self.out_feats))
|
|
176
|
+
|
|
177
|
+
if self.res_fc is not None:
|
|
178
|
+
rst = rst + ops.reshape(self.res_fc(x), (num_nodes, self.num_heads, self.out_feats))
|
|
179
|
+
|
|
180
|
+
if self.concat_out:
|
|
181
|
+
rst = ops.reshape(rst, (num_nodes, self.num_heads * self.out_feats))
|
|
182
|
+
else:
|
|
183
|
+
rst = ops.mean(rst, axis=1)
|
|
184
|
+
|
|
185
|
+
if self.norm is not None:
|
|
186
|
+
rst = self.norm(rst)
|
|
187
|
+
|
|
188
|
+
if self.activation is not None:
|
|
189
|
+
rst = self.activation(rst)
|
|
190
|
+
|
|
191
|
+
return rst
|
|
192
|
+
|
|
193
|
+
|
|
194
|
+
class GraphMAE2GAT(layers.Layer):
|
|
195
|
+
r"""Multi-layer GAT encoder or decoder for GraphMAE2."""
|
|
196
|
+
|
|
197
|
+
def __init__(
|
|
198
|
+
self,
|
|
199
|
+
in_dim: int,
|
|
200
|
+
num_hidden: int,
|
|
201
|
+
out_dim: int,
|
|
202
|
+
num_layers: int,
|
|
203
|
+
nhead: int,
|
|
204
|
+
nhead_out: int,
|
|
205
|
+
activation: Optional[str] = "prelu",
|
|
206
|
+
feat_drop: float = 0.0,
|
|
207
|
+
attn_drop: float = 0.0,
|
|
208
|
+
negative_slope: float = 0.2,
|
|
209
|
+
residual: bool = True,
|
|
210
|
+
norm: Optional[str] = "layernorm",
|
|
211
|
+
concat_out: bool = True,
|
|
212
|
+
encoding: bool = True,
|
|
213
|
+
**kwargs,
|
|
214
|
+
):
|
|
215
|
+
super().__init__(**kwargs)
|
|
216
|
+
self.in_dim = in_dim
|
|
217
|
+
self.num_hidden = num_hidden
|
|
218
|
+
self.out_dim = out_dim
|
|
219
|
+
self.num_layers = num_layers
|
|
220
|
+
self.nhead = nhead
|
|
221
|
+
self.nhead_out = nhead_out
|
|
222
|
+
self.concat_out = concat_out
|
|
223
|
+
self.encoding = encoding
|
|
224
|
+
|
|
225
|
+
self.gat_layers = []
|
|
226
|
+
|
|
227
|
+
last_activation = activation if encoding else None
|
|
228
|
+
last_residual = (encoding and residual)
|
|
229
|
+
last_norm = norm if encoding else None
|
|
230
|
+
|
|
231
|
+
if num_layers == 1:
|
|
232
|
+
self.gat_layers.append(
|
|
233
|
+
GraphMAE2GATConv(
|
|
234
|
+
in_feats=in_dim,
|
|
235
|
+
out_feats=out_dim,
|
|
236
|
+
num_heads=nhead_out,
|
|
237
|
+
feat_drop=feat_drop,
|
|
238
|
+
attn_drop=attn_drop,
|
|
239
|
+
negative_slope=negative_slope,
|
|
240
|
+
residual=last_residual,
|
|
241
|
+
norm=last_norm,
|
|
242
|
+
activation=last_activation,
|
|
243
|
+
concat_out=concat_out,
|
|
244
|
+
)
|
|
245
|
+
)
|
|
246
|
+
else:
|
|
247
|
+
# Layer 0
|
|
248
|
+
self.gat_layers.append(
|
|
249
|
+
GraphMAE2GATConv(
|
|
250
|
+
in_feats=in_dim,
|
|
251
|
+
out_feats=num_hidden,
|
|
252
|
+
num_heads=nhead,
|
|
253
|
+
feat_drop=feat_drop,
|
|
254
|
+
attn_drop=attn_drop,
|
|
255
|
+
negative_slope=negative_slope,
|
|
256
|
+
residual=residual,
|
|
257
|
+
norm=norm,
|
|
258
|
+
activation=activation,
|
|
259
|
+
concat_out=concat_out,
|
|
260
|
+
)
|
|
261
|
+
)
|
|
262
|
+
# Intermediate layers
|
|
263
|
+
for _ in range(1, num_layers - 1):
|
|
264
|
+
self.gat_layers.append(
|
|
265
|
+
GraphMAE2GATConv(
|
|
266
|
+
in_feats=num_hidden * nhead,
|
|
267
|
+
out_feats=num_hidden,
|
|
268
|
+
num_heads=nhead,
|
|
269
|
+
feat_drop=feat_drop,
|
|
270
|
+
attn_drop=attn_drop,
|
|
271
|
+
negative_slope=negative_slope,
|
|
272
|
+
residual=residual,
|
|
273
|
+
norm=norm,
|
|
274
|
+
activation=activation,
|
|
275
|
+
concat_out=concat_out,
|
|
276
|
+
)
|
|
277
|
+
)
|
|
278
|
+
# Output layer
|
|
279
|
+
self.gat_layers.append(
|
|
280
|
+
GraphMAE2GATConv(
|
|
281
|
+
in_feats=num_hidden * nhead,
|
|
282
|
+
out_feats=out_dim,
|
|
283
|
+
num_heads=nhead_out,
|
|
284
|
+
feat_drop=feat_drop,
|
|
285
|
+
attn_drop=attn_drop,
|
|
286
|
+
negative_slope=negative_slope,
|
|
287
|
+
residual=last_residual,
|
|
288
|
+
norm=last_norm,
|
|
289
|
+
activation=last_activation,
|
|
290
|
+
concat_out=concat_out,
|
|
291
|
+
)
|
|
292
|
+
)
|
|
293
|
+
|
|
294
|
+
def build(self, input_shape=None):
|
|
295
|
+
for layer in self.gat_layers:
|
|
296
|
+
if hasattr(layer, "build") and not layer.built:
|
|
297
|
+
layer.build()
|
|
298
|
+
self.built = True
|
|
299
|
+
|
|
300
|
+
def call(self, x, edge_index, training=False):
|
|
301
|
+
h = x
|
|
302
|
+
for layer in self.gat_layers:
|
|
303
|
+
h = layer(h, edge_index, training=training)
|
|
304
|
+
return h
|
|
305
|
+
|
|
306
|
+
|
|
307
|
+
class GraphMAE2(layers.Layer):
|
|
308
|
+
r"""The GraphMAE2 model from `"GraphMAE2: A Decoding-Enhanced Masked
|
|
309
|
+
Self-Supervised Learning Framework for Graphs" <https://arxiv.org/abs/2304.04779>`_.
|
|
310
|
+
|
|
311
|
+
Args:
|
|
312
|
+
in_dim (int): Dimensionality of input node features.
|
|
313
|
+
num_hidden (int): Dimensionality of hidden node representations.
|
|
314
|
+
num_layers (int, optional): Number of encoder layers. (default: ``4``)
|
|
315
|
+
num_dec_layers (int, optional): Number of decoder layers. (default: ``1``)
|
|
316
|
+
num_remasking (int, optional): Number of remasking views in decoder. (default: ``3``)
|
|
317
|
+
nhead (int, optional): Number of attention heads in encoder. (default: ``8``)
|
|
318
|
+
nhead_out (int, optional): Number of attention heads in decoder output. (default: ``1``)
|
|
319
|
+
activation (str, optional): Activation function. (default: ``"prelu"``)
|
|
320
|
+
feat_drop (float, optional): Node feature dropout rate. (default: ``0.2``)
|
|
321
|
+
attn_drop (float, optional): Attention dropout rate. (default: ``0.1``)
|
|
322
|
+
negative_slope (float, optional): LeakyReLU negative slope. (default: ``0.2``)
|
|
323
|
+
residual (bool, optional): Whether to use residual connections. (default: ``True``)
|
|
324
|
+
norm (str, optional): Normalization type (``"layernorm"``, ``"batchnorm"``, or ``None``). (default: ``"layernorm"``)
|
|
325
|
+
mask_rate (float, optional): Fraction of input nodes to mask. (default: ``0.5``)
|
|
326
|
+
remask_rate (float, optional): Fraction of latent nodes to remask. (default: ``0.5``)
|
|
327
|
+
remask_method (str, optional): Remasking method (``"random"`` or ``"fixed"``). (default: ``"random"``)
|
|
328
|
+
loss_fn (str, optional): Reconstruction loss type (``"sce"`` or ``"mse"``). (default: ``"sce"``)
|
|
329
|
+
alpha_l (float, optional): Power exponent in Scaled Cosine Error loss. (default: ``2.0``)
|
|
330
|
+
lam (float, optional): Weight of the latent prediction loss term. (default: ``1.0``)
|
|
331
|
+
momentum (float, optional): Teacher EMA update momentum. (default: ``0.996``)
|
|
332
|
+
delayed_ema_epoch (int, optional): Epoch to begin EMA teacher updates. (default: ``0``)
|
|
333
|
+
|
|
334
|
+
Example:
|
|
335
|
+
```python
|
|
336
|
+
import numpy as np
|
|
337
|
+
from k3_node.models import GraphMAE2
|
|
338
|
+
|
|
339
|
+
x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
|
|
340
|
+
edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
|
|
341
|
+
|
|
342
|
+
model = GraphMAE2(in_dim=8, num_hidden=32, num_layers=2, num_dec_layers=1, nhead=4, nhead_out=1)
|
|
343
|
+
print(tuple(model(x, edge_index).shape)) # (10, 32): node embeddings
|
|
344
|
+
loss = model.loss(x, edge_index) # masked feature reconstruction loss for pre-training
|
|
345
|
+
print(tuple(loss.shape)) # ()
|
|
346
|
+
```
|
|
347
|
+
"""
|
|
348
|
+
|
|
349
|
+
def __init__(
|
|
350
|
+
self,
|
|
351
|
+
in_dim: int,
|
|
352
|
+
num_hidden: int,
|
|
353
|
+
num_layers: int = 4,
|
|
354
|
+
num_dec_layers: int = 1,
|
|
355
|
+
num_remasking: int = 3,
|
|
356
|
+
nhead: int = 8,
|
|
357
|
+
nhead_out: int = 1,
|
|
358
|
+
activation: str = "prelu",
|
|
359
|
+
feat_drop: float = 0.2,
|
|
360
|
+
attn_drop: float = 0.1,
|
|
361
|
+
negative_slope: float = 0.2,
|
|
362
|
+
residual: bool = True,
|
|
363
|
+
norm: Optional[str] = "layernorm",
|
|
364
|
+
mask_rate: float = 0.5,
|
|
365
|
+
remask_rate: float = 0.5,
|
|
366
|
+
remask_method: str = "random",
|
|
367
|
+
loss_fn: str = "sce",
|
|
368
|
+
alpha_l: float = 2.0,
|
|
369
|
+
lam: float = 1.0,
|
|
370
|
+
momentum: float = 0.996,
|
|
371
|
+
delayed_ema_epoch: int = 0,
|
|
372
|
+
**kwargs,
|
|
373
|
+
):
|
|
374
|
+
super().__init__(**kwargs)
|
|
375
|
+
self.in_dim = in_dim
|
|
376
|
+
self.num_hidden = num_hidden
|
|
377
|
+
self.num_layers = num_layers
|
|
378
|
+
self.num_dec_layers = num_dec_layers
|
|
379
|
+
self.num_remasking = num_remasking
|
|
380
|
+
self.nhead = nhead
|
|
381
|
+
self.nhead_out = nhead_out
|
|
382
|
+
self.mask_rate = mask_rate
|
|
383
|
+
self.remask_rate = remask_rate
|
|
384
|
+
self.remask_method = remask_method
|
|
385
|
+
self.loss_fn = loss_fn
|
|
386
|
+
self.alpha_l = alpha_l
|
|
387
|
+
self.lam = lam
|
|
388
|
+
self.momentum = momentum
|
|
389
|
+
self.delayed_ema_epoch = delayed_ema_epoch
|
|
390
|
+
|
|
391
|
+
assert num_hidden % nhead == 0, f"num_hidden ({num_hidden}) must be divisible by nhead ({nhead})"
|
|
392
|
+
assert num_hidden % nhead_out == 0, f"num_hidden ({num_hidden}) must be divisible by nhead_out ({nhead_out})"
|
|
393
|
+
|
|
394
|
+
enc_num_hidden = num_hidden // nhead
|
|
395
|
+
dec_in_dim = num_hidden
|
|
396
|
+
dec_num_hidden = num_hidden // nhead
|
|
397
|
+
|
|
398
|
+
# 1. Student Encoder
|
|
399
|
+
self.encoder = GraphMAE2GAT(
|
|
400
|
+
in_dim=in_dim,
|
|
401
|
+
num_hidden=enc_num_hidden,
|
|
402
|
+
out_dim=enc_num_hidden,
|
|
403
|
+
num_layers=num_layers,
|
|
404
|
+
nhead=nhead,
|
|
405
|
+
nhead_out=nhead,
|
|
406
|
+
activation=activation,
|
|
407
|
+
feat_drop=feat_drop,
|
|
408
|
+
attn_drop=attn_drop,
|
|
409
|
+
negative_slope=negative_slope,
|
|
410
|
+
residual=residual,
|
|
411
|
+
norm=norm,
|
|
412
|
+
concat_out=True,
|
|
413
|
+
encoding=True,
|
|
414
|
+
)
|
|
415
|
+
|
|
416
|
+
# 2. Decoder
|
|
417
|
+
self.decoder = GraphMAE2GAT(
|
|
418
|
+
in_dim=dec_in_dim,
|
|
419
|
+
num_hidden=dec_num_hidden,
|
|
420
|
+
out_dim=in_dim,
|
|
421
|
+
num_layers=num_dec_layers,
|
|
422
|
+
nhead=nhead,
|
|
423
|
+
nhead_out=nhead_out,
|
|
424
|
+
activation=activation,
|
|
425
|
+
feat_drop=feat_drop,
|
|
426
|
+
attn_drop=attn_drop,
|
|
427
|
+
negative_slope=negative_slope,
|
|
428
|
+
residual=residual,
|
|
429
|
+
norm=norm,
|
|
430
|
+
concat_out=True,
|
|
431
|
+
encoding=False,
|
|
432
|
+
)
|
|
433
|
+
|
|
434
|
+
self.encoder_to_decoder = layers.Dense(dec_in_dim, use_bias=False)
|
|
435
|
+
|
|
436
|
+
# 3. Projector & Predictor
|
|
437
|
+
self.projector = keras.Sequential([
|
|
438
|
+
layers.Dense(256),
|
|
439
|
+
layers.PReLU(shared_axes=[1]),
|
|
440
|
+
layers.Dense(num_hidden),
|
|
441
|
+
])
|
|
442
|
+
|
|
443
|
+
self.predictor = keras.Sequential([
|
|
444
|
+
layers.PReLU(shared_axes=[1]),
|
|
445
|
+
layers.Dense(num_hidden),
|
|
446
|
+
])
|
|
447
|
+
|
|
448
|
+
# 4. Teacher EMA networks
|
|
449
|
+
self.encoder_ema = GraphMAE2GAT(
|
|
450
|
+
in_dim=in_dim,
|
|
451
|
+
num_hidden=enc_num_hidden,
|
|
452
|
+
out_dim=enc_num_hidden,
|
|
453
|
+
num_layers=num_layers,
|
|
454
|
+
nhead=nhead,
|
|
455
|
+
nhead_out=nhead,
|
|
456
|
+
activation=activation,
|
|
457
|
+
feat_drop=feat_drop,
|
|
458
|
+
attn_drop=attn_drop,
|
|
459
|
+
negative_slope=negative_slope,
|
|
460
|
+
residual=residual,
|
|
461
|
+
norm=norm,
|
|
462
|
+
concat_out=True,
|
|
463
|
+
encoding=True,
|
|
464
|
+
trainable=False,
|
|
465
|
+
)
|
|
466
|
+
|
|
467
|
+
self.projector_ema = keras.Sequential([
|
|
468
|
+
layers.Dense(256, trainable=False),
|
|
469
|
+
layers.PReLU(shared_axes=[1], trainable=False),
|
|
470
|
+
layers.Dense(num_hidden, trainable=False),
|
|
471
|
+
], trainable=False)
|
|
472
|
+
|
|
473
|
+
self.enc_mask_token = self.add_weight(
|
|
474
|
+
shape=(1, self.in_dim),
|
|
475
|
+
initializer="glorot_normal",
|
|
476
|
+
trainable=True,
|
|
477
|
+
name="enc_mask_token",
|
|
478
|
+
)
|
|
479
|
+
self.dec_mask_token = self.add_weight(
|
|
480
|
+
shape=(1, self.num_hidden),
|
|
481
|
+
initializer="glorot_normal",
|
|
482
|
+
trainable=True,
|
|
483
|
+
name="dec_mask_token",
|
|
484
|
+
)
|
|
485
|
+
|
|
486
|
+
def build(self, input_shape=None):
|
|
487
|
+
if not hasattr(self, "enc_mask_token") or self.enc_mask_token is None:
|
|
488
|
+
self.enc_mask_token = self.add_weight(
|
|
489
|
+
shape=(1, self.in_dim),
|
|
490
|
+
initializer="glorot_normal",
|
|
491
|
+
trainable=True,
|
|
492
|
+
name="enc_mask_token",
|
|
493
|
+
)
|
|
494
|
+
if not hasattr(self, "dec_mask_token") or self.dec_mask_token is None:
|
|
495
|
+
self.dec_mask_token = self.add_weight(
|
|
496
|
+
shape=(1, self.num_hidden),
|
|
497
|
+
initializer="glorot_normal",
|
|
498
|
+
trainable=True,
|
|
499
|
+
name="dec_mask_token",
|
|
500
|
+
)
|
|
501
|
+
|
|
502
|
+
self.encoder.build((None, self.in_dim))
|
|
503
|
+
self.encoder_to_decoder.build((None, self.num_hidden))
|
|
504
|
+
self.decoder.build((None, self.num_hidden))
|
|
505
|
+
self.projector.build((None, self.num_hidden))
|
|
506
|
+
self.predictor.build((None, self.num_hidden))
|
|
507
|
+
self.encoder_ema.build((None, self.in_dim))
|
|
508
|
+
self.projector_ema.build((None, self.num_hidden))
|
|
509
|
+
|
|
510
|
+
# Copy initial weights from student to teacher
|
|
511
|
+
for p_s, p_t in zip(self.encoder.weights, self.encoder_ema.weights):
|
|
512
|
+
p_t.assign(p_s)
|
|
513
|
+
for p_s, p_t in zip(self.projector.weights, self.projector_ema.weights):
|
|
514
|
+
p_t.assign(p_s)
|
|
515
|
+
|
|
516
|
+
self.built = True
|
|
517
|
+
|
|
518
|
+
def embed(self, x, edge_index, training=None):
|
|
519
|
+
r"""Generates node embeddings with the encoder."""
|
|
520
|
+
if not self.built:
|
|
521
|
+
self.build((None, self.in_dim))
|
|
522
|
+
return self.encoder(x, edge_index, training=training)
|
|
523
|
+
|
|
524
|
+
def encoding_mask_noise(self, x, mask_rate: Optional[float] = None, mask_nodes=None):
|
|
525
|
+
r"""Masks node features for encoder input."""
|
|
526
|
+
x = ops.convert_to_tensor(x) # NumPy inputs cannot be mixed with backend tensors
|
|
527
|
+
rate = self.mask_rate if mask_rate is None else mask_rate
|
|
528
|
+
num_nodes = ops.shape(x)[0]
|
|
529
|
+
|
|
530
|
+
if mask_nodes is None:
|
|
531
|
+
perm = np.random.permutation(num_nodes)
|
|
532
|
+
num_mask_nodes = int(rate * num_nodes)
|
|
533
|
+
mask_nodes = ops.convert_to_tensor(perm[:num_mask_nodes], dtype="int32")
|
|
534
|
+
keep_nodes = ops.convert_to_tensor(perm[num_mask_nodes:], dtype="int32")
|
|
535
|
+
else:
|
|
536
|
+
mask_nodes = ops.cast(mask_nodes, "int32")
|
|
537
|
+
all_mask = np.zeros(num_nodes, dtype=bool)
|
|
538
|
+
all_mask[ops.convert_to_numpy(mask_nodes)] = True
|
|
539
|
+
keep_nodes = ops.convert_to_tensor(np.where(~all_mask)[0], dtype="int32")
|
|
540
|
+
|
|
541
|
+
# Replace masked nodes with enc_mask_token
|
|
542
|
+
# Create a zeroed masked version
|
|
543
|
+
mask_vector = np.zeros(num_nodes, dtype=np.float32)
|
|
544
|
+
mask_vector[ops.convert_to_numpy(mask_nodes)] = 1.0
|
|
545
|
+
mask_tensor = ops.expand_dims(ops.convert_to_tensor(mask_vector, dtype=x.dtype), -1)
|
|
546
|
+
|
|
547
|
+
masked_x = x * (1.0 - mask_tensor) + mask_tensor * self.enc_mask_token
|
|
548
|
+
return masked_x, mask_nodes, keep_nodes
|
|
549
|
+
|
|
550
|
+
def random_remask(self, rep, remask_rate: Optional[float] = None, remask_nodes=None):
|
|
551
|
+
r"""Remasks latent representation for decoder input."""
|
|
552
|
+
rate = self.remask_rate if remask_rate is None else remask_rate
|
|
553
|
+
num_nodes = ops.shape(rep)[0]
|
|
554
|
+
|
|
555
|
+
if remask_nodes is None:
|
|
556
|
+
perm = np.random.permutation(num_nodes)
|
|
557
|
+
num_remask_nodes = int(rate * num_nodes)
|
|
558
|
+
remask_nodes = ops.convert_to_tensor(perm[:num_remask_nodes], dtype="int32")
|
|
559
|
+
rekeep_nodes = ops.convert_to_tensor(perm[num_remask_nodes:], dtype="int32")
|
|
560
|
+
else:
|
|
561
|
+
remask_nodes = ops.cast(remask_nodes, "int32")
|
|
562
|
+
all_mask = np.zeros(num_nodes, dtype=bool)
|
|
563
|
+
all_mask[ops.convert_to_numpy(remask_nodes)] = True
|
|
564
|
+
rekeep_nodes = ops.convert_to_tensor(np.where(~all_mask)[0], dtype="int32")
|
|
565
|
+
|
|
566
|
+
remask_vector = np.zeros(num_nodes, dtype=np.float32)
|
|
567
|
+
remask_vector[ops.convert_to_numpy(remask_nodes)] = 1.0
|
|
568
|
+
remask_tensor = ops.expand_dims(ops.convert_to_tensor(remask_vector, dtype=rep.dtype), -1)
|
|
569
|
+
|
|
570
|
+
remasked_rep = rep * (1.0 - remask_tensor) + remask_tensor * self.dec_mask_token
|
|
571
|
+
return remasked_rep, remask_nodes, rekeep_nodes
|
|
572
|
+
|
|
573
|
+
def ema_update(self, momentum: Optional[float] = None):
|
|
574
|
+
r"""Updates teacher EMA parameters."""
|
|
575
|
+
m = self.momentum if momentum is None else momentum
|
|
576
|
+
for p_s, p_t in zip(self.encoder.weights, self.encoder_ema.weights):
|
|
577
|
+
p_t.assign(p_t * m + p_s * (1.0 - m))
|
|
578
|
+
for p_s, p_t in zip(self.projector.weights, self.projector_ema.weights):
|
|
579
|
+
p_t.assign(p_t * m + p_s * (1.0 - m))
|
|
580
|
+
|
|
581
|
+
def loss(
|
|
582
|
+
self,
|
|
583
|
+
x,
|
|
584
|
+
edge_index,
|
|
585
|
+
mask_nodes=None,
|
|
586
|
+
targets=None,
|
|
587
|
+
epoch: int = 0,
|
|
588
|
+
training: bool = True,
|
|
589
|
+
):
|
|
590
|
+
r"""Computes GraphMAE2 loss: attribute reconstruction loss + latent prediction loss."""
|
|
591
|
+
if not self.built:
|
|
592
|
+
self.build((None, self.in_dim))
|
|
593
|
+
|
|
594
|
+
# 1. Masking
|
|
595
|
+
masked_x, mask_nodes, keep_nodes = self.encoding_mask_noise(x, mask_nodes=mask_nodes)
|
|
596
|
+
|
|
597
|
+
# 2. Student encoder
|
|
598
|
+
enc_rep = self.encoder(masked_x, edge_index, training=training)
|
|
599
|
+
|
|
600
|
+
# 3. Teacher EMA target (no gradient)
|
|
601
|
+
teacher_rep = ops.stop_gradient(self.encoder_ema(x, edge_index, training=False))
|
|
602
|
+
if targets is not None:
|
|
603
|
+
latent_target = ops.stop_gradient(self.projector_ema(ops.take(teacher_rep, targets, axis=0)))
|
|
604
|
+
latent_pred = self.predictor(self.projector(ops.take(enc_rep, targets, axis=0)))
|
|
605
|
+
else:
|
|
606
|
+
latent_target = ops.stop_gradient(self.projector_ema(ops.take(teacher_rep, keep_nodes, axis=0)))
|
|
607
|
+
latent_pred = self.predictor(self.projector(ops.take(enc_rep, keep_nodes, axis=0)))
|
|
608
|
+
|
|
609
|
+
loss_latent = sce_loss(latent_pred, latent_target, alpha=1.0)
|
|
610
|
+
|
|
611
|
+
# 4. Decoder attribute reconstruction
|
|
612
|
+
origin_rep = self.encoder_to_decoder(enc_rep)
|
|
613
|
+
|
|
614
|
+
criterion = sce_loss if self.loss_fn == "sce" else (lambda pred, tgt: ops.mean(ops.power(pred - tgt, 2)))
|
|
615
|
+
|
|
616
|
+
loss_rec_all = 0.0
|
|
617
|
+
if self.remask_method == "random":
|
|
618
|
+
for _ in range(self.num_remasking):
|
|
619
|
+
rep, _, _ = self.random_remask(origin_rep)
|
|
620
|
+
recon = self.decoder(rep, edge_index, training=training)
|
|
621
|
+
x_init = ops.take(x, mask_nodes, axis=0)
|
|
622
|
+
x_rec = ops.take(recon, mask_nodes, axis=0)
|
|
623
|
+
loss_rec_all = loss_rec_all + criterion(x_rec, x_init, alpha=self.alpha_l) if self.loss_fn == "sce" else loss_rec_all + criterion(x_rec, x_init)
|
|
624
|
+
loss_rec = loss_rec_all / float(self.num_remasking)
|
|
625
|
+
else:
|
|
626
|
+
# Fixed remasking
|
|
627
|
+
mask_vector = np.zeros(ops.shape(x)[0], dtype=np.float32)
|
|
628
|
+
mask_vector[ops.convert_to_numpy(mask_nodes)] = 1.0
|
|
629
|
+
mask_tensor = ops.expand_dims(ops.convert_to_tensor(mask_vector, dtype=origin_rep.dtype), -1)
|
|
630
|
+
rep = origin_rep * (1.0 - mask_tensor)
|
|
631
|
+
recon = self.decoder(rep, edge_index, training=training)
|
|
632
|
+
x_init = ops.take(x, mask_nodes, axis=0)
|
|
633
|
+
x_rec = ops.take(recon, mask_nodes, axis=0)
|
|
634
|
+
loss_rec = criterion(x_rec, x_init, alpha=self.alpha_l) if self.loss_fn == "sce" else criterion(x_rec, x_init)
|
|
635
|
+
|
|
636
|
+
total_loss = loss_rec + self.lam * loss_latent
|
|
637
|
+
|
|
638
|
+
if epoch >= self.delayed_ema_epoch and training:
|
|
639
|
+
self.ema_update()
|
|
640
|
+
|
|
641
|
+
return total_loss
|
|
642
|
+
|
|
643
|
+
def call(self, x, edge_index, training: bool = False):
|
|
644
|
+
r"""Forward pass: returns node embeddings by default."""
|
|
645
|
+
return self.embed(x, edge_index, training=training)
|
|
646
|
+
|
|
647
|
+
def load_weights_from_checkpoint(
|
|
648
|
+
self,
|
|
649
|
+
checkpoint_path: Optional[str] = None,
|
|
650
|
+
dataset: Optional[str] = None,
|
|
651
|
+
folder: str = "checkpoints",
|
|
652
|
+
download: bool = True,
|
|
653
|
+
):
|
|
654
|
+
r"""Loads weights from a PyTorch state dict checkpoint or Google Drive."""
|
|
655
|
+
return load_graphmae2_weights(
|
|
656
|
+
self,
|
|
657
|
+
checkpoint_path=checkpoint_path,
|
|
658
|
+
dataset=dataset,
|
|
659
|
+
folder=folder,
|
|
660
|
+
download=download,
|
|
661
|
+
)
|
|
662
|
+
|
|
663
|
+
@classmethod
|
|
664
|
+
def from_pretrained(
|
|
665
|
+
cls,
|
|
666
|
+
dataset: str = "ogbn-arxiv",
|
|
667
|
+
folder: str = "checkpoints",
|
|
668
|
+
download: bool = True,
|
|
669
|
+
**kwargs,
|
|
670
|
+
) -> "GraphMAE2":
|
|
671
|
+
r"""Instantiates a GraphMAE2 model with pre-trained weights downloaded from Google Drive.
|
|
672
|
+
|
|
673
|
+
Args:
|
|
674
|
+
dataset (str): Dataset name (``"ogbn-arxiv"``, ``"ogbn-products"``,
|
|
675
|
+
``"mag-scholar-f"``, or ``"ogbn-papers100M"``).
|
|
676
|
+
folder (str, optional): Directory to store/find checkpoints. (default: ``"checkpoints"``)
|
|
677
|
+
download (bool, optional): Whether to download checkpoint if missing locally. (default: ``True``)
|
|
678
|
+
**kwargs: Overrides for model hyperparameters.
|
|
679
|
+
|
|
680
|
+
Returns:
|
|
681
|
+
GraphMAE2: Model instance loaded with pre-trained weights.
|
|
682
|
+
"""
|
|
683
|
+
key = _canonical_dataset_name(dataset)
|
|
684
|
+
if key not in GRAPHMAE2_PRETRAINED:
|
|
685
|
+
raise ValueError(
|
|
686
|
+
f"Unknown dataset '{dataset}'. Available pre-trained models: {list(GRAPHMAE2_PRETRAINED.keys())}"
|
|
687
|
+
)
|
|
688
|
+
cfg = dict(GRAPHMAE2_PRETRAINED[key])
|
|
689
|
+
cfg.pop("id")
|
|
690
|
+
cfg.pop("filename")
|
|
691
|
+
cfg.update(kwargs)
|
|
692
|
+
|
|
693
|
+
model = cls(**cfg)
|
|
694
|
+
load_graphmae2_weights(model, dataset=key, folder=folder, download=download)
|
|
695
|
+
return model
|
|
696
|
+
|
|
697
|
+
def __repr__(self) -> str:
|
|
698
|
+
return (
|
|
699
|
+
f"{self.__class__.__name__}(in_dim={self.in_dim}, "
|
|
700
|
+
f"num_hidden={self.num_hidden}, num_layers={self.num_layers}, "
|
|
701
|
+
f"num_dec_layers={self.num_dec_layers}, nhead={self.nhead})"
|
|
702
|
+
)
|
|
703
|
+
|
|
704
|
+
|
|
705
|
+
GRAPHMAE2_PRETRAINED = {
|
|
706
|
+
"ogbn-arxiv": {
|
|
707
|
+
"id": "1KdU5TbAg0lQwruO7SKenoFiC2MbaZQRr",
|
|
708
|
+
"filename": "gat_gat_1024_4_ogbn-arxiv_0.5_1024_checkpoint.pt",
|
|
709
|
+
"in_dim": 128,
|
|
710
|
+
"num_hidden": 1024,
|
|
711
|
+
"num_layers": 4,
|
|
712
|
+
"num_dec_layers": 1,
|
|
713
|
+
"nhead": 8,
|
|
714
|
+
"nhead_out": 1,
|
|
715
|
+
"activation": "prelu",
|
|
716
|
+
"norm": "layernorm",
|
|
717
|
+
"residual": True,
|
|
718
|
+
},
|
|
719
|
+
"ogbn-products": {
|
|
720
|
+
"id": "1Qk3bgK8H3bee3qmagH_hRJOfQmUD8PCW",
|
|
721
|
+
"filename": "gat_gat_1024_4_ogbn-products_0.5_1024_checkpoint.pt",
|
|
722
|
+
"in_dim": 100,
|
|
723
|
+
"num_hidden": 1024,
|
|
724
|
+
"num_layers": 4,
|
|
725
|
+
"num_dec_layers": 1,
|
|
726
|
+
"nhead": 4,
|
|
727
|
+
"nhead_out": 1,
|
|
728
|
+
"activation": "prelu",
|
|
729
|
+
"norm": "layernorm",
|
|
730
|
+
"residual": True,
|
|
731
|
+
},
|
|
732
|
+
"mag-scholar-f": {
|
|
733
|
+
"id": "1KpQk_OKbbo4qTLQYZ84pAJDy1sh4oZv2",
|
|
734
|
+
"filename": "gat_gat_1024_4_mag-scholar-f_0.5_1024_checkpoint.pt",
|
|
735
|
+
"in_dim": 128,
|
|
736
|
+
"num_hidden": 1024,
|
|
737
|
+
"num_layers": 4,
|
|
738
|
+
"num_dec_layers": 1,
|
|
739
|
+
"nhead": 8,
|
|
740
|
+
"nhead_out": 1,
|
|
741
|
+
"activation": "prelu",
|
|
742
|
+
"norm": "layernorm",
|
|
743
|
+
"residual": True,
|
|
744
|
+
},
|
|
745
|
+
"ogbn-papers100M": {
|
|
746
|
+
"id": "1zCD_vOckLfOXD1dWRY025A30QeuHsA_0",
|
|
747
|
+
"filename": "gat_gat_1024_4_ogbn-papers100M_0.5_1024_checkpoint.pt",
|
|
748
|
+
"in_dim": 128,
|
|
749
|
+
"num_hidden": 1024,
|
|
750
|
+
"num_layers": 4,
|
|
751
|
+
"num_dec_layers": 1,
|
|
752
|
+
"nhead": 8,
|
|
753
|
+
"nhead_out": 1,
|
|
754
|
+
"activation": "prelu",
|
|
755
|
+
"norm": "layernorm",
|
|
756
|
+
"residual": True,
|
|
757
|
+
},
|
|
758
|
+
}
|
|
759
|
+
|
|
760
|
+
DATASET_ALIASES = {
|
|
761
|
+
"arxiv": "ogbn-arxiv",
|
|
762
|
+
"products": "ogbn-products",
|
|
763
|
+
"mag": "mag-scholar-f",
|
|
764
|
+
"mag-scholar": "mag-scholar-f",
|
|
765
|
+
"papers100m": "ogbn-papers100M",
|
|
766
|
+
"ogbn-papers100m": "ogbn-papers100M",
|
|
767
|
+
"papers": "ogbn-papers100M",
|
|
768
|
+
}
|
|
769
|
+
|
|
770
|
+
|
|
771
|
+
def _canonical_dataset_name(name: Optional[str]) -> str:
|
|
772
|
+
if not name:
|
|
773
|
+
return ""
|
|
774
|
+
name_clean = name.strip()
|
|
775
|
+
if name_clean in GRAPHMAE2_PRETRAINED:
|
|
776
|
+
return name_clean
|
|
777
|
+
name_lower = name_clean.lower()
|
|
778
|
+
if name_lower in DATASET_ALIASES:
|
|
779
|
+
return DATASET_ALIASES[name_lower]
|
|
780
|
+
for k in GRAPHMAE2_PRETRAINED:
|
|
781
|
+
if k.lower() == name_lower:
|
|
782
|
+
return k
|
|
783
|
+
for k, v in GRAPHMAE2_PRETRAINED.items():
|
|
784
|
+
if v["filename"] == name_clean:
|
|
785
|
+
return k
|
|
786
|
+
return name_clean
|
|
787
|
+
|
|
788
|
+
|
|
789
|
+
def download_graphmae2_checkpoint(
|
|
790
|
+
dataset: str,
|
|
791
|
+
folder: str = "checkpoints",
|
|
792
|
+
log: bool = True,
|
|
793
|
+
) -> str:
|
|
794
|
+
r"""Downloads a pre-trained GraphMAE2 checkpoint from Google Drive using download_google_url.
|
|
795
|
+
|
|
796
|
+
Google Drive folder: https://drive.google.com/drive/folders/1GiuP0PtIZaYlJWIrjvu73ZQCJGr6kGkh
|
|
797
|
+
|
|
798
|
+
Args:
|
|
799
|
+
dataset (str): Name of dataset (``"ogbn-arxiv"``, ``"ogbn-products"``,
|
|
800
|
+
``"mag-scholar-f"``, or ``"ogbn-papers100M"``).
|
|
801
|
+
folder (str, optional): Target directory to save the checkpoint. (default: ``"checkpoints"``)
|
|
802
|
+
log (bool, optional): Whether to print download progress. (default: ``True``)
|
|
803
|
+
|
|
804
|
+
Returns:
|
|
805
|
+
str: Absolute path to the downloaded checkpoint file.
|
|
806
|
+
"""
|
|
807
|
+
key = _canonical_dataset_name(dataset)
|
|
808
|
+
if key not in GRAPHMAE2_PRETRAINED:
|
|
809
|
+
raise ValueError(
|
|
810
|
+
f"Unknown dataset '{dataset}'. Available pre-trained checkpoints: {list(GRAPHMAE2_PRETRAINED.keys())}"
|
|
811
|
+
)
|
|
812
|
+
info = GRAPHMAE2_PRETRAINED[key]
|
|
813
|
+
|
|
814
|
+
target = os.path.join(folder, info["filename"])
|
|
815
|
+
if os.path.exists(target):
|
|
816
|
+
return target
|
|
817
|
+
|
|
818
|
+
# Check if local file exists in GraphMAE2-main/GraphMAE2_checkpoints when using default folder
|
|
819
|
+
if folder == "checkpoints":
|
|
820
|
+
local_alt = os.path.join("GraphMAE2-main", "GraphMAE2_checkpoints", info["filename"])
|
|
821
|
+
if os.path.exists(local_alt):
|
|
822
|
+
return local_alt
|
|
823
|
+
|
|
824
|
+
return download_google_url(
|
|
825
|
+
id=info["id"],
|
|
826
|
+
folder=folder,
|
|
827
|
+
filename=info["filename"],
|
|
828
|
+
log=log,
|
|
829
|
+
)
|
|
830
|
+
|
|
831
|
+
|
|
832
|
+
def load_graphmae2_weights(
|
|
833
|
+
model: GraphMAE2,
|
|
834
|
+
checkpoint_path: Optional[str] = None,
|
|
835
|
+
dataset: Optional[str] = None,
|
|
836
|
+
folder: str = "checkpoints",
|
|
837
|
+
download: bool = True,
|
|
838
|
+
):
|
|
839
|
+
r"""Loads pre-trained weights from a GraphMAE2 PyTorch checkpoint (.pt).
|
|
840
|
+
|
|
841
|
+
If the checkpoint does not exist locally and download=True, it will be automatically
|
|
842
|
+
downloaded from the official Google Drive folder using download_google_url.
|
|
843
|
+
|
|
844
|
+
Args:
|
|
845
|
+
model (GraphMAE2): The target GraphMAE2 model instance.
|
|
846
|
+
checkpoint_path (str, optional): Local path to .pt file or dataset name.
|
|
847
|
+
dataset (str, optional): Dataset name if downloading from Google Drive.
|
|
848
|
+
folder (str, optional): Directory to store downloaded checkpoints. (default: ``"checkpoints"``)
|
|
849
|
+
download (bool, optional): Whether to download checkpoint if missing locally. (default: ``True``)
|
|
850
|
+
"""
|
|
851
|
+
if checkpoint_path is None and dataset is None:
|
|
852
|
+
raise ValueError("Either checkpoint_path or dataset must be specified.")
|
|
853
|
+
|
|
854
|
+
path_to_load = checkpoint_path
|
|
855
|
+
|
|
856
|
+
candidate_dataset = dataset or (_canonical_dataset_name(checkpoint_path) if checkpoint_path else None)
|
|
857
|
+
if candidate_dataset in GRAPHMAE2_PRETRAINED:
|
|
858
|
+
if checkpoint_path and os.path.isfile(checkpoint_path):
|
|
859
|
+
path_to_load = checkpoint_path
|
|
860
|
+
else:
|
|
861
|
+
filename = GRAPHMAE2_PRETRAINED[candidate_dataset]["filename"]
|
|
862
|
+
alt_local = os.path.join("GraphMAE2-main", "GraphMAE2_checkpoints", filename)
|
|
863
|
+
default_local = os.path.join(folder, filename)
|
|
864
|
+
if os.path.isfile(alt_local):
|
|
865
|
+
path_to_load = alt_local
|
|
866
|
+
elif os.path.isfile(default_local):
|
|
867
|
+
path_to_load = default_local
|
|
868
|
+
elif download:
|
|
869
|
+
path_to_load = download_graphmae2_checkpoint(candidate_dataset, folder=folder)
|
|
870
|
+
else:
|
|
871
|
+
raise FileNotFoundError(f"Checkpoint for '{candidate_dataset}' not found at '{checkpoint_path}'.")
|
|
872
|
+
elif checkpoint_path and not os.path.isfile(checkpoint_path):
|
|
873
|
+
if download and dataset:
|
|
874
|
+
path_to_load = download_graphmae2_checkpoint(dataset, folder=folder)
|
|
875
|
+
else:
|
|
876
|
+
raise FileNotFoundError(f"Checkpoint file '{checkpoint_path}' not found.")
|
|
877
|
+
|
|
878
|
+
import torch
|
|
879
|
+
|
|
880
|
+
state_dict = torch.load(path_to_load, map_location="cpu")
|
|
881
|
+
if not isinstance(state_dict, dict):
|
|
882
|
+
raise ValueError(f"Expected a dict/state_dict in checkpoint, got {type(state_dict)}")
|
|
883
|
+
|
|
884
|
+
if not model.built:
|
|
885
|
+
model.build((None, model.in_dim))
|
|
886
|
+
|
|
887
|
+
def _to_tensor(t):
|
|
888
|
+
if hasattr(t, "detach"):
|
|
889
|
+
t = t.detach()
|
|
890
|
+
if hasattr(t, "numpy"):
|
|
891
|
+
t = t.numpy()
|
|
892
|
+
return ops.convert_to_tensor(np.array(t, dtype=np.float32), dtype="float32")
|
|
893
|
+
|
|
894
|
+
# 1. Mask tokens
|
|
895
|
+
if "enc_mask_token" in state_dict:
|
|
896
|
+
model.enc_mask_token.assign(_to_tensor(state_dict["enc_mask_token"]))
|
|
897
|
+
if "dec_mask_token" in state_dict:
|
|
898
|
+
model.dec_mask_token.assign(_to_tensor(state_dict["dec_mask_token"]))
|
|
899
|
+
|
|
900
|
+
# Helper for GAT module
|
|
901
|
+
def _load_gat(gat_module, prefix):
|
|
902
|
+
for i, layer in enumerate(gat_module.gat_layers):
|
|
903
|
+
p = f"{prefix}.gat_layers.{i}"
|
|
904
|
+
if f"{p}.fc.weight" in state_dict:
|
|
905
|
+
layer.fc.kernel.assign(_to_tensor(state_dict[f"{p}.fc.weight"].t()))
|
|
906
|
+
if f"{p}.attn_l" in state_dict:
|
|
907
|
+
layer.attn_l.assign(_to_tensor(state_dict[f"{p}.attn_l"]))
|
|
908
|
+
if f"{p}.attn_r" in state_dict:
|
|
909
|
+
layer.attn_r.assign(_to_tensor(state_dict[f"{p}.attn_r"]))
|
|
910
|
+
if f"{p}.bias" in state_dict and layer.bias is not None:
|
|
911
|
+
layer.bias.assign(_to_tensor(state_dict[f"{p}.bias"]))
|
|
912
|
+
if f"{p}.res_fc.weight" in state_dict and layer.res_fc is not None:
|
|
913
|
+
layer.res_fc.kernel.assign(_to_tensor(state_dict[f"{p}.res_fc.weight"].t()))
|
|
914
|
+
if f"{p}.activation.weight" in state_dict and hasattr(layer.activation, "alpha"):
|
|
915
|
+
layer.activation.alpha.assign(_to_tensor(state_dict[f"{p}.activation.weight"]))
|
|
916
|
+
if layer.norm is not None:
|
|
917
|
+
if f"{p}.norm.weight" in state_dict and hasattr(layer.norm, "gamma"):
|
|
918
|
+
layer.norm.gamma.assign(_to_tensor(state_dict[f"{p}.norm.weight"]))
|
|
919
|
+
if f"{p}.norm.bias" in state_dict and hasattr(layer.norm, "beta"):
|
|
920
|
+
layer.norm.beta.assign(_to_tensor(state_dict[f"{p}.norm.bias"]))
|
|
921
|
+
|
|
922
|
+
_load_gat(model.encoder, "encoder")
|
|
923
|
+
_load_gat(model.decoder, "decoder")
|
|
924
|
+
_load_gat(model.encoder_ema, "encoder_ema")
|
|
925
|
+
|
|
926
|
+
# 2. encoder_to_decoder
|
|
927
|
+
if "encoder_to_decoder.weight" in state_dict:
|
|
928
|
+
model.encoder_to_decoder.kernel.assign(_to_tensor(state_dict["encoder_to_decoder.weight"].t()))
|
|
929
|
+
|
|
930
|
+
# 3. Projectors
|
|
931
|
+
def _load_projector(proj_module, prefix):
|
|
932
|
+
if f"{prefix}.0.weight" in state_dict:
|
|
933
|
+
proj_module.layers[0].kernel.assign(_to_tensor(state_dict[f"{prefix}.0.weight"].t()))
|
|
934
|
+
if f"{prefix}.0.bias" in state_dict:
|
|
935
|
+
proj_module.layers[0].bias.assign(_to_tensor(state_dict[f"{prefix}.0.bias"]))
|
|
936
|
+
if f"{prefix}.1.weight" in state_dict and hasattr(proj_module.layers[1], "alpha"):
|
|
937
|
+
proj_module.layers[1].alpha.assign(_to_tensor(state_dict[f"{prefix}.1.weight"]))
|
|
938
|
+
if f"{prefix}.2.weight" in state_dict:
|
|
939
|
+
proj_module.layers[2].kernel.assign(_to_tensor(state_dict[f"{prefix}.2.weight"].t()))
|
|
940
|
+
if f"{prefix}.2.bias" in state_dict:
|
|
941
|
+
proj_module.layers[2].bias.assign(_to_tensor(state_dict[f"{prefix}.2.bias"]))
|
|
942
|
+
|
|
943
|
+
_load_projector(model.projector, "projector")
|
|
944
|
+
_load_projector(model.projector_ema, "projector_ema")
|
|
945
|
+
|
|
946
|
+
# 4. Predictor
|
|
947
|
+
if "predictor.0.weight" in state_dict and hasattr(model.predictor.layers[0], "alpha"):
|
|
948
|
+
model.predictor.layers[0].alpha.assign(_to_tensor(state_dict["predictor.0.weight"]))
|
|
949
|
+
if "predictor.1.weight" in state_dict:
|
|
950
|
+
model.predictor.layers[1].kernel.assign(_to_tensor(state_dict["predictor.1.weight"].t()))
|
|
951
|
+
if "predictor.1.bias" in state_dict:
|
|
952
|
+
model.predictor.layers[1].bias.assign(_to_tensor(state_dict["predictor.1.bias"]))
|
|
953
|
+
|
|
954
|
+
return model
|