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,235 @@
1
+ import re
2
+ import warnings
3
+ from typing import Any, Dict, List, Optional
4
+
5
+ import numpy as np
6
+ from keras import ops
7
+
8
+
9
+ x_map: Dict[str, List[Any]] = {
10
+ "atomic_num": list(range(0, 119)),
11
+ "chirality": [
12
+ "CHI_UNSPECIFIED",
13
+ "CHI_TETRAHEDRAL_CW",
14
+ "CHI_TETRAHEDRAL_CCW",
15
+ "CHI_OTHER",
16
+ "CHI_TETRAHEDRAL",
17
+ "CHI_ALLENE",
18
+ "CHI_SQUAREPLANAR",
19
+ "CHI_TRIGONALBIPYRAMIDAL",
20
+ "CHI_OCTAHEDRAL",
21
+ ],
22
+ "degree": list(range(0, 11)),
23
+ "formal_charge": list(range(-5, 7)),
24
+ "num_hs": list(range(0, 9)),
25
+ "num_radical_electrons": list(range(0, 5)),
26
+ "hybridization": [
27
+ "UNSPECIFIED",
28
+ "S",
29
+ "SP",
30
+ "SP2",
31
+ "SP3",
32
+ "SP3D",
33
+ "SP3D2",
34
+ "OTHER",
35
+ ],
36
+ "is_aromatic": [False, True],
37
+ "is_in_ring": [False, True],
38
+ }
39
+
40
+ e_map: Dict[str, List[Any]] = {
41
+ "bond_type": [
42
+ "UNSPECIFIED",
43
+ "SINGLE",
44
+ "DOUBLE",
45
+ "TRIPLE",
46
+ "QUADRUPLE",
47
+ "QUINTUPLE",
48
+ "HEXTUPLE",
49
+ "ONEANDAHALF",
50
+ "TWOANDAHALF",
51
+ "THREEANDAHALF",
52
+ "FOURANDAHALF",
53
+ "FIVEANDAHALF",
54
+ "AROMATIC",
55
+ "IONIC",
56
+ "HYDROGEN",
57
+ "THREECENTER",
58
+ "DATIVEONE",
59
+ "DATIVE",
60
+ "DATIVEL",
61
+ "DATIVER",
62
+ "OTHER",
63
+ "ZERO",
64
+ ],
65
+ "stereo": [
66
+ "STEREONONE",
67
+ "STEREOANY",
68
+ "STEREOZ",
69
+ "STEREOE",
70
+ "STEREOCIS",
71
+ "STEREOTRANS",
72
+ ],
73
+ "is_conjugated": [False, True],
74
+ }
75
+
76
+
77
+ def from_rdmol(mol: Any) -> "Any":
78
+ r"""Converts an :class:`rdkit.Chem.Mol` instance to a :class:`k3_node.data.Data` instance."""
79
+ from rdkit import Chem
80
+ from k3_node.data.data import Data
81
+
82
+ assert isinstance(mol, Chem.Mol)
83
+
84
+ xs: List[List[int]] = []
85
+ for atom in mol.GetAtoms():
86
+ row: List[int] = []
87
+ row.append(x_map["atomic_num"].index(atom.GetAtomicNum()))
88
+ row.append(x_map["chirality"].index(str(atom.GetChiralTag())))
89
+ row.append(x_map["degree"].index(atom.GetTotalDegree()))
90
+ row.append(x_map["formal_charge"].index(atom.GetFormalCharge()))
91
+ row.append(x_map["num_hs"].index(atom.GetTotalNumHs()))
92
+ row.append(x_map["num_radical_electrons"].index(atom.GetNumRadicalElectrons()))
93
+ row.append(x_map["hybridization"].index(str(atom.GetHybridization())))
94
+ row.append(x_map["is_aromatic"].index(atom.GetIsAromatic()))
95
+ row.append(x_map["is_in_ring"].index(atom.IsInRing()))
96
+ xs.append(row)
97
+
98
+ if len(xs) > 0:
99
+ x_np = np.array(xs, dtype=np.int64).reshape(-1, 9)
100
+ else:
101
+ x_np = np.empty((0, 9), dtype=np.int64)
102
+
103
+ edge_indices, edge_attrs = [], []
104
+ for bond in mol.GetBonds():
105
+ i = bond.GetBeginAtomIdx()
106
+ j = bond.GetEndAtomIdx()
107
+
108
+ e = []
109
+ e.append(e_map["bond_type"].index(str(bond.GetBondType())))
110
+ e.append(e_map["stereo"].index(str(bond.GetStereo())))
111
+ e.append(e_map["is_conjugated"].index(bond.GetIsConjugated()))
112
+
113
+ edge_indices += [[i, j], [j, i]]
114
+ edge_attrs += [e, e]
115
+
116
+ if len(edge_indices) > 0:
117
+ edge_index_np = np.array(edge_indices, dtype=np.int64).T.reshape(2, -1)
118
+ edge_attr_np = np.array(edge_attrs, dtype=np.int64).reshape(-1, 3)
119
+
120
+ # Sort indices matching PyG canonical ordering
121
+ perm = (edge_index_np[0] * x_np.shape[0] + edge_index_np[1]).argsort()
122
+ edge_index_np = edge_index_np[:, perm]
123
+ edge_attr_np = edge_attr_np[perm]
124
+ else:
125
+ edge_index_np = np.empty((2, 0), dtype=np.int64)
126
+ edge_attr_np = np.empty((0, 3), dtype=np.int64)
127
+
128
+ x = ops.convert_to_tensor(x_np, dtype="int64")
129
+ edge_index = ops.convert_to_tensor(edge_index_np, dtype="int64")
130
+ edge_attr = ops.convert_to_tensor(edge_attr_np, dtype="int64")
131
+
132
+ return Data(x=x, edge_index=edge_index, edge_attr=edge_attr)
133
+
134
+
135
+ def from_smiles(
136
+ smiles: str,
137
+ with_hydrogen: bool = False,
138
+ kekulize: bool = False,
139
+ ) -> Any:
140
+ r"""Converts a SMILES string to a :class:`k3_node.data.Data` instance."""
141
+ try:
142
+ from rdkit import Chem, RDLogger
143
+ except ImportError as e:
144
+ raise ImportError(
145
+ "from_smiles requires 'rdkit'. Please install it via 'pip install rdkit'."
146
+ ) from e
147
+
148
+ RDLogger.DisableLog("rdApp.*")
149
+
150
+ mol = Chem.MolFromSmiles(smiles)
151
+ if mol is None:
152
+ mol = Chem.MolFromSmiles("")
153
+ if with_hydrogen:
154
+ mol = Chem.AddHs(mol)
155
+ if kekulize:
156
+ Chem.Kekulize(mol)
157
+
158
+ data = from_rdmol(mol)
159
+ data.smiles = smiles
160
+ return data
161
+
162
+
163
+ def to_rdmol(
164
+ data: Any,
165
+ kekulize: bool = False,
166
+ ) -> Any:
167
+ r"""Converts a :class:`k3_node.data.Data` instance to an :class:`rdkit.Chem.Mol` instance."""
168
+ try:
169
+ from rdkit import Chem
170
+ except ImportError as e:
171
+ raise ImportError(
172
+ "to_rdmol requires 'rdkit'. Please install it via 'pip install rdkit'."
173
+ ) from e
174
+
175
+ mol = Chem.RWMol()
176
+
177
+ assert data.x is not None
178
+ assert data.num_nodes is not None
179
+ assert data.edge_index is not None
180
+ assert data.edge_attr is not None
181
+
182
+ x_np = ops.convert_to_numpy(data.x)
183
+ edge_index_np = ops.convert_to_numpy(data.edge_index)
184
+ edge_attr_np = ops.convert_to_numpy(data.edge_attr)
185
+
186
+ for i in range(data.num_nodes):
187
+ atom = Chem.Atom(int(x_np[i, 0]))
188
+ atom.SetChiralTag(Chem.rdchem.ChiralType.values[int(x_np[i, 1])])
189
+ atom.SetFormalCharge(x_map["formal_charge"][int(x_np[i, 3])])
190
+ atom.SetNumExplicitHs(x_map["num_hs"][int(x_np[i, 4])])
191
+ atom.SetNumRadicalElectrons(x_map["num_radical_electrons"][int(x_np[i, 5])])
192
+ atom.SetHybridization(Chem.rdchem.HybridizationType.values[int(x_np[i, 6])])
193
+ atom.SetIsAromatic(bool(x_np[i, 7]))
194
+ mol.AddAtom(atom)
195
+
196
+ edges = [tuple(edge_index_np[:, idx]) for idx in range(edge_index_np.shape[1])]
197
+ visited = set()
198
+
199
+ for idx, (src, dst) in enumerate(edges):
200
+ src, dst = int(src), int(dst)
201
+ if tuple(sorted((src, dst))) in visited:
202
+ continue
203
+
204
+ bond_type = Chem.BondType.values[int(edge_attr_np[idx, 0])]
205
+ mol.AddBond(src, dst, bond_type)
206
+
207
+ stereo = Chem.rdchem.BondStereo.values[int(edge_attr_np[idx, 1])]
208
+ if stereo != Chem.rdchem.BondStereo.STEREONONE:
209
+ db = mol.GetBondBetweenAtoms(src, dst)
210
+ db.SetStereoAtoms(dst, src)
211
+ db.SetStereo(stereo)
212
+
213
+ is_conjugated = bool(edge_attr_np[idx, 2])
214
+ mol.GetBondBetweenAtoms(src, dst).SetIsConjugated(is_conjugated)
215
+
216
+ visited.add(tuple(sorted((src, dst))))
217
+
218
+ mol = mol.GetMol()
219
+ if kekulize:
220
+ Chem.Kekulize(mol)
221
+
222
+ Chem.SanitizeMol(mol)
223
+ Chem.AssignStereochemistry(mol)
224
+ return mol
225
+
226
+
227
+ def to_smiles(
228
+ data: Any,
229
+ kekulize: bool = False,
230
+ ) -> str:
231
+ r"""Converts a :class:`k3_node.data.Data` instance to a SMILES string."""
232
+ from rdkit import Chem
233
+
234
+ mol = to_rdmol(data, kekulize=kekulize)
235
+ return Chem.MolToSmiles(mol, isomericSmiles=True)
@@ -0,0 +1,284 @@
1
+ Metadata-Version: 2.4
2
+ Name: k3-node
3
+ Version: 1.0.0
4
+ Summary: Multi-Backend Graph Neural Networks on Keras 3
5
+ Author: Muhammad Anas Raza
6
+ License: MIT
7
+ Requires-Python: >=3.11
8
+ Description-Content-Type: text/markdown
9
+ License-File: LICENSE
10
+ Requires-Dist: keras>=3.0
11
+ Requires-Dist: scipy
12
+ Requires-Dist: pynndescent
13
+ Requires-Dist: sympy
14
+ Requires-Dist: pandas
15
+ Requires-Dist: huggingface_hub>=0.20.0
16
+ Requires-Dist: onnx>=1.15.0
17
+ Requires-Dist: onnxruntime>=1.17.0
18
+ Requires-Dist: tf2onnx>=1.16.0
19
+ Requires-Dist: onnxscript
20
+ Provides-Extra: examples
21
+ Requires-Dist: scikit-learn; extra == "examples"
22
+ Requires-Dist: rdflib; extra == "examples"
23
+ Requires-Dist: matplotlib; extra == "examples"
24
+ Provides-Extra: test
25
+ Requires-Dist: pytest>=8.0.0; extra == "test"
26
+ Requires-Dist: pytest-cov; extra == "test"
27
+ Requires-Dist: torch>=2.0.0; extra == "test"
28
+ Requires-Dist: torch-geometric>=2.5.0; extra == "test"
29
+ Requires-Dist: networkx; extra == "test"
30
+ Requires-Dist: scikit-learn; extra == "test"
31
+ Requires-Dist: tqdm; extra == "test"
32
+ Requires-Dist: fsspec; extra == "test"
33
+ Requires-Dist: requests; extra == "test"
34
+ Requires-Dist: onnx>=1.15.0; extra == "test"
35
+ Requires-Dist: onnxruntime>=1.17.0; extra == "test"
36
+ Requires-Dist: tf2onnx>=1.16.0; extra == "test"
37
+ Requires-Dist: onnxscript; extra == "test"
38
+ Provides-Extra: docs
39
+ Requires-Dist: mkdocs>=1.5; extra == "docs"
40
+ Requires-Dist: mkdocs-material>=9.5; extra == "docs"
41
+ Requires-Dist: mkdocstrings[python]>=0.25; extra == "docs"
42
+ Requires-Dist: mkdocs-autorefs; extra == "docs"
43
+ Requires-Dist: pymdown-extensions; extra == "docs"
44
+ Requires-Dist: mkdocs-jupyter; extra == "docs"
45
+ Requires-Dist: pygments; extra == "docs"
46
+ Dynamic: license-file
47
+
48
+ # K3-Node: Multi-Backend Graph Neural Networks
49
+
50
+ <p align="center">
51
+ <img src="docs/images/logo.png" alt="K3-Node Logo" width="180"/>
52
+ </p>
53
+
54
+ <p align="center">
55
+ <a href="https://anas-rz.github.io/k3-node/"><img src="https://img.shields.io/badge/docs-GitHub%20Pages-blue.svg" alt="Documentation"></a>
56
+ <a href="https://github.com/anas-rz/k3-node/actions/workflows/test_torch.yml"><img src="https://github.com/anas-rz/k3-node/actions/workflows/test_torch.yml/badge.svg" alt="PyTorch tests"></a>
57
+ <a href="https://github.com/anas-rz/k3-node/actions/workflows/test_tensorflow.yml"><img src="https://github.com/anas-rz/k3-node/actions/workflows/test_tensorflow.yml/badge.svg" alt="TensorFlow tests"></a>
58
+ <a href="https://github.com/anas-rz/k3-node/actions/workflows/test_jax.yml"><img src="https://github.com/anas-rz/k3-node/actions/workflows/test_jax.yml/badge.svg" alt="JAX tests"></a>
59
+ <a href="https://github.com/anas-rz/k3-node/blob/main/LICENSE"><img src="https://img.shields.io/badge/license-MIT-green.svg" alt="License"></a>
60
+ <a href="https://keras.io/keras_3/"><img src="https://img.shields.io/badge/Keras%203-TensorFlow%20%7C%20PyTorch%20%7C%20JAX-orange.svg" alt="Backends"></a>
61
+ <a href="https://github.com/psf/black"><img src="https://img.shields.io/badge/code%20style-black-000000.svg" alt="Code style: black"></a>
62
+ </p>
63
+
64
+ ---
65
+
66
+ **K3-Node** is a next-generation graph neural network (GNN) library built natively on **Keras 3**. Write your GNN models once and execute seamlessly across **TensorFlow**, **PyTorch**, and **JAX** with full hardware acceleration (NVIDIA GPUs, Apple Silicon, Google Cloud TPUs).
67
+
68
+ K3-Node achieves **100% public API parity** with [PyTorch Geometric (PyG)](https://github.com/pyg-team/pytorch_geometric) and incorporates state-of-the-art foundation models and architectures from [Spektral](https://github.com/danielegrattarola/spektral) and [StellarGraph](https://github.com/stellargraph/stellargraph).
69
+
70
+ 📖 **Documentation**: [https://anas-rz.github.io/k3-node/](https://anas-rz.github.io/k3-node/)
71
+ 📋 **Porting Checklist & Parity Status**: [Checklist.md](Checklist.md)
72
+
73
+ ---
74
+
75
+ ## Key Features
76
+
77
+ - 🔄 **True Multi-Backend Freedom**: Switch between PyTorch, TensorFlow, and JAX with a single environment variable (`KERAS_BACKEND=torch|tensorflow|jax`).
78
+ - 🧠 **Pre-trained Foundation Models**: Out-of-the-box architectures and checkpoint loaders for **GraphMAE2**, **Graphormer** (2D & 3D), **GraphGPS**, **GROVER**, and **Mole-BERT**.
79
+ - ⚡ **65+ Convolution Layers**: Full PyG parity (`GCNConv`, `GATv2Conv`, `TransformerConv`, `GPSConv`, `PNAConv`, `SchNet`, `DimeNetPlusPlus`, `ViSNet`, etc.).
80
+ - 📊 **26 Aggregation Operators**: From elementary aggregations (`sum`, `mean`, `max`, `softmax`, `powermean`) to neural aggregations (`SetTransformer`, `GraphMultisetTransformer`, `Set2Set`, `DeepSets`, `LSTMAggregation`).
81
+ - 🌐 **31 Pooling Operators**: Global readouts (`global_add_pool`, `global_mean_pool`), hierarchical coarsening (`TopKPooling`, `SAGPooling`, `ASAPooling`, `EdgePooling`, `ClusterPooling`), and 3D spatial pooling (`voxel_grid`, `fps`, `knn`, `radius`).
82
+ - 🧱 **Dense & Scalable GNNs**: Dense matrix convolutions (`DenseGCNConv`, `DenseGATConv`), spectral pooling (`DMoNPooling`, `dense_diff_pool`, `dense_mincut_pool`), and linear-complexity graph transformers (`SGFormer`, `LPFormer`, `Polynormer`).
83
+ - 🧭 **Knowledge Graph Embeddings**: Multi-relational link prediction with `TransE`, `RotatE`, `DistMult`, `ComplEx`, and framework-agnostic negative sampling loaders.
84
+ - 📦 **Data, Loaders & Transforms**: Full suite of graph data structures (`Data`, `HeteroData`, `Batch`), mini-batch samplers (`NeighborLoader`, `ClusterLoader`, `GraphSAINTSampler`), and 62+ graph and 3D point cloud transforms.
85
+ - ✅ **Rigorous Verification**: 700+ unit tests on every backend, training tests that check each layer's weights actually learn, compiled-vs-eager and cross-backend consistency tests, and numerical parity tests against PyTorch Geometric and reference checkpoints.
86
+
87
+ ---
88
+
89
+ ## Installation
90
+
91
+ ```bash
92
+ # git should be installed
93
+ pip install git+https://github.com/anas-rz/k3-node/
94
+
95
+ # with the extra packages the example notebooks use (scikit-learn, rdflib, matplotlib)
96
+ pip install "k3-node[examples] @ git+https://github.com/anas-rz/k3-node"
97
+ ```
98
+
99
+ ### Selecting your Backend
100
+ Configure your preferred backend before importing `k3_node`:
101
+
102
+ ```bash
103
+ export KERAS_BACKEND="torch" # or "tensorflow" or "jax"
104
+ ```
105
+
106
+ Or programmatically in Python:
107
+
108
+ ```python
109
+ import os
110
+ os.environ["KERAS_BACKEND"] = "torch" # Must be set before importing k3_node / keras
111
+ import k3_node
112
+ ```
113
+
114
+ ---
115
+
116
+ ## Quickstart
117
+
118
+ ### Building a Graph Convolutional Network
119
+
120
+ ```python
121
+ import keras
122
+ from keras import ops
123
+ import k3_node.layers as gnn_layers
124
+ from k3_node.data import Data
125
+
126
+ class GCN(keras.Model):
127
+ def __init__(self, in_channels, hidden_channels, out_channels):
128
+ super().__init__()
129
+ self.conv1 = gnn_layers.GCNConv(in_channels, hidden_channels)
130
+ self.conv2 = gnn_layers.GCNConv(hidden_channels, out_channels)
131
+
132
+ def call(self, x, edge_index):
133
+ x = self.conv1(x, edge_index)
134
+ x = ops.relu(x)
135
+ x = self.conv2(x, edge_index)
136
+ return x
137
+
138
+ # Instantiate model
139
+ model = GCN(in_channels=16, hidden_channels=32, out_channels=7)
140
+
141
+ # Forward pass on graph data
142
+ x = ops.ones((10, 16))
143
+ edge_index = ops.convert_to_tensor([[0, 1, 2, 3], [1, 2, 3, 0]], dtype="int64")
144
+
145
+ out = model(x, edge_index)
146
+ print("Output shape:", out.shape) # (10, 7)
147
+ ```
148
+
149
+ ### Training in a Few Lines
150
+
151
+ The task estimators in `k3_node.tasks` pick the loss, readout and metrics for you:
152
+
153
+ ```python
154
+ from k3_node.datasets import Planetoid
155
+ from k3_node.tasks import NodeClassifier
156
+
157
+ cora = Planetoid("data/Planetoid", name="Cora")[0]
158
+
159
+ classifier = NodeClassifier(backbone="gcn", hidden_channels=64, num_layers=2, dropout=0.5)
160
+ classifier.fit(cora, epochs=100, lr=0.01)
161
+ print(classifier.evaluate(cora, mask="test_mask"))
162
+ ```
163
+
164
+ `GraphClassifier`, `GraphRegressor`, `NodeRegressor` and `LinkPredictor` work the same way.
165
+
166
+ ### Example Notebooks
167
+
168
+ The [`examples/`](examples) folder has 90+ notebooks that follow the architectures of
169
+ [PyG's examples](https://github.com/pyg-team/pytorch_geometric/tree/master/examples), written with
170
+ `keras.Model.fit` and K3-Node's loaders. They cover node, link and graph classification,
171
+ knowledge graphs, molecules (including pre-trained DimeNet, DimeNet++ and SchNet on QM9), point
172
+ clouds, temporal graphs and large-graph mini-batching. Each notebook opens in Colab and runs on
173
+ any backend: change `KERAS_BACKEND` in its first cell. Browse them in the
174
+ [documentation](https://anas-rz.github.io/k3-node/examples/).
175
+
176
+ ---
177
+
178
+ ## Pre-trained Foundation Models
179
+
180
+ K3-Node provides ready-to-use architectures and automated checkpoint loading for state-of-the-art graph foundation models:
181
+
182
+ ### 1. GraphMAE2 (Self-Supervised Masked Autoencoder)
183
+ ```python
184
+ from k3_node.models import GraphMAE2
185
+ from k3_node.models.graphmae2 import load_graphmae2_weights
186
+
187
+ model = GraphMAE2(
188
+ in_dim=100,
189
+ num_hidden=512,
190
+ out_dim=100,
191
+ num_layers=4,
192
+ encoder_type="gat",
193
+ decoder_type="gat"
194
+ )
195
+ # Load reference pre-trained weights
196
+ load_graphmae2_weights(model, "checkpoints/graphmae2_ogbn_arxiv.pt")
197
+ ```
198
+
199
+ ### 2. Graphormer (2D Molecular & 3D Structural Transformer)
200
+ ```python
201
+ from k3_node.models import Graphormer, Graphormer3D
202
+ from k3_node.models.graphormer import load_graphormer_weights
203
+
204
+ # 2D Graphormer (PCQM4Mv2)
205
+ model_2d = Graphormer(num_layers=12, num_heads=32, embed_dim=768)
206
+ load_graphormer_weights(model_2d, "checkpoints/graphormer_pcqm4mv2.pt")
207
+
208
+ # 3D Graphormer (OC20 Catalyst Adsorption & Molecular Conformations)
209
+ model_3d = Graphormer3D(num_layers=12, num_heads=32, embed_dim=768)
210
+ ```
211
+
212
+ ### 3. GraphGPS (Hybrid Local MPNN + Global Transformer)
213
+ ```python
214
+ from k3_node.models import GPSModel
215
+ from k3_node.models.gps_model import load_gps_model_weights
216
+
217
+ model = GPSModel(
218
+ channels=64,
219
+ num_layers=5,
220
+ local_gnn_type="GINE",
221
+ global_model_type="Transformer"
222
+ )
223
+ load_gps_model_weights(model, "checkpoints/graphgps_zinc.pt")
224
+ ```
225
+
226
+ ### 4. GROVER (Self-Supervised Message Passing Transformer)
227
+ ```python
228
+ from k3_node.models import GROVER, GROVEREmbedding
229
+ from k3_node.models.grover import load_grover_weights
230
+
231
+ model = GROVER(hidden_size=128, num_layers=3, num_heads=4)
232
+ load_grover_weights(model, "checkpoints/grover_base.pt")
233
+ ```
234
+
235
+ ### 5. Mole-BERT (Masked Chemical Graph Representation)
236
+ ```python
237
+ from k3_node.models import MoleBERT
238
+ from k3_node.models.mole_bert import load_mole_bert_weights
239
+
240
+ model = MoleBERT(num_layer=5, emb_dim=300, drop_ratio=0.5)
241
+ load_mole_bert_weights(model, "checkpoints/Mole-BERT.pth")
242
+ ```
243
+
244
+ ---
245
+
246
+ ## What's Included
247
+
248
+ | Package | Status | Contents |
249
+ |---|---|---|
250
+ | [`k3_node.layers.conv`](https://anas-rz.github.io/k3-node/api/conv/) | ✅ 65/65 | `GCNConv`, `GATConv`, `GATv2Conv`, `SAGEConv`, `GINConv`, `GPSConv`, `TransformerConv`, `PNAConv`, `SchNet`, `DimeNetPlusPlus`, `ViSNet`, etc. |
251
+ | [`k3_node.layers.pool`](https://anas-rz.github.io/k3-node/api/pool/) | ✅ 31/31 | `global_add_pool`, `global_mean_pool`, `TopKPooling`, `SAGPooling`, `ASAPooling`, `EdgePooling`, `ClusterPooling`, `voxel_grid`, `fps`, `graclus`, etc. |
252
+ | [`k3_node.layers.aggr`](https://anas-rz.github.io/k3-node/api/aggr/) | ✅ 26/26 | `SumAggregation`, `MeanAggregation`, `SoftmaxAggregation`, `PowerMeanAggregation`, `MultiAggregation`, `SetTransformerAggregation`, `Set2Set`, etc. |
253
+ | [`k3_node.layers.norm`](https://anas-rz.github.io/k3-node/api/norm/) | ✅ 11/11 | `GraphNorm`, `PairNorm`, `DiffGroupNorm`, `MessageNorm`, `MeanSubtractionNorm`, `BatchNorm`, `LayerNorm`, `HeteroBatchNorm`, etc. |
254
+ | [`k3_node.layers.dense`](https://anas-rz.github.io/k3-node/api/dense/) | ✅ 11/11 | `DenseGCNConv`, `DenseGATConv`, `DenseGINConv`, `DenseSAGEConv`, `DMoNPooling`, `dense_diff_pool`, `dense_mincut_pool`, `Linear`, etc. |
255
+ | [`k3_node.layers.kge`](https://anas-rz.github.io/k3-node/api/kge/) | ✅ 5/5 | `KGEModel`, `TransE`, `RotatE`, `DistMult`, `ComplEx`, `KGTripletLoader`. |
256
+ | [`k3_node.models`](https://anas-rz.github.io/k3-node/api/models/) | ✅ 46/46 | `MLP`, `GAE`, `VGAE`, `DeepGraphInfomax`, `Node2Vec`, `LabelPropagation`, `LINKX`, `LightGCN`, `SGFormer`, `LPFormer`, `Polynormer`, etc. |
257
+ | **Foundation Models** | ✅ 5/5 | `GraphMAE2`, `Graphormer` (2D/3D), `GPSModel`, `GROVER`, `MoleBERT` with pre-trained weight conversion. |
258
+ | [`k3_node.data`](https://anas-rz.github.io/k3-node/api/data/) | ✅ 19/19 | `Data`, `HeteroData`, `Batch`, `TemporalData`, `HypergraphData`, `InMemoryDataset`, `FeatureStore`, `GraphStore`, etc. |
259
+ | [`k3_node.loader`](https://anas-rz.github.io/k3-node/api/loader/) | ✅ 26/26 | `DataLoader`, `NeighborLoader`, `LinkNeighborLoader`, `ClusterLoader`, `GraphSAINTSampler`, `ShaDowKHopSampler`, etc. |
260
+ | [`k3_node.transforms`](https://anas-rz.github.io/k3-node/api/transforms/) | ✅ 62/62 | Topology rewiring, positional encodings (`LapPE`, `RWPE`, `GPSE`), spectral diffusion (`GDC`), and 3D point cloud transforms. |
261
+
262
+ ---
263
+
264
+
265
+ ## Testing & Verification
266
+
267
+ Run the comprehensive test suite across backends:
268
+
269
+ ```bash
270
+ # Run all unit tests
271
+ pytest k3_node/
272
+
273
+ # Run training tests (each layer's weights learn; slower, not run in CI)
274
+ pytest tests_training/
275
+
276
+ # Run reference parity check against PyTorch implementations
277
+ pytest tests_reference/
278
+ ```
279
+
280
+ ---
281
+
282
+ ## License
283
+
284
+ This project is licensed under the MIT License - see the [LICENSE](LICENSE) file for details.