k3-node 1.0.0__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
Files changed (459) hide show
  1. k3_node/__init__.py +122 -0
  2. k3_node/applications/__init__.py +17 -0
  3. k3_node/applications/bio/__init__.py +21 -0
  4. k3_node/applications/chemistry/__init__.py +155 -0
  5. k3_node/applications/materials/__init__.py +127 -0
  6. k3_node/applications/materials/basis.py +449 -0
  7. k3_node/applications/materials/chgnet.py +360 -0
  8. k3_node/applications/materials/core.py +351 -0
  9. k3_node/applications/materials/grace.py +246 -0
  10. k3_node/applications/materials/io.py +230 -0
  11. k3_node/applications/materials/m3gnet.py +462 -0
  12. k3_node/applications/materials/megnet.py +395 -0
  13. k3_node/applications/materials/qet.py +220 -0
  14. k3_node/applications/materials/readout.py +235 -0
  15. k3_node/applications/materials/so3net.py +234 -0
  16. k3_node/applications/materials/tensornet.py +381 -0
  17. k3_node/applications/materials/test_materials.py +167 -0
  18. k3_node/applications/materials/wrappers.py +95 -0
  19. k3_node/data/__init__.py +47 -0
  20. k3_node/data/batch.py +102 -0
  21. k3_node/data/collate.py +282 -0
  22. k3_node/data/data.py +532 -0
  23. k3_node/data/database.py +154 -0
  24. k3_node/data/dataset.py +182 -0
  25. k3_node/data/download.py +49 -0
  26. k3_node/data/extract.py +45 -0
  27. k3_node/data/feature_store.py +70 -0
  28. k3_node/data/graph_store.py +92 -0
  29. k3_node/data/hetero_data.py +374 -0
  30. k3_node/data/hypergraph_data.py +59 -0
  31. k3_node/data/in_memory_dataset.py +177 -0
  32. k3_node/data/makedirs.py +7 -0
  33. k3_node/data/on_disk_dataset.py +77 -0
  34. k3_node/data/separate.py +115 -0
  35. k3_node/data/storage.py +593 -0
  36. k3_node/data/temporal.py +154 -0
  37. k3_node/data/test_batch.py +67 -0
  38. k3_node/data/test_data.py +68 -0
  39. k3_node/data/test_dataset_and_stores.py +111 -0
  40. k3_node/data/test_hetero_data.py +33 -0
  41. k3_node/data/test_temporal_and_hyper.py +32 -0
  42. k3_node/data/view.py +43 -0
  43. k3_node/datasets/__init__.py +88 -0
  44. k3_node/datasets/actor.py +101 -0
  45. k3_node/datasets/airports.py +84 -0
  46. k3_node/datasets/amazon.py +66 -0
  47. k3_node/datasets/ba2motif_dataset.py +73 -0
  48. k3_node/datasets/ba_shapes.py +81 -0
  49. k3_node/datasets/bitcoin_otc.py +77 -0
  50. k3_node/datasets/citation_full.py +81 -0
  51. k3_node/datasets/coauthor.py +66 -0
  52. k3_node/datasets/dblp.py +106 -0
  53. k3_node/datasets/digits.py +63 -0
  54. k3_node/datasets/email_eu_core.py +60 -0
  55. k3_node/datasets/entities.py +158 -0
  56. k3_node/datasets/explainer_dataset.py +101 -0
  57. k3_node/datasets/facebook.py +51 -0
  58. k3_node/datasets/fake.py +256 -0
  59. k3_node/datasets/freebase.py +90 -0
  60. k3_node/datasets/geometric_shapes.py +69 -0
  61. k3_node/datasets/github.py +51 -0
  62. k3_node/datasets/graph_generator/__init__.py +6 -0
  63. k3_node/datasets/graph_generator/ba_graph.py +20 -0
  64. k3_node/datasets/graph_generator/base.py +29 -0
  65. k3_node/datasets/graph_generator/er_graph.py +21 -0
  66. k3_node/datasets/icews.py +58 -0
  67. k3_node/datasets/imdb.py +96 -0
  68. k3_node/datasets/jodie.py +56 -0
  69. k3_node/datasets/karate.py +56 -0
  70. k3_node/datasets/lastfm_asia.py +51 -0
  71. k3_node/datasets/mesh_correspondence.py +50 -0
  72. k3_node/datasets/molecule_net.py +148 -0
  73. k3_node/datasets/motif_generator/__init__.py +7 -0
  74. k3_node/datasets/motif_generator/base.py +29 -0
  75. k3_node/datasets/motif_generator/custom.py +17 -0
  76. k3_node/datasets/motif_generator/cycle.py +25 -0
  77. k3_node/datasets/motif_generator/house.py +27 -0
  78. k3_node/datasets/movielens.py +55 -0
  79. k3_node/datasets/planetoid.py +137 -0
  80. k3_node/datasets/polblogs.py +63 -0
  81. k3_node/datasets/ppi.py +189 -0
  82. k3_node/datasets/qm7.py +65 -0
  83. k3_node/datasets/qm9.py +132 -0
  84. k3_node/datasets/reddit.py +121 -0
  85. k3_node/datasets/sbm_dataset.py +165 -0
  86. k3_node/datasets/seal.py +74 -0
  87. k3_node/datasets/shape_scenes.py +92 -0
  88. k3_node/datasets/test_datasets.py +322 -0
  89. k3_node/datasets/tu_dataset.py +131 -0
  90. k3_node/datasets/twitch.py +66 -0
  91. k3_node/datasets/webkb.py +102 -0
  92. k3_node/datasets/wikics.py +85 -0
  93. k3_node/datasets/word_net.py +184 -0
  94. k3_node/etl/__init__.py +37 -0
  95. k3_node/etl/encoders.py +248 -0
  96. k3_node/etl/graph_builders.py +270 -0
  97. k3_node/etl/relational_to_graph.py +201 -0
  98. k3_node/etl/table_to_graph.py +244 -0
  99. k3_node/etl/test_etl.py +318 -0
  100. k3_node/export/__init__.py +15 -0
  101. k3_node/export/cross_backend.py +172 -0
  102. k3_node/export/onnx_exporter.py +190 -0
  103. k3_node/export/runtime.py +254 -0
  104. k3_node/export/tensorrt_exporter.py +201 -0
  105. k3_node/export/test_export.py +337 -0
  106. k3_node/export/tflite_exporter.py +112 -0
  107. k3_node/hub/__init__.py +29 -0
  108. k3_node/hub/dataset_hub.py +242 -0
  109. k3_node/hub/hub_mixin.py +599 -0
  110. k3_node/hub/model_card.py +133 -0
  111. k3_node/hub/test_hub.py +419 -0
  112. k3_node/io/__init__.py +22 -0
  113. k3_node/io/fs.py +117 -0
  114. k3_node/io/npz.py +45 -0
  115. k3_node/io/off.py +29 -0
  116. k3_node/io/planetoid.py +98 -0
  117. k3_node/io/tu.py +137 -0
  118. k3_node/io/txt_array.py +58 -0
  119. k3_node/layers/__init__.py +14 -0
  120. k3_node/layers/aggr/__init__.py +70 -0
  121. k3_node/layers/aggr/attention.py +77 -0
  122. k3_node/layers/aggr/base.py +403 -0
  123. k3_node/layers/aggr/basic.py +412 -0
  124. k3_node/layers/aggr/deep_sets.py +65 -0
  125. k3_node/layers/aggr/deepsets.py +29 -0
  126. k3_node/layers/aggr/equilibrium.py +107 -0
  127. k3_node/layers/aggr/fused.py +43 -0
  128. k3_node/layers/aggr/gmt.py +89 -0
  129. k3_node/layers/aggr/gru.py +58 -0
  130. k3_node/layers/aggr/lcm.py +143 -0
  131. k3_node/layers/aggr/lstm.py +58 -0
  132. k3_node/layers/aggr/mlp.py +75 -0
  133. k3_node/layers/aggr/multi.py +154 -0
  134. k3_node/layers/aggr/patch_transformer.py +137 -0
  135. k3_node/layers/aggr/quantile.py +125 -0
  136. k3_node/layers/aggr/resolver.py +68 -0
  137. k3_node/layers/aggr/scaler.py +133 -0
  138. k3_node/layers/aggr/set2set.py +87 -0
  139. k3_node/layers/aggr/set_transformer.py +107 -0
  140. k3_node/layers/aggr/sort.py +68 -0
  141. k3_node/layers/aggr/test_aggr.py +337 -0
  142. k3_node/layers/aggr/utils.py +210 -0
  143. k3_node/layers/aggr/variance_preserving.py +54 -0
  144. k3_node/layers/attention/__init__.py +5 -0
  145. k3_node/layers/attention/pair_attention.py +448 -0
  146. k3_node/layers/attention/performer.py +187 -0
  147. k3_node/layers/attention/polynormer.py +160 -0
  148. k3_node/layers/attention/qformer.py +143 -0
  149. k3_node/layers/attention/sgformer.py +106 -0
  150. k3_node/layers/attention/test_attention.py +68 -0
  151. k3_node/layers/attention/test_pair_attention.py +91 -0
  152. k3_node/layers/conv/__init__.py +149 -0
  153. k3_node/layers/conv/agnn_conv.py +120 -0
  154. k3_node/layers/conv/antisymmetric_conv.py +94 -0
  155. k3_node/layers/conv/appnp.py +105 -0
  156. k3_node/layers/conv/appnp_conv.py +157 -0
  157. k3_node/layers/conv/arma_conv.py +231 -0
  158. k3_node/layers/conv/cg_conv.py +92 -0
  159. k3_node/layers/conv/cheb_conv.py +137 -0
  160. k3_node/layers/conv/cluster_gcn_conv.py +102 -0
  161. k3_node/layers/conv/conv.py +100 -0
  162. k3_node/layers/conv/crystal_conv.py +140 -0
  163. k3_node/layers/conv/cugraph.py +84 -0
  164. k3_node/layers/conv/diffusion_conv.py +144 -0
  165. k3_node/layers/conv/dir_gnn_conv.py +93 -0
  166. k3_node/layers/conv/dna_conv.py +192 -0
  167. k3_node/layers/conv/edge_conv.py +107 -0
  168. k3_node/layers/conv/eg_conv.py +155 -0
  169. k3_node/layers/conv/fa_conv.py +107 -0
  170. k3_node/layers/conv/feast_conv.py +126 -0
  171. k3_node/layers/conv/film_conv.py +143 -0
  172. k3_node/layers/conv/gat_conv.py +244 -0
  173. k3_node/layers/conv/gated_graph_conv.py +136 -0
  174. k3_node/layers/conv/gatv2_conv.py +205 -0
  175. k3_node/layers/conv/gcn.py +144 -0
  176. k3_node/layers/conv/gcn2_conv.py +126 -0
  177. k3_node/layers/conv/gcn_conv.py +135 -0
  178. k3_node/layers/conv/gen_conv.py +163 -0
  179. k3_node/layers/conv/general_conv.py +218 -0
  180. k3_node/layers/conv/gin_conv.py +218 -0
  181. k3_node/layers/conv/gmm_conv.py +172 -0
  182. k3_node/layers/conv/gps_conv.py +153 -0
  183. k3_node/layers/conv/graph_attention.py +262 -0
  184. k3_node/layers/conv/graph_conv.py +84 -0
  185. k3_node/layers/conv/gravnet_conv.py +93 -0
  186. k3_node/layers/conv/han_conv.py +175 -0
  187. k3_node/layers/conv/heat_conv.py +131 -0
  188. k3_node/layers/conv/hetero_conv.py +128 -0
  189. k3_node/layers/conv/hgt_conv.py +218 -0
  190. k3_node/layers/conv/hypergraph_conv.py +182 -0
  191. k3_node/layers/conv/le_conv.py +81 -0
  192. k3_node/layers/conv/lg_conv.py +58 -0
  193. k3_node/layers/conv/meshcnn_conv.py +84 -0
  194. k3_node/layers/conv/message_passing.py +451 -0
  195. k3_node/layers/conv/mf_conv.py +95 -0
  196. k3_node/layers/conv/mixhop_conv.py +108 -0
  197. k3_node/layers/conv/nn_conv.py +110 -0
  198. k3_node/layers/conv/pan_conv.py +100 -0
  199. k3_node/layers/conv/pdn_conv.py +109 -0
  200. k3_node/layers/conv/pna_conv.py +177 -0
  201. k3_node/layers/conv/point_conv.py +101 -0
  202. k3_node/layers/conv/point_gnn_conv.py +90 -0
  203. k3_node/layers/conv/point_transformer_conv.py +132 -0
  204. k3_node/layers/conv/ppf_conv.py +135 -0
  205. k3_node/layers/conv/ppnp.py +89 -0
  206. k3_node/layers/conv/res_gated_graph_conv.py +126 -0
  207. k3_node/layers/conv/rgat_conv.py +251 -0
  208. k3_node/layers/conv/rgcn_conv.py +321 -0
  209. k3_node/layers/conv/sage_conv.py +154 -0
  210. k3_node/layers/conv/sg_conv.py +96 -0
  211. k3_node/layers/conv/signed_conv.py +100 -0
  212. k3_node/layers/conv/simple_conv.py +75 -0
  213. k3_node/layers/conv/spline_conv.py +182 -0
  214. k3_node/layers/conv/ssg_conv.py +101 -0
  215. k3_node/layers/conv/supergat_conv.py +195 -0
  216. k3_node/layers/conv/tag_conv.py +98 -0
  217. k3_node/layers/conv/test_backend_consistency.py +164 -0
  218. k3_node/layers/conv/test_conv.py +176 -0
  219. k3_node/layers/conv/test_conv_pyg.py +566 -0
  220. k3_node/layers/conv/transformer_conv.py +168 -0
  221. k3_node/layers/conv/utils.py +403 -0
  222. k3_node/layers/conv/wl_conv.py +151 -0
  223. k3_node/layers/conv/x_conv.py +187 -0
  224. k3_node/layers/dense/__init__.py +40 -0
  225. k3_node/layers/dense/dense_gat_conv.py +149 -0
  226. k3_node/layers/dense/dense_gcn_conv.py +117 -0
  227. k3_node/layers/dense/dense_gin_conv.py +88 -0
  228. k3_node/layers/dense/dense_graph_conv.py +95 -0
  229. k3_node/layers/dense/dense_sage_conv.py +85 -0
  230. k3_node/layers/dense/diff_pool.py +76 -0
  231. k3_node/layers/dense/dmon_pool.py +223 -0
  232. k3_node/layers/dense/linear.py +327 -0
  233. k3_node/layers/dense/mincut_pool.py +92 -0
  234. k3_node/layers/dense/test_dense.py +377 -0
  235. k3_node/layers/functional/__init__.py +13 -0
  236. k3_node/layers/functional/bro.py +49 -0
  237. k3_node/layers/functional/edge_dropout.py +55 -0
  238. k3_node/layers/functional/gini.py +44 -0
  239. k3_node/layers/functional/test_functional.py +34 -0
  240. k3_node/layers/kge/__init__.py +17 -0
  241. k3_node/layers/kge/base.py +255 -0
  242. k3_node/layers/kge/complex.py +98 -0
  243. k3_node/layers/kge/distmult.py +79 -0
  244. k3_node/layers/kge/loader.py +50 -0
  245. k3_node/layers/kge/rotate.py +103 -0
  246. k3_node/layers/kge/test_kge.py +76 -0
  247. k3_node/layers/kge/transe.py +96 -0
  248. k3_node/layers/norm/__init__.py +23 -0
  249. k3_node/layers/norm/batch_norm.py +328 -0
  250. k3_node/layers/norm/diff_group_norm.py +141 -0
  251. k3_node/layers/norm/graph_norm.py +105 -0
  252. k3_node/layers/norm/graph_size_norm.py +57 -0
  253. k3_node/layers/norm/instance_norm.py +163 -0
  254. k3_node/layers/norm/layer_norm.py +245 -0
  255. k3_node/layers/norm/mean_subtraction_norm.py +57 -0
  256. k3_node/layers/norm/msg_norm.py +58 -0
  257. k3_node/layers/norm/pair_norm.py +94 -0
  258. k3_node/layers/norm/test_norm.py +275 -0
  259. k3_node/layers/pool/__init__.py +83 -0
  260. k3_node/layers/pool/approx_knn.py +101 -0
  261. k3_node/layers/pool/asap.py +173 -0
  262. k3_node/layers/pool/avg_pool.py +165 -0
  263. k3_node/layers/pool/cluster_pool.py +168 -0
  264. k3_node/layers/pool/connect/__init__.py +10 -0
  265. k3_node/layers/pool/connect/base.py +103 -0
  266. k3_node/layers/pool/connect/filter_edges.py +113 -0
  267. k3_node/layers/pool/consecutive.py +30 -0
  268. k3_node/layers/pool/decimation.py +48 -0
  269. k3_node/layers/pool/edge_pool.py +189 -0
  270. k3_node/layers/pool/glob.py +139 -0
  271. k3_node/layers/pool/graclus.py +66 -0
  272. k3_node/layers/pool/knn.py +253 -0
  273. k3_node/layers/pool/max_pool.py +159 -0
  274. k3_node/layers/pool/mem_pool.py +145 -0
  275. k3_node/layers/pool/pan_pool.py +144 -0
  276. k3_node/layers/pool/point_cloud.py +212 -0
  277. k3_node/layers/pool/pool.py +119 -0
  278. k3_node/layers/pool/sag_pool.py +174 -0
  279. k3_node/layers/pool/select/__init__.py +10 -0
  280. k3_node/layers/pool/select/base.py +112 -0
  281. k3_node/layers/pool/select/topk.py +206 -0
  282. k3_node/layers/pool/test_pool.py +456 -0
  283. k3_node/layers/pool/topk_pool.py +103 -0
  284. k3_node/layers/pool/voxel_grid.py +70 -0
  285. k3_node/layers/unpool/__init__.py +9 -0
  286. k3_node/layers/unpool/knn_interpolate.py +57 -0
  287. k3_node/layers/unpool/test_unpool.py +31 -0
  288. k3_node/loader/__init__.py +62 -0
  289. k3_node/loader/base.py +69 -0
  290. k3_node/loader/cache.py +68 -0
  291. k3_node/loader/cluster.py +127 -0
  292. k3_node/loader/data_list_loader.py +45 -0
  293. k3_node/loader/dataloader.py +117 -0
  294. k3_node/loader/dense_data_loader.py +62 -0
  295. k3_node/loader/dynamic_batch_sampler.py +93 -0
  296. k3_node/loader/graph_saint.py +188 -0
  297. k3_node/loader/hgt_loader.py +90 -0
  298. k3_node/loader/imbalanced_sampler.py +87 -0
  299. k3_node/loader/keras_dataset.py +334 -0
  300. k3_node/loader/link_loader.py +179 -0
  301. k3_node/loader/link_neighbor_loader.py +202 -0
  302. k3_node/loader/mixin.py +190 -0
  303. k3_node/loader/neighbor_loader.py +159 -0
  304. k3_node/loader/neighbor_sampler.py +167 -0
  305. k3_node/loader/node_loader.py +185 -0
  306. k3_node/loader/prefetch.py +115 -0
  307. k3_node/loader/random_node_loader.py +89 -0
  308. k3_node/loader/sampler_utils.py +499 -0
  309. k3_node/loader/shadow.py +115 -0
  310. k3_node/loader/temporal_dataloader.py +98 -0
  311. k3_node/loader/test_dataloader.py +113 -0
  312. k3_node/loader/test_keras_dataset.py +221 -0
  313. k3_node/loader/test_neighbor_loader.py +122 -0
  314. k3_node/loader/test_sampler_utils.py +82 -0
  315. k3_node/loader/test_samplers.py +96 -0
  316. k3_node/loader/test_subgraph_loaders.py +89 -0
  317. k3_node/loader/utils.py +232 -0
  318. k3_node/loader/zip_loader.py +88 -0
  319. k3_node/metrics.py +94 -0
  320. k3_node/models/__init__.py +424 -0
  321. k3_node/models/attentive_fp.py +232 -0
  322. k3_node/models/attract_repel.py +108 -0
  323. k3_node/models/autoencoder.py +318 -0
  324. k3_node/models/basic_gnn.py +443 -0
  325. k3_node/models/bio/__init__.py +4 -0
  326. k3_node/models/captum.py +52 -0
  327. k3_node/models/chemistry/__init__.py +4 -0
  328. k3_node/models/correct_and_smooth.py +146 -0
  329. k3_node/models/deep_graph_infomax.py +113 -0
  330. k3_node/models/deepgcn.py +121 -0
  331. k3_node/models/dimenet.py +737 -0
  332. k3_node/models/dimenet_utils.py +153 -0
  333. k3_node/models/gnnff.py +263 -0
  334. k3_node/models/gps_model.py +1122 -0
  335. k3_node/models/gpse.py +638 -0
  336. k3_node/models/graph_unet.py +199 -0
  337. k3_node/models/graphmae2.py +954 -0
  338. k3_node/models/graphormer.py +1258 -0
  339. k3_node/models/graphormer_3d.py +868 -0
  340. k3_node/models/grover.py +1066 -0
  341. k3_node/models/jumping_knowledge.py +200 -0
  342. k3_node/models/label_prop.py +110 -0
  343. k3_node/models/lightgcn.py +171 -0
  344. k3_node/models/linkx.py +181 -0
  345. k3_node/models/lpformer.py +404 -0
  346. k3_node/models/mask_label.py +114 -0
  347. k3_node/models/materials/__init__.py +33 -0
  348. k3_node/models/meta.py +133 -0
  349. k3_node/models/metapath2vec.py +234 -0
  350. k3_node/models/mlp.py +264 -0
  351. k3_node/models/mole_bert.py +379 -0
  352. k3_node/models/neural_fingerprint.py +95 -0
  353. k3_node/models/node2vec.py +213 -0
  354. k3_node/models/pmlp.py +157 -0
  355. k3_node/models/polynormer.py +229 -0
  356. k3_node/models/rect.py +93 -0
  357. k3_node/models/renet.py +221 -0
  358. k3_node/models/rev_gnn.py +128 -0
  359. k3_node/models/schnet.py +484 -0
  360. k3_node/models/sgformer.py +195 -0
  361. k3_node/models/signed_gcn.py +185 -0
  362. k3_node/models/test_attentive_fp.py +32 -0
  363. k3_node/models/test_attract_repel.py +33 -0
  364. k3_node/models/test_autoencoder.py +119 -0
  365. k3_node/models/test_basic_gnn.py +102 -0
  366. k3_node/models/test_correct_and_smooth.py +40 -0
  367. k3_node/models/test_deep_graph_infomax.py +68 -0
  368. k3_node/models/test_deepgcn.py +21 -0
  369. k3_node/models/test_dimenet.py +86 -0
  370. k3_node/models/test_domain_apis.py +138 -0
  371. k3_node/models/test_gnnff.py +24 -0
  372. k3_node/models/test_gps_model.py +271 -0
  373. k3_node/models/test_gpse.py +34 -0
  374. k3_node/models/test_graph_unet.py +26 -0
  375. k3_node/models/test_graphmae2.py +226 -0
  376. k3_node/models/test_graphormer.py +233 -0
  377. k3_node/models/test_graphormer3d.py +163 -0
  378. k3_node/models/test_grover.py +287 -0
  379. k3_node/models/test_jumping_knowledge.py +129 -0
  380. k3_node/models/test_label_prop.py +37 -0
  381. k3_node/models/test_lightgcn.py +38 -0
  382. k3_node/models/test_linkx.py +31 -0
  383. k3_node/models/test_lpformer.py +22 -0
  384. k3_node/models/test_mask_label.py +90 -0
  385. k3_node/models/test_meta.py +159 -0
  386. k3_node/models/test_metapath2vec.py +45 -0
  387. k3_node/models/test_mlp.py +62 -0
  388. k3_node/models/test_mole_bert.py +164 -0
  389. k3_node/models/test_neural_fingerprint.py +13 -0
  390. k3_node/models/test_node2vec.py +57 -0
  391. k3_node/models/test_pmlp.py +81 -0
  392. k3_node/models/test_polynormer.py +104 -0
  393. k3_node/models/test_rect.py +23 -0
  394. k3_node/models/test_renet.py +32 -0
  395. k3_node/models/test_rev_gnn.py +24 -0
  396. k3_node/models/test_schnet.py +43 -0
  397. k3_node/models/test_sgformer.py +48 -0
  398. k3_node/models/test_signed_gcn.py +28 -0
  399. k3_node/models/test_tgn.py +77 -0
  400. k3_node/models/test_unimol.py +179 -0
  401. k3_node/models/test_unimol2.py +114 -0
  402. k3_node/models/test_unimol_plus.py +131 -0
  403. k3_node/models/test_visnet.py +44 -0
  404. k3_node/models/tgn.py +382 -0
  405. k3_node/models/unimol.py +1156 -0
  406. k3_node/models/unimol2.py +616 -0
  407. k3_node/models/unimol_docking_v2.py +301 -0
  408. k3_node/models/unimol_plus.py +456 -0
  409. k3_node/models/utils.py +97 -0
  410. k3_node/models/visnet.py +759 -0
  411. k3_node/ops/__init__.py +4 -0
  412. k3_node/ops/conv.py +56 -0
  413. k3_node/ops/creation.py +43 -0
  414. k3_node/ops/graph.py +27 -0
  415. k3_node/ops/host.py +41 -0
  416. k3_node/ops/matmul.py +49 -0
  417. k3_node/ops/numpy.py +24 -0
  418. k3_node/ops/segment.py +54 -0
  419. k3_node/ops/sparse.py +51 -0
  420. k3_node/rag/__init__.py +49 -0
  421. k3_node/rag/encoders.py +312 -0
  422. k3_node/rag/pipeline.py +192 -0
  423. k3_node/rag/projector.py +184 -0
  424. k3_node/rag/subgraph.py +270 -0
  425. k3_node/rag/test_rag.py +347 -0
  426. k3_node/rag/verbalizer.py +162 -0
  427. k3_node/tasks/__init__.py +19 -0
  428. k3_node/tasks/backbone_resolver.py +125 -0
  429. k3_node/tasks/base.py +67 -0
  430. k3_node/tasks/graph_classification.py +270 -0
  431. k3_node/tasks/graph_regression.py +228 -0
  432. k3_node/tasks/link_prediction.py +306 -0
  433. k3_node/tasks/node_classification.py +194 -0
  434. k3_node/tasks/node_regression.py +138 -0
  435. k3_node/tasks/test_tasks.py +319 -0
  436. k3_node/test_docstring_examples.py +106 -0
  437. k3_node/test_training_forwarding.py +116 -0
  438. k3_node/training.py +115 -0
  439. k3_node/transforms/__init__.py +166 -0
  440. k3_node/transforms/base_transform.py +32 -0
  441. k3_node/transforms/compose.py +58 -0
  442. k3_node/transforms/general.py +676 -0
  443. k3_node/transforms/graph.py +1070 -0
  444. k3_node/transforms/spatial.py +797 -0
  445. k3_node/transforms/test_random_link_split.py +45 -0
  446. k3_node/transforms/test_spatial_transforms.py +65 -0
  447. k3_node/transforms/test_transforms.py +253 -0
  448. k3_node/transforms/utils.py +102 -0
  449. k3_node/utils/__init__.py +5 -0
  450. k3_node/utils/backend_import.py +12 -0
  451. k3_node/utils/graph.py +286 -0
  452. k3_node/utils/keras.py +94 -0
  453. k3_node/utils/random.py +103 -0
  454. k3_node/utils/smiles.py +235 -0
  455. k3_node-1.0.0.dist-info/METADATA +284 -0
  456. k3_node-1.0.0.dist-info/RECORD +459 -0
  457. k3_node-1.0.0.dist-info/WHEEL +5 -0
  458. k3_node-1.0.0.dist-info/licenses/LICENSE +21 -0
  459. k3_node-1.0.0.dist-info/top_level.txt +1 -0
k3_node/data/data.py ADDED
@@ -0,0 +1,532 @@
1
+ import collections
2
+ import copy
3
+ import warnings
4
+ from collections.abc import Mapping, Sequence
5
+ from itertools import chain
6
+ from typing import Any, Callable, Dict, Iterable, Iterator, List, NamedTuple, Optional, Tuple, Union
7
+
8
+ import numpy as np
9
+ from keras import ops
10
+
11
+ from k3_node.data.storage import (
12
+ BaseStorage,
13
+ EdgeStorage,
14
+ GlobalStorage,
15
+ NodeStorage,
16
+ get_shape,
17
+ is_tensor_like,
18
+ recursive_apply,
19
+ recursive_apply_,
20
+ )
21
+ from k3_node.utils.graph import coalesce, contains_isolated_nodes, has_self_loops, is_undirected, subgraph
22
+
23
+
24
+ def size_repr(key: Any, value: Any, indent: int = 0) -> str:
25
+ pad = " " * indent
26
+ if is_tensor_like(value):
27
+ shape = get_shape(value)
28
+ if len(shape) == 0:
29
+ out = str(ops.convert_to_numpy(value).item())
30
+ else:
31
+ out = str(list(shape))
32
+ elif isinstance(value, str):
33
+ out = f"'{value}'"
34
+ elif isinstance(value, (Sequence, set)) and not isinstance(value, str):
35
+ out = str([len(value)])
36
+ elif isinstance(value, Mapping) and len(value) == 0:
37
+ out = "{}"
38
+ elif isinstance(value, Mapping) and len(value) == 1 and not isinstance(list(value.values())[0], Mapping):
39
+ lines = [size_repr(k, v, 0) for k, v in value.items()]
40
+ out = "{ " + ", ".join(lines) + " }"
41
+ elif isinstance(value, Mapping):
42
+ lines = [size_repr(k, v, indent + 2) for k, v in value.items()]
43
+ out = "{\n" + ",\n".join(lines) + ",\n" + pad + "}"
44
+ else:
45
+ out = str(value)
46
+
47
+ key = str(key).replace("'", "")
48
+ return f"{pad}{key}={out}"
49
+
50
+
51
+ class BaseData:
52
+ def __getattr__(self, key: str) -> Any:
53
+ raise NotImplementedError
54
+
55
+ def __setattr__(self, key: str, value: Any):
56
+ raise NotImplementedError
57
+
58
+ def __delattr__(self, key: str):
59
+ raise NotImplementedError
60
+
61
+ def __getitem__(self, key: str) -> Any:
62
+ raise NotImplementedError
63
+
64
+ def __setitem__(self, key: str, value: Any):
65
+ raise NotImplementedError
66
+
67
+ def __delitem__(self, key: str):
68
+ raise NotImplementedError
69
+
70
+ def __copy__(self):
71
+ raise NotImplementedError
72
+
73
+ def __deepcopy__(self, memo=None):
74
+ raise NotImplementedError
75
+
76
+ def __repr__(self) -> str:
77
+ raise NotImplementedError
78
+
79
+ @property
80
+ def stores(self) -> List[BaseStorage]:
81
+ raise NotImplementedError
82
+
83
+ @property
84
+ def node_stores(self) -> List[NodeStorage]:
85
+ raise NotImplementedError
86
+
87
+ @property
88
+ def edge_stores(self) -> List[EdgeStorage]:
89
+ raise NotImplementedError
90
+
91
+ def stores_as(self, data: "BaseData") -> "BaseData":
92
+ raise NotImplementedError
93
+
94
+ def to_dict(self) -> Dict[str, Any]:
95
+ raise NotImplementedError
96
+
97
+ def to_namedtuple(self) -> NamedTuple:
98
+ raise NotImplementedError
99
+
100
+ def to_backend(self, backend: Optional[str] = None) -> "BaseData":
101
+ for store in self.stores:
102
+ store.to_backend(backend)
103
+ return self
104
+
105
+ def update(self, data: "BaseData") -> "BaseData":
106
+ for store, other_store in zip(self.stores, data.stores):
107
+ for key, value in other_store.items():
108
+ store[key] = value
109
+ return self
110
+
111
+ def __len__(self) -> int:
112
+ return len(self.keys())
113
+
114
+ def __contains__(self, key: str) -> bool:
115
+ return key in self.keys()
116
+
117
+ def keys(self, *args: str) -> List[str]:
118
+ out = []
119
+ for store in self.stores:
120
+ out.extend(list(store.keys(*args)))
121
+ return list(set(out))
122
+
123
+ def values(self, *args: str) -> List[Any]:
124
+ return [self[k] for k in self.keys(*args)]
125
+
126
+ def items(self, *args: str) -> List[Tuple[str, Any]]:
127
+ return [(k, self[k]) for k in self.keys(*args)]
128
+
129
+ @property
130
+ def num_nodes(self) -> Optional[int]:
131
+ try:
132
+ return sum([v.num_nodes for v in self.node_stores])
133
+ except TypeError:
134
+ return None
135
+
136
+ @property
137
+ def num_edges(self) -> int:
138
+ return sum([v.num_edges for v in self.edge_stores])
139
+
140
+ def node_attrs(self) -> List[str]:
141
+ return list(set(chain(*[s.node_attrs() for s in self.node_stores])))
142
+
143
+ def edge_attrs(self) -> List[str]:
144
+ return list(set(chain(*[s.edge_attrs() for s in self.edge_stores])))
145
+
146
+
147
+ class Data(BaseData):
148
+ """A data object describing a homogeneous graph."""
149
+
150
+ def __init__(
151
+ self,
152
+ x=None,
153
+ edge_index=None,
154
+ edge_attr=None,
155
+ y=None,
156
+ pos=None,
157
+ **kwargs,
158
+ ):
159
+ self.__dict__["_store"] = GlobalStorage(_parent=self)
160
+ if x is not None:
161
+ self.x = x
162
+ if edge_index is not None:
163
+ self.edge_index = edge_index
164
+ if edge_attr is not None:
165
+ self.edge_attr = edge_attr
166
+ if y is not None:
167
+ self.y = y
168
+ if pos is not None:
169
+ self.pos = pos
170
+ for key, value in kwargs.items():
171
+ setattr(self, key, value)
172
+
173
+ def __getattr__(self, key: str) -> Any:
174
+ if "_store" not in self.__dict__:
175
+ raise AttributeError(f"'{self.__class__.__name__}' object has no attribute '{key}'")
176
+ try:
177
+ return getattr(self._store, key)
178
+ except AttributeError:
179
+ if key in ('x', 'edge_index', 'edge_attr', 'edge_weight', 'y', 'pos', 'face', 'normal', 'batch'):
180
+ return None
181
+ raise AttributeError(f"'{self.__class__.__name__}' object has no attribute '{key}'") from None
182
+
183
+ def __setattr__(self, key: str, value: Any):
184
+ if key == "_store":
185
+ self.__dict__["_store"] = value
186
+ elif "_store" in self.__dict__:
187
+ setattr(self._store, key, value)
188
+ else:
189
+ self.__dict__[key] = value
190
+
191
+ def __delattr__(self, key: str):
192
+ if key == "_store":
193
+ del self.__dict__["_store"]
194
+ elif "_store" in self.__dict__:
195
+ delattr(self._store, key)
196
+ else:
197
+ del self.__dict__[key]
198
+
199
+ def __getitem__(self, key: str) -> Any:
200
+ return self._store[key]
201
+
202
+ def __setitem__(self, key: str, value: Any):
203
+ self._store[key] = value
204
+
205
+ def __delitem__(self, key: str):
206
+ del self._store[key]
207
+
208
+ def __copy__(self):
209
+ out = self.__class__.__new__(self.__class__)
210
+ for k, v in self.__dict__.items():
211
+ out.__dict__[k] = v
212
+ out._store = copy.copy(self._store)
213
+ out._store._parent = weakref_self = None
214
+ out.__dict__["_store"].__dict__["_parent"] = weakref_self
215
+ setattr(out._store, "_parent", out)
216
+ return out
217
+
218
+ def __deepcopy__(self, memo=None):
219
+ out = self.__class__.__new__(self.__class__)
220
+ for k, v in self.__dict__.items():
221
+ if k == "_store":
222
+ out.__dict__[k] = copy.deepcopy(v, memo)
223
+ else:
224
+ out.__dict__[k] = copy.deepcopy(v, memo)
225
+ setattr(out._store, "_parent", out)
226
+ return out
227
+
228
+ def __getstate__(self) -> Dict[str, Any]:
229
+ return self.__dict__.copy()
230
+
231
+ def __setstate__(self, mapping: Dict[str, Any]):
232
+ import weakref
233
+
234
+ for key, value in mapping.items():
235
+ self.__dict__[key] = value
236
+ if "_store" in self.__dict__ and self._store is not None:
237
+ self._store.__dict__["_parent"] = weakref.ref(self)
238
+
239
+ def clone(self) -> "Data":
240
+ return copy.deepcopy(self)
241
+
242
+ def stores_as(self, data: "Data") -> "Data":
243
+ return self
244
+
245
+ @property
246
+ def stores(self) -> List[BaseStorage]:
247
+ return [self._store]
248
+
249
+ @property
250
+ def node_stores(self) -> List[NodeStorage]:
251
+ return [self._store]
252
+
253
+ @property
254
+ def edge_stores(self) -> List[EdgeStorage]:
255
+ return [self._store]
256
+
257
+ def get(self, key: str, default: Any = None) -> Any:
258
+ return self._store.get(key, default)
259
+
260
+ def __call__(self, *args: str) -> Iterator[Tuple[str, Any]]:
261
+ yield from self._store.items(*args)
262
+
263
+ @property
264
+ def num_features(self) -> int:
265
+ return self._store.num_features
266
+
267
+ @property
268
+ def num_node_features(self) -> int:
269
+ return self._store.num_node_features
270
+
271
+ @property
272
+ def num_edge_features(self) -> int:
273
+ return self._store.num_edge_features
274
+
275
+ @property
276
+ def num_node_types(self) -> int:
277
+ node_type = self.get("node_type")
278
+ return int(np.max(ops.convert_to_numpy(node_type))) + 1 if is_tensor_like(node_type) else 1
279
+
280
+ @property
281
+ def num_edge_types(self) -> int:
282
+ edge_type = self.get("edge_type")
283
+ return int(np.max(ops.convert_to_numpy(edge_type))) + 1 if is_tensor_like(edge_type) else 1
284
+
285
+ @property
286
+ def num_classes(self) -> Optional[int]:
287
+ y = self.get("y")
288
+ if y is not None and is_tensor_like(y):
289
+ y_np = ops.convert_to_numpy(y)
290
+ if np.issubdtype(y_np.dtype, np.integer):
291
+ return int(np.max(y_np)) + 1
292
+ return None
293
+
294
+ def is_directed(self) -> bool:
295
+ return self._store.is_directed()
296
+
297
+ def is_undirected(self) -> bool:
298
+ return self._store.is_undirected()
299
+
300
+ def has_self_loops(self) -> bool:
301
+ return self._store.has_self_loops()
302
+
303
+ def has_isolated_nodes(self) -> bool:
304
+ return self._store.has_isolated_nodes()
305
+
306
+ def is_coalesced(self) -> bool:
307
+ return self._store.is_coalesced()
308
+
309
+ def coalesce(self, reduce: str = "add"):
310
+ self._store.coalesce(reduce=reduce)
311
+ return self
312
+
313
+ def __inc__(self, key: str, value: Any, *args, **kwargs) -> Any:
314
+ if "batch" in key:
315
+ return int(value.max()) + 1 if is_tensor_like(value) and value.ndim > 0 and value.shape[0] > 0 else 0
316
+ if "index" in key or "face" in key:
317
+ return self.num_nodes or 0
318
+ return 0
319
+
320
+ def __cat_dim__(self, key: str, value: Any, *args, **kwargs) -> int:
321
+ if key in ("edge_index", "adj_t", "face"): # [2 or 3, num_edges / num_faces]
322
+ return -1
323
+ if is_tensor_like(value) and len(get_shape(value)) == 2 and get_shape(value)[0] == 2 and "index" in key:
324
+ return -1
325
+ return 0
326
+
327
+ def to_dict(self) -> Dict[str, Any]:
328
+ return self._store.to_dict()
329
+
330
+ def to_namedtuple(self) -> NamedTuple:
331
+ fields = sorted(list(self.keys()))
332
+ DataTuple = collections.namedtuple("DataTuple", fields)
333
+ return DataTuple(**{f: self[f] for f in fields})
334
+
335
+ @classmethod
336
+ def from_dict(cls, mapping: Dict[str, Any]) -> "Data":
337
+ return cls(**mapping)
338
+
339
+ def apply(self, func: Callable, *keys: str) -> "Data":
340
+ self._store.apply(func, *keys)
341
+ return self
342
+
343
+ def apply_(self, func: Callable, *keys: str) -> "Data":
344
+ self._store.apply_(func, *keys)
345
+ return self
346
+
347
+ def to(self, *args, **kwargs) -> "Data":
348
+ self._store.to(*args, **kwargs)
349
+ return self
350
+
351
+ def to_backend(self, backend: Optional[str] = None) -> "Data":
352
+ self._store.to_backend(backend)
353
+ return self
354
+
355
+ def cpu(self) -> "Data":
356
+ self._store.cpu()
357
+ return self
358
+
359
+ def cuda(self) -> "Data":
360
+ self._store.cuda()
361
+ return self
362
+
363
+ def requires_grad_(self, *keys: str) -> "Data":
364
+ self._store.requires_grad_(*keys)
365
+ return self
366
+
367
+ def contiguous(self, *keys: str) -> "Data":
368
+ self._store.contiguous(*keys)
369
+ return self
370
+
371
+ def subgraph(self, subset) -> "Data":
372
+ """Returns the induced subgraph for subset nodes."""
373
+ data = copy.copy(self)
374
+ num_nodes = self.num_nodes
375
+ sub_edge_index, sub_edge_attr = subgraph(
376
+ subset,
377
+ self.edge_index,
378
+ edge_attr=self.get("edge_attr"),
379
+ relabel_nodes=True,
380
+ num_nodes=num_nodes,
381
+ )
382
+ data.edge_index = sub_edge_index
383
+ if sub_edge_attr is not None:
384
+ data.edge_attr = sub_edge_attr
385
+ subset_np = ops.convert_to_numpy(subset)
386
+ if subset_np.dtype == bool:
387
+ indices = np.where(subset_np)[0]
388
+ else:
389
+ indices = subset_np
390
+
391
+ for key in self.node_attrs():
392
+ val = self[key]
393
+ if is_tensor_like(val):
394
+ data[key] = ops.take(val, indices, axis=self.__cat_dim__(key, val))
395
+ if "num_nodes" in self._store: # an explicitly stored node count must shrink too
396
+ data.num_nodes = int(len(indices))
397
+ return data
398
+
399
+ def edge_subgraph(self, subset) -> "Data":
400
+ """Returns the graph with only the edges in ``subset`` (a boolean edge mask or edge
401
+ indices). All nodes are kept; every edge-level attribute is filtered.
402
+
403
+ Example:
404
+ ```python
405
+ import numpy as np
406
+ from k3_node.data import Data
407
+
408
+ data = Data(edge_index=np.array([[0, 1, 2], [1, 2, 0]]), edge_type=np.array([0, 1, 0]), num_nodes=3)
409
+ train = data.edge_subgraph(np.array([True, False, True]))
410
+ print(tuple(train.edge_index.shape), train.num_nodes) # (2, 2) 3
411
+ ```
412
+ """
413
+ subset_np = np.asarray(ops.convert_to_numpy(subset))
414
+ indices = np.where(subset_np)[0] if subset_np.dtype == bool else subset_np
415
+ data = copy.copy(self)
416
+ for key in self.edge_attrs():
417
+ val = self[key]
418
+ if is_tensor_like(val):
419
+ data[key] = ops.take(val, indices, axis=self.__cat_dim__(key, val))
420
+ return data
421
+
422
+ def to_heterogeneous(self, node_type: str = "0", edge_type: Tuple[str, str, str] = ("0", "0", "0")):
423
+ from k3_node.data.hetero_data import HeteroData
424
+
425
+ hetero = HeteroData()
426
+ if hasattr(self, "edge_type") and self.edge_type is not None:
427
+ edge_type_np = ops.convert_to_numpy(self.edge_type)
428
+ unique_edge_types = np.unique(edge_type_np)
429
+ hetero[node_type].x = self.x
430
+ for et in unique_edge_types:
431
+ mask = edge_type_np == et
432
+ sub_edge_index = self.edge_index[:, mask]
433
+ hetero[node_type, str(et), node_type].edge_index = sub_edge_index
434
+ if "edge_attr" in self and self.edge_attr is not None:
435
+ hetero[node_type, str(et), node_type].edge_attr = self.edge_attr[mask]
436
+ else:
437
+ for k in self.node_attrs():
438
+ hetero[node_type][k] = self[k]
439
+ for k in self.edge_attrs():
440
+ hetero[edge_type][k] = self[k]
441
+ return hetero
442
+
443
+ def validate(self, raise_on_error: bool = True) -> bool:
444
+ num_nodes = self.num_nodes
445
+ if "edge_index" in self and self.edge_index is not None:
446
+ edge_shape = get_shape(self.edge_index)
447
+ if len(edge_shape) != 2 or edge_shape[0] != 2:
448
+ msg = f"'edge_index' must have shape [2, num_edges], got {edge_shape}"
449
+ if raise_on_error:
450
+ raise ValueError(msg)
451
+ warnings.warn(msg)
452
+ return False
453
+ if num_nodes is not None and edge_shape[1] > 0:
454
+ edge_max = int(np.max(ops.convert_to_numpy(self.edge_index)))
455
+ if edge_max >= num_nodes:
456
+ msg = f"'edge_index' references node {edge_max}, but num_nodes is {num_nodes}"
457
+ if raise_on_error:
458
+ raise ValueError(msg)
459
+ warnings.warn(msg)
460
+ return False
461
+ return True
462
+
463
+ @property
464
+ def inputs(self):
465
+ r"""Returns the tuple of input tensors ``(x, edge_index)`` or ``(x, edge_index, edge_attr)``."""
466
+ if hasattr(self, "edge_attr") and self.edge_attr is not None:
467
+ return (self.x, self.edge_index, self.edge_attr)
468
+ return (self.x, self.edge_index)
469
+
470
+ def to_generator(self, mask: Optional[str] = "train_mask", repeat: bool = True):
471
+ r"""Generates tuples of ((x, edge_index), y, mask) or ((x, edge_index, edge_attr), y, mask)
472
+ ready for training directly with Keras `model.fit()`.
473
+
474
+ Args:
475
+ mask (str, optional): The name of the mask attribute (e.g. ``'train_mask'``,
476
+ ``'val_mask'``, ``'test_mask'``) to use as sample_weight for loss masking.
477
+ If :obj:`None`, no mask is applied. (default: ``'train_mask'``)
478
+ repeat (bool, optional): Whether to yield infinitely for Keras generator training.
479
+ (default: :obj:`True`)
480
+ """
481
+ x = ops.convert_to_tensor(self.x, dtype="float32") if self.x is not None else None
482
+ edge_index = ops.convert_to_tensor(self.edge_index, dtype="int64") if self.edge_index is not None else None
483
+
484
+ inputs = (x, edge_index)
485
+ if hasattr(self, "edge_attr") and self.edge_attr is not None:
486
+ edge_attr = ops.convert_to_tensor(self.edge_attr, dtype="float32")
487
+ inputs = (x, edge_index, edge_attr)
488
+
489
+ y = ops.convert_to_tensor(self.y, dtype="int64") if self.y is not None else None
490
+
491
+ sample_weight = None
492
+ if mask is not None and hasattr(self, mask) and getattr(self, mask) is not None:
493
+ sample_weight = ops.cast(getattr(self, mask), "float32")
494
+
495
+ while True:
496
+ if sample_weight is not None:
497
+ yield inputs, y, sample_weight
498
+ elif y is not None:
499
+ yield inputs, y
500
+ else:
501
+ yield inputs
502
+ if not repeat:
503
+ break
504
+
505
+ def accuracy(self, logits_or_pred: Any, mask: Optional[str] = "test_mask") -> float:
506
+ r"""Convenience method to calculate classification accuracy.
507
+
508
+ Args:
509
+ logits_or_pred: Model prediction output (either class logits or predicted labels).
510
+ mask (str, optional): The mask attribute (e.g. ``'test_mask'``) to evaluate on.
511
+ If :obj:`None`, evaluates across all nodes. (default: ``'test_mask'``)
512
+ """
513
+ pred = logits_or_pred
514
+ shape = ops.shape(pred)
515
+ if len(shape) > 1 and shape[-1] > 1:
516
+ pred = ops.argmax(pred, axis=-1)
517
+
518
+ y = self.y
519
+ if mask is not None and hasattr(self, mask) and getattr(self, mask) is not None:
520
+ m = getattr(self, mask)
521
+ pred = pred[m]
522
+ y = y[m]
523
+
524
+ correct = ops.cast(ops.cast(pred, "int64") == ops.cast(y, "int64"), "float32")
525
+ return float(ops.convert_to_numpy(ops.mean(correct)))
526
+
527
+ def __repr__(self) -> str:
528
+ cls = self.__class__.__name__
529
+ attrs = [size_repr(k, v) for k, v in self._store.items()]
530
+ info = ", ".join(attrs)
531
+ return f"{cls}({info})"
532
+
@@ -0,0 +1,154 @@
1
+ import io
2
+ import pickle
3
+ import sqlite3
4
+ from abc import ABC, abstractmethod
5
+ from typing import Any, Dict, List, Optional, Sequence, Union
6
+
7
+ Schema = Any
8
+
9
+
10
+ class Database(ABC):
11
+ """Base class for key/value and index-based graph databases."""
12
+
13
+ def __init__(self, schema: Schema = object):
14
+ self.schema = schema
15
+
16
+ @abstractmethod
17
+ def connect(self):
18
+ pass
19
+
20
+ @abstractmethod
21
+ def close(self):
22
+ pass
23
+
24
+ @abstractmethod
25
+ def insert(self, index: int, data: Any):
26
+ pass
27
+
28
+ def multi_insert(self, indices: Sequence[int], data_list: Sequence[Any]):
29
+ for idx, d in zip(indices, data_list):
30
+ self.insert(idx, d)
31
+
32
+ @abstractmethod
33
+ def get(self, index: int) -> Any:
34
+ pass
35
+
36
+ def multi_get(self, indices: Sequence[int]) -> List[Any]:
37
+ return [self.get(idx) for idx in indices]
38
+
39
+ @abstractmethod
40
+ def __len__(self) -> int:
41
+ pass
42
+
43
+ def __getitem__(self, idx: Any) -> Any:
44
+ if isinstance(idx, int):
45
+ return self.get(idx)
46
+ elif isinstance(idx, slice):
47
+ start = idx.start or 0
48
+ stop = idx.stop or len(self)
49
+ step = idx.step or 1
50
+ return self.multi_get(range(start, stop, step))
51
+ elif isinstance(idx, (list, tuple)):
52
+ return self.multi_get(idx)
53
+ else:
54
+ return self.get(int(idx))
55
+
56
+ def __setitem__(self, idx: Any, value: Any):
57
+ if isinstance(idx, int):
58
+ self.insert(idx, value)
59
+ elif isinstance(idx, slice):
60
+ start = idx.start or 0
61
+ stop = idx.stop or len(self)
62
+ step = idx.step or 1
63
+ indices = list(range(start, stop, step))
64
+ self.multi_insert(indices, value)
65
+ elif isinstance(idx, (list, tuple)):
66
+ self.multi_insert(idx, value)
67
+ else:
68
+ self.insert(int(idx), value)
69
+
70
+
71
+ class SQLiteDatabase(Database):
72
+ """SQLite-backed persistent database."""
73
+
74
+ def __init__(self, path: str, name: str = "data", schema: Schema = object):
75
+ super().__init__(schema)
76
+ self.path = path
77
+ self.name = name
78
+ self._conn: Optional[sqlite3.Connection] = None
79
+ self.connect()
80
+
81
+ def connect(self):
82
+ self._conn = sqlite3.connect(self.path)
83
+ with self._conn:
84
+ self._conn.execute(
85
+ f"CREATE TABLE IF NOT EXISTS {self.name} (id INTEGER PRIMARY KEY, val BLOB)"
86
+ )
87
+
88
+ def close(self):
89
+ if self._conn is not None:
90
+ self._conn.close()
91
+ self._conn = None
92
+
93
+ def insert(self, index: int, data: Any):
94
+ buf = io.BytesIO()
95
+ pickle.dump(data, buf)
96
+ raw = buf.getvalue()
97
+ with self._conn:
98
+ self._conn.execute(
99
+ f"INSERT OR REPLACE INTO {self.name} (id, val) VALUES (?, ?)",
100
+ (index, raw),
101
+ )
102
+
103
+ def multi_insert(self, indices: Sequence[int], data_list: Sequence[Any]):
104
+ rows = []
105
+ for idx, d in zip(indices, data_list):
106
+ buf = io.BytesIO()
107
+ pickle.dump(d, buf)
108
+ rows.append((idx, buf.getvalue()))
109
+ with self._conn:
110
+ self._conn.executemany(
111
+ f"INSERT OR REPLACE INTO {self.name} (id, val) VALUES (?, ?)",
112
+ rows,
113
+ )
114
+
115
+ def get(self, index: int) -> Any:
116
+ cursor = self._conn.execute(
117
+ f"SELECT val FROM {self.name} WHERE id = ?", (index,)
118
+ )
119
+ row = cursor.fetchone()
120
+ if row is None:
121
+ raise KeyError(f"Index {index} not found in database")
122
+ return pickle.loads(row[0])
123
+
124
+ def multi_get(self, indices: Sequence[int]) -> List[Any]:
125
+ return [self.get(idx) for idx in indices]
126
+
127
+ def __len__(self) -> int:
128
+ cursor = self._conn.execute(f"SELECT COUNT(*) FROM {self.name}")
129
+ return cursor.fetchone()[0]
130
+
131
+
132
+ class RocksDatabase(Database):
133
+ """RocksDB-backed database stub."""
134
+
135
+ def __init__(self, path: str, schema: Schema = object):
136
+ super().__init__(schema)
137
+ self.path = path
138
+ raise NotImplementedError("RocksDatabase requires rocksdb C++ binding; use SQLiteDatabase instead.")
139
+
140
+ def connect(self):
141
+ pass
142
+
143
+ def close(self):
144
+ pass
145
+
146
+ def insert(self, index: int, data: Any):
147
+ pass
148
+
149
+ def get(self, index: int) -> Any:
150
+ pass
151
+
152
+ def __len__(self) -> int:
153
+ return 0
154
+