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,4 @@
1
+ from .graph import *
2
+ from .matmul import *
3
+ from .conv import *
4
+ from .numpy import *
k3_node/ops/conv.py ADDED
@@ -0,0 +1,56 @@
1
+ # ported from spektral
2
+
3
+ import numpy as np
4
+ import warnings, copy
5
+
6
+
7
+ def degree_matrix(A):
8
+ degrees = np.array(A.sum(1)).flatten()
9
+ D = np.diag(degrees)
10
+ return D
11
+
12
+
13
+ def degree_power(A, k):
14
+ with warnings.catch_warnings():
15
+ warnings.simplefilter("ignore")
16
+ degrees = np.power(np.array(A.sum(1)), k).ravel()
17
+ degrees[np.isinf(degrees)] = 0.0
18
+ D = np.diag(degrees)
19
+ return D
20
+
21
+
22
+ def normalized_adjacency(A, symmetric=True):
23
+ if symmetric:
24
+ normalized_D = degree_power(A, -0.5)
25
+ return normalized_D.dot(A).dot(normalized_D)
26
+ else:
27
+ normalized_D = degree_power(A, -1.0)
28
+ return normalized_D.dot(A)
29
+
30
+
31
+ def laplacian(A):
32
+ return degree_matrix(A) - A
33
+
34
+
35
+ def normalized_laplacian(A, symmetric=True):
36
+ I = np.eye(A.shape[-1], dtype=A.dtype)
37
+ normalized_adj = normalized_adjacency(A, symmetric=symmetric)
38
+ return I - normalized_adj
39
+
40
+
41
+ def gcn_filter(A, symmetric=True):
42
+ out = copy.deepcopy(A)
43
+ if isinstance(A, list) or (isinstance(A, np.ndarray) and A.ndim == 3):
44
+ for i in range(len(A)):
45
+ out[i] = A[i]
46
+ out[i][np.diag_indices_from(out[i])] += 1
47
+ out[i] = normalized_adjacency(out[i], symmetric=symmetric)
48
+ else:
49
+ if hasattr(out, "tocsr"):
50
+ out = out.tocsr()
51
+ with warnings.catch_warnings():
52
+ warnings.simplefilter("ignore")
53
+ out[np.diag_indices_from(out)] += 1
54
+ out = normalized_adjacency(out, symmetric=symmetric)
55
+
56
+ return out
@@ -0,0 +1,43 @@
1
+ """Filled-tensor creation that also works on torch's ``meta`` device.
2
+
3
+ ``keras.ops.full`` / ``keras.ops.full_like`` call ``Tensor.item()`` on the torch backend, which
4
+ fails on the ``meta`` device Keras uses to infer shapes before the first training step (the
5
+ fallback then prints a long "unbuilt state" warning). ``ones * value`` avoids that.
6
+ """
7
+ from keras import ops
8
+
9
+
10
+ def full(shape, fill_value, dtype=None):
11
+ r"""Like :func:`keras.ops.full`."""
12
+ return ops.ones(shape, dtype=dtype) * ops.cast(fill_value, dtype or "float32")
13
+
14
+
15
+ def full_like(x, fill_value, dtype=None):
16
+ r"""Like :func:`keras.ops.full_like`."""
17
+ return ops.ones_like(x, dtype=dtype) * ops.cast(fill_value, dtype or x.dtype)
18
+
19
+
20
+ def scatter(indices, values, shape):
21
+ r"""Like :func:`keras.ops.scatter` (values at duplicate indices are summed)."""
22
+ from keras import backend
23
+
24
+ if backend.backend() == "torch":
25
+ import torch
26
+
27
+ values = values if torch.is_tensor(values) else ops.convert_to_tensor(values)
28
+ indices = indices if torch.is_tensor(indices) else ops.convert_to_tensor(indices)
29
+ out = torch.zeros(tuple(int(s) for s in shape), dtype=values.dtype, device=values.device)
30
+ return out.index_put_(tuple(indices.long().T), values, accumulate=True)
31
+ return ops.scatter(indices, values, shape)
32
+
33
+
34
+ def repeat(x, repeats, axis=None):
35
+ r"""Like :func:`keras.ops.repeat`; an integer ``repeats`` also works on torch's ``meta`` device."""
36
+ from keras import backend
37
+
38
+ if backend.backend() == "torch" and isinstance(repeats, int):
39
+ import torch
40
+
41
+ x = x if torch.is_tensor(x) else ops.convert_to_tensor(x)
42
+ return torch.repeat_interleave(x.reshape(-1) if axis is None else x, repeats, dim=axis)
43
+ return ops.repeat(x, repeats, axis=axis)
k3_node/ops/graph.py ADDED
@@ -0,0 +1,27 @@
1
+ # ported from spektral
2
+ from keras import ops as k_ops, backend
3
+
4
+ from k3_node.utils import *
5
+
6
+
7
+ def degrees(A):
8
+ return k_ops.sum(A, axis=-1)
9
+
10
+
11
+ def normalize_A(A):
12
+ D = degrees(A)
13
+ D = k_ops.sqrt(D)[:, None] + backend.epsilon()
14
+ perm = (0, 2, 1) if len(k_ops.shape(A)) == 3 else (1, 0)
15
+ output = (A / D) / k_ops.transpose(D, perm=perm)
16
+ return output
17
+
18
+
19
+ def get_source_target(a):
20
+ if backend.backend() == "tensorflow":
21
+ import tensorflow as tf
22
+ if isinstance(a, tf.sparse.SparseTensor):
23
+ return a.indices[:, 0], a.indices[:, 1]
24
+ else:
25
+ return k_ops.where(a != 0)
26
+ else:
27
+ return k_ops.where(a != 0)
k3_node/ops/host.py ADDED
@@ -0,0 +1,41 @@
1
+ import numpy as np
2
+ from keras import ops
3
+
4
+
5
+ def _is_meta(x) -> bool:
6
+ return getattr(getattr(x, "device", None), "type", None) == "meta"
7
+
8
+
9
+ def to_numpy(x):
10
+ """``keras.ops.convert_to_numpy`` for code that runs on the host (NumPy).
11
+
12
+ During Keras' shape inference on the torch backend, tensors live on the "meta" device and have
13
+ no values. Host code then gets zeros of the right shape and dtype: only the output shapes matter
14
+ there, not the values.
15
+ """
16
+ if x is None or isinstance(x, np.ndarray):
17
+ return x
18
+ if _is_meta(x):
19
+ import torch
20
+
21
+ return np.zeros(tuple(x.shape), dtype=torch.empty(0, dtype=x.dtype).numpy().dtype)
22
+ try:
23
+ return ops.convert_to_numpy(x)
24
+ except Exception:
25
+ # TensorFlow and JAX trace symbolic tensors during Keras' shape inference. Unknown sizes
26
+ # (e.g. the number of nodes) become 1: zeros then describe a valid one-node graph.
27
+ if _in_shape_inference():
28
+ import keras
29
+
30
+ shape = tuple(1 if d is None else d for d in x.shape)
31
+ return np.zeros(shape, dtype=keras.backend.standardize_dtype(x.dtype))
32
+ raise
33
+
34
+
35
+ def _in_shape_inference() -> bool:
36
+ try:
37
+ from keras.src.backend.common.symbolic_scope import in_symbolic_scope
38
+
39
+ return in_symbolic_scope()
40
+ except ImportError:
41
+ return False
k3_node/ops/matmul.py ADDED
@@ -0,0 +1,49 @@
1
+ # ported from spektral
2
+
3
+ from keras import ops as k_ops
4
+
5
+
6
+ def dot(a, b):
7
+ a_ndim = len(k_ops.shape(a))
8
+ b_ndim = len(k_ops.shape(b))
9
+ assert a_ndim == b_ndim, "Expected equal ranks, got {} and {}" "".format(
10
+ a_ndim, b_ndim
11
+ )
12
+ return k_ops.matmul(a, b)
13
+
14
+
15
+ def mixed_mode_dot(a, b):
16
+ a_shp = k_ops.shape(a)
17
+ b_shp = k_ops.shape(b)
18
+
19
+ b_t = k_ops.transpose(b, (1, 2, 0))
20
+ b_t = k_ops.reshape(b_t, k_ops.stack((b_shp[1], -1)))
21
+ output = dot(a, b_t)
22
+ output = k_ops.reshape(output, k_ops.stack((a_shp[0], b_shp[2], -1)))
23
+ output = k_ops.transpose(output, (2, 0, 1))
24
+
25
+ return output
26
+
27
+
28
+ def modal_dot(a, b, transpose_a=False, transpose_b=False):
29
+ a_ndim = len(k_ops.shape(a))
30
+ b_ndim = len(k_ops.shape(b))
31
+ assert a_ndim in (2, 3), "Expected a of rank 2 or 3, got {}".format(a_ndim)
32
+ assert b_ndim in (2, 3), "Expected b of rank 2 or 3, got {}".format(b_ndim)
33
+
34
+ if transpose_a:
35
+ perm = (1, 0) if a_ndim == 2 else (0, 2, 1)
36
+ a = k_ops.transpose(a, perm)
37
+ if transpose_b:
38
+ perm = (1, 0) if b_ndim == 2 else (0, 2, 1)
39
+ b = k_ops.transpose(b, perm)
40
+ if a_ndim == b_ndim:
41
+ # ...ij,...jk->...ik
42
+ return dot(a, b)
43
+ elif a_ndim == 2:
44
+ # ij,bjk->bik
45
+ return mixed_mode_dot(a, b)
46
+ else: # a_ndim == 3
47
+ # bij,jk->bik
48
+ # Immediately fallback to standard dense matmul, no need to reshape
49
+ return k_ops.matmul(a, b)
k3_node/ops/numpy.py ADDED
@@ -0,0 +1,24 @@
1
+ from keras import backend, ops
2
+ from k3_node.utils.backend_import import *
3
+
4
+
5
+ def polyval(p, x):
6
+ p = ops.convert_to_tensor(p)
7
+
8
+ result = ops.zeros_like(x)
9
+
10
+ for i in range(p.shape[0]):
11
+ result = result * x + p[i]
12
+
13
+ return result
14
+
15
+
16
+ def get_unique(inputs):
17
+ if backend.backend() == "tensorflow":
18
+ return tf.unique(inputs)
19
+ elif backend.backend() == "torch":
20
+ return torch.unique(inputs, return_inverse=True)
21
+ elif backend.backend() == "jax":
22
+ return jnp.unique(inputs, return_inverse=True)
23
+ elif backend.backend() == "numpy":
24
+ return np.unique(inputs, return_inverse=True)
k3_node/ops/segment.py ADDED
@@ -0,0 +1,54 @@
1
+ """Segment reductions (``segment_sum`` / ``segment_max`` / ``segment_min`` / ``segment_prod``).
2
+
3
+ These match ``keras.ops.segment_*``. On the torch backend they are reimplemented so they also
4
+ work on the ``meta`` device, which Keras uses to infer output shapes before the first training
5
+ step: ``keras.ops.segment_*`` fails there (it repeats indices by a tensor-valued count), and the
6
+ fallback it triggers runs on uninitialized memory, producing spurious "index out of range"
7
+ warnings and occasionally huge allocations.
8
+ """
9
+ from keras import backend, ops
10
+
11
+
12
+ def _torch_segment(data, segment_ids, reduction, num_segments):
13
+ import torch
14
+
15
+ def as_tensor(x):
16
+ # keras.ops.convert_to_tensor would move tensors to the default device (a meta tensor
17
+ # becomes uninitialized memory), so keep existing tensors where they are.
18
+ if hasattr(x, "value") and not torch.is_tensor(x):
19
+ x = x.value
20
+ return x if torch.is_tensor(x) else ops.convert_to_tensor(x)
21
+
22
+ data, segment_ids = as_tensor(data), as_tensor(segment_ids).long()
23
+ if num_segments is None:
24
+ num_segments = int(segment_ids.max()) + 1 if segment_ids.numel() > 0 else 0
25
+ num_segments = int(num_segments)
26
+
27
+ # Out-of-range ids go to an extra segment that is dropped at the end (as in Keras).
28
+ segment_ids = torch.where((segment_ids >= 0) & (segment_ids < num_segments), segment_ids, num_segments)
29
+ index = segment_ids.reshape((-1,) + (1,) * (data.dim() - 1)).expand(data.shape)
30
+ fill = {"sum": 0.0, "amax": float("-inf"), "amin": float("inf"), "prod": 1.0}[reduction]
31
+ result = torch.full((num_segments + 1,) + tuple(data.shape[1:]), fill, device=data.device)
32
+ result = result.scatter_reduce(0, index, data.float(), reduction)
33
+ return result[:-1].to(data.dtype)
34
+
35
+
36
+ def segment_sum(data, segment_ids, num_segments=None, sorted=False):
37
+ r"""Sums ``data`` over segments; like :func:`keras.ops.segment_sum`."""
38
+ if backend.backend() == "torch":
39
+ return _torch_segment(data, segment_ids, "sum", num_segments)
40
+ return ops.segment_sum(data, segment_ids, num_segments=num_segments, sorted=sorted)
41
+
42
+
43
+ def segment_max(data, segment_ids, num_segments=None, sorted=False):
44
+ r"""Maximum of ``data`` over segments; like :func:`keras.ops.segment_max`."""
45
+ if backend.backend() == "torch":
46
+ return _torch_segment(data, segment_ids, "amax", num_segments)
47
+ return ops.segment_max(data, segment_ids, num_segments=num_segments, sorted=sorted)
48
+
49
+
50
+ def segment_min(data, segment_ids, num_segments=None, sorted=False):
51
+ r"""Minimum of ``data`` over segments; like ``keras.ops.segment_min``."""
52
+ if backend.backend() == "torch":
53
+ return _torch_segment(data, segment_ids, "amin", num_segments)
54
+ return ops.segment_min(data, segment_ids, num_segments=num_segments, sorted=sorted)
k3_node/ops/sparse.py ADDED
@@ -0,0 +1,51 @@
1
+ import keras
2
+ from keras import ops
3
+
4
+ from k3_node.ops.segment import segment_sum
5
+
6
+
7
+ def spmm(index_targets, index_sources, edge_weight, x, num_targets: int):
8
+ r"""Sums weighted source features into targets:
9
+ :math:`\mathbf{out}_i = \sum_{(j, i)} w_{ji} \, \mathbf{x}_j`, i.e. a sparse-dense matrix product.
10
+
11
+ On torch this uses ``torch.sparse.mm``, which never materializes the ``[num_edges, channels]``
12
+ per-edge messages (the dominant memory cost of wide GCN-style layers). Other backends gather and
13
+ sum, which XLA fuses when compiled.
14
+
15
+ Args:
16
+ index_targets: Target node of every edge, shape ``[num_edges]``.
17
+ index_sources: Source node of every edge, shape ``[num_edges]``.
18
+ edge_weight: Edge weights of shape ``[num_edges]``, or ``None`` for weight 1.
19
+ x: Source node features of shape ``[num_sources, channels]``.
20
+ num_targets (int): The number of target nodes.
21
+
22
+ Example:
23
+ ```python
24
+ import numpy as np
25
+ from k3_node.ops.sparse import spmm
26
+
27
+ x = np.array([[1.0], [2.0], [3.0]], dtype="float32") # one feature per node
28
+ targets, sources = np.array([0, 0, 2]), np.array([1, 2, 0]) # edges 1->0, 2->0, 0->2
29
+ weight = np.array([1.0, 0.5, 2.0], dtype="float32")
30
+
31
+ print(np.asarray(spmm(targets, sources, weight, x, 3)).ravel()) # [3.5 0. 2. ]
32
+ ```
33
+ """
34
+ if keras.config.backend() == "torch":
35
+ import torch
36
+
37
+ x_t = ops.convert_to_tensor(x)
38
+ if x_t.device.type != "meta" and len(x_t.shape) == 2:
39
+ index = torch.stack([ops.convert_to_tensor(index_targets).long(),
40
+ ops.convert_to_tensor(index_sources).long()])
41
+ if edge_weight is None:
42
+ values = torch.ones(index.shape[1], dtype=x_t.dtype, device=x_t.device)
43
+ else:
44
+ values = ops.convert_to_tensor(edge_weight).to(x_t.dtype).reshape(-1)
45
+ adj = torch.sparse_coo_tensor(index, values, (int(num_targets), x_t.shape[0]), check_invariants=True)
46
+ return torch.sparse.mm(adj, x_t)
47
+
48
+ messages = ops.take(x, index_sources, axis=0)
49
+ if edge_weight is not None:
50
+ messages = ops.expand_dims(ops.cast(edge_weight, messages.dtype), -1) * messages
51
+ return segment_sum(messages, ops.cast(index_targets, "int32"), num_segments=num_targets)
@@ -0,0 +1,49 @@
1
+ """GraphRAG & KG-LLM Connectors for K3-Node.
2
+
3
+ Enables extracting subgraphs around retrieved entities, verbalizing structured knowledge
4
+ into LLM prompt contexts (Llama 3, Mistral, ChatML), and projecting GNN (RGCN) and KGE (TransE)
5
+ subgraph embeddings into dense prefix vectors for soft-prompt LLM augmentation.
6
+ """
7
+
8
+ from k3_node.rag.subgraph import (
9
+ SubgraphResult,
10
+ extract_subgraph,
11
+ KGEntityRetriever,
12
+ )
13
+ from k3_node.rag.verbalizer import (
14
+ subgraph_to_triples,
15
+ verbalize_subgraph,
16
+ format_llm_prompt,
17
+ )
18
+ from k3_node.rag.encoders import (
19
+ RGCNSubGraphEncoder,
20
+ TransEPrefixEncoder,
21
+ GNNSubGraphEncoder,
22
+ )
23
+ from k3_node.rag.projector import (
24
+ GraphPrefixProjector,
25
+ KGLLMConnector,
26
+ )
27
+ from k3_node.rag.pipeline import (
28
+ GraphRAG,
29
+ )
30
+
31
+ __all__ = [
32
+ # Subgraph Extraction
33
+ "SubgraphResult",
34
+ "extract_subgraph",
35
+ "KGEntityRetriever",
36
+ # Verbalization & Prompts
37
+ "subgraph_to_triples",
38
+ "verbalize_subgraph",
39
+ "format_llm_prompt",
40
+ # GNN & KGE Encoders
41
+ "RGCNSubGraphEncoder",
42
+ "TransEPrefixEncoder",
43
+ "GNNSubGraphEncoder",
44
+ # Connectors & Projectors
45
+ "GraphPrefixProjector",
46
+ "KGLLMConnector",
47
+ # Unified Pipeline
48
+ "GraphRAG",
49
+ ]