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,96 @@
|
|
|
1
|
+
import numpy as np
|
|
2
|
+
import pytest
|
|
3
|
+
|
|
4
|
+
try:
|
|
5
|
+
import torch
|
|
6
|
+
except ImportError:
|
|
7
|
+
torch = None
|
|
8
|
+
|
|
9
|
+
from k3_node.data import Data
|
|
10
|
+
from k3_node.loader import (
|
|
11
|
+
CachedLoader,
|
|
12
|
+
DataLoader,
|
|
13
|
+
DynamicBatchSampler,
|
|
14
|
+
ImbalancedSampler,
|
|
15
|
+
PrefetchLoader,
|
|
16
|
+
ZipLoader,
|
|
17
|
+
)
|
|
18
|
+
|
|
19
|
+
|
|
20
|
+
def create_dummy_data(num_nodes=4, num_features=8, label=0):
|
|
21
|
+
x = np.random.randn(num_nodes, num_features).astype(np.float32)
|
|
22
|
+
edge_index = np.array([[0, 1, 2, 3], [1, 2, 3, 0]], dtype=np.int64)
|
|
23
|
+
y = np.array([label], dtype=np.int64)
|
|
24
|
+
if torch is not None:
|
|
25
|
+
x = torch.from_numpy(x)
|
|
26
|
+
edge_index = torch.from_numpy(edge_index)
|
|
27
|
+
y = torch.from_numpy(y)
|
|
28
|
+
return Data(x=x, edge_index=edge_index, y=y)
|
|
29
|
+
|
|
30
|
+
|
|
31
|
+
def test_dynamic_batch_sampler():
|
|
32
|
+
dataset = [create_dummy_data(num_nodes=i + 2) for i in range(5)]
|
|
33
|
+
sampler = DynamicBatchSampler(dataset, max_num=15, mode='node')
|
|
34
|
+
|
|
35
|
+
loader = DataLoader(dataset, batch_sampler=sampler)
|
|
36
|
+
count = 0
|
|
37
|
+
for batch in loader:
|
|
38
|
+
assert batch.num_nodes <= 15
|
|
39
|
+
count += 1
|
|
40
|
+
assert count > 0
|
|
41
|
+
|
|
42
|
+
with pytest.raises(ValueError, match="length of 'DynamicBatchSampler'"):
|
|
43
|
+
len(sampler)
|
|
44
|
+
assert len(DynamicBatchSampler(dataset, max_num=15, num_steps=2)) == 2
|
|
45
|
+
|
|
46
|
+
|
|
47
|
+
def test_imbalanced_sampler():
|
|
48
|
+
# 8 samples with label 0, 2 samples with label 1
|
|
49
|
+
dataset = [create_dummy_data(label=0) for _ in range(8)] + [create_dummy_data(label=1) for _ in range(2)]
|
|
50
|
+
sampler = ImbalancedSampler(dataset, num_samples=10)
|
|
51
|
+
|
|
52
|
+
sampled_indices = list(sampler)
|
|
53
|
+
assert len(sampled_indices) == 10
|
|
54
|
+
|
|
55
|
+
loader = DataLoader(dataset, batch_size=5, sampler=sampler)
|
|
56
|
+
batches = list(loader)
|
|
57
|
+
assert len(batches) == 2
|
|
58
|
+
|
|
59
|
+
|
|
60
|
+
def test_zip_loader():
|
|
61
|
+
from k3_node.loader import NeighborLoader
|
|
62
|
+
data = create_dummy_data(num_nodes=10)
|
|
63
|
+
|
|
64
|
+
loader1 = NeighborLoader(data, num_neighbors=[2], input_nodes=[0, 1, 2, 3])
|
|
65
|
+
loader2 = NeighborLoader(data, num_neighbors=[2], input_nodes=[4, 5, 6, 7])
|
|
66
|
+
|
|
67
|
+
zip_loader = ZipLoader([loader1, loader2], batch_size=2)
|
|
68
|
+
for batch1, batch2 in zip_loader:
|
|
69
|
+
assert batch1.batch_size == 2
|
|
70
|
+
assert batch2.batch_size == 2
|
|
71
|
+
|
|
72
|
+
|
|
73
|
+
def test_cached_loader():
|
|
74
|
+
dataset = [create_dummy_data(num_nodes=3) for _ in range(4)]
|
|
75
|
+
loader = DataLoader(dataset, batch_size=2)
|
|
76
|
+
|
|
77
|
+
cached_loader = CachedLoader(loader)
|
|
78
|
+
epoch1 = list(cached_loader)
|
|
79
|
+
epoch2 = list(cached_loader)
|
|
80
|
+
|
|
81
|
+
assert len(epoch1) == 2
|
|
82
|
+
assert len(epoch2) == 2
|
|
83
|
+
assert len(cached_loader) == 2
|
|
84
|
+
|
|
85
|
+
cached_loader.clear()
|
|
86
|
+
assert len(cached_loader._cache) == 0
|
|
87
|
+
|
|
88
|
+
|
|
89
|
+
def test_prefetch_loader():
|
|
90
|
+
dataset = [create_dummy_data(num_nodes=3) for _ in range(4)]
|
|
91
|
+
loader = DataLoader(dataset, batch_size=2)
|
|
92
|
+
|
|
93
|
+
prefetch = PrefetchLoader(loader)
|
|
94
|
+
batches = list(prefetch)
|
|
95
|
+
assert len(batches) == 2
|
|
96
|
+
assert len(prefetch) == 2
|
|
@@ -0,0 +1,89 @@
|
|
|
1
|
+
import numpy as np
|
|
2
|
+
try:
|
|
3
|
+
import torch
|
|
4
|
+
except ImportError:
|
|
5
|
+
torch = None
|
|
6
|
+
|
|
7
|
+
from k3_node.data import Data
|
|
8
|
+
from k3_node.loader import (
|
|
9
|
+
ClusterData,
|
|
10
|
+
ClusterLoader,
|
|
11
|
+
GraphSAINTEdgeSampler,
|
|
12
|
+
GraphSAINTNodeSampler,
|
|
13
|
+
GraphSAINTRandomWalkSampler,
|
|
14
|
+
NeighborSampler,
|
|
15
|
+
ShaDowKHopSampler,
|
|
16
|
+
)
|
|
17
|
+
|
|
18
|
+
|
|
19
|
+
def get_graph(num_nodes=12):
|
|
20
|
+
# Circular ring graph
|
|
21
|
+
row = np.arange(num_nodes)
|
|
22
|
+
col = (row + 1) % num_nodes
|
|
23
|
+
edge_index = np.stack([row, col], axis=0).astype(np.int64)
|
|
24
|
+
x = np.random.randn(num_nodes, 8).astype(np.float32)
|
|
25
|
+
y = np.random.randint(0, 2, size=(num_nodes,)).astype(np.int64)
|
|
26
|
+
|
|
27
|
+
if torch is not None:
|
|
28
|
+
edge_index = torch.from_numpy(edge_index)
|
|
29
|
+
x = torch.from_numpy(x)
|
|
30
|
+
y = torch.from_numpy(y)
|
|
31
|
+
|
|
32
|
+
return Data(x=x, edge_index=edge_index, y=y)
|
|
33
|
+
|
|
34
|
+
|
|
35
|
+
def test_cluster_loader():
|
|
36
|
+
data = get_graph(num_nodes=12)
|
|
37
|
+
cluster_data = ClusterData(data, num_parts=3)
|
|
38
|
+
assert len(cluster_data) == 3
|
|
39
|
+
|
|
40
|
+
sub0 = cluster_data[0]
|
|
41
|
+
assert isinstance(sub0, Data)
|
|
42
|
+
assert sub0.num_nodes > 0
|
|
43
|
+
|
|
44
|
+
loader = ClusterLoader(cluster_data, batch_size=2, shuffle=False)
|
|
45
|
+
batches = list(loader)
|
|
46
|
+
assert len(batches) == 2 # 3 parts with batch_size 2 => 2 batches
|
|
47
|
+
|
|
48
|
+
|
|
49
|
+
def test_graph_saint():
|
|
50
|
+
data = get_graph(num_nodes=12)
|
|
51
|
+
|
|
52
|
+
# Node sampler
|
|
53
|
+
node_sampler = GraphSAINTNodeSampler(data, batch_size=4, num_steps=2, sample_coverage=1)
|
|
54
|
+
batch = next(iter(node_sampler))
|
|
55
|
+
assert isinstance(batch, Data)
|
|
56
|
+
assert hasattr(batch, 'node_norm')
|
|
57
|
+
|
|
58
|
+
# Edge sampler
|
|
59
|
+
edge_sampler = GraphSAINTEdgeSampler(data, batch_size=4, num_steps=2)
|
|
60
|
+
batch = next(iter(edge_sampler))
|
|
61
|
+
assert isinstance(batch, Data)
|
|
62
|
+
|
|
63
|
+
# Random walk sampler
|
|
64
|
+
rw_sampler = GraphSAINTRandomWalkSampler(data, batch_size=4, walk_length=2, num_steps=2)
|
|
65
|
+
batch = next(iter(rw_sampler))
|
|
66
|
+
assert isinstance(batch, Data)
|
|
67
|
+
|
|
68
|
+
|
|
69
|
+
def test_shadow_k_hop_sampler():
|
|
70
|
+
data = get_graph(num_nodes=12)
|
|
71
|
+
loader = ShaDowKHopSampler(data, depth=2, num_neighbors=2, batch_size=3)
|
|
72
|
+
|
|
73
|
+
batch = next(iter(loader))
|
|
74
|
+
assert batch.num_graphs == 3
|
|
75
|
+
assert hasattr(batch, 'root_n_id')
|
|
76
|
+
|
|
77
|
+
|
|
78
|
+
def test_legacy_neighbor_sampler():
|
|
79
|
+
data = get_graph(num_nodes=12)
|
|
80
|
+
loader = NeighborSampler(data.edge_index, sizes=[2, 2], batch_size=3)
|
|
81
|
+
|
|
82
|
+
batch_size, n_id, adjs = next(iter(loader))
|
|
83
|
+
assert batch_size == 3
|
|
84
|
+
assert len(n_id) >= 3
|
|
85
|
+
assert len(adjs) == 2
|
|
86
|
+
for adj in adjs:
|
|
87
|
+
assert hasattr(adj, 'edge_index')
|
|
88
|
+
assert hasattr(adj, 'size')
|
|
89
|
+
|
k3_node/loader/utils.py
ADDED
|
@@ -0,0 +1,232 @@
|
|
|
1
|
+
import copy
|
|
2
|
+
from typing import Any, Dict, Optional, Tuple, Union
|
|
3
|
+
|
|
4
|
+
import numpy as np
|
|
5
|
+
try:
|
|
6
|
+
import torch
|
|
7
|
+
except ImportError:
|
|
8
|
+
torch = None
|
|
9
|
+
Tensor = type(None)
|
|
10
|
+
|
|
11
|
+
from k3_node.data import Data, HeteroData
|
|
12
|
+
from k3_node.data.storage import NodeStorage, EdgeStorage
|
|
13
|
+
|
|
14
|
+
|
|
15
|
+
def to_numpy_or_tensor(x):
|
|
16
|
+
if torch is not None and isinstance(x, torch.Tensor):
|
|
17
|
+
return x
|
|
18
|
+
if isinstance(x, np.ndarray):
|
|
19
|
+
return x
|
|
20
|
+
if hasattr(x, '__array__'):
|
|
21
|
+
return np.asarray(x)
|
|
22
|
+
return x
|
|
23
|
+
|
|
24
|
+
|
|
25
|
+
def to_numpy(x, dtype=None):
|
|
26
|
+
r"""Converts any tensor (PyTorch, TensorFlow, JAX) or array to a CPU NumPy ndarray."""
|
|
27
|
+
if hasattr(x, "cpu"):
|
|
28
|
+
x = x.cpu()
|
|
29
|
+
if hasattr(x, "detach"):
|
|
30
|
+
x = x.detach()
|
|
31
|
+
if hasattr(x, "numpy") and callable(x.numpy):
|
|
32
|
+
x = x.numpy()
|
|
33
|
+
return np.asarray(x, dtype=dtype)
|
|
34
|
+
|
|
35
|
+
|
|
36
|
+
def index_select(value: Any, index: Any, dim: int = 0) -> Any:
|
|
37
|
+
r"""Indexes the :obj:`value` tensor along dimension :obj:`dim` using the
|
|
38
|
+
entries in :obj:`index`. Supports PyTorch, TensorFlow, JAX, and NumPy.
|
|
39
|
+
"""
|
|
40
|
+
if torch is not None and isinstance(value, torch.Tensor):
|
|
41
|
+
if not isinstance(index, torch.Tensor):
|
|
42
|
+
index = torch.as_tensor(index, dtype=torch.long, device=value.device)
|
|
43
|
+
else:
|
|
44
|
+
index = index.to(dtype=torch.long, device=value.device)
|
|
45
|
+
return torch.index_select(value, dim, index)
|
|
46
|
+
|
|
47
|
+
# NumPy / Keras / JAX / TensorFlow array:
|
|
48
|
+
if hasattr(value, 'numpy'):
|
|
49
|
+
is_tf_or_jax = True
|
|
50
|
+
np_val = value.numpy()
|
|
51
|
+
elif isinstance(value, np.ndarray):
|
|
52
|
+
is_tf_or_jax = False
|
|
53
|
+
np_val = value
|
|
54
|
+
elif hasattr(value, '__array__'):
|
|
55
|
+
is_tf_or_jax = False
|
|
56
|
+
np_val = np.asarray(value)
|
|
57
|
+
else:
|
|
58
|
+
return value
|
|
59
|
+
|
|
60
|
+
if torch is not None and isinstance(index, torch.Tensor):
|
|
61
|
+
np_idx = index.cpu().numpy()
|
|
62
|
+
else:
|
|
63
|
+
np_idx = np.asarray(index, dtype=np.int64)
|
|
64
|
+
|
|
65
|
+
res = np.take(np_val, np_idx, axis=dim)
|
|
66
|
+
if is_tf_or_jax:
|
|
67
|
+
import keras
|
|
68
|
+
return keras.ops.convert_to_tensor(res)
|
|
69
|
+
return res
|
|
70
|
+
|
|
71
|
+
|
|
72
|
+
def filter_node_store_(store: NodeStorage, out_store: NodeStorage, index: Any):
|
|
73
|
+
for key, value in store.items():
|
|
74
|
+
if key == 'num_nodes':
|
|
75
|
+
numel = index.numel() if hasattr(index, 'numel') else len(index)
|
|
76
|
+
out_store.num_nodes = numel
|
|
77
|
+
elif store.is_node_attr(key):
|
|
78
|
+
dim = 0
|
|
79
|
+
if hasattr(store, '_parent') and store._parent() is not None:
|
|
80
|
+
dim = store._parent().__cat_dim__(key, value, store)
|
|
81
|
+
out_store[key] = index_select(value, index, dim=dim)
|
|
82
|
+
|
|
83
|
+
|
|
84
|
+
def filter_edge_store_(
|
|
85
|
+
store: EdgeStorage,
|
|
86
|
+
out_store: EdgeStorage,
|
|
87
|
+
row: Any,
|
|
88
|
+
col: Any,
|
|
89
|
+
index: Optional[Any],
|
|
90
|
+
perm: Optional[Any] = None,
|
|
91
|
+
):
|
|
92
|
+
for key, value in store.items():
|
|
93
|
+
if key == 'edge_index':
|
|
94
|
+
if torch is not None and isinstance(row, torch.Tensor):
|
|
95
|
+
edge_index = torch.stack([row, col], dim=0)
|
|
96
|
+
else:
|
|
97
|
+
edge_index = np.stack([np.asarray(row), np.asarray(col)], axis=0)
|
|
98
|
+
out_store.edge_index = edge_index
|
|
99
|
+
elif store.is_edge_attr(key):
|
|
100
|
+
if index is None:
|
|
101
|
+
out_store[key] = None
|
|
102
|
+
continue
|
|
103
|
+
dim = 0
|
|
104
|
+
if hasattr(store, '_parent') and store._parent() is not None:
|
|
105
|
+
dim = store._parent().__cat_dim__(key, value, store)
|
|
106
|
+
if perm is None:
|
|
107
|
+
out_store[key] = index_select(value, index, dim=dim)
|
|
108
|
+
else:
|
|
109
|
+
sel_idx = perm[index] if hasattr(perm, '__getitem__') else index
|
|
110
|
+
out_store[key] = index_select(value, sel_idx, dim=dim)
|
|
111
|
+
|
|
112
|
+
|
|
113
|
+
def filter_data(data: Data, node: Any, row: Any, col: Any, edge: Optional[Any] = None, perm: Optional[Any] = None) -> Data:
|
|
114
|
+
out = copy.copy(data)
|
|
115
|
+
out._store = copy.copy(data._store)
|
|
116
|
+
filter_node_store_(data._store, out._store, node)
|
|
117
|
+
filter_edge_store_(data._store, out._store, row, col, edge, perm)
|
|
118
|
+
return out
|
|
119
|
+
|
|
120
|
+
|
|
121
|
+
def filter_hetero_data(
|
|
122
|
+
data: HeteroData,
|
|
123
|
+
node_dict: Dict[str, Any],
|
|
124
|
+
row_dict: Dict[Tuple[str, str, str], Any],
|
|
125
|
+
col_dict: Dict[Tuple[str, str, str], Any],
|
|
126
|
+
edge_dict: Dict[Tuple[str, str, str], Optional[Any]],
|
|
127
|
+
perm_dict: Optional[Dict[Tuple[str, str, str], Optional[Any]]] = None,
|
|
128
|
+
) -> HeteroData:
|
|
129
|
+
out = copy.copy(data)
|
|
130
|
+
out._node_store_dict = {k: copy.copy(v) for k, v in data._node_store_dict.items()}
|
|
131
|
+
out._edge_store_dict = {k: copy.copy(v) for k, v in data._edge_store_dict.items()}
|
|
132
|
+
|
|
133
|
+
for node_type in out.node_types:
|
|
134
|
+
if node_type not in node_dict:
|
|
135
|
+
node_dict[node_type] = torch.empty(0, dtype=torch.long) if torch is not None else np.empty(0, dtype=np.int64)
|
|
136
|
+
filter_node_store_(data[node_type], out[node_type], node_dict[node_type])
|
|
137
|
+
|
|
138
|
+
for edge_type in out.edge_types:
|
|
139
|
+
canonical = data._to_canonical(*edge_type) if hasattr(data, '_to_canonical') else edge_type
|
|
140
|
+
if canonical not in row_dict:
|
|
141
|
+
empty_arr = torch.empty(0, dtype=torch.long) if torch is not None else np.empty(0, dtype=np.int64)
|
|
142
|
+
row_dict[canonical] = empty_arr
|
|
143
|
+
col_dict[canonical] = empty_arr
|
|
144
|
+
edge_dict[canonical] = empty_arr
|
|
145
|
+
|
|
146
|
+
filter_edge_store_(
|
|
147
|
+
data[edge_type],
|
|
148
|
+
out[edge_type],
|
|
149
|
+
row_dict[canonical],
|
|
150
|
+
col_dict[canonical],
|
|
151
|
+
edge_dict[canonical],
|
|
152
|
+
perm_dict.get(canonical, None) if perm_dict else None,
|
|
153
|
+
)
|
|
154
|
+
|
|
155
|
+
return out
|
|
156
|
+
|
|
157
|
+
|
|
158
|
+
def get_input_nodes(
|
|
159
|
+
data: Union[Data, HeteroData],
|
|
160
|
+
input_nodes: Any,
|
|
161
|
+
input_id: Optional[Any] = None,
|
|
162
|
+
) -> Tuple[Optional[str], Any, Optional[Any]]:
|
|
163
|
+
def to_index(nodes, in_id):
|
|
164
|
+
if torch is not None and isinstance(nodes, torch.Tensor):
|
|
165
|
+
if nodes.dtype == torch.bool:
|
|
166
|
+
nodes = nodes.nonzero(as_tuple=False).view(-1)
|
|
167
|
+
in_id = nodes if in_id is None else in_id
|
|
168
|
+
return nodes, in_id
|
|
169
|
+
if isinstance(nodes, np.ndarray) and nodes.dtype == bool:
|
|
170
|
+
nodes = np.nonzero(nodes)[0]
|
|
171
|
+
in_id = nodes if in_id is None else in_id
|
|
172
|
+
return nodes, in_id
|
|
173
|
+
if torch is not None and not isinstance(nodes, torch.Tensor):
|
|
174
|
+
nodes = torch.tensor(nodes, dtype=torch.long)
|
|
175
|
+
elif torch is None:
|
|
176
|
+
nodes = np.asarray(nodes, dtype=np.int64)
|
|
177
|
+
return nodes, in_id
|
|
178
|
+
|
|
179
|
+
if isinstance(data, Data):
|
|
180
|
+
if input_nodes is None:
|
|
181
|
+
nodes = torch.arange(data.num_nodes) if torch is not None else np.arange(data.num_nodes)
|
|
182
|
+
return None, nodes, None
|
|
183
|
+
return None, *to_index(input_nodes, input_id)
|
|
184
|
+
|
|
185
|
+
elif isinstance(data, HeteroData):
|
|
186
|
+
assert input_nodes is not None
|
|
187
|
+
if isinstance(input_nodes, str):
|
|
188
|
+
num_nodes = data[input_nodes].num_nodes
|
|
189
|
+
nodes = torch.arange(num_nodes) if torch is not None else np.arange(num_nodes)
|
|
190
|
+
return input_nodes, nodes, None
|
|
191
|
+
|
|
192
|
+
assert isinstance(input_nodes, (list, tuple)) and len(input_nodes) == 2
|
|
193
|
+
node_type, input_nodes = input_nodes
|
|
194
|
+
if input_nodes is None:
|
|
195
|
+
num_nodes = data[node_type].num_nodes
|
|
196
|
+
nodes = torch.arange(num_nodes) if torch is not None else np.arange(num_nodes)
|
|
197
|
+
return node_type, nodes, None
|
|
198
|
+
return node_type, *to_index(input_nodes, input_id)
|
|
199
|
+
|
|
200
|
+
raise TypeError(f"Invalid data type: {type(data)}")
|
|
201
|
+
|
|
202
|
+
|
|
203
|
+
def get_edge_label_index(
|
|
204
|
+
data: Union[Data, HeteroData],
|
|
205
|
+
edge_label_index: Any,
|
|
206
|
+
) -> Tuple[Optional[Tuple[str, str, str]], Any]:
|
|
207
|
+
if isinstance(data, Data):
|
|
208
|
+
if edge_label_index is None:
|
|
209
|
+
return None, data.edge_index
|
|
210
|
+
return None, edge_label_index
|
|
211
|
+
|
|
212
|
+
if isinstance(data, HeteroData):
|
|
213
|
+
assert edge_label_index is not None
|
|
214
|
+
if isinstance(edge_label_index, (list, tuple)) and len(edge_label_index) == 3 and isinstance(edge_label_index[0], str):
|
|
215
|
+
edge_type = data._to_canonical(*edge_label_index)
|
|
216
|
+
return edge_type, data[edge_type].edge_index
|
|
217
|
+
|
|
218
|
+
assert isinstance(edge_label_index, (list, tuple)) and len(edge_label_index) == 2
|
|
219
|
+
edge_type, edge_index = edge_label_index
|
|
220
|
+
edge_type = data._to_canonical(*edge_type)
|
|
221
|
+
if edge_index is None:
|
|
222
|
+
return edge_type, data[edge_type].edge_index
|
|
223
|
+
return edge_type, edge_index
|
|
224
|
+
|
|
225
|
+
raise TypeError(f"Invalid data type: {type(data)}")
|
|
226
|
+
|
|
227
|
+
|
|
228
|
+
def infer_filter_per_worker(data: Any) -> bool:
|
|
229
|
+
out = True
|
|
230
|
+
if hasattr(data, 'is_cuda') and data.is_cuda:
|
|
231
|
+
out = False
|
|
232
|
+
return out
|
|
@@ -0,0 +1,88 @@
|
|
|
1
|
+
from typing import Any, Iterator, List, Optional, Tuple, Union
|
|
2
|
+
|
|
3
|
+
try:
|
|
4
|
+
import torch
|
|
5
|
+
from torch import Tensor
|
|
6
|
+
BaseDataLoader = torch.utils.data.DataLoader
|
|
7
|
+
except ImportError:
|
|
8
|
+
torch = None
|
|
9
|
+
Tensor = type(None)
|
|
10
|
+
BaseDataLoader = object
|
|
11
|
+
|
|
12
|
+
from k3_node.data import Data, HeteroData
|
|
13
|
+
from k3_node.loader.base import DataLoaderIterator
|
|
14
|
+
from k3_node.loader.utils import infer_filter_per_worker
|
|
15
|
+
|
|
16
|
+
|
|
17
|
+
class ZipLoader(BaseDataLoader):
|
|
18
|
+
r"""A loader that returns a tuple of data objects by sampling from multiple
|
|
19
|
+
loader instances.
|
|
20
|
+
|
|
21
|
+
Args:
|
|
22
|
+
loaders (List[Any]): The loader instances.
|
|
23
|
+
filter_per_worker (bool, optional): If set to :obj:`True`, will filter
|
|
24
|
+
the returned data in each worker's subprocess. (default: :obj:`None`)
|
|
25
|
+
**kwargs (optional): Additional arguments of :class:`torch.utils.data.DataLoader`.
|
|
26
|
+
"""
|
|
27
|
+
def __init__(
|
|
28
|
+
self,
|
|
29
|
+
loaders: List[Any],
|
|
30
|
+
filter_per_worker: Optional[bool] = None,
|
|
31
|
+
**kwargs,
|
|
32
|
+
):
|
|
33
|
+
if filter_per_worker is None:
|
|
34
|
+
first_data = getattr(loaders[0], 'data', None)
|
|
35
|
+
filter_per_worker = infer_filter_per_worker(first_data) if first_data is not None else True
|
|
36
|
+
|
|
37
|
+
kwargs.pop('dataset', None)
|
|
38
|
+
kwargs.pop('collate_fn', None)
|
|
39
|
+
|
|
40
|
+
for loader in loaders:
|
|
41
|
+
if not callable(getattr(loader, 'collate_fn', None)):
|
|
42
|
+
raise ValueError(f"'{loader.__class__.__name__}' does not have a 'collate_fn' method")
|
|
43
|
+
if not callable(getattr(loader, 'filter_fn', None)):
|
|
44
|
+
raise ValueError(f"'{loader.__class__.__name__}' does not have a 'filter_fn' method")
|
|
45
|
+
loader.filter_per_worker = filter_per_worker
|
|
46
|
+
|
|
47
|
+
lens = []
|
|
48
|
+
for loader in loaders:
|
|
49
|
+
if hasattr(loader, 'dataset'):
|
|
50
|
+
lens.append(len(loader.dataset))
|
|
51
|
+
elif hasattr(loader, '__len__'):
|
|
52
|
+
lens.append(len(loader))
|
|
53
|
+
else:
|
|
54
|
+
lens.append(0)
|
|
55
|
+
|
|
56
|
+
iterator = range(min(lens) if lens else 0)
|
|
57
|
+
|
|
58
|
+
self.loaders = loaders
|
|
59
|
+
self.filter_per_worker = filter_per_worker
|
|
60
|
+
|
|
61
|
+
if torch is not None:
|
|
62
|
+
super().__init__(iterator, collate_fn=self.collate_fn, **kwargs)
|
|
63
|
+
else:
|
|
64
|
+
self.dataset = iterator
|
|
65
|
+
self.collate_fn = self.collate_fn
|
|
66
|
+
|
|
67
|
+
def __call__(self, index: Union[Tensor, List[int]]) -> Union[Tuple[Data, ...], Tuple[HeteroData, ...]]:
|
|
68
|
+
out = self.collate_fn(index)
|
|
69
|
+
if not self.filter_per_worker:
|
|
70
|
+
out = self.filter_fn(out)
|
|
71
|
+
return out
|
|
72
|
+
|
|
73
|
+
def collate_fn(self, index: List[int]) -> Tuple[Any, ...]:
|
|
74
|
+
if torch is not None and not isinstance(index, Tensor):
|
|
75
|
+
index = torch.tensor(index, dtype=torch.long)
|
|
76
|
+
return tuple(loader.collate_fn(index) for loader in self.loaders)
|
|
77
|
+
|
|
78
|
+
def filter_fn(self, outs: Tuple[Any, ...]) -> Tuple[Union[Data, HeteroData], ...]:
|
|
79
|
+
return tuple(loader.filter_fn(v) for loader, v in zip(self.loaders, outs))
|
|
80
|
+
|
|
81
|
+
def _get_iterator(self) -> Iterator:
|
|
82
|
+
if self.filter_per_worker:
|
|
83
|
+
return super()._get_iterator()
|
|
84
|
+
return DataLoaderIterator(super()._get_iterator(), self.filter_fn)
|
|
85
|
+
|
|
86
|
+
def __repr__(self) -> str:
|
|
87
|
+
return f'{self.__class__.__name__}(loaders={self.loaders})'
|
|
88
|
+
|
k3_node/metrics.py
ADDED
|
@@ -0,0 +1,94 @@
|
|
|
1
|
+
"""Metrics for graph learning tasks."""
|
|
2
|
+
import numpy as np
|
|
3
|
+
from typing import Optional
|
|
4
|
+
|
|
5
|
+
import keras
|
|
6
|
+
from keras import ops
|
|
7
|
+
|
|
8
|
+
|
|
9
|
+
@keras.saving.register_keras_serializable(package="k3_node")
|
|
10
|
+
class F1Score(keras.metrics.F1Score):
|
|
11
|
+
r"""F1 score that also accepts logits, e.g. for multi-label node classification.
|
|
12
|
+
|
|
13
|
+
Same as :class:`keras.metrics.F1Score`, plus ``from_logits``: when :obj:`True`, a sigmoid is
|
|
14
|
+
applied to the predictions first, so models can output logits (as used with
|
|
15
|
+
``BinaryCrossentropy(from_logits=True)``).
|
|
16
|
+
|
|
17
|
+
Example:
|
|
18
|
+
```python
|
|
19
|
+
import numpy as np
|
|
20
|
+
from k3_node.metrics import F1Score
|
|
21
|
+
|
|
22
|
+
f1 = F1Score(average="micro", from_logits=True)
|
|
23
|
+
y_true = np.array([[1, 0, 1], [0, 1, 0]], dtype="float32")
|
|
24
|
+
logits = np.array([[2.0, -1.0, 0.5], [-3.0, 1.5, 0.2]], dtype="float32")
|
|
25
|
+
f1.update_state(y_true, logits)
|
|
26
|
+
print(round(float(f1.result()), 2)) # 0.86
|
|
27
|
+
```
|
|
28
|
+
"""
|
|
29
|
+
|
|
30
|
+
def __init__(self, average=None, threshold=0.5, from_logits=False, name="f1_score", dtype=None):
|
|
31
|
+
super().__init__(average=average, threshold=threshold, name=name, dtype=dtype)
|
|
32
|
+
self.from_logits = from_logits
|
|
33
|
+
|
|
34
|
+
def update_state(self, y_true, y_pred, sample_weight=None):
|
|
35
|
+
if self.from_logits:
|
|
36
|
+
y_pred = ops.sigmoid(y_pred)
|
|
37
|
+
return super().update_state(y_true, y_pred, sample_weight=sample_weight)
|
|
38
|
+
|
|
39
|
+
def get_config(self):
|
|
40
|
+
return {**super().get_config(), "from_logits": self.from_logits}
|
|
41
|
+
|
|
42
|
+
|
|
43
|
+
def precision_recall_at_k(user_emb, item_emb, train_edge_index, test_edge_index, k: int = 20,
|
|
44
|
+
num_users: Optional[int] = None, batch_size: int = 8192):
|
|
45
|
+
r"""Top-``k`` recommendation quality, as in PyG's LightGCN example.
|
|
46
|
+
|
|
47
|
+
Every user's items are ranked by the dot product of the embeddings, excluding the items the
|
|
48
|
+
user interacted with during training. Returns the precision@k and recall@k averaged over the
|
|
49
|
+
users with at least one test interaction.
|
|
50
|
+
|
|
51
|
+
Args:
|
|
52
|
+
user_emb: User embeddings ``[num_users, dim]``.
|
|
53
|
+
item_emb: Item embeddings ``[num_items, dim]``.
|
|
54
|
+
train_edge_index: Training interactions ``(user, item)``; item ids may be offset by
|
|
55
|
+
``num_users`` (as in a homogeneous user-item graph).
|
|
56
|
+
test_edge_index: Test interactions ``(user, item)``, with the same numbering.
|
|
57
|
+
k (int): The number of recommendations per user. (default: ``20``)
|
|
58
|
+
num_users (int, optional): The item-id offset. (default: ``len(user_emb)``)
|
|
59
|
+
batch_size (int): Users scored at once. (default: ``8192``)
|
|
60
|
+
|
|
61
|
+
Example:
|
|
62
|
+
```python
|
|
63
|
+
import numpy as np
|
|
64
|
+
from k3_node.metrics import precision_recall_at_k
|
|
65
|
+
|
|
66
|
+
users, items = np.eye(2, 4, dtype="float32"), np.eye(3, 4, dtype="float32") # toy embeddings
|
|
67
|
+
train = np.array([[0], [2]]) # user 0 already has item 2 (id 2 + num_users = 4 in the graph)
|
|
68
|
+
test = np.array([[0, 1], [0 + 2, 1 + 2]]) # user 0 likes item 0, user 1 likes item 1
|
|
69
|
+
print(precision_recall_at_k(users, items, train + [[0], [2]], test, k=1)) # (1.0, 1.0)
|
|
70
|
+
```
|
|
71
|
+
"""
|
|
72
|
+
from keras import ops as _ops
|
|
73
|
+
|
|
74
|
+
user_emb, item_emb = (np.asarray(_ops.convert_to_numpy(e)) for e in (user_emb, item_emb))
|
|
75
|
+
num_users = len(user_emb) if num_users is None else num_users
|
|
76
|
+
train = np.asarray(_ops.convert_to_numpy(train_edge_index)).astype(np.int64)
|
|
77
|
+
test = np.asarray(_ops.convert_to_numpy(test_edge_index)).astype(np.int64)
|
|
78
|
+
train = train[:, train[0] < num_users]
|
|
79
|
+
precision = recall = examples = 0.0
|
|
80
|
+
for start in range(0, len(user_emb), batch_size):
|
|
81
|
+
end = min(start + batch_size, len(user_emb))
|
|
82
|
+
logits = user_emb[start:end] @ item_emb.T
|
|
83
|
+
m = (train[0] >= start) & (train[0] < end)
|
|
84
|
+
logits[train[0, m] - start, train[1, m] - num_users] = -np.inf # skip already known items
|
|
85
|
+
truth = np.zeros_like(logits, dtype=bool)
|
|
86
|
+
m = (test[0] >= start) & (test[0] < end)
|
|
87
|
+
truth[test[0, m] - start, test[1, m] - num_users] = True
|
|
88
|
+
count = truth.sum(axis=1)
|
|
89
|
+
top = np.argpartition(-logits, k - 1, axis=1)[:, :k]
|
|
90
|
+
hits = np.take_along_axis(truth, top, axis=1).sum(axis=1)
|
|
91
|
+
precision += float((hits / k)[count > 0].sum())
|
|
92
|
+
recall += float((hits / np.maximum(count, 1e-6))[count > 0].sum())
|
|
93
|
+
examples += int((count > 0).sum())
|
|
94
|
+
return precision / examples, recall / examples
|