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,162 @@
|
|
|
1
|
+
"""Verbalization and LLM prompt formatting for GraphRAG."""
|
|
2
|
+
|
|
3
|
+
from typing import Dict, List, Literal, Optional, Sequence, Tuple, Union
|
|
4
|
+
import numpy as np
|
|
5
|
+
from keras import ops
|
|
6
|
+
|
|
7
|
+
from k3_node.rag.subgraph import SubgraphResult
|
|
8
|
+
|
|
9
|
+
|
|
10
|
+
def subgraph_to_triples(
|
|
11
|
+
subgraph: SubgraphResult,
|
|
12
|
+
id_to_entity: Optional[Dict[int, str]] = None,
|
|
13
|
+
id_to_relation: Optional[Dict[int, str]] = None,
|
|
14
|
+
) -> List[Tuple[str, str, str]]:
|
|
15
|
+
"""Convert a SubgraphResult into a list of (head, relation, tail) string triples.
|
|
16
|
+
|
|
17
|
+
Args:
|
|
18
|
+
subgraph: The extracted SubgraphResult.
|
|
19
|
+
id_to_entity: Optional map from original entity integer ID to entity string.
|
|
20
|
+
id_to_relation: Optional map from relation integer ID to relation string.
|
|
21
|
+
|
|
22
|
+
Returns:
|
|
23
|
+
List of `(head, relation, tail)` string triples.
|
|
24
|
+
"""
|
|
25
|
+
if subgraph.num_edges == 0:
|
|
26
|
+
return []
|
|
27
|
+
|
|
28
|
+
edge_index_np = np.asarray(ops.convert_to_numpy(subgraph.edge_index)).astype(np.int64)
|
|
29
|
+
row, col = edge_index_np[0], edge_index_np[1]
|
|
30
|
+
|
|
31
|
+
# Map subgraph indices back to original node IDs if mapping/nodes available
|
|
32
|
+
if subgraph.nodes is not None and len(subgraph.nodes) > 0:
|
|
33
|
+
orig_row = subgraph.nodes[row]
|
|
34
|
+
orig_col = subgraph.nodes[col]
|
|
35
|
+
else:
|
|
36
|
+
orig_row, orig_col = row, col
|
|
37
|
+
|
|
38
|
+
edge_type_np = None
|
|
39
|
+
if subgraph.edge_type is not None:
|
|
40
|
+
edge_type_np = np.asarray(ops.convert_to_numpy(subgraph.edge_type)).astype(np.int64)
|
|
41
|
+
|
|
42
|
+
triples = []
|
|
43
|
+
for i in range(len(row)):
|
|
44
|
+
h_id = int(orig_row[i])
|
|
45
|
+
t_id = int(orig_col[i])
|
|
46
|
+
h_name = id_to_entity.get(h_id, f"Node_{h_id}") if id_to_entity else f"Node_{h_id}"
|
|
47
|
+
t_name = id_to_entity.get(t_id, f"Node_{t_id}") if id_to_entity else f"Node_{t_id}"
|
|
48
|
+
|
|
49
|
+
if edge_type_np is not None:
|
|
50
|
+
r_id = int(edge_type_np[i])
|
|
51
|
+
r_name = id_to_relation.get(r_id, f"rel_{r_id}") if id_to_relation else f"rel_{r_id}"
|
|
52
|
+
else:
|
|
53
|
+
r_name = "connected_to"
|
|
54
|
+
|
|
55
|
+
triples.append((h_name, r_name, t_name))
|
|
56
|
+
|
|
57
|
+
return triples
|
|
58
|
+
|
|
59
|
+
|
|
60
|
+
def verbalize_subgraph(
|
|
61
|
+
subgraph: SubgraphResult,
|
|
62
|
+
id_to_entity: Optional[Dict[int, str]] = None,
|
|
63
|
+
id_to_relation: Optional[Dict[int, str]] = None,
|
|
64
|
+
format_style: Literal["triples", "markdown", "natural"] = "markdown",
|
|
65
|
+
max_triples: Optional[int] = 50,
|
|
66
|
+
) -> str:
|
|
67
|
+
"""Verbalize an extracted subgraph into textual knowledge for LLM prompt augmentation.
|
|
68
|
+
|
|
69
|
+
Args:
|
|
70
|
+
subgraph: Extracted SubgraphResult around retrieved entities.
|
|
71
|
+
id_to_entity: Optional dictionary mapping node ID to entity name.
|
|
72
|
+
id_to_relation: Optional dictionary mapping relation ID to relation name.
|
|
73
|
+
format_style:
|
|
74
|
+
- `"triples"`: List of `(Head, Relation, Tail)` text triples.
|
|
75
|
+
- `"markdown"`: Markdown bullet list with facts.
|
|
76
|
+
- `"natural"`: Natural language sentences (`"Head relation Tail."`).
|
|
77
|
+
max_triples: Maximum number of facts to include in the context.
|
|
78
|
+
|
|
79
|
+
Returns:
|
|
80
|
+
Formatted textual context string ready to be injected into an LLM prompt.
|
|
81
|
+
"""
|
|
82
|
+
triples = subgraph_to_triples(subgraph, id_to_entity, id_to_relation)
|
|
83
|
+
if not triples:
|
|
84
|
+
return "No relevant knowledge graph facts retrieved."
|
|
85
|
+
|
|
86
|
+
if max_triples is not None and len(triples) > max_triples:
|
|
87
|
+
triples = triples[:max_triples]
|
|
88
|
+
|
|
89
|
+
if format_style == "triples":
|
|
90
|
+
lines = [f"({h}, {r}, {t})" for h, r, t in triples]
|
|
91
|
+
return "\n".join(lines)
|
|
92
|
+
|
|
93
|
+
elif format_style == "natural":
|
|
94
|
+
lines = []
|
|
95
|
+
for h, r, t in triples:
|
|
96
|
+
# Clean relation string (replace underscores with spaces)
|
|
97
|
+
rel_str = r.replace("_", " ")
|
|
98
|
+
lines.append(f"{h} {rel_str} {t}.")
|
|
99
|
+
return " ".join(lines)
|
|
100
|
+
|
|
101
|
+
else: # "markdown"
|
|
102
|
+
lines = ["### Retrieved Knowledge Graph Facts:"]
|
|
103
|
+
for h, r, t in triples:
|
|
104
|
+
lines.append(f"- **{h}** — *{r}* -> **{t}**")
|
|
105
|
+
return "\n".join(lines)
|
|
106
|
+
|
|
107
|
+
|
|
108
|
+
def format_llm_prompt(
|
|
109
|
+
query: str,
|
|
110
|
+
context: str,
|
|
111
|
+
system_prompt: Optional[str] = None,
|
|
112
|
+
model_family: Literal["llama3", "mistral", "chatml", "standard"] = "llama3",
|
|
113
|
+
) -> str:
|
|
114
|
+
"""Format query and retrieved KG context into prompt templates for LLMs.
|
|
115
|
+
|
|
116
|
+
Supported model families include:
|
|
117
|
+
- `"llama3"`: Meta Llama 3 / 3.1 instruct template.
|
|
118
|
+
- `"mistral"`: Mistral / Mixtral instruct template.
|
|
119
|
+
- `"chatml"`: OpenAI / Qwen ChatML template.
|
|
120
|
+
- `"standard"`: General markdown system/user format.
|
|
121
|
+
|
|
122
|
+
Args:
|
|
123
|
+
query: The user's input question or instruction.
|
|
124
|
+
context: The verbalized knowledge graph context.
|
|
125
|
+
system_prompt: Optional system prompt to instruct the LLM.
|
|
126
|
+
model_family: Target model prompt format. (default: "llama3")
|
|
127
|
+
|
|
128
|
+
Returns:
|
|
129
|
+
Formatted prompt string.
|
|
130
|
+
"""
|
|
131
|
+
if system_prompt is None:
|
|
132
|
+
system_prompt = (
|
|
133
|
+
"You are an expert assistant augmented with a Knowledge Graph. "
|
|
134
|
+
"Use the provided Knowledge Graph Context to accurately answer the user's question."
|
|
135
|
+
)
|
|
136
|
+
|
|
137
|
+
if model_family == "llama3":
|
|
138
|
+
return (
|
|
139
|
+
"<|begin_of_text|><|start_header_id|>system<|end_header_id|>\n\n"
|
|
140
|
+
f"{system_prompt}<|eot_id|><|start_header_id|>user<|end_header_id|>\n\n"
|
|
141
|
+
f"Knowledge Graph Context:\n{context}\n\n"
|
|
142
|
+
f"Question: {query}<|eot_id|><|start_header_id|>assistant<|end_header_id|>\n\n"
|
|
143
|
+
)
|
|
144
|
+
elif model_family == "mistral":
|
|
145
|
+
return (
|
|
146
|
+
f"<s>[INST] {system_prompt}\n\n"
|
|
147
|
+
f"Knowledge Graph Context:\n{context}\n\n"
|
|
148
|
+
f"Question: {query} [/INST]"
|
|
149
|
+
)
|
|
150
|
+
elif model_family == "chatml":
|
|
151
|
+
return (
|
|
152
|
+
f"<|im_start|>system\n{system_prompt}<|im_end|>\n"
|
|
153
|
+
f"<|im_start|>user\nKnowledge Graph Context:\n{context}\n\nQuestion: {query}<|im_end|>\n"
|
|
154
|
+
f"<|im_start|>assistant\n"
|
|
155
|
+
)
|
|
156
|
+
else: # "standard"
|
|
157
|
+
return (
|
|
158
|
+
f"System: {system_prompt}\n\n"
|
|
159
|
+
f"Knowledge Graph Context:\n{context}\n\n"
|
|
160
|
+
f"User Question: {query}\n\n"
|
|
161
|
+
f"Assistant:"
|
|
162
|
+
)
|
|
@@ -0,0 +1,19 @@
|
|
|
1
|
+
"""High-level Task APIs for K3-Node."""
|
|
2
|
+
|
|
3
|
+
from k3_node.tasks.base import BaseTask
|
|
4
|
+
from k3_node.tasks.backbone_resolver import resolve_backbone
|
|
5
|
+
from k3_node.tasks.node_classification import NodeClassifier
|
|
6
|
+
from k3_node.tasks.node_regression import NodeRegressor
|
|
7
|
+
from k3_node.tasks.graph_classification import GraphClassifier
|
|
8
|
+
from k3_node.tasks.graph_regression import GraphRegressor
|
|
9
|
+
from k3_node.tasks.link_prediction import LinkPredictor
|
|
10
|
+
|
|
11
|
+
__all__ = [
|
|
12
|
+
"BaseTask",
|
|
13
|
+
"resolve_backbone",
|
|
14
|
+
"NodeClassifier",
|
|
15
|
+
"NodeRegressor",
|
|
16
|
+
"GraphClassifier",
|
|
17
|
+
"GraphRegressor",
|
|
18
|
+
"LinkPredictor",
|
|
19
|
+
]
|
|
@@ -0,0 +1,125 @@
|
|
|
1
|
+
"""Backbone resolver mapping string identifiers to K3-Node models."""
|
|
2
|
+
|
|
3
|
+
from typing import Any, Dict, Optional, Union
|
|
4
|
+
import keras
|
|
5
|
+
|
|
6
|
+
from k3_node import models
|
|
7
|
+
|
|
8
|
+
|
|
9
|
+
BACKBONE_REGISTRY = {
|
|
10
|
+
"gcn": models.GCN,
|
|
11
|
+
"gat": models.GAT,
|
|
12
|
+
"sage": models.GraphSAGE,
|
|
13
|
+
"graphsage": models.GraphSAGE,
|
|
14
|
+
"gin": models.GIN,
|
|
15
|
+
"pna": models.PNA,
|
|
16
|
+
"edge_cnn": models.EdgeCNN,
|
|
17
|
+
"edgecnn": models.EdgeCNN,
|
|
18
|
+
"mlp": models.MLP,
|
|
19
|
+
"linkx": models.LINKX,
|
|
20
|
+
"pmlp": models.PMLP,
|
|
21
|
+
"sgformer": models.SGFormer,
|
|
22
|
+
"polynormer": models.Polynormer,
|
|
23
|
+
"schnet": models.SchNet,
|
|
24
|
+
"dimenet": models.DimeNet,
|
|
25
|
+
"dimenet++": models.DimeNetPlusPlus,
|
|
26
|
+
"dimenetplusplus": models.DimeNetPlusPlus,
|
|
27
|
+
"attentive_fp": models.AttentiveFP,
|
|
28
|
+
"attentivefp": models.AttentiveFP,
|
|
29
|
+
"graph_unet": models.GraphUNet,
|
|
30
|
+
}
|
|
31
|
+
|
|
32
|
+
|
|
33
|
+
def resolve_backbone(
|
|
34
|
+
backbone: Union[str, keras.Model, Any],
|
|
35
|
+
in_channels: int,
|
|
36
|
+
out_channels: int,
|
|
37
|
+
hidden_channels: int = 64,
|
|
38
|
+
num_layers: int = 2,
|
|
39
|
+
dropout: float = 0.5,
|
|
40
|
+
**kwargs,
|
|
41
|
+
) -> keras.Model:
|
|
42
|
+
r"""Resolves a backbone string or instance into a compiled/callable Keras model."""
|
|
43
|
+
if isinstance(backbone, keras.Model):
|
|
44
|
+
return backbone
|
|
45
|
+
|
|
46
|
+
if not isinstance(backbone, str):
|
|
47
|
+
raise TypeError(
|
|
48
|
+
f"Expected backbone to be a string name or keras.Model instance, got {type(backbone)}"
|
|
49
|
+
)
|
|
50
|
+
|
|
51
|
+
name = backbone.lower().strip()
|
|
52
|
+
if name not in BACKBONE_REGISTRY:
|
|
53
|
+
available = ", ".join(sorted(BACKBONE_REGISTRY.keys()))
|
|
54
|
+
raise ValueError(f"Unknown backbone '{backbone}'. Available backbones: {available}")
|
|
55
|
+
|
|
56
|
+
cls = BACKBONE_REGISTRY[name]
|
|
57
|
+
|
|
58
|
+
# Handle models with specialized constructors
|
|
59
|
+
if name in ("schnet",):
|
|
60
|
+
return cls(hidden_channels=hidden_channels, num_filters=hidden_channels, num_interactions=num_layers, **kwargs)
|
|
61
|
+
elif name in ("dimenet", "dimenet++", "dimenetplusplus"):
|
|
62
|
+
return cls(hidden_channels=hidden_channels, out_channels=out_channels, num_blocks=num_layers, **kwargs)
|
|
63
|
+
elif name in ("attentive_fp", "attentivefp"):
|
|
64
|
+
edge_dim = kwargs.pop("edge_dim", in_channels)
|
|
65
|
+
num_timesteps = kwargs.pop("num_timesteps", 2)
|
|
66
|
+
return cls(
|
|
67
|
+
in_channels=in_channels,
|
|
68
|
+
hidden_channels=hidden_channels,
|
|
69
|
+
out_channels=out_channels,
|
|
70
|
+
edge_dim=edge_dim,
|
|
71
|
+
num_layers=num_layers,
|
|
72
|
+
num_timesteps=num_timesteps,
|
|
73
|
+
dropout=dropout,
|
|
74
|
+
**kwargs,
|
|
75
|
+
)
|
|
76
|
+
elif name in ("mlp",):
|
|
77
|
+
channel_list = [in_channels] + [hidden_channels] * (num_layers - 1) + [out_channels]
|
|
78
|
+
return cls(channel_list=channel_list, dropout=dropout, **kwargs)
|
|
79
|
+
elif name in ("linkx",):
|
|
80
|
+
num_nodes = kwargs.pop("num_nodes", 1000)
|
|
81
|
+
return cls(
|
|
82
|
+
num_nodes=num_nodes,
|
|
83
|
+
in_channels=in_channels,
|
|
84
|
+
hidden_channels=hidden_channels,
|
|
85
|
+
out_channels=out_channels,
|
|
86
|
+
num_layers=num_layers,
|
|
87
|
+
dropout=dropout,
|
|
88
|
+
**kwargs,
|
|
89
|
+
)
|
|
90
|
+
elif name in ("sgformer",):
|
|
91
|
+
return cls(
|
|
92
|
+
in_channels=in_channels,
|
|
93
|
+
hidden_channels=hidden_channels,
|
|
94
|
+
out_channels=out_channels,
|
|
95
|
+
num_layers=num_layers,
|
|
96
|
+
dropout=dropout,
|
|
97
|
+
**kwargs,
|
|
98
|
+
)
|
|
99
|
+
elif name in ("polynormer",):
|
|
100
|
+
return cls(
|
|
101
|
+
in_channels=in_channels,
|
|
102
|
+
hidden_channels=hidden_channels,
|
|
103
|
+
out_channels=out_channels,
|
|
104
|
+
num_layers=num_layers,
|
|
105
|
+
dropout=dropout,
|
|
106
|
+
**kwargs,
|
|
107
|
+
)
|
|
108
|
+
elif name in ("graph_unet",):
|
|
109
|
+
return cls(
|
|
110
|
+
in_channels=in_channels,
|
|
111
|
+
hidden_channels=hidden_channels,
|
|
112
|
+
out_channels=out_channels,
|
|
113
|
+
depth=num_layers,
|
|
114
|
+
**kwargs,
|
|
115
|
+
)
|
|
116
|
+
else:
|
|
117
|
+
# Standard BasicGNN (GCN, GAT, SAGE, GIN, PNA, EdgeCNN)
|
|
118
|
+
return cls(
|
|
119
|
+
in_channels=in_channels,
|
|
120
|
+
hidden_channels=hidden_channels,
|
|
121
|
+
num_layers=num_layers,
|
|
122
|
+
out_channels=out_channels,
|
|
123
|
+
dropout=dropout,
|
|
124
|
+
**kwargs,
|
|
125
|
+
)
|
k3_node/tasks/base.py
ADDED
|
@@ -0,0 +1,67 @@
|
|
|
1
|
+
"""Base task abstraction for high-level K3-Node estimators."""
|
|
2
|
+
|
|
3
|
+
from typing import Any, Dict, List, Optional, Tuple, Union
|
|
4
|
+
import keras
|
|
5
|
+
from keras import ops
|
|
6
|
+
|
|
7
|
+
from k3_node.data import BaseData
|
|
8
|
+
from k3_node.hub.hub_mixin import K3NodeHubMixin
|
|
9
|
+
|
|
10
|
+
|
|
11
|
+
class BaseTask(K3NodeHubMixin):
|
|
12
|
+
r"""Abstract base task estimator providing common training, evaluation,
|
|
13
|
+
and serialization workflows.
|
|
14
|
+
"""
|
|
15
|
+
|
|
16
|
+
def __init__(self, model: Optional[keras.Model] = None):
|
|
17
|
+
self.model = model
|
|
18
|
+
self._is_compiled = False
|
|
19
|
+
|
|
20
|
+
def compile(
|
|
21
|
+
self,
|
|
22
|
+
optimizer: Optional[Union[str, keras.optimizers.Optimizer]] = None,
|
|
23
|
+
loss: Optional[Any] = None,
|
|
24
|
+
metrics: Optional[List[Any]] = None,
|
|
25
|
+
**kwargs,
|
|
26
|
+
):
|
|
27
|
+
r"""Configures the task model for training."""
|
|
28
|
+
if self.model is None:
|
|
29
|
+
raise RuntimeError("Model has not been initialized. Call fit() or construct with a model first.")
|
|
30
|
+
|
|
31
|
+
opt = optimizer or keras.optimizers.Adam(learning_rate=0.01)
|
|
32
|
+
self.model.compile(optimizer=opt, loss=loss, metrics=metrics, **kwargs)
|
|
33
|
+
self._is_compiled = True
|
|
34
|
+
return self
|
|
35
|
+
|
|
36
|
+
def summary(self):
|
|
37
|
+
r"""Prints a string summary of the underlying neural network."""
|
|
38
|
+
if self.model is not None:
|
|
39
|
+
return self.model.summary()
|
|
40
|
+
print("Model has not been initialized yet.")
|
|
41
|
+
|
|
42
|
+
def save(self, filepath: str):
|
|
43
|
+
r"""Saves the underlying model weights."""
|
|
44
|
+
if self.model is not None:
|
|
45
|
+
self.model.save(filepath)
|
|
46
|
+
else:
|
|
47
|
+
raise RuntimeError("Cannot save an uninitialized model.")
|
|
48
|
+
|
|
49
|
+
@classmethod
|
|
50
|
+
def load(cls, filepath: str, **kwargs):
|
|
51
|
+
r"""Loads a saved task model from disk."""
|
|
52
|
+
model = keras.models.load_model(filepath, **kwargs)
|
|
53
|
+
instance = cls(model=model)
|
|
54
|
+
instance._is_compiled = True
|
|
55
|
+
return instance
|
|
56
|
+
|
|
57
|
+
def _extract_inputs(self, data: Any):
|
|
58
|
+
r"""Extracts input tensors from a Data object, tuple, or dictionary."""
|
|
59
|
+
if hasattr(data, "inputs"):
|
|
60
|
+
return data.inputs
|
|
61
|
+
elif isinstance(data, (tuple, list)):
|
|
62
|
+
return data
|
|
63
|
+
elif hasattr(data, "x") and hasattr(data, "edge_index"):
|
|
64
|
+
if hasattr(data, "edge_attr") and data.edge_attr is not None:
|
|
65
|
+
return (data.x, data.edge_index, data.edge_attr)
|
|
66
|
+
return (data.x, data.edge_index)
|
|
67
|
+
return data
|
|
@@ -0,0 +1,270 @@
|
|
|
1
|
+
"""High-level Graph Classification 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 GraphClassificationModel(keras.Model):
|
|
15
|
+
r"""Internal wrapper combining a node-level GNN backbone, a global readout
|
|
16
|
+
pooling operation, and a final classification dense head.
|
|
17
|
+
"""
|
|
18
|
+
|
|
19
|
+
def __init__(
|
|
20
|
+
self,
|
|
21
|
+
backbone: keras.Model,
|
|
22
|
+
pooling: str = "mean",
|
|
23
|
+
hidden_channels: int = 64,
|
|
24
|
+
num_classes: int = 2,
|
|
25
|
+
dropout: float = 0.5,
|
|
26
|
+
):
|
|
27
|
+
super().__init__()
|
|
28
|
+
self.backbone = backbone
|
|
29
|
+
self.pooling = pooling
|
|
30
|
+
self.dropout = layers.Dropout(dropout) if dropout > 0 else None
|
|
31
|
+
self.head = layers.Dense(num_classes)
|
|
32
|
+
self.num_graphs = None
|
|
33
|
+
|
|
34
|
+
def call(self, inputs, training=False):
|
|
35
|
+
if isinstance(inputs, (tuple, list)):
|
|
36
|
+
x, edge_index = inputs[0], inputs[1]
|
|
37
|
+
batch = inputs[2] if len(inputs) > 2 else None
|
|
38
|
+
size = inputs[3] if len(inputs) > 3 else self.num_graphs
|
|
39
|
+
else:
|
|
40
|
+
x = inputs
|
|
41
|
+
edge_index = getattr(x, "edge_index", None)
|
|
42
|
+
batch = getattr(x, "batch", None)
|
|
43
|
+
size = getattr(x, "num_graphs", self.num_graphs)
|
|
44
|
+
x = getattr(x, "x", x)
|
|
45
|
+
|
|
46
|
+
if batch is None:
|
|
47
|
+
batch = ops.zeros((ops.shape(x)[0],), dtype="int64")
|
|
48
|
+
|
|
49
|
+
h = self.backbone((x, edge_index), training=training)
|
|
50
|
+
|
|
51
|
+
if self.pooling in ("mean", "global_mean_pool"):
|
|
52
|
+
g = k3_pool.global_mean_pool(h, batch, size=size)
|
|
53
|
+
elif self.pooling in ("add", "sum", "global_add_pool"):
|
|
54
|
+
g = k3_pool.global_add_pool(h, batch, size=size)
|
|
55
|
+
elif self.pooling in ("max", "global_max_pool"):
|
|
56
|
+
g = k3_pool.global_max_pool(h, batch, size=size)
|
|
57
|
+
else:
|
|
58
|
+
g = k3_pool.global_mean_pool(h, batch, size=size)
|
|
59
|
+
|
|
60
|
+
if self.dropout is not None:
|
|
61
|
+
g = self.dropout(g, training=training)
|
|
62
|
+
return self.head(g)
|
|
63
|
+
|
|
64
|
+
|
|
65
|
+
class GraphClassifier(BaseTask):
|
|
66
|
+
r"""High-level estimator for graph classification tasks (e.g., molecular property,
|
|
67
|
+
bioinformatics, social graph classification).
|
|
68
|
+
|
|
69
|
+
Args:
|
|
70
|
+
backbone: GNN architecture (``"gin"``, ``"gcn"``, ``"gat"``, ``"sage"``,
|
|
71
|
+
``"pna"``, etc.) or a custom :class:`keras.Model`. (default: ``"gin"``)
|
|
72
|
+
in_channels (int, optional): Size of input node features.
|
|
73
|
+
hidden_channels (int, optional): Dimensionality of hidden node features. (default: ``64``)
|
|
74
|
+
num_classes (int, optional): Number of graph classes.
|
|
75
|
+
num_layers (int, optional): Number of GNN layers. (default: ``3``)
|
|
76
|
+
pooling (str, optional): Readout pooling (``"mean"``, ``"add"``, ``"max"``). (default: ``"mean"``)
|
|
77
|
+
dropout (float, optional): Dropout probability. (default: ``0.5``)
|
|
78
|
+
**backbone_kwargs: Additional arguments forwarded to the backbone constructor.
|
|
79
|
+
"""
|
|
80
|
+
|
|
81
|
+
def __init__(
|
|
82
|
+
self,
|
|
83
|
+
backbone: Union[str, keras.Model] = "gin",
|
|
84
|
+
in_channels: Optional[int] = None,
|
|
85
|
+
hidden_channels: int = 64,
|
|
86
|
+
num_classes: Optional[int] = None,
|
|
87
|
+
num_layers: int = 3,
|
|
88
|
+
pooling: str = "mean",
|
|
89
|
+
dropout: float = 0.5,
|
|
90
|
+
**backbone_kwargs,
|
|
91
|
+
):
|
|
92
|
+
super().__init__()
|
|
93
|
+
self.backbone = backbone
|
|
94
|
+
self.in_channels = in_channels
|
|
95
|
+
self.hidden_channels = hidden_channels
|
|
96
|
+
self.num_classes = num_classes
|
|
97
|
+
self.num_layers = num_layers
|
|
98
|
+
self.pooling = pooling
|
|
99
|
+
self.dropout = dropout
|
|
100
|
+
self.backbone_kwargs = backbone_kwargs
|
|
101
|
+
|
|
102
|
+
def _init_model(self, sample_data: Any, dataset: Optional[Any] = None):
|
|
103
|
+
in_c = self.in_channels
|
|
104
|
+
if in_c is None:
|
|
105
|
+
if hasattr(sample_data, "num_node_features") and sample_data.num_node_features > 0:
|
|
106
|
+
in_c = sample_data.num_node_features
|
|
107
|
+
elif hasattr(sample_data, "num_features") and sample_data.num_features > 0:
|
|
108
|
+
in_c = sample_data.num_features
|
|
109
|
+
elif hasattr(sample_data, "x") and sample_data.x is not None:
|
|
110
|
+
in_c = int(ops.shape(sample_data.x)[-1])
|
|
111
|
+
else:
|
|
112
|
+
raise ValueError("Could not infer in_channels.")
|
|
113
|
+
|
|
114
|
+
out_c = self.num_classes
|
|
115
|
+
if out_c is None:
|
|
116
|
+
if dataset is not None and hasattr(dataset, "num_classes") and dataset.num_classes is not None:
|
|
117
|
+
out_c = dataset.num_classes
|
|
118
|
+
elif dataset is not None and isinstance(dataset, (list, tuple)):
|
|
119
|
+
max_y = 0
|
|
120
|
+
for g in dataset[:100]:
|
|
121
|
+
if hasattr(g, "y") and g.y is not None:
|
|
122
|
+
max_y = max(max_y, int(ops.convert_to_numpy(ops.max(g.y))))
|
|
123
|
+
out_c = max_y + 1
|
|
124
|
+
elif hasattr(sample_data, "num_classes") and sample_data.num_classes is not None and sample_data.num_classes > 1:
|
|
125
|
+
out_c = sample_data.num_classes
|
|
126
|
+
elif hasattr(sample_data, "y") and sample_data.y is not None:
|
|
127
|
+
out_c = int(ops.convert_to_numpy(ops.max(sample_data.y))) + 1
|
|
128
|
+
else:
|
|
129
|
+
out_c = 2 # default binary
|
|
130
|
+
|
|
131
|
+
out_c = max(int(out_c), 2)
|
|
132
|
+
|
|
133
|
+
self.in_channels = in_c
|
|
134
|
+
self.num_classes = out_c
|
|
135
|
+
|
|
136
|
+
gnn = resolve_backbone(
|
|
137
|
+
self.backbone,
|
|
138
|
+
in_channels=in_c,
|
|
139
|
+
out_channels=self.hidden_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 = GraphClassificationModel(
|
|
147
|
+
backbone=gnn,
|
|
148
|
+
pooling=self.pooling,
|
|
149
|
+
hidden_channels=self.hidden_channels,
|
|
150
|
+
num_classes=out_c,
|
|
151
|
+
dropout=self.dropout,
|
|
152
|
+
)
|
|
153
|
+
|
|
154
|
+
def fit(
|
|
155
|
+
self,
|
|
156
|
+
dataset: Any,
|
|
157
|
+
epochs: int = 20,
|
|
158
|
+
lr: float = 0.01,
|
|
159
|
+
batch_size: int = 32,
|
|
160
|
+
shuffle: bool = True,
|
|
161
|
+
verbose: int = 1,
|
|
162
|
+
callbacks: Optional[List[Any]] = None,
|
|
163
|
+
):
|
|
164
|
+
r"""Trains the graph classifier."""
|
|
165
|
+
if not isinstance(dataset, DataLoader):
|
|
166
|
+
loader = DataLoader(dataset, batch_size=batch_size, shuffle=shuffle)
|
|
167
|
+
sample = dataset[0]
|
|
168
|
+
else:
|
|
169
|
+
loader = dataset
|
|
170
|
+
sample = next(iter(loader))
|
|
171
|
+
|
|
172
|
+
if self.model is None:
|
|
173
|
+
self._init_model(sample, dataset=dataset)
|
|
174
|
+
|
|
175
|
+
if not self._is_compiled:
|
|
176
|
+
self.model.compile(
|
|
177
|
+
optimizer=keras.optimizers.Adam(learning_rate=lr),
|
|
178
|
+
loss=keras.losses.SparseCategoricalCrossentropy(from_logits=True),
|
|
179
|
+
metrics=[keras.metrics.SparseCategoricalAccuracy(name="acc")],
|
|
180
|
+
)
|
|
181
|
+
self._is_compiled = True
|
|
182
|
+
|
|
183
|
+
history = {"loss": [], "acc": []}
|
|
184
|
+
for epoch in range(epochs):
|
|
185
|
+
batch_losses = []
|
|
186
|
+
batch_accs = []
|
|
187
|
+
for batch in loader:
|
|
188
|
+
x = ops.convert_to_tensor(batch.x, dtype="float32")
|
|
189
|
+
edge_index = ops.convert_to_tensor(batch.edge_index, dtype="int64")
|
|
190
|
+
batch_vec = ops.convert_to_tensor(batch.batch, dtype="int64")
|
|
191
|
+
y = ops.convert_to_tensor(batch.y, dtype="int64")
|
|
192
|
+
num_g = int(ops.shape(y)[0])
|
|
193
|
+
if hasattr(self.model, "num_graphs"):
|
|
194
|
+
self.model.num_graphs = num_g
|
|
195
|
+
|
|
196
|
+
if not self.model.built:
|
|
197
|
+
y_pred = self.model((x, edge_index, batch_vec), training=False)
|
|
198
|
+
self.model.built = True
|
|
199
|
+
if hasattr(self.model, "_compile_loss") and self.model._compile_loss is not None:
|
|
200
|
+
self.model._compile_loss.build(y, y_pred)
|
|
201
|
+
if hasattr(self.model, "_compile_metrics") and self.model._compile_metrics is not None:
|
|
202
|
+
self.model._compile_metrics.build(y, y_pred)
|
|
203
|
+
if self.model.optimizer is not None and not self.model.optimizer.built:
|
|
204
|
+
self.model.optimizer.build(self.model.trainable_variables)
|
|
205
|
+
|
|
206
|
+
res = self.model.train_on_batch((x, edge_index, batch_vec), y)
|
|
207
|
+
if isinstance(res, (list, tuple)):
|
|
208
|
+
batch_losses.append(float(res[0]))
|
|
209
|
+
if len(res) > 1:
|
|
210
|
+
batch_accs.append(float(res[1]))
|
|
211
|
+
else:
|
|
212
|
+
batch_losses.append(float(res))
|
|
213
|
+
|
|
214
|
+
avg_loss = float(np.mean(batch_losses)) if batch_losses else 0.0
|
|
215
|
+
avg_acc = float(np.mean(batch_accs)) if batch_accs else 0.0
|
|
216
|
+
history["loss"].append(avg_loss)
|
|
217
|
+
history["acc"].append(avg_acc)
|
|
218
|
+
if verbose:
|
|
219
|
+
print(f"Epoch {epoch + 1}/{epochs} - loss: {avg_loss:.4f} - acc: {avg_acc:.4f}")
|
|
220
|
+
|
|
221
|
+
return history
|
|
222
|
+
|
|
223
|
+
def predict_proba(self, dataset_or_loader: Any, batch_size: int = 32):
|
|
224
|
+
r"""Predicts class probabilities for graphs."""
|
|
225
|
+
if not isinstance(dataset_or_loader, DataLoader):
|
|
226
|
+
loader = DataLoader(dataset_or_loader, batch_size=batch_size, shuffle=False)
|
|
227
|
+
else:
|
|
228
|
+
loader = dataset_or_loader
|
|
229
|
+
|
|
230
|
+
probs = []
|
|
231
|
+
for batch in loader:
|
|
232
|
+
x = ops.convert_to_tensor(batch.x, dtype="float32")
|
|
233
|
+
edge_index = ops.convert_to_tensor(batch.edge_index, dtype="int64")
|
|
234
|
+
batch_vec = ops.convert_to_tensor(batch.batch, dtype="int64")
|
|
235
|
+
num_g = int(ops.convert_to_numpy(ops.max(batch_vec))) + 1 if ops.shape(batch_vec)[0] > 0 else 1
|
|
236
|
+
if hasattr(self.model, "num_graphs"):
|
|
237
|
+
self.model.num_graphs = num_g
|
|
238
|
+
logits = self.model((x, edge_index, batch_vec), training=False)
|
|
239
|
+
probs.append(ops.softmax(logits, axis=-1))
|
|
240
|
+
return ops.concatenate(probs, axis=0)
|
|
241
|
+
|
|
242
|
+
def predict(self, dataset_or_loader: Any, batch_size: int = 32):
|
|
243
|
+
r"""Predicts discrete class labels for graphs."""
|
|
244
|
+
probs = self.predict_proba(dataset_or_loader, batch_size=batch_size)
|
|
245
|
+
return ops.argmax(probs, axis=-1)
|
|
246
|
+
|
|
247
|
+
def evaluate(self, dataset_or_loader: Any, batch_size: int = 32) -> Dict[str, float]:
|
|
248
|
+
r"""Evaluates classification accuracy on the dataset."""
|
|
249
|
+
if not isinstance(dataset_or_loader, DataLoader):
|
|
250
|
+
loader = DataLoader(dataset_or_loader, batch_size=batch_size, shuffle=False)
|
|
251
|
+
else:
|
|
252
|
+
loader = dataset_or_loader
|
|
253
|
+
|
|
254
|
+
correct = 0
|
|
255
|
+
total = 0
|
|
256
|
+
for batch in loader:
|
|
257
|
+
x = ops.convert_to_tensor(batch.x, dtype="float32")
|
|
258
|
+
edge_index = ops.convert_to_tensor(batch.edge_index, dtype="int64")
|
|
259
|
+
batch_vec = ops.convert_to_tensor(batch.batch, dtype="int64")
|
|
260
|
+
num_g = int(ops.convert_to_numpy(ops.max(batch_vec))) + 1 if ops.shape(batch_vec)[0] > 0 else 1
|
|
261
|
+
if hasattr(self.model, "num_graphs"):
|
|
262
|
+
self.model.num_graphs = num_g
|
|
263
|
+
logits = self.model((x, edge_index, batch_vec), training=False)
|
|
264
|
+
pred = ops.argmax(logits, axis=-1)
|
|
265
|
+
pred_np = ops.convert_to_numpy(ops.cast(pred, "int64"))
|
|
266
|
+
y_np = ops.convert_to_numpy(ops.cast(batch.y, "int64"))
|
|
267
|
+
correct += int((pred_np == y_np).sum())
|
|
268
|
+
total += int(y_np.shape[0])
|
|
269
|
+
|
|
270
|
+
return {"accuracy": float(correct / max(total, 1))}
|