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,443 @@
1
+ import copy
2
+ import inspect
3
+ from typing import Any, Callable, Dict, List, Optional, Tuple, Union
4
+
5
+ import keras
6
+ from keras import ops
7
+
8
+ from k3_node.layers.conv import (
9
+ EdgeConv,
10
+ GATConv,
11
+ GATv2Conv,
12
+ GCNConv,
13
+ GINConv,
14
+ MessagePassing,
15
+ PNAConv,
16
+ SAGEConv,
17
+ )
18
+ from k3_node.models.mlp import MLP, _normalization_resolver
19
+ from k3_node.models.jumping_knowledge import JumpingKnowledge
20
+ from k3_node.hub.hub_mixin import K3NodeHubMixin
21
+
22
+
23
+ class BasicGNN(K3NodeHubMixin, keras.Model):
24
+ r"""An abstract base class for implementing basic GNN models.
25
+
26
+ Args:
27
+ in_channels (int or tuple): Size of each input sample.
28
+ hidden_channels (int): Size of each hidden sample.
29
+ num_layers (int): Number of message passing layers.
30
+ out_channels (int, optional): If not set to :obj:`None`, will apply a
31
+ final linear transformation to convert hidden node embeddings to
32
+ output size :obj:`out_channels`. (default: :obj:`None`)
33
+ dropout (float, optional): Dropout probability. (default: :obj:`0.`)
34
+ act (str or Callable, optional): The non-linear activation function to
35
+ use. (default: :obj:`"relu"`)
36
+ act_first (bool, optional): If set to :obj:`True`, activation is
37
+ applied before normalization. (default: :obj:`False`)
38
+ act_kwargs (Dict[str, Any], optional): Arguments passed to the
39
+ respective activation function defined by :obj:`act`.
40
+ (default: :obj:`None`)
41
+ norm (str or Callable, optional): The normalization function to
42
+ use. (default: :obj:`None`)
43
+ norm_kwargs (Dict[str, Any], optional): Arguments passed to the
44
+ respective normalization function defined by :obj:`norm`.
45
+ (default: :obj:`None`)
46
+ jk (str, optional): The Jumping Knowledge mode. If specified, the model
47
+ will additionally apply a final linear transformation to transform
48
+ node embeddings to the expected output feature dimensionality.
49
+ (:obj:`None`, :obj:`"last"`, :obj:`"cat"`, :obj:`"max"`,
50
+ :obj:`"lstm"`). (default: :obj:`None`)
51
+ **kwargs (optional): Additional arguments of the underlying
52
+ :class:`torch_geometric.nn.conv.MessagePassing` layers.
53
+
54
+ Call arguments: ``x``, ``edge_index`` and optionally ``edge_weight``, ``edge_attr`` and
55
+ ``batch``. With mini-batches from :class:`~k3_node.loader.NeighborLoader`, also pass
56
+ ``num_sampled_nodes_per_hop`` and ``num_sampled_edges_per_hop`` (the batch's
57
+ ``num_sampled_nodes`` / ``num_sampled_edges``) to skip the nodes and edges that no longer
58
+ affect the seed nodes after each layer, as in PyG. This trimming needs concrete sizes, so
59
+ compile the model with ``run_eagerly=True`` when using it.
60
+ """
61
+ supports_edge_weight: bool = False
62
+ supports_edge_attr: bool = False
63
+ supports_norm_batch: bool = False
64
+
65
+ def __init__(
66
+ self,
67
+ in_channels: int,
68
+ hidden_channels: int,
69
+ num_layers: int,
70
+ out_channels: Optional[int] = None,
71
+ dropout: float = 0.0,
72
+ act: Union[str, Callable, None] = "relu",
73
+ act_first: bool = False,
74
+ act_kwargs: Optional[Dict[str, Any]] = None,
75
+ norm: Union[str, Callable, None] = None,
76
+ norm_kwargs: Optional[Dict[str, Any]] = None,
77
+ jk: Optional[str] = None,
78
+ **kwargs,
79
+ ):
80
+ super().__init__()
81
+
82
+ self.in_channels = in_channels
83
+ self.hidden_channels = hidden_channels
84
+ self.num_layers = num_layers
85
+ dropout = float(dropout) if dropout is not None else 0.0
86
+ self.dropout_p = dropout
87
+ self.dropout = keras.layers.Dropout(rate=dropout) if dropout > 0 else None
88
+
89
+ if isinstance(act, str):
90
+ self.act = keras.activations.get(act)
91
+ else:
92
+ self.act = act
93
+
94
+ self.jk_mode = jk
95
+ self.act_first = act_first
96
+ self.norm_query = norm
97
+ self.norm_kwargs = norm_kwargs or {}
98
+
99
+ if out_channels is not None:
100
+ self.out_channels = out_channels
101
+ else:
102
+ self.out_channels = hidden_channels
103
+
104
+ self.convs = []
105
+ curr_in = in_channels
106
+ if num_layers > 1:
107
+ self.convs.append(self.init_conv(curr_in, hidden_channels, **kwargs))
108
+ if isinstance(curr_in, (tuple, list)):
109
+ curr_in = (hidden_channels, hidden_channels)
110
+ else:
111
+ curr_in = hidden_channels
112
+
113
+ for _ in range(num_layers - 2):
114
+ self.convs.append(self.init_conv(curr_in, hidden_channels, **kwargs))
115
+ if isinstance(curr_in, (tuple, list)):
116
+ curr_in = (hidden_channels, hidden_channels)
117
+ else:
118
+ curr_in = hidden_channels
119
+
120
+ if out_channels is not None and jk is None:
121
+ self._is_conv_to_out = True
122
+ self.convs.append(self.init_conv(curr_in, out_channels, **kwargs))
123
+ else:
124
+ self.convs.append(self.init_conv(curr_in, hidden_channels, **kwargs))
125
+
126
+ self.norms = []
127
+ self.supports_norm_batch = False
128
+
129
+ for _ in range(num_layers - 1):
130
+ if norm is not None:
131
+ norm_layer = _normalization_resolver(norm, hidden_channels, **self.norm_kwargs)
132
+ self.norms.append(norm_layer)
133
+ if hasattr(norm_layer, "call"):
134
+ sig = inspect.signature(norm_layer.call).parameters
135
+ self.supports_norm_batch = "batch" in sig
136
+ else:
137
+ self.norms.append(None)
138
+
139
+ if jk is not None:
140
+ if norm is not None:
141
+ self.norms.append(_normalization_resolver(norm, hidden_channels, **self.norm_kwargs))
142
+ else:
143
+ self.norms.append(None)
144
+ else:
145
+ self.norms.append(None)
146
+
147
+ if jk is not None and jk != "last":
148
+ self.jk = JumpingKnowledge(jk, hidden_channels, num_layers)
149
+
150
+ if jk is not None:
151
+ if jk == "cat":
152
+ jk_in = num_layers * hidden_channels
153
+ else:
154
+ jk_in = hidden_channels
155
+ self.lin = keras.layers.Dense(self.out_channels)
156
+
157
+ def init_conv(self, in_channels: Union[int, Tuple[int, int]],
158
+ out_channels: int, **kwargs) -> MessagePassing:
159
+ raise NotImplementedError
160
+
161
+ def build(self, input_shape=None):
162
+ self.built = True
163
+
164
+ def reset_parameters(self):
165
+ r"""Resets all learnable parameters of the module."""
166
+ for conv in self.convs:
167
+ if hasattr(conv, "reset_parameters"):
168
+ conv.reset_parameters()
169
+ for norm in self.norms:
170
+ if norm is not None and hasattr(norm, "reset_parameters"):
171
+ norm.reset_parameters()
172
+ if hasattr(self, "jk") and hasattr(self.jk, "reset_parameters"):
173
+ self.jk.reset_parameters()
174
+ if hasattr(self, "lin") and hasattr(self.lin, "reset_parameters"):
175
+ self.lin.reset_parameters()
176
+
177
+ def call(
178
+ self,
179
+ x,
180
+ edge_index=None,
181
+ edge_weight=None,
182
+ edge_attr=None,
183
+ batch=None,
184
+ batch_size=None,
185
+ num_sampled_nodes_per_hop=None,
186
+ num_sampled_edges_per_hop=None,
187
+ training=None,
188
+ ):
189
+ if hasattr(x, "edge_index") and edge_index is None:
190
+ edge_index = getattr(x, "edge_index", None)
191
+ edge_weight = getattr(x, "edge_weight", None) if edge_weight is None else edge_weight
192
+ edge_attr = getattr(x, "edge_attr", None) if edge_attr is None else edge_attr
193
+ batch = getattr(x, "batch", None) if batch is None else batch
194
+ x = x.x
195
+ elif isinstance(x, (tuple, list)) and edge_index is None:
196
+ if len(x) > 1:
197
+ edge_index = x[1]
198
+ if len(x) > 2:
199
+ edge_attr = x[2]
200
+ x = x[0]
201
+
202
+ trim = num_sampled_nodes_per_hop is not None and num_sampled_edges_per_hop is not None
203
+ if trim:
204
+ from k3_node.ops.host import to_numpy
205
+
206
+ nodes_per_hop = [int(n) for n in to_numpy(num_sampled_nodes_per_hop)]
207
+ edges_per_hop = [int(n) for n in to_numpy(num_sampled_edges_per_hop)]
208
+
209
+ xs: List = []
210
+ # `training` is forwarded explicitly: Keras does not propagate it to nested layers on JAX.
211
+ for i, (conv, norm) in enumerate(zip(self.convs, self.norms)):
212
+ if trim and i > 0:
213
+ # Hierarchical neighborhood sampling: the nodes and edges of the outermost hop no
214
+ # longer influence the seed nodes, so drop them (as PyG's `trim_to_layer`).
215
+ x = x[: x.shape[0] - nodes_per_hop[-i]]
216
+ num_edges = edge_index.shape[1] - edges_per_hop[-i]
217
+ edge_index = edge_index[:, :num_edges]
218
+ edge_weight = None if edge_weight is None else edge_weight[:num_edges]
219
+ edge_attr = None if edge_attr is None else edge_attr[:num_edges]
220
+ if self.supports_edge_weight and self.supports_edge_attr:
221
+ x = conv(x, edge_index, edge_weight=edge_weight, edge_attr=edge_attr, training=training)
222
+ elif self.supports_edge_weight:
223
+ x = conv(x, edge_index, edge_weight=edge_weight, training=training)
224
+ elif self.supports_edge_attr:
225
+ x = conv(x, edge_index, edge_attr=edge_attr, training=training)
226
+ else:
227
+ x = conv(x, edge_index, training=training)
228
+
229
+ if i < self.num_layers - 1 or self.jk_mode is not None:
230
+ if self.act is not None and self.act_first:
231
+ x = self.act(x)
232
+ if norm is not None:
233
+ if self.supports_norm_batch and batch is not None:
234
+ x = norm(x, batch=batch, training=training)
235
+ else:
236
+ x = norm(x, training=training)
237
+ if self.act is not None and not self.act_first:
238
+ x = self.act(x)
239
+ if self.dropout is not None:
240
+ x = self.dropout(x, training=training)
241
+ if hasattr(self, "jk"):
242
+ xs.append(x)
243
+
244
+ if hasattr(self, "jk"):
245
+ x = self.jk(xs)
246
+ if hasattr(self, "lin"):
247
+ x = self.lin(x)
248
+
249
+ return x
250
+
251
+ def __repr__(self) -> str:
252
+ return (f'{self.__class__.__name__}({self.in_channels}, '
253
+ f'{self.out_channels}, num_layers={self.num_layers})')
254
+
255
+
256
+ class GCN(BasicGNN):
257
+ r"""The Graph Neural Network from the `"Semi-supervised
258
+ Classification with Graph Convolutional Networks"
259
+ <https://arxiv.org/abs/1609.02907>`_ paper, using the
260
+ :class:`~k3_node.layers.conv.GCNConv` operator for message passing.
261
+
262
+ Example:
263
+ ```python
264
+ import numpy as np
265
+ from k3_node.models import GCN
266
+
267
+ x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
268
+ edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
269
+
270
+ model = GCN(in_channels=8, hidden_channels=16, num_layers=2, out_channels=4)
271
+ out = model(x, edge_index) # e.g. logits for 4 classes per node
272
+ print(tuple(out.shape)) # (10, 4)
273
+ ```
274
+ """
275
+ supports_edge_weight: bool = True
276
+ supports_edge_attr: bool = False
277
+
278
+ def init_conv(self, in_channels: int, out_channels: int, **kwargs) -> MessagePassing:
279
+ return GCNConv(in_channels, out_channels, **kwargs)
280
+
281
+
282
+ class GraphSAGE(BasicGNN):
283
+ r"""The Graph Neural Network from the `"Inductive Representation Learning
284
+ on Large Graphs" <https://arxiv.org/abs/1706.02216>`_ paper, using the
285
+ :class:`~k3_node.layers.conv.SAGEConv` operator for message passing.
286
+
287
+ Example:
288
+ ```python
289
+ import numpy as np
290
+ from k3_node.models import GraphSAGE
291
+
292
+ x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
293
+ edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
294
+
295
+ model = GraphSAGE(in_channels=8, hidden_channels=16, num_layers=2, out_channels=4)
296
+ out = model(x, edge_index) # e.g. logits for 4 classes per node
297
+ print(tuple(out.shape)) # (10, 4)
298
+ ```
299
+ """
300
+ supports_edge_weight: bool = False
301
+ supports_edge_attr: bool = False
302
+
303
+ def init_conv(self, in_channels: Union[int, Tuple[int, int]],
304
+ out_channels: int, **kwargs) -> MessagePassing:
305
+ return SAGEConv(in_channels, out_channels, **kwargs)
306
+
307
+
308
+ class GIN(BasicGNN):
309
+ r"""The Graph Neural Network from the `"How Powerful are Graph Neural
310
+ Networks?" <https://arxiv.org/abs/1810.00826>`_ paper, using the
311
+ :class:`~k3_node.layers.conv.GINConv` operator for message passing.
312
+
313
+ Example:
314
+ ```python
315
+ import numpy as np
316
+ from k3_node.models import GIN
317
+
318
+ x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
319
+ edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
320
+
321
+ model = GIN(in_channels=8, hidden_channels=16, num_layers=2, out_channels=4)
322
+ out = model(x, edge_index) # e.g. logits for 4 classes per node
323
+ print(tuple(out.shape)) # (10, 4)
324
+ ```
325
+ """
326
+ supports_edge_weight: bool = False
327
+ supports_edge_attr: bool = False
328
+
329
+ def init_conv(self, in_channels: int, out_channels: int, **kwargs) -> MessagePassing:
330
+ mlp = MLP(
331
+ [in_channels, out_channels, out_channels],
332
+ act=self.act,
333
+ act_first=self.act_first,
334
+ norm=self.norm_query,
335
+ norm_kwargs=self.norm_kwargs,
336
+ )
337
+ return GINConv(mlp, **kwargs)
338
+
339
+
340
+ class GAT(BasicGNN):
341
+ r"""The Graph Neural Network from `"Graph Attention Networks"
342
+ <https://arxiv.org/abs/1710.10903>`_ or `"How Attentive are Graph Attention
343
+ Networks?" <https://arxiv.org/abs/2105.14491>`_ papers, using the
344
+ :class:`~k3_node.layers.conv.GATConv` or
345
+ :class:`~k3_node.layers.conv.GATv2Conv` operator for message passing.
346
+
347
+ Example:
348
+ ```python
349
+ import numpy as np
350
+ from k3_node.models import GAT
351
+
352
+ x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
353
+ edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
354
+
355
+ model = GAT(in_channels=8, hidden_channels=16, num_layers=2, out_channels=4, heads=2)
356
+ out = model(x, edge_index) # e.g. logits for 4 classes per node
357
+ print(tuple(out.shape)) # (10, 4)
358
+ ```
359
+ """
360
+ supports_edge_weight: bool = False
361
+ supports_edge_attr: bool = True
362
+
363
+ def init_conv(self, in_channels: Union[int, Tuple[int, int]],
364
+ out_channels: int, **kwargs) -> MessagePassing:
365
+ v2 = kwargs.pop('v2', False)
366
+ heads = kwargs.pop('heads', 1)
367
+ concat = kwargs.pop('concat', True)
368
+
369
+ if getattr(self, '_is_conv_to_out', False):
370
+ concat = False
371
+
372
+ if concat and out_channels % heads != 0:
373
+ raise ValueError(f"Ensure that the number of output channels of "
374
+ f"'GATConv' (got '{out_channels}') is divisible "
375
+ f"by the number of heads (got '{heads}')")
376
+
377
+ if concat:
378
+ out_channels = out_channels // heads
379
+
380
+ Conv = GATConv if not v2 else GATv2Conv
381
+ return Conv(in_channels, out_channels, heads=heads, concat=concat,
382
+ dropout=self.dropout_p, **kwargs)
383
+
384
+
385
+ class PNA(BasicGNN):
386
+ r"""The Graph Neural Network from the `"Principal Neighbourhood Aggregation
387
+ for Graph Nets" <https://arxiv.org/abs/2004.05718>`_ paper, using the
388
+ :class:`~k3_node.layers.conv.PNAConv` operator for message passing.
389
+
390
+ Example:
391
+ ```python
392
+ import numpy as np
393
+ from k3_node.models import PNA
394
+
395
+ x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
396
+ edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
397
+
398
+ deg = np.array([0, 2, 4, 3, 1]) # in-degree histogram of the training graphs
399
+ model = PNA(in_channels=8, hidden_channels=16, num_layers=2, out_channels=4,
400
+ aggregators=["mean", "min", "max", "std"],
401
+ scalers=["identity", "amplification", "attenuation"], deg=deg)
402
+ out = model(x, edge_index)
403
+ print(tuple(out.shape)) # (10, 4)
404
+ ```
405
+ """
406
+ supports_edge_weight: bool = False
407
+ supports_edge_attr: bool = True
408
+
409
+ def init_conv(self, in_channels: int, out_channels: int, **kwargs) -> MessagePassing:
410
+ return PNAConv(in_channels, out_channels, **kwargs)
411
+
412
+
413
+ class EdgeCNN(BasicGNN):
414
+ r"""The Graph Neural Network from the `"Dynamic Graph CNN for Learning on
415
+ Point Clouds" <https://arxiv.org/abs/1801.07829>`_ paper, using the
416
+ :class:`~k3_node.layers.conv.EdgeConv` operator for message passing.
417
+
418
+ Example:
419
+ ```python
420
+ import numpy as np
421
+ from k3_node.models import EdgeCNN
422
+
423
+ x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
424
+ edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
425
+
426
+ model = EdgeCNN(in_channels=8, hidden_channels=16, num_layers=2, out_channels=4)
427
+ out = model(x, edge_index) # e.g. logits for 4 classes per node
428
+ print(tuple(out.shape)) # (10, 4)
429
+ ```
430
+ """
431
+ supports_edge_weight: bool = False
432
+ supports_edge_attr: bool = False
433
+
434
+ def init_conv(self, in_channels: int, out_channels: int, **kwargs) -> MessagePassing:
435
+ mlp = MLP(
436
+ [2 * in_channels, out_channels, out_channels],
437
+ act=self.act,
438
+ act_first=self.act_first,
439
+ norm=self.norm_query,
440
+ norm_kwargs=self.norm_kwargs,
441
+ )
442
+ return EdgeConv(mlp, **kwargs)
443
+
@@ -0,0 +1,4 @@
1
+ """Biological and macromolecular graph neural network models (aliased from k3_node.applications.bio)."""
2
+
3
+ from k3_node.applications.bio import *
4
+ from k3_node.applications.bio import __all__
@@ -0,0 +1,52 @@
1
+ r"""Captum model interpretability integration.
2
+
3
+ Note: Captum is a PyTorch-specific model interpretability library.
4
+ In k3-node (a multi-backend Keras 3 graph library supporting TensorFlow,
5
+ PyTorch, and JAX), Captum integrations are provided for PyTorch backend
6
+ compatibility or documented as PyTorch-exclusive.
7
+ """
8
+
9
+ from typing import Optional, Union, Any
10
+
11
+
12
+ def to_captum_model(
13
+ model: Any,
14
+ mask_type: str = "edge",
15
+ output_idx: Optional[int] = None,
16
+ metadata: Optional[Any] = None,
17
+ ):
18
+ r"""Converts a model into a Captum-compatible module.
19
+
20
+ .. note::
21
+ Captum is a PyTorch-exclusive library. This function requires
22
+ PyTorch backend and torch.nn.Module models.
23
+ """
24
+ raise NotImplementedError(
25
+ "Captum integration is PyTorch-specific and requires native torch.nn.Module."
26
+ )
27
+
28
+
29
+ def to_captum_input(
30
+ x: Any,
31
+ edge_index: Any,
32
+ mask_type: str = "edge",
33
+ *args,
34
+ **kwargs,
35
+ ):
36
+ r"""Converts graph inputs into Captum-compatible inputs."""
37
+ raise NotImplementedError(
38
+ "Captum integration is PyTorch-specific and requires native torch.nn.Module."
39
+ )
40
+
41
+
42
+ def captum_output_to_dicts(
43
+ attributions: Any,
44
+ mask_type: str = "edge",
45
+ *args,
46
+ **kwargs,
47
+ ):
48
+ r"""Converts Captum attributions into dictionaries."""
49
+ raise NotImplementedError(
50
+ "Captum integration is PyTorch-specific and requires native torch.nn.Module."
51
+ )
52
+
@@ -0,0 +1,4 @@
1
+ """Chemistry and molecular graph neural network models (aliased from k3_node.applications.chemistry)."""
2
+
3
+ from k3_node.applications.chemistry import *
4
+ from k3_node.applications.chemistry import __all__
@@ -0,0 +1,146 @@
1
+
2
+ import keras
3
+ from keras import ops
4
+
5
+ from k3_node.models.label_prop import LabelPropagation
6
+
7
+
8
+ class CorrectAndSmooth(keras.layers.Layer):
9
+ r"""The correct and smooth (C&S) post-processing model from the
10
+ `"Combining Label Propagation And Simple Models Out-performs Graph Neural
11
+ Networks" <https://arxiv.org/abs/2010.13993>`_ paper.
12
+
13
+ Args:
14
+ num_correction_layers (int): The number of propagations :math:`L_1`.
15
+ correction_alpha (float): The :math:`\alpha_1` coefficient.
16
+ num_smoothing_layers (int): The number of propagations :math:`L_2`.
17
+ smoothing_alpha (float): The :math:`\alpha_2` coefficient.
18
+ autoscale (bool, optional): If set to :obj:`True`, will automatically
19
+ determine the scaling factor :math:`\gamma`. (default: :obj:`True`)
20
+ scale (float, optional): The scaling factor :math:`\gamma`, in case
21
+ :obj:`autoscale = False`. (default: :obj:`1.0`)
22
+
23
+ Example:
24
+ ```python
25
+ import numpy as np
26
+ from k3_node.models import CorrectAndSmooth
27
+
28
+ y_soft = np.random.rand(6, 3).astype("float32") # base model's class probabilities
29
+ y_true = np.array([1, 0, 0, 2, 1, 1])
30
+ train_mask = np.array([True, False, True, False, True, False])
31
+ edge_index = np.array([[0, 1, 1, 2, 4, 5], [1, 0, 2, 1, 5, 4]])
32
+
33
+ model = CorrectAndSmooth(num_correction_layers=2, correction_alpha=0.5,
34
+ num_smoothing_layers=2, smoothing_alpha=0.5)
35
+ y_soft = model.correct(y_soft, y_true[train_mask], train_mask, edge_index) # propagate residual errors
36
+ y_soft = model.smooth(y_soft, y_true[train_mask], train_mask, edge_index) # propagate predictions
37
+ print(tuple(y_soft.shape)) # (6, 3)
38
+ ```
39
+ """
40
+ def __init__(
41
+ self,
42
+ num_correction_layers: int,
43
+ correction_alpha: float,
44
+ num_smoothing_layers: int,
45
+ smoothing_alpha: float,
46
+ autoscale: bool = True,
47
+ scale: float = 1.0,
48
+ **kwargs,
49
+ ):
50
+ super().__init__(**kwargs)
51
+ self.autoscale = autoscale
52
+ self.scale = scale
53
+
54
+ self.prop1 = LabelPropagation(num_correction_layers, correction_alpha)
55
+ self.prop2 = LabelPropagation(num_smoothing_layers, smoothing_alpha)
56
+
57
+ def build(self, input_shape=None):
58
+ self.built = True
59
+
60
+ def call(self, y_soft, *args, **kwargs):
61
+ y_soft = self.correct(y_soft, *args, **kwargs)
62
+ return self.smooth(y_soft, *args, **kwargs)
63
+
64
+ def correct(self, y_soft, y_true, mask, edge_index, edge_weight=None):
65
+ # Plain NumPy inputs cannot be mixed with backend tensors; convert them first.
66
+ y_soft, y_true, mask = ops.convert_to_tensor(y_soft), ops.convert_to_tensor(y_true), ops.convert_to_tensor(mask)
67
+ num_classes = ops.shape(y_soft)[-1]
68
+ y_true_shape = ops.shape(y_true)
69
+ if len(y_true_shape) == 1:
70
+ y_true = ops.one_hot(y_true, num_classes)
71
+ y_true = ops.cast(y_true, y_soft.dtype)
72
+
73
+ mask_shape = ops.shape(mask)
74
+ is_bool_mask = (len(mask_shape) == 1 and 'bool' in str(mask.dtype))
75
+ if is_bool_mask:
76
+ indices = ops.where(mask)[0]
77
+ else:
78
+ indices = mask
79
+
80
+ numel = float(ops.shape(indices)[0])
81
+ if ops.shape(y_true)[0] == ops.shape(y_soft)[0]:
82
+ y_true_sub = ops.take(y_true, indices, axis=0)
83
+ else:
84
+ y_true_sub = y_true
85
+
86
+ error_zeros = ops.zeros_like(y_soft)
87
+ error = ops.scatter_update(
88
+ error_zeros,
89
+ ops.expand_dims(indices, -1),
90
+ y_true_sub - ops.take(y_soft, indices, axis=0)
91
+ )
92
+
93
+ if self.autoscale:
94
+ smoothed_error = self.prop1(
95
+ error, edge_index, edge_weight=edge_weight,
96
+ post_step=lambda x: ops.clip(x, -1.0, 1.0)
97
+ )
98
+
99
+ error_masked = ops.take(error, indices, axis=0)
100
+ sigma = ops.sum(ops.abs(error_masked)) / max(numel, 1.0)
101
+ sum_smoothed = ops.sum(ops.abs(smoothed_error), axis=1, keepdims=True)
102
+ scale = sigma / ops.maximum(sum_smoothed, 1e-12)
103
+ scale = ops.where(scale > 1000.0, ops.ones_like(scale), scale)
104
+ return y_soft + scale * smoothed_error
105
+ else:
106
+ def fix_input(x):
107
+ return ops.scatter_update(x, ops.expand_dims(indices, -1), ops.take(error, indices, axis=0))
108
+
109
+ smoothed_error = self.prop1(
110
+ error, edge_index, edge_weight=edge_weight,
111
+ post_step=fix_input,
112
+ )
113
+ return y_soft + self.scale * smoothed_error
114
+
115
+ def smooth(self, y_soft, y_true, mask, edge_index, edge_weight=None):
116
+ # Plain NumPy inputs cannot be mixed with backend tensors; convert them first.
117
+ y_soft, y_true, mask = ops.convert_to_tensor(y_soft), ops.convert_to_tensor(y_true), ops.convert_to_tensor(mask)
118
+ num_classes = ops.shape(y_soft)[-1]
119
+ y_true_shape = ops.shape(y_true)
120
+ if len(y_true_shape) == 1:
121
+ y_true = ops.one_hot(y_true, num_classes)
122
+ y_true = ops.cast(y_true, y_soft.dtype)
123
+
124
+ mask_shape = ops.shape(mask)
125
+ is_bool_mask = (len(mask_shape) == 1 and 'bool' in str(mask.dtype))
126
+ if is_bool_mask:
127
+ indices = ops.where(mask)[0]
128
+ else:
129
+ indices = mask
130
+
131
+ if ops.shape(y_true)[0] == ops.shape(y_soft)[0]:
132
+ y_true_sub = ops.take(y_true, indices, axis=0)
133
+ else:
134
+ y_true_sub = y_true
135
+
136
+ y_soft = ops.scatter_update(y_soft, ops.expand_dims(indices, -1), y_true_sub)
137
+ return self.prop2(y_soft, edge_index, edge_weight=edge_weight)
138
+
139
+ def __repr__(self):
140
+ L1, alpha1 = self.prop1.num_layers, self.prop1.alpha
141
+ L2, alpha2 = self.prop2.num_layers, self.prop2.alpha
142
+ return (f'{self.__class__.__name__}(\n'
143
+ f' correct: num_layers={L1}, alpha={alpha1}\n'
144
+ f' smooth: num_layers={L2}, alpha={alpha2}\n'
145
+ f' autoscale={self.autoscale}, scale={self.scale}\n'
146
+ ')')