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,255 @@
1
+ from typing import Tuple
2
+
3
+ import numpy as np
4
+ import keras
5
+ from keras import ops
6
+
7
+ from k3_node.layers.kge.loader import KGTripletLoader
8
+ from k3_node.training import no_grad
9
+
10
+ try:
11
+ from tqdm import tqdm
12
+ except ImportError: # pragma: no cover
13
+ def tqdm(iterable, *args, **kwargs):
14
+ return iterable
15
+
16
+
17
+ def normalize(x, p: float = 2.0, axis: int = -1, eps: float = 1e-12):
18
+ r"""Lp-normalizes `x` along `axis`, mirroring `torch.nn.functional.normalize`."""
19
+ norm = ops.power(ops.sum(ops.power(ops.abs(x), p), axis=axis, keepdims=True), 1.0 / p)
20
+ return x / ops.maximum(norm, eps)
21
+
22
+
23
+ def margin_ranking_loss(pos_score, neg_score, margin: float = 1.0):
24
+ r"""Mirrors `torch.nn.functional.margin_ranking_loss` with `target=1`."""
25
+ return ops.mean(ops.relu(margin - pos_score + neg_score))
26
+
27
+
28
+ def binary_cross_entropy_with_logits(logits, target):
29
+ r"""Numerically stable sigmoid cross-entropy, mirroring
30
+ `torch.nn.functional.binary_cross_entropy_with_logits`."""
31
+ zeros = ops.zeros_like(logits)
32
+ cond = logits >= zeros
33
+ relu_logits = ops.where(cond, logits, zeros)
34
+ neg_abs_logits = ops.where(cond, -logits, logits)
35
+ loss = relu_logits - logits * target + ops.log(1.0 + ops.exp(neg_abs_logits))
36
+ return ops.mean(loss)
37
+
38
+
39
+ class KGEModel(keras.layers.Layer):
40
+ r"""An abstract base class for implementing custom KGE models.
41
+
42
+ Args:
43
+ num_nodes (int): The number of nodes/entities in the graph.
44
+ num_relations (int): The number of relations in the graph.
45
+ hidden_channels (int): The hidden embedding size.
46
+ sparse (bool, optional): Kept for API compatibility with PyG; has no
47
+ effect since Keras optimizers do not distinguish sparse
48
+ embedding gradients the way PyTorch does. (default: :obj:`False`)
49
+ """
50
+ def __init__(
51
+ self,
52
+ num_nodes: int,
53
+ num_relations: int,
54
+ hidden_channels: int,
55
+ sparse: bool = False,
56
+ **kwargs,
57
+ ):
58
+ super().__init__(**kwargs)
59
+
60
+ self.num_nodes = num_nodes
61
+ self.num_relations = num_relations
62
+ self.hidden_channels = hidden_channels
63
+ self.sparse = sparse
64
+
65
+ self.node_emb = keras.layers.Embedding(num_nodes, hidden_channels)
66
+ self.rel_emb = keras.layers.Embedding(num_relations, hidden_channels)
67
+ self.node_emb.build((None,))
68
+ self.rel_emb.build((None,))
69
+
70
+ def reset_parameters(self):
71
+ r"""Resets all learnable parameters of the module."""
72
+ self.node_emb.embeddings.assign(
73
+ self.node_emb.embeddings_initializer(ops.shape(self.node_emb.embeddings))
74
+ )
75
+ self.rel_emb.embeddings.assign(
76
+ self.rel_emb.embeddings_initializer(ops.shape(self.rel_emb.embeddings))
77
+ )
78
+
79
+ def call(self, head_index, rel_type, tail_index):
80
+ r"""Returns the score for the given triplet.
81
+
82
+ Args:
83
+ head_index: The head indices.
84
+ rel_type: The relation type.
85
+ tail_index: The tail indices.
86
+ """
87
+ raise NotImplementedError
88
+
89
+ def loss(self, head_index, rel_type, tail_index):
90
+ r"""Returns the loss value for the given triplet."""
91
+ raise NotImplementedError
92
+
93
+ def loader(self, head_index, rel_type, tail_index, **kwargs):
94
+ r"""Returns a mini-batch loader that samples a subset of triplets.
95
+
96
+ Args:
97
+ head_index: The head indices.
98
+ rel_type: The relation type.
99
+ tail_index: The tail indices.
100
+ **kwargs (optional): Additional arguments of
101
+ :class:`k3_node.layers.kge.KGTripletLoader`, such as
102
+ `batch_size`, `shuffle` or `drop_last`.
103
+ """
104
+ return KGTripletLoader(head_index, rel_type, tail_index, **kwargs)
105
+
106
+ def test(
107
+ self,
108
+ head_index,
109
+ rel_type,
110
+ tail_index,
111
+ batch_size: int,
112
+ k: int = 10,
113
+ log: bool = True,
114
+ ) -> Tuple[float, float, float]:
115
+ r"""Evaluates the model quality by computing Mean Rank, MRR and
116
+ Hits@:math:`k` across all possible tail entities.
117
+
118
+ Args:
119
+ head_index: The head indices.
120
+ rel_type: The relation type.
121
+ tail_index: The tail indices.
122
+ batch_size (int): The batch size to use for evaluating.
123
+ k (int, optional): The :math:`k` in Hits @ :math:`k`.
124
+ (default: :obj:`10`)
125
+ log (bool, optional): If set to :obj:`False`, will not print a
126
+ progress bar to the console. (default: :obj:`True`)
127
+ """
128
+ head_index = ops.convert_to_numpy(head_index)
129
+ rel_type = ops.convert_to_numpy(rel_type)
130
+ tail_index = ops.convert_to_numpy(tail_index)
131
+
132
+ arange = range(head_index.shape[0])
133
+ arange = tqdm(arange) if log else arange
134
+
135
+ mean_ranks, reciprocal_ranks, hits_at_k = [], [], []
136
+ for i in arange:
137
+ h, r, t = int(head_index[i]), int(rel_type[i]), int(tail_index[i])
138
+
139
+ scores = []
140
+ tail_indices = np.arange(self.num_nodes)
141
+ for start in range(0, self.num_nodes, batch_size):
142
+ ts = tail_indices[start:start + batch_size]
143
+ hs = np.full_like(ts, h)
144
+ rs = np.full_like(ts, r)
145
+ with no_grad(): # as PyG's @torch.no_grad()
146
+ out = self(
147
+ ops.convert_to_tensor(hs),
148
+ ops.convert_to_tensor(rs),
149
+ ops.convert_to_tensor(ts),
150
+ )
151
+ scores.append(ops.convert_to_numpy(out))
152
+ scores = np.concatenate(scores)
153
+ rank = int(np.nonzero(np.argsort(-scores) == t)[0][0])
154
+
155
+ mean_ranks.append(rank)
156
+ reciprocal_ranks.append(1.0 / (rank + 1))
157
+ hits_at_k.append(rank < k)
158
+
159
+ mean_rank = float(np.mean(mean_ranks))
160
+ mrr = float(np.mean(reciprocal_ranks))
161
+ hits_at_k = float(np.mean(hits_at_k))
162
+
163
+ return mean_rank, mrr, hits_at_k
164
+
165
+ # ---- Keras-style training -------------------------------------------------------------------
166
+ def compile(self, optimizer):
167
+ r"""Sets the optimizer used by :meth:`fit`."""
168
+ self.optimizer = optimizer
169
+
170
+ def fit(self, data, epochs: int = 1, batch_size: int = 1000, validation_data=None,
171
+ validation_batch_size: int = 20000, verbose: int = 1):
172
+ r"""Trains on the triplets of ``data`` (``edge_index`` holds heads and tails, ``edge_type``
173
+ the relations), in shuffled mini-batches.
174
+
175
+ Args:
176
+ data (Data): The training knowledge graph.
177
+ epochs (int): Passes over all triplets. (default: ``1``)
178
+ batch_size (int): Triplets per gradient step. (default: ``1000``)
179
+ validation_data (Data, optional): Evaluated with :meth:`evaluate` after every epoch
180
+ (this ranks all entities for every triplet, so keep it small).
181
+ validation_batch_size (int): Entities scored at once during validation.
182
+ verbose (int): ``0`` is silent, otherwise one line is printed per epoch.
183
+
184
+ Returns:
185
+ dict: The mean loss (and validation metrics) of every epoch.
186
+ """
187
+ from k3_node.training import gradient_step
188
+
189
+ if getattr(self, "optimizer", None) is None:
190
+ raise ValueError("Call `compile(optimizer=...)` before `fit`.")
191
+ head, tail = data.edge_index[0], data.edge_index[1]
192
+ loader = self.loader(head, data.edge_type, tail, batch_size=batch_size, shuffle=True)
193
+ history = {"loss": []}
194
+ for epoch in range(1, epochs + 1):
195
+ total = count = 0
196
+ for h, r, t in loader:
197
+ loss = gradient_step(lambda: self.loss(h, r, t), self.trainable_variables, self.optimizer)
198
+ total += loss * int(h.shape[0])
199
+ count += int(h.shape[0])
200
+ logs = {"loss": total / count}
201
+ if validation_data is not None:
202
+ logs.update({f"val_{k}": v for k, v in
203
+ self.evaluate(validation_data, batch_size=validation_batch_size).items()})
204
+ for key, value in logs.items():
205
+ history.setdefault(key, []).append(value)
206
+ if verbose:
207
+ print(f"Epoch {epoch:03d}: " + ", ".join(f"{k}: {v:.4f}" for k, v in logs.items()))
208
+ return history
209
+
210
+ def evaluate(self, data, batch_size: int = 20000, k: int = 10):
211
+ r"""Ranks the true tail of every triplet in ``data`` among all entities and returns the
212
+ mean rank, the mean reciprocal rank (MRR) and Hits@``k``."""
213
+ mean_rank, mrr, hits = self.test(data.edge_index[0], data.edge_type, data.edge_index[1],
214
+ batch_size=batch_size, k=k, log=False)
215
+ return {"mean_rank": mean_rank, "mrr": mrr, f"hits@{k}": hits}
216
+
217
+ def random_sample(
218
+ self,
219
+ head_index,
220
+ rel_type,
221
+ tail_index,
222
+ ):
223
+ r"""Randomly samples negative triplets by either replacing the head or
224
+ the tail (but not both).
225
+
226
+ Args:
227
+ head_index: The head indices.
228
+ rel_type: The relation type.
229
+ tail_index: The tail indices.
230
+ """
231
+ num_triplets = ops.shape(head_index)[0]
232
+ num_negatives = num_triplets // 2
233
+
234
+ rnd_index = keras.random.randint(
235
+ ops.shape(head_index), 0, self.num_nodes, dtype="int32"
236
+ )
237
+ rnd_index = ops.cast(rnd_index, head_index.dtype)
238
+
239
+ head_index = ops.concatenate(
240
+ [rnd_index[:num_negatives], head_index[num_negatives:]], axis=0
241
+ )
242
+ tail_index = ops.concatenate(
243
+ [tail_index[:num_negatives], rnd_index[num_negatives:]], axis=0
244
+ )
245
+
246
+ return head_index, rel_type, tail_index
247
+
248
+ def __repr__(self) -> str:
249
+ return (
250
+ f"{self.__class__.__name__}({self.num_nodes}, "
251
+ f"num_relations={self.num_relations}, "
252
+ f"hidden_channels={self.hidden_channels})"
253
+ )
254
+
255
+ __str__ = __repr__
@@ -0,0 +1,98 @@
1
+ import keras
2
+ from keras import ops
3
+
4
+ from k3_node.layers.kge.base import KGEModel, binary_cross_entropy_with_logits
5
+
6
+
7
+ def triple_dot(x, y, z):
8
+ return ops.sum(x * y * z, axis=-1)
9
+
10
+
11
+ class ComplEx(KGEModel):
12
+ r"""The ComplEx model from the `"Complex Embeddings for Simple Link
13
+ Prediction" <https://arxiv.org/abs/1606.06357>`_ paper.
14
+
15
+ :class:`ComplEx` models relations as complex-valued bilinear mappings
16
+ between head and tail entities using the Hermetian dot product.
17
+ The entities and relations are embedded in different dimensional spaces,
18
+ resulting in the scoring function:
19
+
20
+ .. math::
21
+ d(h, r, t) = Re(< \mathbf{e}_h, \mathbf{e}_r, \mathbf{e}_t>)
22
+
23
+ Args:
24
+ num_nodes (int): The number of nodes/entities in the graph.
25
+ num_relations (int): The number of relations in the graph.
26
+ hidden_channels (int): The hidden embedding size.
27
+ sparse (bool, optional): Kept for API compatibility. (default: :obj:`False`)
28
+
29
+ Example:
30
+ ```python
31
+ import numpy as np
32
+ from k3_node.layers import ComplEx
33
+
34
+ head = np.random.randint(0, 20, size=(10,)) # 10 (head, relation, tail) triples
35
+ rel = np.random.randint(0, 5, size=(10,))
36
+ tail = np.random.randint(0, 20, size=(10,))
37
+
38
+ model = ComplEx(num_nodes=20, num_relations=5, hidden_channels=8)
39
+ score = model(head, rel, tail) # plausibility score of every triple
40
+ print(tuple(score.shape)) # (10,)
41
+ loss = model.loss(head, rel, tail) # training loss against randomly corrupted triples
42
+ print(tuple(loss.shape)) # (): a scalar
43
+ ```
44
+ """
45
+ def __init__(
46
+ self,
47
+ num_nodes: int,
48
+ num_relations: int,
49
+ hidden_channels: int,
50
+ sparse: bool = False,
51
+ **kwargs,
52
+ ):
53
+ super().__init__(num_nodes, num_relations, hidden_channels, sparse, **kwargs)
54
+
55
+ self.node_emb_im = keras.layers.Embedding(num_nodes, hidden_channels)
56
+ self.rel_emb_im = keras.layers.Embedding(num_relations, hidden_channels)
57
+ self.node_emb_im.build((None,))
58
+ self.rel_emb_im.build((None,))
59
+
60
+ self.reset_parameters()
61
+
62
+ def reset_parameters(self):
63
+ # A new initializer per tensor: a reused unseeded Keras 3 initializer returns the same values on every call.
64
+ glorot = lambda shape: keras.initializers.GlorotUniform()(shape)
65
+ self.node_emb.embeddings.assign(glorot(ops.shape(self.node_emb.embeddings)))
66
+ self.node_emb_im.embeddings.assign(glorot(ops.shape(self.node_emb_im.embeddings)))
67
+ self.rel_emb.embeddings.assign(glorot(ops.shape(self.rel_emb.embeddings)))
68
+ self.rel_emb_im.embeddings.assign(glorot(ops.shape(self.rel_emb_im.embeddings)))
69
+
70
+ def call(self, head_index, rel_type, tail_index):
71
+ head_index = ops.cast(head_index, "int32")
72
+ rel_type = ops.cast(rel_type, "int32")
73
+ tail_index = ops.cast(tail_index, "int32")
74
+
75
+ head_re = self.node_emb(head_index)
76
+ head_im = self.node_emb_im(head_index)
77
+ rel_re = self.rel_emb(rel_type)
78
+ rel_im = self.rel_emb_im(rel_type)
79
+ tail_re = self.node_emb(tail_index)
80
+ tail_im = self.node_emb_im(tail_index)
81
+
82
+ return (
83
+ triple_dot(head_re, rel_re, tail_re)
84
+ + triple_dot(head_im, rel_re, tail_im)
85
+ + triple_dot(head_re, rel_im, tail_im)
86
+ - triple_dot(head_im, rel_im, tail_re)
87
+ )
88
+
89
+ def loss(self, head_index, rel_type, tail_index):
90
+ pos_score = self(head_index, rel_type, tail_index)
91
+ neg_score = self(*self.random_sample(head_index, rel_type, tail_index))
92
+ scores = ops.concatenate([pos_score, neg_score], axis=0)
93
+
94
+ pos_target = ops.ones_like(pos_score)
95
+ neg_target = ops.zeros_like(neg_score)
96
+ target = ops.concatenate([pos_target, neg_target], axis=0)
97
+
98
+ return binary_cross_entropy_with_logits(scores, target)
@@ -0,0 +1,79 @@
1
+ import keras
2
+ from keras import ops
3
+
4
+ from k3_node.layers.kge.base import KGEModel, margin_ranking_loss
5
+
6
+
7
+ class DistMult(KGEModel):
8
+ r"""The DistMult model from the `"Embedding Entities and Relations for
9
+ Learning and Inference in Knowledge Bases"
10
+ <https://arxiv.org/abs/1412.6575>`_ paper.
11
+
12
+ :class:`DistMult` models relations as diagonal matrices, which simplifies
13
+ the bi-linear interaction between the head and tail entities to the score
14
+ function:
15
+
16
+ .. math::
17
+ d(h, r, t) = < \mathbf{e}_h, \mathbf{e}_r, \mathbf{e}_t >
18
+
19
+ Args:
20
+ num_nodes (int): The number of nodes/entities in the graph.
21
+ num_relations (int): The number of relations in the graph.
22
+ hidden_channels (int): The hidden embedding size.
23
+ margin (float, optional): The margin of the ranking loss.
24
+ (default: :obj:`1.0`)
25
+ sparse (bool, optional): Kept for API compatibility. (default: :obj:`False`)
26
+
27
+ Example:
28
+ ```python
29
+ import numpy as np
30
+ from k3_node.layers import DistMult
31
+
32
+ head = np.random.randint(0, 20, size=(10,)) # 10 (head, relation, tail) triples
33
+ rel = np.random.randint(0, 5, size=(10,))
34
+ tail = np.random.randint(0, 20, size=(10,))
35
+
36
+ model = DistMult(num_nodes=20, num_relations=5, hidden_channels=8)
37
+ score = model(head, rel, tail) # plausibility score of every triple
38
+ print(tuple(score.shape)) # (10,)
39
+ loss = model.loss(head, rel, tail) # training loss against randomly corrupted triples
40
+ print(tuple(loss.shape)) # (): a scalar
41
+ ```
42
+ """
43
+ def __init__(
44
+ self,
45
+ num_nodes: int,
46
+ num_relations: int,
47
+ hidden_channels: int,
48
+ margin: float = 1.0,
49
+ sparse: bool = False,
50
+ **kwargs,
51
+ ):
52
+ super().__init__(num_nodes, num_relations, hidden_channels, sparse, **kwargs)
53
+
54
+ self.margin = margin
55
+
56
+ self.reset_parameters()
57
+
58
+ def reset_parameters(self):
59
+ # A new initializer per tensor: a reused unseeded Keras 3 initializer returns the same values on every call.
60
+ glorot = lambda shape: keras.initializers.GlorotUniform()(shape)
61
+ self.node_emb.embeddings.assign(glorot(ops.shape(self.node_emb.embeddings)))
62
+ self.rel_emb.embeddings.assign(glorot(ops.shape(self.rel_emb.embeddings)))
63
+
64
+ def call(self, head_index, rel_type, tail_index):
65
+ head_index = ops.cast(head_index, "int32")
66
+ rel_type = ops.cast(rel_type, "int32")
67
+ tail_index = ops.cast(tail_index, "int32")
68
+
69
+ head = self.node_emb(head_index)
70
+ rel = self.rel_emb(rel_type)
71
+ tail = self.node_emb(tail_index)
72
+
73
+ return ops.sum(head * rel * tail, axis=-1)
74
+
75
+ def loss(self, head_index, rel_type, tail_index):
76
+ pos_score = self(head_index, rel_type, tail_index)
77
+ neg_score = self(*self.random_sample(head_index, rel_type, tail_index))
78
+
79
+ return margin_ranking_loss(pos_score, neg_score, margin=self.margin)
@@ -0,0 +1,50 @@
1
+ import math
2
+
3
+ import numpy as np
4
+ from keras import ops
5
+
6
+
7
+ class KGTripletLoader:
8
+ r"""A minimal, framework-agnostic batching iterator over knowledge-graph
9
+ triplets, mirroring :class:`torch.utils.data.DataLoader` usage in
10
+ :meth:`k3_node.layers.kge.KGEModel.loader`.
11
+
12
+ Args:
13
+ head_index: The head indices.
14
+ rel_type: The relation type.
15
+ tail_index: The tail indices.
16
+ batch_size (int, optional): The batch size. (default: :obj:`1`)
17
+ shuffle (bool, optional): If set to :obj:`True`, shuffles the
18
+ triplets at every epoch. (default: :obj:`False`)
19
+ drop_last (bool, optional): If set to :obj:`True`, drops the last
20
+ incomplete batch. (default: :obj:`False`)
21
+ """
22
+ def __init__(self, head_index, rel_type, tail_index, batch_size: int = 1,
23
+ shuffle: bool = False, drop_last: bool = False, **kwargs):
24
+ self.head_index = ops.convert_to_numpy(head_index)
25
+ self.rel_type = ops.convert_to_numpy(rel_type)
26
+ self.tail_index = ops.convert_to_numpy(tail_index)
27
+ self.batch_size = batch_size
28
+ self.shuffle = shuffle
29
+ self.drop_last = drop_last
30
+ self.num_triplets = self.head_index.shape[0]
31
+
32
+ def __len__(self):
33
+ if self.drop_last:
34
+ return self.num_triplets // self.batch_size
35
+ return math.ceil(self.num_triplets / self.batch_size)
36
+
37
+ def __iter__(self):
38
+ indices = np.arange(self.num_triplets)
39
+ if self.shuffle:
40
+ np.random.shuffle(indices)
41
+
42
+ for start in range(0, self.num_triplets, self.batch_size):
43
+ batch_idx = indices[start:start + self.batch_size]
44
+ if self.drop_last and batch_idx.shape[0] < self.batch_size:
45
+ continue
46
+ yield (
47
+ ops.convert_to_tensor(self.head_index[batch_idx]),
48
+ ops.convert_to_tensor(self.rel_type[batch_idx]),
49
+ ops.convert_to_tensor(self.tail_index[batch_idx]),
50
+ )
@@ -0,0 +1,103 @@
1
+ import math
2
+
3
+ import keras
4
+ from keras import ops
5
+
6
+ from k3_node.layers.kge.base import KGEModel, binary_cross_entropy_with_logits
7
+
8
+
9
+ class RotatE(KGEModel):
10
+ r"""The RotatE model from the `"RotatE: Knowledge Graph Embedding by
11
+ Relational Rotation in Complex Space" <https://arxiv.org/abs/
12
+ 1902.10197>`_ paper.
13
+
14
+ :class:`RotatE` models relations as a rotation in complex space
15
+ from head to tail such that
16
+
17
+ .. math::
18
+ \mathbf{e}_t = \mathbf{e}_h \circ \mathbf{e}_r,
19
+
20
+ resulting in the scoring function
21
+
22
+ .. math::
23
+ d(h, r, t) = - {\| \mathbf{e}_h \circ \mathbf{e}_r - \mathbf{e}_t \|}_p
24
+
25
+ Args:
26
+ num_nodes (int): The number of nodes/entities in the graph.
27
+ num_relations (int): The number of relations in the graph.
28
+ hidden_channels (int): The hidden embedding size.
29
+ margin (float, optional): The margin of the ranking loss.
30
+ (default: :obj:`1.0`)
31
+ sparse (bool, optional): Kept for API compatibility. (default: :obj:`False`)
32
+
33
+ Example:
34
+ ```python
35
+ import numpy as np
36
+ from k3_node.layers import RotatE
37
+
38
+ head = np.random.randint(0, 20, size=(10,)) # 10 (head, relation, tail) triples
39
+ rel = np.random.randint(0, 5, size=(10,))
40
+ tail = np.random.randint(0, 20, size=(10,))
41
+
42
+ model = RotatE(num_nodes=20, num_relations=5, hidden_channels=8)
43
+ score = model(head, rel, tail) # plausibility score of every triple
44
+ print(tuple(score.shape)) # (10,)
45
+ loss = model.loss(head, rel, tail) # training loss against randomly corrupted triples
46
+ print(tuple(loss.shape)) # (): a scalar
47
+ ```
48
+ """
49
+ def __init__(
50
+ self,
51
+ num_nodes: int,
52
+ num_relations: int,
53
+ hidden_channels: int,
54
+ margin: float = 1.0,
55
+ sparse: bool = False,
56
+ **kwargs,
57
+ ):
58
+ super().__init__(num_nodes, num_relations, hidden_channels, sparse, **kwargs)
59
+
60
+ self.margin = margin
61
+ self.node_emb_im = keras.layers.Embedding(num_nodes, hidden_channels)
62
+ self.node_emb_im.build((None,))
63
+
64
+ self.reset_parameters()
65
+
66
+ def reset_parameters(self):
67
+ # A new initializer per tensor: a reused unseeded Keras 3 initializer returns the same values on every call.
68
+ glorot = lambda shape: keras.initializers.GlorotUniform()(shape)
69
+ self.node_emb.embeddings.assign(glorot(ops.shape(self.node_emb.embeddings)))
70
+ self.node_emb_im.embeddings.assign(glorot(ops.shape(self.node_emb_im.embeddings)))
71
+ uniform = keras.initializers.RandomUniform(0, 2 * math.pi)
72
+ self.rel_emb.embeddings.assign(uniform(ops.shape(self.rel_emb.embeddings)))
73
+
74
+ def call(self, head_index, rel_type, tail_index):
75
+ head_index = ops.cast(head_index, "int32")
76
+ rel_type = ops.cast(rel_type, "int32")
77
+ tail_index = ops.cast(tail_index, "int32")
78
+
79
+ head_re = self.node_emb(head_index)
80
+ head_im = self.node_emb_im(head_index)
81
+ tail_re = self.node_emb(tail_index)
82
+ tail_im = self.node_emb_im(tail_index)
83
+
84
+ rel_theta = self.rel_emb(rel_type)
85
+ rel_re, rel_im = ops.cos(rel_theta), ops.sin(rel_theta)
86
+
87
+ re_score = (rel_re * head_re - rel_im * head_im) - tail_re
88
+ im_score = (rel_re * head_im + rel_im * head_re) - tail_im
89
+ complex_score = ops.stack([re_score, im_score], axis=2)
90
+ score = ops.sqrt(ops.sum(ops.square(complex_score), axis=(1, 2)))
91
+
92
+ return self.margin - score
93
+
94
+ def loss(self, head_index, rel_type, tail_index):
95
+ pos_score = self(head_index, rel_type, tail_index)
96
+ neg_score = self(*self.random_sample(head_index, rel_type, tail_index))
97
+ scores = ops.concatenate([pos_score, neg_score], axis=0)
98
+
99
+ pos_target = ops.ones_like(pos_score)
100
+ neg_target = ops.zeros_like(neg_score)
101
+ target = ops.concatenate([pos_target, neg_target], axis=0)
102
+
103
+ return binary_cross_entropy_with_logits(scores, target)
@@ -0,0 +1,76 @@
1
+ from keras import ops
2
+
3
+ from k3_node.layers.kge import TransE, DistMult, ComplEx, RotatE
4
+
5
+
6
+ def _run_model(cls):
7
+ model = cls(num_nodes=10, num_relations=5, hidden_channels=32)
8
+ assert str(model) == f"{cls.__name__}(10, num_relations=5, hidden_channels=32)"
9
+
10
+ head_index = ops.convert_to_tensor([0, 2, 4, 6, 8])
11
+ rel_type = ops.convert_to_tensor([0, 1, 2, 3, 4])
12
+ tail_index = ops.convert_to_tensor([1, 3, 5, 7, 9])
13
+
14
+ loader = model.loader(head_index, rel_type, tail_index, batch_size=5)
15
+ for h, r, t in loader:
16
+ out = model(h, r, t)
17
+ assert ops.shape(out) == (5,)
18
+
19
+ loss = model.loss(h, r, t)
20
+ assert float(ops.convert_to_numpy(loss)) >= 0.0
21
+
22
+ mean_rank, mrr, hits = model.test(h, r, t, batch_size=5, log=False)
23
+ assert 0 <= mean_rank <= 10
24
+ assert 0 < mrr <= 1
25
+ assert hits == 1.0
26
+
27
+
28
+ def test_transe():
29
+ _run_model(TransE)
30
+
31
+
32
+ def test_distmult():
33
+ _run_model(DistMult)
34
+
35
+
36
+ def test_complex():
37
+ _run_model(ComplEx)
38
+
39
+
40
+ def test_rotate():
41
+ _run_model(RotatE)
42
+
43
+
44
+ def test_complex_scoring():
45
+ model = ComplEx(num_nodes=5, num_relations=2, hidden_channels=1)
46
+
47
+ model.node_emb.embeddings.assign(
48
+ ops.convert_to_tensor([[2.0], [3.0], [5.0], [1.0], [2.0]], dtype="float32")
49
+ )
50
+ model.node_emb_im.embeddings.assign(
51
+ ops.convert_to_tensor([[4.0], [1.0], [3.0], [1.0], [2.0]], dtype="float32")
52
+ )
53
+ model.rel_emb.embeddings.assign(ops.convert_to_tensor([[2.0], [3.0]], dtype="float32"))
54
+ model.rel_emb_im.embeddings.assign(ops.convert_to_tensor([[3.0], [1.0]], dtype="float32"))
55
+
56
+ score = model(
57
+ ops.convert_to_tensor([1, 3]),
58
+ ops.convert_to_tensor([1, 0]),
59
+ ops.convert_to_tensor([2, 4]),
60
+ )
61
+ assert ops.convert_to_numpy(score).tolist() == [58.0, 8.0]
62
+
63
+
64
+ def test_same_shape_embeddings_are_initialized_independently():
65
+ # A reused Keras 3 initializer instance repeats its values. For ComplEx, identical real and
66
+ # imaginary parts cancel the asymmetric term, collapsing the model to DistMult.
67
+ import numpy as np
68
+
69
+ complex_model = ComplEx(num_nodes=20, num_relations=5, hidden_channels=8)
70
+ rotate_model = RotatE(num_nodes=20, num_relations=5, hidden_channels=8)
71
+ for a, b in [
72
+ (complex_model.node_emb, complex_model.node_emb_im),
73
+ (complex_model.rel_emb, complex_model.rel_emb_im),
74
+ (rotate_model.node_emb, rotate_model.node_emb_im),
75
+ ]:
76
+ assert not np.allclose(ops.convert_to_numpy(a.embeddings), ops.convert_to_numpy(b.embeddings))