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,188 @@
1
+ from typing import Optional
2
+
3
+ import numpy as np
4
+ from keras import ops
5
+
6
+ try:
7
+ import torch
8
+ import torch.utils.data
9
+ BaseDataLoader = torch.utils.data.DataLoader
10
+ except ImportError:
11
+ torch = None
12
+ BaseDataLoader = object
13
+
14
+ from k3_node.data import Data
15
+ from k3_node.loader.keras_dataset import loader_bases
16
+
17
+
18
+ def _np(x):
19
+ return np.asarray(ops.convert_to_numpy(x))
20
+
21
+
22
+ class _SamplingSteps:
23
+ """The sampler's "dataset": item ``i`` is a freshly sampled ``(node_idx, edge_idx)`` pair."""
24
+
25
+ def __init__(self, sampler):
26
+ self.sampler = sampler
27
+
28
+ def __len__(self):
29
+ return self.sampler.num_steps
30
+
31
+ def __getitem__(self, idx):
32
+ return self.sampler._sample_subgraph()
33
+
34
+
35
+ class GraphSAINTSampler(*loader_bases(BaseDataLoader)):
36
+ r"""The GraphSAINT sampler base class from the `"GraphSAINT: Graph Sampling Based Inductive
37
+ Learning Method" <https://arxiv.org/abs/1907.04931>`_ paper. Every step samples a set of
38
+ nodes and yields the subgraph they induce.
39
+
40
+ With ``sample_coverage > 0``, normalization statistics are estimated beforehand by sampling
41
+ about ``sample_coverage`` times every node: ``node_norm`` (to weight each node's loss) and
42
+ ``edge_norm`` (to weight each edge's message) are added to every subgraph.
43
+
44
+ Args:
45
+ data (Data): The graph data object.
46
+ batch_size (int): The approximate number of samples per batch (see the subclasses).
47
+ num_steps (int, optional): The number of iterations per epoch. (default: ``1``)
48
+ sample_coverage (int): How many samples per node to compute the normalization
49
+ statistics with; ``0`` skips them. (default: ``0``)
50
+ save_dir (str, optional): Unused; kept for API compatibility with PyG.
51
+ log (bool, optional): Unused; kept for API compatibility with PyG.
52
+ **kwargs (optional): Additional arguments of :class:`torch.utils.data.DataLoader`.
53
+ """
54
+
55
+ def __init__(self, data: Data, batch_size: int, num_steps: int = 1, sample_coverage: int = 0,
56
+ save_dir: Optional[str] = None, log: bool = True, **kwargs):
57
+ kwargs.pop('dataset', None)
58
+ kwargs.pop('collate_fn', None)
59
+ kwargs.pop('shuffle', None)
60
+ assert data.edge_index is not None
61
+
62
+ self.num_steps = num_steps
63
+ self._batch_size = batch_size
64
+ self.sample_coverage = sample_coverage
65
+ self.save_dir = save_dir
66
+ self.log = log
67
+ self.data = data
68
+ self.N = data.num_nodes
69
+ edge_index = _np(data.edge_index).astype(np.int64)
70
+ self.E = edge_index.shape[1]
71
+ # CSR by source node: the edges leaving node i are perm[rowptr[i]:rowptr[i+1]]
72
+ self._row, self._col = edge_index
73
+ self._perm = np.argsort(self._row, kind="stable")
74
+ self._rowptr = np.concatenate([[0], np.cumsum(np.bincount(self._row, minlength=self.N))])
75
+
76
+ steps = _SamplingSteps(self)
77
+ if torch is not None:
78
+ super().__init__(steps, batch_size=1, collate_fn=self._collate, **kwargs)
79
+ else:
80
+ self.dataset = steps
81
+ self.batch_size = 1
82
+ self.collate_fn = self._collate
83
+
84
+ if self.sample_coverage > 0:
85
+ self.node_norm, self.edge_norm = self._compute_norm()
86
+
87
+ def _sample_nodes(self, batch_size: int) -> np.ndarray:
88
+ raise NotImplementedError
89
+
90
+ def _sample_subgraph(self):
91
+ node_idx = np.unique(self._sample_nodes(self._batch_size))
92
+ in_sample = np.zeros(self.N, dtype=bool)
93
+ in_sample[node_idx] = True
94
+ edge_idx = np.nonzero(in_sample[self._row] & in_sample[self._col])[0]
95
+ return node_idx, edge_idx
96
+
97
+ def _collate(self, data_list):
98
+ node_idx, edge_idx = data_list[0]
99
+ new_id = np.full(self.N, -1, dtype=np.int64)
100
+ new_id[node_idx] = np.arange(len(node_idx))
101
+
102
+ data = Data()
103
+ data.num_nodes = len(node_idx)
104
+ data.edge_index = ops.convert_to_tensor(
105
+ np.stack([new_id[self._row[edge_idx]], new_id[self._col[edge_idx]]]), dtype="int64")
106
+ for key, item in self.data.items():
107
+ if key in ('edge_index', 'num_nodes'):
108
+ continue
109
+ shape = getattr(item, 'shape', None)
110
+ if shape is not None and len(shape) > 0 and shape[0] == self.N:
111
+ data[key] = ops.take(item, node_idx, axis=0)
112
+ elif shape is not None and len(shape) > 0 and shape[0] == self.E:
113
+ data[key] = ops.take(item, edge_idx, axis=0)
114
+ else:
115
+ data[key] = item
116
+ if self.sample_coverage > 0:
117
+ data.node_norm = ops.convert_to_tensor(self.node_norm[node_idx])
118
+ data.edge_norm = ops.convert_to_tensor(self.edge_norm[edge_idx])
119
+ return data
120
+
121
+ def _compute_norm(self):
122
+ node_count = np.zeros(self.N, dtype=np.float32)
123
+ edge_count = np.zeros(self.E, dtype=np.float32)
124
+ num_samples = total_sampled_nodes = 0
125
+ while total_sampled_nodes < self.N * self.sample_coverage:
126
+ for _ in range(self.num_steps):
127
+ node_idx, edge_idx = self._sample_subgraph()
128
+ node_count[node_idx] += 1
129
+ edge_count[edge_idx] += 1
130
+ total_sampled_nodes += len(node_idx)
131
+ num_samples += self.num_steps
132
+
133
+ with np.errstate(divide='ignore', invalid='ignore'):
134
+ edge_norm = np.clip(node_count[self._row] / edge_count, 0, 1e4)
135
+ edge_norm[np.isnan(edge_norm)] = 0.1
136
+ node_count[node_count == 0] = 0.1
137
+ node_norm = num_samples / node_count / self.N
138
+ return node_norm.astype(np.float32), edge_norm.astype(np.float32)
139
+
140
+ def _random_walk(self, start: np.ndarray, walk_length: int) -> np.ndarray:
141
+ walks, cur = [start], start
142
+ for _ in range(walk_length):
143
+ deg = self._rowptr[cur + 1] - self._rowptr[cur]
144
+ offset = np.floor(np.random.rand(len(cur)) * np.maximum(deg, 1)).astype(np.int64)
145
+ nxt = self._col[self._perm[np.minimum(self._rowptr[cur] + offset, self.E - 1)]]
146
+ cur = np.where(deg > 0, nxt, cur) # nodes without neighbors stay put
147
+ walks.append(cur)
148
+ return np.stack(walks, axis=1)
149
+
150
+
151
+ class GraphSAINTNodeSampler(GraphSAINTSampler):
152
+ r"""The GraphSAINT node sampler: samples ``batch_size`` nodes, each with probability
153
+ proportional to its out-degree."""
154
+
155
+ def _sample_nodes(self, batch_size: int) -> np.ndarray:
156
+ return self._row[np.random.randint(0, self.E, size=batch_size)]
157
+
158
+
159
+ class GraphSAINTEdgeSampler(GraphSAINTSampler):
160
+ r"""The GraphSAINT edge sampler: samples ``batch_size`` edges, each with probability
161
+ proportional to :math:`1 / \deg(u) + 1 / \deg(v)`, and keeps their endpoints."""
162
+
163
+ def _sample_nodes(self, batch_size: int) -> np.ndarray:
164
+ out_deg = np.maximum(np.bincount(self._row, minlength=self.N), 1)
165
+ in_deg = np.maximum(np.bincount(self._col, minlength=self.N), 1)
166
+ prob = 1.0 / in_deg[self._row] + 1.0 / out_deg[self._col]
167
+ # Weighted sampling without replacement (exponential keys, as in PyG)
168
+ keys = np.log(np.random.rand(self.E)) / (prob + 1e-10)
169
+ edge_sample = np.argsort(-keys)[:batch_size]
170
+ return np.concatenate([self._col[edge_sample], self._row[edge_sample]])
171
+
172
+
173
+ class GraphSAINTRandomWalkSampler(GraphSAINTSampler):
174
+ r"""The GraphSAINT random walk sampler: starts ``batch_size`` random walks of length
175
+ ``walk_length`` and keeps the visited nodes.
176
+
177
+ Args:
178
+ walk_length (int): Length of each random walk.
179
+ """
180
+
181
+ def __init__(self, data: Data, batch_size: int, walk_length: int, num_steps: int = 1,
182
+ sample_coverage: int = 0, save_dir: Optional[str] = None, log: bool = True, **kwargs):
183
+ self.walk_length = walk_length
184
+ super().__init__(data, batch_size, num_steps, sample_coverage, save_dir, log, **kwargs)
185
+
186
+ def _sample_nodes(self, batch_size: int) -> np.ndarray:
187
+ start = np.random.randint(0, self.N, size=batch_size)
188
+ return self._random_walk(start, self.walk_length).reshape(-1)
@@ -0,0 +1,90 @@
1
+ from typing import Any, Callable, Dict, List, Optional, Tuple, Union
2
+
3
+
4
+ from k3_node.data import HeteroData
5
+ from k3_node.loader.node_loader import HeteroSamplerOutput, NodeLoader, NodeSamplerInput
6
+ from k3_node.loader.sampler_utils import sample_neighbors_hetero
7
+
8
+
9
+ class InternalHGTSampler:
10
+ r"""HGT balanced neighborhood sampling engine."""
11
+ def __init__(
12
+ self,
13
+ data: HeteroData,
14
+ num_samples: Union[List[int], Dict[str, List[int]]],
15
+ ):
16
+ self.data = data
17
+ self.num_samples = num_samples
18
+ self.edge_permutation = None
19
+
20
+ def sample_from_nodes(self, input_data: NodeSamplerInput) -> HeteroSamplerOutput:
21
+ edge_index_dict = {}
22
+ for edge_type in self.data.edge_types:
23
+ canonical = self.data._to_canonical(*edge_type) if hasattr(self.data, '_to_canonical') else edge_type
24
+ edge_index_dict[canonical] = self.data[edge_type].edge_index
25
+
26
+ node_type = input_data.input_type or self.data.node_types[0]
27
+ seed_dict = {k: None for k in self.data.node_types}
28
+ seed_dict[node_type] = input_data.node
29
+
30
+ # Build num_neighbors dict for hetero sampling
31
+ if isinstance(self.num_samples, dict):
32
+ num_neighbors = {}
33
+ for e in self.data.edge_types:
34
+ dst = e[2]
35
+ can = self.data._to_canonical(*e) if hasattr(self.data, '_to_canonical') else e
36
+ num_neighbors[can] = self.num_samples.get(dst, [10])
37
+ else:
38
+ num_neighbors = self.num_samples
39
+
40
+ node_dict, row_dict, col_dict, edge_dict, n_counts, e_counts = sample_neighbors_hetero(
41
+ edge_index_dict=edge_index_dict,
42
+ seed_nodes_dict=seed_dict,
43
+ num_neighbors=num_neighbors,
44
+ replace=True,
45
+ subgraph_type='directional',
46
+ )
47
+
48
+ return HeteroSamplerOutput(
49
+ node=node_dict,
50
+ row=row_dict,
51
+ col=col_dict,
52
+ edge=edge_dict,
53
+ num_sampled_nodes=n_counts,
54
+ num_sampled_edges=e_counts,
55
+ metadata=(input_data.input_id, input_data.time),
56
+ )
57
+
58
+
59
+ class HGTLoader(NodeLoader):
60
+ r"""The Heterogeneous Graph Sampler from the "Heterogeneous Graph Transformer" paper.
61
+
62
+ Args:
63
+ data (HeteroData): The heterogeneous graph data object.
64
+ num_samples (List[int] or Dict[str, List[int]]): The number of nodes to sample per iteration.
65
+ input_nodes (str or Tuple[str, Tensor]): Seed node type and indices.
66
+ **kwargs (optional): Additional arguments of :class:`torch.utils.data.DataLoader`.
67
+ """
68
+ def __init__(
69
+ self,
70
+ data: HeteroData,
71
+ num_samples: Union[List[int], Dict[str, List[int]]],
72
+ input_nodes: Union[str, Tuple[str, Optional[Any]]],
73
+ is_sorted: bool = False,
74
+ transform: Optional[Callable] = None,
75
+ transform_sampler_output: Optional[Callable] = None,
76
+ filter_per_worker: Optional[bool] = None,
77
+ **kwargs,
78
+ ):
79
+ hgt_sampler = InternalHGTSampler(data, num_samples=num_samples)
80
+
81
+ super().__init__(
82
+ data=data,
83
+ node_sampler=hgt_sampler,
84
+ input_nodes=input_nodes,
85
+ transform=transform,
86
+ transform_sampler_output=transform_sampler_output,
87
+ filter_per_worker=filter_per_worker,
88
+ **kwargs,
89
+ )
90
+
@@ -0,0 +1,87 @@
1
+ from typing import Any, List, Optional, Union
2
+
3
+ import numpy as np
4
+
5
+ try:
6
+ import torch
7
+ from torch import Tensor
8
+ BaseWeightedRandomSampler = torch.utils.data.WeightedRandomSampler
9
+ except ImportError:
10
+ torch = None
11
+ Tensor = type(None)
12
+ BaseWeightedRandomSampler = object
13
+
14
+ from k3_node.data import Data, Dataset, InMemoryDataset
15
+
16
+
17
+ class ImbalancedSampler(BaseWeightedRandomSampler):
18
+ r"""A weighted random sampler that randomly samples elements according to class distribution.
19
+
20
+ Args:
21
+ dataset (Dataset or Data or Tensor): The dataset or class distribution from which to sample.
22
+ input_nodes (Tensor, optional): The indices of nodes used by the corresponding loader. (default: :obj:`None`)
23
+ num_samples (int, optional): The number of samples to draw for a single epoch. (default: :obj:`None`)
24
+ """
25
+ def __init__(
26
+ self,
27
+ dataset: Union[Dataset, Data, List[Data], Any],
28
+ input_nodes: Optional[Any] = None,
29
+ num_samples: Optional[int] = None,
30
+ ):
31
+ if isinstance(dataset, Data):
32
+ y = dataset.y
33
+ if hasattr(y, 'view'):
34
+ y = y.view(-1)
35
+ else:
36
+ y = np.asarray(y).reshape(-1)
37
+ if input_nodes is not None:
38
+ y = y[input_nodes]
39
+
40
+ elif torch is not None and isinstance(dataset, Tensor):
41
+ y = dataset.view(-1)
42
+ if input_nodes is not None:
43
+ y = y[input_nodes]
44
+
45
+ elif isinstance(dataset, InMemoryDataset):
46
+ y = dataset.y
47
+ if hasattr(y, 'view'):
48
+ y = y.view(-1)
49
+ else:
50
+ y = np.asarray(y).reshape(-1)
51
+
52
+ elif isinstance(dataset, (list, tuple)):
53
+ ys = [data.y for data in dataset]
54
+ if torch is not None and isinstance(ys[0], Tensor):
55
+ y = torch.cat(ys, dim=0).view(-1)
56
+ else:
57
+ y = np.concatenate([np.asarray(x).reshape(-1) for x in ys], axis=0)
58
+ else:
59
+ y = np.asarray(dataset).reshape(-1)
60
+
61
+ if torch is not None and not isinstance(y, Tensor):
62
+ y = torch.as_tensor(y, dtype=torch.long)
63
+
64
+ num_samples = (y.numel() if hasattr(y, 'numel') else len(y)) if num_samples is None else num_samples
65
+
66
+ if torch is not None and isinstance(y, Tensor):
67
+ bincount = y.bincount().float()
68
+ class_weight = 1.0 / bincount
69
+ weight = class_weight[y]
70
+ super().__init__(weight, num_samples, replacement=True)
71
+ else:
72
+ classes, counts = np.unique(y, return_counts=True)
73
+ class_weight = {c: 1.0 / count for c, count in zip(classes, counts)}
74
+ weight = np.array([class_weight[int(val)] for val in y], dtype=np.float64)
75
+ weight = weight / weight.sum()
76
+ self.weight = weight
77
+ self.num_samples = num_samples
78
+ self.replacement = True
79
+
80
+ def __iter__(self):
81
+ if torch is not None and hasattr(super(), '__iter__'):
82
+ return super().__iter__()
83
+ indices = np.random.choice(len(self.weight), size=self.num_samples, replace=True, p=self.weight)
84
+ return iter(indices.tolist())
85
+
86
+ def __len__(self):
87
+ return self.num_samples