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,121 @@
1
+ import os
2
+ import os.path as osp
3
+ from typing import Callable, List, Optional
4
+
5
+ import numpy as np
6
+ from keras import ops
7
+
8
+ from k3_node.data import (
9
+ Data,
10
+ InMemoryDataset,
11
+ download_url,
12
+ extract_zip,
13
+ )
14
+ from k3_node.utils import coalesce
15
+
16
+
17
+ class Reddit(InMemoryDataset):
18
+ r"""The Reddit dataset from the `"Inductive Representation Learning on
19
+ Large Graphs" <https://arxiv.org/abs/1706.02216>`_ paper, containing
20
+ Reddit posts belonging to different communities.
21
+
22
+ Args:
23
+ root (str): Root directory where the dataset should be saved.
24
+ transform (callable, optional): A function/transform that takes in an
25
+ :obj:`k3_node.data.Data` object and returns a transformed
26
+ version. The data object will be transformed before every access.
27
+ (default: :obj:`None`)
28
+ pre_transform (callable, optional): A function/transform that takes in
29
+ an :obj:`k3_node.data.Data` object and returns a
30
+ transformed version. The data object will be transformed before
31
+ being saved to disk. (default: :obj:`None`)
32
+ pre_filter (callable, optional): A function that takes in an
33
+ :obj:`k3_node.data.Data` object and returns a boolean
34
+ value, indicating whether the data object should be included in the
35
+ final dataset. (default: :obj:`None`)
36
+ force_reload (bool, optional): Whether to re-process the dataset.
37
+ (default: :obj:`False`)
38
+
39
+ **STATS:**
40
+
41
+ .. list-table::
42
+ :widths: 10 10 10 10
43
+ :header-rows: 1
44
+
45
+ * - #nodes
46
+ - #edges
47
+ - #features
48
+ - #classes
49
+ * - 232,965
50
+ - 114,615,892
51
+ - 602
52
+ - 41
53
+ """
54
+
55
+ url = "https://data.dgl.ai/dataset/reddit.zip"
56
+
57
+ def __init__(
58
+ self,
59
+ root: str,
60
+ transform: Optional[Callable] = None,
61
+ pre_transform: Optional[Callable] = None,
62
+ pre_filter: Optional[Callable] = None,
63
+ force_reload: bool = False,
64
+ ) -> None:
65
+ super().__init__(
66
+ root,
67
+ transform,
68
+ pre_transform,
69
+ pre_filter,
70
+ force_reload=force_reload,
71
+ )
72
+ self.load(self.processed_paths[0])
73
+
74
+ @property
75
+ def raw_file_names(self) -> List[str]:
76
+ return ["reddit_data.npz", "reddit_graph.npz"]
77
+
78
+ @property
79
+ def processed_file_names(self) -> str:
80
+ return "data.pt"
81
+
82
+ @property
83
+ def num_classes(self) -> int:
84
+ return 41
85
+
86
+ def download(self) -> None:
87
+ path = download_url(self.url, self.raw_dir)
88
+ extract_zip(path, self.raw_dir)
89
+ if osp.exists(path):
90
+ os.unlink(path)
91
+
92
+ def process(self) -> None:
93
+ import scipy.sparse as sp
94
+
95
+ data = np.load(osp.join(self.raw_dir, "reddit_data.npz"))
96
+ x = ops.convert_to_tensor(data["feature"], dtype="float32")
97
+ y = ops.convert_to_tensor(data["label"], dtype="int64")
98
+ split = data["node_types"]
99
+
100
+ adj = sp.load_npz(osp.join(self.raw_dir, "reddit_graph.npz"))
101
+ if not hasattr(adj, "row"):
102
+ adj = adj.tocoo()
103
+ row = adj.row.astype(np.int64)
104
+ col = adj.col.astype(np.int64)
105
+ edge_index = np.stack([row, col], axis=0)
106
+ edge_index = ops.convert_to_tensor(edge_index, dtype="int64")
107
+ edge_index, _ = coalesce(edge_index, num_nodes=int(data["feature"].shape[0]))
108
+
109
+ data = Data(x=x, edge_index=edge_index, y=y)
110
+ data.train_mask = ops.convert_to_tensor(split == 1, dtype="bool")
111
+ data.val_mask = ops.convert_to_tensor(split == 2, dtype="bool")
112
+ data.test_mask = ops.convert_to_tensor(split == 3, dtype="bool")
113
+
114
+ if self.pre_filter is not None and not self.pre_filter(data):
115
+ return
116
+
117
+ if self.pre_transform is not None:
118
+ data = self.pre_transform(data)
119
+
120
+ self.save([data], self.processed_paths[0])
121
+
@@ -0,0 +1,165 @@
1
+ import os.path as osp
2
+ from typing import Any, Callable, List, Optional, Union
3
+ import numpy as np
4
+ from keras import ops
5
+
6
+ from k3_node.data import Data, InMemoryDataset
7
+ from k3_node.utils.random import stochastic_blockmodel_graph
8
+
9
+
10
+ class StochasticBlockModelDataset(InMemoryDataset):
11
+ r"""A synthetic graph dataset generated by the stochastic block model.
12
+
13
+ Args:
14
+ root (str): Root directory where the dataset should be saved.
15
+ block_sizes ([int] or array): The sizes of blocks.
16
+ edge_probs ([[float]] or array): The density of edges between blocks.
17
+ num_graphs (int, optional): The number of graphs. (default: 1)
18
+ num_channels (int, optional): The number of node features. (default: None)
19
+ is_undirected (bool, optional): Whether the graph is undirected. (default: True)
20
+ transform (callable, optional): Transform function.
21
+ pre_transform (callable, optional): Pre-transform function.
22
+ force_reload (bool, optional): Whether to re-process the dataset.
23
+ """
24
+
25
+ def __init__(
26
+ self,
27
+ root: str,
28
+ block_sizes: Union[List[int], np.ndarray],
29
+ edge_probs: Union[List[List[float]], np.ndarray],
30
+ num_graphs: int = 1,
31
+ num_channels: Optional[int] = None,
32
+ is_undirected: bool = True,
33
+ transform: Optional[Callable] = None,
34
+ pre_transform: Optional[Callable] = None,
35
+ force_reload: bool = False,
36
+ **kwargs: Any,
37
+ ):
38
+ self.block_sizes = np.array(block_sizes, dtype=np.int64)
39
+ self.edge_probs = np.array(edge_probs, dtype=np.float32)
40
+ assert num_graphs > 0
41
+
42
+ self.num_graphs = num_graphs
43
+ self.num_channels = num_channels
44
+ self.is_undirected = is_undirected
45
+
46
+ self.kwargs = {
47
+ "n_informative": num_channels,
48
+ "n_redundant": 0,
49
+ "flip_y": 0.0,
50
+ "shuffle": False,
51
+ }
52
+ self.kwargs.update(kwargs)
53
+
54
+ super().__init__(root, transform, pre_transform, force_reload=force_reload)
55
+ self.load(self.processed_paths[0])
56
+
57
+ @property
58
+ def processed_dir(self) -> str:
59
+ return osp.join(self.root, self.__class__.__name__, "processed")
60
+
61
+ @property
62
+ def processed_file_names(self) -> str:
63
+ bs = self.block_sizes.flatten().tolist()
64
+ hash1 = "-".join([f"{x:.1f}" for x in bs])
65
+ ep = self.edge_probs.flatten().tolist()
66
+ hash2 = "-".join([f"{x:.1f}" for x in ep])
67
+ return f"data_{self.num_channels}_{hash1}_{hash2}_{self.num_graphs}.pt"
68
+
69
+ def process(self):
70
+ try:
71
+ from sklearn.datasets import make_classification
72
+ except ImportError:
73
+ make_classification = None
74
+
75
+ edge_index = stochastic_blockmodel_graph(
76
+ self.block_sizes, self.edge_probs, directed=not self.is_undirected
77
+ )
78
+
79
+ num_samples = int(self.block_sizes.sum())
80
+ num_classes = len(self.block_sizes)
81
+
82
+ data_list = []
83
+ for _ in range(self.num_graphs):
84
+ x = None
85
+ if self.num_channels is not None:
86
+ if make_classification is not None:
87
+ x_raw, y_not_sorted = make_classification(
88
+ n_samples=num_samples,
89
+ n_features=self.num_channels,
90
+ n_classes=num_classes,
91
+ weights=self.block_sizes / num_samples,
92
+ **self.kwargs,
93
+ )
94
+ x_raw = x_raw[np.argsort(y_not_sorted)]
95
+ x = ops.convert_to_tensor(x_raw.astype(np.float32), dtype="float32")
96
+ else:
97
+ x = ops.convert_to_tensor(
98
+ np.random.randn(num_samples, self.num_channels).astype(np.float32),
99
+ dtype="float32",
100
+ )
101
+
102
+ y_np = np.repeat(np.arange(num_classes, dtype=np.int64), self.block_sizes)
103
+ y = ops.convert_to_tensor(y_np, dtype="int64")
104
+
105
+ data = Data(x=x, edge_index=edge_index, y=y)
106
+ if self.pre_transform is not None:
107
+ data = self.pre_transform(data)
108
+ data_list.append(data)
109
+
110
+ self.save(data_list, self.processed_paths[0])
111
+
112
+
113
+ class RandomPartitionGraphDataset(StochasticBlockModelDataset):
114
+ r"""The random partition graph dataset from the "How to Find Your Friendly
115
+ Neighborhood: Graph Attention Design with Self-Supervision" paper.
116
+ """
117
+
118
+ def __init__(
119
+ self,
120
+ root: str,
121
+ num_classes: int,
122
+ num_nodes_per_class: int,
123
+ node_homophily_ratio: float,
124
+ average_degree: float,
125
+ num_graphs: int = 1,
126
+ num_channels: Optional[int] = None,
127
+ is_undirected: bool = True,
128
+ transform: Optional[Callable] = None,
129
+ pre_transform: Optional[Callable] = None,
130
+ **kwargs: Any,
131
+ ):
132
+ self._num_classes = num_classes
133
+ self.num_nodes_per_class = num_nodes_per_class
134
+ self.node_homophily_ratio = node_homophily_ratio
135
+ self.average_degree = average_degree
136
+
137
+ ec_over_v2 = average_degree / num_nodes_per_class
138
+ p_in = node_homophily_ratio * ec_over_v2
139
+ p_out = (ec_over_v2 - p_in) / (num_classes - 1)
140
+
141
+ block_sizes = [num_nodes_per_class for _ in range(num_classes)]
142
+ edge_probs = [[p_out for _ in range(num_classes)] for _ in range(num_classes)]
143
+ for r in range(num_classes):
144
+ edge_probs[r][r] = p_in
145
+
146
+ super().__init__(
147
+ root,
148
+ block_sizes,
149
+ edge_probs,
150
+ num_graphs=num_graphs,
151
+ num_channels=num_channels,
152
+ is_undirected=is_undirected,
153
+ transform=transform,
154
+ pre_transform=pre_transform,
155
+ **kwargs,
156
+ )
157
+
158
+ @property
159
+ def processed_file_names(self) -> str:
160
+ return (
161
+ f"data_{self.num_channels}_{self._num_classes}_"
162
+ f"{self.num_nodes_per_class}_{self.node_homophily_ratio:.1f}_"
163
+ f"{self.average_degree:.1f}_{self.num_graphs}.pt"
164
+ )
165
+
@@ -0,0 +1,74 @@
1
+ from typing import Optional
2
+
3
+ import numpy as np
4
+ from keras import ops
5
+
6
+ from k3_node.data import Data, InMemoryDataset
7
+
8
+
9
+ class SEALDataset(InMemoryDataset):
10
+ r"""Enclosing subgraphs for link prediction, as in SEAL
11
+ (`"Link Prediction Based on Graph Neural Networks" <https://arxiv.org/abs/1802.09691>`_).
12
+
13
+ Every link to classify becomes a small graph: the ``num_hops``-hop neighborhood of its two
14
+ nodes, without the link itself. Nodes are described only by their double-radius labels
15
+ (their distances to the two nodes, one-hot encoded as ``x``); ``y`` is 1 for a true link and
16
+ 0 for a non-edge. This turns link prediction into graph classification.
17
+
18
+ Args:
19
+ data (Data): One split from :class:`~k3_node.transforms.RandomLinkSplit`: message passing
20
+ edges in ``edge_index`` and the links in ``edge_label_index`` / ``edge_label`` (or
21
+ ``pos_edge_label_index`` / ``neg_edge_label_index``).
22
+ num_hops (int): Size of the neighborhood around each link. (default: ``2``)
23
+ num_labels (int, optional): Size of the one-hot node labels. Use the training set's
24
+ ``num_labels`` for the validation and test sets. (default: the largest label + 1)
25
+
26
+ Example:
27
+ ```python
28
+ import numpy as np
29
+ from k3_node.data import Data
30
+ from k3_node.datasets import SEALDataset
31
+ from k3_node.transforms import RandomLinkSplit
32
+
33
+ data = Data(edge_index=np.random.randint(0, 30, size=(2, 120)), num_nodes=30)
34
+ train_data, val_data, test_data = RandomLinkSplit(num_val=0.1, num_test=0.1)(data)
35
+ train_dataset = SEALDataset(train_data, num_hops=2)
36
+ print(len(train_dataset) == train_data.edge_label_index.shape[1]) # True: one graph per link
37
+ ```
38
+ """
39
+
40
+ def __init__(self, data, num_hops: int = 2, num_labels: Optional[int] = None):
41
+ super().__init__(None)
42
+ from k3_node.utils.graph import drnl_node_labeling, k_hop_subgraph
43
+
44
+ links, labels = self._links(data)
45
+ edge_index = np.asarray(ops.convert_to_numpy(data.edge_index)).astype(np.int64)
46
+ graphs = []
47
+ for (src, dst), y in zip(links.T, labels):
48
+ nodes, sub_edge_index, mapping, _ = k_hop_subgraph(
49
+ [src, dst], num_hops, edge_index, relabel_nodes=True, num_nodes=data.num_nodes)
50
+ s, d = (int(m) for m in mapping)
51
+ keep = ~(((sub_edge_index[0] == s) & (sub_edge_index[1] == d))
52
+ | ((sub_edge_index[0] == d) & (sub_edge_index[1] == s)))
53
+ sub_edge_index = sub_edge_index[:, keep] # hide the link to predict
54
+ z = drnl_node_labeling(sub_edge_index, s, d, num_nodes=len(nodes))
55
+ graphs.append((z, sub_edge_index, y))
56
+
57
+ self.num_labels = num_labels or max(int(z.max()) for z, _, _ in graphs) + 1
58
+ self.data, self.slices = self.collate([
59
+ Data(x=np.eye(self.num_labels, dtype=np.float32)[np.minimum(z, self.num_labels - 1)],
60
+ edge_index=e, y=np.array([y], dtype=np.float32))
61
+ for z, e, y in graphs
62
+ ])
63
+
64
+ @staticmethod
65
+ def _links(data):
66
+ def arr(x):
67
+ return np.asarray(ops.convert_to_numpy(x))
68
+
69
+ if getattr(data, "edge_label_index", None) is not None:
70
+ return arr(data.edge_label_index).astype(np.int64), arr(data.edge_label)
71
+ pos = arr(data.pos_edge_label_index).astype(np.int64)
72
+ neg = getattr(data, "neg_edge_label_index", None)
73
+ neg = np.zeros((2, 0), np.int64) if neg is None else arr(neg).astype(np.int64)
74
+ return np.concatenate([pos, neg], axis=1), np.concatenate([np.ones(pos.shape[1]), np.zeros(neg.shape[1])])
@@ -0,0 +1,92 @@
1
+ from typing import Callable, Optional
2
+
3
+ import numpy as np
4
+
5
+ from k3_node.data import Data, InMemoryDataset
6
+
7
+
8
+ class ShapeScenes(InMemoryDataset):
9
+ r"""Small point cloud scenes for semantic segmentation, built from :class:`GeometricShapes`
10
+ meshes (a compact stand-in for ShapeNet part segmentation).
11
+
12
+ Every scene contains ``shapes_per_scene`` objects, each a cube, sphere, cone or torus that is
13
+ randomly scaled, rotated and placed apart from the others. ``num_points`` points are sampled on
14
+ their surfaces: ``pos`` holds the positions (scaled to the unit ball), ``x`` the surface
15
+ normals, and ``y`` the type of the object each point lies on (4 classes). Training scenes use
16
+ the training meshes, test scenes the test meshes; the scenes are generated with a fixed seed.
17
+
18
+ Args:
19
+ root (str): Directory where GeometricShapes is (or will be) downloaded.
20
+ train (bool, optional): Build training scenes if ``True``, else test scenes.
21
+ num_scenes (int, optional): Number of scenes. (default: ``200`` for training, ``50`` for test)
22
+ shapes_per_scene (int, optional): Objects per scene. (default: ``3``)
23
+ num_points (int, optional): Points per scene. (default: ``1024``)
24
+ transform (callable, optional): A function applied to each scene when it is accessed.
25
+ seed (int, optional): Random seed. (default: ``0``)
26
+ """
27
+
28
+ categories = ["3d_cube", "3d_sphere", "3d_cone", "3d_torus"]
29
+
30
+ def __init__(self, root: str, train: bool = True, num_scenes: Optional[int] = None, shapes_per_scene: int = 3,
31
+ num_points: int = 1024, transform: Optional[Callable] = None, seed: int = 0):
32
+ super().__init__(None, transform)
33
+ from k3_node.datasets.geometric_shapes import GeometricShapes
34
+ from k3_node.transforms.spatial import _np
35
+
36
+ shapes = GeometricShapes(root, train=train)
37
+ names = sorted(__import__("os").listdir(shapes.raw_dir))
38
+ meshes = {}
39
+ for i in range(len(shapes)):
40
+ data = shapes[i]
41
+ name = names[int(_np(data.y)[0])]
42
+ if name in self.categories:
43
+ meshes[self.categories.index(name)] = (_np(data.pos).astype(np.float64), _np(data.face))
44
+
45
+ rng = np.random.default_rng(seed + (0 if train else 1))
46
+ num_scenes = num_scenes or (200 if train else 50)
47
+ self.data, self.slices = self.collate(
48
+ [self._scene(rng, meshes, shapes_per_scene, num_points) for _ in range(num_scenes)])
49
+
50
+ @staticmethod
51
+ def _sample(rng, pos, face, num):
52
+ a, b, c = pos[face[0]], pos[face[1]], pos[face[2]]
53
+ cross = np.cross(b - a, c - a)
54
+ area = np.linalg.norm(cross, axis=1)
55
+ tri = rng.choice(len(area), size=num, p=area / area.sum())
56
+ u, v = rng.random((2, num))
57
+ flip = u + v > 1
58
+ u[flip], v[flip] = 1 - u[flip], 1 - v[flip]
59
+ points = a[tri] + u[:, None] * (b - a)[tri] + v[:, None] * (c - a)[tri]
60
+ normals = cross[tri] / np.maximum(area[tri, None], 1e-12)
61
+ return points, normals
62
+
63
+ @staticmethod
64
+ def _rotation(rng):
65
+ q = rng.normal(size=4)
66
+ q /= np.linalg.norm(q)
67
+ w, x, y, z = q
68
+ return np.array([[1 - 2 * (y * y + z * z), 2 * (x * y - z * w), 2 * (x * z + y * w)],
69
+ [2 * (x * y + z * w), 1 - 2 * (x * x + z * z), 2 * (y * z - x * w)],
70
+ [2 * (x * z - y * w), 2 * (y * z + x * w), 1 - 2 * (x * x + y * y)]])
71
+
72
+ def _scene(self, rng, meshes, shapes_per_scene, num_points):
73
+ counts = np.full(shapes_per_scene, num_points // shapes_per_scene)
74
+ counts[: num_points - counts.sum()] += 1
75
+ angle = rng.random() * 2 * np.pi
76
+ positions, normals, labels = [], [], []
77
+ for i in range(shapes_per_scene):
78
+ label = int(rng.integers(len(self.categories)))
79
+ pos, face = meshes[label]
80
+ points, normal = self._sample(rng, pos, face, counts[i])
81
+ points = points / np.abs(points).max()
82
+ rot = self._rotation(rng)
83
+ theta = angle + 2 * np.pi * i / shapes_per_scene # objects around a circle, apart
84
+ offset = 2.5 * np.array([np.cos(theta), np.sin(theta), 0.0])
85
+ positions.append(points @ rot.T * rng.uniform(0.6, 1.0) + offset)
86
+ normals.append(normal @ rot.T)
87
+ labels.append(np.full(counts[i], label))
88
+ pos = np.concatenate(positions)
89
+ pos = pos - pos.mean(axis=0)
90
+ pos = pos / np.linalg.norm(pos, axis=1).max()
91
+ return Data(pos=pos.astype(np.float32), x=np.concatenate(normals).astype(np.float32),
92
+ y=np.concatenate(labels).astype(np.int64))