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,270 @@
|
|
|
1
|
+
"""Graph topology construction strategies from tabular data."""
|
|
2
|
+
|
|
3
|
+
from typing import Any, Dict, List, Optional, Sequence, Tuple, Union
|
|
4
|
+
import numpy as np
|
|
5
|
+
|
|
6
|
+
|
|
7
|
+
class KNNGraphBuilder:
|
|
8
|
+
r"""Constructs a k-nearest-neighbors graph from a node feature matrix.
|
|
9
|
+
|
|
10
|
+
Args:
|
|
11
|
+
k: Number of nearest neighbors per node. (default: ``5``)
|
|
12
|
+
metric: Distance metric (``"cosine"``, ``"euclidean"``, or ``"manhattan"``).
|
|
13
|
+
(default: ``"cosine"``)
|
|
14
|
+
loop: Whether to include self-loops. (default: ``False``)
|
|
15
|
+
bidirectional: Whether to make the resulting graph undirected. (default: ``True``)
|
|
16
|
+
"""
|
|
17
|
+
|
|
18
|
+
def __init__(
|
|
19
|
+
self,
|
|
20
|
+
k: int = 5,
|
|
21
|
+
metric: str = "cosine",
|
|
22
|
+
loop: bool = False,
|
|
23
|
+
bidirectional: bool = True,
|
|
24
|
+
):
|
|
25
|
+
self.k = k
|
|
26
|
+
self.metric = metric.lower()
|
|
27
|
+
self.loop = loop
|
|
28
|
+
self.bidirectional = bidirectional
|
|
29
|
+
|
|
30
|
+
def __call__(self, x: np.ndarray, **kwargs) -> Tuple[np.ndarray, Optional[np.ndarray]]:
|
|
31
|
+
num_nodes = x.shape[0]
|
|
32
|
+
if num_nodes == 0:
|
|
33
|
+
return np.empty((2, 0), dtype=np.int64), np.empty((0, 1), dtype=np.float32)
|
|
34
|
+
|
|
35
|
+
actual_k = min(self.k if self.loop else self.k + 1, num_nodes)
|
|
36
|
+
|
|
37
|
+
if self.metric == "cosine":
|
|
38
|
+
norms = np.linalg.norm(x, axis=1, keepdims=True)
|
|
39
|
+
norms[norms < 1e-8] = 1.0
|
|
40
|
+
x_norm = x / norms
|
|
41
|
+
sim_matrix = np.dot(x_norm, x_norm.T)
|
|
42
|
+
# Higher similarity is closer
|
|
43
|
+
dist_matrix = 1.0 - sim_matrix
|
|
44
|
+
elif self.metric == "euclidean":
|
|
45
|
+
diff = x[:, np.newaxis, :] - x[np.newaxis, :, :]
|
|
46
|
+
dist_matrix = np.sqrt(np.sum(diff ** 2, axis=-1))
|
|
47
|
+
elif self.metric == "manhattan":
|
|
48
|
+
diff = x[:, np.newaxis, :] - x[np.newaxis, :, :]
|
|
49
|
+
dist_matrix = np.sum(np.abs(diff), axis=-1)
|
|
50
|
+
else:
|
|
51
|
+
raise ValueError(f"Unknown metric '{self.metric}'. Supported: 'cosine', 'euclidean', 'manhattan'.")
|
|
52
|
+
|
|
53
|
+
src_list = []
|
|
54
|
+
dst_list = []
|
|
55
|
+
dist_list = []
|
|
56
|
+
|
|
57
|
+
for i in range(num_nodes):
|
|
58
|
+
row_dists = dist_matrix[i]
|
|
59
|
+
nearest_indices = np.argsort(row_dists)
|
|
60
|
+
count = 0
|
|
61
|
+
for neighbor in nearest_indices:
|
|
62
|
+
if not self.loop and neighbor == i:
|
|
63
|
+
continue
|
|
64
|
+
src_list.append(i)
|
|
65
|
+
dst_list.append(neighbor)
|
|
66
|
+
dist_list.append(row_dists[neighbor])
|
|
67
|
+
count += 1
|
|
68
|
+
if count >= self.k:
|
|
69
|
+
break
|
|
70
|
+
|
|
71
|
+
if self.bidirectional:
|
|
72
|
+
# Add reverse edges if not already present
|
|
73
|
+
edge_set = set(zip(src_list, dst_list))
|
|
74
|
+
for s, d, w in list(zip(src_list, dst_list, dist_list)):
|
|
75
|
+
if (d, s) not in edge_set:
|
|
76
|
+
src_list.append(d)
|
|
77
|
+
dst_list.append(s)
|
|
78
|
+
dist_list.append(w)
|
|
79
|
+
edge_set.add((d, s))
|
|
80
|
+
|
|
81
|
+
edge_index = np.array([src_list, dst_list], dtype=np.int64)
|
|
82
|
+
edge_attr = np.array(dist_list, dtype=np.float32).reshape(-1, 1) if dist_list else np.empty((0, 1), dtype=np.float32)
|
|
83
|
+
return edge_index, edge_attr
|
|
84
|
+
|
|
85
|
+
|
|
86
|
+
class SimilarityGraphBuilder:
|
|
87
|
+
r"""Constructs a graph connecting node pairs whose pairwise similarity exceeds a threshold.
|
|
88
|
+
|
|
89
|
+
Args:
|
|
90
|
+
threshold: Minimum similarity required to create an edge. (default: ``0.7``)
|
|
91
|
+
metric: Similarity function (``"cosine"`` or ``"rbf"``). (default: ``"cosine"``)
|
|
92
|
+
gamma: Bandwidth parameter for RBF kernel. (default: ``1.0``)
|
|
93
|
+
loop: Whether to include self-loops. (default: ``False``)
|
|
94
|
+
"""
|
|
95
|
+
|
|
96
|
+
def __init__(
|
|
97
|
+
self,
|
|
98
|
+
threshold: float = 0.7,
|
|
99
|
+
metric: str = "cosine",
|
|
100
|
+
gamma: float = 1.0,
|
|
101
|
+
loop: bool = False,
|
|
102
|
+
):
|
|
103
|
+
self.threshold = threshold
|
|
104
|
+
self.metric = metric.lower()
|
|
105
|
+
self.gamma = gamma
|
|
106
|
+
self.loop = loop
|
|
107
|
+
|
|
108
|
+
def __call__(self, x: np.ndarray, **kwargs) -> Tuple[np.ndarray, Optional[np.ndarray]]:
|
|
109
|
+
num_nodes = x.shape[0]
|
|
110
|
+
if num_nodes == 0:
|
|
111
|
+
return np.empty((2, 0), dtype=np.int64), np.empty((0, 1), dtype=np.float32)
|
|
112
|
+
|
|
113
|
+
if self.metric == "cosine":
|
|
114
|
+
norms = np.linalg.norm(x, axis=1, keepdims=True)
|
|
115
|
+
norms[norms < 1e-8] = 1.0
|
|
116
|
+
x_norm = x / norms
|
|
117
|
+
sim_matrix = np.dot(x_norm, x_norm.T)
|
|
118
|
+
elif self.metric == "rbf":
|
|
119
|
+
diff = x[:, np.newaxis, :] - x[np.newaxis, :, :]
|
|
120
|
+
sq_dist = np.sum(diff ** 2, axis=-1)
|
|
121
|
+
sim_matrix = np.exp(-self.gamma * sq_dist)
|
|
122
|
+
else:
|
|
123
|
+
raise ValueError(f"Unknown metric '{self.metric}'. Supported: 'cosine', 'rbf'.")
|
|
124
|
+
|
|
125
|
+
if not self.loop:
|
|
126
|
+
np.fill_diagonal(sim_matrix, -np.inf)
|
|
127
|
+
|
|
128
|
+
src, dst = np.where(sim_matrix >= self.threshold)
|
|
129
|
+
weights = sim_matrix[src, dst].astype(np.float32).reshape(-1, 1)
|
|
130
|
+
edge_index = np.stack([src, dst], axis=0).astype(np.int64)
|
|
131
|
+
return edge_index, weights
|
|
132
|
+
|
|
133
|
+
|
|
134
|
+
class SharedEntityGraphBuilder:
|
|
135
|
+
r"""Connects tabular rows that share one or more categorical identifier values.
|
|
136
|
+
(e.g., users sharing the same IP address, device, category, or cluster).
|
|
137
|
+
|
|
138
|
+
Args:
|
|
139
|
+
entity_cols: List of column names to check for shared values.
|
|
140
|
+
max_degree: Maximum number of neighbors created per shared entity (to avoid supernode explosion).
|
|
141
|
+
(default: ``50``)
|
|
142
|
+
loop: Whether to include self-loops. (default: ``False``)
|
|
143
|
+
"""
|
|
144
|
+
|
|
145
|
+
def __init__(
|
|
146
|
+
self,
|
|
147
|
+
entity_cols: Sequence[str],
|
|
148
|
+
max_degree: int = 50,
|
|
149
|
+
loop: bool = False,
|
|
150
|
+
):
|
|
151
|
+
self.entity_cols = list(entity_cols)
|
|
152
|
+
self.max_degree = max_degree
|
|
153
|
+
self.loop = loop
|
|
154
|
+
|
|
155
|
+
def __call__(self, x: np.ndarray, df_or_dict: Optional[Any] = None, **kwargs) -> Tuple[np.ndarray, Optional[np.ndarray]]:
|
|
156
|
+
if df_or_dict is None:
|
|
157
|
+
raise ValueError("SharedEntityGraphBuilder requires 'df_or_dict' containing the entity columns.")
|
|
158
|
+
|
|
159
|
+
num_nodes = x.shape[0]
|
|
160
|
+
src_list = []
|
|
161
|
+
dst_list = []
|
|
162
|
+
|
|
163
|
+
from k3_node.etl.encoders import _get_column_values
|
|
164
|
+
|
|
165
|
+
for col in self.entity_cols:
|
|
166
|
+
vals = _get_column_values(df_or_dict, col)
|
|
167
|
+
val_to_rows: Dict[Any, List[int]] = {}
|
|
168
|
+
for row_idx, val in enumerate(vals):
|
|
169
|
+
if val is None or (isinstance(val, float) and np.isnan(val)) or val == "":
|
|
170
|
+
continue
|
|
171
|
+
val_to_rows.setdefault(val, []).append(row_idx)
|
|
172
|
+
|
|
173
|
+
for val, rows in val_to_rows.items():
|
|
174
|
+
if len(rows) > self.max_degree:
|
|
175
|
+
# Subsample if group is too large
|
|
176
|
+
sampled_rows = np.random.choice(rows, size=self.max_degree, replace=False).tolist()
|
|
177
|
+
else:
|
|
178
|
+
sampled_rows = rows
|
|
179
|
+
|
|
180
|
+
for i in sampled_rows:
|
|
181
|
+
for j in sampled_rows:
|
|
182
|
+
if not self.loop and i == j:
|
|
183
|
+
continue
|
|
184
|
+
src_list.append(i)
|
|
185
|
+
dst_list.append(j)
|
|
186
|
+
|
|
187
|
+
if not src_list:
|
|
188
|
+
return np.empty((2, 0), dtype=np.int64), np.empty((0, 1), dtype=np.float32)
|
|
189
|
+
|
|
190
|
+
edges = list(set(zip(src_list, dst_list)))
|
|
191
|
+
src_arr = np.array([e[0] for e in edges], dtype=np.int64)
|
|
192
|
+
dst_arr = np.array([e[1] for e in edges], dtype=np.int64)
|
|
193
|
+
edge_index = np.stack([src_arr, dst_arr], axis=0)
|
|
194
|
+
edge_attr = np.ones((len(edges), 1), dtype=np.float32)
|
|
195
|
+
return edge_index, edge_attr
|
|
196
|
+
|
|
197
|
+
|
|
198
|
+
class SequentialGraphBuilder:
|
|
199
|
+
r"""Connects tabular rows sequentially in order of an index or timestamp column,
|
|
200
|
+
optionally partitioned by a group entity column.
|
|
201
|
+
|
|
202
|
+
Args:
|
|
203
|
+
order_col: Optional column name used to sort rows (e.g. timestamp or sequence index).
|
|
204
|
+
group_by_col: Optional column name to partition sequences (e.g. user_id or session_id).
|
|
205
|
+
window_size: Number of forward/backward sequential steps to connect. (default: ``1``)
|
|
206
|
+
bidirectional: Whether to create undirected edges. (default: ``True``)
|
|
207
|
+
"""
|
|
208
|
+
|
|
209
|
+
def __init__(
|
|
210
|
+
self,
|
|
211
|
+
order_col: Optional[str] = None,
|
|
212
|
+
group_by_col: Optional[str] = None,
|
|
213
|
+
window_size: int = 1,
|
|
214
|
+
bidirectional: bool = True,
|
|
215
|
+
):
|
|
216
|
+
self.order_col = order_col
|
|
217
|
+
self.group_by_col = group_by_col
|
|
218
|
+
self.window_size = window_size
|
|
219
|
+
self.bidirectional = bidirectional
|
|
220
|
+
|
|
221
|
+
def __call__(self, x: np.ndarray, df_or_dict: Optional[Any] = None, **kwargs) -> Tuple[np.ndarray, Optional[np.ndarray]]:
|
|
222
|
+
num_nodes = x.shape[0]
|
|
223
|
+
if num_nodes == 0:
|
|
224
|
+
return np.empty((2, 0), dtype=np.int64), np.empty((0, 1), dtype=np.float32)
|
|
225
|
+
|
|
226
|
+
from k3_node.etl.encoders import _get_column_values
|
|
227
|
+
|
|
228
|
+
if self.group_by_col is not None and df_or_dict is not None:
|
|
229
|
+
groups = _get_column_values(df_or_dict, self.group_by_col)
|
|
230
|
+
group_to_indices: Dict[Any, List[int]] = {}
|
|
231
|
+
for idx, g in enumerate(groups):
|
|
232
|
+
group_to_indices.setdefault(g, []).append(idx)
|
|
233
|
+
else:
|
|
234
|
+
group_to_indices = {"all": list(range(num_nodes))}
|
|
235
|
+
|
|
236
|
+
if self.order_col is not None and df_or_dict is not None:
|
|
237
|
+
order_vals = _get_column_values(df_or_dict, self.order_col)
|
|
238
|
+
else:
|
|
239
|
+
order_vals = None
|
|
240
|
+
|
|
241
|
+
src_list = []
|
|
242
|
+
dst_list = []
|
|
243
|
+
|
|
244
|
+
for group_name, row_indices in group_to_indices.items():
|
|
245
|
+
if order_vals is not None:
|
|
246
|
+
sorted_indices = sorted(row_indices, key=lambda idx: order_vals[idx])
|
|
247
|
+
else:
|
|
248
|
+
sorted_indices = row_indices
|
|
249
|
+
|
|
250
|
+
n_seq = len(sorted_indices)
|
|
251
|
+
for i in range(n_seq):
|
|
252
|
+
curr_node = sorted_indices[i]
|
|
253
|
+
for step in range(1, self.window_size + 1):
|
|
254
|
+
if i + step < n_seq:
|
|
255
|
+
next_node = sorted_indices[i + step]
|
|
256
|
+
src_list.append(curr_node)
|
|
257
|
+
dst_list.append(next_node)
|
|
258
|
+
if self.bidirectional:
|
|
259
|
+
src_list.append(next_node)
|
|
260
|
+
dst_list.append(curr_node)
|
|
261
|
+
|
|
262
|
+
if not src_list:
|
|
263
|
+
return np.empty((2, 0), dtype=np.int64), np.empty((0, 1), dtype=np.float32)
|
|
264
|
+
|
|
265
|
+
edges = list(set(zip(src_list, dst_list)))
|
|
266
|
+
src_arr = np.array([e[0] for e in edges], dtype=np.int64)
|
|
267
|
+
dst_arr = np.array([e[1] for e in edges], dtype=np.int64)
|
|
268
|
+
edge_index = np.stack([src_arr, dst_arr], axis=0)
|
|
269
|
+
edge_attr = np.ones((len(edges), 1), dtype=np.float32)
|
|
270
|
+
return edge_index, edge_attr
|
|
@@ -0,0 +1,201 @@
|
|
|
1
|
+
"""Relational (multi-table) to Heterogeneous Graph ETL converter."""
|
|
2
|
+
|
|
3
|
+
from typing import Any, Dict, List, Optional, Sequence, Tuple, Union
|
|
4
|
+
import numpy as np
|
|
5
|
+
|
|
6
|
+
from k3_node.data import HeteroData, Data
|
|
7
|
+
from k3_node.etl.encoders import TabularEncoder, _get_column_names, _get_column_values
|
|
8
|
+
|
|
9
|
+
|
|
10
|
+
NodeType = str
|
|
11
|
+
EdgeType = Tuple[str, str, str]
|
|
12
|
+
|
|
13
|
+
|
|
14
|
+
class RelationalToGraph:
|
|
15
|
+
r"""ETL pipeline converting multi-table relational databases / DataFrames into
|
|
16
|
+
a heterogeneous graph :class:`k3_node.data.HeteroData` object.
|
|
17
|
+
|
|
18
|
+
Args:
|
|
19
|
+
id_cols (dict): Mapping from :obj:`NodeType` to the primary key column name.
|
|
20
|
+
(e.g., ``{"user": "user_id", "movie": "movie_id"}``).
|
|
21
|
+
edge_cols (dict): Mapping from :obj:`EdgeType` to the tuple of foreign key column names
|
|
22
|
+
``(source_id_col, target_id_col)``.
|
|
23
|
+
(e.g., ``{("user", "rates", "movie"): ("user_id", "movie_id")}``).
|
|
24
|
+
feature_cols (dict, optional): Mapping from :obj:`NodeType` to a list of feature columns.
|
|
25
|
+
If omitted, all columns except the primary key are encoded.
|
|
26
|
+
edge_attr_cols (dict, optional): Mapping from :obj:`EdgeType` to a list of edge attribute columns.
|
|
27
|
+
node_target_cols (dict, optional): Mapping from :obj:`NodeType` to the target label column.
|
|
28
|
+
edge_target_cols (dict, optional): Mapping from :obj:`EdgeType` to the target edge label column.
|
|
29
|
+
"""
|
|
30
|
+
|
|
31
|
+
def __init__(
|
|
32
|
+
self,
|
|
33
|
+
id_cols: Dict[NodeType, str],
|
|
34
|
+
edge_cols: Dict[EdgeType, Tuple[str, str]],
|
|
35
|
+
feature_cols: Optional[Dict[NodeType, List[str]]] = None,
|
|
36
|
+
edge_attr_cols: Optional[Dict[EdgeType, List[str]]] = None,
|
|
37
|
+
node_target_cols: Optional[Dict[NodeType, str]] = None,
|
|
38
|
+
edge_target_cols: Optional[Dict[EdgeType, str]] = None,
|
|
39
|
+
):
|
|
40
|
+
self.id_cols = id_cols
|
|
41
|
+
self.edge_cols = edge_cols
|
|
42
|
+
self.feature_cols = feature_cols or {}
|
|
43
|
+
self.edge_attr_cols = edge_attr_cols or {}
|
|
44
|
+
self.node_target_cols = node_target_cols or {}
|
|
45
|
+
self.edge_target_cols = edge_target_cols or {}
|
|
46
|
+
|
|
47
|
+
self.node_encoders_: Dict[NodeType, TabularEncoder] = {}
|
|
48
|
+
self.edge_encoders_: Dict[EdgeType, TabularEncoder] = {}
|
|
49
|
+
self.id_maps_: Dict[NodeType, Dict[Any, int]] = {}
|
|
50
|
+
self.inverse_id_maps_: Dict[NodeType, Dict[int, Any]] = {}
|
|
51
|
+
|
|
52
|
+
def fit(self, nodes: Dict[NodeType, Any], edges: Optional[Dict[EdgeType, Any]] = None):
|
|
53
|
+
r"""Fits encoders and builds entity ID mappings across all tables."""
|
|
54
|
+
# 1. Map node IDs and fit node feature encoders
|
|
55
|
+
for node_type, table in nodes.items():
|
|
56
|
+
id_col = self.id_cols[node_type]
|
|
57
|
+
raw_ids = _get_column_values(table, id_col)
|
|
58
|
+
|
|
59
|
+
# Unique contiguous ID mapping
|
|
60
|
+
unique_ids = []
|
|
61
|
+
seen = set()
|
|
62
|
+
for rid in raw_ids:
|
|
63
|
+
if rid not in seen:
|
|
64
|
+
seen.add(rid)
|
|
65
|
+
unique_ids.append(rid)
|
|
66
|
+
|
|
67
|
+
id_map = {rid: i for i, rid in enumerate(unique_ids)}
|
|
68
|
+
inv_map = {i: rid for i, rid in enumerate(unique_ids)}
|
|
69
|
+
self.id_maps_[node_type] = id_map
|
|
70
|
+
self.inverse_id_maps_[node_type] = inv_map
|
|
71
|
+
|
|
72
|
+
# Feature columns
|
|
73
|
+
all_cols = _get_column_names(table)
|
|
74
|
+
target_col = self.node_target_cols.get(node_type)
|
|
75
|
+
ignore_cols = {id_col}
|
|
76
|
+
if target_col:
|
|
77
|
+
ignore_cols.add(target_col)
|
|
78
|
+
|
|
79
|
+
if node_type in self.feature_cols:
|
|
80
|
+
feat_cols = [c for c in self.feature_cols[node_type] if c in all_cols and c not in ignore_cols]
|
|
81
|
+
else:
|
|
82
|
+
feat_cols = [c for c in all_cols if c not in ignore_cols]
|
|
83
|
+
|
|
84
|
+
encoder = TabularEncoder()
|
|
85
|
+
if feat_cols:
|
|
86
|
+
encoder.fit(table, columns=feat_cols)
|
|
87
|
+
self.node_encoders_[node_type] = encoder
|
|
88
|
+
|
|
89
|
+
# 2. Fit edge encoders if edge attributes are specified
|
|
90
|
+
if edges:
|
|
91
|
+
for edge_type, table in edges.items():
|
|
92
|
+
if edge_type in self.edge_attr_cols:
|
|
93
|
+
attr_cols = self.edge_attr_cols[edge_type]
|
|
94
|
+
edge_enc = TabularEncoder()
|
|
95
|
+
edge_enc.fit(table, columns=attr_cols)
|
|
96
|
+
self.edge_encoders_[edge_type] = edge_enc
|
|
97
|
+
|
|
98
|
+
return self
|
|
99
|
+
|
|
100
|
+
def transform(
|
|
101
|
+
self,
|
|
102
|
+
nodes: Dict[NodeType, Any],
|
|
103
|
+
edges: Optional[Dict[EdgeType, Any]] = None,
|
|
104
|
+
) -> HeteroData:
|
|
105
|
+
r"""Constructs a :class:`k3_node.data.HeteroData` instance from relational tables."""
|
|
106
|
+
hetero_data = HeteroData()
|
|
107
|
+
|
|
108
|
+
# 1. Process node tables
|
|
109
|
+
for node_type, table in nodes.items():
|
|
110
|
+
encoder = self.node_encoders_[node_type]
|
|
111
|
+
if encoder.column_order_:
|
|
112
|
+
x = encoder.transform(table)
|
|
113
|
+
hetero_data[node_type].x = x
|
|
114
|
+
else:
|
|
115
|
+
num_nodes = len(self.id_maps_[node_type])
|
|
116
|
+
hetero_data[node_type].num_nodes = num_nodes
|
|
117
|
+
|
|
118
|
+
# Target labels y
|
|
119
|
+
if node_type in self.node_target_cols:
|
|
120
|
+
target_col = self.node_target_cols[node_type]
|
|
121
|
+
raw_y = _get_column_values(table, target_col)
|
|
122
|
+
if all(isinstance(v, (int, np.integer)) for v in raw_y if v is not None):
|
|
123
|
+
hetero_data[node_type].y = np.array(raw_y, dtype=np.int64)
|
|
124
|
+
elif all(isinstance(v, (float, int, np.floating, np.integer)) for v in raw_y if v is not None):
|
|
125
|
+
hetero_data[node_type].y = np.array(raw_y, dtype=np.float32)
|
|
126
|
+
else:
|
|
127
|
+
unique_c = sorted(list(set(raw_y)))
|
|
128
|
+
mapping = {c: i for i, c in enumerate(unique_c)}
|
|
129
|
+
hetero_data[node_type].y = np.array([mapping[c] for c in raw_y], dtype=np.int64)
|
|
130
|
+
|
|
131
|
+
# 2. Process edge tables
|
|
132
|
+
if edges:
|
|
133
|
+
for edge_type, table in edges.items():
|
|
134
|
+
src_type, rel_name, dst_type = edge_type
|
|
135
|
+
src_col, dst_col = self.edge_cols[edge_type]
|
|
136
|
+
|
|
137
|
+
raw_srcs = _get_column_values(table, src_col)
|
|
138
|
+
raw_dsts = _get_column_values(table, dst_col)
|
|
139
|
+
|
|
140
|
+
src_map = self.id_maps_[src_type]
|
|
141
|
+
dst_map = self.id_maps_[dst_type]
|
|
142
|
+
|
|
143
|
+
valid_src = []
|
|
144
|
+
valid_dst = []
|
|
145
|
+
valid_indices = []
|
|
146
|
+
|
|
147
|
+
for row_idx, (s, d) in enumerate(zip(raw_srcs, raw_dsts)):
|
|
148
|
+
if s in src_map and d in dst_map:
|
|
149
|
+
valid_src.append(src_map[s])
|
|
150
|
+
valid_dst.append(dst_map[d])
|
|
151
|
+
valid_indices.append(row_idx)
|
|
152
|
+
|
|
153
|
+
edge_index = np.array([valid_src, valid_dst], dtype=np.int64)
|
|
154
|
+
hetero_data[edge_type].edge_index = edge_index
|
|
155
|
+
|
|
156
|
+
# Edge attributes
|
|
157
|
+
if edge_type in self.edge_encoders_:
|
|
158
|
+
edge_enc = self.edge_encoders_[edge_type]
|
|
159
|
+
full_attrs = edge_enc.transform(table)
|
|
160
|
+
hetero_data[edge_type].edge_attr = full_attrs[valid_indices]
|
|
161
|
+
|
|
162
|
+
# Edge target y
|
|
163
|
+
if edge_type in self.edge_target_cols:
|
|
164
|
+
t_col = self.edge_target_cols[edge_type]
|
|
165
|
+
raw_ey = _get_column_values(table, t_col)
|
|
166
|
+
filtered_ey = [raw_ey[idx] for idx in valid_indices]
|
|
167
|
+
hetero_data[edge_type].edge_label = np.array(filtered_ey, dtype=np.float32)
|
|
168
|
+
|
|
169
|
+
# Attach metadata
|
|
170
|
+
hetero_data.id_maps = self.id_maps_
|
|
171
|
+
hetero_data.inverse_id_maps = self.inverse_id_maps_
|
|
172
|
+
return hetero_data
|
|
173
|
+
|
|
174
|
+
def fit_transform(
|
|
175
|
+
self,
|
|
176
|
+
nodes: Dict[NodeType, Any],
|
|
177
|
+
edges: Optional[Dict[EdgeType, Any]] = None,
|
|
178
|
+
) -> HeteroData:
|
|
179
|
+
r"""Fits encoders and builds the heterogeneous graph in a single call."""
|
|
180
|
+
return self.fit(nodes, edges).transform(nodes, edges)
|
|
181
|
+
|
|
182
|
+
|
|
183
|
+
def relational_to_graph(
|
|
184
|
+
nodes: Dict[NodeType, Any],
|
|
185
|
+
edges: Optional[Dict[EdgeType, Any]] = None,
|
|
186
|
+
id_cols: Optional[Dict[NodeType, str]] = None,
|
|
187
|
+
edge_cols: Optional[Dict[EdgeType, Tuple[str, str]]] = None,
|
|
188
|
+
**kwargs,
|
|
189
|
+
) -> HeteroData:
|
|
190
|
+
r"""Functional shortcut to convert relational tables into a :class:`k3_node.data.HeteroData` graph."""
|
|
191
|
+
if id_cols is None:
|
|
192
|
+
raise ValueError("Must provide 'id_cols' mapping node types to their primary key columns.")
|
|
193
|
+
if edges and edge_cols is None:
|
|
194
|
+
raise ValueError("Must provide 'edge_cols' mapping edge types to (source_col, target_col).")
|
|
195
|
+
|
|
196
|
+
etl = RelationalToGraph(
|
|
197
|
+
id_cols=id_cols,
|
|
198
|
+
edge_cols=edge_cols or {},
|
|
199
|
+
**kwargs,
|
|
200
|
+
)
|
|
201
|
+
return etl.fit_transform(nodes, edges)
|