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,374 @@
1
+ import copy
2
+ from collections import defaultdict
3
+ from typing import Any, Callable, Dict, Iterator, List, NamedTuple, Optional, Sequence, Tuple, Union
4
+
5
+ import numpy as np
6
+ from keras import ops
7
+
8
+ from k3_node.data.data import BaseData, Data, size_repr
9
+ from k3_node.data.storage import (
10
+ BaseStorage,
11
+ EdgeStorage,
12
+ NodeStorage,
13
+ get_shape,
14
+ is_tensor_like,
15
+ )
16
+
17
+ NodeType = str
18
+ EdgeType = Tuple[str, str, str]
19
+
20
+
21
+ class HeteroData(BaseData):
22
+ """A data object describing a heterogeneous graph."""
23
+
24
+ def __init__(self, _mapping: Optional[Dict[str, Any]] = None, **kwargs):
25
+ self.__dict__["_node_store_dict"] = {}
26
+ self.__dict__["_edge_store_dict"] = {}
27
+
28
+ for key, value in (_mapping or {}).items():
29
+ setattr(self, key, value)
30
+ for key, value in kwargs.items():
31
+ setattr(self, key, value)
32
+
33
+ def _to_edge_type(self, key: Any) -> Optional[EdgeType]:
34
+ if isinstance(key, tuple):
35
+ if len(key) == 3:
36
+ return (str(key[0]), str(key[1]), str(key[2]))
37
+ if len(key) == 2:
38
+ # Find matching edge type
39
+ matches = [k for k in self.edge_types if k[0] == key[0] and k[-1] == key[1]]
40
+ if len(matches) == 1:
41
+ return matches[0]
42
+ elif len(matches) > 1:
43
+ raise KeyError(f"Ambiguous edge type for '{key}': {matches}")
44
+ return (str(key[0]), "to", str(key[1]))
45
+ elif isinstance(key, str):
46
+ matches = [k for k in self.edge_types if k[1] == key]
47
+ if len(matches) == 1:
48
+ return matches[0]
49
+ elif len(matches) > 1:
50
+ raise KeyError(f"Ambiguous edge type for rel '{key}': {matches}")
51
+ return None
52
+
53
+ def __getitem__(self, key: Any) -> Any:
54
+ edge_key = self._to_edge_type(key)
55
+ if edge_key is not None and (edge_key in self._edge_store_dict or isinstance(key, tuple)):
56
+ if edge_key not in self._edge_store_dict:
57
+ store = EdgeStorage(_parent=self)
58
+ store.__dict__["_key"] = edge_key
59
+ self._edge_store_dict[edge_key] = store
60
+ return self._edge_store_dict[edge_key]
61
+
62
+ if isinstance(key, str):
63
+ if key not in self._node_store_dict:
64
+ store = NodeStorage(_parent=self)
65
+ store.__dict__["_key"] = key
66
+ self._node_store_dict[key] = store
67
+ return self._node_store_dict[key]
68
+
69
+ raise KeyError(f"Invalid key '{key}' for {self.__class__.__name__}")
70
+
71
+ def __setitem__(self, key: Any, value: Any):
72
+ store = self[key]
73
+ if isinstance(value, BaseStorage):
74
+ for k, v in value.items():
75
+ store[k] = v
76
+ elif isinstance(value, dict):
77
+ for k, v in value.items():
78
+ store[k] = v
79
+ else:
80
+ raise ValueError(f"Value for key '{key}' must be a Storage or dict, got {type(value)}")
81
+
82
+ def __delitem__(self, key: Any):
83
+ edge_key = self._to_edge_type(key)
84
+ if edge_key is not None and edge_key in self._edge_store_dict:
85
+ del self._edge_store_dict[edge_key]
86
+ elif isinstance(key, str) and key in self._node_store_dict:
87
+ del self._node_store_dict[key]
88
+ else:
89
+ raise KeyError(key)
90
+
91
+ def __getattr__(self, key: str) -> Any:
92
+ if key in self.__dict__:
93
+ return self.__dict__[key]
94
+ if "_node_store_dict" in self.__dict__ and key in self._node_store_dict:
95
+ return self._node_store_dict[key]
96
+ if key.endswith("_dict") and "_node_store_dict" in self.__dict__:
97
+ return self.collect(key[:-5])
98
+ raise AttributeError(f"'{self.__class__.__name__}' object has no attribute '{key}'")
99
+
100
+ def collect(self, key: str) -> Dict[Any, Any]:
101
+ r"""Returns the attribute ``key`` of every node and edge type that has it, e.g.
102
+ ``data.collect("x")`` (also available as ``data.x_dict``) or ``data.edge_index_dict``."""
103
+ out = {}
104
+ for stores in (self._node_store_dict, self._edge_store_dict):
105
+ for type_, store in stores.items():
106
+ value = store.get(key) if hasattr(store, "get") else getattr(store, key, None)
107
+ if value is not None:
108
+ out[type_] = value
109
+ return out
110
+
111
+ def __setattr__(self, key: str, value: Any):
112
+ if key in ("_node_store_dict", "_edge_store_dict"):
113
+ self.__dict__[key] = value
114
+ elif isinstance(value, BaseStorage):
115
+ if isinstance(value, NodeStorage):
116
+ self._node_store_dict[key] = value
117
+ elif isinstance(value, EdgeStorage):
118
+ edge_key = self._to_edge_type(key)
119
+ self._edge_store_dict[edge_key or key] = value
120
+ else:
121
+ self.__dict__[key] = value
122
+
123
+ @property
124
+ def node_types(self) -> List[NodeType]:
125
+ return list(self._node_store_dict.keys())
126
+
127
+ @property
128
+ def edge_types(self) -> List[EdgeType]:
129
+ return list(self._edge_store_dict.keys())
130
+
131
+ def metadata(self) -> Tuple[List[NodeType], List[EdgeType]]:
132
+ return self.node_types, self.edge_types
133
+
134
+ def get_node_store(self, key: NodeType) -> NodeStorage:
135
+ r"""Gets the NodeStorage object of a particular node type."""
136
+ out = self._node_store_dict.get(key, None)
137
+ if out is None:
138
+ out = NodeStorage(_parent=self)
139
+ out.__dict__["_key"] = key
140
+ self._node_store_dict[key] = out
141
+ return out
142
+
143
+ def get_edge_store(self, src: str, rel: str, dst: str) -> EdgeStorage:
144
+ r"""Gets the EdgeStorage object of a particular edge type given by (src, rel, dst)."""
145
+ key = (src, rel, dst)
146
+ out = self._edge_store_dict.get(key, None)
147
+ if out is None:
148
+ out = EdgeStorage(_parent=self)
149
+ out.__dict__["_key"] = key
150
+ self._edge_store_dict[key] = out
151
+ return out
152
+
153
+ def stores_as(self, data: "HeteroData") -> "HeteroData":
154
+ for node_type in data.node_types:
155
+ self.get_node_store(node_type)
156
+ for edge_type in data.edge_types:
157
+ self.get_edge_store(*edge_type)
158
+ return self
159
+
160
+ @property
161
+ def stores(self) -> List[BaseStorage]:
162
+ return list(self._node_store_dict.values()) + list(self._edge_store_dict.values())
163
+
164
+ @property
165
+ def node_stores(self) -> List[NodeStorage]:
166
+ return list(self._node_store_dict.values())
167
+
168
+ @property
169
+ def edge_stores(self) -> List[EdgeStorage]:
170
+ return list(self._edge_store_dict.values())
171
+
172
+ @property
173
+ def num_nodes_dict(self) -> Dict[NodeType, int]:
174
+ return {k: v.num_nodes for k, v in self._node_store_dict.items()}
175
+
176
+ @property
177
+ def num_edges_dict(self) -> Dict[EdgeType, int]:
178
+ return {k: v.num_edges for k, v in self._edge_store_dict.items()}
179
+
180
+ def set_value_dict(self, key: str, value_dict: Optional[Dict[Any, Any]]):
181
+ for k, v in (value_dict or {}).items():
182
+ self[k][key] = v
183
+ return self
184
+
185
+ def __copy__(self):
186
+ out = self.__class__.__new__(self.__class__)
187
+ out.__dict__["_node_store_dict"] = {}
188
+ out.__dict__["_edge_store_dict"] = {}
189
+ for k, v in self._node_store_dict.items():
190
+ store = copy.copy(v)
191
+ setattr(store, "_parent", out)
192
+ out._node_store_dict[k] = store
193
+ for k, v in self._edge_store_dict.items():
194
+ store = copy.copy(v)
195
+ setattr(store, "_parent", out)
196
+ out._edge_store_dict[k] = store
197
+ return out
198
+
199
+ def __deepcopy__(self, memo=None):
200
+ out = self.__class__.__new__(self.__class__)
201
+ out.__dict__["_node_store_dict"] = {}
202
+ out.__dict__["_edge_store_dict"] = {}
203
+ for k, v in self._node_store_dict.items():
204
+ store = copy.deepcopy(v, memo)
205
+ setattr(store, "_parent", out)
206
+ out._node_store_dict[k] = store
207
+ for k, v in self._edge_store_dict.items():
208
+ store = copy.deepcopy(v, memo)
209
+ setattr(store, "_parent", out)
210
+ out._edge_store_dict[k] = store
211
+ return out
212
+
213
+ def __getstate__(self) -> Dict[str, Any]:
214
+ return self.__dict__.copy()
215
+
216
+ def __setstate__(self, mapping: Dict[str, Any]):
217
+ import weakref
218
+
219
+ for key, value in mapping.items():
220
+ self.__dict__[key] = value
221
+ for store in self.stores:
222
+ store.__dict__["_parent"] = weakref.ref(self)
223
+
224
+ def clone(self) -> "HeteroData":
225
+ return copy.deepcopy(self)
226
+
227
+ def collect(self, key: str, allow_missing: bool = True) -> Dict[Any, Any]:
228
+ mapping = {}
229
+ for k, store in list(self._node_store_dict.items()) + list(self._edge_store_dict.items()):
230
+ if key in store:
231
+ mapping[k] = store[key]
232
+ elif not allow_missing:
233
+ raise KeyError(f"Key '{key}' not found in store '{k}'")
234
+ return mapping
235
+
236
+ def __inc__(self, key: str, value: Any, store: Optional[BaseStorage] = None, *args, **kwargs) -> Any:
237
+ if "batch" in key:
238
+ return int(value.max()) + 1 if is_tensor_like(value) and value.ndim > 0 and value.shape[0] > 0 else 0
239
+ if "index" in key and store is not None and isinstance(store, EdgeStorage):
240
+ edge_type = store._key
241
+ src, _, dst = edge_type
242
+ src_num = self[src].num_nodes
243
+ dst_num = self[dst].num_nodes
244
+ return np.array([[src_num], [dst_num]])
245
+ return 0
246
+
247
+ def __cat_dim__(self, key: str, value: Any, store: Optional[BaseStorage] = None, *args, **kwargs) -> int:
248
+ if key in ("edge_index", "adj_t"):
249
+ return -1
250
+ if is_tensor_like(value) and len(get_shape(value)) == 2 and get_shape(value)[0] == 2 and "index" in key:
251
+ return -1
252
+ return 0
253
+
254
+ def to_dict(self) -> Dict[str, Any]:
255
+ out = {}
256
+ for k, store in self._node_store_dict.items():
257
+ out[k] = store.to_dict()
258
+ for k, store in self._edge_store_dict.items():
259
+ out[k] = store.to_dict()
260
+ return out
261
+
262
+ def to_namedtuple(self) -> NamedTuple:
263
+ # Build nested namedtuple
264
+ node_fields = sorted(list(self._node_store_dict.keys()))
265
+ node_dict = {k: self._node_store_dict[k].to_dict() for k in node_fields}
266
+ edge_fields = [f"{k[0]}__{k[1]}__{k[2]}" for k in sorted(list(self._edge_store_dict.keys()))]
267
+ edge_dict = {f"{k[0]}__{k[1]}__{k[2]}": self._edge_store_dict[k].to_dict() for k in sorted(list(self._edge_store_dict.keys()))}
268
+ all_fields = node_fields + edge_fields
269
+ HeteroTuple = collections.namedtuple("HeteroTuple", all_fields)
270
+ return HeteroTuple(**node_dict, **edge_dict)
271
+
272
+ def edge_type_subgraph(self, edge_types: List[EdgeType]) -> "HeteroData":
273
+ out = copy.deepcopy(self)
274
+ for et in list(out._edge_store_dict.keys()):
275
+ if et not in edge_types:
276
+ del out._edge_store_dict[et]
277
+ return out
278
+
279
+ def subgraph(self, subset_dict: Dict[NodeType, Any]) -> "HeteroData":
280
+ out = copy.deepcopy(self)
281
+ for node_type, subset in subset_dict.items():
282
+ store = out[node_type]
283
+ subset_np = ops.convert_to_numpy(subset)
284
+ indices = np.where(subset_np)[0] if subset_np.dtype == bool else subset_np
285
+ for key in store.node_attrs():
286
+ val = store[key]
287
+ if is_tensor_like(val):
288
+ store[key] = ops.take(val, indices, axis=self.__cat_dim__(key, val, store))
289
+ return out
290
+
291
+ def to_homogeneous(
292
+ self,
293
+ node_attrs: Optional[List[str]] = None,
294
+ edge_attrs: Optional[List[str]] = None,
295
+ add_node_type: bool = True,
296
+ add_edge_type: bool = True,
297
+ dummy_values: bool = True,
298
+ ) -> Data:
299
+ data = Data()
300
+
301
+ # Compute node offsets and slices
302
+ node_slices = {}
303
+ curr_offset = 0
304
+ node_type_list = []
305
+ for i, node_type in enumerate(self.node_types):
306
+ num_nodes = self[node_type].num_nodes
307
+ node_slices[node_type] = curr_offset
308
+ if add_node_type:
309
+ node_type_list.append(np.full((num_nodes,), i, dtype=np.int64))
310
+ curr_offset += num_nodes
311
+
312
+ data.num_nodes = curr_offset
313
+ if add_node_type and len(node_type_list) > 0:
314
+ node_type_arr = np.concatenate(node_type_list, axis=0)
315
+ data.node_type = ops.convert_to_tensor(node_type_arr, dtype="int64")
316
+
317
+ # Concat node features
318
+ if node_attrs is None:
319
+ # find common node attrs across node types
320
+ all_node_attrs = set()
321
+ for store in self.node_stores:
322
+ all_node_attrs.update(store.node_attrs())
323
+ node_attrs = list(all_node_attrs)
324
+
325
+ for attr in node_attrs:
326
+ attr_vals = []
327
+ for node_type in self.node_types:
328
+ val = self[node_type].get(attr)
329
+ if val is not None:
330
+ attr_vals.append(ops.convert_to_numpy(val))
331
+ elif dummy_values:
332
+ num_nodes = self[node_type].num_nodes
333
+ attr_vals.append(np.zeros((num_nodes, 0), dtype=np.float32))
334
+ if len(attr_vals) > 0:
335
+ concat_val = np.concatenate(attr_vals, axis=0)
336
+ data[attr] = ops.convert_to_tensor(concat_val)
337
+
338
+ # Offsetting edge indices
339
+ edge_indices = []
340
+ edge_type_list = []
341
+ for i, edge_type in enumerate(self.edge_types):
342
+ store = self[edge_type]
343
+ edge_index = store.get("edge_index")
344
+ if edge_index is not None:
345
+ ei_np = ops.convert_to_numpy(edge_index).copy()
346
+ src_offset = node_slices[edge_type[0]]
347
+ dst_offset = node_slices[edge_type[-1]]
348
+ ei_np[0] += src_offset
349
+ ei_np[1] += dst_offset
350
+ edge_indices.append(ei_np)
351
+ if add_edge_type:
352
+ edge_type_list.append(np.full((ei_np.shape[1],), i, dtype=np.int64))
353
+
354
+ if len(edge_indices) > 0:
355
+ concat_ei = np.concatenate(edge_indices, axis=1)
356
+ data.edge_index = ops.convert_to_tensor(concat_ei, dtype="int64")
357
+ if add_edge_type and len(edge_type_list) > 0:
358
+ concat_et = np.concatenate(edge_type_list, axis=0)
359
+ data.edge_type = ops.convert_to_tensor(concat_et, dtype="int64")
360
+
361
+ return data
362
+
363
+ def __repr__(self) -> str:
364
+ cls = self.__class__.__name__
365
+ info_lines = []
366
+ for k, store in self._node_store_dict.items():
367
+ attrs = [size_repr(attr, store[attr]) for attr in store.keys()]
368
+ info_lines.append(f" {k}={{{', '.join(attrs)}}}")
369
+ for k, store in self._edge_store_dict.items():
370
+ attrs = [size_repr(attr, store[attr]) for attr in store.keys()]
371
+ edge_name = f"('{k[0]}', '{k[1]}', '{k[2]}')"
372
+ info_lines.append(f" {edge_name}={{{', '.join(attrs)}}}")
373
+ info = ",\n".join(info_lines)
374
+ return f"{cls}(\n{info}\n)" if info else f"{cls}()"
@@ -0,0 +1,59 @@
1
+ from typing import Any, List, Optional
2
+ import numpy as np
3
+ from keras import ops
4
+
5
+ from k3_node.data.data import Data
6
+ from k3_node.data.storage import is_tensor_like
7
+
8
+
9
+ class HypergraphData(Data):
10
+ """A data object describing a hypergraph."""
11
+
12
+ def __init__(
13
+ self,
14
+ x=None,
15
+ edge_index=None,
16
+ edge_attr=None,
17
+ y=None,
18
+ pos=None,
19
+ **kwargs,
20
+ ):
21
+ super().__init__(
22
+ x=x,
23
+ edge_index=edge_index,
24
+ edge_attr=edge_attr,
25
+ y=y,
26
+ pos=pos,
27
+ **kwargs,
28
+ )
29
+
30
+ @property
31
+ def num_edges(self) -> int:
32
+ if self.edge_index is None:
33
+ return 0
34
+ ei_np = ops.convert_to_numpy(self.edge_index)
35
+ if ei_np.size == 0 or ei_np.shape[1] == 0:
36
+ return 0
37
+ return int(np.max(ei_np[1])) + 1
38
+
39
+ @property
40
+ def num_nodes(self) -> Optional[int]:
41
+ num = super().num_nodes
42
+ if self.edge_index is not None and num == self.num_edges:
43
+ ei_np = ops.convert_to_numpy(self.edge_index)
44
+ if ei_np.size > 0 and ei_np.shape[1] > 0:
45
+ return int(np.max(ei_np[0])) + 1
46
+ return num
47
+
48
+ @num_nodes.setter
49
+ def num_nodes(self, num_nodes: Optional[int]) -> None:
50
+ self._store.num_nodes = num_nodes
51
+
52
+ def __inc__(self, key: str, value: Any, *args, **kwargs) -> Any:
53
+ if key == "edge_index":
54
+ return np.array([[self.num_nodes], [self.num_edges]])
55
+ return super().__inc__(key, value, *args, **kwargs)
56
+
57
+
58
+ HyperGraphData = HypergraphData
59
+
@@ -0,0 +1,177 @@
1
+ import copy
2
+ import os
3
+ import os.path as osp
4
+ import pickle
5
+ from typing import Any, Callable, Dict, List, Optional, Tuple, Union
6
+
7
+ from k3_node.data.collate import collate
8
+ from k3_node.data.data import BaseData, Data
9
+ from k3_node.data.dataset import Dataset
10
+ from k3_node.data.separate import separate
11
+
12
+
13
+ class InMemoryDataset(Dataset):
14
+ """Dataset base class for in-memory graph collections."""
15
+
16
+ def __init__(
17
+ self,
18
+ root: Optional[str] = None,
19
+ transform: Optional[Callable] = None,
20
+ pre_transform: Optional[Callable] = None,
21
+ pre_filter: Optional[Callable] = None,
22
+ log: bool = True,
23
+ force_reload: bool = False,
24
+ ):
25
+ # Must be set before `super().__init__()`, which triggers
26
+ # `self.process()` for datasets processed for the first time --
27
+ # `process()` implementations (e.g. `TUDataset`) commonly call
28
+ # `len(self)` / `self.get(idx)`, both of which read these attributes.
29
+ self._data: Optional[BaseData] = None
30
+ self.slices: Optional[Dict[str, Any]] = None
31
+ self.sizes: Dict[str, Any] = {}
32
+ self._data_list: Optional[List[BaseData]] = None
33
+ super().__init__(root, transform, pre_transform, pre_filter, log, force_reload)
34
+
35
+ @property
36
+ def data(self) -> Optional[BaseData]:
37
+ return self._data
38
+
39
+ @data.setter
40
+ def data(self, value: Optional[BaseData]):
41
+ self._data = value
42
+
43
+ def len(self) -> int:
44
+ if self._data_list is not None:
45
+ return len(self._data_list)
46
+ if self.slices is None:
47
+ return 1 if self._data is not None else 0
48
+ for key, value in self.slices.items():
49
+ if isinstance(value, dict):
50
+ for _, sub_val in value.items():
51
+ return len(sub_val) - 1
52
+ return len(value) - 1
53
+ return 0
54
+
55
+ def get(self, idx: int) -> BaseData:
56
+ if self._data_list is not None:
57
+ return self._data_list[idx]
58
+ if self._data is None:
59
+ raise RuntimeError("Dataset does not contain data. Call 'load' first.")
60
+ if self.slices is None:
61
+ if idx == 0:
62
+ return copy.copy(self._data)
63
+ raise IndexError(f"Index {idx} out of bounds for single graph dataset.")
64
+ return separate(self._data.__class__, self._data, idx, self.slices)
65
+
66
+ @classmethod
67
+ def collate(cls, data_list: List[BaseData]) -> Tuple[BaseData, Optional[Dict[str, Any]]]:
68
+ if len(data_list) == 1:
69
+ return data_list[0], None
70
+ base_cls = data_list[0].__class__
71
+ data, slices, _ = collate(
72
+ base_cls,
73
+ data_list,
74
+ increment=False,
75
+ add_batch=False,
76
+ )
77
+ return data, slices
78
+
79
+ def save(self, data_list: List[BaseData], path: str):
80
+ os.makedirs(osp.dirname(path), exist_ok=True)
81
+ data, slices = self.collate(data_list)
82
+ if path.endswith((".pt", ".pth")):
83
+ try:
84
+ import torch
85
+ data_dict = data.to_dict() if hasattr(data, "to_dict") else dict(data)
86
+ torch.save((data_dict, slices), path)
87
+ return
88
+ except Exception:
89
+ pass
90
+ with open(path, "wb") as f:
91
+ pickle.dump((data, slices), f)
92
+
93
+ def load(self, path: str):
94
+ obj = None
95
+ if path.endswith((".pt", ".pth")):
96
+ try:
97
+ import torch
98
+ obj = torch.load(path, map_location="cpu", weights_only=False)
99
+ except Exception:
100
+ pass
101
+
102
+ if obj is None:
103
+ try:
104
+ with open(path, "rb") as f:
105
+ obj = pickle.load(f)
106
+ except Exception:
107
+ try:
108
+ import torch
109
+ obj = torch.load(path, map_location="cpu", weights_only=False)
110
+ except Exception:
111
+ pass
112
+
113
+ if obj is None:
114
+ if hasattr(self, "process") and callable(self.process):
115
+ self.process()
116
+ try:
117
+ with open(path, "rb") as f:
118
+ obj = pickle.load(f)
119
+ except Exception:
120
+ try:
121
+ import torch
122
+ obj = torch.load(path, map_location="cpu", weights_only=False)
123
+ except Exception:
124
+ pass
125
+ if obj is None:
126
+ raise RuntimeError(f"Cannot load dataset from {path}")
127
+
128
+ if isinstance(obj, tuple):
129
+ if len(obj) == 2:
130
+ data, self.slices = obj
131
+ elif len(obj) == 3:
132
+ data, self.slices, extra = obj
133
+ if isinstance(extra, dict):
134
+ self.sizes = extra
135
+ elif len(obj) >= 4:
136
+ data, self.slices = obj[0], obj[1]
137
+ if isinstance(obj[2], dict):
138
+ self.sizes = obj[2]
139
+ else:
140
+ data = obj[0]
141
+ if isinstance(data, dict):
142
+ data = _from_dict(data)
143
+ self._data = data
144
+ elif isinstance(obj, list):
145
+ self._data_list = obj
146
+ elif isinstance(obj, dict):
147
+ self._data = _from_dict(obj)
148
+ else:
149
+ self._data = obj
150
+
151
+ if self._data is not None and hasattr(self._data, "to_backend"):
152
+ try:
153
+ self._data.to_backend()
154
+ except Exception:
155
+ pass
156
+
157
+ if isinstance(self.slices, dict):
158
+ from k3_node.data.storage import is_tensor_like, to_numpy
159
+ self.slices = {k: to_numpy(v) if is_tensor_like(v) else v for k, v in self.slices.items()}
160
+
161
+
162
+
163
+ def _from_dict(mapping):
164
+ """Rebuilds a saved graph: a ``HeteroData`` if the mapping holds per-type attribute
165
+ dictionaries (node types and ``(src, rel, dst)`` edge types), else a ``Data``."""
166
+ if any(isinstance(v, dict) or not isinstance(k, str) for k, v in mapping.items()):
167
+ from k3_node.data.hetero_data import HeteroData
168
+
169
+ data = HeteroData()
170
+ for key, value in mapping.items():
171
+ if isinstance(value, dict):
172
+ for attr, item in value.items():
173
+ setattr(data[key], attr, item)
174
+ else:
175
+ setattr(data, key, value)
176
+ return data
177
+ return Data(**mapping)
@@ -0,0 +1,7 @@
1
+ import os
2
+
3
+
4
+ def makedirs(path: str):
5
+ """Recursively creates a directory."""
6
+ os.makedirs(path, exist_ok=True)
7
+
@@ -0,0 +1,77 @@
1
+ import os
2
+ import os.path as osp
3
+ from typing import Any, Callable, List, Optional, Union
4
+
5
+ from k3_node.data.data import BaseData
6
+ from k3_node.data.database import Database, RocksDatabase, SQLiteDatabase, Schema
7
+ from k3_node.data.dataset import Dataset
8
+
9
+
10
+ class OnDiskDataset(Dataset):
11
+ """Dataset base class for out-of-core graph datasets using a Database backend."""
12
+
13
+ BACKENDS = {
14
+ "sqlite": SQLiteDatabase,
15
+ "rocksdb": RocksDatabase,
16
+ }
17
+
18
+ def __init__(
19
+ self,
20
+ root: str,
21
+ transform: Optional[Callable] = None,
22
+ pre_filter: Optional[Callable] = None,
23
+ backend: str = "sqlite",
24
+ schema: Schema = object,
25
+ log: bool = True,
26
+ ):
27
+ if backend not in self.BACKENDS:
28
+ raise ValueError(f"Database backend must be one of {set(self.BACKENDS.keys())}, got '{backend}'")
29
+
30
+ self.backend = backend
31
+ self.schema = schema
32
+ self._db: Optional[Database] = None
33
+
34
+ super().__init__(root, transform, pre_filter=pre_filter, log=log)
35
+
36
+ @property
37
+ def processed_file_names(self) -> str:
38
+ return f"{self.backend}.db"
39
+
40
+ @property
41
+ def db(self) -> Database:
42
+ if self._db is not None:
43
+ return self._db
44
+
45
+ cls = self.BACKENDS[self.backend]
46
+ os.makedirs(self.processed_dir, exist_ok=True)
47
+ path = osp.join(self.processed_dir, self.processed_file_names)
48
+ self._db = cls(path=path, schema=self.schema)
49
+ return self._db
50
+
51
+ def close(self):
52
+ if self._db is not None:
53
+ self._db.close()
54
+ self._db = None
55
+
56
+ def serialize(self, data: BaseData) -> Any:
57
+ return data
58
+
59
+ def deserialize(self, data: Any) -> BaseData:
60
+ return data
61
+
62
+ def len(self) -> int:
63
+ return len(self.db)
64
+
65
+ def get(self, idx: int) -> BaseData:
66
+ return self.deserialize(self.db[idx])
67
+
68
+ def append(self, data: BaseData):
69
+ idx = len(self)
70
+ self.db[idx] = self.serialize(data)
71
+
72
+ def extend(self, data_list: List[BaseData]):
73
+ start = len(self)
74
+ indices = list(range(start, start + len(data_list)))
75
+ serialized = [self.serialize(d) for d in data_list]
76
+ self.db.multi_insert(indices, serialized)
77
+