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,182 @@
1
+ import copy
2
+ import os
3
+ import os.path as osp
4
+ from typing import Any, Callable, Iterator, List, Optional, Sequence, Tuple, Union
5
+
6
+ import numpy as np
7
+ from keras import ops
8
+
9
+ from k3_node.data.data import BaseData
10
+ from k3_node.data.storage import is_tensor_like, to_numpy
11
+
12
+
13
+ class Dataset:
14
+ """Dataset base class for creating graph datasets."""
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
+ self.root = root
26
+ self.transform = transform
27
+ self.pre_transform = pre_transform
28
+ self.pre_filter = pre_filter
29
+ self.log = log
30
+ self.force_reload = force_reload
31
+ self._indices: Optional[Sequence] = None
32
+
33
+ if self.has_download:
34
+ self._download()
35
+
36
+ if self.has_process:
37
+ self._process()
38
+
39
+ @property
40
+ def raw_file_names(self) -> Union[str, List[str], Tuple[str, ...]]:
41
+ return []
42
+
43
+ @property
44
+ def processed_file_names(self) -> Union[str, List[str], Tuple[str, ...]]:
45
+ return []
46
+
47
+ @property
48
+ def raw_dir(self) -> str:
49
+ return osp.join(self.root or "", "raw")
50
+
51
+ @property
52
+ def processed_dir(self) -> str:
53
+ return osp.join(self.root or "", "processed")
54
+
55
+ @property
56
+ def raw_paths(self) -> List[str]:
57
+ files = self.raw_file_names
58
+ if isinstance(files, str):
59
+ files = [files]
60
+ return [osp.join(self.raw_dir, f) for f in files]
61
+
62
+ @property
63
+ def processed_paths(self) -> List[str]:
64
+ files = self.processed_file_names
65
+ if isinstance(files, str):
66
+ files = [files]
67
+ return [osp.join(self.processed_dir, f) for f in files]
68
+
69
+ @property
70
+ def has_download(self) -> bool:
71
+ return any("download" in cls.__dict__ for cls in self.__class__.__mro__ if cls is not Dataset)
72
+
73
+ @property
74
+ def has_process(self) -> bool:
75
+ return any("process" in cls.__dict__ for cls in self.__class__.__mro__ if cls is not Dataset)
76
+
77
+ def download(self):
78
+ pass
79
+
80
+ def process(self):
81
+ pass
82
+
83
+ def _download(self):
84
+ if all(osp.exists(p) for p in self.raw_paths if p):
85
+ return
86
+ os.makedirs(self.raw_dir, exist_ok=True)
87
+ self.download()
88
+
89
+ def _process(self):
90
+ if not self.force_reload and len(self.processed_paths) > 0 and all(osp.exists(p) for p in self.processed_paths):
91
+ return
92
+ os.makedirs(self.processed_dir, exist_ok=True)
93
+ self.process()
94
+
95
+ def len(self) -> int:
96
+ raise NotImplementedError
97
+
98
+ def get(self, idx: int) -> BaseData:
99
+ raise NotImplementedError
100
+
101
+ def indices(self) -> Sequence:
102
+ return range(self.len()) if self._indices is None else self._indices
103
+
104
+ def __len__(self) -> int:
105
+ return len(self.indices())
106
+
107
+ def __getitem__(self, idx: Any) -> Any:
108
+ if isinstance(idx, (int, np.integer)):
109
+ data = self.get(self.indices()[idx])
110
+ if hasattr(data, "to_backend"):
111
+ data = data.to_backend()
112
+ data = data if self.transform is None else self.transform(data)
113
+ return data
114
+ else:
115
+ return self.index_select(idx)
116
+
117
+ def index_select(self, idx: Any) -> "Dataset":
118
+ indices = list(self.indices())
119
+ if isinstance(idx, slice):
120
+ # Allow fractional slicing as in PyG, e.g. dataset[:0.9] for the first 90%
121
+ start, stop = idx.start, idx.stop
122
+ start = round(start * len(indices)) if isinstance(start, float) else start
123
+ stop = round(stop * len(indices)) if isinstance(stop, float) else stop
124
+ indices = indices[slice(start, stop, idx.step)]
125
+ elif is_tensor_like(idx):
126
+ idx_np = to_numpy(idx)
127
+ if idx_np.dtype == bool:
128
+ indices = [indices[i] for i in np.where(idx_np)[0]]
129
+ else:
130
+ indices = [indices[int(i)] for i in idx_np]
131
+ elif isinstance(idx, Sequence):
132
+ indices = [indices[i] for i in idx]
133
+ else:
134
+ raise IndexError(f"Invalid index type {type(idx)}")
135
+
136
+ dataset = copy.copy(self)
137
+ dataset._indices = indices
138
+ return dataset
139
+
140
+ def shuffle(self, return_perm: bool = False):
141
+ r"""Randomly shuffles the examples in the dataset (as in PyG).
142
+
143
+ Args:
144
+ return_perm (bool): If :obj:`True`, also returns the permutation used.
145
+ """
146
+ perm = np.random.permutation(len(self))
147
+ dataset = self.index_select(perm.tolist())
148
+ return (dataset, perm) if return_perm else dataset
149
+
150
+ def __iter__(self) -> Iterator[BaseData]:
151
+ for i in range(len(self)):
152
+ yield self[i]
153
+
154
+ @property
155
+ def num_node_features(self) -> int:
156
+ data = self[0]
157
+ return getattr(data, "num_node_features", 0)
158
+
159
+ @property
160
+ def num_features(self) -> int:
161
+ return self.num_node_features
162
+
163
+ @property
164
+ def num_edge_features(self) -> int:
165
+ data = self[0]
166
+ return getattr(data, "num_edge_features", 0)
167
+
168
+ @property
169
+ def num_classes(self) -> int:
170
+ y_list = [d.y for d in self if hasattr(d, "y") and d.y is not None]
171
+ if len(y_list) == 0:
172
+ return 0
173
+ y_np_list = [to_numpy(y) for y in y_list]
174
+ y = np.concatenate(y_np_list, axis=0) if len(y_np_list) > 1 else y_np_list[0]
175
+ if np.issubdtype(y.dtype, np.integer):
176
+ return int(np.max(y)) + 1
177
+ return len(np.unique(y))
178
+
179
+ def __repr__(self) -> str:
180
+ arg_repr = str(len(self)) if len(self) > 1 else ""
181
+ return f"{self.__class__.__name__}({arg_repr})"
182
+
@@ -0,0 +1,49 @@
1
+ import os
2
+ import os.path as osp
3
+ import ssl
4
+ import sys
5
+ import urllib.request
6
+ from typing import Optional
7
+
8
+
9
+ def download_url(
10
+ url: str,
11
+ folder: str,
12
+ log: bool = True,
13
+ filename: Optional[str] = None,
14
+ ) -> str:
15
+ """Downloads the content of a URL to a specific folder."""
16
+ if filename is None:
17
+ filename = url.rpartition("/")[2]
18
+ filename = filename if filename[0] == "?" else filename.split("?")[0]
19
+
20
+ path = osp.join(folder, filename)
21
+ if osp.exists(path):
22
+ return path
23
+
24
+ if log and "PYTEST_CURRENT_TEST" not in os.environ:
25
+ print(f"Downloading {url}", file=sys.stderr)
26
+
27
+ os.makedirs(folder, exist_ok=True)
28
+ context = ssl._create_unverified_context()
29
+ with urllib.request.urlopen(url, context=context) as response:
30
+ with open(path, "wb") as f:
31
+ while True:
32
+ chunk = response.read(10 * 1024 * 1024)
33
+ if not chunk:
34
+ break
35
+ f.write(chunk)
36
+
37
+ return path
38
+
39
+
40
+ def download_google_url(
41
+ id: str,
42
+ folder: str,
43
+ filename: str,
44
+ log: bool = True,
45
+ ) -> str:
46
+ """Downloads the content of a Google Drive ID to a specific folder."""
47
+ url = f"https://drive.usercontent.google.com/download?id={id}&confirm=t"
48
+ return download_url(url, folder, log, filename)
49
+
@@ -0,0 +1,45 @@
1
+ import bz2
2
+ import gzip
3
+ import os
4
+ import os.path as osp
5
+ import sys
6
+ import tarfile
7
+ import zipfile
8
+
9
+
10
+ def maybe_log(path: str, log: bool = True):
11
+ if log and "PYTEST_CURRENT_TEST" not in os.environ:
12
+ print(f"Extracting {path}", file=sys.stderr)
13
+
14
+
15
+ def extract_tar(path: str, folder: str, mode: str = "r:gz", log: bool = True):
16
+ """Extracts a tar archive to a specific folder."""
17
+ maybe_log(path, log)
18
+ with tarfile.open(path, mode) as f:
19
+ f.extractall(folder)
20
+
21
+
22
+ def extract_zip(path: str, folder: str, log: bool = True):
23
+ """Extracts a zip archive to a specific folder."""
24
+ maybe_log(path, log)
25
+ with zipfile.ZipFile(path, "r") as f:
26
+ f.extractall(folder)
27
+
28
+
29
+ def extract_bz2(path: str, folder: str, log: bool = True):
30
+ """Extracts a bz2 archive to a specific folder."""
31
+ maybe_log(path, log)
32
+ path = osp.abspath(path)
33
+ with bz2.open(path, "r") as r:
34
+ with open(osp.join(folder, ".".join(path.split(".")[:-1])), "wb") as w:
35
+ w.write(r.read())
36
+
37
+
38
+ def extract_gz(path: str, folder: str, log: bool = True):
39
+ """Extracts a gz archive to a specific folder."""
40
+ maybe_log(path, log)
41
+ path = osp.abspath(path)
42
+ with gzip.open(path, "r") as r:
43
+ with open(osp.join(folder, ".".join(path.split(".")[:-1])), "wb") as w:
44
+ w.write(r.read())
45
+
@@ -0,0 +1,70 @@
1
+ from abc import ABC, abstractmethod
2
+ from dataclasses import dataclass
3
+ from enum import Enum
4
+ from typing import Any, Dict, List, Optional, Tuple, Union
5
+
6
+
7
+ class _FieldStatus(Enum):
8
+ UNSET = None
9
+
10
+
11
+ @dataclass
12
+ class TensorAttr:
13
+ """Defines the attributes of a FeatureStore tensor."""
14
+
15
+ group_name: Optional[Any] = _FieldStatus.UNSET
16
+ attr_name: Optional[str] = _FieldStatus.UNSET
17
+ index: Optional[Any] = _FieldStatus.UNSET
18
+
19
+ def is_set(self, key: str) -> bool:
20
+ return getattr(self, key) != _FieldStatus.UNSET
21
+
22
+ def is_fully_specified(self) -> bool:
23
+ return all(self.is_set(k) for k in ("group_name", "attr_name", "index"))
24
+
25
+
26
+ class FeatureStore(ABC):
27
+ """Abstract base class for feature stores."""
28
+
29
+ def __init__(self):
30
+ self._feat_dict: Dict[Tuple[Any, str], Any] = {}
31
+
32
+ @abstractmethod
33
+ def _put_tensor(self, tensor: Any, attr: TensorAttr) -> bool:
34
+ pass
35
+
36
+ @abstractmethod
37
+ def _get_tensor(self, attr: TensorAttr) -> Optional[Any]:
38
+ pass
39
+
40
+ @abstractmethod
41
+ def _remove_tensor(self, attr: TensorAttr) -> bool:
42
+ pass
43
+
44
+ def put_tensor(self, tensor: Any, group_name: Any = None, attr_name: Optional[str] = None, index: Any = None) -> bool:
45
+ attr = TensorAttr(group_name=group_name, attr_name=attr_name, index=index)
46
+ return self._put_tensor(tensor, attr)
47
+
48
+ def get_tensor(self, group_name: Any = None, attr_name: Optional[str] = None, index: Any = None) -> Optional[Any]:
49
+ attr = TensorAttr(group_name=group_name, attr_name=attr_name, index=index)
50
+ return self._get_tensor(attr)
51
+
52
+ def remove_tensor(self, group_name: Any = None, attr_name: Optional[str] = None, index: Any = None) -> bool:
53
+ attr = TensorAttr(group_name=group_name, attr_name=attr_name, index=index)
54
+ return self._remove_tensor(attr)
55
+
56
+ def __getitem__(self, key: Any) -> Any:
57
+ if isinstance(key, tuple):
58
+ group_name, attr_name = key[:2]
59
+ index = key[2] if len(key) > 2 else None
60
+ return self.get_tensor(group_name=group_name, attr_name=attr_name, index=index)
61
+ return self.get_tensor(group_name=key)
62
+
63
+ def __setitem__(self, key: Any, value: Any):
64
+ if isinstance(key, tuple):
65
+ group_name, attr_name = key[:2]
66
+ index = key[2] if len(key) > 2 else None
67
+ self.put_tensor(value, group_name=group_name, attr_name=attr_name, index=index)
68
+ else:
69
+ self.put_tensor(value, group_name=key)
70
+
@@ -0,0 +1,92 @@
1
+ from abc import ABC, abstractmethod
2
+ from dataclasses import dataclass
3
+ from enum import Enum
4
+ from typing import Any, Dict, List, Optional, Tuple, Union
5
+
6
+
7
+ class EdgeLayout(Enum):
8
+ COO = "coo"
9
+ CSC = "csc"
10
+ CSR = "csr"
11
+
12
+
13
+ @dataclass
14
+ class EdgeAttr:
15
+ """Defines the attributes of a GraphStore edge."""
16
+
17
+ edge_type: Any
18
+ layout: EdgeLayout
19
+ is_sorted: bool = False
20
+ size: Optional[Tuple[int, int]] = None
21
+
22
+ def __init__(
23
+ self,
24
+ edge_type: Any,
25
+ layout: Union[EdgeLayout, str] = EdgeLayout.COO,
26
+ is_sorted: bool = False,
27
+ size: Optional[Tuple[int, int]] = None,
28
+ ):
29
+ if isinstance(layout, str):
30
+ layout = EdgeLayout(layout.lower())
31
+ self.edge_type = edge_type
32
+ self.layout = layout
33
+ self.is_sorted = is_sorted
34
+ self.size = size
35
+
36
+
37
+ class GraphStore(ABC):
38
+ """Abstract base class for graph edge stores."""
39
+
40
+ @abstractmethod
41
+ def _put_edge_index(self, edge_index: Any, edge_attr: EdgeAttr) -> bool:
42
+ pass
43
+
44
+ @abstractmethod
45
+ def _get_edge_index(self, edge_attr: EdgeAttr) -> Optional[Any]:
46
+ pass
47
+
48
+ @abstractmethod
49
+ def _remove_edge_index(self, edge_attr: EdgeAttr) -> bool:
50
+ pass
51
+
52
+ def put_edge_index(
53
+ self,
54
+ edge_index: Any,
55
+ edge_type: Any,
56
+ layout: Union[EdgeLayout, str] = EdgeLayout.COO,
57
+ is_sorted: bool = False,
58
+ size: Optional[Tuple[int, int]] = None,
59
+ ) -> bool:
60
+ attr = EdgeAttr(edge_type=edge_type, layout=layout, is_sorted=is_sorted, size=size)
61
+ return self._put_edge_index(edge_index, attr)
62
+
63
+ def get_edge_index(
64
+ self,
65
+ edge_type: Any,
66
+ layout: Union[EdgeLayout, str] = EdgeLayout.COO,
67
+ is_sorted: bool = False,
68
+ ) -> Optional[Any]:
69
+ attr = EdgeAttr(edge_type=edge_type, layout=layout, is_sorted=is_sorted)
70
+ return self._get_edge_index(attr)
71
+
72
+ def remove_edge_index(
73
+ self,
74
+ edge_type: Any,
75
+ layout: Union[EdgeLayout, str] = EdgeLayout.COO,
76
+ ) -> bool:
77
+ attr = EdgeAttr(edge_type=edge_type, layout=layout)
78
+ return self._remove_edge_index(attr)
79
+
80
+ def __getitem__(self, key: Any) -> Any:
81
+ if isinstance(key, tuple) and len(key) == 2:
82
+ edge_type, layout = key
83
+ return self.get_edge_index(edge_type=edge_type, layout=layout)
84
+ return self.get_edge_index(edge_type=key)
85
+
86
+ def __setitem__(self, key: Any, value: Any):
87
+ if isinstance(key, tuple) and len(key) == 2:
88
+ edge_type, layout = key
89
+ self.put_edge_index(value, edge_type=edge_type, layout=layout)
90
+ else:
91
+ self.put_edge_index(value, edge_type=key)
92
+