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,244 @@
|
|
|
1
|
+
"""Table-to-Graph ETL converter for single tabular datasets."""
|
|
2
|
+
|
|
3
|
+
from typing import Any, Callable, Dict, List, Optional, Sequence, Union
|
|
4
|
+
import numpy as np
|
|
5
|
+
|
|
6
|
+
from k3_node.data import Data
|
|
7
|
+
from k3_node.etl.encoders import TabularEncoder, _get_column_names, _get_column_values
|
|
8
|
+
from k3_node.etl.graph_builders import (
|
|
9
|
+
KNNGraphBuilder,
|
|
10
|
+
SimilarityGraphBuilder,
|
|
11
|
+
SharedEntityGraphBuilder,
|
|
12
|
+
SequentialGraphBuilder,
|
|
13
|
+
)
|
|
14
|
+
|
|
15
|
+
|
|
16
|
+
class TableToGraph:
|
|
17
|
+
r"""ETL pipeline converting tabular datasets (DataFrames, CSVs, or dictionaries)
|
|
18
|
+
into graph :class:`k3_node.data.Data` objects with node features, graph topology,
|
|
19
|
+
labels, and split masks.
|
|
20
|
+
|
|
21
|
+
Args:
|
|
22
|
+
feature_cols: List of column names used as node features. If :obj:`None`, all
|
|
23
|
+
columns except :obj:`target_col` and :obj:`id_col` are used.
|
|
24
|
+
target_col (str, optional): Target column used to create ground-truth labels :obj:`y`.
|
|
25
|
+
id_col (str, optional): Unique row identifier column (e.g. customer_id, product_id).
|
|
26
|
+
Stored in :obj:`data.node_ids` and :obj:`data.id_to_index`.
|
|
27
|
+
edge_strategy (str or callable): Strategy for constructing edges (``"knn"``,
|
|
28
|
+
``"similarity"``, ``"shared_entity"``, ``"sequential"``, or a custom callable).
|
|
29
|
+
(default: ``"knn"``)
|
|
30
|
+
edge_kwargs (dict, optional): Keyword arguments forwarded to the edge builder.
|
|
31
|
+
(e.g., ``{"k": 5, "metric": "cosine"}`` for KNN).
|
|
32
|
+
column_encoders (dict, optional): Explicit mapping from column names to encoder instances.
|
|
33
|
+
train_ratio (float, optional): Fraction of nodes for training mask. (default: ``None``)
|
|
34
|
+
val_ratio (float, optional): Fraction of nodes for validation mask. (default: ``None``)
|
|
35
|
+
test_ratio (float, optional): Fraction of nodes for test mask. (default: ``None``)
|
|
36
|
+
random_state (int, optional): Random seed for mask splits. (default: ``42``)
|
|
37
|
+
"""
|
|
38
|
+
|
|
39
|
+
def __init__(
|
|
40
|
+
self,
|
|
41
|
+
feature_cols: Optional[List[str]] = None,
|
|
42
|
+
target_col: Optional[str] = None,
|
|
43
|
+
id_col: Optional[str] = None,
|
|
44
|
+
edge_strategy: Union[str, Callable] = "knn",
|
|
45
|
+
edge_kwargs: Optional[Dict[str, Any]] = None,
|
|
46
|
+
column_encoders: Optional[Dict[str, Any]] = None,
|
|
47
|
+
train_ratio: Optional[float] = None,
|
|
48
|
+
val_ratio: Optional[float] = None,
|
|
49
|
+
test_ratio: Optional[float] = None,
|
|
50
|
+
random_state: int = 42,
|
|
51
|
+
):
|
|
52
|
+
self.feature_cols = feature_cols
|
|
53
|
+
self.target_col = target_col
|
|
54
|
+
self.id_col = id_col
|
|
55
|
+
self.edge_strategy = edge_strategy
|
|
56
|
+
self.edge_kwargs = edge_kwargs or {}
|
|
57
|
+
self.column_encoders = column_encoders or {}
|
|
58
|
+
self.train_ratio = train_ratio
|
|
59
|
+
self.val_ratio = val_ratio
|
|
60
|
+
self.test_ratio = test_ratio
|
|
61
|
+
self.random_state = random_state
|
|
62
|
+
|
|
63
|
+
self.encoder_ = TabularEncoder(column_encoders=self.column_encoders)
|
|
64
|
+
self.edge_builder_ = self._resolve_edge_builder()
|
|
65
|
+
|
|
66
|
+
def _resolve_edge_builder(self) -> Callable:
|
|
67
|
+
if callable(self.edge_strategy):
|
|
68
|
+
return self.edge_strategy
|
|
69
|
+
|
|
70
|
+
strategy = str(self.edge_strategy).lower()
|
|
71
|
+
if strategy == "knn":
|
|
72
|
+
return KNNGraphBuilder(**self.edge_kwargs)
|
|
73
|
+
elif strategy in ("similarity", "sim"):
|
|
74
|
+
return SimilarityGraphBuilder(**self.edge_kwargs)
|
|
75
|
+
elif strategy in ("shared_entity", "shared", "bipartite"):
|
|
76
|
+
return SharedEntityGraphBuilder(**self.edge_kwargs)
|
|
77
|
+
elif strategy in ("sequential", "sequence", "temporal"):
|
|
78
|
+
return SequentialGraphBuilder(**self.edge_kwargs)
|
|
79
|
+
else:
|
|
80
|
+
raise ValueError(
|
|
81
|
+
f"Unknown edge_strategy '{self.edge_strategy}'. "
|
|
82
|
+
f"Available: 'knn', 'similarity', 'shared_entity', 'sequential', or custom callable."
|
|
83
|
+
)
|
|
84
|
+
|
|
85
|
+
def fit(self, df_or_dict: Any):
|
|
86
|
+
r"""Fits the tabular feature encoder on the input table."""
|
|
87
|
+
all_cols = _get_column_names(df_or_dict)
|
|
88
|
+
ignore_cols = set()
|
|
89
|
+
if self.target_col:
|
|
90
|
+
ignore_cols.add(self.target_col)
|
|
91
|
+
if self.id_col:
|
|
92
|
+
ignore_cols.add(self.id_col)
|
|
93
|
+
|
|
94
|
+
if self.feature_cols is None:
|
|
95
|
+
fit_cols = [c for c in all_cols if c not in ignore_cols]
|
|
96
|
+
else:
|
|
97
|
+
fit_cols = [c for c in self.feature_cols if c in all_cols and c not in ignore_cols]
|
|
98
|
+
|
|
99
|
+
self.encoder_.fit(df_or_dict, columns=fit_cols)
|
|
100
|
+
return self
|
|
101
|
+
|
|
102
|
+
def transform(self, df_or_dict: Any) -> Data:
|
|
103
|
+
r"""Encodes features, constructs graph edges, and builds a :class:`Data` object."""
|
|
104
|
+
# 1. Node features
|
|
105
|
+
x = self.encoder_.transform(df_or_dict)
|
|
106
|
+
num_nodes = x.shape[0]
|
|
107
|
+
|
|
108
|
+
# 2. Graph topology
|
|
109
|
+
edge_index, edge_attr = self.edge_builder_(x=x, df_or_dict=df_or_dict)
|
|
110
|
+
|
|
111
|
+
# 3. Target labels y
|
|
112
|
+
y = None
|
|
113
|
+
if self.target_col is not None:
|
|
114
|
+
raw_y = _get_column_values(df_or_dict, self.target_col)
|
|
115
|
+
# Check if categorical or numerical
|
|
116
|
+
if all(isinstance(v, (int, np.integer)) for v in raw_y if v is not None):
|
|
117
|
+
y = np.array(raw_y, dtype=np.int64)
|
|
118
|
+
elif all(isinstance(v, (float, int, np.floating, np.integer)) for v in raw_y if v is not None):
|
|
119
|
+
y = np.array(raw_y, dtype=np.float32)
|
|
120
|
+
else:
|
|
121
|
+
# String labels -> label encode to ints
|
|
122
|
+
unique_classes = sorted(list(set(raw_y)))
|
|
123
|
+
mapping = {c: i for i, c in enumerate(unique_classes)}
|
|
124
|
+
y = np.array([mapping[c] for c in raw_y], dtype=np.int64)
|
|
125
|
+
|
|
126
|
+
data = Data(x=x, edge_index=edge_index, edge_attr=edge_attr, y=y)
|
|
127
|
+
|
|
128
|
+
# 4. Optional entity ID tracking
|
|
129
|
+
if self.id_col is not None:
|
|
130
|
+
raw_ids = _get_column_values(df_or_dict, self.id_col)
|
|
131
|
+
data.node_ids = list(raw_ids)
|
|
132
|
+
data.id_to_index = {raw_id: idx for idx, raw_id in enumerate(raw_ids)}
|
|
133
|
+
data.index_to_id = {idx: raw_id for idx, raw_id in enumerate(raw_ids)}
|
|
134
|
+
|
|
135
|
+
# 5. Split masks
|
|
136
|
+
if self.train_ratio is not None:
|
|
137
|
+
np.random.seed(self.random_state)
|
|
138
|
+
indices = np.random.permutation(num_nodes)
|
|
139
|
+
|
|
140
|
+
n_train = int(num_nodes * self.train_ratio)
|
|
141
|
+
val_ratio = self.val_ratio or 0.0
|
|
142
|
+
n_val = int(num_nodes * val_ratio)
|
|
143
|
+
|
|
144
|
+
train_idx = indices[:n_train]
|
|
145
|
+
val_idx = indices[n_train : n_train + n_val]
|
|
146
|
+
test_idx = indices[n_train + n_val :]
|
|
147
|
+
|
|
148
|
+
train_mask = np.zeros(num_nodes, dtype=bool)
|
|
149
|
+
train_mask[train_idx] = True
|
|
150
|
+
data.train_mask = train_mask
|
|
151
|
+
|
|
152
|
+
if val_ratio > 0:
|
|
153
|
+
val_mask = np.zeros(num_nodes, dtype=bool)
|
|
154
|
+
val_mask[val_idx] = True
|
|
155
|
+
data.val_mask = val_mask
|
|
156
|
+
|
|
157
|
+
if self.test_ratio is not None or len(test_idx) > 0:
|
|
158
|
+
test_mask = np.zeros(num_nodes, dtype=bool)
|
|
159
|
+
test_mask[test_idx] = True
|
|
160
|
+
data.test_mask = test_mask
|
|
161
|
+
|
|
162
|
+
return data
|
|
163
|
+
|
|
164
|
+
def fit_transform(self, df_or_dict: Any) -> Data:
|
|
165
|
+
r"""Fits encoders and transforms tabular data into a graph in a single call."""
|
|
166
|
+
return self.fit(df_or_dict).transform(df_or_dict)
|
|
167
|
+
|
|
168
|
+
@classmethod
|
|
169
|
+
def from_dataframe(
|
|
170
|
+
cls,
|
|
171
|
+
df: Any,
|
|
172
|
+
feature_cols: Optional[List[str]] = None,
|
|
173
|
+
target_col: Optional[str] = None,
|
|
174
|
+
id_col: Optional[str] = None,
|
|
175
|
+
edge_strategy: Union[str, Callable] = "knn",
|
|
176
|
+
**kwargs,
|
|
177
|
+
) -> Data:
|
|
178
|
+
r"""Convenience factory method directly converting a pandas DataFrame into a :class:`Data` object."""
|
|
179
|
+
etl = cls(
|
|
180
|
+
feature_cols=feature_cols,
|
|
181
|
+
target_col=target_col,
|
|
182
|
+
id_col=id_col,
|
|
183
|
+
edge_strategy=edge_strategy,
|
|
184
|
+
**kwargs,
|
|
185
|
+
)
|
|
186
|
+
return etl.fit_transform(df)
|
|
187
|
+
|
|
188
|
+
@classmethod
|
|
189
|
+
def from_csv(
|
|
190
|
+
cls,
|
|
191
|
+
filepath: str,
|
|
192
|
+
feature_cols: Optional[List[str]] = None,
|
|
193
|
+
target_col: Optional[str] = None,
|
|
194
|
+
id_col: Optional[str] = None,
|
|
195
|
+
edge_strategy: Union[str, Callable] = "knn",
|
|
196
|
+
**kwargs,
|
|
197
|
+
) -> Data:
|
|
198
|
+
r"""Convenience factory method directly loading a CSV file and converting it into a :class:`Data` object."""
|
|
199
|
+
import csv
|
|
200
|
+
|
|
201
|
+
# Parse CSV into dict of column lists
|
|
202
|
+
with open(filepath, mode="r", encoding="utf-8") as f:
|
|
203
|
+
reader = csv.DictReader(f)
|
|
204
|
+
data_dict: Dict[str, List[Any]] = {field: [] for field in reader.fieldnames or []}
|
|
205
|
+
for row in reader:
|
|
206
|
+
for k, v in row.items():
|
|
207
|
+
# Attempt numeric cast
|
|
208
|
+
try:
|
|
209
|
+
v = float(v) if "." in v else int(v)
|
|
210
|
+
except ValueError:
|
|
211
|
+
pass
|
|
212
|
+
data_dict[k].append(v)
|
|
213
|
+
|
|
214
|
+
return cls.from_dataframe(
|
|
215
|
+
df=data_dict,
|
|
216
|
+
feature_cols=feature_cols,
|
|
217
|
+
target_col=target_col,
|
|
218
|
+
id_col=id_col,
|
|
219
|
+
edge_strategy=edge_strategy,
|
|
220
|
+
**kwargs,
|
|
221
|
+
)
|
|
222
|
+
|
|
223
|
+
|
|
224
|
+
# Alias
|
|
225
|
+
TabularToGraph = TableToGraph
|
|
226
|
+
|
|
227
|
+
|
|
228
|
+
def table_to_graph(
|
|
229
|
+
df_or_dict: Any,
|
|
230
|
+
feature_cols: Optional[List[str]] = None,
|
|
231
|
+
target_col: Optional[str] = None,
|
|
232
|
+
id_col: Optional[str] = None,
|
|
233
|
+
edge_strategy: Union[str, Callable] = "knn",
|
|
234
|
+
**kwargs,
|
|
235
|
+
) -> Data:
|
|
236
|
+
r"""Functional shortcut to convert a table or dictionary into a :class:`k3_node.data.Data` object."""
|
|
237
|
+
etl = TableToGraph(
|
|
238
|
+
feature_cols=feature_cols,
|
|
239
|
+
target_col=target_col,
|
|
240
|
+
id_col=id_col,
|
|
241
|
+
edge_strategy=edge_strategy,
|
|
242
|
+
**kwargs,
|
|
243
|
+
)
|
|
244
|
+
return etl.fit_transform(df_or_dict)
|
k3_node/etl/test_etl.py
ADDED
|
@@ -0,0 +1,318 @@
|
|
|
1
|
+
"""Unit tests for K3-Node Tabular-to-Graph ETL pipelines."""
|
|
2
|
+
|
|
3
|
+
import os
|
|
4
|
+
import tempfile
|
|
5
|
+
import pytest
|
|
6
|
+
import numpy as np
|
|
7
|
+
import pandas as pd
|
|
8
|
+
from keras import ops
|
|
9
|
+
|
|
10
|
+
import k3_node
|
|
11
|
+
from k3_node.data import Data, HeteroData
|
|
12
|
+
from k3_node.etl import (
|
|
13
|
+
NumericalEncoder,
|
|
14
|
+
CategoricalEncoder,
|
|
15
|
+
TabularEncoder,
|
|
16
|
+
KNNGraphBuilder,
|
|
17
|
+
SimilarityGraphBuilder,
|
|
18
|
+
SharedEntityGraphBuilder,
|
|
19
|
+
SequentialGraphBuilder,
|
|
20
|
+
TableToGraph,
|
|
21
|
+
TabularToGraph,
|
|
22
|
+
table_to_graph,
|
|
23
|
+
RelationalToGraph,
|
|
24
|
+
relational_to_graph,
|
|
25
|
+
)
|
|
26
|
+
from k3_node.tasks import NodeClassifier
|
|
27
|
+
|
|
28
|
+
|
|
29
|
+
# ==============================================================================
|
|
30
|
+
# 1. Encoders Unit Tests
|
|
31
|
+
# ==============================================================================
|
|
32
|
+
|
|
33
|
+
def test_numerical_encoder():
|
|
34
|
+
raw = [1.0, 2.0, 3.0, np.nan, 5.0]
|
|
35
|
+
|
|
36
|
+
# Standard scaling
|
|
37
|
+
enc_std = NumericalEncoder(strategy="standard", impute_strategy="mean")
|
|
38
|
+
res_std = enc_std.fit_transform(raw)
|
|
39
|
+
assert res_std.shape == (5, 1)
|
|
40
|
+
assert not np.isnan(res_std).any()
|
|
41
|
+
|
|
42
|
+
# MinMax scaling
|
|
43
|
+
enc_mm = NumericalEncoder(strategy="minmax", impute_strategy="zero")
|
|
44
|
+
res_mm = enc_mm.fit_transform(raw)
|
|
45
|
+
assert res_mm.shape == (5, 1)
|
|
46
|
+
assert res_mm.min() >= 0.0 and res_mm.max() <= 1.0
|
|
47
|
+
|
|
48
|
+
# Log1p
|
|
49
|
+
enc_log = NumericalEncoder(strategy="log1p")
|
|
50
|
+
res_log = enc_log.fit_transform([0.0, 1.0, 10.0])
|
|
51
|
+
assert np.allclose(res_log[0, 0], 0.0)
|
|
52
|
+
|
|
53
|
+
|
|
54
|
+
def test_categorical_encoder():
|
|
55
|
+
raw = ["apple", "banana", "apple", None, "orange"]
|
|
56
|
+
|
|
57
|
+
# One-hot
|
|
58
|
+
enc_oh = CategoricalEncoder(strategy="onehot", handle_unknown="ignore")
|
|
59
|
+
res_oh = enc_oh.fit_transform(raw)
|
|
60
|
+
assert res_oh.shape[0] == 5
|
|
61
|
+
assert res_oh.shape[1] >= 3
|
|
62
|
+
|
|
63
|
+
# Unknown category handling
|
|
64
|
+
transformed = enc_oh.transform(["grape", "banana"])
|
|
65
|
+
assert transformed.shape == (2, res_oh.shape[1])
|
|
66
|
+
assert transformed[0].sum() == 0.0 # unknown ignored
|
|
67
|
+
assert transformed[1].sum() == 1.0 # banana matched
|
|
68
|
+
|
|
69
|
+
# Ordinal
|
|
70
|
+
enc_ord = CategoricalEncoder(strategy="ordinal", unknown_value=-1)
|
|
71
|
+
res_ord = enc_ord.fit_transform(raw)
|
|
72
|
+
assert res_ord.shape == (5, 1)
|
|
73
|
+
assert res_ord.dtype == np.int64
|
|
74
|
+
|
|
75
|
+
# Hash
|
|
76
|
+
enc_hash = CategoricalEncoder(strategy="hash", hash_dim=8)
|
|
77
|
+
res_hash = enc_hash.fit_transform(raw)
|
|
78
|
+
assert res_hash.shape == (5, 8)
|
|
79
|
+
|
|
80
|
+
|
|
81
|
+
def test_tabular_encoder():
|
|
82
|
+
df = pd.DataFrame({
|
|
83
|
+
"age": [25, 30, 35, 40],
|
|
84
|
+
"salary": [50000.0, 60000.0, 75000.0, 90000.0],
|
|
85
|
+
"city": ["NY", "SF", "NY", "LA"],
|
|
86
|
+
})
|
|
87
|
+
|
|
88
|
+
encoder = TabularEncoder()
|
|
89
|
+
x = encoder.fit_transform(df)
|
|
90
|
+
assert x.shape[0] == 4
|
|
91
|
+
# age (1) + salary (1) + city (3 one-hot) = 5
|
|
92
|
+
assert x.shape[1] == 5
|
|
93
|
+
assert x.dtype == np.float32
|
|
94
|
+
|
|
95
|
+
|
|
96
|
+
# ==============================================================================
|
|
97
|
+
# 2. Graph Builders Unit Tests
|
|
98
|
+
# ==============================================================================
|
|
99
|
+
|
|
100
|
+
def test_knn_graph_builder():
|
|
101
|
+
x = np.array([
|
|
102
|
+
[0.0, 0.0],
|
|
103
|
+
[0.1, 0.1],
|
|
104
|
+
[10.0, 10.0],
|
|
105
|
+
[10.1, 10.1],
|
|
106
|
+
], dtype=np.float32)
|
|
107
|
+
|
|
108
|
+
knn = KNNGraphBuilder(k=1, metric="euclidean", loop=False, bidirectional=True)
|
|
109
|
+
edge_index, edge_attr = knn(x)
|
|
110
|
+
|
|
111
|
+
assert edge_index.shape[0] == 2
|
|
112
|
+
assert edge_index.shape[1] > 0
|
|
113
|
+
# Node 0 and Node 1 should be connected
|
|
114
|
+
edges = set(zip(edge_index[0].tolist(), edge_index[1].tolist()))
|
|
115
|
+
assert (0, 1) in edges and (1, 0) in edges
|
|
116
|
+
# Node 2 and Node 3 should be connected
|
|
117
|
+
assert (2, 3) in edges and (3, 2) in edges
|
|
118
|
+
|
|
119
|
+
|
|
120
|
+
def test_similarity_graph_builder():
|
|
121
|
+
x = np.array([
|
|
122
|
+
[1.0, 0.0],
|
|
123
|
+
[0.99, 0.01],
|
|
124
|
+
[0.0, 1.0],
|
|
125
|
+
], dtype=np.float32)
|
|
126
|
+
|
|
127
|
+
sim = SimilarityGraphBuilder(threshold=0.9, metric="cosine", loop=False)
|
|
128
|
+
edge_index, edge_attr = sim(x)
|
|
129
|
+
|
|
130
|
+
edges = set(zip(edge_index[0].tolist(), edge_index[1].tolist()))
|
|
131
|
+
assert (0, 1) in edges and (1, 0) in edges
|
|
132
|
+
assert (0, 2) not in edges
|
|
133
|
+
|
|
134
|
+
|
|
135
|
+
def test_shared_entity_graph_builder():
|
|
136
|
+
df = {
|
|
137
|
+
"user_id": ["u1", "u2", "u3", "u4"],
|
|
138
|
+
"device_id": ["d1", "d1", "d2", "d2"],
|
|
139
|
+
}
|
|
140
|
+
x = np.zeros((4, 2), dtype=np.float32)
|
|
141
|
+
|
|
142
|
+
builder = SharedEntityGraphBuilder(entity_cols=["device_id"], loop=False)
|
|
143
|
+
edge_index, edge_attr = builder(x, df_or_dict=df)
|
|
144
|
+
|
|
145
|
+
edges = set(zip(edge_index[0].tolist(), edge_index[1].tolist()))
|
|
146
|
+
# u1 (0) and u2 (1) share d1
|
|
147
|
+
assert (0, 1) in edges and (1, 0) in edges
|
|
148
|
+
# u3 (2) and u4 (3) share d2
|
|
149
|
+
assert (2, 3) in edges and (3, 2) in edges
|
|
150
|
+
# u1 (0) and u3 (2) do not share
|
|
151
|
+
assert (0, 2) not in edges
|
|
152
|
+
|
|
153
|
+
|
|
154
|
+
def test_sequential_graph_builder():
|
|
155
|
+
df = {
|
|
156
|
+
"timestamp": [100, 102, 101, 200, 201],
|
|
157
|
+
"session": ["s1", "s1", "s1", "s2", "s2"],
|
|
158
|
+
}
|
|
159
|
+
x = np.zeros((5, 2), dtype=np.float32)
|
|
160
|
+
|
|
161
|
+
seq = SequentialGraphBuilder(order_col="timestamp", group_by_col="session", window_size=1)
|
|
162
|
+
edge_index, edge_attr = seq(x, df_or_dict=df)
|
|
163
|
+
|
|
164
|
+
edges = set(zip(edge_index[0].tolist(), edge_index[1].tolist()))
|
|
165
|
+
# s1 sequence ordered: 0 (ts=100) -> 2 (ts=101) -> 1 (ts=102)
|
|
166
|
+
assert (0, 2) in edges
|
|
167
|
+
assert (2, 1) in edges
|
|
168
|
+
# s2 sequence: 3 (ts=200) -> 4 (ts=201)
|
|
169
|
+
assert (3, 4) in edges
|
|
170
|
+
# Cross session edges must not exist
|
|
171
|
+
assert (1, 3) not in edges
|
|
172
|
+
|
|
173
|
+
|
|
174
|
+
# ==============================================================================
|
|
175
|
+
# 3. TableToGraph (Single-Table ETL) Tests
|
|
176
|
+
# ==============================================================================
|
|
177
|
+
|
|
178
|
+
def test_table_to_graph_dataframe():
|
|
179
|
+
df = pd.DataFrame({
|
|
180
|
+
"customer_id": ["c1", "c2", "c3", "c4", "c5", "c6"],
|
|
181
|
+
"age": [20, 25, 30, 45, 50, 55],
|
|
182
|
+
"spend": [100.0, 150.0, 200.0, 500.0, 550.0, 600.0],
|
|
183
|
+
"category": ["retail", "retail", "tech", "retail", "tech", "tech"],
|
|
184
|
+
"churn": [0, 0, 0, 1, 1, 1],
|
|
185
|
+
})
|
|
186
|
+
|
|
187
|
+
etl = TableToGraph(
|
|
188
|
+
target_col="churn",
|
|
189
|
+
id_col="customer_id",
|
|
190
|
+
edge_strategy="knn",
|
|
191
|
+
edge_kwargs={"k": 2},
|
|
192
|
+
train_ratio=0.5,
|
|
193
|
+
test_ratio=0.5,
|
|
194
|
+
)
|
|
195
|
+
data = etl.fit_transform(df)
|
|
196
|
+
|
|
197
|
+
assert isinstance(data, Data)
|
|
198
|
+
assert data.x.shape[0] == 6
|
|
199
|
+
assert data.edge_index.shape[0] == 2
|
|
200
|
+
assert data.edge_index.shape[1] > 0
|
|
201
|
+
assert data.y.shape[0] == 6
|
|
202
|
+
assert data.train_mask.shape[0] == 6
|
|
203
|
+
assert data.test_mask.shape[0] == 6
|
|
204
|
+
assert len(data.node_ids) == 6
|
|
205
|
+
assert data.id_to_index["c1"] == 0
|
|
206
|
+
|
|
207
|
+
|
|
208
|
+
def test_table_to_graph_functional():
|
|
209
|
+
data_dict = {
|
|
210
|
+
"feat1": [1.0, 2.0, 3.0, 4.0],
|
|
211
|
+
"feat2": [4.0, 3.0, 2.0, 1.0],
|
|
212
|
+
"label": ["A", "B", "A", "B"],
|
|
213
|
+
}
|
|
214
|
+
data = table_to_graph(data_dict, target_col="label", edge_strategy="knn", edge_kwargs={"k": 1})
|
|
215
|
+
assert isinstance(data, Data)
|
|
216
|
+
assert data.x.shape == (4, 2)
|
|
217
|
+
assert data.y.dtype == np.int64
|
|
218
|
+
|
|
219
|
+
|
|
220
|
+
def test_table_to_graph_from_csv():
|
|
221
|
+
with tempfile.TemporaryDirectory() as tmpdir:
|
|
222
|
+
csv_path = os.path.join(tmpdir, "sample.csv")
|
|
223
|
+
df = pd.DataFrame({
|
|
224
|
+
"id": ["n1", "n2", "n3", "n4"],
|
|
225
|
+
"v1": [10.0, 20.0, 30.0, 40.0],
|
|
226
|
+
"v2": [1.0, 2.0, 3.0, 4.0],
|
|
227
|
+
"y": [0, 1, 0, 1],
|
|
228
|
+
})
|
|
229
|
+
df.to_csv(csv_path, index=False)
|
|
230
|
+
|
|
231
|
+
data = TableToGraph.from_csv(csv_path, target_col="y", id_col="id", edge_strategy="knn", edge_kwargs={"k": 1})
|
|
232
|
+
assert isinstance(data, Data)
|
|
233
|
+
assert data.x.shape[0] == 4
|
|
234
|
+
assert data.node_ids == ["n1", "n2", "n3", "n4"]
|
|
235
|
+
|
|
236
|
+
|
|
237
|
+
# ==============================================================================
|
|
238
|
+
# 4. RelationalToGraph (Multi-Table ETL) Tests
|
|
239
|
+
# ==============================================================================
|
|
240
|
+
|
|
241
|
+
def test_relational_to_graph():
|
|
242
|
+
users_df = pd.DataFrame({
|
|
243
|
+
"user_id": ["u101", "u102", "u103"],
|
|
244
|
+
"age": [22, 35, 48],
|
|
245
|
+
"segment": ["bronze", "gold", "silver"],
|
|
246
|
+
})
|
|
247
|
+
products_df = pd.DataFrame({
|
|
248
|
+
"prod_id": ["p1", "p2", "p3", "p4"],
|
|
249
|
+
"price": [9.99, 49.99, 19.99, 99.99],
|
|
250
|
+
"department": ["books", "elec", "books", "elec"],
|
|
251
|
+
})
|
|
252
|
+
purchases_df = pd.DataFrame({
|
|
253
|
+
"u_id": ["u101", "u101", "u102", "u103", "unknown_user"],
|
|
254
|
+
"p_id": ["p1", "p2", "p3", "p4", "p1"],
|
|
255
|
+
"rating": [5.0, 4.0, 5.0, 3.0, 1.0],
|
|
256
|
+
})
|
|
257
|
+
|
|
258
|
+
etl = RelationalToGraph(
|
|
259
|
+
id_cols={"user": "user_id", "product": "prod_id"},
|
|
260
|
+
edge_cols={("user", "buys", "product"): ("u_id", "p_id")},
|
|
261
|
+
edge_attr_cols={("user", "buys", "product"): ["rating"]},
|
|
262
|
+
)
|
|
263
|
+
|
|
264
|
+
hetero_data = etl.fit_transform(
|
|
265
|
+
nodes={"user": users_df, "product": products_df},
|
|
266
|
+
edges={("user", "buys", "product"): purchases_df},
|
|
267
|
+
)
|
|
268
|
+
|
|
269
|
+
assert isinstance(hetero_data, HeteroData)
|
|
270
|
+
# Check node tables
|
|
271
|
+
assert hetero_data["user"].x.shape[0] == 3
|
|
272
|
+
assert hetero_data["product"].x.shape[0] == 4
|
|
273
|
+
|
|
274
|
+
# Check edges (unknown_user should be filtered out)
|
|
275
|
+
edge_index = hetero_data["user", "buys", "product"].edge_index
|
|
276
|
+
assert edge_index.shape[0] == 2
|
|
277
|
+
assert edge_index.shape[1] == 4 # 4 valid purchases
|
|
278
|
+
|
|
279
|
+
# Check edge attributes
|
|
280
|
+
edge_attr = hetero_data["user", "buys", "product"].edge_attr
|
|
281
|
+
assert edge_attr.shape == (4, 1)
|
|
282
|
+
|
|
283
|
+
# Check ID mapping consistency
|
|
284
|
+
assert hetero_data.id_maps["user"]["u101"] == 0
|
|
285
|
+
assert hetero_data.inverse_id_maps["user"][0] == "u101"
|
|
286
|
+
|
|
287
|
+
|
|
288
|
+
# ==============================================================================
|
|
289
|
+
# 5. End-to-End ETL to GNN Training Integration
|
|
290
|
+
# ==============================================================================
|
|
291
|
+
|
|
292
|
+
def test_etl_to_node_classifier_integration():
|
|
293
|
+
"""Verify that a DataFrame converted via TableToGraph can directly train a NodeClassifier."""
|
|
294
|
+
df = pd.DataFrame({
|
|
295
|
+
"feat_a": np.random.randn(30).astype("float32"),
|
|
296
|
+
"feat_b": np.random.randn(30).astype("float32"),
|
|
297
|
+
"category": (["cat1", "cat2", "cat3"] * 10),
|
|
298
|
+
"target": (np.random.randint(0, 2, size=30)),
|
|
299
|
+
})
|
|
300
|
+
|
|
301
|
+
# 1. ETL: Tabular -> Graph Data
|
|
302
|
+
data = table_to_graph(
|
|
303
|
+
df,
|
|
304
|
+
target_col="target",
|
|
305
|
+
edge_strategy="knn",
|
|
306
|
+
edge_kwargs={"k": 3},
|
|
307
|
+
train_ratio=0.7,
|
|
308
|
+
test_ratio=0.3,
|
|
309
|
+
)
|
|
310
|
+
|
|
311
|
+
# 2. Train NodeClassifier in 3 lines of code!
|
|
312
|
+
clf = NodeClassifier(backbone="gcn", hidden_channels=16, num_layers=2)
|
|
313
|
+
clf.fit(data, epochs=2, verbose=0)
|
|
314
|
+
|
|
315
|
+
# 3. Predict & evaluate
|
|
316
|
+
metrics = clf.evaluate(data, mask="test_mask")
|
|
317
|
+
assert "accuracy" in metrics
|
|
318
|
+
assert 0.0 <= metrics["accuracy"] <= 1.0
|
|
@@ -0,0 +1,15 @@
|
|
|
1
|
+
"""K3-Node turn-key export and serving module for ONNX, TensorRT, and TensorFlow Lite."""
|
|
2
|
+
|
|
3
|
+
from k3_node.export.onnx_exporter import export_onnx
|
|
4
|
+
from k3_node.export.tflite_exporter import export_tflite
|
|
5
|
+
from k3_node.export.tensorrt_exporter import export_tensorrt, generate_triton_config
|
|
6
|
+
from k3_node.export.runtime import ONNXModel, TFLiteModel
|
|
7
|
+
|
|
8
|
+
__all__ = [
|
|
9
|
+
"export_onnx",
|
|
10
|
+
"export_tflite",
|
|
11
|
+
"export_tensorrt",
|
|
12
|
+
"generate_triton_config",
|
|
13
|
+
"ONNXModel",
|
|
14
|
+
"TFLiteModel",
|
|
15
|
+
]
|