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,131 @@
1
+ import os
2
+ import tempfile
3
+ import numpy as np
4
+ import pytest
5
+ import torch
6
+ import keras
7
+ from keras import ops
8
+
9
+ from k3_node.models import (
10
+ UniMolPlusPCQModel,
11
+ UniMolPlusOC20Model,
12
+ DockingPoseModelV2,
13
+ download_unimol_plus_checkpoint,
14
+ load_unimol_plus_weights,
15
+ download_unimol_docking_checkpoint,
16
+ load_unimol_docking_weights,
17
+ )
18
+
19
+
20
+ def test_unimol_plus_pcq_model():
21
+ bsz = 2
22
+ seq_len = 5
23
+ embed_dim = 32
24
+ pair_dim = 16
25
+ num_heads = 4
26
+ num_layers = 2
27
+
28
+ model = UniMolPlusPCQModel(
29
+ num_layers=num_layers,
30
+ embed_dim=embed_dim,
31
+ pair_dim=pair_dim,
32
+ num_heads=num_heads,
33
+ output_dim=1,
34
+ )
35
+
36
+ atom_types = ops.convert_to_tensor(np.random.randint(0, 64, (bsz, seq_len)), dtype="int32")
37
+ coords = ops.convert_to_tensor(np.random.randn(bsz, seq_len, 3).astype("float32"))
38
+
39
+ # Property prediction
40
+ pred = model(atom_types, coords=coords)
41
+ assert ops.shape(pred) == (bsz, 1)
42
+
43
+ # Return coords
44
+ pred, new_coords = model(atom_types, coords=coords, return_coords=True)
45
+ assert ops.shape(pred) == (bsz, 1)
46
+ assert ops.shape(new_coords) == (bsz, seq_len, 3)
47
+
48
+
49
+ def test_unimol_plus_oc20_model():
50
+ bsz = 2
51
+ seq_len = 6
52
+ embed_dim = 32
53
+ pair_dim = 16
54
+ num_heads = 4
55
+ num_layers = 2
56
+
57
+ model = UniMolPlusOC20Model(
58
+ num_layers=num_layers,
59
+ embed_dim=embed_dim,
60
+ pair_dim=pair_dim,
61
+ num_heads=num_heads,
62
+ output_dim=1,
63
+ )
64
+
65
+ atom_types = ops.convert_to_tensor(np.random.randint(0, 64, (bsz, seq_len)), dtype="int32")
66
+ coords = ops.convert_to_tensor(np.random.randn(bsz, seq_len, 3).astype("float32"))
67
+
68
+ pred = model(atom_types, coords=coords)
69
+ assert ops.shape(pred) == (bsz, 1)
70
+
71
+
72
+ def test_unimol_docking_v2():
73
+ bsz = 2
74
+ n_mol = 4
75
+ n_pkt = 6
76
+ embed_dim = 32
77
+ pair_dim = 16
78
+ num_heads = 4
79
+ num_layers = 2
80
+
81
+ model = DockingPoseModelV2(
82
+ mol_vocab_size=64,
83
+ pocket_vocab_size=64,
84
+ embed_dim=embed_dim,
85
+ pair_dim=pair_dim,
86
+ num_layers=num_layers,
87
+ num_heads=num_heads,
88
+ )
89
+
90
+ mol_tokens = ops.convert_to_tensor(np.random.randint(0, 64, (bsz, n_mol)), dtype="int32")
91
+ pkt_tokens = ops.convert_to_tensor(np.random.randint(0, 64, (bsz, n_pkt)), dtype="int32")
92
+ mol_coords = ops.convert_to_tensor(np.random.randn(bsz, n_mol, 3).astype("float32"))
93
+ pkt_coords = ops.convert_to_tensor(np.random.randn(bsz, n_pkt, 3).astype("float32"))
94
+
95
+ docked_coords, pred_dist = model(
96
+ mol_tokens,
97
+ pkt_tokens,
98
+ mol_coords=mol_coords,
99
+ pocket_coords=pkt_coords,
100
+ )
101
+ assert ops.shape(docked_coords) == (bsz, n_mol, 3)
102
+ assert ops.shape(pred_dist) == (bsz, n_mol + n_pkt, n_mol + n_pkt)
103
+
104
+
105
+ def test_unimol_plus_weight_loading():
106
+ model = UniMolPlusPCQModel(
107
+ num_layers=1,
108
+ embed_dim=16,
109
+ pair_dim=8,
110
+ num_heads=2,
111
+ output_dim=1,
112
+ )
113
+ model.build(None)
114
+
115
+ state_dict = {
116
+ "atom_feature.atom_embed.weight": torch.randn(512, 16),
117
+ "edge_feature.edge_embed.weight": torch.randn(64, 8),
118
+ }
119
+
120
+ with tempfile.NamedTemporaryFile(suffix=".pt", delete=False) as tmp:
121
+ tmp_path = tmp.name
122
+ torch.save({"model": state_dict}, tmp_path)
123
+
124
+ try:
125
+ load_unimol_plus_weights(model, checkpoint_path=tmp_path, download=False)
126
+ w = ops.convert_to_numpy(model.atom_feature.atom_embed.embeddings)
127
+ np.testing.assert_allclose(w, state_dict["atom_feature.atom_embed.weight"].numpy(), atol=1e-5)
128
+ finally:
129
+ if os.path.exists(tmp_path):
130
+ os.remove(tmp_path)
131
+
@@ -0,0 +1,44 @@
1
+ import keras.ops as ops
2
+ from k3_node.models.visnet import ViSNet, CosineCutoff, Sphere
3
+
4
+
5
+ def test_cosine_cutoff():
6
+ cutoff = CosineCutoff(5.0)
7
+ dist = ops.convert_to_tensor([0.0, 2.5, 5.0, 6.0])
8
+ out = cutoff(dist)
9
+ assert float(out[0]) == 1.0
10
+ assert float(out[2]) == 0.0
11
+ assert float(out[3]) == 0.0
12
+
13
+
14
+ def test_sphere():
15
+ sphere1 = Sphere(lmax=1)
16
+ v = ops.convert_to_tensor([[1.0, 2.0, 3.0]])
17
+ sh1 = sphere1(v)
18
+ assert sh1.shape == (1, 3)
19
+
20
+ sphere2 = Sphere(lmax=2)
21
+ sh2 = sphere2(v)
22
+ assert sh2.shape == (1, 8)
23
+
24
+
25
+ def test_visnet():
26
+ z = ops.convert_to_tensor([1, 6, 8, 1])
27
+ pos = ops.convert_to_tensor([
28
+ [0.0, 0.0, 0.0],
29
+ [1.0, 0.0, 0.0],
30
+ [0.0, 1.0, 0.0],
31
+ [1.0, 1.0, 0.0],
32
+ ], dtype="float32")
33
+
34
+ model = ViSNet(
35
+ lmax=1,
36
+ num_heads=2,
37
+ num_layers=2,
38
+ hidden_channels=16,
39
+ num_rbf=8,
40
+ cutoff=5.0,
41
+ )
42
+ y, dy = model(z, pos)
43
+ assert y.shape == (1, 1)
44
+
k3_node/models/tgn.py ADDED
@@ -0,0 +1,382 @@
1
+ from typing import Callable, Dict, List, Optional, Tuple
2
+ import copy
3
+ import numpy as np
4
+ import keras
5
+ from keras import ops
6
+
7
+ from k3_node.layers.aggr import MeanAggregation
8
+
9
+
10
+ class TimeEncoder(keras.layers.Layer):
11
+ """Layer ``TimeEncoder``.
12
+
13
+ Example:
14
+ ```python
15
+ import numpy as np
16
+ from k3_node.models import TimeEncoder
17
+
18
+ t = np.array([1.0, 2.0, 3.0], dtype="float32") # time differences
19
+ print(tuple(TimeEncoder(out_channels=16)(t).shape)) # (3, 16)
20
+ ```
21
+ """
22
+ def __init__(self, out_channels: int, **kwargs):
23
+ super().__init__(**kwargs)
24
+ self.out_channels = out_channels
25
+ self.lin = keras.layers.Dense(out_channels)
26
+
27
+ def reset_parameters(self):
28
+ if self.lin.built:
29
+ self.lin.kernel.assign(keras.initializers.GlorotUniform()(self.lin.kernel.shape))
30
+ if self.lin.bias is not None:
31
+ self.lin.bias.assign(ops.zeros(self.lin.bias.shape))
32
+
33
+ def call(self, t):
34
+ t = ops.reshape(ops.cast(t, "float32"), (-1, 1))
35
+ return ops.cos(self.lin(t))
36
+
37
+
38
+ class IdentityMessage(keras.layers.Layer):
39
+ """Layer ``IdentityMessage``.
40
+
41
+ Example:
42
+ ```python
43
+ import numpy as np
44
+ from k3_node.models import IdentityMessage, LastAggregator, TGNMemory
45
+
46
+ memory = TGNMemory(
47
+ num_nodes=5, raw_msg_dim=8, memory_dim=16, time_dim=16,
48
+ message_module=IdentityMessage(raw_msg_dim=8, memory_dim=16, time_dim=16),
49
+ aggregator_module=LastAggregator(),
50
+ )
51
+ src, dst = np.array([0, 1]), np.array([1, 2]) # two interaction events
52
+ t = np.array([1.0, 2.0], dtype="float32")
53
+ raw_msg = np.random.rand(2, 8).astype("float32")
54
+ memory.update_state(src, dst, t, raw_msg) # update the memory of the involved nodes
55
+
56
+ mem, last_update = memory(np.array([0, 1, 2]))
57
+ print(tuple(mem.shape), tuple(last_update.shape)) # (3, 16) (3,)
58
+ ```
59
+ """
60
+ def __init__(self, raw_msg_dim: int, memory_dim: int, time_dim: int, **kwargs):
61
+ super().__init__(**kwargs)
62
+ self.raw_msg_dim = raw_msg_dim
63
+ self.memory_dim = memory_dim
64
+ self.time_dim = time_dim
65
+ self.out_channels = raw_msg_dim + 2 * memory_dim + time_dim
66
+
67
+ def call(self, z_src, z_dst, raw_msg, t_enc):
68
+ return ops.concatenate([z_src, z_dst, raw_msg, t_enc], axis=-1)
69
+
70
+ def get_config(self):
71
+ config = super().get_config()
72
+ config.update({
73
+ "raw_msg_dim": self.raw_msg_dim,
74
+ "memory_dim": self.memory_dim,
75
+ "time_dim": self.time_dim,
76
+ })
77
+ return config
78
+
79
+
80
+
81
+ class LastAggregator(keras.layers.Layer):
82
+ """Layer ``LastAggregator``.
83
+
84
+ Example:
85
+ ```python
86
+ import numpy as np
87
+ from k3_node.models import LastAggregator
88
+
89
+ msg = np.random.rand(4, 8).astype("float32") # 4 messages
90
+ index = np.array([0, 0, 1, 1]) # destination node of each message
91
+ t = np.array([1.0, 2.0, 1.0, 3.0], dtype="float32") # message timestamps
92
+ out = LastAggregator()(msg, index, t, dim_size=2) # keeps the latest message per node
93
+ print(tuple(out.shape)) # (2, 8)
94
+ ```
95
+ """
96
+ def call(self, msg, index, t, dim_size: int):
97
+ # Which message is the latest per node is decided on the host (it does not depend on the
98
+ # model); the chosen messages are gathered with ops, so gradients reach them.
99
+ from k3_node.ops.host import to_numpy
100
+
101
+ t_np = np.asarray(to_numpy(t)).reshape(-1)
102
+ index_np = np.asarray(to_numpy(index)).astype(np.int64).reshape(-1)
103
+ latest = np.full(dim_size, -1, dtype=np.int64)
104
+ order = np.lexsort((np.arange(len(t_np)), t_np)) # by time, ties: later message wins
105
+ latest[index_np[order]] = order
106
+ has_msg = latest >= 0
107
+ if msg.shape[0] == 0: # no stored messages yet
108
+ return ops.zeros((dim_size, msg.shape[-1]), dtype=msg.dtype)
109
+ out = ops.take(msg, np.maximum(latest, 0), axis=0)
110
+ return out * ops.cast(ops.convert_to_tensor(has_msg[:, None]), out.dtype)
111
+
112
+
113
+ class MeanAggregator(keras.layers.Layer):
114
+ """Layer ``MeanAggregator``.
115
+
116
+ Example:
117
+ ```python
118
+ import numpy as np
119
+ from k3_node.models import MeanAggregator
120
+
121
+ msg = np.random.rand(4, 8).astype("float32") # 4 messages
122
+ index = np.array([0, 0, 1, 1]) # destination node of each message
123
+ t = np.array([1.0, 2.0, 1.0, 3.0], dtype="float32") # message timestamps
124
+ out = MeanAggregator()(msg, index, t, dim_size=2) # averages the messages per node
125
+ print(tuple(out.shape)) # (2, 8)
126
+ ```
127
+ """
128
+ def __init__(self, **kwargs):
129
+ super().__init__(**kwargs)
130
+ self.mean_aggr = MeanAggregation()
131
+
132
+ def build(self, input_shape=None):
133
+ self.built = True
134
+
135
+ def call(self, msg, index, t, dim_size: int):
136
+ return self.mean_aggr(msg, index=index, dim_size=dim_size, dim=0)
137
+
138
+
139
+ class LastNeighborLoader:
140
+ def __init__(self, num_nodes: int, size: int):
141
+ self.num_nodes = num_nodes
142
+ self.size = size
143
+ self.reset_state()
144
+
145
+ def reset_state(self):
146
+ self.cur_e_id = 0
147
+ self.neighbors = np.empty((self.num_nodes, self.size), dtype=np.int64)
148
+ self.e_id = np.full((self.num_nodes, self.size), -1, dtype=np.int64)
149
+
150
+ def __call__(self, n_id):
151
+ n_id_np = ops.convert_to_numpy(n_id).astype(np.int64)
152
+ nbrs = self.neighbors[n_id_np]
153
+ nodes = np.repeat(n_id_np[:, None], self.size, axis=1)
154
+ e_ids = self.e_id[n_id_np]
155
+
156
+ mask = e_ids >= 0
157
+ nbrs = nbrs[mask]
158
+ nodes = nodes[mask]
159
+ e_ids = e_ids[mask]
160
+
161
+ unique_nodes = np.unique(np.concatenate([n_id_np, nbrs]))
162
+ assoc = {node: i for i, node in enumerate(unique_nodes)}
163
+
164
+ mapped_nbrs = np.array([assoc[x] for x in nbrs], dtype=np.int64)
165
+ mapped_nodes = np.array([assoc[x] for x in nodes], dtype=np.int64)
166
+ edge_index = np.stack([mapped_nbrs, mapped_nodes], axis=0) if len(mapped_nbrs) > 0 else np.empty((2, 0), dtype=np.int64)
167
+
168
+ return (
169
+ ops.convert_to_tensor(unique_nodes, dtype="int64"),
170
+ ops.convert_to_tensor(edge_index, dtype="int64"),
171
+ ops.convert_to_tensor(e_ids, dtype="int64"),
172
+ )
173
+
174
+ def insert(self, src, dst):
175
+ src_np = ops.convert_to_numpy(src).astype(np.int64)
176
+ dst_np = ops.convert_to_numpy(dst).astype(np.int64)
177
+
178
+ neighbors = np.concatenate([src_np, dst_np], axis=0)
179
+ nodes = np.concatenate([dst_np, src_np], axis=0)
180
+ num_interactions = len(src_np)
181
+ e_ids = np.repeat(np.arange(self.cur_e_id, self.cur_e_id + num_interactions, dtype=np.int64), 2)
182
+ self.cur_e_id += num_interactions
183
+
184
+ for node, nbr, eid in zip(nodes, neighbors, e_ids):
185
+ # insert into row node
186
+ cur_eids = self.e_id[node]
187
+ cur_nbrs = self.neighbors[node]
188
+ all_eids = np.concatenate([cur_eids, [eid]])
189
+ all_nbrs = np.concatenate([cur_nbrs, [nbr]])
190
+ top_idx = np.argsort(-all_eids)[: self.size]
191
+ self.e_id[node] = all_eids[top_idx]
192
+ self.neighbors[node] = all_nbrs[top_idx]
193
+
194
+
195
+ class TGNMemory(keras.layers.Layer):
196
+ r"""The Temporal Graph Network (TGN) memory model from the
197
+ `"Temporal Graph Networks for Deep Learning on Dynamic Graphs"
198
+ <https://arxiv.org/abs/2006.10637>`_ paper.
199
+
200
+ Args:
201
+ num_nodes (int): The number of nodes to save memories for.
202
+ raw_msg_dim (int): The raw message dimensionality.
203
+ memory_dim (int): The hidden memory dimensionality.
204
+ time_dim (int): The time encoding dimensionality.
205
+ message_module (Callable): Function combining source and destination
206
+ node memory, raw message, and time encoding.
207
+ aggregator_module (Callable): Function aggregating messages to the
208
+ same destination into a single representation.
209
+
210
+ Example:
211
+ ```python
212
+ import numpy as np
213
+ from k3_node.models import IdentityMessage, LastAggregator, TGNMemory
214
+
215
+ memory = TGNMemory(
216
+ num_nodes=5, raw_msg_dim=8, memory_dim=16, time_dim=16,
217
+ message_module=IdentityMessage(raw_msg_dim=8, memory_dim=16, time_dim=16),
218
+ aggregator_module=LastAggregator(),
219
+ )
220
+ src, dst = np.array([0, 1]), np.array([1, 2]) # two interaction events
221
+ t = np.array([1.0, 2.0], dtype="float32")
222
+ raw_msg = np.random.rand(2, 8).astype("float32")
223
+ memory.update_state(src, dst, t, raw_msg) # update the memory of the involved nodes
224
+
225
+ mem, last_update = memory(np.array([0, 1, 2]))
226
+ print(tuple(mem.shape), tuple(last_update.shape)) # (3, 16) (3,)
227
+ ```
228
+ """
229
+ def __init__(
230
+ self,
231
+ num_nodes: int,
232
+ raw_msg_dim: int,
233
+ memory_dim: int,
234
+ time_dim: int,
235
+ message_module: Callable,
236
+ aggregator_module: Callable,
237
+ msg_d_module: Optional[Callable] = None,
238
+ **kwargs,
239
+ ):
240
+ super().__init__(**kwargs)
241
+ self.num_nodes = num_nodes
242
+ self.raw_msg_dim = raw_msg_dim
243
+ self.memory_dim = memory_dim
244
+ self.time_dim = time_dim
245
+
246
+ self.msg_s_module = message_module
247
+ if msg_d_module is not None:
248
+ self.msg_d_module = msg_d_module
249
+ else:
250
+ try:
251
+ self.msg_d_module = message_module.__class__(**message_module.get_config())
252
+ except Exception:
253
+ self.msg_d_module = message_module
254
+ self.aggr_module = aggregator_module
255
+
256
+ self.time_enc = TimeEncoder(time_dim)
257
+ self.gru = keras.layers.GRUCell(memory_dim)
258
+
259
+ self.memory = self.add_weight(
260
+ name="memory",
261
+ shape=(num_nodes, memory_dim),
262
+ initializer="zeros",
263
+ trainable=False,
264
+ )
265
+ self.last_update = self.add_weight(
266
+ name="last_update",
267
+ shape=(num_nodes,),
268
+ initializer="zeros",
269
+ trainable=False,
270
+ dtype="float32",
271
+ )
272
+
273
+ self.msg_s_store = {}
274
+ self.msg_d_store = {}
275
+ self._reset_message_store()
276
+
277
+ def build(self, input_shape=None):
278
+ self.built = True
279
+
280
+ def reset_parameters(self):
281
+ if hasattr(self.msg_s_module, "reset_parameters"):
282
+ self.msg_s_module.reset_parameters()
283
+ if hasattr(self.msg_d_module, "reset_parameters"):
284
+ self.msg_d_module.reset_parameters()
285
+ if hasattr(self.aggr_module, "reset_parameters"):
286
+ self.aggr_module.reset_parameters()
287
+ self.time_enc.reset_parameters()
288
+ self.reset_state()
289
+
290
+ def reset_state(self):
291
+ r"""Starts again from an empty memory and message stores."""
292
+ self.memory.assign(ops.zeros((self.num_nodes, self.memory_dim), dtype=self.memory.dtype))
293
+ self.last_update.assign(ops.zeros((self.num_nodes,), dtype=self.last_update.dtype))
294
+ self._reset_message_store()
295
+
296
+ def _reset_message_store(self):
297
+ empty = (np.zeros(0, np.int64), np.zeros(0, np.int64), np.zeros(0, np.float32),
298
+ np.zeros((0, self.raw_msg_dim), np.float32))
299
+ self.msg_s_store = {j: empty for j in range(self.num_nodes)}
300
+ self.msg_d_store = {j: empty for j in range(self.num_nodes)}
301
+
302
+ # Training mode, as in PyG: in training mode the forward pass computes the memory updated by
303
+ # the stored messages (differentiable); switching to evaluation flushes all pending updates.
304
+ training_mode = True
305
+
306
+ def train(self, mode: bool = True):
307
+ if self.training_mode and not mode:
308
+ self._update_memory(np.arange(self.num_nodes))
309
+ self._reset_message_store()
310
+ self.training_mode = mode
311
+ return self
312
+
313
+ def eval(self):
314
+ return self.train(False)
315
+
316
+ def call(self, n_id):
317
+ r"""Returns the memory and last update time of the nodes ``n_id``."""
318
+ if self.training_mode:
319
+ return self._get_updated_memory(np.asarray(ops.convert_to_numpy(n_id)).astype(np.int64))
320
+ return ops.take(self.memory, n_id, axis=0), ops.take(self.last_update, n_id, axis=0)
321
+
322
+ def update_state(self, src, dst, t, raw_msg):
323
+ r"""Updates the memory with the new events ``(src, dst, t, raw_msg)``."""
324
+ src, dst = (np.asarray(ops.convert_to_numpy(a)).astype(np.int64) for a in (src, dst))
325
+ t = np.asarray(ops.convert_to_numpy(t)).astype(np.float32)
326
+ raw_msg = np.asarray(ops.convert_to_numpy(raw_msg)).astype(np.float32)
327
+ n_id = np.unique(np.concatenate([src, dst]))
328
+ if self.training_mode:
329
+ self._update_memory(n_id)
330
+ self._update_msg_store(src, dst, t, raw_msg, self.msg_s_store)
331
+ self._update_msg_store(dst, src, t, raw_msg, self.msg_d_store)
332
+ else:
333
+ self._update_msg_store(src, dst, t, raw_msg, self.msg_s_store)
334
+ self._update_msg_store(dst, src, t, raw_msg, self.msg_d_store)
335
+ self._update_memory(n_id)
336
+
337
+ def detach(self):
338
+ r"""Kept for PyG compatibility: the stored memory is never differentiated."""
339
+
340
+ @staticmethod
341
+ def _update_msg_store(src, dst, t, raw_msg, store):
342
+ # every node keeps the messages of the latest batch it took part in
343
+ order = np.argsort(src, kind="stable")
344
+ nodes, starts = np.unique(src[order], return_index=True)
345
+ for node, idx in zip(nodes, np.split(order, starts[1:])):
346
+ store[int(node)] = (src[idx], dst[idx], t[idx], raw_msg[idx])
347
+
348
+ def _compute_msg(self, n_id, store, module):
349
+ src, dst, t, raw = (np.concatenate(parts) for parts in zip(*[store[int(i)] for i in n_id]))
350
+ t_rel = ops.convert_to_tensor(t) - ops.take(self.last_update, src, axis=0)
351
+ t_enc = self.time_enc(t_rel)
352
+ msg = module(ops.take(self.memory, src, axis=0), ops.take(self.memory, dst, axis=0),
353
+ ops.convert_to_tensor(raw), t_enc)
354
+ return msg, t, src
355
+
356
+ def _get_updated_memory(self, n_id):
357
+ assoc = np.full(self.num_nodes, -1, dtype=np.int64)
358
+ assoc[n_id] = np.arange(len(n_id))
359
+ msg_s, t_s, src_s = self._compute_msg(n_id, self.msg_s_store, self.msg_s_module)
360
+ msg_d, t_d, src_d = self._compute_msg(n_id, self.msg_d_store, self.msg_d_module)
361
+ idx = np.concatenate([src_s, src_d])
362
+ t = np.concatenate([t_s, t_d])
363
+ aggr = self.aggr_module(ops.concatenate([msg_s, msg_d], axis=0), assoc[idx], t, dim_size=len(n_id))
364
+ memory, _ = self.gru(aggr, [ops.take(self.memory, n_id, axis=0)])
365
+ latest = np.full(self.num_nodes, -np.inf, dtype=np.float32)
366
+ np.maximum.at(latest, idx, t)
367
+ has = np.isfinite(latest[n_id])
368
+ # as PyG: the last update time is the latest message time (0 for nodes without messages)
369
+ last = np.where(has, latest[n_id], 0.0).astype(np.float32)
370
+ return memory, ops.convert_to_tensor(last)
371
+
372
+ def _update_memory(self, n_id):
373
+ n_id = np.asarray(n_id).astype(np.int64)
374
+ if len(n_id) == 0:
375
+ return
376
+ memory, last_update = self._get_updated_memory(n_id)
377
+ memory_np = np.asarray(ops.convert_to_numpy(self.memory)).copy()
378
+ memory_np[n_id] = np.asarray(ops.convert_to_numpy(ops.stop_gradient(memory)))
379
+ self.memory.assign(memory_np)
380
+ last_np = np.asarray(ops.convert_to_numpy(self.last_update)).copy()
381
+ last_np[n_id] = np.asarray(ops.convert_to_numpy(last_update))
382
+ self.last_update.assign(last_np)