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,404 @@
1
+ """LPFormer, ported from PyG (``torch_geometric.nn.models.LPFormer``).
2
+
3
+ The choice of the nodes each candidate link attends to (common neighbors, one-hop neighbors and
4
+ nodes with a high personalized PageRank) depends on the graph only and is computed on the host
5
+ with SciPy sparse matrices; the learnable parts run with Keras ops. The model runs eagerly.
6
+ """
7
+ import math
8
+ from typing import List, Optional
9
+
10
+ import keras
11
+ import numpy as np
12
+ from keras import ops
13
+
14
+ from k3_node.models.basic_gnn import GCN
15
+ from k3_node.ops.segment import segment_sum
16
+
17
+
18
+ def get_ppr(edge_index, num_nodes: int, alpha: float = 0.15, eps: float = 5e-5):
19
+ r"""Approximate personalized PageRank of every node (the Andersen push algorithm used by
20
+ PyG's ``get_ppr``). Returns a ``scipy.sparse.csr_matrix`` whose row ``i`` holds the PPR scores
21
+ of the nodes reachable from ``i``.
22
+
23
+ Example:
24
+ ```python
25
+ import numpy as np
26
+ from k3_node.models.lpformer import get_ppr
27
+
28
+ edge_index = np.array([[0, 1, 1, 2], [1, 0, 2, 1]]) # a path 0 - 1 - 2
29
+ ppr = get_ppr(edge_index, num_nodes=3)
30
+ print(ppr.shape, round(float(ppr[0, 0]), 2)) # (3, 3) 0.19
31
+ ```
32
+ """
33
+ import scipy.sparse as sp
34
+
35
+ from k3_node.ops.host import to_numpy
36
+
37
+ edge_index = np.asarray(to_numpy(edge_index)).astype(np.int64)
38
+ order = np.lexsort((edge_index[1], edge_index[0])) # CSR with sorted columns, as PyG's EdgeIndex
39
+ col = edge_index[1][order]
40
+ rowptr = np.concatenate([[0], np.cumsum(np.bincount(edge_index[0], minlength=num_nodes))])
41
+ cols_list, vals_list = _ppr_push(rowptr, col, alpha, eps)
42
+ rows = np.repeat(np.arange(num_nodes), [len(c) for c in cols_list])
43
+ cols = np.concatenate(cols_list) if cols_list else np.zeros(0, np.int64)
44
+ vals = np.concatenate(vals_list) if vals_list else np.zeros(0)
45
+ return sp.csr_matrix((np.array(vals, dtype=np.float32), (rows, cols)), shape=(num_nodes, num_nodes))
46
+
47
+
48
+ def _ppr_push_python(rowptr, col, alpha, eps):
49
+ """PyG's Andersen push (``torch_geometric.utils.ppr._get_ppr``), one source node at a time."""
50
+ alpha_eps = alpha * eps
51
+ cols_list, vals_list = [], []
52
+ for inode in range(len(rowptr) - 1):
53
+ p, r, q, in_q = {inode: 0.0}, {inode: alpha}, [inode], {inode}
54
+ while q:
55
+ unode = q.pop()
56
+ in_q.discard(unode)
57
+ res = r.get(unode, 0.0)
58
+ p[unode] = p.get(unode, 0.0) + res
59
+ r[unode] = 0.0
60
+ start, end = rowptr[unode], rowptr[unode + 1]
61
+ ucount = end - start
62
+ for vnode in col[start:end]:
63
+ vnode = int(vnode)
64
+ r[vnode] = r.get(vnode, 0.0) + (1 - alpha) * res / ucount
65
+ if r[vnode] >= alpha_eps * (rowptr[vnode + 1] - rowptr[vnode]) and vnode not in in_q:
66
+ q.append(vnode)
67
+ in_q.add(vnode)
68
+ cols_list.append(np.fromiter(p.keys(), dtype=np.int64, count=len(p)))
69
+ vals_list.append(np.fromiter(p.values(), dtype=np.float64, count=len(p)))
70
+ return cols_list, vals_list
71
+
72
+
73
+ _PPR_NUMBA = None
74
+
75
+
76
+ def _ppr_push(rowptr, col, alpha, eps):
77
+ """Runs the push algorithm compiled with numba when it is installed (as PyG does)."""
78
+ global _PPR_NUMBA
79
+ try:
80
+ import numba
81
+ except ImportError:
82
+ return _ppr_push_python(rowptr, col, alpha, eps)
83
+ if _PPR_NUMBA is None:
84
+ def push(rowptr, col, alpha, eps):
85
+ num_nodes = len(rowptr) - 1
86
+ alpha_eps = alpha * eps
87
+ js = [[0]] * num_nodes
88
+ vals = [[0.0]] * num_nodes
89
+ for inode_uint in numba.prange(num_nodes):
90
+ inode = numba.int64(inode_uint)
91
+ p = {inode: 0.0}
92
+ r = {}
93
+ r[inode] = alpha
94
+ q = [inode]
95
+ while len(q) > 0:
96
+ unode = q.pop()
97
+ res = r[unode] if unode in r else 0
98
+ if unode in p:
99
+ p[unode] += res
100
+ else:
101
+ p[unode] = res
102
+ r[unode] = 0
103
+ start, end = rowptr[unode], rowptr[unode + 1]
104
+ ucount = end - start
105
+ for vnode in col[start:end]:
106
+ _val = (1 - alpha) * res / ucount
107
+ if vnode in r:
108
+ r[vnode] += _val
109
+ else:
110
+ r[vnode] = _val
111
+ res_vnode = r[vnode] if vnode in r else 0
112
+ vcount = rowptr[vnode + 1] - rowptr[vnode]
113
+ if res_vnode >= alpha_eps * vcount:
114
+ if vnode not in q:
115
+ q.append(vnode)
116
+ js[inode_uint] = list(p.keys())
117
+ vals[inode_uint] = list(p.values())
118
+ return js, vals
119
+
120
+ _PPR_NUMBA = numba.jit(nopython=True, parallel=True)(push)
121
+ js, vals = _PPR_NUMBA(rowptr, col, alpha, eps)
122
+ return [np.asarray(j, dtype=np.int64) for j in js], [np.asarray(v, dtype=np.float64) for v in vals]
123
+
124
+
125
+ def compute_ppr_matrix(edge_index, num_nodes: int, alpha: float = 0.15, eps: float = 5e-5):
126
+ r"""Alias of :func:`get_ppr`."""
127
+ return get_ppr(edge_index, num_nodes, alpha=alpha, eps=eps)
128
+
129
+
130
+ def _lookup(matrix, rows, cols):
131
+ """``matrix[rows[k], cols[k]]`` for every ``k``, as a 1-D array (also when empty)."""
132
+ if len(rows) == 0:
133
+ return np.zeros(0, dtype=np.float32)
134
+ values = matrix[rows, cols]
135
+ values = values.toarray() if hasattr(values, "toarray") else values
136
+ return np.asarray(values, dtype=np.float32).ravel()
137
+
138
+
139
+ class MLP(keras.layers.Layer):
140
+ r"""The small MLP of LPFormer: linear layers with (layer) normalization, ReLU and dropout
141
+ between them; the last dimension is squeezed if it has size 1."""
142
+
143
+ def __init__(self, in_channels: int, hid_channels: int, out_channels: int, num_layers: int = 2,
144
+ drop: float = 0.0, norm: Optional[str] = "layer", **kwargs):
145
+ super().__init__(**kwargs)
146
+ self.linears = ([keras.layers.Dense(out_channels)] if num_layers == 1 else
147
+ [keras.layers.Dense(hid_channels) for _ in range(num_layers - 1)] + [keras.layers.Dense(out_channels)])
148
+ if norm == "batch":
149
+ self.norm = keras.layers.BatchNormalization(momentum=0.9, epsilon=1e-5)
150
+ elif norm == "layer":
151
+ self.norm = keras.layers.LayerNormalization(epsilon=1e-5)
152
+ else:
153
+ self.norm = None
154
+ self.dropout = keras.layers.Dropout(drop)
155
+
156
+ def call(self, x, training=False):
157
+ for lin in self.linears[:-1]:
158
+ x = lin(x)
159
+ x = self.norm(x, training=training) if self.norm is not None else x
160
+ x = self.dropout(ops.relu(x), training=training)
161
+ x = self.linears[-1](x)
162
+ return ops.squeeze(x, axis=-1) if x.shape[-1] == 1 else x
163
+
164
+
165
+ class LPAttLayer(keras.layers.Layer):
166
+ r"""Attention of every candidate link over its selected nodes (PyG's ``LPAttLayer``).
167
+
168
+ Used inside :class:`LPFormer`. ``edge_index[0]`` is the link and ``edge_index[1]`` a node it
169
+ attends to; ``edge_feats`` holds the two end-node features of every link side by side and
170
+ ``ppr_rpes`` a relative positional encoding for every (link, node) pair.
171
+
172
+ Example:
173
+ ```python
174
+ import numpy as np
175
+ from k3_node.models import LPAttLayer
176
+
177
+ layer = LPAttLayer(in_channels=8, out_channels=8, node_dim=None, num_heads=2, dropout=0.0)
178
+ edge_feats = np.random.rand(4, 16).astype("float32") # 4 links: [x_src | x_dst]
179
+ node_feats = np.random.rand(10, 8).astype("float32") # 10 nodes
180
+ edge_index = np.stack([np.repeat(np.arange(4), 3), np.random.randint(0, 10, size=12)])
181
+ ppr_rpes = np.random.rand(12, 8).astype("float32") # one encoding per (link, node) pair
182
+ out = layer(edge_index, edge_feats, node_feats, ppr_rpes)
183
+ print(tuple(out.shape)) # (4, 16)
184
+ ```
185
+ """
186
+
187
+ def __init__(self, in_channels: int, out_channels: int, node_dim: Optional[int], num_heads: int,
188
+ dropout: float, concat: bool = True, **kwargs):
189
+ super().__init__(**kwargs)
190
+ self.in_channels, self.out_channels, self.heads, self.concat = in_channels, out_channels, num_heads, concat
191
+ self.negative_slope = 0.2
192
+ self.lin_l = keras.layers.Dense(num_heads * out_channels, kernel_initializer="glorot_uniform")
193
+ self.lin_r = keras.layers.Dense(num_heads * out_channels, kernel_initializer="glorot_uniform")
194
+ self.att = self.add_weight(shape=(1, num_heads, out_channels), initializer="glorot_uniform", name="att")
195
+ self.bias = self.add_weight(shape=(num_heads * out_channels if concat else out_channels,),
196
+ initializer="zeros", name="bias")
197
+ self.post_att_norm = keras.layers.LayerNormalization(epsilon=1e-5)
198
+ self.dropout = keras.layers.Dropout(dropout)
199
+
200
+ def call(self, edge_index, edge_feats, node_feats, ppr_rpes, training=False):
201
+ H, C = self.heads, self.out_channels
202
+ pair, node = edge_index[0], edge_index[1] # "target_to_source": link i attends to node j
203
+ num_links = edge_feats.shape[0]
204
+ x_i = ops.take(edge_feats, pair, axis=0)
205
+ x_j = ops.concatenate([ops.take(node_feats, node, axis=0), ppr_rpes], axis=-1)
206
+ x_j = ops.reshape(self.lin_r(x_j), (-1, H, C))
207
+ e1, e2 = ops.split(x_i, 2, axis=-1)
208
+ x = ops.leaky_relu(x_j * (ops.reshape(self.lin_l(e1), (-1, H, C)) + ops.reshape(self.lin_l(e2), (-1, H, C))),
209
+ negative_slope=self.negative_slope)
210
+ alpha = ops.sum(x * self.att, axis=-1) # [K, H]
211
+ from k3_node.layers.conv.utils import softmax
212
+
213
+ alpha = softmax(alpha, pair, num_nodes=num_links)
214
+ out = segment_sum(x_j * ops.expand_dims(alpha, -1), pair, num_segments=num_links) # [B, H, C]
215
+ out = ops.reshape(out, (-1, H * C)) if self.concat else ops.mean(out, axis=1)
216
+ out = self.post_att_norm(out + self.bias)
217
+ return self.dropout(out, training=training)
218
+
219
+
220
+ class LPFormer(keras.Model):
221
+ r"""The LPFormer model from the `"LPFormer: An Adaptive Graph Transformer for Link Prediction"
222
+ <https://arxiv.org/abs/2310.11009>`_ paper, as in PyG.
223
+
224
+ For every candidate link it attends over the common neighbors, the one-hop neighbors and the
225
+ other nodes with a high personalized PageRank (PPR) from both endpoints (``ppr_thresholds``
226
+ for the three kinds), using their PPR scores as relative positional encodings, and combines
227
+ this with counts of each kind of node and with the GCN embeddings of the two endpoints.
228
+
229
+ Args:
230
+ in_channels (int): Input feature dimension.
231
+ hidden_channels (int): Hidden dimension.
232
+ num_gnn_layers (int, optional): Number of GCN layers. (default: ``2``)
233
+ gnn_dropout (float, optional): GCN dropout. (default: ``0.1``)
234
+ num_transformer_layers (int, optional): Number of attention layers. (default: ``1``)
235
+ num_heads (int, optional): Number of attention heads. (default: ``1``)
236
+ transformer_dropout (float, optional): Attention dropout; during training this share of
237
+ the selected nodes is also dropped. (default: ``0.1``)
238
+ ppr_thresholds (list, optional): Minimum PPR of common neighbors, one-hop neighbors and
239
+ other nodes. (default: ``[0, 1e-4, 1e-2]``)
240
+
241
+ Call arguments: ``batch`` (the ``[2, num_links]`` candidate links), ``x`` (node features),
242
+ ``edge_index`` (the graph) and the keyword ``ppr_matrix`` (from :meth:`calc_sparse_ppr`). Returns one logit
243
+ per link. The node selection runs on the host, so train with ``run_eagerly=True`` or
244
+ :func:`~k3_node.training.gradient_step`.
245
+
246
+ Example:
247
+ ```python
248
+ import numpy as np
249
+ from k3_node.models import LPFormer
250
+
251
+ x = np.random.rand(10, 16).astype("float32") # 10 nodes with 16 features each
252
+ edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
253
+ model = LPFormer(in_channels=16, hidden_channels=16)
254
+ ppr = model.calc_sparse_ppr(edge_index, num_nodes=10)
255
+ target_links = np.array([[0, 1], [2, 3]]) # the (source, target) pairs to score
256
+ print(tuple(model(target_links, x, edge_index, ppr_matrix=ppr).shape)) # (2,): one logit per link
257
+ ```
258
+ """
259
+
260
+ def __init__(self, in_channels: int, hidden_channels: int, num_gnn_layers: int = 2, gnn_dropout: float = 0.1,
261
+ num_transformer_layers: int = 1, num_heads: int = 1, transformer_dropout: float = 0.1,
262
+ ppr_thresholds: Optional[List[float]] = None, **kwargs):
263
+ super().__init__(**kwargs)
264
+ ppr_thresholds = [0, 1e-4, 1e-2] if ppr_thresholds is None else ppr_thresholds
265
+ if len(ppr_thresholds) != 3:
266
+ raise ValueError("Argument 'ppr_thresholds' must only be length 3!")
267
+ self.thresh_cn, self.thresh_1hop, self.thresh_non1hop = ppr_thresholds
268
+ self.in_dim, self.hid_dim = in_channels, hidden_channels
269
+ self.trans_drop = transformer_dropout
270
+ self.gnn = GCN(in_channels, hidden_channels, num_gnn_layers, dropout=gnn_dropout, norm="layer_norm")
271
+ self.gnn_norm = keras.layers.LayerNormalization(epsilon=1e-5)
272
+ self.input_dropout = keras.layers.Dropout(gnn_dropout)
273
+ self.att_layers = []
274
+ for il in range(num_transformer_layers):
275
+ if il == 0:
276
+ node_dim = None
277
+ self.out_dim = hidden_channels * 2 if num_transformer_layers > 1 else hidden_channels
278
+ else:
279
+ self.out_dim = node_dim = hidden_channels
280
+ self.att_layers.append(LPAttLayer(hidden_channels, self.out_dim, node_dim, num_heads, transformer_dropout))
281
+ self.elementwise_lin = MLP(hidden_channels, hidden_channels, hidden_channels)
282
+ self.ppr_encoder_cn = MLP(2, hidden_channels, hidden_channels)
283
+ self.ppr_encoder_onehop = MLP(2, hidden_channels, hidden_channels)
284
+ self.ppr_encoder_non1hop = MLP(2, hidden_channels, hidden_channels)
285
+ if self.thresh_non1hop == 1 and self.thresh_1hop == 1:
286
+ self.mask = "cn"
287
+ elif self.thresh_non1hop == 1 and self.thresh_1hop < 1:
288
+ self.mask = "1-hop"
289
+ else:
290
+ self.mask = "all"
291
+ pairwise_dim = hidden_channels * num_heads + 4
292
+ self.pairwise_lin = MLP(pairwise_dim, pairwise_dim, hidden_channels)
293
+ self.score_func = MLP(hidden_channels * 2, hidden_channels * 2, 1, norm=None)
294
+
295
+ @staticmethod
296
+ def calc_sparse_ppr(edge_index, num_nodes: int, alpha: float = 0.15, eps: float = 5e-5):
297
+ r"""The personalized PageRank matrix LPFormer needs (see :func:`get_ppr`)."""
298
+ return get_ppr(edge_index, num_nodes, alpha=alpha, eps=eps)
299
+
300
+ # ---- node selection on the host ----------------------------------------------------------
301
+ def _drop(self, rng, n):
302
+ keep = math.ceil(n * (1 - self.trans_drop))
303
+ return rng.permutation(n)[:keep]
304
+
305
+ def compute_node_mask(self, u, v, adj, ppr, training):
306
+ r"""For the links ``(u, v)``: ``(pair index, node, PPR from u, PPR from v)`` of their common
307
+ neighbors, one-hop neighbors and other high-PPR nodes."""
308
+ pair_adj = (adj[u] * adj[v]) if self.mask == "cn" else (adj[u] + adj[v])
309
+ pair_adj = pair_adj.tocoo()
310
+ order = np.lexsort((pair_adj.col, pair_adj.row))
311
+ rows, cols, node_type = pair_adj.row[order], pair_adj.col[order], pair_adj.data[order]
312
+ src_ppr = _lookup(ppr, u[rows], cols)
313
+ tgt_ppr = _lookup(ppr, v[rows], cols)
314
+ cn_cond = (src_ppr >= self.thresh_cn) & (tgt_ppr >= self.thresh_cn)
315
+ onehop_cond = (src_ppr >= self.thresh_1hop) & (tgt_ppr >= self.thresh_1hop)
316
+ keep = np.where(node_type == 1, onehop_cond, cn_cond) if self.mask != "cn" else np.where(node_type == 0, onehop_cond, cn_cond)
317
+ rows, cols, node_type, src_ppr, tgt_ppr = rows[keep], cols[keep], node_type[keep], src_ppr[keep], tgt_ppr[keep]
318
+
319
+ non1hop = None
320
+ if self.mask == "all":
321
+ non1hop = self._non_1hop(u, v, adj, ppr, training)
322
+ rng = np.random
323
+ if training and self.trans_drop > 0:
324
+ idx = self._drop(rng, len(rows))
325
+ rows, cols, node_type, src_ppr, tgt_ppr = rows[idx], cols[idx], node_type[idx], src_ppr[idx], tgt_ppr[idx]
326
+ if non1hop is not None:
327
+ idx = self._drop(rng, len(non1hop[0]))
328
+ non1hop = tuple(a[idx] for a in non1hop)
329
+ if self.mask == "cn":
330
+ return (rows, cols, src_ppr, tgt_ppr), None, None
331
+ cn = node_type == 2
332
+ one = node_type == 1
333
+ return ((rows[cn], cols[cn], src_ppr[cn], tgt_ppr[cn]), (rows[one], cols[one], src_ppr[one], tgt_ppr[one]),
334
+ non1hop)
335
+
336
+ def _non_1hop(self, u, v, adj, ppr, training):
337
+ import scipy.sparse as sp
338
+
339
+ adj2 = adj
340
+ if training: # the links being predicted are known edges during training
341
+ n = adj.shape[0]
342
+ links = sp.csr_matrix((np.ones(2 * len(u)), (np.concatenate([u, v]), np.concatenate([v, u]))), shape=(n, n))
343
+ adj2 = ((adj + links) > 0).astype(np.float32).tocsr()
344
+ neighbors = ((adj2[u] + adj2[v]) > 0).astype(np.float32)
345
+ src, tgt = ppr[u], ppr[v]
346
+ both = (src >= self.thresh_non1hop).astype(np.float32).multiply((tgt >= self.thresh_non1hop).astype(np.float32))
347
+ both = sp.csr_matrix(both - both.multiply(neighbors)) # high PPR from both ends, not a neighbor
348
+ both.eliminate_zeros()
349
+ both = both.tocoo()
350
+ order = np.lexsort((both.col, both.row))
351
+ rows, cols = both.row[order], both.col[order]
352
+ return rows, cols, _lookup(ppr, u[rows], cols), _lookup(ppr, v[rows], cols)
353
+
354
+ # ---- forward ------------------------------------------------------------------------------------
355
+ def _pos_encoding(self, encoder, s, t, training):
356
+ a = ops.convert_to_tensor(np.stack([s, t], axis=1).astype(np.float32))
357
+ b = ops.convert_to_tensor(np.stack([t, s], axis=1).astype(np.float32))
358
+ return encoder(a, training=training) + encoder(b, training=training)
359
+
360
+ def call(self, batch, x, edge_index, ppr_matrix=None, training=False):
361
+ import scipy.sparse as sp
362
+
363
+ from k3_node.ops.host import to_numpy
364
+
365
+ batch_np = np.asarray(to_numpy(batch)).astype(np.int64)
366
+ edge_np = np.asarray(to_numpy(edge_index)).astype(np.int64)
367
+ num_nodes = x.shape[0]
368
+ if ppr_matrix is None:
369
+ ppr_matrix = self.calc_sparse_ppr(edge_np, num_nodes)
370
+ ppr = sp.csr_matrix(ppr_matrix)
371
+ adj = sp.csr_matrix((np.ones(edge_np.shape[1], np.float32), (edge_np[0], edge_np[1])), shape=(num_nodes, num_nodes))
372
+ adj.data[:] = 1.0 # {0, 1} even with duplicate edges
373
+
374
+ X_node = self.gnn_norm(self.gnn(self.input_dropout(x, training=training), edge_index, training=training))
375
+ u, v = batch_np[0], batch_np[1]
376
+ x_i, x_j = ops.take(X_node, u, axis=0), ops.take(X_node, v, axis=0)
377
+ elementwise = self.elementwise_lin(x_i * x_j, training=training)
378
+
379
+ cn, onehop, non1hop = self.compute_node_mask(u, v, adj, ppr, training)
380
+ groups = [(cn, self.ppr_encoder_cn), (onehop, self.ppr_encoder_onehop), (non1hop, self.ppr_encoder_non1hop)]
381
+ groups = [(g, enc) for g, enc in groups if g is not None]
382
+ rows = np.concatenate([g[0] for g, _ in groups])
383
+ cols = np.concatenate([g[1] for g, _ in groups])
384
+ pes = ops.concatenate([self._pos_encoding(enc, g[2], g[3], training) for g, enc in groups], axis=0)
385
+ all_mask = ops.convert_to_tensor(np.stack([rows, cols]).astype(np.int32))
386
+
387
+ pairwise = ops.concatenate([x_i, x_j], axis=-1)
388
+ for layer in self.att_layers:
389
+ pairwise = layer(all_mask, pairwise, X_node, pes, training=training)
390
+
391
+ B = len(u)
392
+ counts = [np.bincount(cn[0], minlength=B)] # common neighbors (all pass thresh_cn)
393
+ if onehop is not None:
394
+ num_1hop = np.bincount(onehop[0][(onehop[2] >= self.thresh_1hop) & (onehop[3] >= self.thresh_1hop)], minlength=B)
395
+ num_ppr_ones = np.bincount(onehop[0], minlength=B)
396
+ counts += [num_1hop]
397
+ else:
398
+ num_ppr_ones = np.zeros(B)
399
+ counts += [np.zeros(B)]
400
+ counts += [np.bincount(non1hop[0], minlength=B) if non1hop is not None else np.zeros(B), counts[0] + num_ppr_ones]
401
+ counts = ops.convert_to_tensor(np.stack(counts, axis=1).astype(np.float32))
402
+
403
+ pairwise = self.pairwise_lin(ops.concatenate([pairwise, counts], axis=-1), training=training)
404
+ return self.score_func(ops.concatenate([elementwise, pairwise], axis=-1), training=training)
@@ -0,0 +1,114 @@
1
+ from typing import Optional
2
+
3
+ import keras
4
+ from keras import ops
5
+ import numpy as np
6
+
7
+
8
+ class MaskLabel(keras.layers.Layer):
9
+ r"""The label embedding and masking layer from the `"Masked Label
10
+ Prediction: Unified Message Passing Model for Semi-Supervised
11
+ Classification" <https://arxiv.org/abs/2009.03509>`_ paper.
12
+
13
+ Here, node labels :obj:`y` are merged to the initial node features :obj:`x`
14
+ for a subset of their nodes according to :obj:`mask`.
15
+
16
+ Args:
17
+ num_classes (int): The number of classes.
18
+ out_channels (int): Size of each output sample.
19
+ method (str, optional): If set to :obj:`"add"`, label embeddings are
20
+ added to the input. If set to :obj:`"concat"`, label embeddings are
21
+ concatenated. In case :obj:`method="add"`, then :obj:`out_channels`
22
+ needs to be identical to the input dimensionality of node features.
23
+ (default: :obj:`"add"`)
24
+
25
+ Example:
26
+ ```python
27
+ import numpy as np
28
+ from k3_node.models import MaskLabel
29
+
30
+ x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
31
+ edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
32
+
33
+ y = np.random.randint(0, 3, size=(10,)) # node labels
34
+ known = np.random.rand(10) > 0.5 # labels visible to the model
35
+ model = MaskLabel(num_classes=3, out_channels=8) # label embedding added to x
36
+ out = model(x, y, known)
37
+ print(tuple(out.shape)) # (10, 8)
38
+ ```
39
+ """
40
+
41
+ def __init__(
42
+ self,
43
+ num_classes: int,
44
+ out_channels: int,
45
+ method: str = "add",
46
+ **kwargs,
47
+ ):
48
+ super().__init__(**kwargs)
49
+
50
+ self.num_classes = num_classes
51
+ self.out_channels = out_channels
52
+ self.method = method
53
+
54
+ if method not in ["add", "concat"]:
55
+ raise ValueError(
56
+ f"'method' must be either 'add' or 'concat' (got '{method}')"
57
+ )
58
+
59
+ self.emb = keras.layers.Embedding(num_classes, out_channels)
60
+
61
+ def build(self, input_shape=None):
62
+ self.emb.build((None,))
63
+ self.built = True
64
+
65
+ def reset_parameters(self) -> None:
66
+ r"""Resets all learnable parameters of the module."""
67
+ if self.emb.built:
68
+ self.emb.embeddings.assign(
69
+ keras.initializers.Uniform()(self.emb.embeddings.shape)
70
+ )
71
+
72
+ def call(self, x, y, mask):
73
+ """Forward pass.
74
+
75
+ Args:
76
+ x (Tensor): The input node features of shape ``[N, in_channels]``.
77
+ y (Tensor): The node labels of shape ``[N]`` (integer class indices).
78
+ mask (Tensor): Boolean tensor of shape ``[N]`` indicating which
79
+ nodes have ground-truth labels to embed.
80
+ """
81
+ # Embed all labels; then zero-out non-masked entries
82
+ all_emb = self.emb(y) # [N, out_channels]
83
+
84
+ # Build a float mask: 1.0 where mask is True, 0.0 elsewhere
85
+ float_mask = ops.cast(mask, dtype=all_emb.dtype) # [N]
86
+ float_mask = ops.expand_dims(float_mask, axis=-1) # [N, 1]
87
+ masked_emb = all_emb * float_mask # [N, out_channels]
88
+
89
+ if self.method == "concat":
90
+ return ops.concatenate([x, masked_emb], axis=-1)
91
+ else:
92
+ return x + masked_emb
93
+
94
+ @staticmethod
95
+ def ratio_mask(mask, ratio: float):
96
+ r"""Modifies :obj:`mask` by setting :obj:`ratio` of :obj:`True`
97
+ entries to :obj:`False`. Does not operate in-place.
98
+
99
+ Args:
100
+ mask (Tensor): The boolean mask to re-mask.
101
+ ratio (float): The ratio of True entries to keep.
102
+ """
103
+ mask_np = ops.convert_to_numpy(mask).astype(bool)
104
+ n = int(mask_np.sum())
105
+ out_np = mask_np.copy()
106
+ if n > 0:
107
+ keep = np.random.rand(n) < ratio
108
+ true_indices = np.where(mask_np)[0]
109
+ out_np[true_indices] = keep
110
+ return ops.convert_to_tensor(out_np, dtype="bool")
111
+
112
+ def __repr__(self) -> str:
113
+ return f'{self.__class__.__name__}()'
114
+
@@ -0,0 +1,33 @@
1
+ """Materials and crystal models (aliased from k3_node.applications.materials)."""
2
+
3
+ import sys
4
+ from k3_node.applications.materials import *
5
+ from k3_node.applications.materials import (
6
+ basis,
7
+ core,
8
+ readout,
9
+ wrappers,
10
+ io,
11
+ megnet,
12
+ m3gnet,
13
+ tensornet,
14
+ chgnet,
15
+ so3net,
16
+ grace,
17
+ qet,
18
+ )
19
+ from k3_node.applications.materials import __all__
20
+
21
+ # Alias submodules in sys.modules for full backward compatibility
22
+ sys.modules[__name__ + ".basis"] = basis
23
+ sys.modules[__name__ + ".core"] = core
24
+ sys.modules[__name__ + ".readout"] = readout
25
+ sys.modules[__name__ + ".wrappers"] = wrappers
26
+ sys.modules[__name__ + ".io"] = io
27
+ sys.modules[__name__ + ".megnet"] = megnet
28
+ sys.modules[__name__ + ".m3gnet"] = m3gnet
29
+ sys.modules[__name__ + ".tensornet"] = tensornet
30
+ sys.modules[__name__ + ".chgnet"] = chgnet
31
+ sys.modules[__name__ + ".so3net"] = so3net
32
+ sys.modules[__name__ + ".grace"] = grace
33
+ sys.modules[__name__ + ".qet"] = qet
k3_node/models/meta.py ADDED
@@ -0,0 +1,133 @@
1
+ from typing import Optional, Tuple
2
+
3
+ import keras
4
+ from keras import ops
5
+
6
+
7
+ class MetaLayer(keras.layers.Layer):
8
+ r"""A meta layer for building any kind of graph network, inspired by the
9
+ `"Relational Inductive Biases, Deep Learning, and Graph Networks"
10
+ <https://arxiv.org/abs/1806.01261>`_ paper.
11
+
12
+ A graph network takes a graph as input and returns an updated graph as
13
+ output (with same connectivity). The input graph has node features :obj:`x`,
14
+ edge features :obj:`edge_attr` as well as graph-level features :obj:`u`.
15
+ The output graph has the same structure, but updated features.
16
+
17
+ Edge features, node features as well as global features are updated by
18
+ calling the modules :obj:`edge_model`, :obj:`node_model` and
19
+ :obj:`global_model`, respectively.
20
+
21
+ To allow for batch-wise graph processing, all callable functions take an
22
+ additional argument :obj:`batch`, which determines the assignment of
23
+ edges or nodes to their specific graphs.
24
+
25
+ Args:
26
+ edge_model (callable, optional): A callable which updates a graph's
27
+ edge features based on its source and target node features, its
28
+ current edge features and its global features.
29
+ (default: :obj:`None`)
30
+ node_model (callable, optional): A callable which updates a graph's
31
+ node features based on its current node features, its graph
32
+ connectivity, its edge features and its global features.
33
+ (default: :obj:`None`)
34
+ global_model (callable, optional): A callable which updates a graph's
35
+ global features based on its node features, its graph connectivity,
36
+ its edge features and its current global features.
37
+ (default: :obj:`None`)
38
+
39
+ Example:
40
+ ```python
41
+ import numpy as np
42
+ import keras
43
+ from k3_node.models import MetaLayer
44
+
45
+ x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
46
+ edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
47
+ edge_attr = np.random.rand(30, 4).astype("float32")
48
+
49
+ edge_mlp = keras.layers.Dense(8)
50
+ node_mlp = keras.layers.Dense(8)
51
+
52
+ def edge_model(src, dst, edge_attr, u, batch): # update every edge from its endpoints
53
+ return edge_mlp(keras.ops.concatenate([src, dst, edge_attr], axis=-1))
54
+
55
+ def node_model(x, edge_index, edge_attr, u, batch): # update nodes from incoming edges
56
+ incoming = keras.ops.segment_sum(edge_attr, edge_index[1], num_segments=10)
57
+ return node_mlp(keras.ops.concatenate([x, incoming], axis=-1))
58
+
59
+ model = MetaLayer(edge_model=edge_model, node_model=node_model)
60
+ x_out, edge_attr_out, u_out = model(x, edge_index, edge_attr=edge_attr)
61
+ print(tuple(x_out.shape), tuple(edge_attr_out.shape)) # (10, 8) (30, 8)
62
+ ```
63
+ """
64
+
65
+ def __init__(
66
+ self,
67
+ edge_model=None,
68
+ node_model=None,
69
+ global_model=None,
70
+ **kwargs,
71
+ ):
72
+ super().__init__(**kwargs)
73
+ self.edge_model = edge_model
74
+ self.node_model = node_model
75
+ self.global_model = global_model
76
+
77
+ self.reset_parameters()
78
+
79
+ def reset_parameters(self) -> None:
80
+ r"""Resets all learnable parameters of the module."""
81
+ for item in [self.node_model, self.edge_model, self.global_model]:
82
+ if hasattr(item, 'reset_parameters'):
83
+ item.reset_parameters()
84
+
85
+ def call(
86
+ self,
87
+ x,
88
+ edge_index,
89
+ edge_attr=None,
90
+ u=None,
91
+ batch=None,
92
+ ) -> Tuple:
93
+ r"""Forward pass.
94
+
95
+ Args:
96
+ x (Tensor): The node features of shape ``[N, F_x]``.
97
+ edge_index (Tensor): The edge indices of shape ``[2, E]``.
98
+ edge_attr (Tensor, optional): The edge features of shape
99
+ ``[E, F_e]``. (default: :obj:`None`)
100
+ u (Tensor, optional): The global graph features of shape
101
+ ``[B, F_u]``. (default: :obj:`None`)
102
+ batch (Tensor, optional): The batch vector
103
+ :math:`\mathbf{b} \in {\{ 0, \ldots, B-1\}}^N`.
104
+ (default: :obj:`None`)
105
+ """
106
+ row = edge_index[0]
107
+ col = edge_index[1]
108
+
109
+ if self.edge_model is not None:
110
+ edge_batch = batch if batch is None else ops.take(batch, row, axis=0)
111
+ edge_attr = self.edge_model(
112
+ ops.take(x, row, axis=0),
113
+ ops.take(x, col, axis=0),
114
+ edge_attr,
115
+ u,
116
+ edge_batch,
117
+ )
118
+
119
+ if self.node_model is not None:
120
+ x = self.node_model(x, edge_index, edge_attr, u, batch)
121
+
122
+ if self.global_model is not None:
123
+ u = self.global_model(x, edge_index, edge_attr, u, batch)
124
+
125
+ return x, edge_attr, u
126
+
127
+ def __repr__(self) -> str:
128
+ return (f'{self.__class__.__name__}(\n'
129
+ f' edge_model={self.edge_model},\n'
130
+ f' node_model={self.node_model},\n'
131
+ f' global_model={self.global_model}\n'
132
+ f')')
133
+