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,96 @@
1
+ import math
2
+
3
+ import keras
4
+ from keras import ops
5
+
6
+ from k3_node.layers.kge.base import KGEModel, margin_ranking_loss, normalize
7
+
8
+
9
+ class TransE(KGEModel):
10
+ r"""The TransE model from the `"Translating Embeddings for Modeling
11
+ Multi-Relational Data" <https://proceedings.neurips.cc/paper/2013/file/
12
+ 1cecc7a77928ca8133fa24680a88d2f9-Paper.pdf>`_ paper.
13
+
14
+ :class:`TransE` models relations as a translation from head to tail
15
+ entities such that
16
+
17
+ .. math::
18
+ \mathbf{e}_h + \mathbf{e}_r \approx \mathbf{e}_t,
19
+
20
+ resulting in the scoring function:
21
+
22
+ .. math::
23
+ d(h, r, t) = - {\| \mathbf{e}_h + \mathbf{e}_r - \mathbf{e}_t \|}_p
24
+
25
+ Args:
26
+ num_nodes (int): The number of nodes/entities in the graph.
27
+ num_relations (int): The number of relations in the graph.
28
+ hidden_channels (int): The hidden embedding size.
29
+ margin (float, optional): The margin of the ranking loss.
30
+ (default: :obj:`1.0`)
31
+ p_norm (float, optional): The order embedding and distance
32
+ normalization. (default: :obj:`1.0`)
33
+ sparse (bool, optional): Kept for API compatibility. (default: :obj:`False`)
34
+
35
+ Example:
36
+ ```python
37
+ import numpy as np
38
+ from k3_node.layers import TransE
39
+
40
+ head = np.random.randint(0, 20, size=(10,)) # 10 (head, relation, tail) triples
41
+ rel = np.random.randint(0, 5, size=(10,))
42
+ tail = np.random.randint(0, 20, size=(10,))
43
+
44
+ model = TransE(num_nodes=20, num_relations=5, hidden_channels=8)
45
+ score = model(head, rel, tail) # plausibility score of every triple
46
+ print(tuple(score.shape)) # (10,)
47
+ loss = model.loss(head, rel, tail) # training loss against randomly corrupted triples
48
+ print(tuple(loss.shape)) # (): a scalar
49
+ ```
50
+ """
51
+ def __init__(
52
+ self,
53
+ num_nodes: int,
54
+ num_relations: int,
55
+ hidden_channels: int,
56
+ margin: float = 1.0,
57
+ p_norm: float = 1.0,
58
+ sparse: bool = False,
59
+ **kwargs,
60
+ ):
61
+ super().__init__(num_nodes, num_relations, hidden_channels, sparse, **kwargs)
62
+
63
+ self.p_norm = p_norm
64
+ self.margin = margin
65
+
66
+ self.reset_parameters()
67
+
68
+ def reset_parameters(self):
69
+ bound = 6.0 / math.sqrt(self.hidden_channels)
70
+ # A new initializer per tensor: a reused unseeded Keras 3 initializer returns the same values on every call.
71
+ uniform = lambda shape: keras.initializers.RandomUniform(-bound, bound)(shape)
72
+ self.node_emb.embeddings.assign(uniform(ops.shape(self.node_emb.embeddings)))
73
+ self.rel_emb.embeddings.assign(uniform(ops.shape(self.rel_emb.embeddings)))
74
+ self.rel_emb.embeddings.assign(normalize(self.rel_emb.embeddings, p=self.p_norm, axis=-1))
75
+
76
+ def call(self, head_index, rel_type, tail_index):
77
+ head_index = ops.cast(head_index, "int32")
78
+ rel_type = ops.cast(rel_type, "int32")
79
+ tail_index = ops.cast(tail_index, "int32")
80
+
81
+ head = self.node_emb(head_index)
82
+ rel = self.rel_emb(rel_type)
83
+ tail = self.node_emb(tail_index)
84
+
85
+ head = normalize(head, p=self.p_norm, axis=-1)
86
+ tail = normalize(tail, p=self.p_norm, axis=-1)
87
+
88
+ # Calculate *negative* TransE norm:
89
+ diff = (head + rel) - tail
90
+ return -ops.power(ops.sum(ops.power(ops.abs(diff), self.p_norm), axis=-1), 1.0 / self.p_norm)
91
+
92
+ def loss(self, head_index, rel_type, tail_index):
93
+ pos_score = self(head_index, rel_type, tail_index)
94
+ neg_score = self(*self.random_sample(head_index, rel_type, tail_index))
95
+
96
+ return margin_ranking_loss(pos_score, neg_score, margin=self.margin)
@@ -0,0 +1,23 @@
1
+ from .batch_norm import BatchNorm, HeteroBatchNorm
2
+ from .diff_group_norm import DiffGroupNorm
3
+ from .graph_norm import GraphNorm
4
+ from .graph_size_norm import GraphSizeNorm
5
+ from .instance_norm import InstanceNorm
6
+ from .layer_norm import HeteroLayerNorm, LayerNorm
7
+ from .mean_subtraction_norm import MeanSubtractionNorm
8
+ from .msg_norm import MessageNorm
9
+ from .pair_norm import PairNorm
10
+
11
+ __all__ = [
12
+ "BatchNorm",
13
+ "HeteroBatchNorm",
14
+ "InstanceNorm",
15
+ "LayerNorm",
16
+ "HeteroLayerNorm",
17
+ "GraphNorm",
18
+ "GraphSizeNorm",
19
+ "PairNorm",
20
+ "MeanSubtractionNorm",
21
+ "MessageNorm",
22
+ "DiffGroupNorm",
23
+ ]
@@ -0,0 +1,328 @@
1
+ from typing import Optional
2
+ from keras import layers, ops
3
+ from k3_node.ops.segment import segment_sum
4
+
5
+
6
+ class BatchNorm(layers.Layer):
7
+ r"""Applies batch normalization over a batch of features as described in
8
+ the `"Batch Normalization: Accelerating Deep Network Training by
9
+ Reducing Internal Covariate Shift" <https://arxiv.org/abs/1502.03167>`_
10
+ paper.
11
+
12
+ .. math::
13
+ \mathbf{x}^{\prime}_i = \frac{\mathbf{x} -
14
+ \textrm{E}[\mathbf{x}]}{\sqrt{\textrm{Var}[\mathbf{x}] + \epsilon}}
15
+ \odot \gamma + \beta
16
+
17
+ Args:
18
+ in_channels (int): Size of each input sample.
19
+ eps (float, optional): A value added to the denominator for numerical
20
+ stability. (default: :obj:`1e-5`)
21
+ momentum (float, optional): The value used for the running mean and
22
+ running variance computation. (default: :obj:`0.1`)
23
+ affine (bool, optional): If set to :obj:`True`, this module has
24
+ learnable affine parameters :math:`\gamma` and :math:`\beta`.
25
+ (default: :obj:`True`)
26
+ track_running_stats (bool, optional): If set to :obj:`True`, this
27
+ module tracks the running mean and variance, and when set to
28
+ :obj:`False`, this module does not track such statistics and always
29
+ uses batch statistics in both training and eval modes.
30
+ (default: :obj:`True`)
31
+ allow_single_element (bool, optional): If set to :obj:`True`, batches
32
+ with only a single element will work as during in evaluation.
33
+ That is the running mean and variance will be used.
34
+ Requires :obj:`track_running_stats=True`. (default: :obj:`False`)
35
+
36
+ Example:
37
+ ```python
38
+ import numpy as np
39
+ from k3_node.layers import BatchNorm
40
+
41
+ x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
42
+
43
+ layer = BatchNorm(in_channels=8)
44
+ out = layer(x, training=True) # uses batch statistics while training
45
+ print(tuple(out.shape)) # (10, 8)
46
+ ```
47
+ """
48
+ def __init__(
49
+ self,
50
+ in_channels: int,
51
+ eps: float = 1e-5,
52
+ momentum: Optional[float] = 0.1,
53
+ affine: bool = True,
54
+ track_running_stats: bool = True,
55
+ allow_single_element: bool = False,
56
+ **kwargs
57
+ ):
58
+ super().__init__(**kwargs)
59
+ if allow_single_element and not track_running_stats:
60
+ raise ValueError("'allow_single_element' requires "
61
+ "'track_running_stats' to be set to `True`")
62
+
63
+ self.in_channels = in_channels
64
+ self.eps = eps
65
+ self.momentum = momentum
66
+ self.affine = affine
67
+ self.track_running_stats = track_running_stats
68
+ self.allow_single_element = allow_single_element
69
+
70
+ if affine:
71
+ self.weight = self.add_weight(
72
+ shape=(in_channels,),
73
+ initializer="ones",
74
+ trainable=True,
75
+ name="weight",
76
+ )
77
+ self.bias = self.add_weight(
78
+ shape=(in_channels,),
79
+ initializer="zeros",
80
+ trainable=True,
81
+ name="bias",
82
+ )
83
+ else:
84
+ self.weight = None
85
+ self.bias = None
86
+
87
+ if track_running_stats:
88
+ self.running_mean = self.add_weight(
89
+ shape=(in_channels,),
90
+ initializer="zeros",
91
+ trainable=False,
92
+ name="running_mean",
93
+ )
94
+ self.running_var = self.add_weight(
95
+ shape=(in_channels,),
96
+ initializer="ones",
97
+ trainable=False,
98
+ name="running_var",
99
+ )
100
+ self.num_batches_tracked = self.add_weight(
101
+ shape=(),
102
+ initializer="zeros",
103
+ dtype="int32",
104
+ trainable=False,
105
+ name="num_batches_tracked",
106
+ )
107
+ else:
108
+ self.running_mean = None
109
+ self.running_var = None
110
+ self.num_batches_tracked = None
111
+
112
+ def reset_running_stats(self):
113
+ if self.track_running_stats:
114
+ self.running_mean.assign(ops.zeros(self.running_mean.shape, dtype=self.running_mean.dtype))
115
+ self.running_var.assign(ops.ones(self.running_var.shape, dtype=self.running_var.dtype))
116
+ self.num_batches_tracked.assign(ops.cast(0, "int32"))
117
+
118
+ def reset_parameters(self):
119
+ self.reset_running_stats()
120
+ if self.affine:
121
+ self.weight.assign(ops.ones(self.weight.shape, dtype=self.weight.dtype))
122
+ self.bias.assign(ops.zeros(self.bias.shape, dtype=self.bias.dtype))
123
+
124
+ def call(self, x, training=None):
125
+ num_samples_static = x.shape[0]
126
+ # Keras semantics: `training=None` means inference; fit() passes training=True via the call context.
127
+ is_training = bool(training) if training is not None else False
128
+
129
+ if is_training:
130
+ if num_samples_static is not None and num_samples_static <= 1:
131
+ if not self.allow_single_element:
132
+ raise ValueError(f"Expected more than 1 value per channel when training, got input size {ops.shape(x)}")
133
+ # Evaluation behavior with running stats
134
+ mean = self.running_mean
135
+ var = self.running_var
136
+ else:
137
+ mean = ops.mean(x, axis=0)
138
+ var = ops.var(x, axis=0)
139
+
140
+ if self.track_running_stats:
141
+ n = ops.cast(ops.shape(x)[0], dtype=x.dtype)
142
+ unbiased_var = var * n / ops.maximum(n - 1.0, 1.0)
143
+ if self.momentum is None:
144
+ count = ops.cast(self.num_batches_tracked + 1, dtype=x.dtype)
145
+ m = 1.0 / count
146
+ else:
147
+ m = self.momentum
148
+
149
+ new_running_mean = (1.0 - m) * self.running_mean + m * mean
150
+ new_running_var = (1.0 - m) * self.running_var + m * unbiased_var
151
+ self.running_mean.assign(new_running_mean)
152
+ self.running_var.assign(new_running_var)
153
+ self.num_batches_tracked.assign(self.num_batches_tracked + 1)
154
+ else:
155
+ if self.track_running_stats:
156
+ mean = self.running_mean
157
+ var = self.running_var
158
+ else:
159
+ mean = ops.mean(x, axis=0)
160
+ var = ops.var(x, axis=0)
161
+
162
+ out = (x - mean) / ops.sqrt(var + self.eps)
163
+
164
+ if self.affine:
165
+ out = out * self.weight + self.bias
166
+
167
+ return out
168
+
169
+ def compute_output_shape(self, input_shape):
170
+ return input_shape
171
+
172
+
173
+ class HeteroBatchNorm(layers.Layer):
174
+ r"""Applies batch normalization over a batch of heterogeneous features as
175
+ described in the `"Batch Normalization: Accelerating Deep Network Training
176
+ by Reducing Internal Covariate Shift" <https://arxiv.org/abs/1502.03167>`_
177
+ paper.
178
+ Compared to :class:`BatchNorm`, :class:`HeteroBatchNorm` applies
179
+ normalization individually for each node or edge type.
180
+
181
+ Args:
182
+ in_channels (int): Size of each input sample.
183
+ num_types (int): The number of types.
184
+ eps (float, optional): A value added to the denominator for numerical
185
+ stability. (default: :obj:`1e-5`)
186
+ momentum (float, optional): The value used for the running mean and
187
+ running variance computation. (default: :obj:`0.1`)
188
+ affine (bool, optional): If set to :obj:`True`, this module has
189
+ learnable affine parameters :math:`\gamma` and :math:`\beta`.
190
+ (default: :obj:`True`)
191
+ track_running_stats (bool, optional): If set to :obj:`True`, this
192
+ module tracks the running mean and variance, and when set to
193
+ :obj:`False`, this module does not track such statistics and always
194
+ uses batch statistics in both training and eval modes.
195
+ (default: :obj:`True`)
196
+
197
+ Example:
198
+ ```python
199
+ import numpy as np
200
+ from k3_node.layers import HeteroBatchNorm
201
+
202
+ x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
203
+ node_type = np.random.randint(0, 3, size=(10,)) # type of each node
204
+
205
+ layer = HeteroBatchNorm(in_channels=8, num_types=3)
206
+ out = layer(x, node_type, training=True) # separate statistics per node type
207
+ print(tuple(out.shape)) # (10, 8)
208
+ ```
209
+ """
210
+ def __init__(
211
+ self,
212
+ in_channels: int,
213
+ num_types: int,
214
+ eps: float = 1e-5,
215
+ momentum: Optional[float] = 0.1,
216
+ affine: bool = True,
217
+ track_running_stats: bool = True,
218
+ **kwargs
219
+ ):
220
+ super().__init__(**kwargs)
221
+ self.in_channels = in_channels
222
+ self.num_types = num_types
223
+ self.eps = eps
224
+ self.momentum = momentum
225
+ self.affine = affine
226
+ self.track_running_stats = track_running_stats
227
+
228
+ if affine:
229
+ self.weight = self.add_weight(
230
+ shape=(num_types, in_channels),
231
+ initializer="ones",
232
+ trainable=True,
233
+ name="weight",
234
+ )
235
+ self.bias = self.add_weight(
236
+ shape=(num_types, in_channels),
237
+ initializer="zeros",
238
+ trainable=True,
239
+ name="bias",
240
+ )
241
+ else:
242
+ self.weight = None
243
+ self.bias = None
244
+
245
+ if track_running_stats:
246
+ self.running_mean = self.add_weight(
247
+ shape=(num_types, in_channels),
248
+ initializer="zeros",
249
+ trainable=False,
250
+ name="running_mean",
251
+ )
252
+ self.running_var = self.add_weight(
253
+ shape=(num_types, in_channels),
254
+ initializer="ones",
255
+ trainable=False,
256
+ name="running_var",
257
+ )
258
+ self.num_batches_tracked = self.add_weight(
259
+ shape=(),
260
+ initializer="zeros",
261
+ dtype="int32",
262
+ trainable=False,
263
+ name="num_batches_tracked",
264
+ )
265
+ else:
266
+ self.running_mean = None
267
+ self.running_var = None
268
+ self.num_batches_tracked = None
269
+
270
+ def reset_running_stats(self):
271
+ if self.track_running_stats:
272
+ self.running_mean.assign(ops.zeros(self.running_mean.shape, dtype=self.running_mean.dtype))
273
+ self.running_var.assign(ops.ones(self.running_var.shape, dtype=self.running_var.dtype))
274
+ self.num_batches_tracked.assign(ops.cast(0, "int32"))
275
+
276
+ def reset_parameters(self):
277
+ self.reset_running_stats()
278
+ if self.affine:
279
+ self.weight.assign(ops.ones(self.weight.shape, dtype=self.weight.dtype))
280
+ self.bias.assign(ops.zeros(self.bias.shape, dtype=self.bias.dtype))
281
+
282
+ def call(self, x, type_vec=None, training=None):
283
+ if type_vec is None and isinstance(x, (tuple, list)):
284
+ x, type_vec = x
285
+
286
+ # Keras semantics: `training=None` means inference; fit() passes training=True via the call context.
287
+ is_training = bool(training) if training is not None else False
288
+ type_vec = ops.cast(type_vec, "int32")
289
+
290
+ if not is_training and self.track_running_stats:
291
+ mean = self.running_mean
292
+ var = self.running_var
293
+ else:
294
+ ones = ops.ones((ops.shape(x)[0], 1), dtype=x.dtype)
295
+ counts = ops.maximum(segment_sum(ones, type_vec, num_segments=self.num_types), 1.0)
296
+ mean = segment_sum(x, type_vec, num_segments=self.num_types) / counts
297
+ x_c = x - ops.take(mean, type_vec, axis=0)
298
+ var = segment_sum(ops.power(x_c, 2), type_vec, num_segments=self.num_types) / counts
299
+
300
+ if is_training and self.track_running_stats:
301
+ if self.momentum is None:
302
+ count = ops.cast(self.num_batches_tracked + 1, dtype=x.dtype)
303
+ exp_avg_factor = 1.0 / count
304
+ else:
305
+ exp_avg_factor = self.momentum
306
+
307
+ new_running_mean = (1.0 - exp_avg_factor) * self.running_mean + exp_avg_factor * mean
308
+ new_running_var = (1.0 - exp_avg_factor) * self.running_var + exp_avg_factor * var
309
+ self.running_mean.assign(new_running_mean)
310
+ self.running_var.assign(new_running_var)
311
+ self.num_batches_tracked.assign(self.num_batches_tracked + 1)
312
+
313
+ std = ops.sqrt(var + self.eps)
314
+ mean_taken = ops.take(mean, type_vec, axis=0)
315
+ std_taken = ops.take(std, type_vec, axis=0)
316
+ out = (x - mean_taken) / std_taken
317
+
318
+ if self.affine:
319
+ w_taken = ops.take(self.weight, type_vec, axis=0)
320
+ b_taken = ops.take(self.bias, type_vec, axis=0)
321
+ out = out * w_taken + b_taken
322
+
323
+ return out
324
+
325
+ def compute_output_shape(self, input_shape):
326
+ if isinstance(input_shape, (tuple, list)):
327
+ return input_shape[0]
328
+ return input_shape
@@ -0,0 +1,141 @@
1
+ import numpy as np
2
+ from scipy.spatial.distance import cdist
3
+ from keras import initializers, layers, ops
4
+
5
+ from .batch_norm import BatchNorm
6
+
7
+
8
+ class DiffGroupNorm(layers.Layer):
9
+ r"""The differentiable group normalization layer from the `"Towards Deeper
10
+ Graph Neural Networks with Differentiable Group Normalization"
11
+ <https://arxiv.org/abs/2006.06972>`_ paper, which normalizes node features
12
+ group-wise via a learnable soft cluster assignment.
13
+
14
+ .. math::
15
+
16
+ \mathbf{S} = \text{softmax} (\mathbf{X} \mathbf{W})
17
+
18
+ where :math:`\mathbf{W} \in \mathbb{R}^{F \times G}` denotes a trainable
19
+ weight matrix mapping each node into one of :math:`G` clusters.
20
+ Normalization is then performed group-wise via:
21
+
22
+ .. math::
23
+
24
+ \mathbf{X}^{\prime} = \mathbf{X} + \lambda \sum_{i = 1}^G
25
+ \text{BatchNorm}(\mathbf{S}[:, i] \odot \mathbf{X})
26
+
27
+ Args:
28
+ in_channels (int): Size of each input sample :math:`F`.
29
+ groups (int): The number of groups :math:`G`.
30
+ lamda (float, optional): The balancing factor :math:`\lambda` between
31
+ input embeddings and normalized embeddings. (default: :obj:`0.01`)
32
+ eps (float, optional): A value added to the denominator for numerical
33
+ stability. (default: :obj:`1e-5`)
34
+ momentum (float, optional): The value used for the running mean and
35
+ running variance computation. (default: :obj:`0.1`)
36
+ affine (bool, optional): If set to :obj:`True`, this module has
37
+ learnable affine parameters :math:`\gamma` and :math:`\beta`.
38
+ (default: :obj:`True`)
39
+ track_running_stats (bool, optional): If set to :obj:`True`, this
40
+ module tracks the running mean and variance, and when set to
41
+ :obj:`False`, this module does not track such statistics and always
42
+ uses batch statistics in both training and eval modes.
43
+ (default: :obj:`True`)
44
+
45
+ Example:
46
+ ```python
47
+ import numpy as np
48
+ from k3_node.layers import DiffGroupNorm
49
+
50
+ x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
51
+
52
+ layer = DiffGroupNorm(in_channels=8, groups=2)
53
+ out = layer(x, training=True) # uses batch statistics while training
54
+ print(tuple(out.shape)) # (10, 8)
55
+ ```
56
+ """
57
+ def __init__(
58
+ self,
59
+ in_channels: int,
60
+ groups: int,
61
+ lamda: float = 0.01,
62
+ eps: float = 1e-5,
63
+ momentum: float = 0.1,
64
+ affine: bool = True,
65
+ track_running_stats: bool = True,
66
+ **kwargs
67
+ ):
68
+ super().__init__(**kwargs)
69
+ self.in_channels = in_channels
70
+ self.groups = groups
71
+ self.lamda = lamda
72
+ self.eps = eps
73
+ self.momentum = momentum
74
+ self.affine = affine
75
+ self.track_running_stats = track_running_stats
76
+
77
+ self.lin_weight = self.add_weight(
78
+ shape=(in_channels, groups),
79
+ initializer="glorot_uniform",
80
+ trainable=True,
81
+ name="lin_weight",
82
+ )
83
+ self.norm = BatchNorm(
84
+ groups * in_channels,
85
+ eps=eps,
86
+ momentum=momentum,
87
+ affine=affine,
88
+ track_running_stats=track_running_stats,
89
+ name="norm",
90
+ )
91
+
92
+ def reset_parameters(self):
93
+ self.lin_weight.assign(
94
+ initializers.GlorotUniform()(shape=self.lin_weight.shape, dtype=self.lin_weight.dtype)
95
+ )
96
+ self.norm.reset_parameters()
97
+
98
+ def call(self, x, training=None):
99
+ F, G = self.in_channels, self.groups
100
+
101
+ s = ops.softmax(ops.matmul(x, self.lin_weight), axis=-1) # [N, G]
102
+ out = ops.expand_dims(s, axis=-1) * ops.expand_dims(x, axis=-2) # [N, G, F]
103
+ out_flat = ops.reshape(out, (-1, G * F))
104
+ out_norm = self.norm(out_flat, training=training)
105
+ out = ops.sum(ops.reshape(out_norm, (-1, G, F)), axis=-2) # [N, F]
106
+
107
+ return x + self.lamda * out
108
+
109
+ @staticmethod
110
+ def group_distance_ratio(x, y, eps: float = 1e-5) -> float:
111
+ r"""Measures the ratio of inter-group distance over intra-group
112
+ distance.
113
+ """
114
+ x_np = ops.convert_to_numpy(x)
115
+ y_np = ops.convert_to_numpy(y).astype(np.int64)
116
+
117
+ num_classes = int(y_np.max()) + 1
118
+
119
+ numerator = 0.0
120
+ for i in range(num_classes):
121
+ mask = (y_np == i)
122
+ if not np.any(mask) or np.all(mask):
123
+ continue
124
+ dist = cdist(x_np[mask], x_np[~mask])
125
+ numerator += (1.0 / dist.size) * float(dist.sum())
126
+ numerator *= 1.0 / ((num_classes - 1) ** 2)
127
+
128
+ denominator = 0.0
129
+ for i in range(num_classes):
130
+ mask = (y_np == i)
131
+ if not np.any(mask):
132
+ continue
133
+ dist = cdist(x_np[mask], x_np[mask])
134
+ denominator += (1.0 / dist.size) * float(dist.sum())
135
+ denominator *= 1.0 / num_classes
136
+
137
+ return float(numerator / (denominator + eps))
138
+
139
+ def compute_output_shape(self, input_shape):
140
+ return input_shape
141
+
@@ -0,0 +1,105 @@
1
+ from keras import layers, ops
2
+ from k3_node.layers.conv.utils import is_tracing
3
+ from k3_node.ops.segment import segment_sum
4
+
5
+
6
+ class GraphNorm(layers.Layer):
7
+ r"""Applies graph normalization over individual graphs as described in the
8
+ `"GraphNorm: A Principled Approach to Accelerating Graph Neural Network
9
+ Training" <https://arxiv.org/abs/2009.03294>`_ paper.
10
+
11
+ .. math::
12
+ \mathbf{x}^{\prime}_i = \frac{\mathbf{x} - \alpha \odot
13
+ \textrm{E}[\mathbf{x}]}
14
+ {\sqrt{\textrm{Var}[\mathbf{x} - \alpha \odot \textrm{E}[\mathbf{x}]]
15
+ + \epsilon}} \odot \gamma + \beta
16
+
17
+ where :math:`\alpha` denotes parameters that learn how much information
18
+ to keep in the mean.
19
+
20
+ Args:
21
+ in_channels (int): Size of each input sample.
22
+ eps (float, optional): A value added to the denominator for numerical
23
+ stability. (default: :obj:`1e-5`)
24
+
25
+ Example:
26
+ ```python
27
+ import numpy as np
28
+ from k3_node.layers import GraphNorm
29
+
30
+ x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
31
+ batch = np.repeat([0, 1], 5) # nodes 0-4 belong to graph 0, nodes 5-9 to graph 1
32
+
33
+ layer = GraphNorm(in_channels=8)
34
+ out = layer(x, batch) # normalizes each graph separately
35
+ print(tuple(out.shape)) # (10, 8)
36
+ ```
37
+ """
38
+ def __init__(self, in_channels: int, eps: float = 1e-5, **kwargs):
39
+ super().__init__(**kwargs)
40
+ self.in_channels = in_channels
41
+ self.eps = eps
42
+
43
+ self.weight = self.add_weight(
44
+ shape=(in_channels,),
45
+ initializer="ones",
46
+ trainable=True,
47
+ name="weight",
48
+ )
49
+ self.bias = self.add_weight(
50
+ shape=(in_channels,),
51
+ initializer="zeros",
52
+ trainable=True,
53
+ name="bias",
54
+ )
55
+ self.mean_scale = self.add_weight(
56
+ shape=(in_channels,),
57
+ initializer="ones",
58
+ trainable=True,
59
+ name="mean_scale",
60
+ )
61
+
62
+ def reset_parameters(self):
63
+ self.weight.assign(ops.ones(self.weight.shape, dtype=self.weight.dtype))
64
+ self.bias.assign(ops.zeros(self.bias.shape, dtype=self.bias.dtype))
65
+ self.mean_scale.assign(ops.ones(self.mean_scale.shape, dtype=self.mean_scale.dtype))
66
+
67
+ def call(self, x, batch=None, batch_size=None):
68
+ if batch is None and isinstance(x, (tuple, list)):
69
+ if len(x) == 2:
70
+ x, batch = x
71
+ elif len(x) == 3:
72
+ x, batch, batch_size = x
73
+
74
+ if batch is None:
75
+ mean = ops.mean(x, axis=0, keepdims=True)
76
+ out = x - mean * self.mean_scale
77
+ var = ops.mean(ops.power(out, 2), axis=0, keepdims=True)
78
+ std = ops.sqrt(var + self.eps)
79
+ return self.weight * out / std + self.bias
80
+
81
+ if batch_size is None:
82
+ if not is_tracing(batch):
83
+ try:
84
+ batch_size = int(ops.max(batch)) + 1
85
+ except Exception:
86
+ batch_size = None
87
+ elif not isinstance(batch_size, int):
88
+ try:
89
+ batch_size = int(batch_size)
90
+ except Exception:
91
+ pass
92
+
93
+ batch = ops.cast(batch, "int32")
94
+ ones = ops.ones((ops.shape(x)[0], 1), dtype=x.dtype)
95
+ counts = ops.maximum(segment_sum(ones, batch, num_segments=batch_size), 1.0)
96
+ mean = segment_sum(x, batch, num_segments=batch_size) / counts
97
+ out = x - ops.take(mean, batch, axis=0) * self.mean_scale
98
+ var = segment_sum(ops.power(out, 2), batch, num_segments=batch_size) / counts
99
+ std = ops.take(ops.sqrt(var + self.eps), batch, axis=0)
100
+ return self.weight * out / std + self.bias
101
+
102
+ def compute_output_shape(self, input_shape):
103
+ if isinstance(input_shape, (tuple, list)) and isinstance(input_shape[0], (tuple, list)):
104
+ return input_shape[0]
105
+ return input_shape