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,121 @@
|
|
|
1
|
+
import os
|
|
2
|
+
import os.path as osp
|
|
3
|
+
from typing import Callable, List, Optional
|
|
4
|
+
|
|
5
|
+
import numpy as np
|
|
6
|
+
from keras import ops
|
|
7
|
+
|
|
8
|
+
from k3_node.data import (
|
|
9
|
+
Data,
|
|
10
|
+
InMemoryDataset,
|
|
11
|
+
download_url,
|
|
12
|
+
extract_zip,
|
|
13
|
+
)
|
|
14
|
+
from k3_node.utils import coalesce
|
|
15
|
+
|
|
16
|
+
|
|
17
|
+
class Reddit(InMemoryDataset):
|
|
18
|
+
r"""The Reddit dataset from the `"Inductive Representation Learning on
|
|
19
|
+
Large Graphs" <https://arxiv.org/abs/1706.02216>`_ paper, containing
|
|
20
|
+
Reddit posts belonging to different communities.
|
|
21
|
+
|
|
22
|
+
Args:
|
|
23
|
+
root (str): Root directory where the dataset should be saved.
|
|
24
|
+
transform (callable, optional): A function/transform that takes in an
|
|
25
|
+
:obj:`k3_node.data.Data` object and returns a transformed
|
|
26
|
+
version. The data object will be transformed before every access.
|
|
27
|
+
(default: :obj:`None`)
|
|
28
|
+
pre_transform (callable, optional): A function/transform that takes in
|
|
29
|
+
an :obj:`k3_node.data.Data` object and returns a
|
|
30
|
+
transformed version. The data object will be transformed before
|
|
31
|
+
being saved to disk. (default: :obj:`None`)
|
|
32
|
+
pre_filter (callable, optional): A function that takes in an
|
|
33
|
+
:obj:`k3_node.data.Data` object and returns a boolean
|
|
34
|
+
value, indicating whether the data object should be included in the
|
|
35
|
+
final dataset. (default: :obj:`None`)
|
|
36
|
+
force_reload (bool, optional): Whether to re-process the dataset.
|
|
37
|
+
(default: :obj:`False`)
|
|
38
|
+
|
|
39
|
+
**STATS:**
|
|
40
|
+
|
|
41
|
+
.. list-table::
|
|
42
|
+
:widths: 10 10 10 10
|
|
43
|
+
:header-rows: 1
|
|
44
|
+
|
|
45
|
+
* - #nodes
|
|
46
|
+
- #edges
|
|
47
|
+
- #features
|
|
48
|
+
- #classes
|
|
49
|
+
* - 232,965
|
|
50
|
+
- 114,615,892
|
|
51
|
+
- 602
|
|
52
|
+
- 41
|
|
53
|
+
"""
|
|
54
|
+
|
|
55
|
+
url = "https://data.dgl.ai/dataset/reddit.zip"
|
|
56
|
+
|
|
57
|
+
def __init__(
|
|
58
|
+
self,
|
|
59
|
+
root: str,
|
|
60
|
+
transform: Optional[Callable] = None,
|
|
61
|
+
pre_transform: Optional[Callable] = None,
|
|
62
|
+
pre_filter: Optional[Callable] = None,
|
|
63
|
+
force_reload: bool = False,
|
|
64
|
+
) -> None:
|
|
65
|
+
super().__init__(
|
|
66
|
+
root,
|
|
67
|
+
transform,
|
|
68
|
+
pre_transform,
|
|
69
|
+
pre_filter,
|
|
70
|
+
force_reload=force_reload,
|
|
71
|
+
)
|
|
72
|
+
self.load(self.processed_paths[0])
|
|
73
|
+
|
|
74
|
+
@property
|
|
75
|
+
def raw_file_names(self) -> List[str]:
|
|
76
|
+
return ["reddit_data.npz", "reddit_graph.npz"]
|
|
77
|
+
|
|
78
|
+
@property
|
|
79
|
+
def processed_file_names(self) -> str:
|
|
80
|
+
return "data.pt"
|
|
81
|
+
|
|
82
|
+
@property
|
|
83
|
+
def num_classes(self) -> int:
|
|
84
|
+
return 41
|
|
85
|
+
|
|
86
|
+
def download(self) -> None:
|
|
87
|
+
path = download_url(self.url, self.raw_dir)
|
|
88
|
+
extract_zip(path, self.raw_dir)
|
|
89
|
+
if osp.exists(path):
|
|
90
|
+
os.unlink(path)
|
|
91
|
+
|
|
92
|
+
def process(self) -> None:
|
|
93
|
+
import scipy.sparse as sp
|
|
94
|
+
|
|
95
|
+
data = np.load(osp.join(self.raw_dir, "reddit_data.npz"))
|
|
96
|
+
x = ops.convert_to_tensor(data["feature"], dtype="float32")
|
|
97
|
+
y = ops.convert_to_tensor(data["label"], dtype="int64")
|
|
98
|
+
split = data["node_types"]
|
|
99
|
+
|
|
100
|
+
adj = sp.load_npz(osp.join(self.raw_dir, "reddit_graph.npz"))
|
|
101
|
+
if not hasattr(adj, "row"):
|
|
102
|
+
adj = adj.tocoo()
|
|
103
|
+
row = adj.row.astype(np.int64)
|
|
104
|
+
col = adj.col.astype(np.int64)
|
|
105
|
+
edge_index = np.stack([row, col], axis=0)
|
|
106
|
+
edge_index = ops.convert_to_tensor(edge_index, dtype="int64")
|
|
107
|
+
edge_index, _ = coalesce(edge_index, num_nodes=int(data["feature"].shape[0]))
|
|
108
|
+
|
|
109
|
+
data = Data(x=x, edge_index=edge_index, y=y)
|
|
110
|
+
data.train_mask = ops.convert_to_tensor(split == 1, dtype="bool")
|
|
111
|
+
data.val_mask = ops.convert_to_tensor(split == 2, dtype="bool")
|
|
112
|
+
data.test_mask = ops.convert_to_tensor(split == 3, dtype="bool")
|
|
113
|
+
|
|
114
|
+
if self.pre_filter is not None and not self.pre_filter(data):
|
|
115
|
+
return
|
|
116
|
+
|
|
117
|
+
if self.pre_transform is not None:
|
|
118
|
+
data = self.pre_transform(data)
|
|
119
|
+
|
|
120
|
+
self.save([data], self.processed_paths[0])
|
|
121
|
+
|
|
@@ -0,0 +1,165 @@
|
|
|
1
|
+
import os.path as osp
|
|
2
|
+
from typing import Any, Callable, List, Optional, Union
|
|
3
|
+
import numpy as np
|
|
4
|
+
from keras import ops
|
|
5
|
+
|
|
6
|
+
from k3_node.data import Data, InMemoryDataset
|
|
7
|
+
from k3_node.utils.random import stochastic_blockmodel_graph
|
|
8
|
+
|
|
9
|
+
|
|
10
|
+
class StochasticBlockModelDataset(InMemoryDataset):
|
|
11
|
+
r"""A synthetic graph dataset generated by the stochastic block model.
|
|
12
|
+
|
|
13
|
+
Args:
|
|
14
|
+
root (str): Root directory where the dataset should be saved.
|
|
15
|
+
block_sizes ([int] or array): The sizes of blocks.
|
|
16
|
+
edge_probs ([[float]] or array): The density of edges between blocks.
|
|
17
|
+
num_graphs (int, optional): The number of graphs. (default: 1)
|
|
18
|
+
num_channels (int, optional): The number of node features. (default: None)
|
|
19
|
+
is_undirected (bool, optional): Whether the graph is undirected. (default: True)
|
|
20
|
+
transform (callable, optional): Transform function.
|
|
21
|
+
pre_transform (callable, optional): Pre-transform function.
|
|
22
|
+
force_reload (bool, optional): Whether to re-process the dataset.
|
|
23
|
+
"""
|
|
24
|
+
|
|
25
|
+
def __init__(
|
|
26
|
+
self,
|
|
27
|
+
root: str,
|
|
28
|
+
block_sizes: Union[List[int], np.ndarray],
|
|
29
|
+
edge_probs: Union[List[List[float]], np.ndarray],
|
|
30
|
+
num_graphs: int = 1,
|
|
31
|
+
num_channels: Optional[int] = None,
|
|
32
|
+
is_undirected: bool = True,
|
|
33
|
+
transform: Optional[Callable] = None,
|
|
34
|
+
pre_transform: Optional[Callable] = None,
|
|
35
|
+
force_reload: bool = False,
|
|
36
|
+
**kwargs: Any,
|
|
37
|
+
):
|
|
38
|
+
self.block_sizes = np.array(block_sizes, dtype=np.int64)
|
|
39
|
+
self.edge_probs = np.array(edge_probs, dtype=np.float32)
|
|
40
|
+
assert num_graphs > 0
|
|
41
|
+
|
|
42
|
+
self.num_graphs = num_graphs
|
|
43
|
+
self.num_channels = num_channels
|
|
44
|
+
self.is_undirected = is_undirected
|
|
45
|
+
|
|
46
|
+
self.kwargs = {
|
|
47
|
+
"n_informative": num_channels,
|
|
48
|
+
"n_redundant": 0,
|
|
49
|
+
"flip_y": 0.0,
|
|
50
|
+
"shuffle": False,
|
|
51
|
+
}
|
|
52
|
+
self.kwargs.update(kwargs)
|
|
53
|
+
|
|
54
|
+
super().__init__(root, transform, pre_transform, force_reload=force_reload)
|
|
55
|
+
self.load(self.processed_paths[0])
|
|
56
|
+
|
|
57
|
+
@property
|
|
58
|
+
def processed_dir(self) -> str:
|
|
59
|
+
return osp.join(self.root, self.__class__.__name__, "processed")
|
|
60
|
+
|
|
61
|
+
@property
|
|
62
|
+
def processed_file_names(self) -> str:
|
|
63
|
+
bs = self.block_sizes.flatten().tolist()
|
|
64
|
+
hash1 = "-".join([f"{x:.1f}" for x in bs])
|
|
65
|
+
ep = self.edge_probs.flatten().tolist()
|
|
66
|
+
hash2 = "-".join([f"{x:.1f}" for x in ep])
|
|
67
|
+
return f"data_{self.num_channels}_{hash1}_{hash2}_{self.num_graphs}.pt"
|
|
68
|
+
|
|
69
|
+
def process(self):
|
|
70
|
+
try:
|
|
71
|
+
from sklearn.datasets import make_classification
|
|
72
|
+
except ImportError:
|
|
73
|
+
make_classification = None
|
|
74
|
+
|
|
75
|
+
edge_index = stochastic_blockmodel_graph(
|
|
76
|
+
self.block_sizes, self.edge_probs, directed=not self.is_undirected
|
|
77
|
+
)
|
|
78
|
+
|
|
79
|
+
num_samples = int(self.block_sizes.sum())
|
|
80
|
+
num_classes = len(self.block_sizes)
|
|
81
|
+
|
|
82
|
+
data_list = []
|
|
83
|
+
for _ in range(self.num_graphs):
|
|
84
|
+
x = None
|
|
85
|
+
if self.num_channels is not None:
|
|
86
|
+
if make_classification is not None:
|
|
87
|
+
x_raw, y_not_sorted = make_classification(
|
|
88
|
+
n_samples=num_samples,
|
|
89
|
+
n_features=self.num_channels,
|
|
90
|
+
n_classes=num_classes,
|
|
91
|
+
weights=self.block_sizes / num_samples,
|
|
92
|
+
**self.kwargs,
|
|
93
|
+
)
|
|
94
|
+
x_raw = x_raw[np.argsort(y_not_sorted)]
|
|
95
|
+
x = ops.convert_to_tensor(x_raw.astype(np.float32), dtype="float32")
|
|
96
|
+
else:
|
|
97
|
+
x = ops.convert_to_tensor(
|
|
98
|
+
np.random.randn(num_samples, self.num_channels).astype(np.float32),
|
|
99
|
+
dtype="float32",
|
|
100
|
+
)
|
|
101
|
+
|
|
102
|
+
y_np = np.repeat(np.arange(num_classes, dtype=np.int64), self.block_sizes)
|
|
103
|
+
y = ops.convert_to_tensor(y_np, dtype="int64")
|
|
104
|
+
|
|
105
|
+
data = Data(x=x, edge_index=edge_index, y=y)
|
|
106
|
+
if self.pre_transform is not None:
|
|
107
|
+
data = self.pre_transform(data)
|
|
108
|
+
data_list.append(data)
|
|
109
|
+
|
|
110
|
+
self.save(data_list, self.processed_paths[0])
|
|
111
|
+
|
|
112
|
+
|
|
113
|
+
class RandomPartitionGraphDataset(StochasticBlockModelDataset):
|
|
114
|
+
r"""The random partition graph dataset from the "How to Find Your Friendly
|
|
115
|
+
Neighborhood: Graph Attention Design with Self-Supervision" paper.
|
|
116
|
+
"""
|
|
117
|
+
|
|
118
|
+
def __init__(
|
|
119
|
+
self,
|
|
120
|
+
root: str,
|
|
121
|
+
num_classes: int,
|
|
122
|
+
num_nodes_per_class: int,
|
|
123
|
+
node_homophily_ratio: float,
|
|
124
|
+
average_degree: float,
|
|
125
|
+
num_graphs: int = 1,
|
|
126
|
+
num_channels: Optional[int] = None,
|
|
127
|
+
is_undirected: bool = True,
|
|
128
|
+
transform: Optional[Callable] = None,
|
|
129
|
+
pre_transform: Optional[Callable] = None,
|
|
130
|
+
**kwargs: Any,
|
|
131
|
+
):
|
|
132
|
+
self._num_classes = num_classes
|
|
133
|
+
self.num_nodes_per_class = num_nodes_per_class
|
|
134
|
+
self.node_homophily_ratio = node_homophily_ratio
|
|
135
|
+
self.average_degree = average_degree
|
|
136
|
+
|
|
137
|
+
ec_over_v2 = average_degree / num_nodes_per_class
|
|
138
|
+
p_in = node_homophily_ratio * ec_over_v2
|
|
139
|
+
p_out = (ec_over_v2 - p_in) / (num_classes - 1)
|
|
140
|
+
|
|
141
|
+
block_sizes = [num_nodes_per_class for _ in range(num_classes)]
|
|
142
|
+
edge_probs = [[p_out for _ in range(num_classes)] for _ in range(num_classes)]
|
|
143
|
+
for r in range(num_classes):
|
|
144
|
+
edge_probs[r][r] = p_in
|
|
145
|
+
|
|
146
|
+
super().__init__(
|
|
147
|
+
root,
|
|
148
|
+
block_sizes,
|
|
149
|
+
edge_probs,
|
|
150
|
+
num_graphs=num_graphs,
|
|
151
|
+
num_channels=num_channels,
|
|
152
|
+
is_undirected=is_undirected,
|
|
153
|
+
transform=transform,
|
|
154
|
+
pre_transform=pre_transform,
|
|
155
|
+
**kwargs,
|
|
156
|
+
)
|
|
157
|
+
|
|
158
|
+
@property
|
|
159
|
+
def processed_file_names(self) -> str:
|
|
160
|
+
return (
|
|
161
|
+
f"data_{self.num_channels}_{self._num_classes}_"
|
|
162
|
+
f"{self.num_nodes_per_class}_{self.node_homophily_ratio:.1f}_"
|
|
163
|
+
f"{self.average_degree:.1f}_{self.num_graphs}.pt"
|
|
164
|
+
)
|
|
165
|
+
|
k3_node/datasets/seal.py
ADDED
|
@@ -0,0 +1,74 @@
|
|
|
1
|
+
from typing import Optional
|
|
2
|
+
|
|
3
|
+
import numpy as np
|
|
4
|
+
from keras import ops
|
|
5
|
+
|
|
6
|
+
from k3_node.data import Data, InMemoryDataset
|
|
7
|
+
|
|
8
|
+
|
|
9
|
+
class SEALDataset(InMemoryDataset):
|
|
10
|
+
r"""Enclosing subgraphs for link prediction, as in SEAL
|
|
11
|
+
(`"Link Prediction Based on Graph Neural Networks" <https://arxiv.org/abs/1802.09691>`_).
|
|
12
|
+
|
|
13
|
+
Every link to classify becomes a small graph: the ``num_hops``-hop neighborhood of its two
|
|
14
|
+
nodes, without the link itself. Nodes are described only by their double-radius labels
|
|
15
|
+
(their distances to the two nodes, one-hot encoded as ``x``); ``y`` is 1 for a true link and
|
|
16
|
+
0 for a non-edge. This turns link prediction into graph classification.
|
|
17
|
+
|
|
18
|
+
Args:
|
|
19
|
+
data (Data): One split from :class:`~k3_node.transforms.RandomLinkSplit`: message passing
|
|
20
|
+
edges in ``edge_index`` and the links in ``edge_label_index`` / ``edge_label`` (or
|
|
21
|
+
``pos_edge_label_index`` / ``neg_edge_label_index``).
|
|
22
|
+
num_hops (int): Size of the neighborhood around each link. (default: ``2``)
|
|
23
|
+
num_labels (int, optional): Size of the one-hot node labels. Use the training set's
|
|
24
|
+
``num_labels`` for the validation and test sets. (default: the largest label + 1)
|
|
25
|
+
|
|
26
|
+
Example:
|
|
27
|
+
```python
|
|
28
|
+
import numpy as np
|
|
29
|
+
from k3_node.data import Data
|
|
30
|
+
from k3_node.datasets import SEALDataset
|
|
31
|
+
from k3_node.transforms import RandomLinkSplit
|
|
32
|
+
|
|
33
|
+
data = Data(edge_index=np.random.randint(0, 30, size=(2, 120)), num_nodes=30)
|
|
34
|
+
train_data, val_data, test_data = RandomLinkSplit(num_val=0.1, num_test=0.1)(data)
|
|
35
|
+
train_dataset = SEALDataset(train_data, num_hops=2)
|
|
36
|
+
print(len(train_dataset) == train_data.edge_label_index.shape[1]) # True: one graph per link
|
|
37
|
+
```
|
|
38
|
+
"""
|
|
39
|
+
|
|
40
|
+
def __init__(self, data, num_hops: int = 2, num_labels: Optional[int] = None):
|
|
41
|
+
super().__init__(None)
|
|
42
|
+
from k3_node.utils.graph import drnl_node_labeling, k_hop_subgraph
|
|
43
|
+
|
|
44
|
+
links, labels = self._links(data)
|
|
45
|
+
edge_index = np.asarray(ops.convert_to_numpy(data.edge_index)).astype(np.int64)
|
|
46
|
+
graphs = []
|
|
47
|
+
for (src, dst), y in zip(links.T, labels):
|
|
48
|
+
nodes, sub_edge_index, mapping, _ = k_hop_subgraph(
|
|
49
|
+
[src, dst], num_hops, edge_index, relabel_nodes=True, num_nodes=data.num_nodes)
|
|
50
|
+
s, d = (int(m) for m in mapping)
|
|
51
|
+
keep = ~(((sub_edge_index[0] == s) & (sub_edge_index[1] == d))
|
|
52
|
+
| ((sub_edge_index[0] == d) & (sub_edge_index[1] == s)))
|
|
53
|
+
sub_edge_index = sub_edge_index[:, keep] # hide the link to predict
|
|
54
|
+
z = drnl_node_labeling(sub_edge_index, s, d, num_nodes=len(nodes))
|
|
55
|
+
graphs.append((z, sub_edge_index, y))
|
|
56
|
+
|
|
57
|
+
self.num_labels = num_labels or max(int(z.max()) for z, _, _ in graphs) + 1
|
|
58
|
+
self.data, self.slices = self.collate([
|
|
59
|
+
Data(x=np.eye(self.num_labels, dtype=np.float32)[np.minimum(z, self.num_labels - 1)],
|
|
60
|
+
edge_index=e, y=np.array([y], dtype=np.float32))
|
|
61
|
+
for z, e, y in graphs
|
|
62
|
+
])
|
|
63
|
+
|
|
64
|
+
@staticmethod
|
|
65
|
+
def _links(data):
|
|
66
|
+
def arr(x):
|
|
67
|
+
return np.asarray(ops.convert_to_numpy(x))
|
|
68
|
+
|
|
69
|
+
if getattr(data, "edge_label_index", None) is not None:
|
|
70
|
+
return arr(data.edge_label_index).astype(np.int64), arr(data.edge_label)
|
|
71
|
+
pos = arr(data.pos_edge_label_index).astype(np.int64)
|
|
72
|
+
neg = getattr(data, "neg_edge_label_index", None)
|
|
73
|
+
neg = np.zeros((2, 0), np.int64) if neg is None else arr(neg).astype(np.int64)
|
|
74
|
+
return np.concatenate([pos, neg], axis=1), np.concatenate([np.ones(pos.shape[1]), np.zeros(neg.shape[1])])
|
|
@@ -0,0 +1,92 @@
|
|
|
1
|
+
from typing import Callable, Optional
|
|
2
|
+
|
|
3
|
+
import numpy as np
|
|
4
|
+
|
|
5
|
+
from k3_node.data import Data, InMemoryDataset
|
|
6
|
+
|
|
7
|
+
|
|
8
|
+
class ShapeScenes(InMemoryDataset):
|
|
9
|
+
r"""Small point cloud scenes for semantic segmentation, built from :class:`GeometricShapes`
|
|
10
|
+
meshes (a compact stand-in for ShapeNet part segmentation).
|
|
11
|
+
|
|
12
|
+
Every scene contains ``shapes_per_scene`` objects, each a cube, sphere, cone or torus that is
|
|
13
|
+
randomly scaled, rotated and placed apart from the others. ``num_points`` points are sampled on
|
|
14
|
+
their surfaces: ``pos`` holds the positions (scaled to the unit ball), ``x`` the surface
|
|
15
|
+
normals, and ``y`` the type of the object each point lies on (4 classes). Training scenes use
|
|
16
|
+
the training meshes, test scenes the test meshes; the scenes are generated with a fixed seed.
|
|
17
|
+
|
|
18
|
+
Args:
|
|
19
|
+
root (str): Directory where GeometricShapes is (or will be) downloaded.
|
|
20
|
+
train (bool, optional): Build training scenes if ``True``, else test scenes.
|
|
21
|
+
num_scenes (int, optional): Number of scenes. (default: ``200`` for training, ``50`` for test)
|
|
22
|
+
shapes_per_scene (int, optional): Objects per scene. (default: ``3``)
|
|
23
|
+
num_points (int, optional): Points per scene. (default: ``1024``)
|
|
24
|
+
transform (callable, optional): A function applied to each scene when it is accessed.
|
|
25
|
+
seed (int, optional): Random seed. (default: ``0``)
|
|
26
|
+
"""
|
|
27
|
+
|
|
28
|
+
categories = ["3d_cube", "3d_sphere", "3d_cone", "3d_torus"]
|
|
29
|
+
|
|
30
|
+
def __init__(self, root: str, train: bool = True, num_scenes: Optional[int] = None, shapes_per_scene: int = 3,
|
|
31
|
+
num_points: int = 1024, transform: Optional[Callable] = None, seed: int = 0):
|
|
32
|
+
super().__init__(None, transform)
|
|
33
|
+
from k3_node.datasets.geometric_shapes import GeometricShapes
|
|
34
|
+
from k3_node.transforms.spatial import _np
|
|
35
|
+
|
|
36
|
+
shapes = GeometricShapes(root, train=train)
|
|
37
|
+
names = sorted(__import__("os").listdir(shapes.raw_dir))
|
|
38
|
+
meshes = {}
|
|
39
|
+
for i in range(len(shapes)):
|
|
40
|
+
data = shapes[i]
|
|
41
|
+
name = names[int(_np(data.y)[0])]
|
|
42
|
+
if name in self.categories:
|
|
43
|
+
meshes[self.categories.index(name)] = (_np(data.pos).astype(np.float64), _np(data.face))
|
|
44
|
+
|
|
45
|
+
rng = np.random.default_rng(seed + (0 if train else 1))
|
|
46
|
+
num_scenes = num_scenes or (200 if train else 50)
|
|
47
|
+
self.data, self.slices = self.collate(
|
|
48
|
+
[self._scene(rng, meshes, shapes_per_scene, num_points) for _ in range(num_scenes)])
|
|
49
|
+
|
|
50
|
+
@staticmethod
|
|
51
|
+
def _sample(rng, pos, face, num):
|
|
52
|
+
a, b, c = pos[face[0]], pos[face[1]], pos[face[2]]
|
|
53
|
+
cross = np.cross(b - a, c - a)
|
|
54
|
+
area = np.linalg.norm(cross, axis=1)
|
|
55
|
+
tri = rng.choice(len(area), size=num, p=area / area.sum())
|
|
56
|
+
u, v = rng.random((2, num))
|
|
57
|
+
flip = u + v > 1
|
|
58
|
+
u[flip], v[flip] = 1 - u[flip], 1 - v[flip]
|
|
59
|
+
points = a[tri] + u[:, None] * (b - a)[tri] + v[:, None] * (c - a)[tri]
|
|
60
|
+
normals = cross[tri] / np.maximum(area[tri, None], 1e-12)
|
|
61
|
+
return points, normals
|
|
62
|
+
|
|
63
|
+
@staticmethod
|
|
64
|
+
def _rotation(rng):
|
|
65
|
+
q = rng.normal(size=4)
|
|
66
|
+
q /= np.linalg.norm(q)
|
|
67
|
+
w, x, y, z = q
|
|
68
|
+
return np.array([[1 - 2 * (y * y + z * z), 2 * (x * y - z * w), 2 * (x * z + y * w)],
|
|
69
|
+
[2 * (x * y + z * w), 1 - 2 * (x * x + z * z), 2 * (y * z - x * w)],
|
|
70
|
+
[2 * (x * z - y * w), 2 * (y * z + x * w), 1 - 2 * (x * x + y * y)]])
|
|
71
|
+
|
|
72
|
+
def _scene(self, rng, meshes, shapes_per_scene, num_points):
|
|
73
|
+
counts = np.full(shapes_per_scene, num_points // shapes_per_scene)
|
|
74
|
+
counts[: num_points - counts.sum()] += 1
|
|
75
|
+
angle = rng.random() * 2 * np.pi
|
|
76
|
+
positions, normals, labels = [], [], []
|
|
77
|
+
for i in range(shapes_per_scene):
|
|
78
|
+
label = int(rng.integers(len(self.categories)))
|
|
79
|
+
pos, face = meshes[label]
|
|
80
|
+
points, normal = self._sample(rng, pos, face, counts[i])
|
|
81
|
+
points = points / np.abs(points).max()
|
|
82
|
+
rot = self._rotation(rng)
|
|
83
|
+
theta = angle + 2 * np.pi * i / shapes_per_scene # objects around a circle, apart
|
|
84
|
+
offset = 2.5 * np.array([np.cos(theta), np.sin(theta), 0.0])
|
|
85
|
+
positions.append(points @ rot.T * rng.uniform(0.6, 1.0) + offset)
|
|
86
|
+
normals.append(normal @ rot.T)
|
|
87
|
+
labels.append(np.full(counts[i], label))
|
|
88
|
+
pos = np.concatenate(positions)
|
|
89
|
+
pos = pos - pos.mean(axis=0)
|
|
90
|
+
pos = pos / np.linalg.norm(pos, axis=1).max()
|
|
91
|
+
return Data(pos=pos.astype(np.float32), x=np.concatenate(normals).astype(np.float32),
|
|
92
|
+
y=np.concatenate(labels).astype(np.int64))
|