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,108 @@
|
|
|
1
|
+
import keras
|
|
2
|
+
from keras import ops
|
|
3
|
+
|
|
4
|
+
|
|
5
|
+
class ARLinkPredictor(keras.layers.Layer):
|
|
6
|
+
r"""Link predictor using Attract-Repel embeddings from the paper
|
|
7
|
+
`"Pseudo-Euclidean Attract-Repel Embeddings for Undirected Graphs"
|
|
8
|
+
<https://arxiv.org/abs/2106.09671>`_.
|
|
9
|
+
|
|
10
|
+
This model splits node embeddings into: attract and repel.
|
|
11
|
+
The edge prediction score is computed as the dot product of attract
|
|
12
|
+
components minus the dot product of repel components.
|
|
13
|
+
|
|
14
|
+
Args:
|
|
15
|
+
in_channels (int): Size of each input sample.
|
|
16
|
+
hidden_channels (int): Size of hidden embeddings.
|
|
17
|
+
out_channels (int, optional): Size of output embeddings. If set to
|
|
18
|
+
`None`, will default to `hidden_channels`. (default: `None`)
|
|
19
|
+
num_layers (int): Number of message passing layers. (default: `2`)
|
|
20
|
+
dropout (float): Dropout probability. (default: `0.0`)
|
|
21
|
+
attract_ratio (float): Ratio to use for attract component. Must be
|
|
22
|
+
between 0 and 1. (default: `0.5`)
|
|
23
|
+
|
|
24
|
+
Example:
|
|
25
|
+
```python
|
|
26
|
+
import numpy as np
|
|
27
|
+
from k3_node.models import ARLinkPredictor
|
|
28
|
+
|
|
29
|
+
x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
|
|
30
|
+
edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
|
|
31
|
+
|
|
32
|
+
model = ARLinkPredictor(in_channels=8, hidden_channels=16, num_layers=2)
|
|
33
|
+
scores = model(x, edge_index) # one link score per edge in edge_index
|
|
34
|
+
print(tuple(scores.shape)) # (30,)
|
|
35
|
+
```
|
|
36
|
+
"""
|
|
37
|
+
def __init__(self, in_channels, hidden_channels, out_channels=None,
|
|
38
|
+
num_layers=2, dropout=0.0, attract_ratio=0.5, **kwargs):
|
|
39
|
+
super().__init__(**kwargs)
|
|
40
|
+
|
|
41
|
+
if out_channels is None:
|
|
42
|
+
out_channels = hidden_channels
|
|
43
|
+
|
|
44
|
+
self.in_channels = in_channels
|
|
45
|
+
self.hidden_channels = hidden_channels
|
|
46
|
+
self.out_channels = out_channels
|
|
47
|
+
self.num_layers = num_layers
|
|
48
|
+
self.dropout_rate = dropout
|
|
49
|
+
|
|
50
|
+
if not 0 <= attract_ratio <= 1:
|
|
51
|
+
raise ValueError(f"attract_ratio must be between 0 and 1, got {attract_ratio}")
|
|
52
|
+
|
|
53
|
+
self.attract_ratio = attract_ratio
|
|
54
|
+
self.attract_dim = int(out_channels * attract_ratio)
|
|
55
|
+
self.repel_dim = out_channels - self.attract_dim
|
|
56
|
+
|
|
57
|
+
self.lins = [keras.layers.Dense(hidden_channels)]
|
|
58
|
+
for _ in range(num_layers - 2):
|
|
59
|
+
self.lins.append(keras.layers.Dense(hidden_channels))
|
|
60
|
+
|
|
61
|
+
self.lin_attract = keras.layers.Dense(self.attract_dim)
|
|
62
|
+
self.lin_repel = keras.layers.Dense(self.repel_dim)
|
|
63
|
+
self.dropout = keras.layers.Dropout(dropout) if dropout > 0.0 else None
|
|
64
|
+
|
|
65
|
+
self.lins[0].build((None, in_channels))
|
|
66
|
+
for lin in self.lins[1:]:
|
|
67
|
+
lin.build((None, hidden_channels))
|
|
68
|
+
self.lin_attract.build((None, hidden_channels))
|
|
69
|
+
self.lin_repel.build((None, hidden_channels))
|
|
70
|
+
|
|
71
|
+
def encode(self, x, *args, training=None, **kwargs):
|
|
72
|
+
r"""Encode node features into attract-repel embeddings."""
|
|
73
|
+
for lin in self.lins:
|
|
74
|
+
x = lin(x)
|
|
75
|
+
x = ops.relu(x)
|
|
76
|
+
if self.dropout is not None:
|
|
77
|
+
x = self.dropout(x, training=training)
|
|
78
|
+
|
|
79
|
+
attract_x = self.lin_attract(x)
|
|
80
|
+
repel_x = self.lin_repel(x)
|
|
81
|
+
|
|
82
|
+
return attract_x, repel_x
|
|
83
|
+
|
|
84
|
+
def decode(self, attract_z, repel_z, edge_index):
|
|
85
|
+
r"""Decode edge scores from attract-repel embeddings."""
|
|
86
|
+
row, col = edge_index[0], edge_index[1]
|
|
87
|
+
attract_z_row = ops.take(attract_z, row, axis=0)
|
|
88
|
+
attract_z_col = ops.take(attract_z, col, axis=0)
|
|
89
|
+
repel_z_row = ops.take(repel_z, row, axis=0)
|
|
90
|
+
repel_z_col = ops.take(repel_z, col, axis=0)
|
|
91
|
+
|
|
92
|
+
attract_score = ops.sum(attract_z_row * attract_z_col, axis=1)
|
|
93
|
+
repel_score = ops.sum(repel_z_row * repel_z_col, axis=1)
|
|
94
|
+
|
|
95
|
+
return attract_score - repel_score
|
|
96
|
+
|
|
97
|
+
def call(self, x, edge_index, training=None):
|
|
98
|
+
attract_z, repel_z = self.encode(x, training=training)
|
|
99
|
+
return ops.sigmoid(self.decode(attract_z, repel_z, edge_index))
|
|
100
|
+
|
|
101
|
+
def calculate_r_fraction(self, attract_z, repel_z):
|
|
102
|
+
r"""Calculate the R-fraction (proportion of energy in repel space)."""
|
|
103
|
+
attract_norm_squared = ops.sum(ops.square(attract_z))
|
|
104
|
+
repel_norm_squared = ops.sum(ops.square(repel_z))
|
|
105
|
+
|
|
106
|
+
r_fraction = repel_norm_squared / (attract_norm_squared + repel_norm_squared + 1e-10)
|
|
107
|
+
|
|
108
|
+
return float(ops.convert_to_numpy(r_fraction))
|
|
@@ -0,0 +1,318 @@
|
|
|
1
|
+
import keras
|
|
2
|
+
from keras import ops
|
|
3
|
+
|
|
4
|
+
from k3_node.models.utils import negative_sampling, reset
|
|
5
|
+
|
|
6
|
+
EPS = 1e-15
|
|
7
|
+
MAX_LOGSTD = 10
|
|
8
|
+
|
|
9
|
+
|
|
10
|
+
def _randn_like(x):
|
|
11
|
+
return keras.random.normal(ops.shape(x), dtype=x.dtype)
|
|
12
|
+
|
|
13
|
+
|
|
14
|
+
class InnerProductDecoder:
|
|
15
|
+
r"""The inner product decoder from the `"Variational Graph Auto-Encoders"
|
|
16
|
+
<https://arxiv.org/abs/1611.07308>`_ paper.
|
|
17
|
+
|
|
18
|
+
.. math::
|
|
19
|
+
\sigma(\mathbf{Z}\mathbf{Z}^{\top})
|
|
20
|
+
|
|
21
|
+
where :math:`\mathbf{Z} \in \mathbb{R}^{N \times d}` denotes the latent
|
|
22
|
+
space produced by the encoder.
|
|
23
|
+
"""
|
|
24
|
+
def __call__(self, z, edge_index, sigmoid: bool = True):
|
|
25
|
+
r"""Decodes the latent variables `z` into edge probabilities for
|
|
26
|
+
the given node-pairs `edge_index`."""
|
|
27
|
+
row, col = edge_index[0], edge_index[1]
|
|
28
|
+
value = ops.sum(ops.take(z, row, axis=0) * ops.take(z, col, axis=0), axis=1)
|
|
29
|
+
return ops.sigmoid(value) if sigmoid else value
|
|
30
|
+
|
|
31
|
+
def forward_all(self, z, sigmoid: bool = True):
|
|
32
|
+
r"""Decodes the latent variables `z` into a probabilistic dense
|
|
33
|
+
adjacency matrix."""
|
|
34
|
+
adj = ops.matmul(z, ops.transpose(z))
|
|
35
|
+
return ops.sigmoid(adj) if sigmoid else adj
|
|
36
|
+
|
|
37
|
+
|
|
38
|
+
class GAE:
|
|
39
|
+
r"""The Graph Auto-Encoder model from the
|
|
40
|
+
`"Variational Graph Auto-Encoders" <https://arxiv.org/abs/1611.07308>`_
|
|
41
|
+
paper based on user-defined encoder and decoder models.
|
|
42
|
+
|
|
43
|
+
Args:
|
|
44
|
+
encoder: The encoder module.
|
|
45
|
+
decoder (optional): The decoder module. If set to `None`, will
|
|
46
|
+
default to `InnerProductDecoder`. (default: `None`)
|
|
47
|
+
"""
|
|
48
|
+
def __init__(self, encoder, decoder=None):
|
|
49
|
+
self.encoder = encoder
|
|
50
|
+
self.decoder = InnerProductDecoder() if decoder is None else decoder
|
|
51
|
+
self.reset_parameters()
|
|
52
|
+
|
|
53
|
+
def reset_parameters(self):
|
|
54
|
+
r"""Resets all learnable parameters of the module."""
|
|
55
|
+
reset(self.encoder)
|
|
56
|
+
reset(self.decoder)
|
|
57
|
+
|
|
58
|
+
def __call__(self, *args, **kwargs):
|
|
59
|
+
r"""Alias for `encode`."""
|
|
60
|
+
return self.encoder(*args, **kwargs)
|
|
61
|
+
|
|
62
|
+
def encode(self, *args, **kwargs):
|
|
63
|
+
r"""Runs the encoder and computes node-wise latent variables."""
|
|
64
|
+
return self.encoder(*args, **kwargs)
|
|
65
|
+
|
|
66
|
+
def decode(self, *args, **kwargs):
|
|
67
|
+
r"""Runs the decoder and computes edge probabilities."""
|
|
68
|
+
return self.decoder(*args, **kwargs)
|
|
69
|
+
|
|
70
|
+
def eval(self):
|
|
71
|
+
self.training = False
|
|
72
|
+
return self
|
|
73
|
+
|
|
74
|
+
def train(self):
|
|
75
|
+
self.training = True
|
|
76
|
+
return self
|
|
77
|
+
|
|
78
|
+
def to(self, *args, **kwargs):
|
|
79
|
+
return self
|
|
80
|
+
|
|
81
|
+
def recon_loss(self, z, pos_edge_index, neg_edge_index=None):
|
|
82
|
+
r"""Given latent variables `z`, computes the binary cross entropy
|
|
83
|
+
loss for positive edges `pos_edge_index` and negative sampled
|
|
84
|
+
edges."""
|
|
85
|
+
pos_loss = -ops.mean(ops.log(self.decoder(z, pos_edge_index, sigmoid=True) + EPS))
|
|
86
|
+
|
|
87
|
+
if neg_edge_index is None:
|
|
88
|
+
neg_edge_index = negative_sampling(pos_edge_index, ops.shape(z)[0])
|
|
89
|
+
neg_loss = -ops.mean(ops.log(1 - self.decoder(z, neg_edge_index, sigmoid=True) + EPS))
|
|
90
|
+
|
|
91
|
+
return pos_loss + neg_loss
|
|
92
|
+
|
|
93
|
+
def test(self, z, pos_edge_index, neg_edge_index):
|
|
94
|
+
r"""Given latent variables `z`, positive edges `pos_edge_index` and
|
|
95
|
+
negative edges `neg_edge_index`, computes area under the ROC curve
|
|
96
|
+
(AUC) and average precision (AP) scores."""
|
|
97
|
+
from sklearn.metrics import average_precision_score, roc_auc_score
|
|
98
|
+
|
|
99
|
+
pos_y = ops.ones((ops.shape(pos_edge_index)[1],))
|
|
100
|
+
neg_y = ops.zeros((ops.shape(neg_edge_index)[1],))
|
|
101
|
+
y = ops.concatenate([pos_y, neg_y], axis=0)
|
|
102
|
+
|
|
103
|
+
pos_pred = self.decoder(z, pos_edge_index, sigmoid=True)
|
|
104
|
+
neg_pred = self.decoder(z, neg_edge_index, sigmoid=True)
|
|
105
|
+
pred = ops.concatenate([pos_pred, neg_pred], axis=0)
|
|
106
|
+
|
|
107
|
+
y, pred = ops.convert_to_numpy(y), ops.convert_to_numpy(pred)
|
|
108
|
+
|
|
109
|
+
return roc_auc_score(y, pred), average_precision_score(y, pred)
|
|
110
|
+
|
|
111
|
+
# ---- Keras-style training -------------------------------------------------------------------
|
|
112
|
+
def compile(self, optimizer, discriminator_optimizer=None, discriminator_steps: int = 5):
|
|
113
|
+
r"""Sets the optimizers used by :meth:`fit`.
|
|
114
|
+
|
|
115
|
+
Args:
|
|
116
|
+
optimizer (keras.optimizers.Optimizer): Trains the encoder (and decoder).
|
|
117
|
+
discriminator_optimizer (keras.optimizers.Optimizer, optional): Trains the
|
|
118
|
+
discriminator of adversarial models (:class:`ARGA`, :class:`ARGVA`).
|
|
119
|
+
discriminator_steps (int): Discriminator updates per encoder update. (default: ``5``)
|
|
120
|
+
"""
|
|
121
|
+
self.optimizer = optimizer
|
|
122
|
+
self.discriminator_optimizer = discriminator_optimizer
|
|
123
|
+
self.discriminator_steps = discriminator_steps
|
|
124
|
+
|
|
125
|
+
def _encoder_loss(self, data):
|
|
126
|
+
z = self.encode(data.x, data.edge_index, training=True)
|
|
127
|
+
loss = self.recon_loss(z, data.pos_edge_label_index)
|
|
128
|
+
if isinstance(self, ARGA):
|
|
129
|
+
loss = loss + self.reg_loss(z)
|
|
130
|
+
if hasattr(self, "kl_loss"):
|
|
131
|
+
loss = loss + (1 / data.num_nodes) * self.kl_loss()
|
|
132
|
+
return loss
|
|
133
|
+
|
|
134
|
+
def train_step(self, data):
|
|
135
|
+
r"""Runs one training step on ``data`` and returns the loss."""
|
|
136
|
+
from k3_node.training import gradient_step
|
|
137
|
+
|
|
138
|
+
self.train()
|
|
139
|
+
if isinstance(self, ARGA):
|
|
140
|
+
z = ops.stop_gradient(self.encode(data.x, data.edge_index, training=True))
|
|
141
|
+
for _ in range(self.discriminator_steps):
|
|
142
|
+
gradient_step(lambda: self.discriminator_loss(z), self.discriminator.trainable_variables,
|
|
143
|
+
self.discriminator_optimizer)
|
|
144
|
+
return gradient_step(lambda: self._encoder_loss(data), self._trainable_variables(), self.optimizer)
|
|
145
|
+
|
|
146
|
+
def _trainable_variables(self):
|
|
147
|
+
variables = list(self.encoder.trainable_variables)
|
|
148
|
+
return variables + list(getattr(self.decoder, "trainable_variables", []))
|
|
149
|
+
|
|
150
|
+
def fit(self, data, epochs: int = 1, validation_data=None, verbose: int = 1):
|
|
151
|
+
r"""Trains the model on one graph for ``epochs`` full-graph steps.
|
|
152
|
+
|
|
153
|
+
Args:
|
|
154
|
+
data (Data): The training graph with node features ``x``, the message passing edges
|
|
155
|
+
``edge_index`` and the edges to reconstruct ``pos_edge_label_index``, as created
|
|
156
|
+
by :class:`~k3_node.transforms.RandomLinkSplit` with ``split_labels=True``.
|
|
157
|
+
epochs (int): The number of training steps. (default: ``1``)
|
|
158
|
+
validation_data (Data, optional): A graph with ``pos_edge_label_index`` and
|
|
159
|
+
``neg_edge_label_index`` on which AUC and average precision are reported.
|
|
160
|
+
verbose (int): ``0`` is silent, otherwise one line is printed per epoch.
|
|
161
|
+
|
|
162
|
+
Returns:
|
|
163
|
+
dict: The loss (and validation metrics) of every epoch.
|
|
164
|
+
"""
|
|
165
|
+
if getattr(self, "optimizer", None) is None:
|
|
166
|
+
raise ValueError("Call `compile(optimizer=...)` before `fit`.")
|
|
167
|
+
if isinstance(self, ARGA) and self.discriminator_optimizer is None:
|
|
168
|
+
raise ValueError("Adversarial models need `compile(..., discriminator_optimizer=...)`.")
|
|
169
|
+
# Create the variables before the first gradient step
|
|
170
|
+
z = self.encode(data.x, data.edge_index)
|
|
171
|
+
if isinstance(self, ARGA):
|
|
172
|
+
self.discriminator(z)
|
|
173
|
+
|
|
174
|
+
history = {"loss": []}
|
|
175
|
+
for epoch in range(1, epochs + 1):
|
|
176
|
+
logs = {"loss": self.train_step(data)}
|
|
177
|
+
if validation_data is not None:
|
|
178
|
+
logs.update({f"val_{k}": v for k, v in self.evaluate(validation_data).items()})
|
|
179
|
+
for key, value in logs.items():
|
|
180
|
+
history.setdefault(key, []).append(value)
|
|
181
|
+
if verbose:
|
|
182
|
+
print(f"Epoch {epoch:03d}: " + ", ".join(f"{k}: {v:.4f}" for k, v in logs.items()))
|
|
183
|
+
return history
|
|
184
|
+
|
|
185
|
+
def evaluate(self, data):
|
|
186
|
+
r"""Returns the link prediction AUC and average precision on ``data`` (which needs
|
|
187
|
+
``pos_edge_label_index`` and ``neg_edge_label_index``)."""
|
|
188
|
+
z = self.embed(data)
|
|
189
|
+
auc, ap = self.test(z, data.pos_edge_label_index, data.neg_edge_label_index)
|
|
190
|
+
return {"auc": float(auc), "ap": float(ap)}
|
|
191
|
+
|
|
192
|
+
def embed(self, data):
|
|
193
|
+
r"""Returns the node embeddings of ``data`` (without sampling noise)."""
|
|
194
|
+
from k3_node.training import no_grad
|
|
195
|
+
|
|
196
|
+
self.eval()
|
|
197
|
+
with no_grad():
|
|
198
|
+
return self.encode(data.x, data.edge_index, training=False)
|
|
199
|
+
|
|
200
|
+
|
|
201
|
+
class VGAE(GAE):
|
|
202
|
+
r"""The Variational Graph Auto-Encoder model from the
|
|
203
|
+
`"Variational Graph Auto-Encoders" <https://arxiv.org/abs/1611.07308>`_
|
|
204
|
+
paper.
|
|
205
|
+
|
|
206
|
+
Args:
|
|
207
|
+
encoder: The encoder module to compute :math:`\mu` and
|
|
208
|
+
:math:`\log\sigma^2`.
|
|
209
|
+
decoder (optional): The decoder module. If set to `None`, will
|
|
210
|
+
default to `InnerProductDecoder`. (default: `None`)
|
|
211
|
+
"""
|
|
212
|
+
def __init__(self, encoder, decoder=None):
|
|
213
|
+
super().__init__(encoder, decoder)
|
|
214
|
+
self.training = True
|
|
215
|
+
|
|
216
|
+
def reparametrize(self, mu, logstd):
|
|
217
|
+
if self.training:
|
|
218
|
+
return mu + _randn_like(logstd) * ops.exp(logstd)
|
|
219
|
+
return mu
|
|
220
|
+
|
|
221
|
+
def encode(self, *args, **kwargs):
|
|
222
|
+
self._mu, self._logstd = self.encoder(*args, **kwargs)
|
|
223
|
+
self._logstd = ops.minimum(self._logstd, MAX_LOGSTD)
|
|
224
|
+
z = self.reparametrize(self._mu, self._logstd)
|
|
225
|
+
return z
|
|
226
|
+
|
|
227
|
+
def kl_loss(self, mu=None, logstd=None):
|
|
228
|
+
r"""Computes the KL loss, either for the passed arguments `mu` and
|
|
229
|
+
`logstd`, or based on latent variables from last encoding."""
|
|
230
|
+
mu = self._mu if mu is None else mu
|
|
231
|
+
logstd = self._logstd if logstd is None else ops.minimum(logstd, MAX_LOGSTD)
|
|
232
|
+
return -0.5 * ops.mean(
|
|
233
|
+
ops.sum(1 + 2 * logstd - ops.square(mu) - ops.square(ops.exp(logstd)), axis=1)
|
|
234
|
+
)
|
|
235
|
+
|
|
236
|
+
def eval(self):
|
|
237
|
+
self.training = False
|
|
238
|
+
|
|
239
|
+
def train(self):
|
|
240
|
+
self.training = True
|
|
241
|
+
|
|
242
|
+
|
|
243
|
+
class ARGA(GAE):
|
|
244
|
+
r"""The Adversarially Regularized Graph Auto-Encoder model from the
|
|
245
|
+
`"Adversarially Regularized Graph Autoencoder for Graph Embedding"
|
|
246
|
+
<https://arxiv.org/abs/1802.04407>`_ paper.
|
|
247
|
+
|
|
248
|
+
Args:
|
|
249
|
+
encoder: The encoder module.
|
|
250
|
+
discriminator: The discriminator module.
|
|
251
|
+
decoder (optional): The decoder module. If set to `None`, will
|
|
252
|
+
default to `InnerProductDecoder`. (default: `None`)
|
|
253
|
+
"""
|
|
254
|
+
def __init__(self, encoder, discriminator, decoder=None):
|
|
255
|
+
super().__init__(encoder, decoder)
|
|
256
|
+
self.discriminator = discriminator
|
|
257
|
+
reset(self.discriminator)
|
|
258
|
+
|
|
259
|
+
def reset_parameters(self):
|
|
260
|
+
super().reset_parameters()
|
|
261
|
+
reset(getattr(self, "discriminator", None))
|
|
262
|
+
|
|
263
|
+
def reg_loss(self, z):
|
|
264
|
+
r"""Computes the regularization loss of the encoder."""
|
|
265
|
+
real = ops.sigmoid(self.discriminator(z))
|
|
266
|
+
return -ops.mean(ops.log(real + EPS))
|
|
267
|
+
|
|
268
|
+
def discriminator_loss(self, z):
|
|
269
|
+
r"""Computes the loss of the discriminator."""
|
|
270
|
+
real = ops.sigmoid(self.discriminator(_randn_like(z)))
|
|
271
|
+
fake = ops.sigmoid(self.discriminator(ops.stop_gradient(z)))
|
|
272
|
+
real_loss = -ops.mean(ops.log(real + EPS))
|
|
273
|
+
fake_loss = -ops.mean(ops.log(1 - fake + EPS))
|
|
274
|
+
return real_loss + fake_loss
|
|
275
|
+
|
|
276
|
+
|
|
277
|
+
class ARGVA(ARGA):
|
|
278
|
+
r"""The Adversarially Regularized Variational Graph Auto-Encoder model
|
|
279
|
+
from the `"Adversarially Regularized Graph Autoencoder for Graph
|
|
280
|
+
Embedding" <https://arxiv.org/abs/1802.04407>`_ paper.
|
|
281
|
+
|
|
282
|
+
Args:
|
|
283
|
+
encoder: The encoder module to compute :math:`\mu` and
|
|
284
|
+
:math:`\log\sigma^2`.
|
|
285
|
+
discriminator: The discriminator module.
|
|
286
|
+
decoder (optional): The decoder module. If set to `None`, will
|
|
287
|
+
default to `InnerProductDecoder`. (default: `None`)
|
|
288
|
+
"""
|
|
289
|
+
def __init__(self, encoder, discriminator, decoder=None):
|
|
290
|
+
super().__init__(encoder, discriminator, decoder)
|
|
291
|
+
self.vgae = VGAE(encoder, decoder)
|
|
292
|
+
|
|
293
|
+
@property
|
|
294
|
+
def _mu(self):
|
|
295
|
+
return self.vgae._mu
|
|
296
|
+
|
|
297
|
+
@property
|
|
298
|
+
def _logstd(self):
|
|
299
|
+
return self.vgae._logstd
|
|
300
|
+
|
|
301
|
+
def reparametrize(self, mu, logstd):
|
|
302
|
+
return self.vgae.reparametrize(mu, logstd)
|
|
303
|
+
|
|
304
|
+
def encode(self, *args, **kwargs):
|
|
305
|
+
return self.vgae.encode(*args, **kwargs)
|
|
306
|
+
|
|
307
|
+
def kl_loss(self, mu=None, logstd=None):
|
|
308
|
+
return self.vgae.kl_loss(mu, logstd)
|
|
309
|
+
|
|
310
|
+
def eval(self):
|
|
311
|
+
self.training = False
|
|
312
|
+
self.vgae.eval()
|
|
313
|
+
return self
|
|
314
|
+
|
|
315
|
+
def train(self):
|
|
316
|
+
self.training = True
|
|
317
|
+
self.vgae.train()
|
|
318
|
+
return self
|