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,200 @@
|
|
|
1
|
+
from typing import Dict, List, Optional
|
|
2
|
+
|
|
3
|
+
import keras
|
|
4
|
+
from keras import ops
|
|
5
|
+
|
|
6
|
+
|
|
7
|
+
class JumpingKnowledge(keras.layers.Layer):
|
|
8
|
+
r"""The Jumping Knowledge layer aggregation module from the
|
|
9
|
+
`"Representation Learning on Graphs with Jumping Knowledge Networks"
|
|
10
|
+
<https://arxiv.org/abs/1806.03536>`_ paper.
|
|
11
|
+
|
|
12
|
+
Jumping knowledge is performed based on either **concatenation**
|
|
13
|
+
(:obj:`"cat"`)
|
|
14
|
+
|
|
15
|
+
.. math::
|
|
16
|
+
|
|
17
|
+
\mathbf{x}_v^{(1)} \, \Vert \, \ldots \, \Vert \, \mathbf{x}_v^{(T)},
|
|
18
|
+
|
|
19
|
+
**max pooling** (:obj:`"max"`)
|
|
20
|
+
|
|
21
|
+
.. math::
|
|
22
|
+
|
|
23
|
+
\max \left( \mathbf{x}_v^{(1)}, \ldots, \mathbf{x}_v^{(T)} \right),
|
|
24
|
+
|
|
25
|
+
or **weighted summation**
|
|
26
|
+
|
|
27
|
+
.. math::
|
|
28
|
+
|
|
29
|
+
\sum_{t=1}^T \alpha_v^{(t)} \mathbf{x}_v^{(t)}
|
|
30
|
+
|
|
31
|
+
with attention scores :math:`\alpha_v^{(t)}` obtained from a bi-directional
|
|
32
|
+
LSTM (:obj:`"lstm"`).
|
|
33
|
+
|
|
34
|
+
Args:
|
|
35
|
+
mode (str): The aggregation scheme to use
|
|
36
|
+
(:obj:`"cat"`, :obj:`"max"` or :obj:`"lstm"`).
|
|
37
|
+
channels (int, optional): The number of channels per representation.
|
|
38
|
+
Needs to be only set for LSTM-style aggregation.
|
|
39
|
+
(default: :obj:`None`)
|
|
40
|
+
num_layers (int, optional): The number of layers to aggregate. Needs to
|
|
41
|
+
be only set for LSTM-style aggregation. (default: :obj:`None`)
|
|
42
|
+
|
|
43
|
+
Example:
|
|
44
|
+
```python
|
|
45
|
+
import numpy as np
|
|
46
|
+
from k3_node.models import JumpingKnowledge
|
|
47
|
+
|
|
48
|
+
# Node representations from 3 GNN layers
|
|
49
|
+
xs = [np.random.rand(10, 16).astype("float32") for _ in range(3)]
|
|
50
|
+
print(tuple(JumpingKnowledge("cat")(xs).shape)) # (10, 48): concatenate all layers
|
|
51
|
+
print(tuple(JumpingKnowledge("max")(xs).shape)) # (10, 16): element-wise max
|
|
52
|
+
print(tuple(JumpingKnowledge("lstm", channels=16, num_layers=3)(xs).shape)) # (10, 16): attention over layers
|
|
53
|
+
```
|
|
54
|
+
"""
|
|
55
|
+
|
|
56
|
+
def __init__(
|
|
57
|
+
self,
|
|
58
|
+
mode: str,
|
|
59
|
+
channels: Optional[int] = None,
|
|
60
|
+
num_layers: Optional[int] = None,
|
|
61
|
+
**kwargs,
|
|
62
|
+
) -> None:
|
|
63
|
+
super().__init__(**kwargs)
|
|
64
|
+
self.mode = mode.lower()
|
|
65
|
+
assert self.mode in ['cat', 'max', 'lstm'], \
|
|
66
|
+
f"mode must be 'cat', 'max', or 'lstm', got '{mode}'"
|
|
67
|
+
|
|
68
|
+
self.channels = channels
|
|
69
|
+
self.num_layers = num_layers
|
|
70
|
+
|
|
71
|
+
if self.mode == 'lstm':
|
|
72
|
+
assert channels is not None, 'channels cannot be None for lstm'
|
|
73
|
+
assert num_layers is not None, 'num_layers cannot be None for lstm'
|
|
74
|
+
lstm_units = (num_layers * channels) // 2
|
|
75
|
+
self.lstm = keras.layers.Bidirectional(
|
|
76
|
+
keras.layers.LSTM(lstm_units, return_sequences=True),
|
|
77
|
+
merge_mode='concat',
|
|
78
|
+
)
|
|
79
|
+
self.att = keras.layers.Dense(1, use_bias=False)
|
|
80
|
+
else:
|
|
81
|
+
self.lstm = None
|
|
82
|
+
self.att = None
|
|
83
|
+
|
|
84
|
+
def reset_parameters(self) -> None:
|
|
85
|
+
r"""Resets all learnable parameters of the module."""
|
|
86
|
+
# Keras layers reinitialize on next call; explicit rebuild if needed
|
|
87
|
+
if self.lstm is not None and self.lstm.built:
|
|
88
|
+
for layer in self.lstm.layers:
|
|
89
|
+
if hasattr(layer, 'kernel') and layer.kernel is not None:
|
|
90
|
+
layer.kernel.assign(
|
|
91
|
+
keras.initializers.GlorotUniform()(layer.kernel.shape)
|
|
92
|
+
)
|
|
93
|
+
if hasattr(layer, 'recurrent_kernel') and layer.recurrent_kernel is not None:
|
|
94
|
+
layer.recurrent_kernel.assign(
|
|
95
|
+
keras.initializers.Orthogonal()(layer.recurrent_kernel.shape)
|
|
96
|
+
)
|
|
97
|
+
if hasattr(layer, 'bias') and layer.bias is not None:
|
|
98
|
+
layer.bias.assign(ops.zeros(layer.bias.shape))
|
|
99
|
+
if self.att is not None and self.att.built:
|
|
100
|
+
if self.att.kernel is not None:
|
|
101
|
+
self.att.kernel.assign(
|
|
102
|
+
keras.initializers.GlorotUniform()(self.att.kernel.shape)
|
|
103
|
+
)
|
|
104
|
+
|
|
105
|
+
def call(self, xs: List) -> object:
|
|
106
|
+
r"""Forward pass.
|
|
107
|
+
|
|
108
|
+
Args:
|
|
109
|
+
xs (List[Tensor]): List containing the layer-wise representations.
|
|
110
|
+
"""
|
|
111
|
+
if self.mode == 'cat':
|
|
112
|
+
return ops.concatenate(xs, axis=-1)
|
|
113
|
+
elif self.mode == 'max':
|
|
114
|
+
return ops.max(ops.stack(xs, axis=-1), axis=-1)
|
|
115
|
+
else: # lstm
|
|
116
|
+
assert self.lstm is not None and self.att is not None
|
|
117
|
+
x = ops.stack(xs, axis=1) # [num_nodes, num_layers, num_channels]
|
|
118
|
+
alpha = self.lstm(x) # [num_nodes, num_layers, 2*lstm_units]
|
|
119
|
+
alpha = self.att(alpha) # [num_nodes, num_layers, 1]
|
|
120
|
+
alpha = ops.squeeze(alpha, axis=-1) # [num_nodes, num_layers]
|
|
121
|
+
alpha = ops.softmax(alpha, axis=-1)
|
|
122
|
+
return ops.sum(x * ops.expand_dims(alpha, axis=-1), axis=1)
|
|
123
|
+
|
|
124
|
+
def __repr__(self) -> str:
|
|
125
|
+
if self.mode == 'lstm':
|
|
126
|
+
return (f'{self.__class__.__name__}({self.mode}, '
|
|
127
|
+
f'channels={self.channels}, layers={self.num_layers})')
|
|
128
|
+
return f'{self.__class__.__name__}({self.mode})'
|
|
129
|
+
|
|
130
|
+
|
|
131
|
+
class HeteroJumpingKnowledge(keras.layers.Layer):
|
|
132
|
+
r"""A heterogeneous version of the :class:`JumpingKnowledge` module.
|
|
133
|
+
|
|
134
|
+
Args:
|
|
135
|
+
types (List[str]): The keys of the input dictionary.
|
|
136
|
+
mode (str): The aggregation scheme to use
|
|
137
|
+
(:obj:`"cat"`, :obj:`"max"` or :obj:`"lstm"`).
|
|
138
|
+
channels (int, optional): The number of channels per representation.
|
|
139
|
+
Needs to be only set for LSTM-style aggregation.
|
|
140
|
+
(default: :obj:`None`)
|
|
141
|
+
num_layers (int, optional): The number of layers to aggregate. Needs to
|
|
142
|
+
be only set for LSTM-style aggregation. (default: :obj:`None`)
|
|
143
|
+
|
|
144
|
+
Example:
|
|
145
|
+
```python
|
|
146
|
+
import numpy as np
|
|
147
|
+
from k3_node.models import HeteroJumpingKnowledge
|
|
148
|
+
|
|
149
|
+
# Per node type: representations from 3 GNN layers
|
|
150
|
+
xs_dict = {
|
|
151
|
+
"author": [np.random.rand(3, 16).astype("float32") for _ in range(3)],
|
|
152
|
+
"paper": [np.random.rand(4, 16).astype("float32") for _ in range(3)],
|
|
153
|
+
}
|
|
154
|
+
model = HeteroJumpingKnowledge(["author", "paper"], mode="cat")
|
|
155
|
+
out_dict = model(xs_dict)
|
|
156
|
+
print(tuple(out_dict["author"].shape), tuple(out_dict["paper"].shape)) # (3, 48) (4, 48)
|
|
157
|
+
```
|
|
158
|
+
"""
|
|
159
|
+
|
|
160
|
+
def __init__(
|
|
161
|
+
self,
|
|
162
|
+
types: List[str],
|
|
163
|
+
mode: str,
|
|
164
|
+
channels: Optional[int] = None,
|
|
165
|
+
num_layers: Optional[int] = None,
|
|
166
|
+
**kwargs,
|
|
167
|
+
) -> None:
|
|
168
|
+
super().__init__(**kwargs)
|
|
169
|
+
self.mode = mode.lower()
|
|
170
|
+
self.types = list(types)
|
|
171
|
+
|
|
172
|
+
self.jk_dict = {
|
|
173
|
+
key: JumpingKnowledge(mode, channels, num_layers)
|
|
174
|
+
for key in types
|
|
175
|
+
}
|
|
176
|
+
|
|
177
|
+
def reset_parameters(self) -> None:
|
|
178
|
+
r"""Resets all learnable parameters of the module."""
|
|
179
|
+
for jk in self.jk_dict.values():
|
|
180
|
+
jk.reset_parameters()
|
|
181
|
+
|
|
182
|
+
def call(self, xs_dict: Dict[str, List]) -> Dict[str, object]:
|
|
183
|
+
r"""Forward pass.
|
|
184
|
+
|
|
185
|
+
Args:
|
|
186
|
+
xs_dict (Dict[str, List[Tensor]]): A dictionary holding a
|
|
187
|
+
list of layer-wise representation for each type.
|
|
188
|
+
"""
|
|
189
|
+
return {key: self.jk_dict[key](xs_dict[key]) for key in self.types}
|
|
190
|
+
|
|
191
|
+
def __repr__(self) -> str:
|
|
192
|
+
if self.mode == 'lstm':
|
|
193
|
+
jk = next(iter(self.jk_dict.values()))
|
|
194
|
+
return (f'{self.__class__.__name__}('
|
|
195
|
+
f'num_types={len(self.jk_dict)}, '
|
|
196
|
+
f'mode={self.mode}, channels={jk.channels}, '
|
|
197
|
+
f'layers={jk.num_layers})')
|
|
198
|
+
return (f'{self.__class__.__name__}(num_types={len(self.jk_dict)}, '
|
|
199
|
+
f'mode={self.mode})')
|
|
200
|
+
|
|
@@ -0,0 +1,110 @@
|
|
|
1
|
+
from typing import Callable, Optional
|
|
2
|
+
|
|
3
|
+
import keras
|
|
4
|
+
from keras import ops
|
|
5
|
+
|
|
6
|
+
from k3_node.layers.conv.message_passing import MessagePassing
|
|
7
|
+
from k3_node.layers.conv.utils import gcn_norm
|
|
8
|
+
|
|
9
|
+
|
|
10
|
+
class LabelPropagation(MessagePassing):
|
|
11
|
+
r"""The label propagation operator, firstly introduced in the
|
|
12
|
+
`"Learning from Labeled and Unlabeled Data with Label Propagation"
|
|
13
|
+
<http://mlg.eng.cam.ac.uk/zoubin/papers/CMU-CALD-02-107.pdf>`_ paper.
|
|
14
|
+
|
|
15
|
+
.. math::
|
|
16
|
+
\mathbf{Y}^{\prime} = \alpha \cdot \mathbf{D}^{-1/2} \mathbf{A}
|
|
17
|
+
\mathbf{D}^{-1/2} \mathbf{Y} + (1 - \alpha) \mathbf{Y},
|
|
18
|
+
|
|
19
|
+
where unlabeled data is inferred by labeled data via propagation.
|
|
20
|
+
This concrete implementation here is derived from the `"Combining Label
|
|
21
|
+
Propagation And Simple Models Out-performs Graph Neural Networks"
|
|
22
|
+
<https://arxiv.org/abs/2010.13993>`_ paper.
|
|
23
|
+
|
|
24
|
+
Args:
|
|
25
|
+
num_layers (int): The number of propagations.
|
|
26
|
+
alpha (float): The :math:`\alpha` coefficient.
|
|
27
|
+
|
|
28
|
+
Example:
|
|
29
|
+
```python
|
|
30
|
+
import numpy as np
|
|
31
|
+
from k3_node.models import LabelPropagation
|
|
32
|
+
|
|
33
|
+
y = np.array([0, 1, 2, 0, 1, 2]) # node labels
|
|
34
|
+
train_mask = np.array([True, True, True, False, False, False]) # labels known for 3 nodes
|
|
35
|
+
edge_index = np.array([[0, 1, 2, 3, 4, 5], [3, 4, 5, 0, 1, 2]])
|
|
36
|
+
model = LabelPropagation(num_layers=3, alpha=0.9)
|
|
37
|
+
out = model(y, edge_index, train_mask) # soft labels for every node
|
|
38
|
+
print(tuple(out.shape)) # (6, 3)
|
|
39
|
+
```
|
|
40
|
+
"""
|
|
41
|
+
def __init__(self, num_layers: int, alpha: float, **kwargs):
|
|
42
|
+
super().__init__(aggr='sum', **kwargs)
|
|
43
|
+
self.num_layers = num_layers
|
|
44
|
+
self.alpha = alpha
|
|
45
|
+
|
|
46
|
+
def build(self, input_shape=None):
|
|
47
|
+
self.built = True
|
|
48
|
+
|
|
49
|
+
# `mask` selects the labeled nodes; it is not a Keras sequence mask, so the output gets none
|
|
50
|
+
# (this also stops Keras from warning that the layer drops the mask).
|
|
51
|
+
supports_masking = True
|
|
52
|
+
|
|
53
|
+
def compute_mask(self, *args, **kwargs):
|
|
54
|
+
return None
|
|
55
|
+
|
|
56
|
+
def call(
|
|
57
|
+
self,
|
|
58
|
+
y,
|
|
59
|
+
edge_index,
|
|
60
|
+
mask=None,
|
|
61
|
+
edge_weight=None,
|
|
62
|
+
post_step: Optional[Callable] = None,
|
|
63
|
+
):
|
|
64
|
+
shape = ops.shape(y)
|
|
65
|
+
if len(shape) == 1:
|
|
66
|
+
num_classes = int(ops.max(y)) + 1
|
|
67
|
+
y = ops.one_hot(y, num_classes)
|
|
68
|
+
|
|
69
|
+
y = ops.cast(y, dtype="float32")
|
|
70
|
+
out = y
|
|
71
|
+
if mask is not None:
|
|
72
|
+
mask_shape = ops.shape(mask)
|
|
73
|
+
# standardize_dtype: torch.bool does not compare equal to the string "bool"
|
|
74
|
+
if len(mask_shape) == 1 and keras.backend.standardize_dtype(mask.dtype) == "bool":
|
|
75
|
+
mask_expanded = ops.expand_dims(mask, axis=-1)
|
|
76
|
+
out = ops.where(mask_expanded, y, ops.zeros_like(y))
|
|
77
|
+
else:
|
|
78
|
+
out_zeros = ops.zeros_like(y)
|
|
79
|
+
out = ops.scatter_update(out_zeros, ops.expand_dims(mask, -1), ops.take(y, mask, axis=0))
|
|
80
|
+
|
|
81
|
+
if edge_weight is None:
|
|
82
|
+
num_nodes = ops.shape(y)[0]
|
|
83
|
+
edge_index, edge_weight = gcn_norm(
|
|
84
|
+
edge_index,
|
|
85
|
+
edge_weight=None,
|
|
86
|
+
num_nodes=num_nodes,
|
|
87
|
+
add_self_loops=False,
|
|
88
|
+
dtype=y.dtype,
|
|
89
|
+
)
|
|
90
|
+
|
|
91
|
+
res = (1.0 - self.alpha) * out
|
|
92
|
+
for _ in range(self.num_layers):
|
|
93
|
+
out = self.propagate(edge_index, x=out, edge_weight=edge_weight)
|
|
94
|
+
out = self.alpha * out + res
|
|
95
|
+
if post_step is not None:
|
|
96
|
+
out = post_step(out)
|
|
97
|
+
else:
|
|
98
|
+
out = ops.clip(out, 0.0, 1.0)
|
|
99
|
+
|
|
100
|
+
return out
|
|
101
|
+
|
|
102
|
+
def message(self, x_j, edge_weight=None):
|
|
103
|
+
if edge_weight is None:
|
|
104
|
+
return x_j
|
|
105
|
+
return ops.expand_dims(edge_weight, axis=-1) * x_j
|
|
106
|
+
|
|
107
|
+
def __repr__(self) -> str:
|
|
108
|
+
return (f'{self.__class__.__name__}(num_layers={self.num_layers}, '
|
|
109
|
+
f'alpha={self.alpha})')
|
|
110
|
+
|
|
@@ -0,0 +1,171 @@
|
|
|
1
|
+
from typing import Optional, Union
|
|
2
|
+
|
|
3
|
+
import keras
|
|
4
|
+
from keras import ops
|
|
5
|
+
import numpy as np
|
|
6
|
+
|
|
7
|
+
from k3_node.layers.conv import LGConv
|
|
8
|
+
|
|
9
|
+
|
|
10
|
+
class BPRLoss:
|
|
11
|
+
r"""The Bayesian Personalized Ranking (BPR) loss."""
|
|
12
|
+
def __init__(self, lambda_reg: float = 0.0):
|
|
13
|
+
self.lambda_reg = lambda_reg
|
|
14
|
+
|
|
15
|
+
def __call__(self, positives, negatives, parameters=None):
|
|
16
|
+
diff = positives - negatives
|
|
17
|
+
# log(sigmoid(x)) = -log(1 + exp(-x)) or ops.log_sigmoid if available
|
|
18
|
+
log_prob = ops.mean(-ops.softplus(-diff))
|
|
19
|
+
|
|
20
|
+
regularization = 0.0
|
|
21
|
+
if self.lambda_reg != 0.0 and parameters is not None:
|
|
22
|
+
regularization = self.lambda_reg * ops.sum(ops.square(parameters))
|
|
23
|
+
regularization = regularization / ops.cast(ops.shape(positives)[0], "float32")
|
|
24
|
+
|
|
25
|
+
return -log_prob + regularization
|
|
26
|
+
|
|
27
|
+
|
|
28
|
+
class LightGCN(keras.Model):
|
|
29
|
+
r"""The LightGCN model from the `"LightGCN: Simplifying and Powering
|
|
30
|
+
Graph Convolution Network for Recommendation"
|
|
31
|
+
<https://arxiv.org/abs/2002.02126>`_ paper.
|
|
32
|
+
|
|
33
|
+
Args:
|
|
34
|
+
num_nodes (int): The number of nodes in the graph.
|
|
35
|
+
embedding_dim (int): The dimensionality of node embeddings.
|
|
36
|
+
num_layers (int): The number of :class:`LGConv` layers.
|
|
37
|
+
alpha (float or Tensor, optional): The scalar or vector specifying
|
|
38
|
+
the re-weighting coefficients for aggregating the final embedding.
|
|
39
|
+
(default: :obj:`None`)
|
|
40
|
+
|
|
41
|
+
Example:
|
|
42
|
+
```python
|
|
43
|
+
import numpy as np
|
|
44
|
+
from k3_node.models import LightGCN
|
|
45
|
+
|
|
46
|
+
edge_index = np.array([[0, 1, 2, 3, 4, 5, 6, 7], [1, 2, 3, 4, 5, 6, 7, 0]]) # user-item interactions
|
|
47
|
+
edge_label_index = np.array([[0, 1, 2, 3], [4, 5, 6, 7]]) # pairs to score
|
|
48
|
+
model = LightGCN(num_nodes=50, embedding_dim=16, num_layers=2)
|
|
49
|
+
scores = model(edge_index, edge_label_index)
|
|
50
|
+
print(tuple(scores.shape)) # (4,)
|
|
51
|
+
print(tuple(model.recommend(edge_index, k=2).shape)) # (50, 2): top-2 recommendations per node
|
|
52
|
+
```
|
|
53
|
+
"""
|
|
54
|
+
def __init__(
|
|
55
|
+
self,
|
|
56
|
+
num_nodes: int,
|
|
57
|
+
embedding_dim: int,
|
|
58
|
+
num_layers: int,
|
|
59
|
+
alpha: Optional[Union[float, list]] = None,
|
|
60
|
+
**kwargs,
|
|
61
|
+
):
|
|
62
|
+
super().__init__()
|
|
63
|
+
|
|
64
|
+
self.num_nodes = num_nodes
|
|
65
|
+
self.embedding_dim = embedding_dim
|
|
66
|
+
self.num_layers = num_layers
|
|
67
|
+
|
|
68
|
+
if alpha is None:
|
|
69
|
+
alpha = [1.0 / (num_layers + 1)] * (num_layers + 1)
|
|
70
|
+
elif isinstance(alpha, (int, float)):
|
|
71
|
+
alpha = [float(alpha)] * (num_layers + 1)
|
|
72
|
+
self.alpha_list = list(alpha)
|
|
73
|
+
|
|
74
|
+
self.embedding = keras.layers.Embedding(
|
|
75
|
+
num_nodes,
|
|
76
|
+
embedding_dim,
|
|
77
|
+
embeddings_initializer=keras.initializers.GlorotUniform(),
|
|
78
|
+
)
|
|
79
|
+
self.convs = [LGConv(**kwargs) for _ in range(num_layers)]
|
|
80
|
+
|
|
81
|
+
def build(self, input_shape=None):
|
|
82
|
+
self.embedding.build((None,))
|
|
83
|
+
self.built = True
|
|
84
|
+
|
|
85
|
+
def reset_parameters(self):
|
|
86
|
+
r"""Resets all learnable parameters of the module."""
|
|
87
|
+
if self.embedding.built:
|
|
88
|
+
self.embedding.embeddings.assign(
|
|
89
|
+
keras.initializers.GlorotUniform()(shape=(self.num_nodes, self.embedding_dim))
|
|
90
|
+
)
|
|
91
|
+
for conv in self.convs:
|
|
92
|
+
if hasattr(conv, "reset_parameters"):
|
|
93
|
+
conv.reset_parameters()
|
|
94
|
+
|
|
95
|
+
def get_embedding(self, edge_index, edge_weight=None):
|
|
96
|
+
r"""Returns the embedding of nodes in the graph."""
|
|
97
|
+
if not self.embedding.built:
|
|
98
|
+
self.embedding.build((None,))
|
|
99
|
+
# Embedding weights: shape [num_nodes, embedding_dim]
|
|
100
|
+
x = self.embedding.weights[0]
|
|
101
|
+
out = x * self.alpha_list[0]
|
|
102
|
+
|
|
103
|
+
for i in range(self.num_layers):
|
|
104
|
+
x = self.convs[i](x, edge_index, edge_weight=edge_weight)
|
|
105
|
+
out = out + x * self.alpha_list[i + 1]
|
|
106
|
+
|
|
107
|
+
return out
|
|
108
|
+
|
|
109
|
+
def call(self, edge_index, edge_label_index=None, edge_weight=None):
|
|
110
|
+
r"""Computes rankings for pairs of nodes."""
|
|
111
|
+
if edge_label_index is None:
|
|
112
|
+
edge_label_index = edge_index
|
|
113
|
+
|
|
114
|
+
out = self.get_embedding(edge_index, edge_weight)
|
|
115
|
+
|
|
116
|
+
out_src = ops.take(out, edge_label_index[0], axis=0)
|
|
117
|
+
out_dst = ops.take(out, edge_label_index[1], axis=0)
|
|
118
|
+
|
|
119
|
+
return ops.sum(out_src * out_dst, axis=-1)
|
|
120
|
+
|
|
121
|
+
def predict_link(
|
|
122
|
+
self,
|
|
123
|
+
edge_index,
|
|
124
|
+
edge_label_index=None,
|
|
125
|
+
edge_weight=None,
|
|
126
|
+
prob: bool = False,
|
|
127
|
+
):
|
|
128
|
+
pred = ops.sigmoid(self(edge_index, edge_label_index, edge_weight))
|
|
129
|
+
return pred if prob else ops.round(pred)
|
|
130
|
+
|
|
131
|
+
def recommend(
|
|
132
|
+
self,
|
|
133
|
+
edge_index,
|
|
134
|
+
edge_weight=None,
|
|
135
|
+
src_index=None,
|
|
136
|
+
dst_index=None,
|
|
137
|
+
k: int = 1,
|
|
138
|
+
sorted: bool = True,
|
|
139
|
+
):
|
|
140
|
+
out = self.get_embedding(edge_index, edge_weight)
|
|
141
|
+
out_src = ops.take(out, src_index, axis=0) if src_index is not None else out
|
|
142
|
+
out_dst = ops.take(out, dst_index, axis=0) if dst_index is not None else out
|
|
143
|
+
|
|
144
|
+
pred = out_src @ ops.transpose(out_dst)
|
|
145
|
+
top_indices = ops.top_k(pred, k=k, sorted=sorted)[1]
|
|
146
|
+
|
|
147
|
+
if dst_index is not None:
|
|
148
|
+
top_indices = ops.take(dst_index, top_indices, axis=0)
|
|
149
|
+
|
|
150
|
+
return top_indices
|
|
151
|
+
|
|
152
|
+
def link_pred_loss(self, pred, edge_label):
|
|
153
|
+
loss_fn = keras.losses.BinaryCrossentropy(from_logits=True)
|
|
154
|
+
return loss_fn(edge_label, pred)
|
|
155
|
+
|
|
156
|
+
def recommendation_loss(
|
|
157
|
+
self,
|
|
158
|
+
pos_edge_rank,
|
|
159
|
+
neg_edge_rank,
|
|
160
|
+
node_id=None,
|
|
161
|
+
lambda_reg: float = 1e-4,
|
|
162
|
+
):
|
|
163
|
+
loss_fn = BPRLoss(lambda_reg)
|
|
164
|
+
emb = self.embedding.weights[0]
|
|
165
|
+
emb = emb if node_id is None else ops.take(emb, node_id, axis=0)
|
|
166
|
+
return loss_fn(pos_edge_rank, neg_edge_rank, emb)
|
|
167
|
+
|
|
168
|
+
def __repr__(self) -> str:
|
|
169
|
+
return (f'{self.__class__.__name__}({self.num_nodes}, '
|
|
170
|
+
f'{self.embedding_dim}, num_layers={self.num_layers})')
|
|
171
|
+
|
k3_node/models/linkx.py
ADDED
|
@@ -0,0 +1,181 @@
|
|
|
1
|
+
import math
|
|
2
|
+
from typing import Optional
|
|
3
|
+
|
|
4
|
+
import keras
|
|
5
|
+
from keras import ops
|
|
6
|
+
|
|
7
|
+
from k3_node.layers.conv.message_passing import MessagePassing
|
|
8
|
+
from k3_node.layers.conv.utils import scatter
|
|
9
|
+
from k3_node.layers.norm import BatchNorm
|
|
10
|
+
from k3_node.models.mlp import MLP
|
|
11
|
+
|
|
12
|
+
|
|
13
|
+
class SparseLinear(keras.layers.Layer):
|
|
14
|
+
r"""A sparse linear transformation operator computing :math:`\mathbf{A}\mathbf{W} + \mathbf{b}`.
|
|
15
|
+
|
|
16
|
+
Example:
|
|
17
|
+
```python
|
|
18
|
+
import numpy as np
|
|
19
|
+
from k3_node.models import SparseLinear
|
|
20
|
+
|
|
21
|
+
x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
|
|
22
|
+
edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
|
|
23
|
+
|
|
24
|
+
layer = SparseLinear(in_channels=10, out_channels=16) # in_channels = number of nodes
|
|
25
|
+
out = layer(edge_index) # multiplies the adjacency matrix with a learned weight matrix
|
|
26
|
+
print(tuple(out.shape)) # (10, 16)
|
|
27
|
+
```
|
|
28
|
+
"""
|
|
29
|
+
def __init__(self, in_channels: int, out_channels: int, bias: bool = True, **kwargs):
|
|
30
|
+
super().__init__(**kwargs)
|
|
31
|
+
self.in_channels = in_channels
|
|
32
|
+
self.out_channels = out_channels
|
|
33
|
+
self.use_bias = bias
|
|
34
|
+
|
|
35
|
+
self.weight = self.add_weight(
|
|
36
|
+
shape=(in_channels, out_channels),
|
|
37
|
+
initializer=keras.initializers.GlorotUniform(),
|
|
38
|
+
trainable=True,
|
|
39
|
+
name="weight",
|
|
40
|
+
)
|
|
41
|
+
if bias:
|
|
42
|
+
self.bias = self.add_weight(
|
|
43
|
+
shape=(out_channels,),
|
|
44
|
+
initializer="zeros",
|
|
45
|
+
trainable=True,
|
|
46
|
+
name="bias",
|
|
47
|
+
)
|
|
48
|
+
else:
|
|
49
|
+
self.bias = None
|
|
50
|
+
|
|
51
|
+
def build(self, input_shape=None):
|
|
52
|
+
self.built = True
|
|
53
|
+
|
|
54
|
+
def reset_parameters(self):
|
|
55
|
+
self.weight.assign(
|
|
56
|
+
keras.initializers.GlorotUniform()(shape=(self.in_channels, self.out_channels))
|
|
57
|
+
)
|
|
58
|
+
if self.use_bias and self.bias is not None:
|
|
59
|
+
self.bias.assign(ops.zeros((self.out_channels,)))
|
|
60
|
+
|
|
61
|
+
def call(self, edge_index, edge_weight=None):
|
|
62
|
+
row, col = edge_index[0], edge_index[1]
|
|
63
|
+
weight_j = ops.take(self.weight, row, axis=0)
|
|
64
|
+
if edge_weight is not None:
|
|
65
|
+
weight_j = ops.expand_dims(edge_weight, -1) * weight_j
|
|
66
|
+
|
|
67
|
+
out = scatter(weight_j, col, dim=0, dim_size=self.in_channels, reduce="sum")
|
|
68
|
+
if self.use_bias and self.bias is not None:
|
|
69
|
+
out = out + self.bias
|
|
70
|
+
return out
|
|
71
|
+
|
|
72
|
+
|
|
73
|
+
class LINKX(keras.Model):
|
|
74
|
+
r"""The LINKX model from the `"Large Scale Learning on Non-Homophilous
|
|
75
|
+
Graphs: New Benchmarks and Strong Simple Methods"
|
|
76
|
+
<https://arxiv.org/abs/2110.14446>`_ paper.
|
|
77
|
+
|
|
78
|
+
Args:
|
|
79
|
+
num_nodes (int): The number of nodes in the graph.
|
|
80
|
+
in_channels (int): Size of each input sample.
|
|
81
|
+
hidden_channels (int): Size of each hidden sample.
|
|
82
|
+
out_channels (int): Size of each output sample.
|
|
83
|
+
num_layers (int): Number of layers of :math:`\textrm{MLP}_{f}`.
|
|
84
|
+
num_edge_layers (int, optional): Number of layers of
|
|
85
|
+
:math:`\textrm{MLP}_{\mathbf{A}}`. (default: :obj:`1`)
|
|
86
|
+
num_node_layers (int, optional): Number of layers of
|
|
87
|
+
:math:`\textrm{MLP}_{\mathbf{X}}`. (default: :obj:`1`)
|
|
88
|
+
dropout (float, optional): Dropout probability. (default: :obj:`0.0`)
|
|
89
|
+
|
|
90
|
+
Example:
|
|
91
|
+
```python
|
|
92
|
+
import numpy as np
|
|
93
|
+
from k3_node.models import LINKX
|
|
94
|
+
|
|
95
|
+
x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
|
|
96
|
+
edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
|
|
97
|
+
|
|
98
|
+
# LINKX learns from node features and adjacency separately, which suits heterophilic graphs
|
|
99
|
+
model = LINKX(num_nodes=10, in_channels=8, hidden_channels=32, out_channels=4, num_layers=2)
|
|
100
|
+
out = model(x, edge_index)
|
|
101
|
+
print(tuple(out.shape)) # (10, 4)
|
|
102
|
+
```
|
|
103
|
+
"""
|
|
104
|
+
def __init__(
|
|
105
|
+
self,
|
|
106
|
+
num_nodes: int,
|
|
107
|
+
in_channels: int,
|
|
108
|
+
hidden_channels: int,
|
|
109
|
+
out_channels: int,
|
|
110
|
+
num_layers: int,
|
|
111
|
+
num_edge_layers: int = 1,
|
|
112
|
+
num_node_layers: int = 1,
|
|
113
|
+
dropout: float = 0.0,
|
|
114
|
+
**kwargs,
|
|
115
|
+
):
|
|
116
|
+
super().__init__(**kwargs)
|
|
117
|
+
|
|
118
|
+
self.num_nodes = num_nodes
|
|
119
|
+
self.in_channels = in_channels
|
|
120
|
+
self.hidden_channels = hidden_channels
|
|
121
|
+
self.out_channels = out_channels
|
|
122
|
+
self.num_edge_layers = num_edge_layers
|
|
123
|
+
|
|
124
|
+
self.edge_lin = SparseLinear(num_nodes, hidden_channels)
|
|
125
|
+
|
|
126
|
+
if num_edge_layers > 1:
|
|
127
|
+
self.edge_norm = BatchNorm(hidden_channels)
|
|
128
|
+
channels = [hidden_channels] * num_edge_layers
|
|
129
|
+
self.edge_mlp = MLP(channels, dropout=0.0, act_first=True)
|
|
130
|
+
else:
|
|
131
|
+
self.edge_norm = None
|
|
132
|
+
self.edge_mlp = None
|
|
133
|
+
|
|
134
|
+
channels = [in_channels] + [hidden_channels] * num_node_layers
|
|
135
|
+
self.node_mlp = MLP(channels, dropout=0.0, act_first=True)
|
|
136
|
+
|
|
137
|
+
self.cat_lin1 = keras.layers.Dense(hidden_channels)
|
|
138
|
+
self.cat_lin2 = keras.layers.Dense(hidden_channels)
|
|
139
|
+
|
|
140
|
+
channels = [hidden_channels] * num_layers + [out_channels]
|
|
141
|
+
self.final_mlp = MLP(channels, dropout=dropout, act_first=True)
|
|
142
|
+
|
|
143
|
+
def build(self, input_shape=None):
|
|
144
|
+
self.edge_lin.build()
|
|
145
|
+
self.built = True
|
|
146
|
+
|
|
147
|
+
def reset_parameters(self):
|
|
148
|
+
r"""Resets all learnable parameters of the module."""
|
|
149
|
+
self.edge_lin.reset_parameters()
|
|
150
|
+
if self.edge_norm is not None and hasattr(self.edge_norm, "reset_parameters"):
|
|
151
|
+
self.edge_norm.reset_parameters()
|
|
152
|
+
if self.edge_mlp is not None and hasattr(self.edge_mlp, "reset_parameters"):
|
|
153
|
+
self.edge_mlp.reset_parameters()
|
|
154
|
+
self.node_mlp.reset_parameters()
|
|
155
|
+
self.final_mlp.reset_parameters()
|
|
156
|
+
|
|
157
|
+
def call(self, x, edge_index=None, edge_weight=None, training=None):
|
|
158
|
+
if edge_index is None and isinstance(x, (tuple, list)):
|
|
159
|
+
if len(x) >= 2:
|
|
160
|
+
x, edge_index = x[0], x[1]
|
|
161
|
+
out = self.edge_lin(edge_index, edge_weight)
|
|
162
|
+
|
|
163
|
+
if self.edge_norm is not None and self.edge_mlp is not None:
|
|
164
|
+
out = ops.relu(out)
|
|
165
|
+
out = self.edge_norm(out, training=training)
|
|
166
|
+
out = self.edge_mlp(out, training=training)
|
|
167
|
+
|
|
168
|
+
out = out + self.cat_lin1(out)
|
|
169
|
+
|
|
170
|
+
if x is not None:
|
|
171
|
+
x = self.node_mlp(x, training=training)
|
|
172
|
+
out = out + x
|
|
173
|
+
out = out + self.cat_lin2(x)
|
|
174
|
+
|
|
175
|
+
return self.final_mlp(ops.relu(out), training=training)
|
|
176
|
+
|
|
177
|
+
def __repr__(self) -> str:
|
|
178
|
+
return (f'{self.__class__.__name__}(num_nodes={self.num_nodes}, '
|
|
179
|
+
f'in_channels={self.in_channels}, '
|
|
180
|
+
f'out_channels={self.out_channels})')
|
|
181
|
+
|