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,403 @@
|
|
|
1
|
+
import keras
|
|
2
|
+
from typing import Optional, Tuple
|
|
3
|
+
from keras import layers, ops
|
|
4
|
+
import numpy as np
|
|
5
|
+
from k3_node.ops.segment import segment_max, segment_sum
|
|
6
|
+
from k3_node.ops.creation import scatter, full
|
|
7
|
+
|
|
8
|
+
|
|
9
|
+
def ptr2index(ptr):
|
|
10
|
+
r"""Converts a pointer tensor into an index tensor.
|
|
11
|
+
|
|
12
|
+
Example:
|
|
13
|
+
```python
|
|
14
|
+
import numpy as np
|
|
15
|
+
from k3_node.layers import ptr2index
|
|
16
|
+
|
|
17
|
+
ptr = np.array([0, 3, 5]) # CSR pointer: set 0 has 3 elements, set 1 has 2
|
|
18
|
+
print(tuple(ptr2index(ptr).shape)) # (5,): set id of every element
|
|
19
|
+
```
|
|
20
|
+
"""
|
|
21
|
+
ptr_np = ops.convert_to_numpy(ptr).astype(np.int64)
|
|
22
|
+
counts = ptr_np[1:] - ptr_np[:-1]
|
|
23
|
+
index_np = np.repeat(np.arange(len(counts), dtype=np.int64), counts)
|
|
24
|
+
return ops.convert_to_tensor(index_np, dtype="int32")
|
|
25
|
+
|
|
26
|
+
|
|
27
|
+
def to_dense_batch(
|
|
28
|
+
x,
|
|
29
|
+
index,
|
|
30
|
+
dim_size: Optional[int] = None,
|
|
31
|
+
fill_value: float = 0.0,
|
|
32
|
+
max_num_elements: Optional[int] = None,
|
|
33
|
+
) -> Tuple[any, any]:
|
|
34
|
+
r"""Transforms a batched feature tensor into a dense representation
|
|
35
|
+
of shape `(batch_size, max_nodes, *dims)`.
|
|
36
|
+
|
|
37
|
+
Example:
|
|
38
|
+
```python
|
|
39
|
+
import numpy as np
|
|
40
|
+
from k3_node.layers import to_dense_batch
|
|
41
|
+
|
|
42
|
+
x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
|
|
43
|
+
batch = np.repeat([0, 1], 5) # nodes 0-4 belong to graph 0, nodes 5-9 to graph 1
|
|
44
|
+
|
|
45
|
+
x_dense, mask = to_dense_batch(x, batch) # [num_graphs, max_nodes, features] + validity mask
|
|
46
|
+
print(tuple(x_dense.shape), tuple(mask.shape)) # (2, 5, 8) (2, 5)
|
|
47
|
+
```
|
|
48
|
+
"""
|
|
49
|
+
from k3_node.layers.conv.utils import is_tracing
|
|
50
|
+
from k3_node.ops.host import _in_shape_inference
|
|
51
|
+
|
|
52
|
+
# Keras' shape inference traces too, but only needs shapes: the host path then sees zeros.
|
|
53
|
+
if not is_tracing(index) or _in_shape_inference():
|
|
54
|
+
try:
|
|
55
|
+
# `index` is purely structural (never differentiated), so plain
|
|
56
|
+
# numpy is fine for it. `x` itself is placed into the dense
|
|
57
|
+
# tensor with the differentiable `ops.scatter` below -- a numpy
|
|
58
|
+
# round-trip on `x` would silently detach it from the graph and
|
|
59
|
+
# stop gradients from flowing back into whatever produced it.
|
|
60
|
+
from k3_node.ops.host import to_numpy # zeros during Keras' shape inference
|
|
61
|
+
|
|
62
|
+
index_np = np.asarray(to_numpy(index)).astype(np.int64)
|
|
63
|
+
N = len(index_np)
|
|
64
|
+
|
|
65
|
+
B = int(np.max(index_np)) + 1 if N > 0 else 0
|
|
66
|
+
if dim_size is not None:
|
|
67
|
+
B = max(B, int(dim_size))
|
|
68
|
+
|
|
69
|
+
# Compute local index for each node in its graph. `index` is
|
|
70
|
+
# required to be sorted (see `assert_sorted_index`), so the
|
|
71
|
+
# first occurrence of each value is its graph's start offset.
|
|
72
|
+
counts = np.bincount(index_np, minlength=B)
|
|
73
|
+
max_nodes = int(np.max(counts)) if len(counts) > 0 else 0
|
|
74
|
+
if max_num_elements is not None:
|
|
75
|
+
max_nodes = max(max_nodes, int(max_num_elements))
|
|
76
|
+
|
|
77
|
+
local_index = np.arange(N) - np.searchsorted(index_np, index_np, side="left")
|
|
78
|
+
valid = local_index < max_nodes
|
|
79
|
+
|
|
80
|
+
mask_np = np.zeros((B, max_nodes), dtype=bool)
|
|
81
|
+
mask_np[index_np[valid], local_index[valid]] = True
|
|
82
|
+
|
|
83
|
+
feat_shape = tuple(ops.shape(x)[1:])
|
|
84
|
+
scatter_idx = np.stack([index_np[valid], local_index[valid]], axis=1)
|
|
85
|
+
x_valid = x if bool(valid.all()) else ops.take(x, np.nonzero(valid)[0], axis=0)
|
|
86
|
+
out = scatter(scatter_idx, x_valid, shape=(B, max_nodes, *feat_shape))
|
|
87
|
+
|
|
88
|
+
if fill_value != 0.0:
|
|
89
|
+
mask_t = ops.convert_to_tensor(mask_np)
|
|
90
|
+
mask_expanded = ops.reshape(mask_t, (B, max_nodes) + (1,) * len(feat_shape))
|
|
91
|
+
fill = full((B, max_nodes, *feat_shape), fill_value, dtype=x.dtype)
|
|
92
|
+
out = ops.where(mask_expanded, out, fill)
|
|
93
|
+
|
|
94
|
+
return out, ops.convert_to_tensor(mask_np, dtype="bool")
|
|
95
|
+
except Exception:
|
|
96
|
+
pass
|
|
97
|
+
|
|
98
|
+
# Pure ops implementation for symbolic tracing / graph execution. `index` is sorted, so a
|
|
99
|
+
# node's position in its graph is its offset from the graph's first node (O(N) memory).
|
|
100
|
+
from k3_node.ops.segment import segment_sum
|
|
101
|
+
|
|
102
|
+
N = ops.shape(x)[0]
|
|
103
|
+
index = ops.cast(index, "int32")
|
|
104
|
+
static_size = isinstance(dim_size, (int, np.integer)) # a tensor while tracing
|
|
105
|
+
num_segments = int(dim_size) if static_size else N
|
|
106
|
+
counts = segment_sum(ops.ones_like(index), index, num_segments=num_segments)
|
|
107
|
+
starts = ops.cumsum(counts) - counts
|
|
108
|
+
local_idx = ops.arange(N, dtype="int32") - ops.take(starts, index, axis=0)
|
|
109
|
+
|
|
110
|
+
if max_num_elements is not None:
|
|
111
|
+
max_nodes = int(max_num_elements)
|
|
112
|
+
elif hasattr(x, "shape") and x.shape[0] is not None:
|
|
113
|
+
max_nodes = int(x.shape[0])
|
|
114
|
+
else:
|
|
115
|
+
max_nodes = ops.max(local_idx) + 1
|
|
116
|
+
|
|
117
|
+
B = int(dim_size) if static_size else ops.max(index) + 1
|
|
118
|
+
from k3_node.layers.conv.utils import is_tracing
|
|
119
|
+
|
|
120
|
+
if not static_size and keras.config.backend() == "jax" and is_tracing(index):
|
|
121
|
+
raise ValueError(
|
|
122
|
+
"Compiled JAX code needs the number of graphs as a Python int to build the dense batch. "
|
|
123
|
+
"Pass it explicitly, e.g. `batch_size=data.num_graphs` (or `dim_size=` for "
|
|
124
|
+
"`to_dense_batch`), or train with `run_eagerly=True`."
|
|
125
|
+
)
|
|
126
|
+
|
|
127
|
+
dense_x = full((B, max_nodes, *ops.shape(x)[1:]), fill_value, dtype=x.dtype)
|
|
128
|
+
mask = ops.zeros((B, max_nodes), dtype="bool")
|
|
129
|
+
|
|
130
|
+
scatter_indices = ops.stack([ops.cast(index, "int32"), ops.cast(local_idx, "int32")], axis=1)
|
|
131
|
+
dense_x = ops.scatter_update(dense_x, scatter_indices, x)
|
|
132
|
+
mask = ops.scatter_update(mask, scatter_indices, ops.ones((N,), dtype="bool"))
|
|
133
|
+
return dense_x, mask
|
|
134
|
+
|
|
135
|
+
|
|
136
|
+
def from_dense_batch(x_dense, index):
|
|
137
|
+
r"""The inverse of :func:`to_dense_batch`: turns ``[batch_size, max_nodes, *dims]`` back into
|
|
138
|
+
one row per node, in the order given by the sorted ``index``. Static shapes, differentiable.
|
|
139
|
+
|
|
140
|
+
Example:
|
|
141
|
+
```python
|
|
142
|
+
import numpy as np
|
|
143
|
+
from k3_node.layers import from_dense_batch, to_dense_batch
|
|
144
|
+
|
|
145
|
+
x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
|
|
146
|
+
batch = np.array([0, 0, 0, 1, 1, 1, 1, 1, 2, 2]) # three graphs of different sizes
|
|
147
|
+
|
|
148
|
+
x_dense, mask = to_dense_batch(x, batch)
|
|
149
|
+
print(np.allclose(from_dense_batch(x_dense, batch), x)) # True
|
|
150
|
+
```
|
|
151
|
+
"""
|
|
152
|
+
from k3_node.ops.segment import segment_sum
|
|
153
|
+
|
|
154
|
+
index = ops.cast(ops.convert_to_tensor(index), "int32")
|
|
155
|
+
num_graphs, max_nodes = ops.shape(x_dense)[0], ops.shape(x_dense)[1]
|
|
156
|
+
counts = segment_sum(ops.ones_like(index), index, num_segments=num_graphs)
|
|
157
|
+
starts = ops.cumsum(counts) - counts
|
|
158
|
+
position = ops.arange(ops.shape(index)[0], dtype="int32") - ops.take(starts, index, axis=0)
|
|
159
|
+
flat = ops.reshape(x_dense, (-1,) + tuple(x_dense.shape[2:]))
|
|
160
|
+
return ops.take(flat, index * max_nodes + position, axis=0)
|
|
161
|
+
|
|
162
|
+
|
|
163
|
+
def to_dense_adj(edge_index, batch=None, edge_attr=None, max_num_nodes: Optional[int] = None,
|
|
164
|
+
batch_size: Optional[int] = None):
|
|
165
|
+
r"""Converts a batch of graphs into dense adjacency matrices of shape
|
|
166
|
+
``[num_graphs, max_nodes, max_nodes]`` (or ``[..., edge_features]`` with ``edge_attr``), as in PyG.
|
|
167
|
+
|
|
168
|
+
The node dimension matches :func:`to_dense_batch` for the same ``batch`` vector, so the two can
|
|
169
|
+
be used together. Duplicate edges are summed.
|
|
170
|
+
|
|
171
|
+
Example:
|
|
172
|
+
```python
|
|
173
|
+
import numpy as np
|
|
174
|
+
from k3_node.layers import to_dense_adj
|
|
175
|
+
|
|
176
|
+
edge_index = np.array([[0, 1, 2, 3, 4], [1, 2, 0, 4, 3]])
|
|
177
|
+
batch = np.array([0, 0, 0, 1, 1]) # two graphs with 3 and 2 nodes
|
|
178
|
+
adj = to_dense_adj(edge_index, batch)
|
|
179
|
+
print(tuple(adj.shape)) # (2, 3, 3)
|
|
180
|
+
```
|
|
181
|
+
"""
|
|
182
|
+
from k3_node.layers.conv.utils import is_tracing
|
|
183
|
+
from k3_node.ops.host import _in_shape_inference
|
|
184
|
+
|
|
185
|
+
edge_index = ops.cast(ops.convert_to_tensor(edge_index), "int32")
|
|
186
|
+
if batch is None:
|
|
187
|
+
num_nodes = max_num_nodes or (int(ops.convert_to_numpy(ops.max(edge_index))) + 1)
|
|
188
|
+
batch = ops.zeros((num_nodes,), dtype="int32")
|
|
189
|
+
batch = ops.cast(ops.convert_to_tensor(batch), "int32")
|
|
190
|
+
|
|
191
|
+
if not is_tracing(batch) or (_in_shape_inference() and None not in tuple(batch.shape)):
|
|
192
|
+
from k3_node.ops.host import to_numpy # zeros during Keras' shape inference
|
|
193
|
+
|
|
194
|
+
index_np = np.asarray(to_numpy(batch)).astype(np.int64)
|
|
195
|
+
num_graphs = int(index_np.max()) + 1 if len(index_np) else 0
|
|
196
|
+
num_graphs = max(num_graphs, int(batch_size or 0))
|
|
197
|
+
local = np.arange(len(index_np)) - np.searchsorted(index_np, index_np, side="left")
|
|
198
|
+
max_nodes = int(np.bincount(index_np, minlength=num_graphs).max()) if len(index_np) else 0
|
|
199
|
+
max_nodes = max(max_nodes, int(max_num_nodes or 0))
|
|
200
|
+
local = ops.convert_to_tensor(local.astype("int32"))
|
|
201
|
+
else: # compiled: the same local indices and padding as to_dense_batch's compiled path
|
|
202
|
+
from k3_node.ops.segment import segment_sum
|
|
203
|
+
|
|
204
|
+
n = ops.shape(batch)[0]
|
|
205
|
+
static = isinstance(batch_size, (int, np.integer))
|
|
206
|
+
counts = segment_sum(ops.ones_like(batch), batch, num_segments=int(batch_size) if static else n)
|
|
207
|
+
local = ops.arange(n, dtype="int32") - ops.take(ops.cumsum(counts) - counts, batch, axis=0)
|
|
208
|
+
num_graphs = batch_size if static else ops.max(batch) + 1
|
|
209
|
+
if max_num_nodes is not None:
|
|
210
|
+
max_nodes = int(max_num_nodes)
|
|
211
|
+
elif batch.shape[0] is not None:
|
|
212
|
+
max_nodes = int(batch.shape[0])
|
|
213
|
+
else:
|
|
214
|
+
max_nodes = ops.max(local) + 1
|
|
215
|
+
|
|
216
|
+
src, dst = edge_index[0], edge_index[1]
|
|
217
|
+
indices = ops.stack([ops.take(batch, src), ops.take(local, src), ops.take(local, dst)], axis=1)
|
|
218
|
+
num_edges = ops.shape(edge_index)[1]
|
|
219
|
+
values = ops.ones((num_edges,), dtype="float32") if edge_attr is None else ops.convert_to_tensor(edge_attr)
|
|
220
|
+
extra = tuple(values.shape[1:])
|
|
221
|
+
return scatter(indices, values, (num_graphs, max_nodes, max_nodes) + extra)
|
|
222
|
+
|
|
223
|
+
|
|
224
|
+
class Aggregation(layers.Layer):
|
|
225
|
+
r"""An abstract base class for implementing custom aggregations."""
|
|
226
|
+
|
|
227
|
+
def __init__(self, **kwargs):
|
|
228
|
+
super().__init__(**kwargs)
|
|
229
|
+
self.built = True
|
|
230
|
+
|
|
231
|
+
def build(self, input_shape=None):
|
|
232
|
+
self.built = True
|
|
233
|
+
|
|
234
|
+
def reset_parameters(self):
|
|
235
|
+
r"""Resets all learnable parameters of the module."""
|
|
236
|
+
pass
|
|
237
|
+
|
|
238
|
+
def __call__(
|
|
239
|
+
self,
|
|
240
|
+
x,
|
|
241
|
+
index: Optional[any] = None,
|
|
242
|
+
ptr: Optional[any] = None,
|
|
243
|
+
dim_size: Optional[int] = None,
|
|
244
|
+
dim: int = -2,
|
|
245
|
+
max_num_elements: Optional[int] = None,
|
|
246
|
+
**kwargs,
|
|
247
|
+
):
|
|
248
|
+
# Plain NumPy inputs cannot be mixed with backend tensors (e.g. `ndarray - torch.Tensor`).
|
|
249
|
+
if isinstance(x, np.ndarray):
|
|
250
|
+
x = ops.convert_to_tensor(x)
|
|
251
|
+
dim_total = len(x.shape) if hasattr(x, "shape") and x.shape is not None else len(ops.shape(x))
|
|
252
|
+
if dim >= dim_total or dim < -dim_total:
|
|
253
|
+
raise ValueError(
|
|
254
|
+
f"Encountered invalid dimension '{dim}' of source tensor with "
|
|
255
|
+
f"{dim_total} dimensions"
|
|
256
|
+
)
|
|
257
|
+
|
|
258
|
+
if index is None and ptr is None:
|
|
259
|
+
N = x.shape[dim] if hasattr(x, "shape") and x.shape[dim] is not None else ops.shape(x)[dim]
|
|
260
|
+
index = ops.zeros((N,), dtype="int32")
|
|
261
|
+
|
|
262
|
+
if ptr is not None and index is None:
|
|
263
|
+
index = ptr2index(ptr)
|
|
264
|
+
|
|
265
|
+
if ptr is not None:
|
|
266
|
+
ptr_len = ptr.shape[0] if hasattr(ptr, "shape") and ptr.shape[0] is not None else ops.shape(ptr)[0]
|
|
267
|
+
if dim_size is None:
|
|
268
|
+
dim_size = ptr_len - 1
|
|
269
|
+
elif dim_size != ptr_len - 1:
|
|
270
|
+
raise ValueError(
|
|
271
|
+
f"Encountered invalid 'dim_size' (got '{dim_size}' but "
|
|
272
|
+
f"expected '{ptr_len - 1}')"
|
|
273
|
+
)
|
|
274
|
+
|
|
275
|
+
if index is not None and dim_size is None:
|
|
276
|
+
dim_size = ops.max(index) + 1
|
|
277
|
+
try:
|
|
278
|
+
dim_size = int(dim_size)
|
|
279
|
+
except Exception:
|
|
280
|
+
pass
|
|
281
|
+
|
|
282
|
+
# Handle positional / keyword call to call()
|
|
283
|
+
return self.call(
|
|
284
|
+
x,
|
|
285
|
+
index=index,
|
|
286
|
+
ptr=ptr,
|
|
287
|
+
dim_size=dim_size,
|
|
288
|
+
dim=dim,
|
|
289
|
+
max_num_elements=max_num_elements,
|
|
290
|
+
**kwargs,
|
|
291
|
+
)
|
|
292
|
+
|
|
293
|
+
def call(
|
|
294
|
+
self,
|
|
295
|
+
x,
|
|
296
|
+
index: Optional[any] = None,
|
|
297
|
+
ptr: Optional[any] = None,
|
|
298
|
+
dim_size: Optional[int] = None,
|
|
299
|
+
dim: int = -2,
|
|
300
|
+
max_num_elements: Optional[int] = None,
|
|
301
|
+
):
|
|
302
|
+
raise NotImplementedError
|
|
303
|
+
|
|
304
|
+
def assert_index_present(self, index: Optional[any]):
|
|
305
|
+
if index is None:
|
|
306
|
+
raise NotImplementedError("Aggregation requires 'index' to be specified")
|
|
307
|
+
|
|
308
|
+
def assert_sorted_index(self, index: Optional[any]):
|
|
309
|
+
if index is not None:
|
|
310
|
+
from k3_node.layers.conv.utils import is_tracing
|
|
311
|
+
if is_tracing(index):
|
|
312
|
+
return
|
|
313
|
+
idx_np = ops.convert_to_numpy(index)
|
|
314
|
+
if not np.all(idx_np[:-1] <= idx_np[1:]):
|
|
315
|
+
raise ValueError(
|
|
316
|
+
"Can not perform aggregation since the 'index' tensor is not sorted. "
|
|
317
|
+
"Specifically, if you use this aggregation as part of 'MessagePassing', "
|
|
318
|
+
"ensure that 'edge_index' is sorted by destination nodes."
|
|
319
|
+
)
|
|
320
|
+
|
|
321
|
+
def assert_two_dimensional_input(self, x, dim: int = -2):
|
|
322
|
+
if len(ops.shape(x)) != 2:
|
|
323
|
+
raise ValueError(
|
|
324
|
+
f"Aggregation requires two-dimensional inputs (got '{len(ops.shape(x))}')"
|
|
325
|
+
)
|
|
326
|
+
if dim not in [-2, 0]:
|
|
327
|
+
raise ValueError(
|
|
328
|
+
f"Aggregation needs to perform aggregation in first dimension (got '{dim}')"
|
|
329
|
+
)
|
|
330
|
+
|
|
331
|
+
def reduce(
|
|
332
|
+
self,
|
|
333
|
+
x,
|
|
334
|
+
index: Optional[any] = None,
|
|
335
|
+
ptr: Optional[any] = None,
|
|
336
|
+
dim_size: Optional[int] = None,
|
|
337
|
+
dim: int = -2,
|
|
338
|
+
reduce: str = "sum",
|
|
339
|
+
):
|
|
340
|
+
r"""Reduces features along groups specified by `index` or `ptr`."""
|
|
341
|
+
if ptr is not None and index is None:
|
|
342
|
+
index = ptr2index(ptr)
|
|
343
|
+
|
|
344
|
+
if index is None:
|
|
345
|
+
raise RuntimeError("Aggregation requires 'index' to be specified")
|
|
346
|
+
|
|
347
|
+
index = ops.cast(index, dtype="int32")
|
|
348
|
+
if dim_size is None:
|
|
349
|
+
dim_size = int(ops.max(index)) + 1 if ops.shape(index)[0] > 0 else 0
|
|
350
|
+
|
|
351
|
+
if reduce in ["sum", "add"]:
|
|
352
|
+
return segment_sum(x, index, num_segments=dim_size)
|
|
353
|
+
elif reduce == "mean":
|
|
354
|
+
sum_val = segment_sum(x, index, num_segments=dim_size)
|
|
355
|
+
ones = ops.ones_like(x)
|
|
356
|
+
count = segment_sum(ones, index, num_segments=dim_size)
|
|
357
|
+
return sum_val / ops.maximum(count, 1.0)
|
|
358
|
+
elif reduce == "max":
|
|
359
|
+
val = segment_max(x, index, num_segments=dim_size)
|
|
360
|
+
ones = ops.ones_like(x)
|
|
361
|
+
count = segment_sum(ones, index, num_segments=dim_size)
|
|
362
|
+
return ops.where(ops.greater(count, 0), val, ops.zeros_like(val))
|
|
363
|
+
elif reduce == "min":
|
|
364
|
+
val = -segment_max(-x, index, num_segments=dim_size)
|
|
365
|
+
ones = ops.ones_like(x)
|
|
366
|
+
count = segment_sum(ones, index, num_segments=dim_size)
|
|
367
|
+
return ops.where(ops.greater(count, 0), val, ops.zeros_like(val))
|
|
368
|
+
elif reduce == "mul":
|
|
369
|
+
log_abs = ops.log(ops.maximum(ops.abs(x), 1e-7))
|
|
370
|
+
sum_log = segment_sum(log_abs, index, num_segments=dim_size)
|
|
371
|
+
neg_count = segment_sum(ops.cast(ops.less(x, 0.0), dtype=x.dtype), index, num_segments=dim_size)
|
|
372
|
+
sign = ops.cos(ops.cast(3.141592653589793, dtype=x.dtype) * neg_count)
|
|
373
|
+
return ops.exp(sum_log) * sign
|
|
374
|
+
else:
|
|
375
|
+
raise ValueError(f"Unsupported reduction '{reduce}'")
|
|
376
|
+
|
|
377
|
+
def to_dense_batch(
|
|
378
|
+
self,
|
|
379
|
+
x,
|
|
380
|
+
index: Optional[any] = None,
|
|
381
|
+
ptr: Optional[any] = None,
|
|
382
|
+
dim_size: Optional[int] = None,
|
|
383
|
+
dim: int = -2,
|
|
384
|
+
fill_value: float = 0.0,
|
|
385
|
+
max_num_elements: Optional[int] = None,
|
|
386
|
+
) -> Tuple[any, any]:
|
|
387
|
+
if ptr is not None and index is None:
|
|
388
|
+
index = ptr2index(ptr)
|
|
389
|
+
|
|
390
|
+
self.assert_index_present(index)
|
|
391
|
+
self.assert_sorted_index(index)
|
|
392
|
+
self.assert_two_dimensional_input(x, dim)
|
|
393
|
+
|
|
394
|
+
return to_dense_batch(
|
|
395
|
+
x,
|
|
396
|
+
index,
|
|
397
|
+
dim_size=dim_size,
|
|
398
|
+
fill_value=fill_value,
|
|
399
|
+
max_num_elements=max_num_elements,
|
|
400
|
+
)
|
|
401
|
+
|
|
402
|
+
def __repr__(self) -> str:
|
|
403
|
+
return f"{self.__class__.__name__}()"
|