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.
Files changed (459) hide show
  1. k3_node/__init__.py +122 -0
  2. k3_node/applications/__init__.py +17 -0
  3. k3_node/applications/bio/__init__.py +21 -0
  4. k3_node/applications/chemistry/__init__.py +155 -0
  5. k3_node/applications/materials/__init__.py +127 -0
  6. k3_node/applications/materials/basis.py +449 -0
  7. k3_node/applications/materials/chgnet.py +360 -0
  8. k3_node/applications/materials/core.py +351 -0
  9. k3_node/applications/materials/grace.py +246 -0
  10. k3_node/applications/materials/io.py +230 -0
  11. k3_node/applications/materials/m3gnet.py +462 -0
  12. k3_node/applications/materials/megnet.py +395 -0
  13. k3_node/applications/materials/qet.py +220 -0
  14. k3_node/applications/materials/readout.py +235 -0
  15. k3_node/applications/materials/so3net.py +234 -0
  16. k3_node/applications/materials/tensornet.py +381 -0
  17. k3_node/applications/materials/test_materials.py +167 -0
  18. k3_node/applications/materials/wrappers.py +95 -0
  19. k3_node/data/__init__.py +47 -0
  20. k3_node/data/batch.py +102 -0
  21. k3_node/data/collate.py +282 -0
  22. k3_node/data/data.py +532 -0
  23. k3_node/data/database.py +154 -0
  24. k3_node/data/dataset.py +182 -0
  25. k3_node/data/download.py +49 -0
  26. k3_node/data/extract.py +45 -0
  27. k3_node/data/feature_store.py +70 -0
  28. k3_node/data/graph_store.py +92 -0
  29. k3_node/data/hetero_data.py +374 -0
  30. k3_node/data/hypergraph_data.py +59 -0
  31. k3_node/data/in_memory_dataset.py +177 -0
  32. k3_node/data/makedirs.py +7 -0
  33. k3_node/data/on_disk_dataset.py +77 -0
  34. k3_node/data/separate.py +115 -0
  35. k3_node/data/storage.py +593 -0
  36. k3_node/data/temporal.py +154 -0
  37. k3_node/data/test_batch.py +67 -0
  38. k3_node/data/test_data.py +68 -0
  39. k3_node/data/test_dataset_and_stores.py +111 -0
  40. k3_node/data/test_hetero_data.py +33 -0
  41. k3_node/data/test_temporal_and_hyper.py +32 -0
  42. k3_node/data/view.py +43 -0
  43. k3_node/datasets/__init__.py +88 -0
  44. k3_node/datasets/actor.py +101 -0
  45. k3_node/datasets/airports.py +84 -0
  46. k3_node/datasets/amazon.py +66 -0
  47. k3_node/datasets/ba2motif_dataset.py +73 -0
  48. k3_node/datasets/ba_shapes.py +81 -0
  49. k3_node/datasets/bitcoin_otc.py +77 -0
  50. k3_node/datasets/citation_full.py +81 -0
  51. k3_node/datasets/coauthor.py +66 -0
  52. k3_node/datasets/dblp.py +106 -0
  53. k3_node/datasets/digits.py +63 -0
  54. k3_node/datasets/email_eu_core.py +60 -0
  55. k3_node/datasets/entities.py +158 -0
  56. k3_node/datasets/explainer_dataset.py +101 -0
  57. k3_node/datasets/facebook.py +51 -0
  58. k3_node/datasets/fake.py +256 -0
  59. k3_node/datasets/freebase.py +90 -0
  60. k3_node/datasets/geometric_shapes.py +69 -0
  61. k3_node/datasets/github.py +51 -0
  62. k3_node/datasets/graph_generator/__init__.py +6 -0
  63. k3_node/datasets/graph_generator/ba_graph.py +20 -0
  64. k3_node/datasets/graph_generator/base.py +29 -0
  65. k3_node/datasets/graph_generator/er_graph.py +21 -0
  66. k3_node/datasets/icews.py +58 -0
  67. k3_node/datasets/imdb.py +96 -0
  68. k3_node/datasets/jodie.py +56 -0
  69. k3_node/datasets/karate.py +56 -0
  70. k3_node/datasets/lastfm_asia.py +51 -0
  71. k3_node/datasets/mesh_correspondence.py +50 -0
  72. k3_node/datasets/molecule_net.py +148 -0
  73. k3_node/datasets/motif_generator/__init__.py +7 -0
  74. k3_node/datasets/motif_generator/base.py +29 -0
  75. k3_node/datasets/motif_generator/custom.py +17 -0
  76. k3_node/datasets/motif_generator/cycle.py +25 -0
  77. k3_node/datasets/motif_generator/house.py +27 -0
  78. k3_node/datasets/movielens.py +55 -0
  79. k3_node/datasets/planetoid.py +137 -0
  80. k3_node/datasets/polblogs.py +63 -0
  81. k3_node/datasets/ppi.py +189 -0
  82. k3_node/datasets/qm7.py +65 -0
  83. k3_node/datasets/qm9.py +132 -0
  84. k3_node/datasets/reddit.py +121 -0
  85. k3_node/datasets/sbm_dataset.py +165 -0
  86. k3_node/datasets/seal.py +74 -0
  87. k3_node/datasets/shape_scenes.py +92 -0
  88. k3_node/datasets/test_datasets.py +322 -0
  89. k3_node/datasets/tu_dataset.py +131 -0
  90. k3_node/datasets/twitch.py +66 -0
  91. k3_node/datasets/webkb.py +102 -0
  92. k3_node/datasets/wikics.py +85 -0
  93. k3_node/datasets/word_net.py +184 -0
  94. k3_node/etl/__init__.py +37 -0
  95. k3_node/etl/encoders.py +248 -0
  96. k3_node/etl/graph_builders.py +270 -0
  97. k3_node/etl/relational_to_graph.py +201 -0
  98. k3_node/etl/table_to_graph.py +244 -0
  99. k3_node/etl/test_etl.py +318 -0
  100. k3_node/export/__init__.py +15 -0
  101. k3_node/export/cross_backend.py +172 -0
  102. k3_node/export/onnx_exporter.py +190 -0
  103. k3_node/export/runtime.py +254 -0
  104. k3_node/export/tensorrt_exporter.py +201 -0
  105. k3_node/export/test_export.py +337 -0
  106. k3_node/export/tflite_exporter.py +112 -0
  107. k3_node/hub/__init__.py +29 -0
  108. k3_node/hub/dataset_hub.py +242 -0
  109. k3_node/hub/hub_mixin.py +599 -0
  110. k3_node/hub/model_card.py +133 -0
  111. k3_node/hub/test_hub.py +419 -0
  112. k3_node/io/__init__.py +22 -0
  113. k3_node/io/fs.py +117 -0
  114. k3_node/io/npz.py +45 -0
  115. k3_node/io/off.py +29 -0
  116. k3_node/io/planetoid.py +98 -0
  117. k3_node/io/tu.py +137 -0
  118. k3_node/io/txt_array.py +58 -0
  119. k3_node/layers/__init__.py +14 -0
  120. k3_node/layers/aggr/__init__.py +70 -0
  121. k3_node/layers/aggr/attention.py +77 -0
  122. k3_node/layers/aggr/base.py +403 -0
  123. k3_node/layers/aggr/basic.py +412 -0
  124. k3_node/layers/aggr/deep_sets.py +65 -0
  125. k3_node/layers/aggr/deepsets.py +29 -0
  126. k3_node/layers/aggr/equilibrium.py +107 -0
  127. k3_node/layers/aggr/fused.py +43 -0
  128. k3_node/layers/aggr/gmt.py +89 -0
  129. k3_node/layers/aggr/gru.py +58 -0
  130. k3_node/layers/aggr/lcm.py +143 -0
  131. k3_node/layers/aggr/lstm.py +58 -0
  132. k3_node/layers/aggr/mlp.py +75 -0
  133. k3_node/layers/aggr/multi.py +154 -0
  134. k3_node/layers/aggr/patch_transformer.py +137 -0
  135. k3_node/layers/aggr/quantile.py +125 -0
  136. k3_node/layers/aggr/resolver.py +68 -0
  137. k3_node/layers/aggr/scaler.py +133 -0
  138. k3_node/layers/aggr/set2set.py +87 -0
  139. k3_node/layers/aggr/set_transformer.py +107 -0
  140. k3_node/layers/aggr/sort.py +68 -0
  141. k3_node/layers/aggr/test_aggr.py +337 -0
  142. k3_node/layers/aggr/utils.py +210 -0
  143. k3_node/layers/aggr/variance_preserving.py +54 -0
  144. k3_node/layers/attention/__init__.py +5 -0
  145. k3_node/layers/attention/pair_attention.py +448 -0
  146. k3_node/layers/attention/performer.py +187 -0
  147. k3_node/layers/attention/polynormer.py +160 -0
  148. k3_node/layers/attention/qformer.py +143 -0
  149. k3_node/layers/attention/sgformer.py +106 -0
  150. k3_node/layers/attention/test_attention.py +68 -0
  151. k3_node/layers/attention/test_pair_attention.py +91 -0
  152. k3_node/layers/conv/__init__.py +149 -0
  153. k3_node/layers/conv/agnn_conv.py +120 -0
  154. k3_node/layers/conv/antisymmetric_conv.py +94 -0
  155. k3_node/layers/conv/appnp.py +105 -0
  156. k3_node/layers/conv/appnp_conv.py +157 -0
  157. k3_node/layers/conv/arma_conv.py +231 -0
  158. k3_node/layers/conv/cg_conv.py +92 -0
  159. k3_node/layers/conv/cheb_conv.py +137 -0
  160. k3_node/layers/conv/cluster_gcn_conv.py +102 -0
  161. k3_node/layers/conv/conv.py +100 -0
  162. k3_node/layers/conv/crystal_conv.py +140 -0
  163. k3_node/layers/conv/cugraph.py +84 -0
  164. k3_node/layers/conv/diffusion_conv.py +144 -0
  165. k3_node/layers/conv/dir_gnn_conv.py +93 -0
  166. k3_node/layers/conv/dna_conv.py +192 -0
  167. k3_node/layers/conv/edge_conv.py +107 -0
  168. k3_node/layers/conv/eg_conv.py +155 -0
  169. k3_node/layers/conv/fa_conv.py +107 -0
  170. k3_node/layers/conv/feast_conv.py +126 -0
  171. k3_node/layers/conv/film_conv.py +143 -0
  172. k3_node/layers/conv/gat_conv.py +244 -0
  173. k3_node/layers/conv/gated_graph_conv.py +136 -0
  174. k3_node/layers/conv/gatv2_conv.py +205 -0
  175. k3_node/layers/conv/gcn.py +144 -0
  176. k3_node/layers/conv/gcn2_conv.py +126 -0
  177. k3_node/layers/conv/gcn_conv.py +135 -0
  178. k3_node/layers/conv/gen_conv.py +163 -0
  179. k3_node/layers/conv/general_conv.py +218 -0
  180. k3_node/layers/conv/gin_conv.py +218 -0
  181. k3_node/layers/conv/gmm_conv.py +172 -0
  182. k3_node/layers/conv/gps_conv.py +153 -0
  183. k3_node/layers/conv/graph_attention.py +262 -0
  184. k3_node/layers/conv/graph_conv.py +84 -0
  185. k3_node/layers/conv/gravnet_conv.py +93 -0
  186. k3_node/layers/conv/han_conv.py +175 -0
  187. k3_node/layers/conv/heat_conv.py +131 -0
  188. k3_node/layers/conv/hetero_conv.py +128 -0
  189. k3_node/layers/conv/hgt_conv.py +218 -0
  190. k3_node/layers/conv/hypergraph_conv.py +182 -0
  191. k3_node/layers/conv/le_conv.py +81 -0
  192. k3_node/layers/conv/lg_conv.py +58 -0
  193. k3_node/layers/conv/meshcnn_conv.py +84 -0
  194. k3_node/layers/conv/message_passing.py +451 -0
  195. k3_node/layers/conv/mf_conv.py +95 -0
  196. k3_node/layers/conv/mixhop_conv.py +108 -0
  197. k3_node/layers/conv/nn_conv.py +110 -0
  198. k3_node/layers/conv/pan_conv.py +100 -0
  199. k3_node/layers/conv/pdn_conv.py +109 -0
  200. k3_node/layers/conv/pna_conv.py +177 -0
  201. k3_node/layers/conv/point_conv.py +101 -0
  202. k3_node/layers/conv/point_gnn_conv.py +90 -0
  203. k3_node/layers/conv/point_transformer_conv.py +132 -0
  204. k3_node/layers/conv/ppf_conv.py +135 -0
  205. k3_node/layers/conv/ppnp.py +89 -0
  206. k3_node/layers/conv/res_gated_graph_conv.py +126 -0
  207. k3_node/layers/conv/rgat_conv.py +251 -0
  208. k3_node/layers/conv/rgcn_conv.py +321 -0
  209. k3_node/layers/conv/sage_conv.py +154 -0
  210. k3_node/layers/conv/sg_conv.py +96 -0
  211. k3_node/layers/conv/signed_conv.py +100 -0
  212. k3_node/layers/conv/simple_conv.py +75 -0
  213. k3_node/layers/conv/spline_conv.py +182 -0
  214. k3_node/layers/conv/ssg_conv.py +101 -0
  215. k3_node/layers/conv/supergat_conv.py +195 -0
  216. k3_node/layers/conv/tag_conv.py +98 -0
  217. k3_node/layers/conv/test_backend_consistency.py +164 -0
  218. k3_node/layers/conv/test_conv.py +176 -0
  219. k3_node/layers/conv/test_conv_pyg.py +566 -0
  220. k3_node/layers/conv/transformer_conv.py +168 -0
  221. k3_node/layers/conv/utils.py +403 -0
  222. k3_node/layers/conv/wl_conv.py +151 -0
  223. k3_node/layers/conv/x_conv.py +187 -0
  224. k3_node/layers/dense/__init__.py +40 -0
  225. k3_node/layers/dense/dense_gat_conv.py +149 -0
  226. k3_node/layers/dense/dense_gcn_conv.py +117 -0
  227. k3_node/layers/dense/dense_gin_conv.py +88 -0
  228. k3_node/layers/dense/dense_graph_conv.py +95 -0
  229. k3_node/layers/dense/dense_sage_conv.py +85 -0
  230. k3_node/layers/dense/diff_pool.py +76 -0
  231. k3_node/layers/dense/dmon_pool.py +223 -0
  232. k3_node/layers/dense/linear.py +327 -0
  233. k3_node/layers/dense/mincut_pool.py +92 -0
  234. k3_node/layers/dense/test_dense.py +377 -0
  235. k3_node/layers/functional/__init__.py +13 -0
  236. k3_node/layers/functional/bro.py +49 -0
  237. k3_node/layers/functional/edge_dropout.py +55 -0
  238. k3_node/layers/functional/gini.py +44 -0
  239. k3_node/layers/functional/test_functional.py +34 -0
  240. k3_node/layers/kge/__init__.py +17 -0
  241. k3_node/layers/kge/base.py +255 -0
  242. k3_node/layers/kge/complex.py +98 -0
  243. k3_node/layers/kge/distmult.py +79 -0
  244. k3_node/layers/kge/loader.py +50 -0
  245. k3_node/layers/kge/rotate.py +103 -0
  246. k3_node/layers/kge/test_kge.py +76 -0
  247. k3_node/layers/kge/transe.py +96 -0
  248. k3_node/layers/norm/__init__.py +23 -0
  249. k3_node/layers/norm/batch_norm.py +328 -0
  250. k3_node/layers/norm/diff_group_norm.py +141 -0
  251. k3_node/layers/norm/graph_norm.py +105 -0
  252. k3_node/layers/norm/graph_size_norm.py +57 -0
  253. k3_node/layers/norm/instance_norm.py +163 -0
  254. k3_node/layers/norm/layer_norm.py +245 -0
  255. k3_node/layers/norm/mean_subtraction_norm.py +57 -0
  256. k3_node/layers/norm/msg_norm.py +58 -0
  257. k3_node/layers/norm/pair_norm.py +94 -0
  258. k3_node/layers/norm/test_norm.py +275 -0
  259. k3_node/layers/pool/__init__.py +83 -0
  260. k3_node/layers/pool/approx_knn.py +101 -0
  261. k3_node/layers/pool/asap.py +173 -0
  262. k3_node/layers/pool/avg_pool.py +165 -0
  263. k3_node/layers/pool/cluster_pool.py +168 -0
  264. k3_node/layers/pool/connect/__init__.py +10 -0
  265. k3_node/layers/pool/connect/base.py +103 -0
  266. k3_node/layers/pool/connect/filter_edges.py +113 -0
  267. k3_node/layers/pool/consecutive.py +30 -0
  268. k3_node/layers/pool/decimation.py +48 -0
  269. k3_node/layers/pool/edge_pool.py +189 -0
  270. k3_node/layers/pool/glob.py +139 -0
  271. k3_node/layers/pool/graclus.py +66 -0
  272. k3_node/layers/pool/knn.py +253 -0
  273. k3_node/layers/pool/max_pool.py +159 -0
  274. k3_node/layers/pool/mem_pool.py +145 -0
  275. k3_node/layers/pool/pan_pool.py +144 -0
  276. k3_node/layers/pool/point_cloud.py +212 -0
  277. k3_node/layers/pool/pool.py +119 -0
  278. k3_node/layers/pool/sag_pool.py +174 -0
  279. k3_node/layers/pool/select/__init__.py +10 -0
  280. k3_node/layers/pool/select/base.py +112 -0
  281. k3_node/layers/pool/select/topk.py +206 -0
  282. k3_node/layers/pool/test_pool.py +456 -0
  283. k3_node/layers/pool/topk_pool.py +103 -0
  284. k3_node/layers/pool/voxel_grid.py +70 -0
  285. k3_node/layers/unpool/__init__.py +9 -0
  286. k3_node/layers/unpool/knn_interpolate.py +57 -0
  287. k3_node/layers/unpool/test_unpool.py +31 -0
  288. k3_node/loader/__init__.py +62 -0
  289. k3_node/loader/base.py +69 -0
  290. k3_node/loader/cache.py +68 -0
  291. k3_node/loader/cluster.py +127 -0
  292. k3_node/loader/data_list_loader.py +45 -0
  293. k3_node/loader/dataloader.py +117 -0
  294. k3_node/loader/dense_data_loader.py +62 -0
  295. k3_node/loader/dynamic_batch_sampler.py +93 -0
  296. k3_node/loader/graph_saint.py +188 -0
  297. k3_node/loader/hgt_loader.py +90 -0
  298. k3_node/loader/imbalanced_sampler.py +87 -0
  299. k3_node/loader/keras_dataset.py +334 -0
  300. k3_node/loader/link_loader.py +179 -0
  301. k3_node/loader/link_neighbor_loader.py +202 -0
  302. k3_node/loader/mixin.py +190 -0
  303. k3_node/loader/neighbor_loader.py +159 -0
  304. k3_node/loader/neighbor_sampler.py +167 -0
  305. k3_node/loader/node_loader.py +185 -0
  306. k3_node/loader/prefetch.py +115 -0
  307. k3_node/loader/random_node_loader.py +89 -0
  308. k3_node/loader/sampler_utils.py +499 -0
  309. k3_node/loader/shadow.py +115 -0
  310. k3_node/loader/temporal_dataloader.py +98 -0
  311. k3_node/loader/test_dataloader.py +113 -0
  312. k3_node/loader/test_keras_dataset.py +221 -0
  313. k3_node/loader/test_neighbor_loader.py +122 -0
  314. k3_node/loader/test_sampler_utils.py +82 -0
  315. k3_node/loader/test_samplers.py +96 -0
  316. k3_node/loader/test_subgraph_loaders.py +89 -0
  317. k3_node/loader/utils.py +232 -0
  318. k3_node/loader/zip_loader.py +88 -0
  319. k3_node/metrics.py +94 -0
  320. k3_node/models/__init__.py +424 -0
  321. k3_node/models/attentive_fp.py +232 -0
  322. k3_node/models/attract_repel.py +108 -0
  323. k3_node/models/autoencoder.py +318 -0
  324. k3_node/models/basic_gnn.py +443 -0
  325. k3_node/models/bio/__init__.py +4 -0
  326. k3_node/models/captum.py +52 -0
  327. k3_node/models/chemistry/__init__.py +4 -0
  328. k3_node/models/correct_and_smooth.py +146 -0
  329. k3_node/models/deep_graph_infomax.py +113 -0
  330. k3_node/models/deepgcn.py +121 -0
  331. k3_node/models/dimenet.py +737 -0
  332. k3_node/models/dimenet_utils.py +153 -0
  333. k3_node/models/gnnff.py +263 -0
  334. k3_node/models/gps_model.py +1122 -0
  335. k3_node/models/gpse.py +638 -0
  336. k3_node/models/graph_unet.py +199 -0
  337. k3_node/models/graphmae2.py +954 -0
  338. k3_node/models/graphormer.py +1258 -0
  339. k3_node/models/graphormer_3d.py +868 -0
  340. k3_node/models/grover.py +1066 -0
  341. k3_node/models/jumping_knowledge.py +200 -0
  342. k3_node/models/label_prop.py +110 -0
  343. k3_node/models/lightgcn.py +171 -0
  344. k3_node/models/linkx.py +181 -0
  345. k3_node/models/lpformer.py +404 -0
  346. k3_node/models/mask_label.py +114 -0
  347. k3_node/models/materials/__init__.py +33 -0
  348. k3_node/models/meta.py +133 -0
  349. k3_node/models/metapath2vec.py +234 -0
  350. k3_node/models/mlp.py +264 -0
  351. k3_node/models/mole_bert.py +379 -0
  352. k3_node/models/neural_fingerprint.py +95 -0
  353. k3_node/models/node2vec.py +213 -0
  354. k3_node/models/pmlp.py +157 -0
  355. k3_node/models/polynormer.py +229 -0
  356. k3_node/models/rect.py +93 -0
  357. k3_node/models/renet.py +221 -0
  358. k3_node/models/rev_gnn.py +128 -0
  359. k3_node/models/schnet.py +484 -0
  360. k3_node/models/sgformer.py +195 -0
  361. k3_node/models/signed_gcn.py +185 -0
  362. k3_node/models/test_attentive_fp.py +32 -0
  363. k3_node/models/test_attract_repel.py +33 -0
  364. k3_node/models/test_autoencoder.py +119 -0
  365. k3_node/models/test_basic_gnn.py +102 -0
  366. k3_node/models/test_correct_and_smooth.py +40 -0
  367. k3_node/models/test_deep_graph_infomax.py +68 -0
  368. k3_node/models/test_deepgcn.py +21 -0
  369. k3_node/models/test_dimenet.py +86 -0
  370. k3_node/models/test_domain_apis.py +138 -0
  371. k3_node/models/test_gnnff.py +24 -0
  372. k3_node/models/test_gps_model.py +271 -0
  373. k3_node/models/test_gpse.py +34 -0
  374. k3_node/models/test_graph_unet.py +26 -0
  375. k3_node/models/test_graphmae2.py +226 -0
  376. k3_node/models/test_graphormer.py +233 -0
  377. k3_node/models/test_graphormer3d.py +163 -0
  378. k3_node/models/test_grover.py +287 -0
  379. k3_node/models/test_jumping_knowledge.py +129 -0
  380. k3_node/models/test_label_prop.py +37 -0
  381. k3_node/models/test_lightgcn.py +38 -0
  382. k3_node/models/test_linkx.py +31 -0
  383. k3_node/models/test_lpformer.py +22 -0
  384. k3_node/models/test_mask_label.py +90 -0
  385. k3_node/models/test_meta.py +159 -0
  386. k3_node/models/test_metapath2vec.py +45 -0
  387. k3_node/models/test_mlp.py +62 -0
  388. k3_node/models/test_mole_bert.py +164 -0
  389. k3_node/models/test_neural_fingerprint.py +13 -0
  390. k3_node/models/test_node2vec.py +57 -0
  391. k3_node/models/test_pmlp.py +81 -0
  392. k3_node/models/test_polynormer.py +104 -0
  393. k3_node/models/test_rect.py +23 -0
  394. k3_node/models/test_renet.py +32 -0
  395. k3_node/models/test_rev_gnn.py +24 -0
  396. k3_node/models/test_schnet.py +43 -0
  397. k3_node/models/test_sgformer.py +48 -0
  398. k3_node/models/test_signed_gcn.py +28 -0
  399. k3_node/models/test_tgn.py +77 -0
  400. k3_node/models/test_unimol.py +179 -0
  401. k3_node/models/test_unimol2.py +114 -0
  402. k3_node/models/test_unimol_plus.py +131 -0
  403. k3_node/models/test_visnet.py +44 -0
  404. k3_node/models/tgn.py +382 -0
  405. k3_node/models/unimol.py +1156 -0
  406. k3_node/models/unimol2.py +616 -0
  407. k3_node/models/unimol_docking_v2.py +301 -0
  408. k3_node/models/unimol_plus.py +456 -0
  409. k3_node/models/utils.py +97 -0
  410. k3_node/models/visnet.py +759 -0
  411. k3_node/ops/__init__.py +4 -0
  412. k3_node/ops/conv.py +56 -0
  413. k3_node/ops/creation.py +43 -0
  414. k3_node/ops/graph.py +27 -0
  415. k3_node/ops/host.py +41 -0
  416. k3_node/ops/matmul.py +49 -0
  417. k3_node/ops/numpy.py +24 -0
  418. k3_node/ops/segment.py +54 -0
  419. k3_node/ops/sparse.py +51 -0
  420. k3_node/rag/__init__.py +49 -0
  421. k3_node/rag/encoders.py +312 -0
  422. k3_node/rag/pipeline.py +192 -0
  423. k3_node/rag/projector.py +184 -0
  424. k3_node/rag/subgraph.py +270 -0
  425. k3_node/rag/test_rag.py +347 -0
  426. k3_node/rag/verbalizer.py +162 -0
  427. k3_node/tasks/__init__.py +19 -0
  428. k3_node/tasks/backbone_resolver.py +125 -0
  429. k3_node/tasks/base.py +67 -0
  430. k3_node/tasks/graph_classification.py +270 -0
  431. k3_node/tasks/graph_regression.py +228 -0
  432. k3_node/tasks/link_prediction.py +306 -0
  433. k3_node/tasks/node_classification.py +194 -0
  434. k3_node/tasks/node_regression.py +138 -0
  435. k3_node/tasks/test_tasks.py +319 -0
  436. k3_node/test_docstring_examples.py +106 -0
  437. k3_node/test_training_forwarding.py +116 -0
  438. k3_node/training.py +115 -0
  439. k3_node/transforms/__init__.py +166 -0
  440. k3_node/transforms/base_transform.py +32 -0
  441. k3_node/transforms/compose.py +58 -0
  442. k3_node/transforms/general.py +676 -0
  443. k3_node/transforms/graph.py +1070 -0
  444. k3_node/transforms/spatial.py +797 -0
  445. k3_node/transforms/test_random_link_split.py +45 -0
  446. k3_node/transforms/test_spatial_transforms.py +65 -0
  447. k3_node/transforms/test_transforms.py +253 -0
  448. k3_node/transforms/utils.py +102 -0
  449. k3_node/utils/__init__.py +5 -0
  450. k3_node/utils/backend_import.py +12 -0
  451. k3_node/utils/graph.py +286 -0
  452. k3_node/utils/keras.py +94 -0
  453. k3_node/utils/random.py +103 -0
  454. k3_node/utils/smiles.py +235 -0
  455. k3_node-1.0.0.dist-info/METADATA +284 -0
  456. k3_node-1.0.0.dist-info/RECORD +459 -0
  457. k3_node-1.0.0.dist-info/WHEEL +5 -0
  458. k3_node-1.0.0.dist-info/licenses/LICENSE +21 -0
  459. k3_node-1.0.0.dist-info/top_level.txt +1 -0
@@ -0,0 +1,57 @@
1
+ from keras import layers, ops
2
+ from k3_node.ops.segment import segment_sum
3
+
4
+
5
+ class GraphSizeNorm(layers.Layer):
6
+ r"""Applies Graph Size Normalization over each individual graph in a batch
7
+ of node features:
8
+
9
+ .. math::
10
+ \mathbf{x}^{\prime}_i = \frac{\mathbf{x}_i}{\sqrt{|\mathcal{V}|}}
11
+
12
+ Example:
13
+ ```python
14
+ import numpy as np
15
+ from k3_node.layers import GraphSizeNorm
16
+
17
+ x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
18
+ batch = np.repeat([0, 1], 5) # nodes 0-4 belong to graph 0, nodes 5-9 to graph 1
19
+
20
+ layer = GraphSizeNorm()
21
+ out = layer(x, batch) # normalizes each graph separately
22
+ print(tuple(out.shape)) # (10, 8)
23
+ ```
24
+ """
25
+ def __init__(self, **kwargs):
26
+ super().__init__(**kwargs)
27
+
28
+ def call(self, x, batch=None, batch_size=None):
29
+ if batch is None and isinstance(x, (tuple, list)):
30
+ if len(x) == 2:
31
+ x, batch = x
32
+ elif len(x) == 3:
33
+ x, batch, batch_size = x
34
+
35
+ if batch is None:
36
+ num_nodes = ops.cast(ops.shape(x)[0], dtype=x.dtype)
37
+ return x * ops.power(num_nodes, -0.5)
38
+
39
+ if batch_size is not None and not isinstance(batch_size, int):
40
+ try:
41
+ batch_size = int(batch_size)
42
+ except Exception:
43
+ pass
44
+ elif batch_size is None:
45
+ batch_size = ops.cast(ops.max(batch), "int32") + 1
46
+
47
+ batch = ops.cast(batch, "int32")
48
+ ones = ops.ones((ops.shape(x)[0], 1), dtype=x.dtype)
49
+ deg = segment_sum(ones, batch, num_segments=batch_size)
50
+ inv_sqrt_deg = ops.power(deg, -0.5)
51
+ scale = ops.take(inv_sqrt_deg, batch, axis=0)
52
+ return x * scale
53
+
54
+ def compute_output_shape(self, input_shape):
55
+ if isinstance(input_shape, (tuple, list)) and isinstance(input_shape[0], (tuple, list)):
56
+ return input_shape[0]
57
+ return input_shape
@@ -0,0 +1,163 @@
1
+ from keras import layers, ops
2
+ from k3_node.ops.segment import segment_sum
3
+
4
+
5
+ class InstanceNorm(layers.Layer):
6
+ r"""Applies instance normalization over each individual example in a batch
7
+ of node features as described in the `"Instance Normalization: The Missing
8
+ Ingredient for Fast Stylization" <https://arxiv.org/abs/1607.06450>`_
9
+ paper.
10
+
11
+ .. math::
12
+ \mathbf{x}^{\prime}_i = \frac{\mathbf{x} -
13
+ \textrm{E}[\mathbf{x}]}{\sqrt{\textrm{Var}[\mathbf{x}] + \epsilon}}
14
+ \odot \gamma + \beta
15
+
16
+ Args:
17
+ in_channels (int): Size of each input sample.
18
+ eps (float, optional): A value added to the denominator for numerical
19
+ stability. (default: :obj:`1e-5`)
20
+ momentum (float, optional): The value used for the running mean and
21
+ running variance computation. (default: :obj:`0.1`)
22
+ affine (bool, optional): If set to :obj:`True`, this module has
23
+ learnable affine parameters :math:`\gamma` and :math:`\beta`.
24
+ (default: :obj:`False`)
25
+ track_running_stats (bool, optional): If set to :obj:`True`, this
26
+ module tracks the running mean and variance, and when set to
27
+ :obj:`False`, this module does not track such statistics and always
28
+ uses instance statistics in both training and eval modes.
29
+ (default: :obj:`False`)
30
+
31
+ Example:
32
+ ```python
33
+ import numpy as np
34
+ from k3_node.layers import InstanceNorm
35
+
36
+ x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
37
+ batch = np.repeat([0, 1], 5) # nodes 0-4 belong to graph 0, nodes 5-9 to graph 1
38
+
39
+ layer = InstanceNorm(in_channels=8)
40
+ out = layer(x, batch) # normalizes each graph separately
41
+ print(tuple(out.shape)) # (10, 8)
42
+ ```
43
+ """
44
+ def __init__(
45
+ self,
46
+ in_channels: int,
47
+ eps: float = 1e-5,
48
+ momentum: float = 0.1,
49
+ affine: bool = False,
50
+ track_running_stats: bool = False,
51
+ **kwargs
52
+ ):
53
+ super().__init__(**kwargs)
54
+ self.in_channels = in_channels
55
+ self.eps = eps
56
+ self.momentum = momentum
57
+ self.affine = affine
58
+ self.track_running_stats = track_running_stats
59
+
60
+ if affine:
61
+ self.weight = self.add_weight(
62
+ shape=(in_channels,),
63
+ initializer="ones",
64
+ trainable=True,
65
+ name="weight",
66
+ )
67
+ self.bias = self.add_weight(
68
+ shape=(in_channels,),
69
+ initializer="zeros",
70
+ trainable=True,
71
+ name="bias",
72
+ )
73
+ else:
74
+ self.weight = None
75
+ self.bias = None
76
+
77
+ if track_running_stats:
78
+ self.running_mean = self.add_weight(
79
+ shape=(in_channels,),
80
+ initializer="zeros",
81
+ trainable=False,
82
+ name="running_mean",
83
+ )
84
+ self.running_var = self.add_weight(
85
+ shape=(in_channels,),
86
+ initializer="ones",
87
+ trainable=False,
88
+ name="running_var",
89
+ )
90
+ else:
91
+ self.running_mean = None
92
+ self.running_var = None
93
+
94
+ def reset_running_stats(self):
95
+ if self.track_running_stats:
96
+ self.running_mean.assign(ops.zeros(self.running_mean.shape, dtype=self.running_mean.dtype))
97
+ self.running_var.assign(ops.ones(self.running_var.shape, dtype=self.running_var.dtype))
98
+
99
+ def reset_parameters(self):
100
+ self.reset_running_stats()
101
+ if self.affine:
102
+ self.weight.assign(ops.ones(self.weight.shape, dtype=self.weight.dtype))
103
+ self.bias.assign(ops.zeros(self.bias.shape, dtype=self.bias.dtype))
104
+
105
+ def call(self, x, batch=None, batch_size=None, training=None):
106
+ if batch is None and isinstance(x, (tuple, list)):
107
+ if len(x) == 2:
108
+ x, batch = x
109
+ elif len(x) == 3:
110
+ x, batch, batch_size = x
111
+
112
+ # Keras semantics: `training=None` means inference; fit() passes training=True via the call context.
113
+ is_training = bool(training) if training is not None else False
114
+
115
+ if batch is None:
116
+ batch = ops.zeros((ops.shape(x)[0],), dtype="int32")
117
+ batch_size = 1
118
+ elif batch_size is not None and not isinstance(batch_size, int):
119
+ try:
120
+ batch_size = int(batch_size)
121
+ except Exception:
122
+ pass
123
+ elif batch_size is None:
124
+ batch_size = ops.cast(ops.max(batch), "int32") + 1
125
+
126
+ batch = ops.cast(batch, "int32")
127
+
128
+ if is_training or not self.track_running_stats:
129
+ ones = ops.ones((ops.shape(x)[0], 1), dtype=x.dtype)
130
+ counts = ops.maximum(segment_sum(ones, batch, num_segments=batch_size), 1.0)
131
+ unbiased_counts = ops.maximum(counts - 1.0, 1.0)
132
+
133
+ mean = segment_sum(x, batch, num_segments=batch_size) / counts
134
+ x_c = x - ops.take(mean, batch, axis=0)
135
+ sq_diff = segment_sum(ops.power(x_c, 2), batch, num_segments=batch_size)
136
+ var = sq_diff / counts
137
+ unbiased_var = sq_diff / unbiased_counts
138
+
139
+ if is_training and self.track_running_stats:
140
+ m = self.momentum
141
+ cur_mean = ops.mean(mean, axis=0)
142
+ cur_var = ops.mean(unbiased_var, axis=0)
143
+ new_running_mean = (1.0 - m) * self.running_mean + m * cur_mean
144
+ new_running_var = (1.0 - m) * self.running_var + m * cur_var
145
+ self.running_mean.assign(new_running_mean)
146
+ self.running_var.assign(new_running_var)
147
+
148
+ std_x = ops.take(ops.sqrt(var + self.eps), batch, axis=0)
149
+ out = x_c / std_x
150
+ else:
151
+ x_c = x - self.running_mean
152
+ std = ops.sqrt(self.running_var + self.eps)
153
+ out = x_c / std
154
+
155
+ if self.affine:
156
+ out = out * self.weight + self.bias
157
+
158
+ return out
159
+
160
+ def compute_output_shape(self, input_shape):
161
+ if isinstance(input_shape, (tuple, list)) and isinstance(input_shape[0], (tuple, list)):
162
+ return input_shape[0]
163
+ return input_shape
@@ -0,0 +1,245 @@
1
+ from typing import List, Optional, Union
2
+ from keras import layers, ops
3
+ from k3_node.ops.segment import segment_sum
4
+
5
+
6
+ class LayerNorm(layers.Layer):
7
+ r"""Applies layer normalization over each individual example in a batch
8
+ of features as described in the `"Layer Normalization"
9
+ <https://arxiv.org/abs/1607.06450>`_ paper.
10
+
11
+ .. math::
12
+ \mathbf{x}^{\prime}_i = \frac{\mathbf{x} -
13
+ \textrm{E}[\mathbf{x}]}{\sqrt{\textrm{Var}[\mathbf{x}] + \epsilon}}
14
+ \odot \gamma + \beta
15
+
16
+ Args:
17
+ in_channels (int): Size of each input sample.
18
+ eps (float, optional): A value added to the denominator for numerical
19
+ stability. (default: :obj:`1e-5`)
20
+ affine (bool, optional): If set to :obj:`True`, this module has
21
+ learnable affine parameters :math:`\gamma` and :math:`\beta`.
22
+ (default: :obj:`True`)
23
+ mode (str, optional): The normalization mode to use for layer
24
+ normalization (:obj:`"graph"` or :obj:`"node"`). If :obj:`"graph"`
25
+ is used, each graph will be considered as an element to be
26
+ normalized. If `"node"` is used, each node will be considered as
27
+ an element to be normalized. (default: :obj:`"graph"`)
28
+
29
+ Example:
30
+ ```python
31
+ import numpy as np
32
+ from k3_node.layers import LayerNorm
33
+
34
+ x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
35
+ batch = np.repeat([0, 1], 5) # nodes 0-4 belong to graph 0, nodes 5-9 to graph 1
36
+
37
+ layer = LayerNorm(in_channels=8, mode="graph")
38
+ out = layer(x, batch) # normalizes each graph separately
39
+ print(tuple(out.shape)) # (10, 8)
40
+ ```
41
+ """
42
+ def __init__(
43
+ self,
44
+ in_channels: int,
45
+ eps: float = 1e-5,
46
+ affine: bool = True,
47
+ mode: str = 'graph',
48
+ **kwargs
49
+ ):
50
+ super().__init__(**kwargs)
51
+ if mode not in ('graph', 'node'):
52
+ raise ValueError(f"Unknown normalization mode: {mode}")
53
+
54
+ self.in_channels = in_channels
55
+ self.eps = eps
56
+ self.affine = affine
57
+ self.mode = mode
58
+
59
+ if affine:
60
+ self.weight = self.add_weight(
61
+ shape=(in_channels,),
62
+ initializer="ones",
63
+ trainable=True,
64
+ name="weight",
65
+ )
66
+ self.bias = self.add_weight(
67
+ shape=(in_channels,),
68
+ initializer="zeros",
69
+ trainable=True,
70
+ name="bias",
71
+ )
72
+ else:
73
+ self.weight = None
74
+ self.bias = None
75
+
76
+ def reset_parameters(self):
77
+ if self.affine:
78
+ self.weight.assign(ops.ones(self.weight.shape, dtype=self.weight.dtype))
79
+ self.bias.assign(ops.zeros(self.bias.shape, dtype=self.bias.dtype))
80
+
81
+ def call(self, x, batch=None, batch_size=None):
82
+ if batch is None and isinstance(x, (tuple, list)):
83
+ if len(x) == 2:
84
+ x, batch = x
85
+ elif len(x) == 3:
86
+ x, batch, batch_size = x
87
+
88
+ if self.mode == 'graph':
89
+ if batch is None:
90
+ mean = ops.mean(x)
91
+ var = ops.mean(ops.power(x - mean, 2))
92
+ out = (x - mean) / ops.sqrt(var + self.eps)
93
+ else:
94
+ if batch_size is not None and not isinstance(batch_size, int):
95
+ try:
96
+ batch_size = int(batch_size)
97
+ except Exception:
98
+ pass
99
+ elif batch_size is None:
100
+ batch_size = ops.cast(ops.max(batch), "int32") + 1
101
+
102
+ batch = ops.cast(batch, "int32")
103
+ in_channels = ops.cast(ops.shape(x)[-1], dtype=x.dtype)
104
+ ones = ops.ones((ops.shape(x)[0], 1), dtype=x.dtype)
105
+ node_counts = ops.maximum(segment_sum(ones, batch, num_segments=batch_size), 1.0)
106
+ total_count = node_counts * in_channels
107
+
108
+ sum_x = ops.sum(segment_sum(x, batch, num_segments=batch_size), axis=-1, keepdims=True)
109
+ mean = sum_x / total_count
110
+ x_centered = x - ops.take(mean, batch, axis=0)
111
+
112
+ sum_sq = ops.sum(segment_sum(ops.power(x_centered, 2), batch, num_segments=batch_size), axis=-1, keepdims=True)
113
+ var = sum_sq / total_count
114
+ std_x = ops.take(ops.sqrt(var + self.eps), batch, axis=0)
115
+ out = x_centered / std_x
116
+
117
+ if self.affine:
118
+ out = out * self.weight + self.bias
119
+ return out
120
+
121
+ elif self.mode == 'node':
122
+ mean = ops.mean(x, axis=-1, keepdims=True)
123
+ var = ops.var(x, axis=-1, keepdims=True)
124
+ out = (x - mean) / ops.sqrt(var + self.eps)
125
+ if self.affine:
126
+ out = out * self.weight + self.bias
127
+ return out
128
+
129
+ def compute_output_shape(self, input_shape):
130
+ if isinstance(input_shape, (tuple, list)) and isinstance(input_shape[0], (tuple, list)):
131
+ return input_shape[0]
132
+ return input_shape
133
+
134
+
135
+ class HeteroLayerNorm(layers.Layer):
136
+ r"""Applies layer normalization over each individual example in a batch
137
+ of heterogeneous features as described in the `"Layer Normalization"
138
+ <https://arxiv.org/abs/1607.06450>`_ paper.
139
+ Compared to :class:`LayerNorm`, :class:`HeteroLayerNorm` applies
140
+ normalization individually for each node or edge type.
141
+
142
+ Args:
143
+ in_channels (int): Size of each input sample.
144
+ num_types (int): The number of types.
145
+ eps (float, optional): A value added to the denominator for numerical
146
+ stability. (default: :obj:`1e-5`)
147
+ affine (bool, optional): If set to :obj:`True`, this module has
148
+ learnable affine parameters :math:`\gamma` and :math:`\beta`.
149
+ (default: :obj:`True`)
150
+ mode (str, optional): The normalization mode to use for layer
151
+ normalization (:obj:`"node"`). (default: :obj:`"node"`)
152
+
153
+ Example:
154
+ ```python
155
+ import numpy as np
156
+ from k3_node.layers import HeteroLayerNorm
157
+
158
+ x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
159
+ node_type = np.random.randint(0, 3, size=(10,)) # type of each node
160
+
161
+ layer = HeteroLayerNorm(in_channels=8, num_types=3)
162
+ out = layer(x, node_type) # separate statistics per node type
163
+ print(tuple(out.shape)) # (10, 8)
164
+ ```
165
+ """
166
+ def __init__(
167
+ self,
168
+ in_channels: int,
169
+ num_types: int,
170
+ eps: float = 1e-5,
171
+ affine: bool = True,
172
+ mode: str = 'node',
173
+ **kwargs
174
+ ):
175
+ super().__init__(**kwargs)
176
+ if mode != 'node':
177
+ raise ValueError(f"HeteroLayerNorm only supports mode='node' (got '{mode}')")
178
+
179
+ self.in_channels = in_channels
180
+ self.num_types = num_types
181
+ self.eps = eps
182
+ self.affine = affine
183
+ self.mode = mode
184
+
185
+ if affine:
186
+ self.weight = self.add_weight(
187
+ shape=(num_types, in_channels),
188
+ initializer="ones",
189
+ trainable=True,
190
+ name="weight",
191
+ )
192
+ self.bias = self.add_weight(
193
+ shape=(num_types, in_channels),
194
+ initializer="zeros",
195
+ trainable=True,
196
+ name="bias",
197
+ )
198
+ else:
199
+ self.weight = None
200
+ self.bias = None
201
+
202
+ def reset_parameters(self):
203
+ if self.affine:
204
+ self.weight.assign(ops.ones(self.weight.shape, dtype=self.weight.dtype))
205
+ self.bias.assign(ops.zeros(self.bias.shape, dtype=self.bias.dtype))
206
+
207
+ def call(
208
+ self,
209
+ x,
210
+ type_vec=None,
211
+ type_ptr: Optional[Union[list, tuple]] = None,
212
+ ):
213
+ if type_vec is None and isinstance(x, (tuple, list)):
214
+ if len(x) == 2:
215
+ x, type_vec = x
216
+ elif len(x) == 3:
217
+ x, type_vec, type_ptr = x
218
+
219
+ if type_vec is None and type_ptr is None:
220
+ raise ValueError("Either 'type_vec' or 'type_ptr' must be given")
221
+
222
+ mean = ops.mean(x, axis=-1, keepdims=True)
223
+ var = ops.var(x, axis=-1, keepdims=True)
224
+ out = (x - mean) / ops.sqrt(var + self.eps)
225
+
226
+ if self.affine:
227
+ if type_ptr is not None:
228
+ parts = []
229
+ for i in range(len(type_ptr) - 1):
230
+ s, e = type_ptr[i], type_ptr[i + 1]
231
+ part = out[s:e] * self.weight[i] + self.bias[i]
232
+ parts.append(part)
233
+ out = ops.concatenate(parts, axis=0)
234
+ else:
235
+ type_vec = ops.cast(type_vec, "int32")
236
+ w = ops.take(self.weight, type_vec, axis=0)
237
+ b = ops.take(self.bias, type_vec, axis=0)
238
+ out = out * w + b
239
+
240
+ return out
241
+
242
+ def compute_output_shape(self, input_shape):
243
+ if isinstance(input_shape, (tuple, list)) and isinstance(input_shape[0], (tuple, list)):
244
+ return input_shape[0]
245
+ return input_shape
@@ -0,0 +1,57 @@
1
+ from keras import layers, ops
2
+ from k3_node.ops.segment import segment_sum
3
+
4
+
5
+ class MeanSubtractionNorm(layers.Layer):
6
+ r"""Applies layer normalization by subtracting the mean from the inputs
7
+ as described in the `"Revisiting 'Over-smoothing' in Deep GCNs"
8
+ <https://arxiv.org/abs/2003.13663>`_ paper.
9
+
10
+ .. math::
11
+ \mathbf{x}_i = \mathbf{x}_i - \frac{1}{|\mathcal{V}|}
12
+ \sum_{j \in \mathcal{V}} \mathbf{x}_j
13
+
14
+ Example:
15
+ ```python
16
+ import numpy as np
17
+ from k3_node.layers import MeanSubtractionNorm
18
+
19
+ x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
20
+ batch = np.repeat([0, 1], 5) # nodes 0-4 belong to graph 0, nodes 5-9 to graph 1
21
+
22
+ layer = MeanSubtractionNorm()
23
+ out = layer(x, batch) # normalizes each graph separately
24
+ print(tuple(out.shape)) # (10, 8)
25
+ ```
26
+ """
27
+ def __init__(self, **kwargs):
28
+ super().__init__(**kwargs)
29
+
30
+ def call(self, x, batch=None, dim_size=None):
31
+ if batch is None and isinstance(x, (tuple, list)):
32
+ if len(x) == 2:
33
+ x, batch = x
34
+ elif len(x) == 3:
35
+ x, batch, dim_size = x
36
+
37
+ if batch is None:
38
+ return x - ops.mean(x, axis=0, keepdims=True)
39
+
40
+ if dim_size is not None and not isinstance(dim_size, int):
41
+ try:
42
+ dim_size = int(dim_size)
43
+ except Exception:
44
+ pass
45
+ elif dim_size is None:
46
+ dim_size = ops.cast(ops.max(batch), "int32") + 1
47
+
48
+ batch = ops.cast(batch, "int32")
49
+ ones = ops.ones((ops.shape(x)[0], 1), dtype=x.dtype)
50
+ counts = ops.maximum(segment_sum(ones, batch, num_segments=dim_size), 1.0)
51
+ mean = segment_sum(x, batch, num_segments=dim_size) / counts
52
+ return x - ops.take(mean, batch, axis=0)
53
+
54
+ def compute_output_shape(self, input_shape):
55
+ if isinstance(input_shape, (tuple, list)) and isinstance(input_shape[0], (tuple, list)):
56
+ return input_shape[0]
57
+ return input_shape
@@ -0,0 +1,58 @@
1
+ from keras import layers, ops
2
+
3
+
4
+ class MessageNorm(layers.Layer):
5
+ r"""Applies message normalization over the aggregated messages as described
6
+ in the `"DeeperGCNs: All You Need to Train Deeper GCNs"
7
+ <https://arxiv.org/abs/2006.07739>`_ paper.
8
+
9
+ .. math::
10
+
11
+ \mathbf{x}_i^{\prime} = \mathbf{x}_{i} + s \cdot
12
+ {\| \mathbf{x}_i \|}_2 \cdot
13
+ \frac{\mathbf{m}_{i}}{{\|\mathbf{m}_i\|}_2}
14
+
15
+ Args:
16
+ learn_scale (bool, optional): If set to :obj:`True`, will learn the
17
+ scaling factor :math:`s` of message normalization.
18
+ (default: :obj:`False`)
19
+
20
+ Example:
21
+ ```python
22
+ import numpy as np
23
+ from k3_node.layers import MessageNorm
24
+
25
+ x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
26
+
27
+ msg = np.random.rand(10, 8).astype("float32") # aggregated messages for each node
28
+ layer = MessageNorm(learn_scale=True)
29
+ out = layer(x, msg) # rescales messages to the norm of x
30
+ print(tuple(out.shape)) # (10, 8)
31
+ ```
32
+ """
33
+ def __init__(self, learn_scale: bool = False, **kwargs):
34
+ super().__init__(**kwargs)
35
+ self.learn_scale = learn_scale
36
+ self.scale = self.add_weight(
37
+ shape=(1,),
38
+ initializer="ones",
39
+ trainable=learn_scale,
40
+ name="scale",
41
+ )
42
+
43
+ def reset_parameters(self):
44
+ self.scale.assign(ops.ones(self.scale.shape, dtype=self.scale.dtype))
45
+
46
+ def call(self, x, msg=None, p=2.0):
47
+ if msg is None and isinstance(x, (tuple, list)):
48
+ x, msg = x
49
+
50
+ msg_norm = ops.maximum(ops.norm(msg, ord=p, axis=-1, keepdims=True), 1e-12)
51
+ msg_normalized = msg / msg_norm
52
+ x_norm = ops.norm(x, ord=p, axis=-1, keepdims=True)
53
+ return msg_normalized * x_norm * self.scale
54
+
55
+ def compute_output_shape(self, input_shape):
56
+ if isinstance(input_shape, (tuple, list)):
57
+ return input_shape[0]
58
+ return input_shape
@@ -0,0 +1,94 @@
1
+ from keras import layers, ops
2
+ from k3_node.ops.segment import segment_sum
3
+
4
+
5
+ class PairNorm(layers.Layer):
6
+ r"""Applies pair normalization over node features as described in the
7
+ `"PairNorm: Tackling Oversmoothing in GNNs"
8
+ <https://arxiv.org/abs/1909.12223>`_ paper.
9
+
10
+ .. math::
11
+ \mathbf{x}_i^c &= \mathbf{x}_i - \frac{1}{n}
12
+ \sum_{i=1}^n \mathbf{x}_i \\
13
+
14
+ \mathbf{x}_i^{\prime} &= s \cdot
15
+ \frac{\mathbf{x}_i^c}{\sqrt{\frac{1}{n} \sum_{i=1}^n
16
+ {\| \mathbf{x}_i^c \|}^2_2}}
17
+
18
+ Args:
19
+ scale (float, optional): Scaling factor :math:`s` of normalization.
20
+ (default: :obj:`1.0`)
21
+ scale_individually (bool, optional): If set to :obj:`True`, will
22
+ compute the scaling step as :math:`\mathbf{x}^{\prime}_i = s \cdot
23
+ \frac{\mathbf{x}_i^c}{{\| \mathbf{x}_i^c \|}_2}`.
24
+ (default: :obj:`False`)
25
+ eps (float, optional): A value added to the denominator for numerical
26
+ stability. (default: :obj:`1e-5`)
27
+
28
+ Example:
29
+ ```python
30
+ import numpy as np
31
+ from k3_node.layers import PairNorm
32
+
33
+ x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
34
+ batch = np.repeat([0, 1], 5) # nodes 0-4 belong to graph 0, nodes 5-9 to graph 1
35
+
36
+ layer = PairNorm()
37
+ out = layer(x, batch) # normalizes each graph separately
38
+ print(tuple(out.shape)) # (10, 8)
39
+ ```
40
+ """
41
+ def __init__(self, scale: float = 1.0, scale_individually: bool = False,
42
+ eps: float = 1e-5, **kwargs):
43
+ super().__init__(**kwargs)
44
+ self.scale = scale
45
+ self.scale_individually = scale_individually
46
+ self.eps = eps
47
+
48
+ def call(self, x, batch=None, batch_size=None):
49
+ if batch is None and isinstance(x, (tuple, list)):
50
+ if len(x) == 2:
51
+ x, batch = x
52
+ elif len(x) == 3:
53
+ x, batch, batch_size = x
54
+
55
+ scale = self.scale
56
+
57
+ if batch is None:
58
+ x = x - ops.mean(x, axis=0, keepdims=True)
59
+
60
+ if not self.scale_individually:
61
+ mean_sq = ops.mean(ops.sum(ops.power(x, 2), axis=-1))
62
+ return scale * x / ops.sqrt(self.eps + mean_sq)
63
+ else:
64
+ norm = ops.sqrt(ops.sum(ops.power(x, 2), axis=-1, keepdims=True))
65
+ return scale * x / (self.eps + norm)
66
+
67
+ if batch_size is not None and not isinstance(batch_size, int):
68
+ try:
69
+ batch_size = int(batch_size)
70
+ except Exception:
71
+ pass
72
+ elif batch_size is None:
73
+ batch_size = ops.cast(ops.max(batch), "int32") + 1
74
+
75
+ batch = ops.cast(batch, "int32")
76
+ ones = ops.ones((ops.shape(x)[0], 1), dtype=x.dtype)
77
+ counts = ops.maximum(segment_sum(ones, batch, num_segments=batch_size), 1.0)
78
+ mean = segment_sum(x, batch, num_segments=batch_size) / counts
79
+ x = x - ops.take(mean, batch, axis=0)
80
+
81
+ if not self.scale_individually:
82
+ sq_sum = ops.sum(ops.power(x, 2), axis=-1, keepdims=True)
83
+ mean_sq = segment_sum(sq_sum, batch, num_segments=batch_size) / counts
84
+ denom = ops.sqrt(self.eps + ops.take(mean_sq, batch, axis=0))
85
+ return scale * x / denom
86
+ else:
87
+ norm = ops.sqrt(ops.sum(ops.power(x, 2), axis=-1, keepdims=True))
88
+ return scale * x / (self.eps + norm)
89
+
90
+ def compute_output_shape(self, input_shape):
91
+ if isinstance(input_shape, (tuple, list)) and isinstance(input_shape[0], (tuple, list)):
92
+ return input_shape[0]
93
+ return input_shape
94
+