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,234 @@
|
|
|
1
|
+
from typing import Dict, List, Optional, Tuple
|
|
2
|
+
import numpy as np
|
|
3
|
+
import keras
|
|
4
|
+
from keras import ops
|
|
5
|
+
|
|
6
|
+
EdgeType = Tuple[str, str, str]
|
|
7
|
+
NodeType = str
|
|
8
|
+
|
|
9
|
+
|
|
10
|
+
class MetaPath2Vec(keras.layers.Layer):
|
|
11
|
+
r"""The MetaPath2Vec model from the `"metapath2vec: Scalable Representation
|
|
12
|
+
Learning for Heterogeneous Networks"
|
|
13
|
+
<https://ericdongyx.github.io/papers/
|
|
14
|
+
KDD17-dong-chawla-swami-metapath2vec.pdf>`_ paper where random walks based
|
|
15
|
+
on a given :obj:`metapath` are sampled in a heterogeneous graph, and node
|
|
16
|
+
embeddings are learned via negative sampling optimization.
|
|
17
|
+
|
|
18
|
+
Args:
|
|
19
|
+
edge_index_dict (Dict[Tuple[str, str, str], Tensor]): Dictionary
|
|
20
|
+
holding edge indices for each edge type.
|
|
21
|
+
embedding_dim (int): The size of each embedding vector.
|
|
22
|
+
metapath (List[Tuple[str, str, str]]): The sequence of edge types
|
|
23
|
+
denoting the metapath.
|
|
24
|
+
walk_length (int): The walk length.
|
|
25
|
+
context_size (int): The context size considered for positive samples.
|
|
26
|
+
walks_per_node (int, optional): The number of walks to sample for each node.
|
|
27
|
+
(default: :obj:`1`)
|
|
28
|
+
num_negative_samples (int, optional): The number of negative samples.
|
|
29
|
+
(default: :obj:`1`)
|
|
30
|
+
num_nodes_dict (Dict[str, int], optional): The number of nodes for each
|
|
31
|
+
node type. (default: :obj:`None`)
|
|
32
|
+
|
|
33
|
+
Example:
|
|
34
|
+
```python
|
|
35
|
+
import numpy as np
|
|
36
|
+
from k3_node.models import MetaPath2Vec
|
|
37
|
+
|
|
38
|
+
edge_index_dict = {
|
|
39
|
+
("author", "writes", "paper"): np.array([[0, 1, 1], [0, 0, 1]]),
|
|
40
|
+
("paper", "written_by", "author"): np.array([[0, 0, 1], [0, 1, 1]]),
|
|
41
|
+
}
|
|
42
|
+
metapath = [("author", "writes", "paper"), ("paper", "written_by", "author")]
|
|
43
|
+
model = MetaPath2Vec(edge_index_dict, embedding_dim=16, metapath=metapath,
|
|
44
|
+
walk_length=2, context_size=2, walks_per_node=2)
|
|
45
|
+
print(tuple(model("author").shape)) # (2, 16): embeddings of all authors
|
|
46
|
+
print(tuple(model("paper").shape)) # (2, 16)
|
|
47
|
+
```
|
|
48
|
+
"""
|
|
49
|
+
def __init__(
|
|
50
|
+
self,
|
|
51
|
+
edge_index_dict: Dict[EdgeType, any],
|
|
52
|
+
embedding_dim: int,
|
|
53
|
+
metapath: List[EdgeType],
|
|
54
|
+
walk_length: int,
|
|
55
|
+
context_size: int,
|
|
56
|
+
walks_per_node: int = 1,
|
|
57
|
+
num_negative_samples: int = 1,
|
|
58
|
+
num_nodes_dict: Optional[Dict[NodeType, int]] = None,
|
|
59
|
+
**kwargs,
|
|
60
|
+
):
|
|
61
|
+
super().__init__(**kwargs)
|
|
62
|
+
|
|
63
|
+
if num_nodes_dict is None:
|
|
64
|
+
num_nodes_dict = {}
|
|
65
|
+
for keys, edge_index in edge_index_dict.items():
|
|
66
|
+
e_np = ops.convert_to_numpy(edge_index)
|
|
67
|
+
key = keys[0]
|
|
68
|
+
N = int(e_np[0].max() + 1) if e_np.size > 0 else 0
|
|
69
|
+
num_nodes_dict[key] = max(N, num_nodes_dict.get(key, N))
|
|
70
|
+
|
|
71
|
+
key = keys[-1]
|
|
72
|
+
N = int(e_np[1].max() + 1) if e_np.size > 0 else 0
|
|
73
|
+
num_nodes_dict[key] = max(N, num_nodes_dict.get(key, N))
|
|
74
|
+
|
|
75
|
+
# Build adjacency dictionaries
|
|
76
|
+
self.adj_dict = {}
|
|
77
|
+
for keys, edge_index in edge_index_dict.items():
|
|
78
|
+
src_type, _, dst_type = keys
|
|
79
|
+
num_src = num_nodes_dict[src_type]
|
|
80
|
+
adj = [[] for _ in range(num_src)]
|
|
81
|
+
e_np = ops.convert_to_numpy(edge_index).astype(np.int64)
|
|
82
|
+
if e_np.size > 0:
|
|
83
|
+
for s, d in zip(e_np[0], e_np[1]):
|
|
84
|
+
adj[s].append(int(d))
|
|
85
|
+
self.adj_dict[keys] = adj
|
|
86
|
+
|
|
87
|
+
for edge_type1, edge_type2 in zip(metapath[:-1], metapath[1:]):
|
|
88
|
+
if edge_type1[-1] != edge_type2[0]:
|
|
89
|
+
raise ValueError(
|
|
90
|
+
"Found invalid metapath. Ensure that the destination node "
|
|
91
|
+
"type matches with the source node type across all "
|
|
92
|
+
"consecutive edge types."
|
|
93
|
+
)
|
|
94
|
+
|
|
95
|
+
assert walk_length + 1 >= context_size
|
|
96
|
+
|
|
97
|
+
self.embedding_dim = embedding_dim
|
|
98
|
+
self.metapath = metapath
|
|
99
|
+
self.walk_length = walk_length
|
|
100
|
+
self.context_size = context_size
|
|
101
|
+
self.walks_per_node = walks_per_node
|
|
102
|
+
self.num_negative_samples = num_negative_samples
|
|
103
|
+
self.num_nodes_dict = num_nodes_dict
|
|
104
|
+
self.EPS = 1e-15
|
|
105
|
+
|
|
106
|
+
types = sorted(list({x[0] for x in metapath} | {x[-1] for x in metapath}))
|
|
107
|
+
|
|
108
|
+
count = 0
|
|
109
|
+
self.start, self.end = {}, {}
|
|
110
|
+
for key in types:
|
|
111
|
+
self.start[key] = count
|
|
112
|
+
count += num_nodes_dict[key]
|
|
113
|
+
self.end[key] = count
|
|
114
|
+
|
|
115
|
+
offset = [self.start[metapath[0][0]]]
|
|
116
|
+
offset += [self.start[keys[-1]] for keys in metapath] * int(
|
|
117
|
+
(walk_length / len(metapath)) + 1
|
|
118
|
+
)
|
|
119
|
+
offset = offset[: walk_length + 1]
|
|
120
|
+
self.offset = np.array(offset, dtype=np.int64)
|
|
121
|
+
|
|
122
|
+
self.dummy_idx = count
|
|
123
|
+
self.embedding = keras.layers.Embedding(count + 1, embedding_dim)
|
|
124
|
+
|
|
125
|
+
def reset_parameters(self):
|
|
126
|
+
if self.embedding.built:
|
|
127
|
+
self.embedding.embeddings.assign(
|
|
128
|
+
keras.initializers.GlorotUniform()(self.embedding.embeddings.shape)
|
|
129
|
+
)
|
|
130
|
+
|
|
131
|
+
def __call__(self, node_type: str, batch: Optional[any] = None, **kwargs):
|
|
132
|
+
return self.call(node_type, batch)
|
|
133
|
+
|
|
134
|
+
def forward(self, node_type: str, batch: Optional[any] = None, **kwargs):
|
|
135
|
+
return self.call(node_type, batch)
|
|
136
|
+
|
|
137
|
+
def call(self, node_type: str, batch: Optional[any] = None):
|
|
138
|
+
r"""Returns the embeddings for the nodes in :obj:`batch` of type
|
|
139
|
+
:obj:`node_type`.
|
|
140
|
+
"""
|
|
141
|
+
start = self.start[node_type]
|
|
142
|
+
end = self.end[node_type]
|
|
143
|
+
if batch is None:
|
|
144
|
+
batch = ops.arange(start, end, dtype="int64")
|
|
145
|
+
else:
|
|
146
|
+
batch = ops.cast(batch, "int64") + start
|
|
147
|
+
return self.embedding(batch)
|
|
148
|
+
|
|
149
|
+
def _pos_sample(self, batch):
|
|
150
|
+
batch_np = ops.convert_to_numpy(batch).astype(np.int64)
|
|
151
|
+
repeated = np.repeat(batch_np, self.walks_per_node)
|
|
152
|
+
|
|
153
|
+
all_walks = []
|
|
154
|
+
for node in repeated:
|
|
155
|
+
walk = [int(node)]
|
|
156
|
+
for i in range(self.walk_length):
|
|
157
|
+
edge_type = self.metapath[i % len(self.metapath)]
|
|
158
|
+
cur = walk[-1]
|
|
159
|
+
nbrs = (
|
|
160
|
+
self.adj_dict[edge_type][cur]
|
|
161
|
+
if cur < len(self.adj_dict[edge_type])
|
|
162
|
+
else []
|
|
163
|
+
)
|
|
164
|
+
if len(nbrs) > 0:
|
|
165
|
+
walk.append(nbrs[np.random.randint(len(nbrs))])
|
|
166
|
+
else:
|
|
167
|
+
walk.append(self.dummy_idx)
|
|
168
|
+
all_walks.append(walk)
|
|
169
|
+
|
|
170
|
+
rw = np.array(all_walks, dtype=np.int64)
|
|
171
|
+
rw = rw + self.offset[None, :]
|
|
172
|
+
rw[rw > self.dummy_idx] = self.dummy_idx
|
|
173
|
+
|
|
174
|
+
walks = []
|
|
175
|
+
num_walks_per_rw = 1 + self.walk_length + 1 - self.context_size
|
|
176
|
+
for j in range(num_walks_per_rw):
|
|
177
|
+
walks.append(rw[:, j : j + self.context_size])
|
|
178
|
+
out = np.concatenate(walks, axis=0) if len(walks) > 0 else rw
|
|
179
|
+
return ops.convert_to_tensor(out, dtype="int64")
|
|
180
|
+
|
|
181
|
+
def _neg_sample(self, batch):
|
|
182
|
+
batch_np = ops.convert_to_numpy(batch).astype(np.int64)
|
|
183
|
+
repeated = np.repeat(
|
|
184
|
+
batch_np, self.walks_per_node * self.num_negative_samples
|
|
185
|
+
)
|
|
186
|
+
|
|
187
|
+
rws = [repeated]
|
|
188
|
+
for i in range(self.walk_length):
|
|
189
|
+
keys = self.metapath[i % len(self.metapath)]
|
|
190
|
+
num_nodes = self.num_nodes_dict[keys[-1]]
|
|
191
|
+
rand = np.random.randint(0, max(num_nodes, 1), size=len(repeated))
|
|
192
|
+
rws.append(rand)
|
|
193
|
+
|
|
194
|
+
rw = np.stack(rws, axis=-1)
|
|
195
|
+
rw = rw + self.offset[None, :]
|
|
196
|
+
|
|
197
|
+
walks = []
|
|
198
|
+
num_walks_per_rw = 1 + self.walk_length + 1 - self.context_size
|
|
199
|
+
for j in range(num_walks_per_rw):
|
|
200
|
+
walks.append(rw[:, j : j + self.context_size])
|
|
201
|
+
out = np.concatenate(walks, axis=0) if len(walks) > 0 else rw
|
|
202
|
+
return ops.convert_to_tensor(out, dtype="int64")
|
|
203
|
+
|
|
204
|
+
def loss(self, pos_rw, neg_rw):
|
|
205
|
+
r"""Computes the loss given positive and negative random walks."""
|
|
206
|
+
# Positive loss
|
|
207
|
+
start = pos_rw[:, 0]
|
|
208
|
+
rest = pos_rw[:, 1:]
|
|
209
|
+
pos_b = ops.shape(pos_rw)[0]
|
|
210
|
+
|
|
211
|
+
h_start = ops.reshape(self.embedding(start), (pos_b, 1, self.embedding_dim))
|
|
212
|
+
h_rest = ops.reshape(
|
|
213
|
+
self.embedding(ops.reshape(rest, (-1,))),
|
|
214
|
+
(pos_b, -1, self.embedding_dim),
|
|
215
|
+
)
|
|
216
|
+
|
|
217
|
+
out = ops.reshape(ops.sum(h_start * h_rest, axis=-1), (-1,))
|
|
218
|
+
pos_loss = -ops.mean(ops.log(ops.sigmoid(out) + self.EPS))
|
|
219
|
+
|
|
220
|
+
# Negative loss
|
|
221
|
+
start = neg_rw[:, 0]
|
|
222
|
+
rest = neg_rw[:, 1:]
|
|
223
|
+
neg_b = ops.shape(neg_rw)[0]
|
|
224
|
+
|
|
225
|
+
h_start = ops.reshape(self.embedding(start), (neg_b, 1, self.embedding_dim))
|
|
226
|
+
h_rest = ops.reshape(
|
|
227
|
+
self.embedding(ops.reshape(rest, (-1,))),
|
|
228
|
+
(neg_b, -1, self.embedding_dim),
|
|
229
|
+
)
|
|
230
|
+
|
|
231
|
+
out = ops.reshape(ops.sum(h_start * h_rest, axis=-1), (-1,))
|
|
232
|
+
neg_loss = -ops.mean(ops.log(1.0 - ops.sigmoid(out) + self.EPS))
|
|
233
|
+
|
|
234
|
+
return pos_loss + neg_loss
|
k3_node/models/mlp.py
ADDED
|
@@ -0,0 +1,264 @@
|
|
|
1
|
+
import re
|
|
2
|
+
import inspect
|
|
3
|
+
from typing import List, Optional, Union
|
|
4
|
+
|
|
5
|
+
import keras
|
|
6
|
+
from keras import ops
|
|
7
|
+
|
|
8
|
+
import k3_node.layers.norm as norm_module
|
|
9
|
+
|
|
10
|
+
|
|
11
|
+
def _normalize_string(s: str) -> str:
|
|
12
|
+
return re.sub(r"[_\-\s]", "", s).lower()
|
|
13
|
+
|
|
14
|
+
|
|
15
|
+
def _normalization_resolver(query, *args, **kwargs):
|
|
16
|
+
if query is None:
|
|
17
|
+
return None
|
|
18
|
+
if not isinstance(query, str):
|
|
19
|
+
return query
|
|
20
|
+
|
|
21
|
+
norms = {
|
|
22
|
+
_normalize_string(name): cls
|
|
23
|
+
for name, cls in vars(norm_module).items()
|
|
24
|
+
if isinstance(cls, type)
|
|
25
|
+
}
|
|
26
|
+
key = _normalize_string(query)
|
|
27
|
+
if key in norms:
|
|
28
|
+
return norms[key](*args, **kwargs)
|
|
29
|
+
if key + "norm" in norms:
|
|
30
|
+
return norms[key + "norm"](*args, **kwargs)
|
|
31
|
+
if key.endswith("norm") and key[:-4] in norms:
|
|
32
|
+
return norms[key[:-4]](*args, **kwargs)
|
|
33
|
+
raise ValueError(f"Could not resolve normalization layer '{query}'")
|
|
34
|
+
|
|
35
|
+
|
|
36
|
+
def _activation_resolver(act, **kwargs):
|
|
37
|
+
"""Resolves Keras and PyG-style activation names ("relu", "leaky_relu", "LeakyReLU", ...)."""
|
|
38
|
+
if act is None:
|
|
39
|
+
return None
|
|
40
|
+
if isinstance(act, str):
|
|
41
|
+
name = act.lower().replace("_", "")
|
|
42
|
+
aliases = {"leakyrelu": "leaky_relu", "hardswish": "hard_swish", "hardsigmoid": "hard_sigmoid"}
|
|
43
|
+
act = keras.activations.get(aliases.get(name, act.lower()))
|
|
44
|
+
if kwargs:
|
|
45
|
+
import functools
|
|
46
|
+
|
|
47
|
+
return functools.partial(act, **kwargs)
|
|
48
|
+
return act
|
|
49
|
+
|
|
50
|
+
|
|
51
|
+
class MLP(keras.Model):
|
|
52
|
+
r"""A Multi-Layer Perceptron (MLP) model.
|
|
53
|
+
|
|
54
|
+
There exists two ways to instantiate an `MLP`:
|
|
55
|
+
|
|
56
|
+
1. By specifying explicit channel sizes, e.g., `MLP([16, 32, 64, 128])`
|
|
57
|
+
creates a three-layer MLP with **differently** sized hidden layers.
|
|
58
|
+
|
|
59
|
+
2. By specifying fixed hidden channel sizes over a number of layers,
|
|
60
|
+
e.g., `MLP(in_channels=16, hidden_channels=32, out_channels=128,
|
|
61
|
+
num_layers=3)` creates a three-layer MLP with **equally** sized
|
|
62
|
+
hidden layers.
|
|
63
|
+
|
|
64
|
+
Args:
|
|
65
|
+
channel_list (List[int] or int, optional): List of input,
|
|
66
|
+
intermediate and output channels such that
|
|
67
|
+
`len(channel_list) - 1` denotes the number of layers of the
|
|
68
|
+
MLP. (default: `None`)
|
|
69
|
+
in_channels (int, optional): Size of each input sample. Will
|
|
70
|
+
override `channel_list`. (default: `None`)
|
|
71
|
+
hidden_channels (int, optional): Size of each hidden sample. Will
|
|
72
|
+
override `channel_list`. (default: `None`)
|
|
73
|
+
out_channels (int, optional): Size of each output sample. Will
|
|
74
|
+
override `channel_list`. (default: `None`)
|
|
75
|
+
num_layers (int, optional): The number of layers. Will override
|
|
76
|
+
`channel_list`. (default: `None`)
|
|
77
|
+
dropout (float or List[float], optional): Dropout probability of
|
|
78
|
+
each hidden embedding. (default: `0.`)
|
|
79
|
+
act (str or Callable, optional): The non-linear activation function
|
|
80
|
+
to use. (default: `"relu"`)
|
|
81
|
+
act_first (bool, optional): If set to `True`, activation is applied
|
|
82
|
+
before normalization. (default: `False`)
|
|
83
|
+
act_kwargs (dict, optional): Arguments passed to the activation function, e.g.
|
|
84
|
+
``{"negative_slope": 0.2}`` for ``"leaky_relu"``. (default: `None`)
|
|
85
|
+
norm (str or Callable, optional): The normalization function to
|
|
86
|
+
use. (default: `"batch_norm"`)
|
|
87
|
+
norm_kwargs (dict, optional): Arguments passed to the respective
|
|
88
|
+
normalization function. (default: `None`)
|
|
89
|
+
plain_last (bool, optional): If set to `False`, will apply
|
|
90
|
+
non-linearity, normalization and dropout to the last layer as
|
|
91
|
+
well. (default: `True`)
|
|
92
|
+
bias (bool or List[bool], optional): If set to `False`, the module
|
|
93
|
+
will not learn additive biases. (default: `True`)
|
|
94
|
+
|
|
95
|
+
Example:
|
|
96
|
+
```python
|
|
97
|
+
import numpy as np
|
|
98
|
+
from k3_node.models import MLP
|
|
99
|
+
|
|
100
|
+
x = np.random.rand(10, 16).astype("float32")
|
|
101
|
+
mlp = MLP([16, 32, 32, 4]) # channel sizes: input, hidden, hidden, output
|
|
102
|
+
print(tuple(mlp(x).shape)) # (10, 4)
|
|
103
|
+
|
|
104
|
+
mlp = MLP(in_channels=16, hidden_channels=32, out_channels=4, num_layers=3, dropout=0.1)
|
|
105
|
+
print(tuple(mlp(x).shape)) # (10, 4)
|
|
106
|
+
```
|
|
107
|
+
"""
|
|
108
|
+
def __init__(
|
|
109
|
+
self,
|
|
110
|
+
channel_list: Optional[Union[List[int], int]] = None,
|
|
111
|
+
*args,
|
|
112
|
+
in_channels: Optional[int] = None,
|
|
113
|
+
hidden_channels: Optional[int] = None,
|
|
114
|
+
out_channels: Optional[int] = None,
|
|
115
|
+
num_layers: Optional[int] = None,
|
|
116
|
+
dropout: Union[float, List[float]] = 0.0,
|
|
117
|
+
act="relu",
|
|
118
|
+
act_first: bool = False,
|
|
119
|
+
act_kwargs: Optional[dict] = None,
|
|
120
|
+
norm="batch_norm",
|
|
121
|
+
norm_kwargs: Optional[dict] = None,
|
|
122
|
+
plain_last: bool = True,
|
|
123
|
+
bias: Union[bool, List[bool]] = True,
|
|
124
|
+
**kwargs,
|
|
125
|
+
):
|
|
126
|
+
super().__init__(**kwargs)
|
|
127
|
+
|
|
128
|
+
if len(args) > 0:
|
|
129
|
+
if isinstance(channel_list, int):
|
|
130
|
+
channel_list = [channel_list] + list(args)
|
|
131
|
+
elif isinstance(channel_list, (list, tuple)):
|
|
132
|
+
channel_list = list(channel_list) + list(args)
|
|
133
|
+
|
|
134
|
+
if isinstance(channel_list, int):
|
|
135
|
+
in_channels = channel_list
|
|
136
|
+
channel_list = None
|
|
137
|
+
|
|
138
|
+
if in_channels is not None:
|
|
139
|
+
if num_layers is None:
|
|
140
|
+
raise ValueError("Argument `num_layers` must be given")
|
|
141
|
+
if num_layers > 1 and hidden_channels is None:
|
|
142
|
+
raise ValueError(
|
|
143
|
+
f"Argument `hidden_channels` must be given for `num_layers={num_layers}`"
|
|
144
|
+
)
|
|
145
|
+
if out_channels is None:
|
|
146
|
+
raise ValueError("Argument `out_channels` must be given")
|
|
147
|
+
|
|
148
|
+
channel_list = [hidden_channels] * (num_layers - 1)
|
|
149
|
+
channel_list = [in_channels] + channel_list + [out_channels]
|
|
150
|
+
|
|
151
|
+
assert isinstance(channel_list, (tuple, list))
|
|
152
|
+
assert len(channel_list) >= 2
|
|
153
|
+
self.channel_list = list(channel_list)
|
|
154
|
+
self.in_channels = self.channel_list[0]
|
|
155
|
+
self.out_channels = self.channel_list[-1]
|
|
156
|
+
|
|
157
|
+
self.act = _activation_resolver(act, **(act_kwargs or {}))
|
|
158
|
+
self.act_first = act_first
|
|
159
|
+
self.plain_last = plain_last
|
|
160
|
+
|
|
161
|
+
if isinstance(dropout, float):
|
|
162
|
+
dropout = [dropout] * (len(channel_list) - 1)
|
|
163
|
+
if plain_last:
|
|
164
|
+
dropout[-1] = 0.0
|
|
165
|
+
if len(dropout) != len(channel_list) - 1:
|
|
166
|
+
raise ValueError(
|
|
167
|
+
f"Number of dropout values provided ({len(dropout)}) does not "
|
|
168
|
+
f"match the number of layers specified ({len(channel_list) - 1})"
|
|
169
|
+
)
|
|
170
|
+
self.dropout_rate = dropout
|
|
171
|
+
|
|
172
|
+
if isinstance(bias, bool):
|
|
173
|
+
bias = [bias] * (len(channel_list) - 1)
|
|
174
|
+
if len(bias) != len(channel_list) - 1:
|
|
175
|
+
raise ValueError(
|
|
176
|
+
f"Number of bias values provided ({len(bias)}) does not match "
|
|
177
|
+
f"the number of layers specified ({len(channel_list) - 1})"
|
|
178
|
+
)
|
|
179
|
+
|
|
180
|
+
self.lins = []
|
|
181
|
+
for in_c, out_c, _bias in zip(channel_list[:-1], channel_list[1:], bias):
|
|
182
|
+
lin = keras.layers.Dense(out_c, use_bias=_bias)
|
|
183
|
+
lin.build((None, in_c))
|
|
184
|
+
self.lins.append(lin)
|
|
185
|
+
|
|
186
|
+
self.norms = []
|
|
187
|
+
iterator = channel_list[1:-1] if plain_last else channel_list[1:]
|
|
188
|
+
for hc in iterator:
|
|
189
|
+
norm_layer = _normalization_resolver(norm, hc, **(norm_kwargs or {}))
|
|
190
|
+
if norm_layer is not None and hasattr(norm_layer, "build") and not norm_layer.built:
|
|
191
|
+
norm_layer.build((None, hc))
|
|
192
|
+
self.norms.append(norm_layer)
|
|
193
|
+
|
|
194
|
+
self.dropouts = [keras.layers.Dropout(p) if p > 0.0 else None for p in self.dropout_rate]
|
|
195
|
+
|
|
196
|
+
self.supports_norm_batch = False
|
|
197
|
+
if len(self.norms) > 0 and self.norms[0] is not None:
|
|
198
|
+
norm_params = inspect.signature(self.norms[0].call).parameters
|
|
199
|
+
self.supports_norm_batch = "batch" in norm_params
|
|
200
|
+
# Norms such as BatchNorm treat ``training=None`` as training mode, so pass it explicitly.
|
|
201
|
+
self.supports_norm_training = False
|
|
202
|
+
if len(self.norms) > 0 and self.norms[0] is not None:
|
|
203
|
+
self.supports_norm_training = "training" in inspect.signature(self.norms[0].call).parameters
|
|
204
|
+
|
|
205
|
+
@property
|
|
206
|
+
def num_layers(self) -> int:
|
|
207
|
+
r"""The number of layers."""
|
|
208
|
+
return len(self.channel_list) - 1
|
|
209
|
+
|
|
210
|
+
def build(self, input_shape=None):
|
|
211
|
+
for lin in self.lins:
|
|
212
|
+
if hasattr(lin, "built") and not lin.built:
|
|
213
|
+
lin.build(input_shape)
|
|
214
|
+
norm_iter = self.channel_list[1:-1] if self.plain_last else self.channel_list[1:]
|
|
215
|
+
for norm, hc in zip(self.norms, norm_iter):
|
|
216
|
+
if norm is not None and hasattr(norm, "built") and not norm.built:
|
|
217
|
+
norm.build((None, hc))
|
|
218
|
+
for drop in self.dropouts:
|
|
219
|
+
if drop is not None and hasattr(drop, "built") and not drop.built:
|
|
220
|
+
drop.build(input_shape)
|
|
221
|
+
self.built = True
|
|
222
|
+
|
|
223
|
+
def reset_parameters(self):
|
|
224
|
+
r"""Resets all learnable parameters of the module."""
|
|
225
|
+
for lin in self.lins:
|
|
226
|
+
if hasattr(lin, "kernel_initializer") and lin.kernel is not None:
|
|
227
|
+
lin.kernel.assign(lin.kernel_initializer(ops.shape(lin.kernel)))
|
|
228
|
+
if lin.bias is not None:
|
|
229
|
+
lin.bias.assign(lin.bias_initializer(ops.shape(lin.bias)))
|
|
230
|
+
for norm in self.norms:
|
|
231
|
+
if hasattr(norm, "reset_parameters"):
|
|
232
|
+
norm.reset_parameters()
|
|
233
|
+
|
|
234
|
+
def call(self, x, batch=None, batch_size=None, return_emb=None, training=None):
|
|
235
|
+
emb = None
|
|
236
|
+
|
|
237
|
+
# If `plain_last=True`, `len(norms) == len(lins) - 1`, thus skipping
|
|
238
|
+
# execution of the last layer inside the loop.
|
|
239
|
+
for i, (lin, norm) in enumerate(zip(self.lins, self.norms)):
|
|
240
|
+
x = lin(x)
|
|
241
|
+
if self.act is not None and self.act_first:
|
|
242
|
+
x = self.act(x)
|
|
243
|
+
if norm is not None:
|
|
244
|
+
norm_kwargs = {"training": training} if self.supports_norm_training else {}
|
|
245
|
+
if self.supports_norm_batch:
|
|
246
|
+
x = norm(x, batch, batch_size, **norm_kwargs)
|
|
247
|
+
else:
|
|
248
|
+
x = norm(x, **norm_kwargs)
|
|
249
|
+
if self.act is not None and not self.act_first:
|
|
250
|
+
x = self.act(x)
|
|
251
|
+
if self.dropouts[i] is not None:
|
|
252
|
+
x = self.dropouts[i](x, training=training)
|
|
253
|
+
if isinstance(return_emb, bool) and return_emb is True:
|
|
254
|
+
emb = x
|
|
255
|
+
|
|
256
|
+
if self.plain_last:
|
|
257
|
+
x = self.lins[-1](x)
|
|
258
|
+
if self.dropouts[-1] is not None:
|
|
259
|
+
x = self.dropouts[-1](x, training=training)
|
|
260
|
+
|
|
261
|
+
return (x, emb) if isinstance(return_emb, bool) else x
|
|
262
|
+
|
|
263
|
+
def __repr__(self) -> str:
|
|
264
|
+
return f"{self.__class__.__name__}({str(self.channel_list)[1:-1]})"
|