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
k3_node/models/pmlp.py ADDED
@@ -0,0 +1,157 @@
1
+ from typing import Optional
2
+
3
+ import keras
4
+ from keras import ops
5
+
6
+ from k3_node.layers.conv import SimpleConv
7
+ from k3_node.layers.norm import BatchNorm
8
+
9
+
10
+ class PMLP(keras.Model):
11
+ r"""The P(ropagational)MLP model from the `"Graph Neural Networks are
12
+ Inherently Good Generalizers: Insights by Bridging GNNs and MLPs"
13
+ <https://arxiv.org/abs/2212.09034>`_ paper.
14
+
15
+ :class:`PMLP` is identical to a standard MLP during training, but then
16
+ adopts a GNN architecture during testing.
17
+
18
+ Args:
19
+ in_channels (int): Size of each input sample.
20
+ hidden_channels (int): Size of each hidden sample.
21
+ out_channels (int): Size of each output sample.
22
+ num_layers (int): The number of layers.
23
+ dropout (float, optional): Dropout probability of each hidden
24
+ embedding. (default: :obj:`0.`)
25
+ norm (bool, optional): If set to :obj:`False`, will not apply batch
26
+ normalization. (default: :obj:`True`)
27
+ bias (bool, optional): If set to :obj:`False`, the module will not
28
+ learn additive biases. (default: :obj:`True`)
29
+
30
+ Example:
31
+ ```python
32
+ import numpy as np
33
+ from k3_node.models import PMLP
34
+
35
+ x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
36
+ edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
37
+
38
+ model = PMLP(in_channels=8, hidden_channels=32, out_channels=4, num_layers=2)
39
+ print(tuple(model(x, edge_index).shape)) # (10, 4): message passing is used at inference
40
+ ```
41
+ """
42
+
43
+ def __init__(
44
+ self,
45
+ in_channels: int,
46
+ hidden_channels: int,
47
+ out_channels: int,
48
+ num_layers: int,
49
+ dropout: float = 0.,
50
+ norm: bool = True,
51
+ bias: bool = True,
52
+ **kwargs,
53
+ ):
54
+ super().__init__(**kwargs)
55
+
56
+ self.in_channels = in_channels
57
+ self.hidden_channels = hidden_channels
58
+ self.out_channels = out_channels
59
+ self.num_layers = num_layers
60
+ self.dropout_rate = dropout
61
+ self.use_norm = norm
62
+ self.use_bias = bias
63
+ # Instance-level training flag matching PyG's API: pmlp.training = False
64
+ self.training = True
65
+
66
+ # Build weight dimensions
67
+ dims = [in_channels] + [hidden_channels] * (num_layers - 1) + [out_channels]
68
+ self._weight_shapes = [(dims[i], dims[i + 1]) for i in range(num_layers)]
69
+
70
+ # We store weights as raw keras Variables so they are always available
71
+ self._weights_list = []
72
+ self._biases_list = []
73
+ for i, (in_d, out_d) in enumerate(self._weight_shapes):
74
+ w = self.add_weight(
75
+ shape=(in_d, out_d),
76
+ initializer=keras.initializers.GlorotUniform(),
77
+ trainable=True,
78
+ name=f"weight_{i}",
79
+ )
80
+ self._weights_list.append(w)
81
+ if bias:
82
+ b = self.add_weight(
83
+ shape=(out_d,),
84
+ initializer="zeros",
85
+ trainable=True,
86
+ name=f"bias_{i}",
87
+ )
88
+ self._biases_list.append(b)
89
+ else:
90
+ self._biases_list.append(None)
91
+
92
+ self._norm_layer = None
93
+ if norm:
94
+ self._norm_layer = BatchNorm(
95
+ hidden_channels,
96
+ affine=False,
97
+ track_running_stats=False,
98
+ )
99
+
100
+ self.conv = SimpleConv(aggr='mean', combine_root='self_loop')
101
+ self._dropout = keras.layers.Dropout(dropout)
102
+
103
+ def build(self, input_shape=None):
104
+ self.built = True
105
+
106
+ def reset_parameters(self) -> None:
107
+ r"""Resets all learnable parameters of the module."""
108
+ for i, (in_d, out_d) in enumerate(self._weight_shapes):
109
+ self._weights_list[i].assign(
110
+ keras.initializers.GlorotUniform()(shape=(in_d, out_d))
111
+ )
112
+ if self.use_bias and self._biases_list[i] is not None:
113
+ self._biases_list[i].assign(ops.zeros((out_d,)))
114
+
115
+ def call(self, x, edge_index=None, training=None):
116
+ """Forward pass.
117
+
118
+ Args:
119
+ x (Tensor): The node features of shape ``[N, in_channels]``.
120
+ edge_index (Tensor, optional): The edge indices. Required during
121
+ inference. (default: :obj:`None`)
122
+ training (bool, optional): Override the instance-level
123
+ ``self.training`` flag. (default: :obj:`None`)
124
+ """
125
+ if edge_index is None and isinstance(x, (tuple, list)):
126
+ if len(x) >= 2:
127
+ x, edge_index = x[0], x[1]
128
+ # Respect both call-time kwarg and instance-level flag (PyG compat)
129
+ is_training = training if training is not None else self.training
130
+
131
+ if not is_training and edge_index is None:
132
+ raise ValueError(
133
+ f"'edge_index' needs to be present during inference "
134
+ f"in '{self.__class__.__name__}'"
135
+ )
136
+
137
+ for i in range(self.num_layers):
138
+ # Apply weight multiplication (like x @ W^T where W is [in, out])
139
+ x = x @ self._weights_list[i]
140
+
141
+ if not is_training:
142
+ x = self.conv(x, edge_index)
143
+
144
+ if self.use_bias and self._biases_list[i] is not None:
145
+ x = x + self._biases_list[i]
146
+
147
+ if i != self.num_layers - 1:
148
+ if self._norm_layer is not None:
149
+ x = self._norm_layer(x, training=is_training)
150
+ x = ops.relu(x)
151
+ x = self._dropout(x, training=is_training)
152
+
153
+ return x
154
+
155
+ def __repr__(self) -> str:
156
+ return (f'{self.__class__.__name__}({self.in_channels}, '
157
+ f'{self.out_channels}, num_layers={self.num_layers})')
@@ -0,0 +1,229 @@
1
+ from typing import Optional
2
+
3
+ import keras
4
+ from keras import ops
5
+
6
+ from k3_node.layers.aggr.base import from_dense_batch, to_dense_batch
7
+ import numpy as np
8
+
9
+ from k3_node.layers.conv import GATConv, GCNConv
10
+ from k3_node.layers.attention import PolynormerAttention
11
+
12
+
13
+ class Polynormer(keras.layers.Layer):
14
+ r"""The Polynormer module from the `"Polynormer: polynomial-expressive
15
+ graph transformer in linear time"
16
+ <https://arxiv.org/abs/2403.01232>`_ paper.
17
+
18
+ Args:
19
+ in_channels (int): Input channels.
20
+ hidden_channels (int): Hidden channels.
21
+ out_channels (int): Output channels.
22
+ local_layers (int): The number of local attention layers.
23
+ (default: :obj:`7`)
24
+ global_layers (int): The number of global attention layers.
25
+ (default: :obj:`2`)
26
+ in_dropout (float): Input dropout rate.
27
+ (default: :obj:`0.15`)
28
+ dropout (float): Dropout rate.
29
+ (default: :obj:`0.5`)
30
+ global_dropout (float): Global dropout rate.
31
+ (default: :obj:`0.5`)
32
+ heads (int): The number of heads.
33
+ (default: :obj:`1`)
34
+ beta (float): Aggregate type.
35
+ (default: :obj:`0.9`)
36
+ qk_shared (bool, optional): Whether weight of query and key are shared.
37
+ (default: :obj:`True`)
38
+ pre_ln (bool): Pre layer normalization.
39
+ (default: :obj:`False`)
40
+ post_bn (bool): Post batch normalization.
41
+ (default: :obj:`True`)
42
+ local_attn (bool): Whether use local attention (GATConv vs GCNConv).
43
+ (default: :obj:`False`)
44
+
45
+ Example:
46
+ ```python
47
+ import numpy as np
48
+ from k3_node.models import Polynormer
49
+
50
+ x = np.random.rand(10, 8).astype("float32") # 10 nodes with 8 features each
51
+ edge_index = np.random.randint(0, 10, size=(2, 30)) # 30 random edges
52
+
53
+ batch = np.repeat([0, 1], 5) # two graphs with 5 nodes each
54
+ model = Polynormer(in_channels=8, hidden_channels=32, out_channels=4,
55
+ local_layers=2, global_layers=1, heads=2)
56
+ out = model(x, edge_index, batch)
57
+ print(tuple(out.shape)) # (10, 4)
58
+ ```
59
+ """
60
+
61
+ def __init__(
62
+ self,
63
+ in_channels: int,
64
+ hidden_channels: int,
65
+ out_channels: int,
66
+ local_layers: int = 7,
67
+ global_layers: int = 2,
68
+ in_dropout: float = 0.15,
69
+ dropout: float = 0.5,
70
+ global_dropout: float = 0.5,
71
+ heads: int = 1,
72
+ beta: float = 0.9,
73
+ qk_shared: bool = False,
74
+ pre_ln: bool = False,
75
+ post_bn: bool = True,
76
+ local_attn: bool = False,
77
+ **kwargs,
78
+ ) -> None:
79
+ super().__init__(**kwargs)
80
+
81
+ self._global = False
82
+ self.in_drop = in_dropout
83
+ self.dropout = dropout
84
+ self.pre_ln = pre_ln
85
+ self.post_bn = post_bn
86
+ self.beta = beta
87
+ self.heads = heads
88
+ self.hidden_channels = hidden_channels
89
+ self.local_attn = local_attn
90
+
91
+ inner_channels = heads * hidden_channels
92
+
93
+ self.h_lins = []
94
+ self.local_convs = []
95
+ self.lins = []
96
+ self.lns = []
97
+ self.pre_lns = [] if pre_ln else None
98
+ self.post_bns = [] if post_bn else None
99
+
100
+ # ---- First local layer ----
101
+ self.h_lins.append(keras.layers.Dense(inner_channels))
102
+ if local_attn:
103
+ self.local_convs.append(
104
+ GATConv(in_channels, hidden_channels, heads=heads, concat=True,
105
+ add_self_loops=False, bias=False)
106
+ )
107
+ else:
108
+ self.local_convs.append(
109
+ GCNConv(in_channels, inner_channels, cached=False, normalize=True)
110
+ )
111
+ self.lins.append(keras.layers.Dense(inner_channels))
112
+ self.lns.append(keras.layers.LayerNormalization(epsilon=1e-5))
113
+ if pre_ln:
114
+ self.pre_lns.append(keras.layers.LayerNormalization(epsilon=1e-5))
115
+ if post_bn:
116
+ self.post_bns.append(
117
+ keras.layers.BatchNormalization(
118
+ center=True, scale=True, momentum=0.9, epsilon=1e-5,
119
+ )
120
+ )
121
+
122
+ # ---- Subsequent local layers ----
123
+ for _ in range(local_layers - 1):
124
+ self.h_lins.append(keras.layers.Dense(inner_channels))
125
+ if local_attn:
126
+ self.local_convs.append(
127
+ GATConv(inner_channels, hidden_channels, heads=heads,
128
+ concat=True, add_self_loops=False, bias=False)
129
+ )
130
+ else:
131
+ self.local_convs.append(
132
+ GCNConv(inner_channels, inner_channels, cached=False,
133
+ normalize=True)
134
+ )
135
+ self.lins.append(keras.layers.Dense(inner_channels))
136
+ self.lns.append(keras.layers.LayerNormalization(epsilon=1e-5))
137
+ if pre_ln:
138
+ self.pre_lns.append(keras.layers.LayerNormalization(epsilon=1e-5))
139
+ if post_bn:
140
+ self.post_bns.append(
141
+ keras.layers.BatchNormalization(
142
+ center=True, scale=True, momentum=0.9, epsilon=1e-5,
143
+ )
144
+ )
145
+
146
+ self.lin_in = keras.layers.Dense(inner_channels)
147
+ self.ln = keras.layers.LayerNormalization(epsilon=1e-5)
148
+
149
+ self.global_attn = [
150
+ PolynormerAttention(
151
+ channels=hidden_channels,
152
+ heads=heads,
153
+ head_channels=hidden_channels,
154
+ beta=beta,
155
+ dropout=global_dropout,
156
+ qk_shared=qk_shared,
157
+ )
158
+ for _ in range(global_layers)
159
+ ]
160
+
161
+ self.pred_local = keras.layers.Dense(out_channels)
162
+ self.pred_global = keras.layers.Dense(out_channels)
163
+
164
+ self._in_dropout = keras.layers.Dropout(in_dropout)
165
+ self._dropout = keras.layers.Dropout(dropout)
166
+
167
+ def build(self, input_shape=None):
168
+ self.built = True
169
+
170
+ def reset_parameters(self) -> None:
171
+ r"""Resets all learnable parameters of the module."""
172
+ # Keras layers reinitialize on next forward; no-op for unbuilt layers.
173
+ pass
174
+
175
+ def call(self, x, edge_index, batch: Optional[object] = None, training=None):
176
+ r"""Forward pass.
177
+
178
+ Args:
179
+ x (Tensor): The input node features.
180
+ edge_index (Tensor): The edge indices.
181
+ batch (Tensor, optional): The batch vector assigning each node to
182
+ a graph. (default: :obj:`None`)
183
+ training (bool, optional): Whether in training mode.
184
+ (default: :obj:`None`)
185
+ """
186
+ x = self._in_dropout(x, training=training)
187
+
188
+ # ---- Equivariant local attention ----
189
+ x_local = 0
190
+ for i, local_conv in enumerate(self.local_convs):
191
+ if self.pre_ln:
192
+ x = self.pre_lns[i](x)
193
+ h = self.h_lins[i](x)
194
+ h = ops.relu(h)
195
+ x = local_conv(x, edge_index) + self.lins[i](x)
196
+ if self.post_bn:
197
+ x = self.post_bns[i](x, training=training)
198
+ x = ops.relu(x)
199
+ x = self._dropout(x, training=training)
200
+ x = (1 - self.beta) * self.lns[i](h * x) + self.beta * x
201
+ x_local = x_local + x
202
+
203
+ # ---- Equivariant global attention ----
204
+ if self._global:
205
+ # Sort nodes by batch assignment (required by to_dense_batch)
206
+ batch_i = ops.cast(batch, "int32")
207
+ indices = ops.argsort(batch_i)
208
+ rev_perm = ops.argsort(indices)
209
+ batch_sorted = ops.take(batch_i, indices, axis=0)
210
+ x_local_sorted = self.ln(ops.take(x_local, indices, axis=0))
211
+
212
+ x_global, mask = to_dense_batch(x_local_sorted, batch_sorted)
213
+ for attn in self.global_attn:
214
+ x_global = attn(x_global, mask=mask, training=training)
215
+
216
+ # Flatten and undo the sort
217
+ x = ops.take(from_dense_batch(x_global, batch_sorted), rev_perm, axis=0)
218
+ x = self.pred_global(x)
219
+ else:
220
+ x = self.pred_local(x_local)
221
+
222
+ return ops.log_softmax(x, axis=-1)
223
+
224
+ def __repr__(self) -> str:
225
+ return (f'{self.__class__.__name__}('
226
+ f'in_channels={self.hidden_channels}, '
227
+ f'hidden_channels={self.hidden_channels}, '
228
+ f'heads={self.heads})')
229
+
k3_node/models/rect.py ADDED
@@ -0,0 +1,93 @@
1
+ from typing import Optional
2
+
3
+ import keras
4
+ from keras import ops
5
+
6
+ from k3_node.layers.conv import GCNConv
7
+ from k3_node.layers.conv.utils import scatter
8
+
9
+
10
+ class RECT_L(keras.layers.Layer):
11
+ r"""The RECT model, *i.e.* its supervised RECT-L part, from the
12
+ `"Network Embedding with Completely-imbalanced Labels"
13
+ <https://arxiv.org/abs/2007.03545>`_ paper.
14
+
15
+ Args:
16
+ in_channels (int): Size of each input sample.
17
+ hidden_channels (int): Intermediate size of each sample.
18
+ normalize (bool, optional): Whether to add self-loops and compute
19
+ symmetric normalization coefficients on-the-fly.
20
+ (default: :obj:`True`)
21
+ dropout (float, optional): The dropout probability.
22
+ (default: :obj:`0.0`)
23
+
24
+ Example:
25
+ ```python
26
+ import numpy as np
27
+ from k3_node.models import RECT_L
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
+ model = RECT_L(in_channels=8, hidden_channels=16)
33
+ out = model(x, edge_index) # reconstructs the (semantic) input features
34
+ print(tuple(out.shape)) # (10, 8)
35
+ print(tuple(model.embed(x, edge_index).shape)) # (10, 16): node embeddings
36
+ ```
37
+ """
38
+ def __init__(
39
+ self,
40
+ in_channels: int,
41
+ hidden_channels: int,
42
+ normalize: bool = True,
43
+ dropout: float = 0.0,
44
+ **kwargs,
45
+ ):
46
+ super().__init__(**kwargs)
47
+ self.in_channels = in_channels
48
+ self.hidden_channels = hidden_channels
49
+ self.dropout = dropout
50
+
51
+ self.conv = GCNConv(in_channels, hidden_channels, normalize=normalize)
52
+ self.lin = keras.layers.Dense(in_channels)
53
+ self._dropout = keras.layers.Dropout(dropout) if dropout > 0 else None
54
+
55
+ def build(self, input_shape=None):
56
+ self.built = True
57
+
58
+ def reset_parameters(self):
59
+ r"""Resets all learnable parameters of the module."""
60
+ self.conv.reset_parameters()
61
+ if self.lin.built:
62
+ self.lin.kernel.assign(
63
+ keras.initializers.GlorotUniform()(self.lin.kernel.shape)
64
+ )
65
+ if self.lin.bias is not None:
66
+ self.lin.bias.assign(ops.zeros(self.lin.bias.shape))
67
+
68
+ def call(self, x, edge_index, edge_weight=None, training=None):
69
+ x = self.conv(x, edge_index, edge_weight=edge_weight)
70
+ if self._dropout is not None:
71
+ x = self._dropout(x, training=training)
72
+ return self.lin(x)
73
+
74
+ def embed(self, x, edge_index, edge_weight=None):
75
+ return self.conv(x, edge_index, edge_weight=edge_weight)
76
+
77
+ def get_semantic_labels(self, x, y, mask):
78
+ r"""Replaces the original labels by their class-centers."""
79
+ mask_shape = ops.shape(mask)
80
+ if len(mask_shape) == 1 and 'bool' in str(mask.dtype):
81
+ y_sub = y[mask]
82
+ x_sub = x[mask]
83
+ else:
84
+ y_sub = ops.take(y, mask, axis=0)
85
+ x_sub = ops.take(x, mask, axis=0)
86
+
87
+ num_classes = int(ops.max(y_sub)) + 1
88
+ mean = scatter(x_sub, y_sub, dim=0, dim_size=num_classes, reduce='mean')
89
+ return ops.take(mean, y_sub, axis=0)
90
+
91
+ def __repr__(self) -> str:
92
+ return (f'{self.__class__.__name__}({self.in_channels}, '
93
+ f'{self.hidden_channels})')
@@ -0,0 +1,221 @@
1
+ from typing import Callable, List, Optional, Tuple
2
+ import math
3
+ import numpy as np
4
+ import keras
5
+ from keras import ops
6
+
7
+ from k3_node.layers.aggr import MeanAggregation
8
+ from k3_node.ops.creation import repeat
9
+
10
+
11
+ class RENet(keras.Model):
12
+ r"""The Recurrent Event Network model from the `"Recurrent Event Network
13
+ for Reasoning over Temporal Knowledge Graphs"
14
+ <https://arxiv.org/abs/1904.05530>`_ paper.
15
+
16
+ Args:
17
+ num_nodes (int): The number of nodes in the knowledge graph.
18
+ num_rels (int): The number of relations in the knowledge graph.
19
+ hidden_channels (int): Hidden size of node and relation embeddings.
20
+ seq_len (int): The sequence length of past events.
21
+ num_layers (int, optional): The number of recurrent layers.
22
+ (default: :obj:`1`)
23
+ dropout (float, optional): Dropout rate before final prediction.
24
+ (default: :obj:`0.0`)
25
+ bias (bool, optional): If set to :obj:`False`, all layers will not
26
+ learn an additive bias. (default: :obj:`True`)
27
+
28
+ Example:
29
+ ```python
30
+ import numpy as np
31
+ from k3_node.models import RENet
32
+
33
+ model = RENet(num_nodes=5, num_rels=4, hidden_channels=16, seq_len=3)
34
+ sub, rel, obj = np.array([0, 1]), np.array([0, 1]), np.array([2, 3]) # queries at time t
35
+ # Neighbor histories of subjects and objects: neighbor id, timestep and query index
36
+ h_sub, h_sub_t, h_sub_batch = np.array([0, 1, 2]), np.array([0, 1, 0]), np.array([0, 0, 1])
37
+ h_obj, h_obj_t, h_obj_batch = np.array([1, 2, 3]), np.array([1, 2, 0]), np.array([0, 0, 1])
38
+ log_prob_obj, log_prob_sub = model(sub, rel, obj, h_sub, h_sub_t, h_sub_batch, h_obj, h_obj_t, h_obj_batch)
39
+ print(tuple(log_prob_obj.shape)) # (2, 5): scores over all entities for each query
40
+ ```
41
+ """
42
+ def __init__(
43
+ self,
44
+ num_nodes: int,
45
+ num_rels: int,
46
+ hidden_channels: int,
47
+ seq_len: int,
48
+ num_layers: int = 1,
49
+ dropout: float = 0.0,
50
+ bias: bool = True,
51
+ **kwargs,
52
+ ):
53
+ super().__init__(**kwargs)
54
+
55
+ self.num_nodes = num_nodes
56
+ self.num_rels = num_rels
57
+ self.hidden_channels = hidden_channels
58
+ self.seq_len = seq_len
59
+ self.dropout_rate = dropout
60
+ self.num_layers = num_layers
61
+
62
+ self.ent = self.add_weight(
63
+ name="ent",
64
+ shape=(num_nodes, hidden_channels),
65
+ initializer=keras.initializers.GlorotUniform(),
66
+ )
67
+ self.rel = self.add_weight(
68
+ name="rel",
69
+ shape=(num_rels, hidden_channels),
70
+ initializer=keras.initializers.GlorotUniform(),
71
+ )
72
+
73
+ self.sub_gru = keras.layers.GRU(
74
+ hidden_channels,
75
+ return_sequences=False,
76
+ use_bias=bias,
77
+ )
78
+ self.obj_gru = keras.layers.GRU(
79
+ hidden_channels,
80
+ return_sequences=False,
81
+ use_bias=bias,
82
+ )
83
+
84
+ self.sub_lin = keras.layers.Dense(num_nodes, use_bias=bias)
85
+ self.obj_lin = keras.layers.Dense(num_nodes, use_bias=bias)
86
+ self.drop = keras.layers.Dropout(dropout)
87
+ self.mean_aggr = MeanAggregation()
88
+
89
+ def reset_parameters(self):
90
+ self.ent.assign(
91
+ keras.initializers.GlorotUniform()(self.ent.shape)
92
+ )
93
+ self.rel.assign(
94
+ keras.initializers.GlorotUniform()(self.rel.shape)
95
+ )
96
+
97
+ def build(self, input_shape=None):
98
+ self.built = True
99
+
100
+
101
+ def call(
102
+ self,
103
+ sub,
104
+ rel,
105
+ obj,
106
+ h_sub,
107
+ h_sub_t,
108
+ h_sub_batch,
109
+ h_obj,
110
+ h_obj_t,
111
+ h_obj_batch,
112
+ training=False,
113
+ ):
114
+ batch_size = ops.shape(sub)[0]
115
+ seq_len = self.seq_len
116
+
117
+ h_sub_t = h_sub_t + h_sub_batch * seq_len
118
+ h_obj_t = h_obj_t + h_obj_batch * seq_len
119
+
120
+ ent_h_sub = ops.take(self.ent, h_sub, axis=0)
121
+ ent_h_obj = ops.take(self.ent, h_obj, axis=0)
122
+
123
+ h_sub_scatter = self.mean_aggr(
124
+ ent_h_sub, index=h_sub_t, dim_size=batch_size * seq_len, dim=0
125
+ )
126
+ h_sub = ops.reshape(h_sub_scatter, (-1, seq_len, self.hidden_channels)) # static feature size
127
+
128
+ h_obj_scatter = self.mean_aggr(
129
+ ent_h_obj, index=h_obj_t, dim_size=batch_size * seq_len, dim=0
130
+ )
131
+ h_obj = ops.reshape(h_obj_scatter, (-1, seq_len, self.hidden_channels)) # static feature size
132
+
133
+ sub_emb = ops.take(self.ent, sub, axis=0)
134
+ rel_emb = ops.take(self.rel, rel, axis=0)
135
+ obj_emb = ops.take(self.ent, obj, axis=0)
136
+
137
+ sub_rep = repeat(ops.expand_dims(sub_emb, 1), seq_len, axis=1)
138
+ rel_rep = repeat(ops.expand_dims(rel_emb, 1), seq_len, axis=1)
139
+ obj_rep = repeat(ops.expand_dims(obj_emb, 1), seq_len, axis=1)
140
+
141
+ gru_sub_in = ops.concatenate([sub_rep, h_sub, rel_rep], axis=-1)
142
+ gru_obj_in = ops.concatenate([obj_rep, h_obj, rel_rep], axis=-1)
143
+
144
+ h_sub = self.sub_gru(gru_sub_in, training=training)
145
+ h_obj = self.obj_gru(gru_obj_in, training=training)
146
+
147
+ h_sub = ops.concatenate([sub_emb, h_sub, rel_emb], axis=-1)
148
+ h_obj = ops.concatenate([obj_emb, h_obj, rel_emb], axis=-1)
149
+
150
+ h_sub = self.drop(h_sub, training=training)
151
+ h_obj = self.drop(h_obj, training=training)
152
+
153
+ log_prob_obj = ops.log_softmax(self.sub_lin(h_sub), axis=-1)
154
+ log_prob_sub = ops.log_softmax(self.obj_lin(h_obj), axis=-1)
155
+
156
+ return log_prob_obj, log_prob_sub
157
+
158
+ @staticmethod
159
+ def pre_transform(seq_len: int) -> Callable:
160
+ r"""Returns a pre-transform that adds to every event (processed in time order) the history
161
+ of its subject and object: the entities they were linked to by the same relation in each of
162
+ the last ``seq_len`` time steps (``h_sub`` / ``h_obj``, with the step in ``h_sub_t`` /
163
+ ``h_obj_t``), as in PyG."""
164
+
165
+ class PreTransform:
166
+ def __init__(self, seq_len):
167
+ self.seq_len = seq_len
168
+ self.t_last = 0
169
+ self.sub_hist, self.obj_hist = {}, {} # node -> list of seq_len + 1 steps of (node, rel)
170
+
171
+ def _hist(self, hist, node):
172
+ if node not in hist:
173
+ hist[node] = [[] for _ in range(self.seq_len + 1)]
174
+ return hist[node]
175
+
176
+ def _history(self, hist, node, rel):
177
+ steps = self._hist(hist, node)
178
+ nodes, ts = [], []
179
+ for s in range(self.seq_len):
180
+ for other, r in steps[s]:
181
+ if r == rel:
182
+ nodes.append(other)
183
+ ts.append(s)
184
+ return np.array(nodes, dtype=np.int64), np.array(ts, dtype=np.int64)
185
+
186
+ def __call__(self, data):
187
+ sub, rel, obj, t = int(data.sub), int(data.rel), int(data.obj), int(data.t)
188
+ if t > self.t_last: # a new time step: forget the oldest one
189
+ for hist in (self.sub_hist, self.obj_hist):
190
+ for steps in hist.values():
191
+ steps.pop(0)
192
+ steps.append([])
193
+ self.t_last = t
194
+ data.h_sub, data.h_sub_t = self._history(self.sub_hist, sub, rel)
195
+ data.h_obj, data.h_obj_t = self._history(self.obj_hist, obj, rel)
196
+ self._hist(self.sub_hist, sub)[-1].append((obj, rel))
197
+ self._hist(self.obj_hist, obj)[-1].append((sub, rel))
198
+ return data
199
+
200
+ def __repr__(self):
201
+ return f"{self.__class__.__name__}(seq_len={self.seq_len})"
202
+
203
+ return PreTransform(seq_len)
204
+
205
+ def test(self, logits, y):
206
+ r"""Given ground-truth :obj:`y`, computes Mean Reciprocal Rank (MRR)
207
+ and Hits at 1/3/10.
208
+ """
209
+ logits_np = ops.convert_to_numpy(logits)
210
+ y_np = ops.convert_to_numpy(y).reshape(-1, 1)
211
+
212
+ perm = np.argsort(-logits_np, axis=1)
213
+ mask = y_np == perm
214
+
215
+ rows, cols = np.nonzero(mask)
216
+ mrr = float(np.mean(1.0 / (cols + 1.0)))
217
+ hits1 = float(np.sum(cols < 1) / len(y_np))
218
+ hits3 = float(np.sum(cols < 3) / len(y_np))
219
+ hits10 = float(np.sum(cols < 10) / len(y_np))
220
+
221
+ return ops.convert_to_tensor([mrr, hits1, hits3, hits10], dtype="float32")