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,159 @@
1
+ import pytest
2
+ import numpy as np
3
+ import keras
4
+ from keras import ops
5
+
6
+ from k3_node.models import MetaLayer
7
+ from k3_node.layers.conv.utils import scatter
8
+
9
+
10
+ # Global counter to verify model call counts
11
+ _count = 0
12
+
13
+
14
+ def test_meta_layer_repr():
15
+ """MetaLayer should have the correct string representation."""
16
+ assert str(MetaLayer()) == (
17
+ 'MetaLayer(\n'
18
+ ' edge_model=None,\n'
19
+ ' node_model=None,\n'
20
+ ' global_model=None\n'
21
+ ')'
22
+ )
23
+
24
+
25
+ def test_meta_layer_none_models():
26
+ """All None models should return unchanged tensors."""
27
+ x = ops.ones((20, 10))
28
+ edge_index = ops.convert_to_tensor([[0, 1, 2], [1, 2, 0]])
29
+
30
+ model = MetaLayer()
31
+ x_out, edge_attr_out, u_out = model(x, edge_index)
32
+
33
+ assert ops.shape(x_out) == (20, 10)
34
+ assert edge_attr_out is None
35
+ assert u_out is None
36
+
37
+
38
+ def test_meta_layer_call_counting():
39
+ """Verify that the correct submodels are called for each combination."""
40
+ global _count
41
+ _count = 0
42
+
43
+ def dummy_model(*args):
44
+ global _count
45
+ _count += 1
46
+ return None
47
+
48
+ x = ops.ones((20, 10))
49
+ edge_index = ops.convert_to_tensor([[0, 1, 2, 3], [1, 2, 3, 0]])
50
+
51
+ for edge_model in (dummy_model, None):
52
+ for node_model in (dummy_model, None):
53
+ for global_model in (dummy_model, None):
54
+ model = MetaLayer(edge_model, node_model, global_model)
55
+ out = model(x, edge_index)
56
+ assert isinstance(out, tuple) and len(out) == 3
57
+
58
+ assert _count == 12 # 3 models × 4 combinations with at least 1 non-None
59
+
60
+
61
+ def test_meta_layer_with_edge_model():
62
+ """Edge model should receive correct src/dst/edge_attr/u/batch args."""
63
+ received = {}
64
+
65
+ def edge_model(src, dst, edge_attr, u, batch):
66
+ received['src_shape'] = ops.shape(src)
67
+ received['dst_shape'] = ops.shape(dst)
68
+ received['edge_attr_shape'] = ops.shape(edge_attr) if edge_attr is not None else None
69
+ return edge_attr # pass-through
70
+
71
+ x = ops.ones((5, 4))
72
+ edge_index = ops.convert_to_tensor([[0, 1, 2], [1, 2, 3]])
73
+ edge_attr = ops.zeros((3, 7))
74
+
75
+ model = MetaLayer(edge_model=edge_model)
76
+ model(x, edge_index, edge_attr=edge_attr)
77
+
78
+ assert received['src_shape'] == (3, 4)
79
+ assert received['dst_shape'] == (3, 4)
80
+ assert received['edge_attr_shape'] == (3, 7)
81
+
82
+
83
+ def test_meta_layer_full_example():
84
+ """Full graph network test matching PyG's test_meta_layer_example."""
85
+
86
+ class EdgeModel(keras.layers.Layer):
87
+ def __init__(self):
88
+ super().__init__()
89
+ self.mlp = keras.Sequential([
90
+ keras.layers.Dense(5),
91
+ keras.layers.ReLU(),
92
+ keras.layers.Dense(5),
93
+ ])
94
+
95
+ def call(self, src, dst, edge_attr, u, batch):
96
+ assert edge_attr is not None
97
+ assert u is not None
98
+ assert batch is not None
99
+ out = ops.concatenate([src, dst, edge_attr, ops.take(u, batch, axis=0)], axis=1)
100
+ return self.mlp(out)
101
+
102
+ class NodeModel(keras.layers.Layer):
103
+ def __init__(self):
104
+ super().__init__()
105
+ self.mlp1 = keras.Sequential([
106
+ keras.layers.Dense(10),
107
+ keras.layers.ReLU(),
108
+ keras.layers.Dense(10),
109
+ ])
110
+ self.mlp2 = keras.Sequential([
111
+ keras.layers.Dense(10),
112
+ keras.layers.ReLU(),
113
+ keras.layers.Dense(10),
114
+ ])
115
+
116
+ def call(self, x, edge_index, edge_attr, u, batch):
117
+ assert edge_attr is not None
118
+ assert u is not None
119
+ assert batch is not None
120
+ row = edge_index[0]
121
+ col = edge_index[1]
122
+ out = ops.concatenate([ops.take(x, row, axis=0), edge_attr], axis=1)
123
+ out = self.mlp1(out)
124
+ out = scatter(out, col, dim_size=ops.shape(x)[0], reduce="mean")
125
+ out = ops.concatenate([x, out, ops.take(u, batch, axis=0)], axis=1)
126
+ return self.mlp2(out)
127
+
128
+ class GlobalModel(keras.layers.Layer):
129
+ def __init__(self):
130
+ super().__init__()
131
+ self.mlp = keras.Sequential([
132
+ keras.layers.Dense(20),
133
+ keras.layers.ReLU(),
134
+ keras.layers.Dense(20),
135
+ ])
136
+
137
+ def call(self, x, edge_index, edge_attr, u, batch):
138
+ assert u is not None
139
+ assert batch is not None
140
+ batch_size = ops.shape(u)[0]
141
+ x_mean = scatter(x, batch, dim_size=batch_size, reduce="mean")
142
+ out = ops.concatenate([u, x_mean], axis=1)
143
+ return self.mlp(out)
144
+
145
+ op = MetaLayer(EdgeModel(), NodeModel(), GlobalModel())
146
+
147
+ x = ops.ones((20, 10))
148
+ edge_attr = ops.ones((40, 5))
149
+ u = ops.ones((2, 20))
150
+ batch = ops.convert_to_tensor([0] * 10 + [1] * 10)
151
+ row_idx = list(range(20)) + list(range(20))
152
+ col_idx = list(range(1, 20)) + [0] + list(range(1, 20)) + [0]
153
+ edge_index = ops.convert_to_tensor([row_idx, col_idx])
154
+
155
+ x_out, edge_attr_out, u_out = op(x, edge_index, edge_attr, u, batch)
156
+ assert ops.shape(x_out) == (20, 10)
157
+ assert ops.shape(edge_attr_out) == (40, 5)
158
+ assert ops.shape(u_out) == (2, 20)
159
+
@@ -0,0 +1,45 @@
1
+ from keras import ops
2
+ from k3_node.models import MetaPath2Vec
3
+
4
+
5
+ def test_metapath2vec():
6
+ edge_index_dict = {
7
+ ("author", "writes", "paper"): ops.convert_to_tensor([[0, 1, 1], [0, 0, 1]], dtype="int64"),
8
+ ("paper", "written_by", "author"): ops.convert_to_tensor([[0, 0, 1], [0, 1, 1]], dtype="int64"),
9
+ }
10
+ metapath = [
11
+ ("author", "writes", "paper"),
12
+ ("paper", "written_by", "author"),
13
+ ]
14
+
15
+ model = MetaPath2Vec(
16
+ edge_index_dict=edge_index_dict,
17
+ embedding_dim=16,
18
+ metapath=metapath,
19
+ walk_length=2,
20
+ context_size=2,
21
+ walks_per_node=2,
22
+ )
23
+
24
+ # Test forward
25
+ out_author = model("author")
26
+ assert out_author.shape == (2, 16)
27
+
28
+ out_paper = model("paper")
29
+ assert out_paper.shape == (2, 16)
30
+
31
+ batch = ops.convert_to_tensor([0], dtype="int64")
32
+ out_batch = model("author", batch)
33
+ assert out_batch.shape == (1, 16)
34
+
35
+ # Test sampling
36
+ pos_rw = model._pos_sample(batch)
37
+ neg_rw = model._neg_sample(batch)
38
+ assert pos_rw.shape[1] == 2
39
+ assert neg_rw.shape[1] == 2
40
+
41
+ # Test loss
42
+ loss = model.loss(pos_rw, neg_rw)
43
+ assert loss.shape == ()
44
+ assert float(ops.convert_to_numpy(loss)) > 0
45
+
@@ -0,0 +1,62 @@
1
+ import pytest
2
+ from keras import ops
3
+
4
+ from k3_node.models import MLP
5
+
6
+
7
+ @pytest.mark.parametrize("norm", ["batch_norm", None])
8
+ @pytest.mark.parametrize("act_first", [False, True])
9
+ @pytest.mark.parametrize("plain_last", [False, True])
10
+ def test_mlp(norm, act_first, plain_last):
11
+ x = ops.ones((4, 16))
12
+
13
+ mlp = MLP([16, 32, 32, 64], norm=norm, act_first=act_first, plain_last=plain_last)
14
+ assert str(mlp) == "MLP(16, 32, 32, 64)"
15
+ out = mlp(x)
16
+ assert ops.shape(out) == (4, 64)
17
+
18
+ mlp2 = MLP(
19
+ 16,
20
+ hidden_channels=32,
21
+ out_channels=64,
22
+ num_layers=3,
23
+ norm=norm,
24
+ act_first=act_first,
25
+ plain_last=plain_last,
26
+ )
27
+ assert ops.shape(mlp2(x)) == (4, 64)
28
+
29
+
30
+ @pytest.mark.parametrize("norm", ["BatchNorm", "GraphNorm", "InstanceNorm", "LayerNorm"])
31
+ def test_batch(norm):
32
+ x = ops.ones((3, 8))
33
+ batch = ops.convert_to_tensor([0, 0, 1], dtype="int64")
34
+
35
+ model = MLP(8, hidden_channels=16, out_channels=32, num_layers=2, norm=norm)
36
+ assert model.supports_norm_batch == (norm != "BatchNorm")
37
+
38
+ out = model(x, batch=batch)
39
+ assert ops.shape(out) == (3, 32)
40
+
41
+
42
+ def test_mlp_return_emb():
43
+ x = ops.ones((4, 16))
44
+
45
+ mlp = MLP([16, 32, 1])
46
+
47
+ out, emb = mlp(x, return_emb=True)
48
+ assert ops.shape(out) == (4, 1)
49
+ assert ops.shape(emb) == (4, 32)
50
+
51
+ out, emb = mlp(x, return_emb=False)
52
+ assert ops.shape(out) == (4, 1)
53
+ assert emb is None
54
+
55
+ out = mlp(x)
56
+ assert ops.shape(out) == (4, 1)
57
+
58
+
59
+ @pytest.mark.parametrize("plain_last", [False, True])
60
+ def test_fine_grained_mlp(plain_last):
61
+ mlp = MLP([16, 32, 32, 64], dropout=[0.1, 0.2, 0.3], bias=[False, True, False], plain_last=plain_last)
62
+ assert ops.shape(mlp(ops.ones((4, 16)))) == (4, 64)
@@ -0,0 +1,164 @@
1
+ import os
2
+ import numpy as np
3
+ import pytest
4
+ import keras.ops as ops
5
+ try:
6
+ import torch
7
+ except ImportError:
8
+ torch = None
9
+
10
+ from k3_node.models.mole_bert import (
11
+ MoleBERT,
12
+ MoleBERTGNN,
13
+ MoleBERTGINConv,
14
+ load_mole_bert_weights,
15
+ download_mole_bert_checkpoint,
16
+ )
17
+
18
+
19
+ def _make_dummy_graph(num_nodes=5, num_edges=8, emb_dim=300):
20
+ np.random.seed(42)
21
+ x = np.stack(
22
+ [
23
+ np.random.randint(0, 119, size=num_nodes),
24
+ np.random.randint(0, 3, size=num_nodes),
25
+ ],
26
+ axis=1,
27
+ )
28
+ src = np.random.randint(0, num_nodes, size=num_edges)
29
+ dst = np.random.randint(0, num_nodes, size=num_edges)
30
+ edge_index = np.stack([src, dst], axis=0)
31
+ edge_attr = np.stack(
32
+ [
33
+ np.random.randint(0, 5, size=num_edges),
34
+ np.random.randint(0, 3, size=num_edges),
35
+ ],
36
+ axis=1,
37
+ )
38
+ batch = np.array([0, 0, 0, 1, 1], dtype=np.int64)
39
+
40
+ return (
41
+ ops.convert_to_tensor(x, dtype="int64"),
42
+ ops.convert_to_tensor(edge_index, dtype="int64"),
43
+ ops.convert_to_tensor(edge_attr, dtype="int64"),
44
+ ops.convert_to_tensor(batch, dtype="int64"),
45
+ )
46
+
47
+
48
+ def test_mole_bert_gin_conv():
49
+ emb_dim = 64
50
+ conv = MoleBERTGINConv(emb_dim=emb_dim)
51
+ conv.build(None)
52
+
53
+ num_nodes = 4
54
+ x = ops.convert_to_tensor(np.random.randn(num_nodes, emb_dim).astype(np.float32))
55
+ edge_index = ops.convert_to_tensor(
56
+ np.array([[0, 1, 1, 2], [1, 0, 2, 1]]), dtype="int64"
57
+ )
58
+ edge_attr = ops.convert_to_tensor(
59
+ np.array([[0, 0], [0, 0], [1, 0], [1, 0]]), dtype="int64"
60
+ )
61
+
62
+ out = conv(x, edge_index, edge_attr)
63
+ assert ops.shape(out) == (num_nodes, emb_dim)
64
+
65
+
66
+ def test_mole_bert_gnn_jk_modes():
67
+ x, edge_index, edge_attr, _ = _make_dummy_graph(num_nodes=5, num_edges=8, emb_dim=32)
68
+ emb_dim = 32
69
+
70
+ # 1. JK = 'last'
71
+ gnn_last = MoleBERTGNN(num_layer=3, emb_dim=emb_dim, JK="last")
72
+ gnn_last.build(None)
73
+ out_last = gnn_last(x, edge_index, edge_attr)
74
+ assert ops.shape(out_last) == (5, emb_dim)
75
+
76
+ # 2. JK = 'concat'
77
+ gnn_concat = MoleBERTGNN(num_layer=3, emb_dim=emb_dim, JK="concat")
78
+ gnn_concat.build(None)
79
+ out_concat = gnn_concat(x, edge_index, edge_attr)
80
+ assert ops.shape(out_concat) == (5, 4 * emb_dim)
81
+
82
+ # 3. JK = 'sum'
83
+ gnn_sum = MoleBERTGNN(num_layer=3, emb_dim=emb_dim, JK="sum")
84
+ gnn_sum.build(None)
85
+ out_sum = gnn_sum(x, edge_index, edge_attr)
86
+ assert ops.shape(out_sum) == (5, emb_dim)
87
+
88
+ # 4. JK = 'max'
89
+ gnn_max = MoleBERTGNN(num_layer=3, emb_dim=emb_dim, JK="max")
90
+ gnn_max.build(None)
91
+ out_max = gnn_max(x, edge_index, edge_attr)
92
+ assert ops.shape(out_max) == (5, emb_dim)
93
+
94
+
95
+ def test_mole_bert_pooling_and_prediction():
96
+ x, edge_index, edge_attr, batch = _make_dummy_graph(num_nodes=5, num_edges=8, emb_dim=32)
97
+ emb_dim = 32
98
+
99
+ # Multi-task graph classification (e.g. ClinTox: 2 tasks, Tox21: 12 tasks)
100
+ for pooling in ["mean", "sum", "max"]:
101
+ model = MoleBERT(
102
+ num_layer=3,
103
+ emb_dim=emb_dim,
104
+ num_tasks=2,
105
+ JK="last",
106
+ graph_pooling=pooling,
107
+ )
108
+ model.build(None)
109
+
110
+ logits, node_rep = model((x, edge_index, edge_attr, batch))
111
+ assert ops.shape(logits) == (2, 2)
112
+ assert ops.shape(node_rep) == (5, emb_dim)
113
+
114
+
115
+ def test_mole_bert_synthetic_checkpoint_load(tmp_path):
116
+ if torch is None:
117
+ pytest.skip("PyTorch is required for checkpoint loading test")
118
+ emb_dim = 16
119
+ num_layer = 3
120
+ model = MoleBERT(num_layer=num_layer, emb_dim=emb_dim, num_tasks=1)
121
+ model.build(None)
122
+
123
+ state_dict = {
124
+ "x_embedding1.weight": torch.randn(120, emb_dim),
125
+ "x_embedding2.weight": torch.randn(3, emb_dim),
126
+ }
127
+ for l in range(num_layer):
128
+ state_dict[f"gnns.{l}.edge_embedding1.weight"] = torch.randn(6, emb_dim)
129
+ state_dict[f"gnns.{l}.edge_embedding2.weight"] = torch.randn(3, emb_dim)
130
+ state_dict[f"gnns.{l}.mlp.0.weight"] = torch.randn(2 * emb_dim, emb_dim)
131
+ state_dict[f"gnns.{l}.mlp.0.bias"] = torch.zeros(2 * emb_dim)
132
+ state_dict[f"gnns.{l}.mlp.2.weight"] = torch.randn(emb_dim, 2 * emb_dim)
133
+ state_dict[f"gnns.{l}.mlp.2.bias"] = torch.zeros(emb_dim)
134
+ state_dict[f"batch_norms.{l}.weight"] = torch.ones(emb_dim)
135
+ state_dict[f"batch_norms.{l}.bias"] = torch.zeros(emb_dim)
136
+ state_dict[f"batch_norms.{l}.running_mean"] = torch.zeros(emb_dim)
137
+ state_dict[f"batch_norms.{l}.running_var"] = torch.ones(emb_dim)
138
+
139
+ ckpt_file = str(tmp_path / "synthetic_mole_bert.pth")
140
+ torch.save(state_dict, ckpt_file)
141
+
142
+ load_mole_bert_weights(model, ckpt_file)
143
+
144
+ x, edge_index, edge_attr, batch = _make_dummy_graph(num_nodes=5, num_edges=8, emb_dim=emb_dim)
145
+ logits, node_rep = model((x, edge_index, edge_attr, batch))
146
+ assert ops.shape(logits) == (2, 1)
147
+ assert ops.shape(node_rep) == (5, emb_dim)
148
+
149
+
150
+ def test_mole_bert_official_checkpoint_load():
151
+ torch = pytest.importorskip("torch")
152
+ ckpt_path = "Mole-BERT/model_gin/Mole-BERT.pth"
153
+ if not os.path.exists(ckpt_path):
154
+ ckpt_path = download_mole_bert_checkpoint()
155
+
156
+ model = MoleBERT(num_layer=5, emb_dim=300, num_tasks=1)
157
+ model.build(None)
158
+ load_mole_bert_weights(model, ckpt_path)
159
+
160
+ x, edge_index, edge_attr, batch = _make_dummy_graph(num_nodes=5, num_edges=8, emb_dim=300)
161
+ logits, node_rep = model((x, edge_index, edge_attr, batch))
162
+ assert ops.shape(logits) == (2, 1)
163
+ assert ops.shape(node_rep) == (5, 300)
164
+
@@ -0,0 +1,13 @@
1
+ from keras import ops, random
2
+ from k3_node.models import NeuralFingerprint
3
+
4
+
5
+ def test_neural_fingerprint():
6
+ model = NeuralFingerprint(in_channels=16, hidden_channels=32, out_channels=8, num_layers=3)
7
+ x = random.normal((6, 16))
8
+ edge_index = ops.convert_to_tensor([[0, 1, 2, 3, 4], [1, 2, 0, 4, 5]], dtype="int64")
9
+ batch = ops.convert_to_tensor([0, 0, 0, 1, 1, 1], dtype="int64")
10
+
11
+ out = model(x, edge_index, batch)
12
+ assert out.shape == (2, 8)
13
+
@@ -0,0 +1,57 @@
1
+ from keras import ops
2
+ from k3_node.models import Node2Vec
3
+
4
+
5
+ def test_node2vec():
6
+ edge_index = ops.convert_to_tensor([
7
+ [0, 1, 2, 3, 0, 2],
8
+ [1, 2, 3, 0, 2, 0],
9
+ ], dtype="int64")
10
+
11
+ model = Node2Vec(
12
+ edge_index,
13
+ embedding_dim=16,
14
+ walk_length=4,
15
+ context_size=3,
16
+ walks_per_node=2,
17
+ num_negative_samples=1,
18
+ )
19
+
20
+ # Test forward
21
+ out = model()
22
+ assert out.shape == (4, 16)
23
+
24
+ batch = ops.convert_to_tensor([0, 1], dtype="int64")
25
+ out_batch = model(batch)
26
+ assert out_batch.shape == (2, 16)
27
+
28
+ # Test pos and neg sampling
29
+ pos_rw = model.pos_sample(batch)
30
+ assert pos_rw.shape[1] == 3
31
+
32
+ neg_rw = model.neg_sample(batch)
33
+ assert neg_rw.shape[1] == 3
34
+
35
+ # Test loss
36
+ loss = model.loss(pos_rw, neg_rw)
37
+ assert loss.shape == ()
38
+ assert float(ops.convert_to_numpy(loss)) > 0
39
+
40
+
41
+
42
+ def test_node2vec_fit_and_biased_walks():
43
+ import numpy as np
44
+ from k3_node.models import Node2Vec
45
+
46
+ rng = np.random.default_rng(0)
47
+ edge_index = rng.integers(0, 20, size=(2, 80))
48
+ edge_index = np.concatenate([edge_index, edge_index[::-1]], axis=1)
49
+ import keras
50
+ model = Node2Vec(edge_index, embedding_dim=8, walk_length=5, context_size=3, walks_per_node=2, num_nodes=20)
51
+ model.compile(keras.optimizers.Adam(0.01))
52
+ history = model.fit(epochs=2, batch_size=8, verbose=0)
53
+ assert len(history["loss"]) == 2 and np.isfinite(history["loss"]).all()
54
+ # With a tiny p the walk (almost) always steps back to where it came from
55
+ model.p, model.q = 1e-6, 1.0
56
+ walk = model._random_walk(int(edge_index[0, 0]))
57
+ assert all(walk[i] == walk[i - 2] for i in range(2, len(walk)))
@@ -0,0 +1,81 @@
1
+ import pytest
2
+ from keras import ops
3
+
4
+ from k3_node.models import PMLP
5
+
6
+
7
+ def test_pmlp():
8
+ x = ops.ones((4, 16))
9
+ edge_index = ops.convert_to_tensor([[0, 1, 1, 2], [1, 0, 2, 1]])
10
+
11
+ pmlp = PMLP(in_channels=16, hidden_channels=32, out_channels=2, num_layers=4)
12
+ assert str(pmlp) == 'PMLP(16, 2, num_layers=4)'
13
+
14
+ pmlp.training = True
15
+ out = pmlp(x)
16
+ assert ops.shape(out) == (4, 2)
17
+
18
+ pmlp.training = False
19
+ out = pmlp(x, edge_index)
20
+ assert ops.shape(out) == (4, 2)
21
+
22
+
23
+ def test_pmlp_raises_without_edge_index():
24
+ """Should raise ValueError when edge_index is missing during inference."""
25
+ x = ops.ones((4, 16))
26
+ pmlp = PMLP(in_channels=16, hidden_channels=32, out_channels=2, num_layers=4)
27
+
28
+ with pytest.raises(ValueError, match="'edge_index' needs to be present"):
29
+ pmlp.training = False
30
+ pmlp(x)
31
+
32
+
33
+ def test_pmlp_call_training_override():
34
+ """The training kwarg at call-time should override self.training."""
35
+ x = ops.ones((4, 16))
36
+ edge_index = ops.convert_to_tensor([[0, 1], [1, 0]])
37
+ pmlp = PMLP(in_channels=16, hidden_channels=32, out_channels=4, num_layers=2)
38
+
39
+ # Instance says training=True but call says training=False (requires edge_index)
40
+ pmlp.training = True
41
+ out = pmlp(x, edge_index, training=False)
42
+ assert ops.shape(out) == (4, 4)
43
+
44
+ # Instance says training=False but call says training=True (no edge_index needed)
45
+ pmlp.training = False
46
+ out = pmlp(x, training=True)
47
+ assert ops.shape(out) == (4, 4)
48
+
49
+
50
+ @pytest.mark.parametrize("norm", [True, False])
51
+ @pytest.mark.parametrize("bias", [True, False])
52
+ @pytest.mark.parametrize("dropout", [0.0, 0.3])
53
+ def test_pmlp_configs(norm, bias, dropout):
54
+ x = ops.ones((8, 10))
55
+ pmlp = PMLP(
56
+ in_channels=10,
57
+ hidden_channels=20,
58
+ out_channels=5,
59
+ num_layers=3,
60
+ norm=norm,
61
+ bias=bias,
62
+ dropout=dropout,
63
+ )
64
+ pmlp.training = True
65
+ out = pmlp(x)
66
+ assert ops.shape(out) == (8, 5)
67
+
68
+
69
+ def test_pmlp_single_layer():
70
+ """Edge case: num_layers=1 should work (no hidden layers, direct in→out)."""
71
+ x = ops.ones((4, 8))
72
+ edge_index = ops.convert_to_tensor([[0, 1], [1, 0]])
73
+ pmlp = PMLP(in_channels=8, hidden_channels=16, out_channels=3, num_layers=1)
74
+ pmlp.training = True
75
+ out = pmlp(x)
76
+ assert ops.shape(out) == (4, 3)
77
+
78
+ pmlp.training = False
79
+ out = pmlp(x, edge_index)
80
+ assert ops.shape(out) == (4, 3)
81
+
@@ -0,0 +1,104 @@
1
+ import pytest
2
+ from keras import ops
3
+
4
+ from k3_node.models import Polynormer
5
+
6
+
7
+ @pytest.mark.parametrize('local_attn', [True, False])
8
+ @pytest.mark.parametrize('qk_shared', [True, False])
9
+ @pytest.mark.parametrize('pre_ln', [True, False])
10
+ @pytest.mark.parametrize('post_bn', [True, False])
11
+ def test_polynormer_local(local_attn, qk_shared, pre_ln, post_bn):
12
+ """Local-mode Polynormer (default _global=False) should output [N, out_channels]."""
13
+ x = ops.ones((10, 16))
14
+ edge_index = ops.convert_to_tensor([
15
+ [0, 1, 2, 3, 4, 5, 6, 7, 8, 9],
16
+ [1, 2, 3, 4, 0, 6, 7, 8, 9, 5],
17
+ ])
18
+ batch = ops.convert_to_tensor([0, 0, 0, 0, 1, 1, 1, 1, 1, 1])
19
+
20
+ model = Polynormer(
21
+ in_channels=16,
22
+ hidden_channels=32,
23
+ out_channels=40,
24
+ local_layers=2,
25
+ global_layers=1,
26
+ qk_shared=qk_shared,
27
+ pre_ln=pre_ln,
28
+ post_bn=post_bn,
29
+ local_attn=local_attn,
30
+ heads=2,
31
+ )
32
+ out = model(x, edge_index, batch)
33
+ assert ops.shape(out) == (10, 40)
34
+
35
+
36
+ @pytest.mark.parametrize('local_attn', [True, False])
37
+ @pytest.mark.parametrize('qk_shared', [True, False])
38
+ def test_polynormer_global(local_attn, qk_shared):
39
+ """Global-mode Polynormer (_global=True) should output [N, out_channels]."""
40
+ x = ops.ones((10, 16))
41
+ edge_index = ops.convert_to_tensor([
42
+ [0, 1, 2, 3, 4, 5, 6, 7, 8, 9],
43
+ [1, 2, 3, 4, 0, 6, 7, 8, 9, 5],
44
+ ])
45
+ batch = ops.convert_to_tensor([0, 0, 0, 0, 1, 1, 1, 1, 1, 1])
46
+
47
+ model = Polynormer(
48
+ in_channels=16,
49
+ hidden_channels=32,
50
+ out_channels=40,
51
+ local_layers=2,
52
+ global_layers=1,
53
+ qk_shared=qk_shared,
54
+ local_attn=local_attn,
55
+ )
56
+ model._global = True
57
+ out = model(x, edge_index, batch)
58
+ assert ops.shape(out) == (10, 40)
59
+
60
+
61
+ def test_polynormer_log_softmax_sums():
62
+ """Output of Polynormer is log_softmax, so exp(out).sum(axis=-1) ≈ 1."""
63
+ import numpy as np
64
+ x = ops.ones((6, 8))
65
+ edge_index = ops.convert_to_tensor([[0, 1, 2], [1, 2, 0]])
66
+ batch = ops.convert_to_tensor([0, 0, 0, 0, 0, 0])
67
+
68
+ model = Polynormer(
69
+ in_channels=8,
70
+ hidden_channels=16,
71
+ out_channels=5,
72
+ local_layers=1,
73
+ global_layers=1,
74
+ )
75
+ out = model(x, edge_index, batch)
76
+ assert ops.shape(out) == (6, 5)
77
+
78
+ out_np = ops.convert_to_numpy(out)
79
+ sums = np.exp(out_np).sum(axis=-1)
80
+ assert np.allclose(sums, 1.0, atol=1e-5)
81
+
82
+
83
+ def test_polynormer_global_toggle():
84
+ """Switching between local and global modes on the same model instance."""
85
+ x = ops.ones((8, 16))
86
+ edge_index = ops.convert_to_tensor([[0, 1, 2, 3], [1, 2, 3, 0]])
87
+ batch = ops.convert_to_tensor([0, 0, 0, 0, 1, 1, 1, 1])
88
+
89
+ model = Polynormer(
90
+ in_channels=16,
91
+ hidden_channels=16,
92
+ out_channels=10,
93
+ local_layers=1,
94
+ global_layers=1,
95
+ )
96
+
97
+ model._global = False
98
+ out_local = model(x, edge_index, batch)
99
+ assert ops.shape(out_local) == (8, 10)
100
+
101
+ model._global = True
102
+ out_global = model(x, edge_index, batch)
103
+ assert ops.shape(out_global) == (8, 10)
104
+