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,31 @@
1
+ from keras import ops
2
+
3
+ from k3_node.layers.unpool import knn_interpolate
4
+
5
+
6
+ def test_knn_interpolate():
7
+ x = ops.convert_to_tensor(
8
+ [[1.0], [10.0], [100.0], [-1.0], [-10.0], [-100.0]], dtype="float32"
9
+ )
10
+ pos_x = ops.convert_to_tensor([
11
+ [-1.0, 0.0], [0.0, 0.0], [1.0, 0.0],
12
+ [-2.0, 0.0], [0.0, 0.0], [2.0, 0.0],
13
+ ], dtype="float32")
14
+ pos_y = ops.convert_to_tensor([
15
+ [-1.0, -1.0], [1.0, 1.0], [-2.0, -2.0], [2.0, 2.0],
16
+ ], dtype="float32")
17
+ batch_x = ops.convert_to_tensor([0, 0, 0, 1, 1, 1], dtype="int64")
18
+ batch_y = ops.convert_to_tensor([0, 0, 1, 1], dtype="int64")
19
+
20
+ y = knn_interpolate(x, pos_x, pos_y, batch_x, batch_y, k=2)
21
+ assert ops.shape(y) == (4, 1)
22
+ assert ops.convert_to_numpy(y).tolist() == [[4.0], [70.0], [-4.0], [-70.0]]
23
+
24
+
25
+ def test_knn_interpolate_no_batch():
26
+ x = ops.convert_to_tensor([[1.0], [10.0], [100.0]], dtype="float32")
27
+ pos_x = ops.convert_to_tensor([[-1.0, 0.0], [0.0, 0.0], [1.0, 0.0]], dtype="float32")
28
+ pos_y = ops.convert_to_tensor([[-1.0, -1.0], [1.0, 1.0]], dtype="float32")
29
+
30
+ y = knn_interpolate(x, pos_x, pos_y, k=2)
31
+ assert ops.shape(y) == (2, 1)
@@ -0,0 +1,62 @@
1
+ from .base import DataLoaderIterator
2
+ from .cache import CachedLoader
3
+ from .cluster import ClusterData, ClusterLoader
4
+ from .data_list_loader import DataListLoader
5
+ from .dataloader import Collater, DataLoader
6
+ from .dense_data_loader import DenseDataLoader
7
+ from .dynamic_batch_sampler import DynamicBatchSampler
8
+ from .keras_dataset import FullGraphDataset
9
+ from .graph_saint import (
10
+ GraphSAINTEdgeSampler,
11
+ GraphSAINTNodeSampler,
12
+ GraphSAINTRandomWalkSampler,
13
+ GraphSAINTSampler,
14
+ )
15
+ from .hgt_loader import HGTLoader
16
+ from .imbalanced_sampler import ImbalancedSampler
17
+ from .link_loader import EdgeSamplerInput, LinkLoader
18
+ from .link_neighbor_loader import LinkNeighborLoader
19
+ from .mixin import AffinityMixin, LogMemoryMixin, MultithreadingMixin
20
+ from .neighbor_loader import NeighborLoader
21
+ from .neighbor_sampler import Adj, EdgeIndex, NeighborSampler
22
+ from .node_loader import HeteroSamplerOutput, NodeLoader, NodeSamplerInput, SamplerOutput
23
+ from .prefetch import DeviceHelper, PrefetchLoader
24
+ from .random_node_loader import RandomNodeLoader
25
+ from .shadow import ShaDowKHopSampler
26
+ from .temporal_dataloader import TemporalDataLoader
27
+ from .zip_loader import ZipLoader
28
+ from .utils import to_numpy
29
+
30
+ __all__ = [
31
+ 'to_numpy',
32
+ 'DataLoader',
33
+ 'NodeLoader',
34
+ 'LinkLoader',
35
+ 'NeighborLoader',
36
+ 'LinkNeighborLoader',
37
+ 'HGTLoader',
38
+ 'ClusterData',
39
+ 'ClusterLoader',
40
+ 'GraphSAINTSampler',
41
+ 'GraphSAINTNodeSampler',
42
+ 'GraphSAINTEdgeSampler',
43
+ 'GraphSAINTRandomWalkSampler',
44
+ 'ShaDowKHopSampler',
45
+ 'RandomNodeLoader',
46
+ 'ZipLoader',
47
+ 'DataListLoader',
48
+ 'DenseDataLoader',
49
+ 'TemporalDataLoader',
50
+ 'NeighborSampler',
51
+ 'ImbalancedSampler',
52
+ 'DynamicBatchSampler',
53
+ 'PrefetchLoader',
54
+ 'CachedLoader',
55
+ 'AffinityMixin',
56
+ 'MultithreadingMixin',
57
+ 'LogMemoryMixin',
58
+ 'Collater',
59
+ 'DataLoaderIterator',
60
+ 'FullGraphDataset',
61
+ ]
62
+
k3_node/loader/base.py ADDED
@@ -0,0 +1,69 @@
1
+ from typing import Any, Callable
2
+
3
+ try:
4
+ import torch
5
+ BaseDataLoader = torch.utils.data.DataLoader
6
+ except ImportError:
7
+ class BaseDataLoader:
8
+ r"""Fallback BaseDataLoader when PyTorch is not installed."""
9
+ def __init__(self, dataset=None, batch_size=1, shuffle=False, **kwargs):
10
+ self.dataset = dataset
11
+ self.batch_size = batch_size
12
+ self.shuffle = shuffle
13
+
14
+ def __iter__(self):
15
+ collate_fn = getattr(self, 'collate_fn', None) or (lambda x: x)
16
+ dataset = getattr(self, 'dataset', [])
17
+ batch_size = getattr(self, 'batch_size', 1)
18
+ shuffle = getattr(self, 'shuffle', False)
19
+ indices = list(range(len(dataset)))
20
+ if shuffle:
21
+ import random
22
+ random.shuffle(indices)
23
+ for i in range(0, len(indices), batch_size):
24
+ batch_indices = indices[i:i + batch_size]
25
+ batch = [dataset[idx] for idx in batch_indices]
26
+ yield collate_fn(batch)
27
+
28
+ def __len__(self):
29
+ dataset = getattr(self, 'dataset', [])
30
+ batch_size = getattr(self, 'batch_size', 1)
31
+ return (len(dataset) + batch_size - 1) // batch_size if len(dataset) > 0 else 0
32
+
33
+ try:
34
+ from torch.utils.data.dataloader import (
35
+ _BaseDataLoaderIter,
36
+ _MultiProcessingDataLoaderIter,
37
+ )
38
+ except ImportError:
39
+ class _BaseDataLoaderIter:
40
+ pass
41
+
42
+ class _MultiProcessingDataLoaderIter:
43
+ pass
44
+
45
+
46
+ class DataLoaderIterator:
47
+ r"""A data loader iterator extended by a post transformation function
48
+ :meth:`transform_fn`.
49
+ """
50
+ def __init__(self, iterator: Any, transform_fn: Callable):
51
+ self.iterator = iterator
52
+ self.transform_fn = transform_fn
53
+
54
+ def __iter__(self) -> 'DataLoaderIterator':
55
+ return self
56
+
57
+ def _reset(self, loader: Any, first_iter: bool = False):
58
+ if hasattr(self.iterator, '_reset'):
59
+ self.iterator._reset(loader, first_iter)
60
+
61
+ def __len__(self) -> int:
62
+ return len(self.iterator)
63
+
64
+ def __next__(self) -> Any:
65
+ return self.transform_fn(next(self.iterator))
66
+
67
+ def __del__(self) -> Any:
68
+ if isinstance(self.iterator, _MultiProcessingDataLoaderIter):
69
+ self.iterator.__del__()
@@ -0,0 +1,68 @@
1
+ from collections.abc import Mapping, Sequence
2
+ from typing import Any, Callable, List, Optional
3
+
4
+ try:
5
+ import torch
6
+ from torch.utils.data import DataLoader
7
+ except ImportError:
8
+ torch = None
9
+ DataLoader = object
10
+
11
+
12
+ def to_device(inputs: Any, device: Optional[Any] = None) -> Any:
13
+ if device is None:
14
+ return inputs
15
+ if hasattr(inputs, 'to'):
16
+ return inputs.to(device)
17
+ elif isinstance(inputs, Mapping):
18
+ return {key: to_device(value, device) for key, value in inputs.items()}
19
+ elif isinstance(inputs, tuple) and hasattr(inputs, '_fields'):
20
+ return type(inputs)(*(to_device(s, device) for s in zip(*inputs)))
21
+ elif isinstance(inputs, Sequence) and not isinstance(inputs, str):
22
+ return [to_device(s, device) for s in zip(*inputs)]
23
+ return inputs
24
+
25
+
26
+ class CachedLoader:
27
+ r"""A loader to cache mini-batch outputs in memory across epochs.
28
+
29
+ Args:
30
+ loader (DataLoader): The data loader.
31
+ device (torch.device, optional): The device to load the data to. (default: :obj:`None`)
32
+ transform (callable, optional): A function that takes in a sampled mini-batch and returns a transformed version. (default: :obj:`None`)
33
+ """
34
+ def __init__(
35
+ self,
36
+ loader: Any,
37
+ device: Optional[Any] = None,
38
+ transform: Optional[Callable] = None,
39
+ ):
40
+ self.loader = loader
41
+ self.device = device
42
+ self.transform = transform
43
+ self._cache: List[Any] = []
44
+
45
+ def clear(self):
46
+ r"""Clears the cache."""
47
+ self._cache = []
48
+
49
+ def __iter__(self) -> Any:
50
+ if len(self._cache) > 0:
51
+ for batch in self._cache:
52
+ yield batch
53
+ return
54
+
55
+ for batch in self.loader:
56
+ if self.transform is not None:
57
+ batch = self.transform(batch)
58
+
59
+ batch = to_device(batch, self.device)
60
+ self._cache.append(batch)
61
+ yield batch
62
+
63
+ def __len__(self) -> int:
64
+ return len(self.loader)
65
+
66
+ def __repr__(self) -> str:
67
+ return f'{self.__class__.__name__}({self.loader})'
68
+
@@ -0,0 +1,127 @@
1
+ import copy
2
+ from typing import Any, List, Optional, Union
3
+
4
+ import numpy as np
5
+
6
+ try:
7
+ import torch
8
+ import torch.utils.data
9
+ from torch import Tensor
10
+ BaseDataset = torch.utils.data.Dataset
11
+ BaseDataLoader = torch.utils.data.DataLoader
12
+ except ImportError:
13
+ torch = None
14
+ Tensor = type(None)
15
+ BaseDataset = object
16
+ BaseDataLoader = object
17
+
18
+ from k3_node.data import Data
19
+ from k3_node.loader.sampler_utils import partition_graph
20
+ from k3_node.loader.keras_dataset import loader_bases
21
+
22
+
23
+ class ClusterData(BaseDataset):
24
+ r"""Clusters/partitions a graph data object into multiple subgraphs, as
25
+ motivated by the "Cluster-GCN" paper.
26
+
27
+ Args:
28
+ data (Data): The graph data object.
29
+ num_parts (int): The number of partitions.
30
+ recursive (bool, optional): Multilevel recursive bisection if True. (default: :obj:`False`)
31
+ save_dir (str, optional): Directory to save partitioned data. (default: :obj:`None`)
32
+ log (bool, optional): If set to :obj:`False`, will not log. (default: :obj:`True`)
33
+ keep_inter_cluster_edges (bool, optional): Keep inter-cluster connections. (default: :obj:`False`)
34
+ """
35
+ def __init__(
36
+ self,
37
+ data: Data,
38
+ num_parts: int,
39
+ recursive: bool = False,
40
+ save_dir: Optional[str] = None,
41
+ filename: Optional[str] = None,
42
+ log: bool = True,
43
+ keep_inter_cluster_edges: bool = False,
44
+ sparse_format: str = 'csr',
45
+ ):
46
+ assert data.edge_index is not None
47
+
48
+ self.num_parts = num_parts
49
+ self.recursive = recursive
50
+ self.keep_inter_cluster_edges = keep_inter_cluster_edges
51
+ self.sparse_format = sparse_format
52
+ self.data = data
53
+
54
+ loaded = False
55
+ if save_dir is not None:
56
+ import os.path as osp
57
+ recursive_str = '_recursive' if recursive else ''
58
+ root_dir = osp.join(save_dir, f'part_{num_parts}{recursive_str}')
59
+ path = osp.join(root_dir, filename or 'metis.pt')
60
+ if osp.exists(path):
61
+ try:
62
+ import torch
63
+ part = torch.load(path, map_location="cpu", weights_only=False)
64
+ if hasattr(part, "partptr") and hasattr(part, "node_perm"):
65
+ partptr = np.asarray(part.partptr)
66
+ node_perm = np.asarray(part.node_perm)
67
+ self.part_nodes = [node_perm[partptr[i]:partptr[i+1]] for i in range(num_parts)]
68
+ self.cluster = np.zeros(data.num_nodes, dtype=np.int64)
69
+ for i in range(num_parts):
70
+ self.cluster[self.part_nodes[i]] = i
71
+ loaded = True
72
+ except Exception:
73
+ pass
74
+
75
+ if not loaded:
76
+ self.cluster = partition_graph(data.edge_index, data.num_nodes, num_parts)
77
+ from k3_node.loader.utils import to_numpy
78
+ cluster_np = to_numpy(self.cluster)
79
+ sort_idx = np.argsort(cluster_np)
80
+ sorted_cluster = cluster_np[sort_idx]
81
+ split_idx = np.searchsorted(sorted_cluster, np.arange(num_parts + 1))
82
+ self.part_nodes = [sort_idx[split_idx[i]:split_idx[i+1]] for i in range(num_parts)]
83
+
84
+ def __len__(self) -> int:
85
+ return self.num_parts
86
+
87
+ def __getitem__(self, idx: int) -> Data:
88
+ nodes = self.part_nodes[idx]
89
+ return self.data.subgraph(nodes)
90
+
91
+ def __repr__(self) -> str:
92
+ return f'{self.__class__.__name__}({self.num_parts})'
93
+
94
+
95
+ class ClusterLoader(*loader_bases(BaseDataLoader)):
96
+ r"""The data loader scheme from Cluster-GCN which merges partitioned
97
+ subgraphs to form a mini-batch.
98
+
99
+ Args:
100
+ cluster_data (ClusterData): The already partitioned data object.
101
+ **kwargs (optional): Additional arguments of :class:`torch.utils.data.DataLoader`.
102
+ """
103
+ def __init__(self, cluster_data: ClusterData, **kwargs):
104
+ self.cluster_data = cluster_data
105
+ kwargs.pop('collate_fn', None)
106
+ iterator = range(len(cluster_data))
107
+
108
+ if torch is not None:
109
+ super().__init__(iterator, collate_fn=self._collate, **kwargs)
110
+ else:
111
+ self.dataset = iterator
112
+ self.collate_fn = self._collate
113
+
114
+ def _collate(self, batch: List[int]) -> Data:
115
+ all_nodes = []
116
+ is_torch = torch is not None and isinstance(self.cluster_data.cluster, Tensor)
117
+
118
+ for part_id in batch:
119
+ all_nodes.append(self.cluster_data.part_nodes[part_id])
120
+
121
+ if is_torch and all_nodes and isinstance(all_nodes[0], Tensor):
122
+ nodes = torch.cat(all_nodes, dim=0)
123
+ else:
124
+ nodes = np.concatenate([np.asarray(x) for x in all_nodes], axis=0)
125
+
126
+ return self.cluster_data.data.subgraph(nodes)
127
+
@@ -0,0 +1,45 @@
1
+ from typing import List, Union
2
+
3
+ try:
4
+ import torch
5
+ import torch.utils.data
6
+ BaseDataLoader = torch.utils.data.DataLoader
7
+ except ImportError:
8
+ torch = None
9
+ BaseDataLoader = object
10
+
11
+ from k3_node.data import Dataset
12
+ from k3_node.data.data import BaseData
13
+
14
+
15
+ def collate_fn(data_list):
16
+ return data_list
17
+
18
+
19
+ class DataListLoader(BaseDataLoader):
20
+ r"""A data loader which batches data objects from a
21
+ :class:`k3_node.data.Dataset` to a Python list without batch collation.
22
+ """
23
+ def __init__(
24
+ self,
25
+ dataset: Union[Dataset, List[BaseData]],
26
+ batch_size: int = 1,
27
+ shuffle: bool = False,
28
+ **kwargs,
29
+ ):
30
+ kwargs.pop('collate_fn', None)
31
+
32
+ if torch is not None:
33
+ super().__init__(
34
+ dataset,
35
+ batch_size=batch_size,
36
+ shuffle=shuffle,
37
+ collate_fn=collate_fn,
38
+ **kwargs,
39
+ )
40
+ else:
41
+ self.dataset = dataset
42
+ self.batch_size = batch_size
43
+ self.shuffle = shuffle
44
+ self.collate_fn = collate_fn
45
+
@@ -0,0 +1,117 @@
1
+ from collections.abc import Mapping, Sequence
2
+ from typing import Any, List, Optional, Union
3
+
4
+ import numpy as np
5
+
6
+ try:
7
+ import torch
8
+ import torch.utils.data
9
+ from torch.utils.data.dataloader import default_collate
10
+ except ImportError:
11
+ torch = None
12
+ default_collate = None
13
+
14
+ from k3_node.data import Batch
15
+ from k3_node.data.data import BaseData
16
+ from k3_node.data.dataset import Dataset
17
+ from k3_node.loader.keras_dataset import loader_bases
18
+
19
+
20
+ class Collater:
21
+ r"""Collates a list of graph data objects or primitives into a mini-batch."""
22
+ def __init__(
23
+ self,
24
+ dataset: Optional[Union[Dataset, Sequence[BaseData]]] = None,
25
+ follow_batch: Optional[List[str]] = None,
26
+ exclude_keys: Optional[List[str]] = None,
27
+ ):
28
+ self.dataset = dataset
29
+ self.follow_batch = follow_batch
30
+ self.exclude_keys = exclude_keys
31
+
32
+ def __call__(self, batch: List[Any]) -> Any:
33
+ elem = batch[0]
34
+ if isinstance(elem, BaseData):
35
+ return Batch.from_data_list(
36
+ batch,
37
+ follow_batch=self.follow_batch,
38
+ exclude_keys=self.exclude_keys,
39
+ )
40
+ elif torch is not None and isinstance(elem, torch.Tensor):
41
+ return default_collate(batch)
42
+ elif isinstance(elem, np.ndarray):
43
+ return np.stack(batch, axis=0)
44
+ elif hasattr(elem, '__array__') and not isinstance(elem, (str, bytes)):
45
+ import keras
46
+ np_batch = np.stack([np.asarray(x) for x in batch], axis=0)
47
+ return keras.ops.convert_to_tensor(np_batch)
48
+ elif isinstance(elem, float):
49
+ if torch is not None:
50
+ return torch.tensor(batch, dtype=torch.float)
51
+ return np.array(batch, dtype=np.float32)
52
+ elif isinstance(elem, int):
53
+ if torch is not None:
54
+ return torch.tensor(batch, dtype=torch.long)
55
+ return np.array(batch, dtype=np.int64)
56
+ elif isinstance(elem, str):
57
+ return batch
58
+ elif isinstance(elem, Mapping):
59
+ return {key: self([data[key] for data in batch]) for key in elem}
60
+ elif isinstance(elem, tuple) and hasattr(elem, '_fields'):
61
+ return type(elem)(*(self(s) for s in zip(*batch)))
62
+ elif isinstance(elem, Sequence) and not isinstance(elem, str):
63
+ return [self(s) for s in zip(*batch)]
64
+
65
+ raise TypeError(f"DataLoader found invalid type: '{type(elem)}'")
66
+
67
+
68
+ BaseDataLoader = torch.utils.data.DataLoader if torch is not None else object
69
+
70
+
71
+ class DataLoader(*loader_bases(BaseDataLoader)):
72
+ r"""A data loader which merges data objects from a
73
+ :class:`k3_node.data.Dataset` to a mini-batch.
74
+ Data objects can be either of type :class:`~k3_node.data.Data` or
75
+ :class:`~k3_node.data.HeteroData`.
76
+
77
+ Args:
78
+ dataset (Dataset): The dataset from which to load the data.
79
+ batch_size (int, optional): How many samples per batch to load.
80
+ (default: :obj:`1`)
81
+ shuffle (bool, optional): If set to :obj:`True`, the data will be
82
+ reshuffled at every epoch. (default: :obj:`False`)
83
+ follow_batch (List[str], optional): Creates assignment batch
84
+ vectors for each key in the list. (default: :obj:`None`)
85
+ exclude_keys (List[str], optional): Will exclude each key in the
86
+ list. (default: :obj:`None`)
87
+ **kwargs (optional): Additional arguments of
88
+ :class:`torch.utils.data.DataLoader`.
89
+ """
90
+ def __init__(
91
+ self,
92
+ dataset: Union[Dataset, Sequence[BaseData]],
93
+ batch_size: int = 1,
94
+ shuffle: bool = False,
95
+ follow_batch: Optional[List[str]] = None,
96
+ exclude_keys: Optional[List[str]] = None,
97
+ **kwargs,
98
+ ):
99
+ kwargs.pop('collate_fn', None)
100
+
101
+ self.follow_batch = follow_batch
102
+ self.exclude_keys = exclude_keys
103
+
104
+ if torch is not None:
105
+ super().__init__(
106
+ dataset,
107
+ batch_size=batch_size,
108
+ shuffle=shuffle,
109
+ collate_fn=Collater(dataset, follow_batch, exclude_keys),
110
+ **kwargs,
111
+ )
112
+ else:
113
+ self.dataset = dataset
114
+ self.batch_size = batch_size
115
+ self.shuffle = shuffle
116
+ self.collate_fn = Collater(dataset, follow_batch, exclude_keys)
117
+
@@ -0,0 +1,62 @@
1
+ from typing import List, Union
2
+
3
+ import numpy as np
4
+
5
+ try:
6
+ import torch
7
+ import torch.utils.data
8
+ from torch.utils.data.dataloader import default_collate
9
+ BaseDataLoader = torch.utils.data.DataLoader
10
+ except ImportError:
11
+ torch = None
12
+ default_collate = None
13
+ BaseDataLoader = object
14
+
15
+ from k3_node.data import Batch, Data, Dataset
16
+ from k3_node.loader.keras_dataset import loader_bases
17
+
18
+
19
+ def collate_fn(data_list: List[Data]) -> Batch:
20
+ batch = Batch()
21
+ for key in data_list[0].keys():
22
+ vals = [data[key] for data in data_list]
23
+ if default_collate is not None and isinstance(vals[0], torch.Tensor):
24
+ batch[key] = default_collate(vals)
25
+ elif isinstance(vals[0], np.ndarray):
26
+ batch[key] = np.stack(vals, axis=0)
27
+ elif hasattr(vals[0], '__array__'):
28
+ import keras
29
+ batch[key] = keras.ops.stack(vals, axis=0)
30
+ else:
31
+ batch[key] = vals
32
+ return batch
33
+
34
+
35
+ class DenseDataLoader(*loader_bases(BaseDataLoader)):
36
+ r"""A data loader which batches data objects from a
37
+ :class:`k3_node.data.Dataset` to a :class:`k3_node.data.Batch`
38
+ object by stacking all attributes in a new dimension.
39
+ """
40
+ def __init__(
41
+ self,
42
+ dataset: Union[Dataset, List[Data]],
43
+ batch_size: int = 1,
44
+ shuffle: bool = False,
45
+ **kwargs,
46
+ ):
47
+ kwargs.pop('collate_fn', None)
48
+
49
+ if torch is not None:
50
+ super().__init__(
51
+ dataset,
52
+ batch_size=batch_size,
53
+ shuffle=shuffle,
54
+ collate_fn=collate_fn,
55
+ **kwargs,
56
+ )
57
+ else:
58
+ self.dataset = dataset
59
+ self.batch_size = batch_size
60
+ self.shuffle = shuffle
61
+ self.collate_fn = collate_fn
62
+
@@ -0,0 +1,93 @@
1
+ from typing import Iterator, List, Optional
2
+
3
+ import numpy as np
4
+
5
+ try:
6
+ import torch
7
+ import torch.utils.data.sampler
8
+ BaseSampler = torch.utils.data.sampler.Sampler
9
+ except ImportError:
10
+ torch = None
11
+ BaseSampler = object
12
+
13
+ from k3_node.data import Dataset
14
+
15
+
16
+ class DynamicBatchSampler(BaseSampler):
17
+ r"""Dynamically adds samples to a mini-batch up to a maximum size (either
18
+ based on number of nodes or number of edges).
19
+
20
+ Args:
21
+ dataset (Dataset): Dataset to sample from.
22
+ max_num (int): Size of mini-batch to aim for in number of nodes or edges.
23
+ mode (str, optional): :obj:`"node"` or :obj:`"edge"` to measure batch size. (default: :obj:`"node"`)
24
+ shuffle (bool, optional): If set to :obj:`True`, will have the data reshuffled at every epoch. (default: :obj:`False`)
25
+ skip_too_big (bool, optional): If set to :obj:`True`, skip samples which cannot fit in a batch by itself. (default: :obj:`False`)
26
+ num_steps (int, optional): The number of mini-batches to draw for a single epoch. (default: :obj:`None`)
27
+ """
28
+ def __init__(
29
+ self,
30
+ dataset: Dataset,
31
+ max_num: int,
32
+ mode: str = 'node',
33
+ shuffle: bool = False,
34
+ skip_too_big: bool = False,
35
+ num_steps: Optional[int] = None,
36
+ ):
37
+ if max_num <= 0:
38
+ raise ValueError(f"`max_num` should be a positive integer value (got {max_num})")
39
+ if mode not in ['node', 'edge']:
40
+ raise ValueError(f"`mode` choice should be either 'node' or 'edge' (got '{mode}')")
41
+
42
+ self.dataset = dataset
43
+ self.max_num = max_num
44
+ self.mode = mode
45
+ self.shuffle = shuffle
46
+ self.skip_too_big = skip_too_big
47
+ self.num_steps = num_steps
48
+ self.max_steps = num_steps or len(dataset)
49
+
50
+ def __iter__(self) -> Iterator[List[int]]:
51
+ if self.shuffle:
52
+ if torch is not None:
53
+ indices = torch.randperm(len(self.dataset)).tolist()
54
+ else:
55
+ indices = np.random.permutation(len(self.dataset)).tolist()
56
+ else:
57
+ indices = list(range(len(self.dataset)))
58
+
59
+ samples: List[int] = []
60
+ current_num: int = 0
61
+ num_steps: int = 0
62
+ num_processed: int = 0
63
+
64
+ while num_processed < len(self.dataset) and num_steps < self.max_steps:
65
+ for i in indices[num_processed:]:
66
+ data = self.dataset[i]
67
+ num = data.num_nodes if self.mode == 'node' else data.num_edges
68
+
69
+ if current_num + num > self.max_num:
70
+ if current_num == 0:
71
+ if self.skip_too_big:
72
+ num_processed += 1
73
+ continue
74
+ else:
75
+ break
76
+
77
+ samples.append(i)
78
+ num_processed += 1
79
+ current_num += num
80
+
81
+ yield samples
82
+ samples = []
83
+ current_num = 0
84
+ num_steps += 1
85
+
86
+ def __len__(self) -> int:
87
+ if self.num_steps is None:
88
+ raise ValueError(
89
+ f"The length of '{self.__class__.__name__}' is undefined since the number of steps per epoch "
90
+ f"is ambiguous. Either specify `num_steps` or use a static batch sampler."
91
+ )
92
+ return self.num_steps
93
+