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,228 @@
|
|
|
1
|
+
"""High-level Graph / Molecular Regression Task."""
|
|
2
|
+
|
|
3
|
+
from typing import Any, Dict, List, Optional, Union
|
|
4
|
+
import keras
|
|
5
|
+
from keras import layers, ops
|
|
6
|
+
import numpy as np
|
|
7
|
+
|
|
8
|
+
from k3_node.tasks.base import BaseTask
|
|
9
|
+
from k3_node.tasks.backbone_resolver import resolve_backbone
|
|
10
|
+
from k3_node.layers import pool as k3_pool
|
|
11
|
+
from k3_node.loader import DataLoader
|
|
12
|
+
|
|
13
|
+
|
|
14
|
+
class GraphRegressionModel(keras.Model):
|
|
15
|
+
r"""Combines GNN backbone, pooling, and a continuous regression head."""
|
|
16
|
+
|
|
17
|
+
def __init__(
|
|
18
|
+
self,
|
|
19
|
+
backbone: keras.Model,
|
|
20
|
+
pooling: str = "add",
|
|
21
|
+
out_channels: int = 1,
|
|
22
|
+
dropout: float = 0.0,
|
|
23
|
+
):
|
|
24
|
+
super().__init__()
|
|
25
|
+
self.backbone = backbone
|
|
26
|
+
self.pooling = pooling
|
|
27
|
+
self.dropout = layers.Dropout(dropout) if dropout > 0 else None
|
|
28
|
+
self.head = layers.Dense(out_channels)
|
|
29
|
+
self.num_graphs = None
|
|
30
|
+
|
|
31
|
+
def call(self, inputs, training=False):
|
|
32
|
+
if isinstance(inputs, (tuple, list)):
|
|
33
|
+
x, edge_index = inputs[0], inputs[1]
|
|
34
|
+
batch = inputs[2] if len(inputs) > 2 else None
|
|
35
|
+
size = inputs[3] if len(inputs) > 3 else self.num_graphs
|
|
36
|
+
else:
|
|
37
|
+
x = inputs
|
|
38
|
+
edge_index = getattr(x, "edge_index", None)
|
|
39
|
+
batch = getattr(x, "batch", None)
|
|
40
|
+
size = getattr(x, "num_graphs", self.num_graphs)
|
|
41
|
+
x = getattr(x, "x", x)
|
|
42
|
+
|
|
43
|
+
if batch is None:
|
|
44
|
+
batch = ops.zeros((ops.shape(x)[0],), dtype="int64")
|
|
45
|
+
|
|
46
|
+
h = self.backbone((x, edge_index), training=training)
|
|
47
|
+
|
|
48
|
+
if self.pooling in ("add", "sum", "global_add_pool"):
|
|
49
|
+
g = k3_pool.global_add_pool(h, batch, size=size)
|
|
50
|
+
elif self.pooling in ("mean", "global_mean_pool"):
|
|
51
|
+
g = k3_pool.global_mean_pool(h, batch, size=size)
|
|
52
|
+
elif self.pooling in ("max", "global_max_pool"):
|
|
53
|
+
g = k3_pool.global_max_pool(h, batch, size=size)
|
|
54
|
+
else:
|
|
55
|
+
g = k3_pool.global_add_pool(h, batch, size=size)
|
|
56
|
+
|
|
57
|
+
if self.dropout is not None:
|
|
58
|
+
g = self.dropout(g, training=training)
|
|
59
|
+
return self.head(g)
|
|
60
|
+
|
|
61
|
+
|
|
62
|
+
class GraphRegressor(BaseTask):
|
|
63
|
+
r"""High-level estimator for graph and molecular property regression tasks.
|
|
64
|
+
|
|
65
|
+
Args:
|
|
66
|
+
backbone: Architecture (``"schnet"``, ``"dimenet++"``, ``"attentive_fp"``,
|
|
67
|
+
``"pna"``, ``"gin"``, ``"gcn"``, etc.) or custom model. (default: ``"gin"``)
|
|
68
|
+
in_channels (int, optional): Size of input node features.
|
|
69
|
+
hidden_channels (int, optional): Hidden feature dimension. (default: ``64``)
|
|
70
|
+
out_channels (int, optional): Number of continuous target variables. (default: ``1``)
|
|
71
|
+
num_layers (int, optional): Number of GNN layers. (default: ``3``)
|
|
72
|
+
pooling (str, optional): Readout pooling (``"add"``, ``"mean"``, ``"max"``). (default: ``"add"``)
|
|
73
|
+
loss (str, optional): Loss function (``"mae"`` or ``"mse"``). (default: ``"mae"``)
|
|
74
|
+
**backbone_kwargs: Additional arguments forwarded to the backbone constructor.
|
|
75
|
+
"""
|
|
76
|
+
|
|
77
|
+
def __init__(
|
|
78
|
+
self,
|
|
79
|
+
backbone: Union[str, keras.Model] = "gin",
|
|
80
|
+
in_channels: Optional[int] = None,
|
|
81
|
+
hidden_channels: int = 64,
|
|
82
|
+
out_channels: int = 1,
|
|
83
|
+
num_layers: int = 3,
|
|
84
|
+
pooling: str = "add",
|
|
85
|
+
loss: str = "mae",
|
|
86
|
+
**backbone_kwargs,
|
|
87
|
+
):
|
|
88
|
+
super().__init__()
|
|
89
|
+
self.backbone = backbone
|
|
90
|
+
self.in_channels = in_channels
|
|
91
|
+
self.hidden_channels = hidden_channels
|
|
92
|
+
self.out_channels = out_channels
|
|
93
|
+
self.num_layers = num_layers
|
|
94
|
+
self.pooling = pooling
|
|
95
|
+
self.loss_name = loss
|
|
96
|
+
self.backbone_kwargs = backbone_kwargs
|
|
97
|
+
|
|
98
|
+
def _init_model(self, sample_data: Any):
|
|
99
|
+
name = str(self.backbone).lower()
|
|
100
|
+
is_direct_graph_model = any(k in name for k in ("schnet", "dimenet", "attentive"))
|
|
101
|
+
|
|
102
|
+
in_c = self.in_channels or getattr(sample_data, "num_node_features", None) or getattr(sample_data, "num_features", None) or 16
|
|
103
|
+
self.in_channels = in_c
|
|
104
|
+
|
|
105
|
+
resolved = resolve_backbone(
|
|
106
|
+
self.backbone,
|
|
107
|
+
in_channels=in_c,
|
|
108
|
+
out_channels=self.out_channels,
|
|
109
|
+
hidden_channels=self.hidden_channels,
|
|
110
|
+
num_layers=self.num_layers,
|
|
111
|
+
**self.backbone_kwargs,
|
|
112
|
+
)
|
|
113
|
+
|
|
114
|
+
if is_direct_graph_model:
|
|
115
|
+
self.model = resolved
|
|
116
|
+
else:
|
|
117
|
+
self.model = GraphRegressionModel(
|
|
118
|
+
backbone=resolved,
|
|
119
|
+
pooling=self.pooling,
|
|
120
|
+
out_channels=self.out_channels,
|
|
121
|
+
)
|
|
122
|
+
|
|
123
|
+
def fit(
|
|
124
|
+
self,
|
|
125
|
+
dataset: Any,
|
|
126
|
+
epochs: int = 20,
|
|
127
|
+
lr: float = 0.001,
|
|
128
|
+
batch_size: int = 32,
|
|
129
|
+
shuffle: bool = True,
|
|
130
|
+
verbose: int = 1,
|
|
131
|
+
callbacks: Optional[List[Any]] = None,
|
|
132
|
+
):
|
|
133
|
+
r"""Trains the graph regressor."""
|
|
134
|
+
if not isinstance(dataset, DataLoader):
|
|
135
|
+
loader = DataLoader(dataset, batch_size=batch_size, shuffle=shuffle)
|
|
136
|
+
sample = dataset[0]
|
|
137
|
+
else:
|
|
138
|
+
loader = dataset
|
|
139
|
+
sample = next(iter(loader))
|
|
140
|
+
|
|
141
|
+
if self.model is None:
|
|
142
|
+
self._init_model(sample)
|
|
143
|
+
|
|
144
|
+
if not self._is_compiled:
|
|
145
|
+
loss = keras.losses.MeanAbsoluteError() if self.loss_name == "mae" else keras.losses.MeanSquaredError()
|
|
146
|
+
self.model.compile(
|
|
147
|
+
optimizer=keras.optimizers.Adam(learning_rate=lr),
|
|
148
|
+
loss=loss,
|
|
149
|
+
metrics=[keras.metrics.MeanAbsoluteError(name="mae")],
|
|
150
|
+
)
|
|
151
|
+
self._is_compiled = True
|
|
152
|
+
|
|
153
|
+
history = {"loss": [], "mae": []}
|
|
154
|
+
for epoch in range(epochs):
|
|
155
|
+
batch_losses = []
|
|
156
|
+
batch_maes = []
|
|
157
|
+
for batch in loader:
|
|
158
|
+
x = ops.convert_to_tensor(batch.x, dtype="float32")
|
|
159
|
+
edge_index = ops.convert_to_tensor(batch.edge_index, dtype="int64")
|
|
160
|
+
batch_vec = ops.convert_to_tensor(batch.batch, dtype="int64")
|
|
161
|
+
y = ops.cast(ops.convert_to_tensor(batch.y), "float32")
|
|
162
|
+
if len(ops.shape(y)) == 1:
|
|
163
|
+
y = ops.expand_dims(y, axis=-1)
|
|
164
|
+
num_g = int(ops.shape(y)[0])
|
|
165
|
+
if hasattr(self.model, "num_graphs"):
|
|
166
|
+
self.model.num_graphs = num_g
|
|
167
|
+
|
|
168
|
+
if not self.model.built:
|
|
169
|
+
y_pred = self.model((x, edge_index, batch_vec), training=False)
|
|
170
|
+
self.model.built = True
|
|
171
|
+
if hasattr(self.model, "_compile_loss") and self.model._compile_loss is not None:
|
|
172
|
+
self.model._compile_loss.build(y, y_pred)
|
|
173
|
+
if hasattr(self.model, "_compile_metrics") and self.model._compile_metrics is not None:
|
|
174
|
+
self.model._compile_metrics.build(y, y_pred)
|
|
175
|
+
if self.model.optimizer is not None and not self.model.optimizer.built:
|
|
176
|
+
self.model.optimizer.build(self.model.trainable_variables)
|
|
177
|
+
|
|
178
|
+
res = self.model.train_on_batch((x, edge_index, batch_vec), y)
|
|
179
|
+
if isinstance(res, (list, tuple)):
|
|
180
|
+
batch_losses.append(float(res[0]))
|
|
181
|
+
if len(res) > 1:
|
|
182
|
+
batch_maes.append(float(res[1]))
|
|
183
|
+
else:
|
|
184
|
+
batch_losses.append(float(res))
|
|
185
|
+
|
|
186
|
+
avg_loss = float(np.mean(batch_losses)) if batch_losses else 0.0
|
|
187
|
+
avg_mae = float(np.mean(batch_maes)) if batch_maes else 0.0
|
|
188
|
+
history["loss"].append(avg_loss)
|
|
189
|
+
history["mae"].append(avg_mae)
|
|
190
|
+
if verbose:
|
|
191
|
+
print(f"Epoch {epoch + 1}/{epochs} - loss: {avg_loss:.4f} - mae: {avg_mae:.4f}")
|
|
192
|
+
|
|
193
|
+
return history
|
|
194
|
+
|
|
195
|
+
def predict(self, dataset_or_loader: Any, batch_size: int = 32):
|
|
196
|
+
r"""Returns continuous predictions for graphs."""
|
|
197
|
+
if not isinstance(dataset_or_loader, DataLoader):
|
|
198
|
+
loader = DataLoader(dataset_or_loader, batch_size=batch_size, shuffle=False)
|
|
199
|
+
else:
|
|
200
|
+
loader = dataset_or_loader
|
|
201
|
+
|
|
202
|
+
preds = []
|
|
203
|
+
for batch in loader:
|
|
204
|
+
x = ops.convert_to_tensor(batch.x, dtype="float32")
|
|
205
|
+
edge_index = ops.convert_to_tensor(batch.edge_index, dtype="int64")
|
|
206
|
+
batch_vec = ops.convert_to_tensor(batch.batch, dtype="int64")
|
|
207
|
+
num_g = int(ops.convert_to_numpy(ops.max(batch_vec))) + 1 if ops.shape(batch_vec)[0] > 0 else 1
|
|
208
|
+
if hasattr(self.model, "num_graphs"):
|
|
209
|
+
self.model.num_graphs = num_g
|
|
210
|
+
pred = self.model((x, edge_index, batch_vec), training=False)
|
|
211
|
+
preds.append(pred)
|
|
212
|
+
return ops.concatenate(preds, axis=0)
|
|
213
|
+
|
|
214
|
+
def evaluate(self, dataset_or_loader: Any, batch_size: int = 32) -> Dict[str, float]:
|
|
215
|
+
r"""Evaluates MAE on the dataset."""
|
|
216
|
+
preds = self.predict(dataset_or_loader, batch_size=batch_size)
|
|
217
|
+
ys = []
|
|
218
|
+
loader = dataset_or_loader if isinstance(dataset_or_loader, DataLoader) else DataLoader(dataset_or_loader, batch_size=batch_size, shuffle=False)
|
|
219
|
+
for batch in loader:
|
|
220
|
+
y_b = ops.cast(batch.y, "float32")
|
|
221
|
+
if len(ops.shape(y_b)) == 1:
|
|
222
|
+
y_b = ops.expand_dims(y_b, axis=-1)
|
|
223
|
+
ys.append(y_b)
|
|
224
|
+
y_all = ops.concatenate(ys, axis=0)
|
|
225
|
+
diff = ops.abs(preds - y_all)
|
|
226
|
+
mae = float(ops.convert_to_numpy(ops.mean(diff)))
|
|
227
|
+
mse = float(ops.convert_to_numpy(ops.mean(ops.square(diff))))
|
|
228
|
+
return {"mae": mae, "mse": mse, "loss": mae}
|
|
@@ -0,0 +1,306 @@
|
|
|
1
|
+
"""High-level Link Prediction Task."""
|
|
2
|
+
|
|
3
|
+
from typing import Any, Dict, List, Optional, Tuple, Union
|
|
4
|
+
import keras
|
|
5
|
+
from keras import layers, ops
|
|
6
|
+
import numpy as np
|
|
7
|
+
|
|
8
|
+
from k3_node.tasks.base import BaseTask
|
|
9
|
+
from k3_node.tasks.backbone_resolver import resolve_backbone
|
|
10
|
+
from k3_node.models.utils import negative_sampling
|
|
11
|
+
|
|
12
|
+
|
|
13
|
+
class LinkPredictionModel(keras.Model):
|
|
14
|
+
r"""Internal neural network module bundling encoder and edge decoder."""
|
|
15
|
+
|
|
16
|
+
def __init__(
|
|
17
|
+
self,
|
|
18
|
+
encoder: keras.Model,
|
|
19
|
+
decoder_type: str = "inner_product",
|
|
20
|
+
hidden_channels: int = 64,
|
|
21
|
+
**kwargs,
|
|
22
|
+
):
|
|
23
|
+
kwargs.setdefault("name", "link_prediction_model")
|
|
24
|
+
super().__init__(**kwargs)
|
|
25
|
+
self.encoder = encoder
|
|
26
|
+
self.decoder_type = decoder_type.lower()
|
|
27
|
+
if self.decoder_type == "mlp":
|
|
28
|
+
self.decoder_mlp = keras.Sequential([
|
|
29
|
+
layers.Dense(hidden_channels, activation="relu"),
|
|
30
|
+
layers.Dense(1),
|
|
31
|
+
])
|
|
32
|
+
else:
|
|
33
|
+
self.decoder_mlp = None
|
|
34
|
+
|
|
35
|
+
def encode(self, inputs):
|
|
36
|
+
if hasattr(inputs, "inputs"):
|
|
37
|
+
inputs = inputs.inputs
|
|
38
|
+
return self.encoder(inputs)
|
|
39
|
+
|
|
40
|
+
def decode(self, z, edge_label_index):
|
|
41
|
+
edge_label_index = ops.convert_to_tensor(edge_label_index, dtype="int32")
|
|
42
|
+
src_idx = edge_label_index[0]
|
|
43
|
+
dst_idx = edge_label_index[1]
|
|
44
|
+
src = ops.take(z, src_idx, axis=0)
|
|
45
|
+
dst = ops.take(z, dst_idx, axis=0)
|
|
46
|
+
|
|
47
|
+
if self.decoder_type in ("inner_product", "dot"):
|
|
48
|
+
return ops.sum(src * dst, axis=-1)
|
|
49
|
+
elif self.decoder_type == "cosine":
|
|
50
|
+
src_norm = ops.sqrt(ops.maximum(ops.sum(ops.square(src), axis=-1, keepdims=True), 1e-8))
|
|
51
|
+
dst_norm = ops.sqrt(ops.maximum(ops.sum(ops.square(dst), axis=-1, keepdims=True), 1e-8))
|
|
52
|
+
return ops.sum((src / src_norm) * (dst / dst_norm), axis=-1)
|
|
53
|
+
elif self.decoder_type == "mlp":
|
|
54
|
+
feat = ops.concatenate([src, dst], axis=-1)
|
|
55
|
+
return ops.squeeze(self.decoder_mlp(feat), axis=-1)
|
|
56
|
+
else:
|
|
57
|
+
raise ValueError(f"Unknown decoder type '{self.decoder_type}'. Supported: 'inner_product', 'cosine', 'mlp'.")
|
|
58
|
+
|
|
59
|
+
def call(self, inputs, training=None):
|
|
60
|
+
r"""Executes forward pass.
|
|
61
|
+
inputs can be either:
|
|
62
|
+
- A tuple of (graph_inputs, edge_label_index)
|
|
63
|
+
- Just graph_inputs (in which case node embeddings z are returned)
|
|
64
|
+
"""
|
|
65
|
+
if isinstance(inputs, (tuple, list)) and len(inputs) == 2 and (isinstance(inputs[1], np.ndarray) or ops.is_tensor(inputs[1])):
|
|
66
|
+
graph_inputs, edge_label_index = inputs
|
|
67
|
+
z = self.encode(graph_inputs)
|
|
68
|
+
return self.decode(z, edge_label_index)
|
|
69
|
+
else:
|
|
70
|
+
return self.encode(inputs)
|
|
71
|
+
|
|
72
|
+
|
|
73
|
+
class LinkPredictor(BaseTask):
|
|
74
|
+
r"""High-level estimator for link prediction tasks on graphs.
|
|
75
|
+
|
|
76
|
+
Args:
|
|
77
|
+
backbone: Model architecture string (``"gcn"``, ``"gat"``, ``"sage"``,
|
|
78
|
+
``"gin"``, ``"pna"``, ``"mlp"``, etc.) or custom :class:`keras.Model`.
|
|
79
|
+
(default: ``"gcn"``)
|
|
80
|
+
in_channels (int, optional): Size of input node features. If not specified,
|
|
81
|
+
it is automatically inferred from the dataset during :meth:`fit`.
|
|
82
|
+
hidden_channels (int, optional): Dimensionality of hidden node features.
|
|
83
|
+
(default: ``64``)
|
|
84
|
+
out_channels (int, optional): Dimensionality of output node embeddings
|
|
85
|
+
used for link scoring. (default: ``64``)
|
|
86
|
+
num_layers (int, optional): Number of message passing layers. (default: ``2``)
|
|
87
|
+
decoder (str, optional): Type of edge score decoder (``"inner_product"``,
|
|
88
|
+
``"cosine"``, or ``"mlp"``). (default: ``"inner_product"``)
|
|
89
|
+
dropout (float, optional): Dropout probability. (default: ``0.0``)
|
|
90
|
+
**backbone_kwargs: Additional arguments forwarded to the backbone constructor.
|
|
91
|
+
"""
|
|
92
|
+
|
|
93
|
+
def __init__(
|
|
94
|
+
self,
|
|
95
|
+
backbone: Union[str, keras.Model] = "gcn",
|
|
96
|
+
in_channels: Optional[int] = None,
|
|
97
|
+
hidden_channels: int = 64,
|
|
98
|
+
out_channels: int = 64,
|
|
99
|
+
num_layers: int = 2,
|
|
100
|
+
decoder: str = "inner_product",
|
|
101
|
+
dropout: float = 0.0,
|
|
102
|
+
**backbone_kwargs,
|
|
103
|
+
):
|
|
104
|
+
super().__init__()
|
|
105
|
+
self.backbone = backbone
|
|
106
|
+
self.in_channels = in_channels
|
|
107
|
+
self.hidden_channels = hidden_channels
|
|
108
|
+
self.out_channels = out_channels
|
|
109
|
+
self.num_layers = num_layers
|
|
110
|
+
self.decoder = decoder
|
|
111
|
+
self.dropout = dropout
|
|
112
|
+
self.backbone_kwargs = backbone_kwargs
|
|
113
|
+
|
|
114
|
+
if isinstance(backbone, keras.Model):
|
|
115
|
+
self.model = LinkPredictionModel(
|
|
116
|
+
encoder=backbone,
|
|
117
|
+
decoder_type=decoder,
|
|
118
|
+
hidden_channels=hidden_channels,
|
|
119
|
+
)
|
|
120
|
+
|
|
121
|
+
def _init_model(self, data: Any):
|
|
122
|
+
r"""Infers missing dimensions and instantiates the link prediction module."""
|
|
123
|
+
in_c = self.in_channels
|
|
124
|
+
if in_c is None:
|
|
125
|
+
if hasattr(data, "num_node_features") and data.num_node_features > 0:
|
|
126
|
+
in_c = data.num_node_features
|
|
127
|
+
elif hasattr(data, "num_features") and data.num_features > 0:
|
|
128
|
+
in_c = data.num_features
|
|
129
|
+
elif hasattr(data, "x") and data.x is not None:
|
|
130
|
+
in_c = int(ops.shape(data.x)[-1])
|
|
131
|
+
else:
|
|
132
|
+
raise ValueError("Could not automatically infer in_channels from data. Please specify in_channels.")
|
|
133
|
+
|
|
134
|
+
self.in_channels = in_c
|
|
135
|
+
|
|
136
|
+
encoder = resolve_backbone(
|
|
137
|
+
self.backbone,
|
|
138
|
+
in_channels=in_c,
|
|
139
|
+
out_channels=self.out_channels,
|
|
140
|
+
hidden_channels=self.hidden_channels,
|
|
141
|
+
num_layers=self.num_layers,
|
|
142
|
+
dropout=self.dropout,
|
|
143
|
+
**self.backbone_kwargs,
|
|
144
|
+
)
|
|
145
|
+
|
|
146
|
+
self.model = LinkPredictionModel(
|
|
147
|
+
encoder=encoder,
|
|
148
|
+
decoder_type=self.decoder,
|
|
149
|
+
hidden_channels=self.hidden_channels,
|
|
150
|
+
)
|
|
151
|
+
|
|
152
|
+
def fit(
|
|
153
|
+
self,
|
|
154
|
+
data: Any,
|
|
155
|
+
edge_label_index: Optional[Any] = None,
|
|
156
|
+
edge_label: Optional[Any] = None,
|
|
157
|
+
epochs: int = 20,
|
|
158
|
+
lr: float = 0.01,
|
|
159
|
+
weight_decay: float = 0.0,
|
|
160
|
+
neg_ratio: float = 1.0,
|
|
161
|
+
verbose: int = 1,
|
|
162
|
+
callbacks: Optional[List[Any]] = None,
|
|
163
|
+
):
|
|
164
|
+
r"""Trains the link predictor on graph connectivity."""
|
|
165
|
+
if self.model is None:
|
|
166
|
+
self._init_model(data)
|
|
167
|
+
|
|
168
|
+
if not self._is_compiled:
|
|
169
|
+
opt = keras.optimizers.Adam(learning_rate=lr, weight_decay=weight_decay)
|
|
170
|
+
loss = keras.losses.BinaryCrossentropy(from_logits=True)
|
|
171
|
+
metrics = [keras.metrics.BinaryAccuracy(name="acc", threshold=0.0)]
|
|
172
|
+
self.model.compile(optimizer=opt, loss=loss, metrics=metrics)
|
|
173
|
+
self._is_compiled = True
|
|
174
|
+
|
|
175
|
+
graph_inputs = self._extract_inputs(data)
|
|
176
|
+
|
|
177
|
+
# Check for pre-split edge labels on data or passed explicitly
|
|
178
|
+
has_labels = edge_label_index is not None and edge_label is not None
|
|
179
|
+
if not has_labels:
|
|
180
|
+
if hasattr(data, "train_edge_label_index") and hasattr(data, "train_edge_label"):
|
|
181
|
+
edge_label_index = data.train_edge_label_index
|
|
182
|
+
edge_label = data.train_edge_label
|
|
183
|
+
has_labels = True
|
|
184
|
+
elif hasattr(data, "edge_label_index") and hasattr(data, "edge_label"):
|
|
185
|
+
edge_label_index = data.edge_label_index
|
|
186
|
+
edge_label = data.edge_label
|
|
187
|
+
has_labels = True
|
|
188
|
+
|
|
189
|
+
num_nodes = None
|
|
190
|
+
if hasattr(data, "num_nodes") and data.num_nodes is not None:
|
|
191
|
+
num_nodes = data.num_nodes
|
|
192
|
+
elif hasattr(data, "x") and data.x is not None:
|
|
193
|
+
num_nodes = int(ops.shape(data.x)[0])
|
|
194
|
+
|
|
195
|
+
if not has_labels:
|
|
196
|
+
pos_edge_index = data.edge_index
|
|
197
|
+
pos_np = ops.convert_to_numpy(pos_edge_index).astype(np.int32)
|
|
198
|
+
num_pos = pos_np.shape[1]
|
|
199
|
+
num_neg = int(num_pos * neg_ratio)
|
|
200
|
+
|
|
201
|
+
history = {"loss": [], "acc": []}
|
|
202
|
+
for epoch in range(epochs):
|
|
203
|
+
if has_labels:
|
|
204
|
+
total_edges = ops.convert_to_tensor(edge_label_index, dtype="int32")
|
|
205
|
+
labels = ops.cast(ops.convert_to_tensor(edge_label), "float32")
|
|
206
|
+
else:
|
|
207
|
+
neg_np = negative_sampling(pos_np, num_nodes=num_nodes, num_neg_samples=num_neg)
|
|
208
|
+
total_edges = ops.convert_to_tensor(np.concatenate([pos_np, neg_np], axis=1), dtype="int32")
|
|
209
|
+
labels = ops.convert_to_tensor(
|
|
210
|
+
np.concatenate([np.ones(num_pos, dtype=np.float32), np.zeros(num_neg, dtype=np.float32)]),
|
|
211
|
+
dtype="float32",
|
|
212
|
+
)
|
|
213
|
+
|
|
214
|
+
if not self.model.built:
|
|
215
|
+
y_pred = self.model((graph_inputs, total_edges), training=False)
|
|
216
|
+
self.model.built = True
|
|
217
|
+
if hasattr(self.model, "_compile_loss") and self.model._compile_loss is not None:
|
|
218
|
+
self.model._compile_loss.build(labels, y_pred)
|
|
219
|
+
if hasattr(self.model, "_compile_metrics") and self.model._compile_metrics is not None:
|
|
220
|
+
self.model._compile_metrics.build(labels, y_pred)
|
|
221
|
+
if self.model.optimizer is not None and not self.model.optimizer.built:
|
|
222
|
+
self.model.optimizer.build(self.model.trainable_variables)
|
|
223
|
+
|
|
224
|
+
res = self.model.train_on_batch((graph_inputs, total_edges), labels)
|
|
225
|
+
if isinstance(res, (list, tuple)):
|
|
226
|
+
l, a = float(res[0]), float(res[1]) if len(res) > 1 else 0.0
|
|
227
|
+
else:
|
|
228
|
+
l, a = float(res), 0.0
|
|
229
|
+
history["loss"].append(l)
|
|
230
|
+
history["acc"].append(a)
|
|
231
|
+
if verbose:
|
|
232
|
+
print(f"Epoch {epoch + 1}/{epochs} - loss: {l:.4f} - acc: {a:.4f}")
|
|
233
|
+
|
|
234
|
+
return history
|
|
235
|
+
|
|
236
|
+
def encode(self, data: Any):
|
|
237
|
+
r"""Computes latent node representations for the input graph."""
|
|
238
|
+
if self.model is None:
|
|
239
|
+
raise RuntimeError("Model is not initialized. Call fit() or construct with a model first.")
|
|
240
|
+
graph_inputs = self._extract_inputs(data)
|
|
241
|
+
return self.model.encode(graph_inputs)
|
|
242
|
+
|
|
243
|
+
def predict_proba(self, data: Any, edge_label_index: Optional[Any] = None):
|
|
244
|
+
r"""Predicts link existence probabilities for edge pairs."""
|
|
245
|
+
if self.model is None:
|
|
246
|
+
raise RuntimeError("Model is not initialized. Fit or load a model first.")
|
|
247
|
+
|
|
248
|
+
if edge_label_index is None:
|
|
249
|
+
if hasattr(data, "test_edge_label_index"):
|
|
250
|
+
edge_label_index = data.test_edge_label_index
|
|
251
|
+
elif hasattr(data, "edge_label_index"):
|
|
252
|
+
edge_label_index = data.edge_label_index
|
|
253
|
+
elif hasattr(data, "edge_index"):
|
|
254
|
+
edge_label_index = data.edge_index
|
|
255
|
+
else:
|
|
256
|
+
raise ValueError("No edge_label_index provided and none found on data object.")
|
|
257
|
+
|
|
258
|
+
z = self.encode(data)
|
|
259
|
+
logits = self.model.decode(z, edge_label_index)
|
|
260
|
+
return ops.sigmoid(logits)
|
|
261
|
+
|
|
262
|
+
def predict(self, data: Any, edge_label_index: Optional[Any] = None, threshold: float = 0.5):
|
|
263
|
+
r"""Predicts binary link presence (0 or 1) for edge pairs."""
|
|
264
|
+
probs = self.predict_proba(data, edge_label_index=edge_label_index)
|
|
265
|
+
return ops.cast(probs >= threshold, "int64")
|
|
266
|
+
|
|
267
|
+
def evaluate(
|
|
268
|
+
self,
|
|
269
|
+
data: Any,
|
|
270
|
+
edge_label_index: Optional[Any] = None,
|
|
271
|
+
edge_label: Optional[Any] = None,
|
|
272
|
+
) -> Dict[str, float]:
|
|
273
|
+
r"""Evaluates link prediction performance (AUC, AP, Accuracy)."""
|
|
274
|
+
if edge_label_index is None and edge_label is None:
|
|
275
|
+
if hasattr(data, "test_edge_label_index") and hasattr(data, "test_edge_label"):
|
|
276
|
+
edge_label_index = data.test_edge_label_index
|
|
277
|
+
edge_label = data.test_edge_label
|
|
278
|
+
elif hasattr(data, "edge_label_index") and hasattr(data, "edge_label"):
|
|
279
|
+
edge_label_index = data.edge_label_index
|
|
280
|
+
edge_label = data.edge_label
|
|
281
|
+
else:
|
|
282
|
+
# Sample negative edges against edge_index
|
|
283
|
+
pos_edges = ops.convert_to_numpy(data.edge_index).astype(np.int32)
|
|
284
|
+
num_nodes = data.num_nodes if hasattr(data, "num_nodes") else int(ops.shape(data.x)[0])
|
|
285
|
+
neg_edges = negative_sampling(pos_edges, num_nodes=num_nodes, num_neg_samples=pos_edges.shape[1])
|
|
286
|
+
edge_label_index = np.concatenate([pos_edges, neg_edges], axis=1)
|
|
287
|
+
edge_label = np.concatenate([np.ones(pos_edges.shape[1]), np.zeros(neg_edges.shape[1])])
|
|
288
|
+
|
|
289
|
+
probs = self.predict_proba(data, edge_label_index=edge_label_index)
|
|
290
|
+
probs_np = ops.convert_to_numpy(probs).flatten()
|
|
291
|
+
y_np = ops.convert_to_numpy(edge_label).flatten()
|
|
292
|
+
|
|
293
|
+
metrics = {}
|
|
294
|
+
# Accuracy
|
|
295
|
+
preds_bin = (probs_np >= 0.5).astype(np.float32)
|
|
296
|
+
metrics["accuracy"] = float(np.mean(preds_bin == y_np))
|
|
297
|
+
|
|
298
|
+
# AUC and AP via scikit-learn when available
|
|
299
|
+
try:
|
|
300
|
+
from sklearn.metrics import roc_auc_score, average_precision_score
|
|
301
|
+
metrics["auc"] = float(roc_auc_score(y_np, probs_np))
|
|
302
|
+
metrics["ap"] = float(average_precision_score(y_np, probs_np))
|
|
303
|
+
except ImportError:
|
|
304
|
+
pass
|
|
305
|
+
|
|
306
|
+
return metrics
|