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,379 @@
1
+ import os
2
+ import os.path as osp
3
+ import shutil
4
+ from typing import Optional, Tuple, Union
5
+
6
+ import keras
7
+ from keras import layers, ops
8
+
9
+ from k3_node.layers.conv.message_passing import MessagePassing
10
+ from k3_node.layers.conv.utils import add_self_loops
11
+ from k3_node.layers.pool import global_add_pool, global_max_pool, global_mean_pool
12
+ from k3_node.data.download import download_url
13
+ from k3_node.ops.creation import full
14
+
15
+
16
+ MOLE_BERT_URL = (
17
+ "https://github.com/junxia97/Mole-BERT/raw/refs/heads/main/model_gin/Mole-BERT.pth"
18
+ )
19
+
20
+
21
+ class MoleBERTGINConv(MessagePassing):
22
+ """Extension of GIN aggregation to incorporate categorical edge information with self-loops.
23
+
24
+ Matches the GINConv variant from Hu et al. used in Mole-BERT.
25
+
26
+ Example:
27
+ ```python
28
+ import numpy as np
29
+ from k3_node.models import MoleBERTGINConv
30
+
31
+ x = np.random.rand(4, 32).astype("float32") # atom embeddings
32
+ edge_index = np.array([[0, 1, 1, 2], [1, 0, 2, 1]])
33
+ edge_attr = np.array([[0, 0], [0, 0], [1, 0], [1, 0]]) # bond type, bond direction
34
+ conv = MoleBERTGINConv(emb_dim=32)
35
+ print(tuple(conv(x, edge_index, edge_attr).shape)) # (4, 32)
36
+ ```
37
+ """
38
+
39
+ def __init__(
40
+ self,
41
+ emb_dim: int = 300,
42
+ out_dim: Optional[int] = None,
43
+ num_bond_type: int = 6,
44
+ num_bond_direction: int = 3,
45
+ **kwargs,
46
+ ):
47
+ super().__init__(aggr="add", **kwargs)
48
+ self.emb_dim = emb_dim
49
+ self.out_dim = emb_dim if out_dim is None else out_dim
50
+ self.num_bond_type = num_bond_type
51
+ self.num_bond_direction = num_bond_direction
52
+
53
+ self.mlp = keras.Sequential(
54
+ [
55
+ layers.Dense(2 * emb_dim, activation="relu", name="mlp_0"),
56
+ layers.Dense(self.out_dim, name="mlp_2"),
57
+ ],
58
+ name="mlp",
59
+ )
60
+ self.edge_embedding1 = layers.Embedding(
61
+ num_bond_type, emb_dim, name="edge_embedding1"
62
+ )
63
+ self.edge_embedding2 = layers.Embedding(
64
+ num_bond_direction, emb_dim, name="edge_embedding2"
65
+ )
66
+
67
+ def build(self, input_shape=None):
68
+ self.mlp.build((None, self.emb_dim))
69
+ self.edge_embedding1.build((None,))
70
+ self.edge_embedding2.build((None,))
71
+ super().build(input_shape)
72
+
73
+ def call(self, x, edge_index, edge_attr):
74
+ num_nodes = ops.shape(x)[0]
75
+
76
+ # Add self-loops to edge space
77
+ edge_index, _ = add_self_loops(edge_index, num_nodes=num_nodes)
78
+
79
+ # Add features corresponding to self-loop edges: [4, 0]
80
+ self_loop_attr = ops.stack(
81
+ [
82
+ full((num_nodes,), 4, dtype=edge_attr.dtype),
83
+ ops.zeros((num_nodes,), dtype=edge_attr.dtype),
84
+ ],
85
+ axis=1,
86
+ )
87
+ edge_attr = ops.concatenate([edge_attr, self_loop_attr], axis=0)
88
+
89
+ edge_embeddings = self.edge_embedding1(edge_attr[:, 0]) + self.edge_embedding2(
90
+ edge_attr[:, 1]
91
+ )
92
+
93
+ return self.propagate(
94
+ edge_index, x=x, edge_attr=edge_embeddings, size=(num_nodes, num_nodes)
95
+ )
96
+
97
+ def message(self, x_j, edge_attr):
98
+ return x_j + edge_attr
99
+
100
+ def update(self, aggr_out):
101
+ return self.mlp(aggr_out)
102
+
103
+
104
+ class MoleBERTGNN(layers.Layer):
105
+ """5-layer GIN encoder backbone of Mole-BERT with Jumping Knowledge.
106
+
107
+ Example:
108
+ ```python
109
+ import numpy as np
110
+ from k3_node.models import MoleBERTGNN
111
+
112
+ x = np.stack([np.random.randint(0, 119, size=5), np.random.randint(0, 3, size=5)], axis=1) # atom type, chirality
113
+ edge_index = np.array([[0, 1, 1, 2, 3, 4], [1, 0, 2, 1, 4, 3]])
114
+ edge_attr = np.stack([np.random.randint(0, 5, size=6), np.random.randint(0, 3, size=6)], axis=1) # bond type, direction
115
+ batch = np.array([0, 0, 0, 1, 1]) # two molecules
116
+
117
+ gnn = MoleBERTGNN(num_layer=3, emb_dim=32, JK="last")
118
+ print(tuple(gnn(x, edge_index, edge_attr).shape)) # (5, 32)
119
+ ```
120
+ """
121
+
122
+ def __init__(
123
+ self,
124
+ num_layer: int = 5,
125
+ emb_dim: int = 300,
126
+ num_atom_type: int = 120,
127
+ num_chirality_tag: int = 3,
128
+ JK: str = "last",
129
+ drop_ratio: float = 0.0,
130
+ **kwargs,
131
+ ):
132
+ super().__init__(**kwargs)
133
+ if num_layer < 2:
134
+ raise ValueError("Number of GNN layers must be greater than 1.")
135
+ self.num_layer = num_layer
136
+ self.emb_dim = emb_dim
137
+ self.num_atom_type = num_atom_type
138
+ self.num_chirality_tag = num_chirality_tag
139
+ self.JK = JK
140
+ self.drop_ratio = drop_ratio
141
+
142
+ self.x_embedding1 = layers.Embedding(
143
+ num_atom_type, emb_dim, name="x_embedding1"
144
+ )
145
+ self.x_embedding2 = layers.Embedding(
146
+ num_chirality_tag, emb_dim, name="x_embedding2"
147
+ )
148
+
149
+ self.gnns = [
150
+ MoleBERTGINConv(emb_dim=emb_dim, name=f"gnns_{i}")
151
+ for i in range(num_layer)
152
+ ]
153
+ self.batch_norms = [
154
+ layers.BatchNormalization(
155
+ axis=-1, epsilon=1e-5, momentum=0.9, name=f"batch_norms_{i}"
156
+ )
157
+ for i in range(num_layer)
158
+ ]
159
+ self.dropout_layer = layers.Dropout(drop_ratio)
160
+
161
+ def build(self, input_shape=None):
162
+ self.x_embedding1.build((None,))
163
+ self.x_embedding2.build((None,))
164
+ for i in range(self.num_layer):
165
+ self.gnns[i].build(None)
166
+ self.batch_norms[i].build((None, self.emb_dim))
167
+ super().build(input_shape)
168
+
169
+ def call(self, x, edge_index, edge_attr, training: bool = False):
170
+ h = self.x_embedding1(x[:, 0]) + self.x_embedding2(x[:, 1])
171
+ h_list = [h]
172
+
173
+ for layer in range(self.num_layer):
174
+ h = self.gnns[layer](h_list[layer], edge_index, edge_attr)
175
+ h = self.batch_norms[layer](h, training=training)
176
+ if layer == self.num_layer - 1:
177
+ # Remove ReLU for the last layer
178
+ h = self.dropout_layer(h, training=training)
179
+ else:
180
+ h = self.dropout_layer(ops.relu(h), training=training)
181
+ h_list.append(h)
182
+
183
+ if self.JK == "concat":
184
+ node_representation = ops.concatenate(h_list, axis=1)
185
+ elif self.JK == "last":
186
+ node_representation = h_list[-1]
187
+ elif self.JK == "max":
188
+ node_representation = ops.max(ops.stack(h_list, axis=0), axis=0)
189
+ elif self.JK == "sum":
190
+ node_representation = ops.sum(ops.stack(h_list, axis=0), axis=0)
191
+ else:
192
+ raise ValueError(f"Unknown JK mode: {self.JK}")
193
+
194
+ return node_representation
195
+
196
+
197
+ class MoleBERT(keras.Model):
198
+ """Complete Mole-BERT Model with graph-level pooling and property prediction head.
199
+
200
+ Example:
201
+ ```python
202
+ import numpy as np
203
+ from k3_node.models import MoleBERT
204
+
205
+ x = np.stack([np.random.randint(0, 119, size=5), np.random.randint(0, 3, size=5)], axis=1) # atom type, chirality
206
+ edge_index = np.array([[0, 1, 1, 2, 3, 4], [1, 0, 2, 1, 4, 3]])
207
+ edge_attr = np.stack([np.random.randint(0, 5, size=6), np.random.randint(0, 3, size=6)], axis=1) # bond type, direction
208
+ batch = np.array([0, 0, 0, 1, 1]) # two molecules
209
+
210
+ model = MoleBERT(num_layer=3, emb_dim=32, num_tasks=2, graph_pooling="mean")
211
+ logits, node_rep = model((x, edge_index, edge_attr, batch))
212
+ print(tuple(logits.shape), tuple(node_rep.shape)) # (2, 2) (5, 32): per-molecule predictions, per-atom embeddings
213
+ ```
214
+ """
215
+
216
+ def __init__(
217
+ self,
218
+ num_layer: int = 5,
219
+ emb_dim: int = 300,
220
+ num_tasks: Optional[int] = None,
221
+ JK: str = "last",
222
+ drop_ratio: float = 0.0,
223
+ graph_pooling: str = "mean",
224
+ **kwargs,
225
+ ):
226
+ super().__init__(**kwargs)
227
+ self.num_layer = num_layer
228
+ self.emb_dim = emb_dim
229
+ self.num_tasks = num_tasks
230
+ self.JK = JK
231
+ self.drop_ratio = drop_ratio
232
+ self.graph_pooling = graph_pooling
233
+
234
+ self.gnn = MoleBERTGNN(
235
+ num_layer=num_layer,
236
+ emb_dim=emb_dim,
237
+ JK=JK,
238
+ drop_ratio=drop_ratio,
239
+ name="gnn",
240
+ )
241
+
242
+ if graph_pooling in ("sum", "add"):
243
+ self.pool_fn = global_add_pool
244
+ elif graph_pooling == "mean":
245
+ self.pool_fn = global_mean_pool
246
+ elif graph_pooling == "max":
247
+ self.pool_fn = global_max_pool
248
+ else:
249
+ raise ValueError(f"Invalid graph pooling type: '{graph_pooling}'")
250
+
251
+ if num_tasks is not None:
252
+ mult = (num_layer + 1) if JK == "concat" else 1
253
+ self.graph_pred_linear = layers.Dense(
254
+ num_tasks, name="graph_pred_linear"
255
+ )
256
+ else:
257
+ self.graph_pred_linear = None
258
+
259
+ def build(self, input_shape=None):
260
+ self.gnn.build(None)
261
+ if self.graph_pred_linear is not None:
262
+ mult = (self.num_layer + 1) if self.JK == "concat" else 1
263
+ self.graph_pred_linear.build((None, mult * self.emb_dim))
264
+ super().build(input_shape)
265
+
266
+ def call(self, inputs, training: bool = False):
267
+ """Call MoleBERT model.
268
+
269
+ inputs can be a tuple: `(x, edge_index, edge_attr)` or `(x, edge_index, edge_attr, batch)`.
270
+ """
271
+ if isinstance(inputs, (list, tuple)):
272
+ if len(inputs) == 3:
273
+ x, edge_index, edge_attr = inputs
274
+ batch = None
275
+ elif len(inputs) == 4:
276
+ x, edge_index, edge_attr, batch = inputs
277
+ else:
278
+ raise ValueError("Expected 3 or 4 input tensors.")
279
+ elif isinstance(inputs, dict):
280
+ x = inputs["x"]
281
+ edge_index = inputs["edge_index"]
282
+ edge_attr = inputs["edge_attr"]
283
+ batch = inputs.get("batch", None)
284
+ else:
285
+ raise ValueError("inputs must be a tuple, list, or dict.")
286
+
287
+ node_rep = self.gnn(x, edge_index, edge_attr, training=training)
288
+
289
+ if batch is not None:
290
+ graph_rep = self.pool_fn(node_rep, batch)
291
+ if self.graph_pred_linear is not None:
292
+ logits = self.graph_pred_linear(graph_rep)
293
+ return logits, node_rep
294
+ return graph_rep, node_rep
295
+
296
+ if self.graph_pred_linear is not None:
297
+ # If no batch given, default all nodes to batch 0
298
+ num_nodes = ops.shape(node_rep)[0]
299
+ batch_zeros = ops.zeros((num_nodes,), dtype="int32")
300
+ graph_rep = self.pool_fn(node_rep, batch_zeros)
301
+ logits = self.graph_pred_linear(graph_rep)
302
+ return logits, node_rep
303
+
304
+ return node_rep
305
+
306
+
307
+ def load_mole_bert_weights(model: Union[MoleBERT, MoleBERTGNN], checkpoint_path: str):
308
+ """Loads official PyTorch Mole-BERT.pth checkpoint state dict into Keras 3 MoleBERT model."""
309
+ import torch
310
+
311
+ try:
312
+ state_dict = torch.load(checkpoint_path, map_location="cpu", weights_only=False)
313
+ except Exception:
314
+ state_dict = torch.load(checkpoint_path, map_location="cpu")
315
+
316
+ if "state_dict" in state_dict:
317
+ state_dict = state_dict["state_dict"]
318
+
319
+ def get_np(key):
320
+ return state_dict[key].detach().cpu().float().numpy()
321
+
322
+ if not model.built:
323
+ model.build(None)
324
+
325
+ gnn = model.gnn if hasattr(model, "gnn") else model
326
+
327
+ # Embeddings
328
+ gnn.x_embedding1.set_weights([get_np("x_embedding1.weight")])
329
+ gnn.x_embedding2.set_weights([get_np("x_embedding2.weight")])
330
+
331
+ for l in range(gnn.num_layer):
332
+ conv = gnn.gnns[l]
333
+ bn = gnn.batch_norms[l]
334
+
335
+ # GINConv edge embeddings
336
+ conv.edge_embedding1.set_weights([get_np(f"gnns.{l}.edge_embedding1.weight")])
337
+ conv.edge_embedding2.set_weights([get_np(f"gnns.{l}.edge_embedding2.weight")])
338
+
339
+ # GINConv MLP layers
340
+ w0 = get_np(f"gnns.{l}.mlp.0.weight").T
341
+ b0 = get_np(f"gnns.{l}.mlp.0.bias")
342
+ conv.mlp.layers[0].set_weights([w0, b0])
343
+
344
+ w2 = get_np(f"gnns.{l}.mlp.2.weight").T
345
+ b2 = get_np(f"gnns.{l}.mlp.2.bias")
346
+ conv.mlp.layers[1].set_weights([w2, b2])
347
+
348
+ # Batch Normalization
349
+ gamma = get_np(f"batch_norms.{l}.weight")
350
+ beta = get_np(f"batch_norms.{l}.bias")
351
+ mean = get_np(f"batch_norms.{l}.running_mean")
352
+ var = get_np(f"batch_norms.{l}.running_var")
353
+ bn.set_weights([gamma, beta, mean, var])
354
+
355
+
356
+ def download_mole_bert_checkpoint(cache_dir: Optional[str] = None) -> str:
357
+ """Downloads official Mole-BERT.pth checkpoint from GitHub."""
358
+ if cache_dir is None:
359
+ cache_dir = osp.expanduser("~/.cache/k3_node/mole_bert")
360
+
361
+ os.makedirs(cache_dir, exist_ok=True)
362
+ target_path = osp.join(cache_dir, "Mole-BERT.pth")
363
+
364
+ if osp.exists(target_path) and osp.getsize(target_path) > 1000:
365
+ return target_path
366
+
367
+ # Check local path in Mole-BERT/model_gin/Mole-BERT.pth
368
+ local_path = osp.join(
369
+ osp.dirname(osp.dirname(osp.dirname(osp.abspath(__file__)))),
370
+ "Mole-BERT",
371
+ "model_gin",
372
+ "Mole-BERT.pth",
373
+ )
374
+ if osp.exists(local_path) and osp.getsize(local_path) > 1000:
375
+ shutil.copyfile(local_path, target_path)
376
+ return target_path
377
+
378
+ return download_url(MOLE_BERT_URL, cache_dir, filename="Mole-BERT.pth")
379
+
@@ -0,0 +1,95 @@
1
+ from typing import Optional
2
+
3
+ import keras
4
+ from keras import ops
5
+
6
+ from k3_node.layers.conv import MFConv
7
+ from k3_node.layers.pool import global_add_pool
8
+
9
+
10
+ class NeuralFingerprint(keras.layers.Layer):
11
+ r"""The Neural Fingerprint model from the
12
+ `"Convolutional Networks on Graphs for Learning Molecular Fingerprints"
13
+ <https://arxiv.org/abs/1509.09292>`__ paper to generate fingerprints
14
+ of molecules.
15
+
16
+ Args:
17
+ in_channels (int): Size of each input sample.
18
+ hidden_channels (int): Size of each hidden sample.
19
+ out_channels (int): Size of each output fingerprint.
20
+ num_layers (int): Number of layers.
21
+ **kwargs (optional): Additional arguments of
22
+ :class:`~k3_node.layers.conv.MFConv`.
23
+
24
+ Example:
25
+ ```python
26
+ import numpy as np
27
+ from k3_node.models import NeuralFingerprint
28
+
29
+ x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
30
+ edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
31
+
32
+ batch = np.repeat([0, 1], 5) # two molecules with 5 atoms each
33
+ model = NeuralFingerprint(in_channels=8, hidden_channels=32, out_channels=16, num_layers=3)
34
+ fingerprint = model(x, edge_index, batch) # one learned fingerprint per molecule
35
+ print(tuple(fingerprint.shape)) # (2, 16)
36
+ ```
37
+ """
38
+ def __init__(
39
+ self,
40
+ in_channels: int,
41
+ hidden_channels: int,
42
+ out_channels: int,
43
+ num_layers: int,
44
+ **kwargs,
45
+ ):
46
+ super().__init__()
47
+
48
+ self.in_channels = in_channels
49
+ self.hidden_channels = hidden_channels
50
+ self.out_channels = out_channels
51
+ self.num_layers = num_layers
52
+
53
+ self.convs = []
54
+ for i in range(self.num_layers):
55
+ in_c = self.in_channels if i == 0 else self.hidden_channels
56
+ self.convs.append(MFConv(in_c, hidden_channels, **kwargs))
57
+
58
+ self.lins = []
59
+ for _ in range(self.num_layers):
60
+ self.lins.append(keras.layers.Dense(out_channels, use_bias=False))
61
+
62
+ def build(self, input_shape=None):
63
+ self.built = True
64
+
65
+ def reset_parameters(self):
66
+ r"""Resets all learnable parameters of the module."""
67
+ for conv in self.convs:
68
+ if hasattr(conv, "reset_parameters"):
69
+ conv.reset_parameters()
70
+ for lin in self.lins:
71
+ if lin.built:
72
+ lin.kernel.assign(keras.initializers.GlorotUniform()(lin.kernel.shape))
73
+
74
+ def call(
75
+ self,
76
+ x,
77
+ edge_index,
78
+ batch: Optional[any] = None,
79
+ batch_size: Optional[int] = None,
80
+ ):
81
+ outs = []
82
+ for conv, lin in zip(self.convs, self.lins):
83
+ x = ops.sigmoid(conv(x, edge_index))
84
+ y = ops.softmax(lin(x), axis=-1)
85
+ outs.append(global_add_pool(y, batch, size=batch_size))
86
+
87
+ out = outs[0]
88
+ for item in outs[1:]:
89
+ out = out + item
90
+ return out
91
+
92
+ def __repr__(self) -> str:
93
+ return (f'{self.__class__.__name__}({self.in_channels}, '
94
+ f'{self.out_channels}, num_layers={self.num_layers})')
95
+
@@ -0,0 +1,213 @@
1
+ from typing import Optional
2
+ import numpy as np
3
+ import keras
4
+ from keras import ops
5
+
6
+
7
+ class Node2Vec(keras.layers.Layer):
8
+ r"""The Node2Vec model from the
9
+ `"node2vec: Scalable Feature Learning for Networks"
10
+ <https://arxiv.org/abs/1607.00653>`_ paper where random walks of
11
+ length :obj:`walk_length` are sampled in a given graph, and node embeddings
12
+ are learned via negative sampling optimization.
13
+
14
+ Args:
15
+ edge_index: The edge indices.
16
+ embedding_dim (int): The size of each embedding vector.
17
+ walk_length (int): The walk length.
18
+ context_size (int): The actual context size which is considered for
19
+ positive samples.
20
+ walks_per_node (int, optional): The number of walks to sample for each
21
+ node. (default: :obj:`1`)
22
+ p (float, optional): Likelihood of immediately revisiting a node in the
23
+ walk. (default: :obj:`1.0`)
24
+ q (float, optional): Control parameter to interpolate between
25
+ breadth-first strategy and depth-first strategy. (default: :obj:`1.0`)
26
+ num_negative_samples (int, optional): The number of negative samples to
27
+ use for each positive sample. (default: :obj:`1`)
28
+ num_nodes (int, optional): The number of nodes. (default: :obj:`None`)
29
+
30
+ Example:
31
+ ```python
32
+ import numpy as np
33
+ from k3_node.models import Node2Vec
34
+
35
+ edge_index = np.array([[0, 1, 2, 3, 0, 2], [1, 2, 3, 0, 2, 0]])
36
+ model = Node2Vec(edge_index, embedding_dim=16, walk_length=4, context_size=3, walks_per_node=2)
37
+ print(tuple(model().shape)) # (4, 16): embeddings of all nodes
38
+
39
+ batch = np.array([0, 1])
40
+ loss = model.loss(model.pos_sample(batch), model.neg_sample(batch)) # skip-gram loss on random walks
41
+ print(tuple(loss.shape)) # (): a scalar
42
+ ```
43
+ """
44
+ def __init__(
45
+ self,
46
+ edge_index,
47
+ embedding_dim: int,
48
+ walk_length: int,
49
+ context_size: int,
50
+ walks_per_node: int = 1,
51
+ p: float = 1.0,
52
+ q: float = 1.0,
53
+ num_negative_samples: int = 1,
54
+ num_nodes: Optional[int] = None,
55
+ **kwargs,
56
+ ):
57
+ super().__init__(**kwargs)
58
+
59
+ edge_index_np = ops.convert_to_numpy(edge_index).astype(np.int64)
60
+ if num_nodes is None:
61
+ num_nodes = int(edge_index_np.max()) + 1 if edge_index_np.size > 0 else 0
62
+
63
+ self.num_nodes = num_nodes
64
+ self.embedding_dim = embedding_dim
65
+ self.walk_length = walk_length - 1
66
+ self.context_size = context_size
67
+ self.walks_per_node = walks_per_node
68
+ self.p = p
69
+ self.q = q
70
+ self.num_negative_samples = num_negative_samples
71
+ self.EPS = 1e-15
72
+
73
+ # Build adjacency list for fast random walks
74
+ self.adj = [[] for _ in range(num_nodes)]
75
+ if edge_index_np.size > 0:
76
+ for src, dst in zip(edge_index_np[0], edge_index_np[1]):
77
+ self.adj[src].append(int(dst))
78
+ self.adj_sets = [set(nbrs) for nbrs in self.adj]
79
+
80
+ self.embedding = keras.layers.Embedding(num_nodes, embedding_dim)
81
+
82
+ def reset_parameters(self):
83
+ if self.embedding.built:
84
+ self.embedding.embeddings.assign(
85
+ keras.initializers.GlorotUniform()(self.embedding.embeddings.shape)
86
+ )
87
+
88
+ def call(self, batch: Optional[any] = None):
89
+ """Returns the embeddings for the nodes in :obj:`batch`."""
90
+ if batch is None:
91
+ batch = ops.arange(self.num_nodes, dtype="int64")
92
+ return self.embedding(batch)
93
+
94
+ def pos_sample(self, batch):
95
+ batch_np = ops.convert_to_numpy(batch).astype(np.int64)
96
+ repeated = np.repeat(batch_np, self.walks_per_node)
97
+
98
+ # Sample random walks
99
+ all_walks = [self._random_walk(int(node)) for node in repeated]
100
+
101
+ rw = np.array(all_walks, dtype=np.int64)
102
+ walks = []
103
+ num_walks_per_rw = 1 + self.walk_length + 1 - self.context_size
104
+ for j in range(num_walks_per_rw):
105
+ walks.append(rw[:, j : j + self.context_size])
106
+ out = np.concatenate(walks, axis=0) if len(walks) > 0 else rw
107
+ return ops.convert_to_tensor(out, dtype="int64")
108
+
109
+ def _random_walk(self, start: int):
110
+ """A node2vec walk: from ``cur`` (reached from ``prev``) the next node is weighted by 1/p
111
+ if it returns to ``prev``, 1 if it is also a neighbor of ``prev`` and 1/q otherwise."""
112
+ walk = [start]
113
+ for _ in range(self.walk_length):
114
+ cur = walk[-1]
115
+ nbrs = self.adj[cur]
116
+ if not nbrs:
117
+ walk.append(cur)
118
+ continue
119
+ if (self.p == 1.0 and self.q == 1.0) or len(walk) == 1:
120
+ walk.append(nbrs[np.random.randint(len(nbrs))])
121
+ continue
122
+ prev = walk[-2]
123
+ prev_nbrs = self.adj_sets[prev]
124
+ weights = np.array([1.0 / self.p if n == prev else (1.0 if n in prev_nbrs else 1.0 / self.q)
125
+ for n in nbrs])
126
+ walk.append(nbrs[np.random.choice(len(nbrs), p=weights / weights.sum())])
127
+ return walk
128
+
129
+ def loader(self, batch_size: int = 128, shuffle: bool = True):
130
+ r"""Yields ``(pos_rw, neg_rw)`` random-walk batches, starting from ``batch_size`` nodes
131
+ at a time, as PyG's ``Node2Vec.loader``."""
132
+ nodes = np.random.permutation(self.num_nodes) if shuffle else np.arange(self.num_nodes)
133
+ for start in range(0, self.num_nodes, batch_size):
134
+ batch = nodes[start:start + batch_size]
135
+ yield self.pos_sample(batch), self.neg_sample(batch)
136
+
137
+ def compile(self, optimizer):
138
+ r"""Sets the optimizer used by :meth:`fit`."""
139
+ self.optimizer = optimizer
140
+
141
+ def fit(self, epochs: int = 1, batch_size: int = 128, verbose: int = 1):
142
+ r"""Trains the embeddings on freshly sampled random walks for ``epochs`` passes over
143
+ all nodes; returns the mean loss of every epoch."""
144
+ from k3_node.training import gradient_step
145
+
146
+ if getattr(self, "optimizer", None) is None:
147
+ raise ValueError("Call `compile(optimizer=...)` before `fit`.")
148
+ self(ops.arange(1)) # create the embeddings
149
+ history = {"loss": []}
150
+ for epoch in range(1, epochs + 1):
151
+ losses = [gradient_step(lambda: self.loss(pos_rw, neg_rw), self.trainable_variables, self.optimizer)
152
+ for pos_rw, neg_rw in self.loader(batch_size, shuffle=True)]
153
+ history["loss"].append(float(np.mean(losses)))
154
+ if verbose:
155
+ print(f"Epoch {epoch:03d}: loss: {history['loss'][-1]:.4f}")
156
+ return history
157
+
158
+ def test(self, train_z, train_y, test_z, test_y, solver: str = "lbfgs", *args, **kwargs):
159
+ r"""Evaluates the embeddings with a logistic regression classifier; returns its accuracy."""
160
+ from sklearn.linear_model import LogisticRegression
161
+
162
+ clf = LogisticRegression(*args, solver=solver, **kwargs).fit(
163
+ ops.convert_to_numpy(train_z), ops.convert_to_numpy(train_y))
164
+ return clf.score(ops.convert_to_numpy(test_z), ops.convert_to_numpy(test_y))
165
+
166
+ def neg_sample(self, batch):
167
+ batch_np = ops.convert_to_numpy(batch).astype(np.int64)
168
+ repeated = np.repeat(
169
+ batch_np, self.walks_per_node * self.num_negative_samples
170
+ )
171
+ rand_rest = np.random.randint(
172
+ 0, max(self.num_nodes, 1), size=(len(repeated), self.walk_length)
173
+ )
174
+ rw = np.concatenate([repeated[:, None], rand_rest], axis=-1)
175
+
176
+ walks = []
177
+ num_walks_per_rw = 1 + self.walk_length + 1 - self.context_size
178
+ for j in range(num_walks_per_rw):
179
+ walks.append(rw[:, j : j + self.context_size])
180
+ out = np.concatenate(walks, axis=0) if len(walks) > 0 else rw
181
+ return ops.convert_to_tensor(out, dtype="int64")
182
+
183
+ def loss(self, pos_rw, neg_rw):
184
+ r"""Computes the loss given positive and negative random walks."""
185
+ # Positive loss
186
+ start = pos_rw[:, 0]
187
+ rest = pos_rw[:, 1:]
188
+ pos_b = ops.shape(pos_rw)[0]
189
+
190
+ h_start = ops.reshape(self.embedding(start), (pos_b, 1, self.embedding_dim))
191
+ h_rest = ops.reshape(
192
+ self.embedding(ops.reshape(rest, (-1,))),
193
+ (pos_b, -1, self.embedding_dim),
194
+ )
195
+
196
+ out = ops.reshape(ops.sum(h_start * h_rest, axis=-1), (-1,))
197
+ pos_loss = -ops.mean(ops.log(ops.sigmoid(out) + self.EPS))
198
+
199
+ # Negative loss
200
+ start = neg_rw[:, 0]
201
+ rest = neg_rw[:, 1:]
202
+ neg_b = ops.shape(neg_rw)[0]
203
+
204
+ h_start = ops.reshape(self.embedding(start), (neg_b, 1, self.embedding_dim))
205
+ h_rest = ops.reshape(
206
+ self.embedding(ops.reshape(rest, (-1,))),
207
+ (neg_b, -1, self.embedding_dim),
208
+ )
209
+
210
+ out = ops.reshape(ops.sum(h_start * h_rest, axis=-1), (-1,))
211
+ neg_loss = -ops.mean(ops.log(1.0 - ops.sigmoid(out) + self.EPS))
212
+
213
+ return pos_loss + neg_loss