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,172 @@
|
|
|
1
|
+
"""Cross-backend export bridge for PyTorch and JAX Keras backends."""
|
|
2
|
+
|
|
3
|
+
import importlib
|
|
4
|
+
import json
|
|
5
|
+
import os
|
|
6
|
+
from pathlib import Path
|
|
7
|
+
import subprocess
|
|
8
|
+
import sys
|
|
9
|
+
import tempfile
|
|
10
|
+
from typing import Any, Dict, List, Optional, Union
|
|
11
|
+
import numpy as np
|
|
12
|
+
|
|
13
|
+
import keras
|
|
14
|
+
from keras import ops
|
|
15
|
+
|
|
16
|
+
|
|
17
|
+
def is_tensorflow_backend() -> bool:
|
|
18
|
+
r"""Checks whether the currently active Keras backend is TensorFlow."""
|
|
19
|
+
return keras.backend.backend() == "tensorflow"
|
|
20
|
+
|
|
21
|
+
|
|
22
|
+
def export_via_tf_subprocess(
|
|
23
|
+
exporter_type: str, # "onnx" or "tflite"
|
|
24
|
+
model_or_task: Any,
|
|
25
|
+
output_path: Union[str, Path],
|
|
26
|
+
dummy_inputs: Optional[Any] = None,
|
|
27
|
+
**kwargs: Any,
|
|
28
|
+
) -> Path:
|
|
29
|
+
r"""Bridges model export to a TensorFlow worker subprocess when running under PyTorch or JAX.
|
|
30
|
+
|
|
31
|
+
Args:
|
|
32
|
+
exporter_type: "onnx" or "tflite".
|
|
33
|
+
model_or_task: Model or task instance.
|
|
34
|
+
output_path: Target path for the exported model file.
|
|
35
|
+
dummy_inputs: Optional dummy inputs.
|
|
36
|
+
**kwargs: Extra arguments forwarded to the exporter.
|
|
37
|
+
|
|
38
|
+
Returns:
|
|
39
|
+
Path to the exported model file.
|
|
40
|
+
"""
|
|
41
|
+
out_file = Path(output_path).resolve()
|
|
42
|
+
out_file.parent.mkdir(parents=True, exist_ok=True)
|
|
43
|
+
|
|
44
|
+
with tempfile.TemporaryDirectory() as tmpdir:
|
|
45
|
+
tmp_path = Path(tmpdir)
|
|
46
|
+
model_dir = tmp_path / "model"
|
|
47
|
+
model_dir.mkdir(parents=True, exist_ok=True)
|
|
48
|
+
|
|
49
|
+
# 1. Ensure model is built before saving weights
|
|
50
|
+
raw_model = getattr(model_or_task, "model", None) or model_or_task
|
|
51
|
+
if hasattr(raw_model, "built") and not raw_model.built:
|
|
52
|
+
from k3_node.export.onnx_exporter import _extract_model_and_inputs
|
|
53
|
+
try:
|
|
54
|
+
_, extracted_inputs, _ = _extract_model_and_inputs(model_or_task, dummy_inputs)
|
|
55
|
+
raw_model(*extracted_inputs)
|
|
56
|
+
except Exception:
|
|
57
|
+
pass
|
|
58
|
+
|
|
59
|
+
# 2. Save model or task weights & config
|
|
60
|
+
if hasattr(model_or_task, "save_pretrained"):
|
|
61
|
+
model_or_task.save_pretrained(str(model_dir))
|
|
62
|
+
cls_module = model_or_task.__class__.__module__
|
|
63
|
+
cls_name = model_or_task.__class__.__name__
|
|
64
|
+
elif hasattr(model_or_task, "model") and hasattr(model_or_task.model, "save_pretrained"):
|
|
65
|
+
model_or_task.model.save_pretrained(str(model_dir))
|
|
66
|
+
cls_module = model_or_task.model.__class__.__module__
|
|
67
|
+
cls_name = model_or_task.model.__class__.__name__
|
|
68
|
+
else:
|
|
69
|
+
raise TypeError(
|
|
70
|
+
f"Model or task of type {type(model_or_task)} must support `save_pretrained` "
|
|
71
|
+
"for cross-backend export."
|
|
72
|
+
)
|
|
73
|
+
|
|
74
|
+
# 2. Serialize dummy inputs if provided
|
|
75
|
+
dummy_path = tmp_path / "dummy.npz"
|
|
76
|
+
dummy_meta_path = tmp_path / "dummy_meta.json"
|
|
77
|
+
has_dummy = False
|
|
78
|
+
|
|
79
|
+
if dummy_inputs is not None:
|
|
80
|
+
arrays: Dict[str, np.ndarray] = {}
|
|
81
|
+
meta: Dict[str, Any] = {}
|
|
82
|
+
|
|
83
|
+
if hasattr(dummy_inputs, "x") and hasattr(dummy_inputs, "edge_index"):
|
|
84
|
+
meta["type"] = "data"
|
|
85
|
+
arrays["x"] = np.asarray(ops.convert_to_numpy(dummy_inputs.x))
|
|
86
|
+
arrays["edge_index"] = np.asarray(ops.convert_to_numpy(dummy_inputs.edge_index))
|
|
87
|
+
if hasattr(dummy_inputs, "batch") and dummy_inputs.batch is not None:
|
|
88
|
+
arrays["batch"] = np.asarray(ops.convert_to_numpy(dummy_inputs.batch))
|
|
89
|
+
elif hasattr(dummy_inputs, "z") and hasattr(dummy_inputs, "pos"):
|
|
90
|
+
meta["type"] = "data"
|
|
91
|
+
arrays["z"] = np.asarray(ops.convert_to_numpy(dummy_inputs.z))
|
|
92
|
+
arrays["pos"] = np.asarray(ops.convert_to_numpy(dummy_inputs.pos))
|
|
93
|
+
if hasattr(dummy_inputs, "batch") and dummy_inputs.batch is not None:
|
|
94
|
+
arrays["batch"] = np.asarray(ops.convert_to_numpy(dummy_inputs.batch))
|
|
95
|
+
elif isinstance(dummy_inputs, (tuple, list)):
|
|
96
|
+
meta["type"] = "tuple"
|
|
97
|
+
for i, elem in enumerate(dummy_inputs):
|
|
98
|
+
arrays[f"arr_{i}"] = np.asarray(ops.convert_to_numpy(elem))
|
|
99
|
+
elif isinstance(dummy_inputs, dict):
|
|
100
|
+
meta["type"] = "dict"
|
|
101
|
+
for k, v in dummy_inputs.items():
|
|
102
|
+
if v is not None:
|
|
103
|
+
arrays[str(k)] = np.asarray(ops.convert_to_numpy(v))
|
|
104
|
+
else:
|
|
105
|
+
meta["type"] = "single"
|
|
106
|
+
arrays["arr_0"] = np.asarray(ops.convert_to_numpy(dummy_inputs))
|
|
107
|
+
|
|
108
|
+
np.savez(str(dummy_path), **arrays)
|
|
109
|
+
with open(dummy_meta_path, "w", encoding="utf-8") as f:
|
|
110
|
+
json.dump(meta, f)
|
|
111
|
+
has_dummy = True
|
|
112
|
+
|
|
113
|
+
# 3. Build worker command
|
|
114
|
+
worker_code = f"""
|
|
115
|
+
import os
|
|
116
|
+
os.environ["KERAS_BACKEND"] = "tensorflow"
|
|
117
|
+
import importlib
|
|
118
|
+
import json
|
|
119
|
+
from pathlib import Path
|
|
120
|
+
import numpy as np
|
|
121
|
+
|
|
122
|
+
import k3_node
|
|
123
|
+
from k3_node.data import Data
|
|
124
|
+
from k3_node.export.onnx_exporter import export_onnx
|
|
125
|
+
from k3_node.export.tflite_exporter import export_tflite
|
|
126
|
+
|
|
127
|
+
# Load model / task
|
|
128
|
+
mod = importlib.import_module("{cls_module}")
|
|
129
|
+
cls = getattr(mod, "{cls_name}")
|
|
130
|
+
model_or_task = cls.from_pretrained(r"{model_dir}")
|
|
131
|
+
|
|
132
|
+
# Load dummy inputs
|
|
133
|
+
dummy = None
|
|
134
|
+
has_dummy = {has_dummy}
|
|
135
|
+
if has_dummy:
|
|
136
|
+
with open(r"{dummy_meta_path}", "r") as f:
|
|
137
|
+
meta = json.load(f)
|
|
138
|
+
npz = np.load(r"{dummy_path}")
|
|
139
|
+
t = meta["type"]
|
|
140
|
+
if t == "data":
|
|
141
|
+
kwargs = {{k: npz[k] for k in npz.files}}
|
|
142
|
+
dummy = Data(**kwargs)
|
|
143
|
+
elif t == "tuple":
|
|
144
|
+
dummy = tuple(npz[f"arr_{{i}}"] for i in range(len(npz.files)))
|
|
145
|
+
elif t == "dict":
|
|
146
|
+
dummy = {{k: npz[k] for k in npz.files}}
|
|
147
|
+
else:
|
|
148
|
+
dummy = npz["arr_0"]
|
|
149
|
+
|
|
150
|
+
extra_kwargs = json.loads(r'''{json.dumps(kwargs)}''')
|
|
151
|
+
if "{exporter_type}" == "onnx":
|
|
152
|
+
export_onnx(model_or_task, r"{out_file}", dummy_inputs=dummy, **extra_kwargs)
|
|
153
|
+
elif "{exporter_type}" == "tflite":
|
|
154
|
+
export_tflite(model_or_task, r"{out_file}", dummy_inputs=dummy, **extra_kwargs)
|
|
155
|
+
"""
|
|
156
|
+
|
|
157
|
+
env = dict(os.environ, KERAS_BACKEND="tensorflow")
|
|
158
|
+
res = subprocess.run(
|
|
159
|
+
[sys.executable, "-c", worker_code],
|
|
160
|
+
capture_output=True,
|
|
161
|
+
text=True,
|
|
162
|
+
env=env,
|
|
163
|
+
)
|
|
164
|
+
|
|
165
|
+
if res.returncode != 0:
|
|
166
|
+
raise RuntimeError(
|
|
167
|
+
f"Cross-backend export to {exporter_type.upper()} failed with exit code {res.returncode}:\n"
|
|
168
|
+
f"STDOUT: {res.stdout}\n"
|
|
169
|
+
f"STDERR: {res.stderr}"
|
|
170
|
+
)
|
|
171
|
+
|
|
172
|
+
return out_file
|
|
@@ -0,0 +1,190 @@
|
|
|
1
|
+
"""Turn-key ONNX exporter for K3-Node models and tasks."""
|
|
2
|
+
|
|
3
|
+
import os
|
|
4
|
+
from pathlib import Path
|
|
5
|
+
from typing import Any, Dict, List, Optional, Tuple, Union
|
|
6
|
+
import numpy as np
|
|
7
|
+
|
|
8
|
+
import keras
|
|
9
|
+
from keras import ops
|
|
10
|
+
|
|
11
|
+
|
|
12
|
+
def _extract_model_and_inputs(
|
|
13
|
+
model_or_task: Any,
|
|
14
|
+
dummy_inputs: Optional[Any] = None,
|
|
15
|
+
) -> Tuple[Any, Tuple[Any, ...], List[Any]]:
|
|
16
|
+
r"""Extracts the underlying neural network model and resolves dummy inputs / signatures."""
|
|
17
|
+
# 1. Resolve model
|
|
18
|
+
if hasattr(model_or_task, "model") and model_or_task.model is not None:
|
|
19
|
+
raw_model = model_or_task.model
|
|
20
|
+
else:
|
|
21
|
+
raw_model = model_or_task
|
|
22
|
+
|
|
23
|
+
# 2. Extract inputs from dummy_inputs if provided
|
|
24
|
+
if dummy_inputs is not None:
|
|
25
|
+
if hasattr(dummy_inputs, "x") and hasattr(dummy_inputs, "edge_index"):
|
|
26
|
+
inputs = [dummy_inputs.x, dummy_inputs.edge_index]
|
|
27
|
+
names = ["x", "edge_index"]
|
|
28
|
+
if hasattr(dummy_inputs, "batch") and dummy_inputs.batch is not None:
|
|
29
|
+
inputs.append(dummy_inputs.batch)
|
|
30
|
+
names.append("batch")
|
|
31
|
+
return raw_model, tuple(inputs), names
|
|
32
|
+
elif hasattr(dummy_inputs, "z") and hasattr(dummy_inputs, "pos"):
|
|
33
|
+
inputs = [dummy_inputs.z, dummy_inputs.pos]
|
|
34
|
+
names = ["z", "pos"]
|
|
35
|
+
if hasattr(dummy_inputs, "batch") and dummy_inputs.batch is not None:
|
|
36
|
+
inputs.append(dummy_inputs.batch)
|
|
37
|
+
names.append("batch")
|
|
38
|
+
return raw_model, tuple(inputs), names
|
|
39
|
+
elif isinstance(dummy_inputs, (tuple, list)):
|
|
40
|
+
return raw_model, tuple(dummy_inputs), None
|
|
41
|
+
elif isinstance(dummy_inputs, dict):
|
|
42
|
+
return raw_model, (dummy_inputs,), None
|
|
43
|
+
else:
|
|
44
|
+
return raw_model, (dummy_inputs,), None
|
|
45
|
+
|
|
46
|
+
# 3. Auto-infer dummy inputs from model hyperparameters
|
|
47
|
+
cls_name = model_or_task.__class__.__name__
|
|
48
|
+
in_channels = (
|
|
49
|
+
getattr(model_or_task, "in_channels", None)
|
|
50
|
+
or getattr(raw_model, "in_channels", None)
|
|
51
|
+
or 16
|
|
52
|
+
)
|
|
53
|
+
|
|
54
|
+
if cls_name == "NodeClassifier":
|
|
55
|
+
x = ops.zeros((4, in_channels), dtype="float32")
|
|
56
|
+
edge_index = ops.convert_to_tensor([[0, 1, 2, 3], [1, 2, 3, 0]], dtype="int64")
|
|
57
|
+
return raw_model, (x, edge_index), ["x", "edge_index"]
|
|
58
|
+
elif cls_name in ("GraphClassifier", "GraphRegressor"):
|
|
59
|
+
x = ops.zeros((4, in_channels), dtype="float32")
|
|
60
|
+
edge_index = ops.convert_to_tensor([[0, 1, 2, 3], [1, 2, 3, 0]], dtype="int64")
|
|
61
|
+
batch = ops.convert_to_tensor([0, 0, 1, 1], dtype="int64")
|
|
62
|
+
if hasattr(raw_model, "num_graphs"):
|
|
63
|
+
raw_model.num_graphs = 2
|
|
64
|
+
return raw_model, (x, edge_index, batch), ["x", "edge_index", "batch"]
|
|
65
|
+
elif cls_name == "LinkPredictor":
|
|
66
|
+
x = ops.zeros((4, in_channels), dtype="float32")
|
|
67
|
+
edge_index = ops.convert_to_tensor([[0, 1, 2, 3], [1, 2, 3, 0]], dtype="int64")
|
|
68
|
+
label_idx = ops.convert_to_tensor([[0, 1], [1, 2]], dtype="int64")
|
|
69
|
+
return raw_model, ((x, edge_index), label_idx), ["edge_tuple", "edge_label_index"]
|
|
70
|
+
elif cls_name in ("SchNet", "DimeNet", "DimeNetPlusPlus", "ViSNet", "GNNFF"):
|
|
71
|
+
z = ops.convert_to_tensor([1, 6, 8, 1], dtype="int32")
|
|
72
|
+
pos = ops.zeros((4, 3), dtype="float32")
|
|
73
|
+
return raw_model, (z, pos), ["z", "pos"]
|
|
74
|
+
else:
|
|
75
|
+
# Default standard GNN: (x, edge_index)
|
|
76
|
+
x = ops.zeros((4, in_channels), dtype="float32")
|
|
77
|
+
edge_index = ops.convert_to_tensor([[0, 1, 2, 3], [1, 2, 3, 0]], dtype="int64")
|
|
78
|
+
return raw_model, (x, edge_index), ["x", "edge_index"]
|
|
79
|
+
|
|
80
|
+
|
|
81
|
+
def export_onnx(
|
|
82
|
+
model_or_task: Any,
|
|
83
|
+
output_path: Union[str, Path],
|
|
84
|
+
dummy_inputs: Optional[Any] = None,
|
|
85
|
+
opset: int = 17,
|
|
86
|
+
dynamic_axes: bool = True,
|
|
87
|
+
input_names: Optional[List[str]] = None,
|
|
88
|
+
output_names: Optional[List[str]] = None,
|
|
89
|
+
verbose: bool = False,
|
|
90
|
+
) -> Path:
|
|
91
|
+
r"""Exports a K3-Node GNN model or task to high-performance ONNX format.
|
|
92
|
+
|
|
93
|
+
Supports arbitrary Graph Neural Networks (GCN, GAT, GraphSAGE, GIN, SchNet,
|
|
94
|
+
materials models, and task estimators) with dynamic graph sizing (varying numbers
|
|
95
|
+
of nodes and edges).
|
|
96
|
+
|
|
97
|
+
Args:
|
|
98
|
+
model_or_task: A K3-Node task instance (e.g. `NodeClassifier`, `GraphClassifier`)
|
|
99
|
+
or model instance (e.g. `GCN`, `SchNet`, `CHGNet`).
|
|
100
|
+
output_path: Target path for the `.onnx` file.
|
|
101
|
+
dummy_inputs: Optional sample input data (e.g., PyG `Data` object, tuple of tensors).
|
|
102
|
+
If `None`, automatically generated based on model topology.
|
|
103
|
+
opset: ONNX operator set version. (default: `17`)
|
|
104
|
+
dynamic_axes: Whether node and edge dimensions should be dynamic. (default: `True`)
|
|
105
|
+
input_names: Optional custom names for input tensors.
|
|
106
|
+
output_names: Optional custom names for output tensors.
|
|
107
|
+
verbose: Whether to print verbose export progress. (default: `False`)
|
|
108
|
+
|
|
109
|
+
Returns:
|
|
110
|
+
Path object pointing to the generated `.onnx` file.
|
|
111
|
+
"""
|
|
112
|
+
out_file = Path(output_path)
|
|
113
|
+
out_file.parent.mkdir(parents=True, exist_ok=True)
|
|
114
|
+
|
|
115
|
+
from k3_node.export.cross_backend import is_tensorflow_backend, export_via_tf_subprocess
|
|
116
|
+
|
|
117
|
+
if not is_tensorflow_backend():
|
|
118
|
+
return export_via_tf_subprocess(
|
|
119
|
+
exporter_type="onnx",
|
|
120
|
+
model_or_task=model_or_task,
|
|
121
|
+
output_path=output_path,
|
|
122
|
+
dummy_inputs=dummy_inputs,
|
|
123
|
+
opset=opset,
|
|
124
|
+
dynamic_axes=dynamic_axes,
|
|
125
|
+
input_names=input_names,
|
|
126
|
+
output_names=output_names,
|
|
127
|
+
verbose=verbose,
|
|
128
|
+
)
|
|
129
|
+
|
|
130
|
+
try:
|
|
131
|
+
import tf2onnx
|
|
132
|
+
import tensorflow as tf
|
|
133
|
+
except ImportError:
|
|
134
|
+
raise ImportError(
|
|
135
|
+
"The `tf2onnx` and `tensorflow` packages are required for ONNX export. "
|
|
136
|
+
"Install them via `pip install tf2onnx tensorflow onnx`."
|
|
137
|
+
)
|
|
138
|
+
|
|
139
|
+
model, inputs, inferred_names = _extract_model_and_inputs(model_or_task, dummy_inputs)
|
|
140
|
+
names = input_names or inferred_names or [f"input_{i}" for i in range(len(inputs))]
|
|
141
|
+
|
|
142
|
+
# Build input signature with dynamic axes if requested
|
|
143
|
+
signature = []
|
|
144
|
+
for inp, name in zip(inputs, names):
|
|
145
|
+
inp_np = ops.convert_to_numpy(inp)
|
|
146
|
+
dtype = tf.as_dtype(inp_np.dtype)
|
|
147
|
+
if dynamic_axes:
|
|
148
|
+
if inp_np.ndim == 2 and inp_np.shape[0] == 2 and (
|
|
149
|
+
np.issubdtype(inp_np.dtype, np.integer) or "edge" in name.lower()
|
|
150
|
+
):
|
|
151
|
+
# Edge index: (2, num_edges) -> dynamic num_edges
|
|
152
|
+
shape = (2, None)
|
|
153
|
+
elif inp_np.ndim == 2:
|
|
154
|
+
# Node feature: (num_nodes, in_channels) -> dynamic num_nodes
|
|
155
|
+
shape = (None, inp_np.shape[1])
|
|
156
|
+
elif inp_np.ndim == 1:
|
|
157
|
+
# Vector (batch or z): (num_nodes,) -> dynamic num_nodes
|
|
158
|
+
shape = (None,)
|
|
159
|
+
else:
|
|
160
|
+
shape = tuple(None if i == 0 else s for i, s in enumerate(inp_np.shape))
|
|
161
|
+
else:
|
|
162
|
+
shape = inp_np.shape
|
|
163
|
+
signature.append(tf.TensorSpec(shape=shape, dtype=dtype, name=name))
|
|
164
|
+
|
|
165
|
+
# Define trace function
|
|
166
|
+
@tf.function(input_signature=signature)
|
|
167
|
+
def forward_fn(*tensors):
|
|
168
|
+
return model(*tensors)
|
|
169
|
+
|
|
170
|
+
if verbose:
|
|
171
|
+
print(f"Exporting model to ONNX with input signature: {signature}")
|
|
172
|
+
|
|
173
|
+
# Convert using tf2onnx
|
|
174
|
+
model_proto, _ = tf2onnx.convert.from_function(
|
|
175
|
+
forward_fn,
|
|
176
|
+
input_signature=signature,
|
|
177
|
+
output_path=str(out_file),
|
|
178
|
+
opset=opset,
|
|
179
|
+
)
|
|
180
|
+
|
|
181
|
+
# Validate ONNX graph
|
|
182
|
+
try:
|
|
183
|
+
import onnx
|
|
184
|
+
onnx_model = onnx.load(str(out_file))
|
|
185
|
+
onnx.checker.check_model(onnx_model)
|
|
186
|
+
except Exception as e:
|
|
187
|
+
if verbose:
|
|
188
|
+
print(f"Warning: ONNX validation check returned: {e}")
|
|
189
|
+
|
|
190
|
+
return out_file
|
|
@@ -0,0 +1,254 @@
|
|
|
1
|
+
"""Lightweight inference runtime engines for serving ONNX and TFLite models."""
|
|
2
|
+
|
|
3
|
+
from pathlib import Path
|
|
4
|
+
from typing import Any, Dict, List, Optional, Union
|
|
5
|
+
import numpy as np
|
|
6
|
+
|
|
7
|
+
|
|
8
|
+
class ONNXModel:
|
|
9
|
+
r"""High-performance serving wrapper for exported ONNX GNN models.
|
|
10
|
+
|
|
11
|
+
Requires only `onnxruntime` and `numpy`. Completely decoupled from Keras,
|
|
12
|
+
PyTorch, and TensorFlow for lightweight production microservices.
|
|
13
|
+
|
|
14
|
+
Example:
|
|
15
|
+
```python
|
|
16
|
+
from k3_node.export import ONNXModel
|
|
17
|
+
model = ONNXModel("cora_gcn.onnx")
|
|
18
|
+
preds = model.predict(graph_data)
|
|
19
|
+
```
|
|
20
|
+
"""
|
|
21
|
+
|
|
22
|
+
def __init__(
|
|
23
|
+
self,
|
|
24
|
+
model_path: Union[str, Path],
|
|
25
|
+
providers: Optional[List[str]] = None,
|
|
26
|
+
session_options: Optional[Any] = None,
|
|
27
|
+
):
|
|
28
|
+
r"""Initializes the ONNX runtime inference session.
|
|
29
|
+
|
|
30
|
+
Args:
|
|
31
|
+
model_path: Path to the `.onnx` model file.
|
|
32
|
+
providers: Execution providers list (e.g. `['CUDAExecutionProvider', 'CPUExecutionProvider']`).
|
|
33
|
+
If `None`, automatically picks the fastest available provider.
|
|
34
|
+
session_options: Optional custom ONNX Runtime SessionOptions.
|
|
35
|
+
"""
|
|
36
|
+
try:
|
|
37
|
+
import onnxruntime as ort
|
|
38
|
+
except ImportError:
|
|
39
|
+
raise ImportError(
|
|
40
|
+
"The `onnxruntime` package is required to load and serve ONNX models. "
|
|
41
|
+
"Install it via `pip install onnxruntime` (or `onnxruntime-gpu` for CUDA/TensorRT)."
|
|
42
|
+
)
|
|
43
|
+
|
|
44
|
+
self.model_path = Path(model_path)
|
|
45
|
+
if not self.model_path.exists():
|
|
46
|
+
raise FileNotFoundError(f"ONNX model file not found at: {self.model_path}")
|
|
47
|
+
|
|
48
|
+
if providers is None:
|
|
49
|
+
available = ort.get_available_providers()
|
|
50
|
+
# Prioritize TensorRT, CUDA, CoreML, DirectML, CPU
|
|
51
|
+
priority = [
|
|
52
|
+
"TensorrtExecutionProvider",
|
|
53
|
+
"CUDAExecutionProvider",
|
|
54
|
+
"CoreMLExecutionProvider",
|
|
55
|
+
"DmlExecutionProvider",
|
|
56
|
+
"CPUExecutionProvider",
|
|
57
|
+
]
|
|
58
|
+
providers = [p for p in priority if p in available]
|
|
59
|
+
|
|
60
|
+
self.session = ort.InferenceSession(
|
|
61
|
+
str(self.model_path),
|
|
62
|
+
sess_options=session_options,
|
|
63
|
+
providers=providers,
|
|
64
|
+
)
|
|
65
|
+
self.input_names = [inp.name for inp in self.session.get_inputs()]
|
|
66
|
+
self.output_names = [out.name for out in self.session.get_outputs()]
|
|
67
|
+
|
|
68
|
+
def predict(self, data: Any = None, *args: Any, **kwargs: Any) -> np.ndarray:
|
|
69
|
+
r"""Runs low-latency inference on graph data.
|
|
70
|
+
|
|
71
|
+
Accepts:
|
|
72
|
+
- PyG / K3-Node `Data` object
|
|
73
|
+
- Dictionary of tensors
|
|
74
|
+
- Positional numpy arrays
|
|
75
|
+
|
|
76
|
+
Args:
|
|
77
|
+
data: Input graph Data, dictionary, or array.
|
|
78
|
+
*args: Additional positional inputs.
|
|
79
|
+
|
|
80
|
+
Returns:
|
|
81
|
+
Numpy array containing model predictions or logits.
|
|
82
|
+
"""
|
|
83
|
+
feed_dict = {}
|
|
84
|
+
|
|
85
|
+
# 1. PyG/K3 Data object
|
|
86
|
+
if hasattr(data, "x") and hasattr(data, "edge_index"):
|
|
87
|
+
feed_dict = self._match_inputs({
|
|
88
|
+
"x": np.asarray(data.x, dtype=np.float32),
|
|
89
|
+
"edge_index": np.asarray(data.edge_index, dtype=np.int64),
|
|
90
|
+
"batch": np.asarray(getattr(data, "batch", None), dtype=np.int64) if getattr(data, "batch", None) is not None else None,
|
|
91
|
+
})
|
|
92
|
+
elif hasattr(data, "z") and hasattr(data, "pos"):
|
|
93
|
+
feed_dict = self._match_inputs({
|
|
94
|
+
"z": np.asarray(data.z, dtype=np.int32),
|
|
95
|
+
"pos": np.asarray(data.pos, dtype=np.float32),
|
|
96
|
+
"batch": np.asarray(getattr(data, "batch", None), dtype=np.int32) if getattr(data, "batch", None) is not None else None,
|
|
97
|
+
})
|
|
98
|
+
elif isinstance(data, dict):
|
|
99
|
+
feed_dict = self._match_inputs(data)
|
|
100
|
+
elif isinstance(data, (tuple, list)):
|
|
101
|
+
for name, val in zip(self.input_names, data):
|
|
102
|
+
feed_dict[name] = np.asarray(val)
|
|
103
|
+
elif data is not None:
|
|
104
|
+
all_args = [data, *args]
|
|
105
|
+
for name, val in zip(self.input_names, all_args):
|
|
106
|
+
feed_dict[name] = np.asarray(val)
|
|
107
|
+
|
|
108
|
+
outputs = self.session.run(self.output_names, feed_dict)
|
|
109
|
+
return outputs[0] if len(outputs) == 1 else tuple(outputs)
|
|
110
|
+
|
|
111
|
+
def _match_inputs(self, named_inputs: Dict[str, Any]) -> Dict[str, np.ndarray]:
|
|
112
|
+
matched = {}
|
|
113
|
+
valid_inputs = {k: np.asarray(v) for k, v in named_inputs.items() if v is not None}
|
|
114
|
+
lower_inputs = {k.lower(): v for k, v in valid_inputs.items()}
|
|
115
|
+
|
|
116
|
+
used_keys = set()
|
|
117
|
+
session_inputs = self.session.get_inputs()
|
|
118
|
+
|
|
119
|
+
# Step 1: Match exact or known semantic names
|
|
120
|
+
for sess_inp in session_inputs:
|
|
121
|
+
name = sess_inp.name
|
|
122
|
+
clean = name.split(":")[0].lower()
|
|
123
|
+
|
|
124
|
+
target_val = None
|
|
125
|
+
matched_key = None
|
|
126
|
+
|
|
127
|
+
if clean in lower_inputs:
|
|
128
|
+
target_val = lower_inputs[clean]
|
|
129
|
+
matched_key = clean
|
|
130
|
+
elif "edge" in clean and "edge_index" in lower_inputs:
|
|
131
|
+
target_val = lower_inputs["edge_index"]
|
|
132
|
+
matched_key = "edge_index"
|
|
133
|
+
elif clean == "x" and "x" in lower_inputs:
|
|
134
|
+
target_val = lower_inputs["x"]
|
|
135
|
+
matched_key = "x"
|
|
136
|
+
elif clean in ("batch", "batch_idx") and "batch" in lower_inputs:
|
|
137
|
+
target_val = lower_inputs["batch"]
|
|
138
|
+
matched_key = "batch"
|
|
139
|
+
elif clean in ("z", "atomic_numbers") and "z" in lower_inputs:
|
|
140
|
+
target_val = lower_inputs["z"]
|
|
141
|
+
matched_key = "z"
|
|
142
|
+
elif clean in ("pos", "positions", "coord") and "pos" in lower_inputs:
|
|
143
|
+
target_val = lower_inputs["pos"]
|
|
144
|
+
matched_key = "pos"
|
|
145
|
+
|
|
146
|
+
if target_val is not None:
|
|
147
|
+
matched[name] = target_val
|
|
148
|
+
used_keys.add(matched_key)
|
|
149
|
+
|
|
150
|
+
# Step 2: Positional fallback for remaining inputs
|
|
151
|
+
unmatched_session = [inp for inp in session_inputs if inp.name not in matched]
|
|
152
|
+
unmatched_keys = [k for k in valid_inputs.keys() if k.lower() not in used_keys]
|
|
153
|
+
|
|
154
|
+
if unmatched_session and unmatched_keys:
|
|
155
|
+
for sess_inp, k in zip(unmatched_session, unmatched_keys):
|
|
156
|
+
matched[sess_inp.name] = valid_inputs[k]
|
|
157
|
+
|
|
158
|
+
# Step 3: Align data types with expected ONNX session tensor types
|
|
159
|
+
type_map = {
|
|
160
|
+
"tensor(float)": np.float32,
|
|
161
|
+
"tensor(float16)": np.float16,
|
|
162
|
+
"tensor(double)": np.float64,
|
|
163
|
+
"tensor(int64)": np.int64,
|
|
164
|
+
"tensor(int32)": np.int32,
|
|
165
|
+
"tensor(int8)": np.int8,
|
|
166
|
+
"tensor(uint8)": np.uint8,
|
|
167
|
+
"tensor(bool)": np.bool_,
|
|
168
|
+
}
|
|
169
|
+
for sess_inp in session_inputs:
|
|
170
|
+
if sess_inp.name in matched:
|
|
171
|
+
expected_np_type = type_map.get(sess_inp.type)
|
|
172
|
+
if expected_np_type and matched[sess_inp.name].dtype != expected_np_type:
|
|
173
|
+
matched[sess_inp.name] = matched[sess_inp.name].astype(expected_np_type)
|
|
174
|
+
|
|
175
|
+
return matched
|
|
176
|
+
|
|
177
|
+
|
|
178
|
+
class TFLiteModel:
|
|
179
|
+
r"""Lightweight serving wrapper for TensorFlow Lite flatbuffer GNN models.
|
|
180
|
+
|
|
181
|
+
Requires only standard `tensorflow` or `tflite_runtime`. Ideal for mobile,
|
|
182
|
+
Raspberry Pi, and edge embedded devices.
|
|
183
|
+
|
|
184
|
+
Example:
|
|
185
|
+
```python
|
|
186
|
+
from k3_node.export import TFLiteModel
|
|
187
|
+
model = TFLiteModel("model.tflite")
|
|
188
|
+
preds = model.predict(graph_data)
|
|
189
|
+
```
|
|
190
|
+
"""
|
|
191
|
+
|
|
192
|
+
def __init__(self, model_path: Union[str, Path]):
|
|
193
|
+
r"""Initializes the TFLite interpreter.
|
|
194
|
+
|
|
195
|
+
Args:
|
|
196
|
+
model_path: Path to the `.tflite` model file.
|
|
197
|
+
"""
|
|
198
|
+
self.model_path = Path(model_path)
|
|
199
|
+
if not self.model_path.exists():
|
|
200
|
+
raise FileNotFoundError(f"TFLite model file not found at: {self.model_path}")
|
|
201
|
+
|
|
202
|
+
try:
|
|
203
|
+
import tensorflow as tf
|
|
204
|
+
self.interpreter = tf.lite.Interpreter(model_path=str(self.model_path))
|
|
205
|
+
except ImportError:
|
|
206
|
+
try:
|
|
207
|
+
import tflite_runtime.interpreter as tflite
|
|
208
|
+
self.interpreter = tflite.Interpreter(model_path=str(self.model_path))
|
|
209
|
+
except ImportError:
|
|
210
|
+
raise ImportError(
|
|
211
|
+
"Either `tensorflow` or `tflite_runtime` is required to run TFLite models. "
|
|
212
|
+
"Install via `pip install tflite-runtime` or `pip install tensorflow`."
|
|
213
|
+
)
|
|
214
|
+
|
|
215
|
+
self.interpreter.allocate_tensors()
|
|
216
|
+
self.input_details = self.interpreter.get_input_details()
|
|
217
|
+
self.output_details = self.interpreter.get_output_details()
|
|
218
|
+
|
|
219
|
+
def predict(self, data: Any = None, *args: Any, **kwargs: Any) -> np.ndarray:
|
|
220
|
+
r"""Runs inference using the TFLite interpreter.
|
|
221
|
+
|
|
222
|
+
Args:
|
|
223
|
+
data: Input graph Data, dict, or numpy array.
|
|
224
|
+
*args: Additional positional inputs.
|
|
225
|
+
|
|
226
|
+
Returns:
|
|
227
|
+
Numpy array containing prediction results.
|
|
228
|
+
"""
|
|
229
|
+
inputs = []
|
|
230
|
+
if hasattr(data, "x") and hasattr(data, "edge_index"):
|
|
231
|
+
inputs = [np.asarray(data.x, dtype=np.float32), np.asarray(data.edge_index, dtype=np.int64)]
|
|
232
|
+
if hasattr(data, "batch") and data.batch is not None:
|
|
233
|
+
inputs.append(np.asarray(data.batch, dtype=np.int64))
|
|
234
|
+
elif hasattr(data, "z") and hasattr(data, "pos"):
|
|
235
|
+
inputs = [np.asarray(data.z, dtype=np.int32), np.asarray(data.pos, dtype=np.float32)]
|
|
236
|
+
if hasattr(data, "batch") and data.batch is not None:
|
|
237
|
+
inputs.append(np.asarray(data.batch, dtype=np.int32))
|
|
238
|
+
elif isinstance(data, (tuple, list)):
|
|
239
|
+
inputs = [np.asarray(x) for x in data]
|
|
240
|
+
elif isinstance(data, dict):
|
|
241
|
+
inputs = [np.asarray(v) for v in data.values()]
|
|
242
|
+
elif data is not None:
|
|
243
|
+
inputs = [np.asarray(data)] + [np.asarray(a) for a in args]
|
|
244
|
+
|
|
245
|
+
# Feed tensors
|
|
246
|
+
for detail, inp in zip(self.input_details, inputs):
|
|
247
|
+
# Cast dtype to match interpreter expectation
|
|
248
|
+
target_dtype = detail["dtype"]
|
|
249
|
+
inp_cast = inp.astype(target_dtype) if inp.dtype != target_dtype else inp
|
|
250
|
+
self.interpreter.set_tensor(detail["index"], inp_cast)
|
|
251
|
+
|
|
252
|
+
self.interpreter.invoke()
|
|
253
|
+
outputs = [self.interpreter.get_tensor(d["index"]) for d in self.output_details]
|
|
254
|
+
return outputs[0] if len(outputs) == 1 else tuple(outputs)
|