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,270 @@
1
+ """Subgraph extraction utilities for GraphRAG and Knowledge Graph LLM integration."""
2
+
3
+ from dataclasses import dataclass
4
+ from typing import Dict, List, Optional, Sequence, Tuple, Union
5
+
6
+ import numpy as np
7
+ from keras import ops
8
+
9
+ from k3_node.data import Data
10
+
11
+
12
+ @dataclass
13
+ class SubgraphResult:
14
+ """Structured result of a subgraph extraction around retrieved entities.
15
+
16
+ Attributes:
17
+ edge_index: Tensor of shape `(2, num_edges)` with subgraph edges.
18
+ edge_type: Optional tensor of shape `(num_edges,)` with relation types.
19
+ edge_attr: Optional tensor of shape `(num_edges, edge_dim)` with edge features.
20
+ x: Optional tensor of shape `(num_nodes, in_channels)` with node features.
21
+ nodes: 1D numpy array of original node indices in the full graph.
22
+ center_nodes: 1D numpy array of center entity indices in the relabeled subgraph.
23
+ mapping: Dictionary mapping original node index to subgraph node index.
24
+ edge_mask: 1D boolean array indicating edges kept from original graph.
25
+ num_nodes: Number of nodes in the extracted subgraph.
26
+ num_edges: Number of edges in the extracted subgraph.
27
+ """
28
+
29
+ edge_index: any
30
+ edge_type: Optional[any] = None
31
+ edge_attr: Optional[any] = None
32
+ x: Optional[any] = None
33
+ nodes: Optional[np.ndarray] = None
34
+ center_nodes: Optional[np.ndarray] = None
35
+ mapping: Optional[Dict[int, int]] = None
36
+ edge_mask: Optional[np.ndarray] = None
37
+ num_nodes: int = 0
38
+ num_edges: int = 0
39
+
40
+ def to_data(self) -> Data:
41
+ """Convert extracted subgraph into a K3-Node Data object."""
42
+ return Data(
43
+ x=self.x,
44
+ edge_index=self.edge_index,
45
+ edge_type=self.edge_type,
46
+ edge_attr=self.edge_attr,
47
+ center_nodes=ops.convert_to_tensor(self.center_nodes, dtype="int32")
48
+ if self.center_nodes is not None
49
+ else None,
50
+ original_nodes=ops.convert_to_tensor(self.nodes, dtype="int32")
51
+ if self.nodes is not None
52
+ else None,
53
+ )
54
+
55
+
56
+ def extract_subgraph(
57
+ entities: Union[int, Sequence[int], np.ndarray, any],
58
+ edge_index: Union[np.ndarray, any],
59
+ edge_type: Optional[Union[np.ndarray, any]] = None,
60
+ edge_attr: Optional[Union[np.ndarray, any]] = None,
61
+ x: Optional[Union[np.ndarray, any]] = None,
62
+ num_hops: int = 2,
63
+ max_nodes_per_hop: Optional[int] = None,
64
+ directed: bool = False,
65
+ relabel_nodes: bool = True,
66
+ num_nodes: Optional[int] = None,
67
+ ) -> SubgraphResult:
68
+ """Extract multi-hop enclosing or ego-subgraphs around retrieved entities.
69
+
70
+ Given a knowledge graph or relational graph, this function expands `num_hops`
71
+ around seed `entities`, keeping all induced edges, relation types, and node/edge
72
+ attributes.
73
+
74
+ Args:
75
+ entities: Single node index or list/array of seed entity indices.
76
+ edge_index: Graph connectivity tensor of shape `(2, num_edges)`.
77
+ edge_type: Optional 1D relation type tensor of shape `(num_edges,)`.
78
+ edge_attr: Optional edge attribute tensor of shape `(num_edges, edge_dim)`.
79
+ x: Optional node feature tensor of shape `(total_nodes, in_channels)`.
80
+ num_hops: Number of hops to expand around retrieved entities. (default: 2)
81
+ max_nodes_per_hop: Optional maximum number of neighboring nodes to keep
82
+ per hop (useful to restrict explosion on hub entities).
83
+ directed: If True, only follows outgoing edges. If False, follows edges
84
+ in both directions (standard for KG context expansion). (default: False)
85
+ relabel_nodes: If True, relabels subgraph node IDs to `0..num_subgraph_nodes-1`.
86
+ (default: True)
87
+ num_nodes: Optional total number of nodes in graph. Inferred if not given.
88
+
89
+ Returns:
90
+ `SubgraphResult` containing relabeled edge_index, edge_type, features,
91
+ and entity mappings.
92
+ """
93
+ edge_index_np = np.asarray(ops.convert_to_numpy(edge_index)).astype(np.int64)
94
+ if num_nodes is None:
95
+ num_nodes = int(edge_index_np.max()) + 1 if edge_index_np.size > 0 else 0
96
+
97
+ entities_arr = np.atleast_1d(np.asarray(ops.convert_to_numpy(entities))).astype(np.int64)
98
+ if entities_arr.size == 0:
99
+ empty_ei = ops.convert_to_tensor(np.zeros((2, 0), dtype=np.int64), dtype=edge_index.dtype)
100
+ return SubgraphResult(
101
+ edge_index=empty_ei,
102
+ nodes=np.array([], dtype=np.int64),
103
+ center_nodes=np.array([], dtype=np.int64),
104
+ mapping={},
105
+ num_nodes=0,
106
+ num_edges=0,
107
+ )
108
+
109
+ row, col = edge_index_np[0], edge_index_np[1]
110
+ subsets = [entities_arr]
111
+ visited = set(entities_arr.tolist())
112
+
113
+ for _ in range(num_hops):
114
+ current_frontier = subsets[-1]
115
+ if len(current_frontier) == 0:
116
+ break
117
+
118
+ mask_frontier = np.zeros(num_nodes, dtype=bool)
119
+ mask_frontier[current_frontier] = True
120
+
121
+ # If undirected (or GraphRAG bidirectional expansion), collect neighbors in both directions
122
+ if not directed:
123
+ neighbors_out = col[mask_frontier[row]]
124
+ neighbors_in = row[mask_frontier[col]]
125
+ new_neighbors = np.concatenate([neighbors_out, neighbors_in])
126
+ else:
127
+ new_neighbors = col[mask_frontier[row]]
128
+
129
+ if new_neighbors.size > 0:
130
+ unique_new = np.unique(new_neighbors)
131
+ unseen = [n for n in unique_new if n not in visited]
132
+ if max_nodes_per_hop is not None and len(unseen) > max_nodes_per_hop:
133
+ unseen = unseen[:max_nodes_per_hop]
134
+ visited.update(unseen)
135
+ subsets.append(np.array(unseen, dtype=np.int64))
136
+ else:
137
+ break
138
+
139
+ subgraph_nodes = np.unique(np.concatenate([s for s in subsets if len(s) > 0]))
140
+
141
+ # Induced edge mask: both endpoints must be in subgraph_nodes
142
+ node_mask = np.zeros(num_nodes, dtype=bool)
143
+ node_mask[subgraph_nodes] = True
144
+ edge_mask = node_mask[row] & node_mask[col]
145
+
146
+ sub_edge_index = edge_index_np[:, edge_mask]
147
+
148
+ mapping = {int(orig): int(new_idx) for new_idx, orig in enumerate(subgraph_nodes)}
149
+ center_indices = np.array([mapping[int(e)] for e in entities_arr if int(e) in mapping], dtype=np.int64)
150
+
151
+ if relabel_nodes:
152
+ new_id_table = np.full(num_nodes, -1, dtype=np.int64)
153
+ new_id_table[subgraph_nodes] = np.arange(len(subgraph_nodes), dtype=np.int64)
154
+ sub_edge_index = new_id_table[sub_edge_index]
155
+
156
+ # Preserve tensors in their original backend / dtype
157
+ sub_edge_index_tensor = ops.convert_to_tensor(sub_edge_index, dtype=edge_index.dtype)
158
+
159
+ sub_edge_type_tensor = None
160
+ if edge_type is not None:
161
+ et_np = np.asarray(ops.convert_to_numpy(edge_type))[edge_mask]
162
+ sub_edge_type_tensor = ops.convert_to_tensor(et_np, dtype=edge_type.dtype)
163
+
164
+ sub_edge_attr_tensor = None
165
+ if edge_attr is not None:
166
+ ea_np = np.asarray(ops.convert_to_numpy(edge_attr))[edge_mask]
167
+ sub_edge_attr_tensor = ops.convert_to_tensor(ea_np, dtype=edge_attr.dtype)
168
+
169
+ sub_x_tensor = None
170
+ if x is not None:
171
+ x_np = np.asarray(ops.convert_to_numpy(x))[subgraph_nodes]
172
+ sub_x_tensor = ops.convert_to_tensor(x_np, dtype=x.dtype)
173
+
174
+ return SubgraphResult(
175
+ edge_index=sub_edge_index_tensor,
176
+ edge_type=sub_edge_type_tensor,
177
+ edge_attr=sub_edge_attr_tensor,
178
+ x=sub_x_tensor,
179
+ nodes=subgraph_nodes,
180
+ center_nodes=center_indices,
181
+ mapping=mapping,
182
+ edge_mask=edge_mask,
183
+ num_nodes=len(subgraph_nodes),
184
+ num_edges=int(sub_edge_index.shape[1]),
185
+ )
186
+
187
+
188
+ class KGEntityRetriever:
189
+ """Knowledge Graph Entity and Subgraph Retriever for GraphRAG.
190
+
191
+ Maintains entity name dictionaries and relation mappings, extracts seed entities
192
+ from text queries, and retrieves enclosing multi-hop subgraphs.
193
+
194
+ Args:
195
+ entity_to_id: Dictionary mapping entity strings to node integer IDs.
196
+ relation_to_id: Dictionary mapping relation strings to edge_type integer IDs.
197
+ edge_index: Graph connectivity tensor of shape `(2, num_edges)`.
198
+ edge_type: Optional relation type tensor of shape `(num_edges,)`.
199
+ edge_attr: Optional edge feature tensor.
200
+ x: Optional node feature tensor.
201
+ """
202
+
203
+ def __init__(
204
+ self,
205
+ entity_to_id: Dict[str, int],
206
+ relation_to_id: Dict[str, int],
207
+ edge_index: Union[np.ndarray, any],
208
+ edge_type: Optional[Union[np.ndarray, any]] = None,
209
+ edge_attr: Optional[Union[np.ndarray, any]] = None,
210
+ x: Optional[Union[np.ndarray, any]] = None,
211
+ ):
212
+ self.entity_to_id = entity_to_id
213
+ self.relation_to_id = relation_to_id
214
+ self.id_to_entity = {v: k for k, v in entity_to_id.items()}
215
+ self.id_to_relation = {v: k for k, v in relation_to_id.items()}
216
+
217
+ self.edge_index = edge_index
218
+ self.edge_type = edge_type
219
+ self.edge_attr = edge_attr
220
+ self.x = x
221
+
222
+ def get_entity_id(self, name: str) -> Optional[int]:
223
+ """Look up entity ID by exact name (case-insensitive fallback)."""
224
+ if name in self.entity_to_id:
225
+ return self.entity_to_id[name]
226
+ # Case-insensitive fallback
227
+ lower_map = {k.lower(): v for k, v in self.entity_to_id.items()}
228
+ return lower_map.get(name.lower(), None)
229
+
230
+ def find_entities_in_text(self, text: str) -> List[str]:
231
+ """Find matching known entity names mentioned in a query text."""
232
+ text_lower = text.lower()
233
+ matched = []
234
+ # Sort by length descending to match longest phrases first
235
+ for name in sorted(self.entity_to_id.keys(), key=lambda s: len(s), reverse=True):
236
+ if name.lower() in text_lower:
237
+ matched.append(name)
238
+ return matched
239
+
240
+ def retrieve_subgraph(
241
+ self,
242
+ entities: Union[Sequence[Union[str, int]], str, int],
243
+ num_hops: int = 2,
244
+ max_nodes_per_hop: Optional[int] = None,
245
+ directed: bool = False,
246
+ ) -> SubgraphResult:
247
+ """Extract multi-hop subgraph around the specified entity names or IDs."""
248
+ if isinstance(entities, (str, int)):
249
+ entities = [entities]
250
+
251
+ entity_ids = []
252
+ for e in entities:
253
+ if isinstance(e, str):
254
+ eid = self.get_entity_id(e)
255
+ if eid is not None:
256
+ entity_ids.append(eid)
257
+ else:
258
+ entity_ids.append(int(e))
259
+
260
+ return extract_subgraph(
261
+ entities=entity_ids,
262
+ edge_index=self.edge_index,
263
+ edge_type=self.edge_type,
264
+ edge_attr=self.edge_attr,
265
+ x=self.x,
266
+ num_hops=num_hops,
267
+ max_nodes_per_hop=max_nodes_per_hop,
268
+ directed=directed,
269
+ relabel_nodes=True,
270
+ )
@@ -0,0 +1,347 @@
1
+ """Tests for k3_node.rag GraphRAG and KG-LLM connectors."""
2
+
3
+ import numpy as np
4
+ import pytest
5
+ from keras import ops
6
+
7
+ import k3_node as k3
8
+ from k3_node.layers.kge import TransE
9
+ from k3_node.rag import (
10
+ GraphPrefixProjector,
11
+ GraphRAG,
12
+ KGEntityRetriever,
13
+ KGLLMConnector,
14
+ RGCNSubGraphEncoder,
15
+ SubgraphResult,
16
+ TransEPrefixEncoder,
17
+ extract_subgraph,
18
+ format_llm_prompt,
19
+ subgraph_to_triples,
20
+ verbalize_subgraph,
21
+ )
22
+
23
+
24
+ @pytest.fixture
25
+ def sample_kg():
26
+ """Create a sample knowledge graph for testing.
27
+
28
+ Graph:
29
+ 0 (Aspirin) --[0: treats]--> 1 (Headache)
30
+ 0 (Aspirin) --[1: inhibits]--> 2 (COX-1)
31
+ 2 (COX-1) --[2: produces]--> 3 (Prostaglandin)
32
+ 3 (Prostaglandin) --[3: causes]--> 1 (Headache)
33
+ 4 (Ibuprofen) --[0: treats]--> 1 (Headache)
34
+ 4 (Ibuprofen) --[1: inhibits]--> 2 (COX-1)
35
+ """
36
+ edges = np.array(
37
+ [
38
+ [0, 0, 2, 3, 4, 4],
39
+ [1, 2, 3, 1, 1, 2],
40
+ ],
41
+ dtype=np.int64,
42
+ )
43
+ edge_type = np.array([0, 1, 2, 3, 0, 1], dtype=np.int64)
44
+ x = np.random.randn(5, 16).astype(np.float32)
45
+
46
+ entity_to_id = {
47
+ "Aspirin": 0,
48
+ "Headache": 1,
49
+ "COX-1": 2,
50
+ "Prostaglandin": 3,
51
+ "Ibuprofen": 4,
52
+ }
53
+ relation_to_id = {
54
+ "treats": 0,
55
+ "inhibits": 1,
56
+ "produces": 2,
57
+ "causes": 3,
58
+ }
59
+
60
+ return {
61
+ "edge_index": ops.convert_to_tensor(edges, dtype="int32"),
62
+ "edge_type": ops.convert_to_tensor(edge_type, dtype="int32"),
63
+ "x": ops.convert_to_tensor(x, dtype="float32"),
64
+ "entity_to_id": entity_to_id,
65
+ "relation_to_id": relation_to_id,
66
+ "num_nodes": 5,
67
+ "num_relations": 4,
68
+ }
69
+
70
+
71
+ def test_extract_subgraph_1hop(sample_kg):
72
+ sub = extract_subgraph(
73
+ entities=[0], # Aspirin
74
+ edge_index=sample_kg["edge_index"],
75
+ edge_type=sample_kg["edge_type"],
76
+ x=sample_kg["x"],
77
+ num_hops=1,
78
+ directed=False,
79
+ )
80
+
81
+ assert isinstance(sub, SubgraphResult)
82
+ # Neighbors of 0 within 1 hop: 0, 1, 2
83
+ assert set(sub.nodes.tolist()) == {0, 1, 2}
84
+ assert sub.center_nodes.tolist() == [sub.mapping[0]]
85
+ assert sub.num_nodes == 3
86
+ assert sub.num_edges > 0
87
+ assert sub.x is not None
88
+ assert ops.shape(sub.x)[0] == 3
89
+
90
+ # Check to_data()
91
+ data = sub.to_data()
92
+ assert isinstance(data, k3.data.Data)
93
+ assert ops.shape(data.edge_index)[0] == 2
94
+
95
+
96
+ def test_extract_subgraph_2hop_multiseed(sample_kg):
97
+ sub = extract_subgraph(
98
+ entities=[0, 3], # Aspirin & Prostaglandin
99
+ edge_index=sample_kg["edge_index"],
100
+ edge_type=sample_kg["edge_type"],
101
+ num_hops=2,
102
+ )
103
+ assert sub.num_nodes == 5
104
+ assert len(sub.center_nodes) == 2
105
+
106
+
107
+ def test_kg_entity_retriever(sample_kg):
108
+ retriever = KGEntityRetriever(
109
+ entity_to_id=sample_kg["entity_to_id"],
110
+ relation_to_id=sample_kg["relation_to_id"],
111
+ edge_index=sample_kg["edge_index"],
112
+ edge_type=sample_kg["edge_type"],
113
+ x=sample_kg["x"],
114
+ )
115
+
116
+ # Lookup
117
+ assert retriever.get_entity_id("Aspirin") == 0
118
+ assert retriever.get_entity_id("aspirin") == 0 # case-insensitive
119
+ assert retriever.get_entity_id("Unknown") is None
120
+
121
+ # Text entity extraction
122
+ query = "Does Aspirin or Ibuprofen treat headache?"
123
+ found = retriever.find_entities_in_text(query)
124
+ assert "Aspirin" in found
125
+ assert "Ibuprofen" in found
126
+ assert "Headache" in found
127
+
128
+ # Subgraph retrieval
129
+ sub = retriever.retrieve_subgraph(["Aspirin"], num_hops=1)
130
+ assert sub.num_nodes >= 2
131
+
132
+
133
+ def test_verbalization(sample_kg):
134
+ sub = extract_subgraph(
135
+ entities=[0],
136
+ edge_index=sample_kg["edge_index"],
137
+ edge_type=sample_kg["edge_type"],
138
+ num_hops=1,
139
+ )
140
+
141
+ id_to_e = {v: k for k, v in sample_kg["entity_to_id"].items()}
142
+ id_to_r = {v: k for k, v in sample_kg["relation_to_id"].items()}
143
+
144
+ # Triples list
145
+ triples = subgraph_to_triples(sub, id_to_e, id_to_r)
146
+ assert len(triples) > 0
147
+ assert ("Aspirin", "treats", "Headache") in triples
148
+
149
+ # Markdown format
150
+ md_text = verbalize_subgraph(sub, id_to_e, id_to_r, format_style="markdown")
151
+ assert "**Aspirin**" in md_text
152
+ assert "*treats*" in md_text
153
+
154
+ # Natural language format
155
+ natural_text = verbalize_subgraph(sub, id_to_e, id_to_r, format_style="natural")
156
+ assert "Aspirin treats Headache." in natural_text
157
+
158
+
159
+ def test_prompt_formatting():
160
+ context = "- **Aspirin** — *treats* -> **Headache**"
161
+ query = "What treats headache?"
162
+
163
+ # Llama 3
164
+ llama3_prompt = format_llm_prompt(query, context, model_family="llama3")
165
+ assert "<|start_header_id|>system<|end_header_id|>" in llama3_prompt
166
+ assert "<|start_header_id|>user<|end_header_id|>" in llama3_prompt
167
+ assert "Knowledge Graph Context:" in llama3_prompt
168
+ assert query in llama3_prompt
169
+
170
+ # Mistral
171
+ mistral_prompt = format_llm_prompt(query, context, model_family="mistral")
172
+ assert "<s>[INST]" in mistral_prompt
173
+ assert "[/INST]" in mistral_prompt
174
+
175
+ # ChatML
176
+ chatml_prompt = format_llm_prompt(query, context, model_family="chatml")
177
+ assert "<|im_start|>system" in chatml_prompt
178
+
179
+
180
+ def test_rgcn_subgraph_encoder(sample_kg):
181
+ encoder = RGCNSubGraphEncoder(
182
+ in_channels=16,
183
+ hidden_channels=32,
184
+ out_channels=64,
185
+ num_relations=sample_kg["num_relations"],
186
+ num_layers=2,
187
+ pooling="center",
188
+ )
189
+
190
+ sub = extract_subgraph(
191
+ entities=[0],
192
+ edge_index=sample_kg["edge_index"],
193
+ edge_type=sample_kg["edge_type"],
194
+ x=sample_kg["x"],
195
+ num_hops=1,
196
+ )
197
+
198
+ emb = encoder.encode_subgraph(sub)
199
+ assert ops.shape(emb) == (1, 64)
200
+
201
+ # Test other pooling options
202
+ encoder_mean = RGCNSubGraphEncoder(
203
+ in_channels=16,
204
+ hidden_channels=32,
205
+ out_channels=64,
206
+ num_relations=sample_kg["num_relations"],
207
+ pooling="mean",
208
+ )
209
+ emb_mean = encoder_mean.encode_subgraph(sub)
210
+ assert ops.shape(emb_mean) == (1, 64)
211
+
212
+
213
+ def test_transe_prefix_encoder(sample_kg):
214
+ transe = TransE(
215
+ num_nodes=sample_kg["num_nodes"],
216
+ num_relations=sample_kg["num_relations"],
217
+ hidden_channels=32,
218
+ )
219
+
220
+ kge_encoder = TransEPrefixEncoder(
221
+ kge_model=transe,
222
+ out_channels=64,
223
+ pooling="center",
224
+ )
225
+
226
+ sub = extract_subgraph(
227
+ entities=[0, 1],
228
+ edge_index=sample_kg["edge_index"],
229
+ edge_type=sample_kg["edge_type"],
230
+ num_hops=1,
231
+ )
232
+
233
+ emb = kge_encoder.encode_subgraph(sub)
234
+ assert ops.shape(emb) == (1, 64)
235
+
236
+ # Test standalone initialization without pretrained model
237
+ standalone_encoder = TransEPrefixEncoder(
238
+ num_nodes=10,
239
+ num_relations=4,
240
+ embedding_dim=32,
241
+ out_channels=64,
242
+ )
243
+ emb_standalone = standalone_encoder.encode_subgraph(sub)
244
+ assert ops.shape(emb_standalone) == (1, 64)
245
+
246
+
247
+ def test_graph_prefix_projector():
248
+ projector = GraphPrefixProjector(
249
+ in_channels=64,
250
+ llm_dim=256, # test dimension
251
+ num_prefix_tokens=4,
252
+ projector_type="mlp",
253
+ )
254
+
255
+ graph_emb = ops.convert_to_tensor(np.random.randn(1, 64).astype(np.float32))
256
+ prefix = projector(graph_emb)
257
+ assert ops.shape(prefix) == (1, 4, 256)
258
+
259
+ # Linear projector
260
+ lin_projector = GraphPrefixProjector(
261
+ in_channels=64,
262
+ llm_dim=256,
263
+ num_prefix_tokens=4,
264
+ projector_type="linear",
265
+ )
266
+ prefix_lin = lin_projector(graph_emb)
267
+ assert ops.shape(prefix_lin) == (1, 4, 256)
268
+
269
+
270
+ def test_kg_llm_connector(sample_kg):
271
+ encoder = RGCNSubGraphEncoder(
272
+ in_channels=16,
273
+ hidden_channels=32,
274
+ out_channels=64,
275
+ num_relations=sample_kg["num_relations"],
276
+ )
277
+ projector = GraphPrefixProjector(
278
+ in_channels=64,
279
+ llm_dim=512,
280
+ num_prefix_tokens=4,
281
+ )
282
+ connector = KGLLMConnector(encoder=encoder, projector=projector)
283
+
284
+ sub = extract_subgraph(
285
+ entities=[0],
286
+ edge_index=sample_kg["edge_index"],
287
+ edge_type=sample_kg["edge_type"],
288
+ x=sample_kg["x"],
289
+ num_hops=1,
290
+ )
291
+
292
+ prefix_tokens = connector.encode_subgraph(sub)
293
+ assert ops.shape(prefix_tokens) == (1, 4, 512)
294
+
295
+ # Test prefix injection into LLM text tokens
296
+ text_tokens = ops.convert_to_tensor(np.random.randn(1, 10, 512).astype(np.float32))
297
+ augmented = connector.inject_prefix(text_tokens, prefix_tokens)
298
+ assert ops.shape(augmented) == (1, 14, 512)
299
+
300
+ # Test attention mask extension
301
+ attn_mask = ops.ones((1, 10), dtype="int32")
302
+ extended_mask = connector.extend_attention_mask(attn_mask, num_prefix_tokens=4)
303
+ assert ops.shape(extended_mask) == (1, 14)
304
+
305
+
306
+ def test_graph_rag_pipeline(sample_kg):
307
+ # Pipeline with RGCN
308
+ rag_rgcn = GraphRAG(
309
+ edge_index=sample_kg["edge_index"],
310
+ edge_type=sample_kg["edge_type"],
311
+ entity_to_id=sample_kg["entity_to_id"],
312
+ relation_to_id=sample_kg["relation_to_id"],
313
+ x=sample_kg["x"],
314
+ encoder_type="rgcn",
315
+ hidden_dim=32,
316
+ encoder_out_dim=64,
317
+ llm_dim=256,
318
+ num_prefix_tokens=4,
319
+ )
320
+
321
+ # Retrieval
322
+ sub = rag_rgcn.retrieve(["Aspirin"], num_hops=1)
323
+ assert sub.num_nodes >= 2
324
+
325
+ # Prompt builder
326
+ prompt = rag_rgcn.build_prompt("What does Aspirin treat?", model_family="llama3")
327
+ assert "Aspirin" in prompt
328
+ assert "<|begin_of_text|>" in prompt
329
+
330
+ # Prefix encoding
331
+ prefix = rag_rgcn.encode_prefix(sub)
332
+ assert ops.shape(prefix) == (1, 4, 256)
333
+
334
+ # Pipeline with TransE
335
+ rag_transe = GraphRAG(
336
+ edge_index=sample_kg["edge_index"],
337
+ edge_type=sample_kg["edge_type"],
338
+ entity_to_id=sample_kg["entity_to_id"],
339
+ relation_to_id=sample_kg["relation_to_id"],
340
+ encoder_type="transe",
341
+ hidden_dim=32,
342
+ encoder_out_dim=64,
343
+ llm_dim=256,
344
+ num_prefix_tokens=4,
345
+ )
346
+ prefix_transe = rag_transe.encode_prefix(sub)
347
+ assert ops.shape(prefix_transe) == (1, 4, 256)