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,499 @@
1
+ import math
2
+ from typing import Any, Dict, List, Optional, Tuple, Union
3
+
4
+ import numpy as np
5
+
6
+ try:
7
+ import torch
8
+ from torch import Tensor
9
+ except ImportError:
10
+ torch = None
11
+ Tensor = type(None)
12
+
13
+
14
+ class FastGraph:
15
+ r"""Fast CSR-based adjacency index for neighbor lookups in pure Python / NumPy."""
16
+ def __init__(self, edge_index: Any, num_nodes: Optional[int] = None):
17
+ if torch is not None and isinstance(edge_index, Tensor):
18
+ np_edge_index = edge_index.detach().cpu().numpy()
19
+ else:
20
+ np_edge_index = np.asarray(edge_index)
21
+
22
+ row, col = np_edge_index[0], np_edge_index[1]
23
+ if num_nodes is None:
24
+ num_nodes = int(max(np.max(row), np.max(col)) + 1) if len(row) > 0 else 0
25
+
26
+ self.num_nodes = num_nodes
27
+ self.num_edges = len(row)
28
+
29
+ # Build CSC layout (target -> sources) for incoming neighbor sampling
30
+ # col is target, row is source
31
+ order = np.argsort(col, kind='mergesort')
32
+ sorted_col = col[order]
33
+ self.sorted_row = row[order]
34
+ self.sorted_edge_id = order.astype(np.int64)
35
+
36
+ # Compute indptr
37
+ counts = np.bincount(sorted_col, minlength=num_nodes)
38
+ self.indptr = np.zeros(num_nodes + 1, dtype=np.int64)
39
+ np.cumsum(counts, out=self.indptr[1:])
40
+
41
+ def local_map(self, size: int) -> np.ndarray:
42
+ """A reusable global-to-local node id array filled with -1 (callers must reset what they set)."""
43
+ local = getattr(self, "_local", None)
44
+ if local is None or len(local) < size:
45
+ local = np.full(size, -1, dtype=np.int64)
46
+ self._local = local
47
+ return local
48
+
49
+ def get_neighbors(self, node: int) -> Tuple[np.ndarray, np.ndarray]:
50
+ if node >= self.num_nodes:
51
+ return np.empty(0, dtype=np.int64), np.empty(0, dtype=np.int64)
52
+ start = self.indptr[node]
53
+ end = self.indptr[node + 1]
54
+ return self.sorted_row[start:end], self.sorted_edge_id[start:end]
55
+
56
+
57
+ def _csr_ranges(starts: np.ndarray, counts: np.ndarray) -> Tuple[np.ndarray, np.ndarray]:
58
+ """Flattens the ranges ``[starts[i], starts[i] + counts[i])``; returns (segment id, position)."""
59
+ counts = counts.astype(np.int64)
60
+ segment = np.repeat(np.arange(len(counts), dtype=np.int64), counts)
61
+ offsets = np.arange(int(counts.sum()), dtype=np.int64) - np.repeat(np.cumsum(counts) - counts, counts)
62
+ return segment, starts.astype(np.int64)[segment] + offsets
63
+
64
+
65
+ def _sample_positions(
66
+ starts: np.ndarray, counts: np.ndarray, k: int, replace: bool
67
+ ) -> Tuple[np.ndarray, np.ndarray]:
68
+ """Picks up to ``k`` CSR positions per segment (all of them if ``k == -1`` or ``count <= k``).
69
+
70
+ Returns (segment id, position), grouped by segment in ascending segment order.
71
+ """
72
+ take_all = (counts <= k) | (k == -1)
73
+ seg_all, pos_all = _csr_ranges(starts[take_all], counts[take_all])
74
+ seg_all = np.nonzero(take_all)[0][seg_all]
75
+
76
+ sampled = np.nonzero(~take_all)[0]
77
+ if len(sampled) == 0:
78
+ return seg_all, pos_all
79
+ if replace:
80
+ seg_s = np.repeat(sampled, k)
81
+ pos_s = starts[seg_s] + np.floor(np.random.random(len(seg_s)) * counts[seg_s]).astype(np.int64)
82
+ else:
83
+ # Random keys per candidate; the k smallest keys of each segment form a uniform subset.
84
+ seg_c, pos_c = _csr_ranges(starts[sampled], counts[sampled])
85
+ order = np.lexsort((np.random.random(len(seg_c)), seg_c))
86
+ seg_c, pos_c = seg_c[order], pos_c[order]
87
+ first = np.searchsorted(seg_c, seg_c, side="left")
88
+ keep = (np.arange(len(seg_c)) - first) < k
89
+ seg_s, pos_s = sampled[seg_c[keep]], pos_c[keep]
90
+
91
+ segment = np.concatenate([seg_all, seg_s])
92
+ position = np.concatenate([pos_all, pos_s])
93
+ order = np.argsort(segment, kind="stable")
94
+ return segment[order], position[order]
95
+
96
+
97
+ def sample_neighbors_homo(
98
+ edge_index: Any,
99
+ seed_nodes: Any,
100
+ num_neighbors: List[int],
101
+ num_nodes: Optional[int] = None,
102
+ replace: bool = False,
103
+ subgraph_type: str = 'directional',
104
+ disjoint: bool = False,
105
+ graph: Optional[FastGraph] = None,
106
+ ) -> Tuple[Any, Any, Any, Any, List[int], List[int]]:
107
+ r"""Samples multi-hop neighborhoods on a homogeneous graph.
108
+
109
+ All steps are vectorized with NumPy. Pass a prebuilt ``graph`` (a :class:`FastGraph` of
110
+ ``edge_index``) to avoid rebuilding the CSR index for every batch.
111
+
112
+ Returns:
113
+ (sampled_nodes, local_row, local_col, edge_ids, num_sampled_nodes, num_sampled_edges)
114
+ """
115
+ is_torch = torch is not None and isinstance(seed_nodes, Tensor)
116
+ if is_torch:
117
+ device = seed_nodes.device
118
+ np_seeds = seed_nodes.detach().cpu().numpy()
119
+ else:
120
+ device = None
121
+ np_seeds = np.asarray(seed_nodes)
122
+ np_seeds = np_seeds.astype(np.int64).reshape(-1)
123
+
124
+ if graph is None:
125
+ graph = FastGraph(edge_index, num_nodes=num_nodes)
126
+
127
+ # Global -> local id map; only the entries touched here are reset afterwards.
128
+ size = max(graph.num_nodes, int(np_seeds.max()) + 1 if len(np_seeds) else 0)
129
+ local = graph.local_map(size)
130
+
131
+ _, first_idx = np.unique(np_seeds, return_index=True)
132
+ first_idx = np.sort(first_idx)
133
+ frontier = np_seeds[first_idx]
134
+ node_chunks = [frontier]
135
+ local[frontier] = np.arange(len(frontier))
136
+ num_total = len(frontier)
137
+ batch_ids = first_idx.astype(np.int64) if disjoint else None
138
+
139
+ num_sampled_nodes = [len(frontier)]
140
+ num_sampled_edges = []
141
+ hop_src, hop_dst, hop_eid = [], [], []
142
+ try:
143
+ for k in num_neighbors:
144
+ in_graph = frontier < graph.num_nodes
145
+ starts = np.zeros(len(frontier), dtype=np.int64)
146
+ counts = np.zeros(len(frontier), dtype=np.int64)
147
+ starts[in_graph] = graph.indptr[frontier[in_graph]]
148
+ counts[in_graph] = graph.indptr[frontier[in_graph] + 1] - starts[in_graph]
149
+
150
+ segment, position = _sample_positions(starts, counts, k, replace)
151
+ srcs = graph.sorted_row[position].astype(np.int64)
152
+ dsts = frontier[segment]
153
+ hop_src.append(srcs)
154
+ hop_dst.append(dsts)
155
+ hop_eid.append(graph.sorted_edge_id[position])
156
+ num_sampled_edges.append(len(srcs))
157
+
158
+ # Unvisited sources become new nodes, in order of first appearance.
159
+ is_new = local[srcs] == -1
160
+ new_srcs = srcs[is_new]
161
+ uniq, first = np.unique(new_srcs, return_index=True)
162
+ order = np.argsort(first, kind="stable")
163
+ new_nodes = uniq[order]
164
+ local[new_nodes] = num_total + np.arange(len(new_nodes))
165
+ if disjoint:
166
+ discoverer = dsts[is_new][first[order]]
167
+ batch_ids = np.concatenate([batch_ids, batch_ids[local[discoverer]]])
168
+ num_total += len(new_nodes)
169
+ node_chunks.append(new_nodes)
170
+ num_sampled_nodes.append(len(new_nodes))
171
+ frontier = new_nodes
172
+
173
+ nodes = np.concatenate(node_chunks).astype(np.int64)
174
+
175
+ if subgraph_type == 'induced':
176
+ # Every original edge whose endpoints were both visited, in original edge order.
177
+ targets = nodes[nodes < graph.num_nodes]
178
+ t_starts = graph.indptr[targets]
179
+ segment, position = _csr_ranges(t_starts, graph.indptr[targets + 1] - t_starts)
180
+ srcs = graph.sorted_row[position].astype(np.int64)
181
+ keep = local[srcs] != -1
182
+ edge_ids = graph.sorted_edge_id[position][keep]
183
+ order = np.argsort(edge_ids, kind="stable")
184
+ edge_ids = edge_ids[order]
185
+ local_row = local[srcs[keep]][order]
186
+ local_col = local[targets[segment[keep]]][order]
187
+ else:
188
+ srcs = np.concatenate(hop_src) if hop_src else np.empty(0, dtype=np.int64)
189
+ dsts = np.concatenate(hop_dst) if hop_dst else np.empty(0, dtype=np.int64)
190
+ edge_ids = np.concatenate(hop_eid) if hop_eid else np.empty(0, dtype=np.int64)
191
+ local_row, local_col = local[srcs], local[dsts]
192
+ if subgraph_type == 'bidirectional':
193
+ # Each sampled edge followed by its reverse.
194
+ local_row, local_col = (
195
+ np.stack([local_row, local_col], axis=1).reshape(-1),
196
+ np.stack([local_col, local_row], axis=1).reshape(-1),
197
+ )
198
+ edge_ids = np.repeat(edge_ids, 2)
199
+ finally:
200
+ local[np.concatenate(node_chunks)] = -1
201
+
202
+ local_row = local_row.astype(np.int64)
203
+ local_col = local_col.astype(np.int64)
204
+ edge_ids = edge_ids.astype(np.int64)
205
+
206
+ if is_torch:
207
+ nodes = torch.from_numpy(nodes).to(device=device, dtype=torch.long)
208
+ local_row = torch.from_numpy(local_row).to(device=device, dtype=torch.long)
209
+ local_col = torch.from_numpy(local_col).to(device=device, dtype=torch.long)
210
+ edge_ids = torch.from_numpy(edge_ids).to(device=device, dtype=torch.long)
211
+
212
+ return nodes, local_row, local_col, edge_ids, num_sampled_nodes, num_sampled_edges
213
+
214
+
215
+ def sample_neighbors_hetero(
216
+ edge_index_dict: Dict[Tuple[str, str, str], Any],
217
+ seed_nodes_dict: Dict[str, Any],
218
+ num_neighbors: Union[List[int], Dict[Tuple[str, str, str], List[int]]],
219
+ num_nodes_dict: Optional[Dict[str, int]] = None,
220
+ replace: bool = False,
221
+ subgraph_type: str = 'directional',
222
+ ) -> Tuple[Dict[str, Any], Dict[Tuple[str, str, str], Any], Dict[Tuple[str, str, str], Any], Dict[Tuple[str, str, str], Any], Dict[str, List[int]], Dict[Tuple[str, str, str], List[int]]]:
223
+ r"""Samples multi-hop neighborhoods on a heterogeneous graph."""
224
+ # Build fast graphs per edge type
225
+ graphs = {}
226
+ for edge_type, edge_index in edge_index_dict.items():
227
+ graphs[edge_type] = FastGraph(edge_index)
228
+
229
+ # Determine num_hops
230
+ if isinstance(num_neighbors, dict):
231
+ first_key = list(num_neighbors.keys())[0]
232
+ num_hops = len(num_neighbors[first_key])
233
+ else:
234
+ num_hops = len(num_neighbors)
235
+
236
+ is_torch = torch is not None and any(isinstance(v, Tensor) for v in seed_nodes_dict.values())
237
+ device = None
238
+ if is_torch:
239
+ for v in seed_nodes_dict.values():
240
+ if isinstance(v, Tensor):
241
+ device = v.device
242
+ break
243
+
244
+ nodes_dict = {k: [] for k in seed_nodes_dict.keys()}
245
+ visited_dict = {k: {} for k in seed_nodes_dict.keys()}
246
+
247
+ for node_type, seeds in seed_nodes_dict.items():
248
+ if seeds is None:
249
+ continue
250
+ np_seeds = seeds.detach().cpu().numpy() if (torch is not None and isinstance(seeds, Tensor)) else np.asarray(seeds)
251
+ for s in np_seeds:
252
+ s = int(s)
253
+ if s not in visited_dict[node_type]:
254
+ visited_dict[node_type][s] = len(nodes_dict[node_type])
255
+ nodes_dict[node_type].append(s)
256
+
257
+ frontier_dict = {k: list(v) for k, v in nodes_dict.items()}
258
+ num_sampled_nodes = {k: [len(v)] for k, v in nodes_dict.items()}
259
+ num_sampled_edges = {k: [] for k in edge_index_dict.keys()}
260
+ sampled_edges_dict = {k: [] for k in edge_index_dict.keys()}
261
+
262
+ for hop in range(num_hops):
263
+ next_frontier_dict = {k: [] for k in visited_dict.keys()}
264
+
265
+ for edge_type, graph in graphs.items():
266
+ src_type, rel, dst_type = edge_type
267
+ if dst_type not in frontier_dict:
268
+ continue
269
+
270
+ if isinstance(num_neighbors, dict):
271
+ k = num_neighbors[edge_type][hop]
272
+ else:
273
+ k = num_neighbors[hop]
274
+
275
+ hop_edges = 0
276
+ for target in frontier_dict[dst_type]:
277
+ srcs, e_ids = graph.get_neighbors(target)
278
+ count = len(srcs)
279
+ if count == 0:
280
+ continue
281
+
282
+ if k == -1 or k >= count:
283
+ chosen_idx = np.arange(count)
284
+ else:
285
+ chosen_idx = np.random.choice(count, size=k, replace=replace)
286
+
287
+ for s, e in zip(srcs[chosen_idx], e_ids[chosen_idx]):
288
+ s = int(s)
289
+ e = int(e)
290
+ sampled_edges_dict[edge_type].append((s, target, e))
291
+ hop_edges += 1
292
+
293
+ if src_type not in visited_dict:
294
+ visited_dict[src_type] = {}
295
+ nodes_dict[src_type] = []
296
+ next_frontier_dict[src_type] = []
297
+ num_sampled_nodes[src_type] = [0] * (hop + 1)
298
+
299
+ if s not in visited_dict[src_type]:
300
+ visited_dict[src_type][s] = len(nodes_dict[src_type])
301
+ nodes_dict[src_type].append(s)
302
+ next_frontier_dict[src_type].append(s)
303
+
304
+ num_sampled_edges[edge_type].append(hop_edges)
305
+
306
+ frontier_dict = next_frontier_dict
307
+ for k in nodes_dict.keys():
308
+ count = len(next_frontier_dict.get(k, []))
309
+ if k in num_sampled_nodes:
310
+ num_sampled_nodes[k].append(count)
311
+
312
+ # Build local rows, cols, edge_ids
313
+ out_row = {}
314
+ out_col = {}
315
+ out_edge = {}
316
+
317
+ for edge_type, edge_list in sampled_edges_dict.items():
318
+ src_type, rel, dst_type = edge_type
319
+ rows, cols, eids = [], [], []
320
+ for u, v, e in edge_list:
321
+ rows.append(visited_dict[src_type][u])
322
+ cols.append(visited_dict[dst_type][v])
323
+ eids.append(e)
324
+
325
+ if len(rows) > 0:
326
+ r = np.array(rows, dtype=np.int64)
327
+ c = np.array(cols, dtype=np.int64)
328
+ e = np.array(eids, dtype=np.int64)
329
+ else:
330
+ r = np.empty(0, dtype=np.int64)
331
+ c = np.empty(0, dtype=np.int64)
332
+ e = np.empty(0, dtype=np.int64)
333
+
334
+ if is_torch:
335
+ out_row[edge_type] = torch.from_numpy(r).to(device=device, dtype=torch.long)
336
+ out_col[edge_type] = torch.from_numpy(c).to(device=device, dtype=torch.long)
337
+ out_edge[edge_type] = torch.from_numpy(e).to(device=device, dtype=torch.long)
338
+ else:
339
+ out_row[edge_type] = r
340
+ out_col[edge_type] = c
341
+ out_edge[edge_type] = e
342
+
343
+ out_nodes = {}
344
+ for k, v in nodes_dict.items():
345
+ arr = np.array(v, dtype=np.int64)
346
+ if is_torch:
347
+ out_nodes[k] = torch.from_numpy(arr).to(device=device, dtype=torch.long)
348
+ else:
349
+ out_nodes[k] = arr
350
+
351
+ return out_nodes, out_row, out_col, out_edge, num_sampled_nodes, num_sampled_edges
352
+
353
+
354
+ def random_walk(edge_index: Any, start_nodes: Any, walk_length: int, num_nodes: Optional[int] = None) -> Any:
355
+ r"""Executes random walks from start_nodes."""
356
+ graph = FastGraph(edge_index, num_nodes=num_nodes)
357
+ is_torch = torch is not None and isinstance(start_nodes, Tensor)
358
+ np_starts = start_nodes.detach().cpu().numpy() if is_torch else np.asarray(start_nodes)
359
+
360
+ walks = []
361
+ for s in np_starts:
362
+ curr = int(s)
363
+ walk = [curr]
364
+ for _ in range(walk_length):
365
+ srcs, _ = graph.get_neighbors(curr)
366
+ if len(srcs) == 0:
367
+ walk.append(curr)
368
+ else:
369
+ curr = int(np.random.choice(srcs))
370
+ walk.append(curr)
371
+ walks.append(walk)
372
+
373
+ walks_arr = np.array(walks, dtype=np.int64)
374
+ if is_torch:
375
+ return torch.from_numpy(walks_arr).to(device=start_nodes.device, dtype=torch.long)
376
+ return walks_arr
377
+
378
+
379
+ def partition_graph(edge_index: Any, num_nodes: int, num_parts: int) -> Any:
380
+ r"""Partitions graph nodes into num_parts clusters."""
381
+ if num_parts <= 1:
382
+ cluster = np.zeros(num_nodes, dtype=np.int64)
383
+ if torch is not None and isinstance(edge_index, Tensor):
384
+ return torch.from_numpy(cluster).to(device=edge_index.device, dtype=torch.long)
385
+ return cluster
386
+
387
+ # Try Metis if installed
388
+ try:
389
+ import torch_geometric.typing as pyg_typing
390
+ if hasattr(pyg_typing, 'WITH_TORCH_SPARSE') and pyg_typing.WITH_TORCH_SPARSE:
391
+ from torch_geometric.index import index2ptr
392
+ from torch_geometric.utils import sort_edge_index
393
+ row, col = sort_edge_index(edge_index, num_nodes=num_nodes)
394
+ indptr = index2ptr(row, size=num_nodes)
395
+ return torch.ops.torch_sparse.partition(indptr.cpu(), col.cpu(), None, num_parts, False).to(edge_index.device)
396
+ except Exception:
397
+ pass
398
+
399
+ # Fast BFS / linear partition fallback
400
+ cluster = np.full(num_nodes, -1, dtype=np.int64)
401
+ part_size = math.ceil(num_nodes / num_parts)
402
+
403
+ graph = FastGraph(edge_index, num_nodes=num_nodes)
404
+ unassigned = set(range(num_nodes))
405
+
406
+ current_part = 0
407
+ while unassigned and current_part < num_parts:
408
+ seed = next(iter(unassigned))
409
+ queue = [seed]
410
+ unassigned.remove(seed)
411
+ cluster[seed] = current_part
412
+ assigned_in_part = 1
413
+
414
+ while queue and assigned_in_part < part_size:
415
+ v = queue.pop(0)
416
+ srcs, _ = graph.get_neighbors(v)
417
+ for u in srcs:
418
+ u = int(u)
419
+ if u in unassigned:
420
+ unassigned.remove(u)
421
+ cluster[u] = current_part
422
+ queue.append(u)
423
+ assigned_in_part += 1
424
+ if assigned_in_part >= part_size:
425
+ break
426
+
427
+ current_part += 1
428
+
429
+ # Any remaining nodes get distributed
430
+ if unassigned:
431
+ for idx, u in enumerate(unassigned):
432
+ cluster[u] = idx % num_parts
433
+
434
+ if torch is not None and isinstance(edge_index, Tensor):
435
+ return torch.from_numpy(cluster).to(device=edge_index.device, dtype=torch.long)
436
+ return cluster
437
+
438
+
439
+
440
+ def sample_neighbors_disjoint(graph: FastGraph, seed_nodes: Any, num_neighbors: List[int], replace: bool = False):
441
+ r"""Samples a separate multi-hop neighborhood for every seed node, as PyG's ``disjoint=True``.
442
+
443
+ A node reached from two seeds appears twice, once in each seed's subgraph. The sampled edges
444
+ point from sources to the nodes that sampled them ("directional").
445
+
446
+ Returns:
447
+ (nodes, local_row, local_col, edge_ids, batch, num_sampled_nodes, num_sampled_edges), where
448
+ ``batch[i]`` is the seed (subgraph) that local node ``i`` belongs to.
449
+ """
450
+ seeds = np.asarray(seed_nodes).astype(np.int64).reshape(-1)
451
+ num_nodes = graph.num_nodes
452
+ frontier_nodes, frontier_ids = seeds, np.arange(len(seeds), dtype=np.int64)
453
+ node_chunks, batch_chunks = [seeds], [np.arange(len(seeds), dtype=np.int64)]
454
+ keys = np.arange(len(seeds), dtype=np.int64) * num_nodes + seeds # (seed, node) of every local node
455
+ sorted_order = np.argsort(keys)
456
+ num_total = len(seeds)
457
+ rows, cols, eids = [], [], []
458
+ num_sampled_nodes, num_sampled_edges = [len(seeds)], []
459
+ for k in num_neighbors:
460
+ in_graph = frontier_nodes < num_nodes
461
+ starts = np.zeros(len(frontier_nodes), dtype=np.int64)
462
+ counts = np.zeros(len(frontier_nodes), dtype=np.int64)
463
+ starts[in_graph] = graph.indptr[frontier_nodes[in_graph]]
464
+ counts[in_graph] = graph.indptr[frontier_nodes[in_graph] + 1] - starts[in_graph]
465
+ segment, position = _sample_positions(starts, counts, k, replace)
466
+ srcs = graph.sorted_row[position].astype(np.int64)
467
+ dst_ids = frontier_ids[segment]
468
+ src_batch = np.concatenate(batch_chunks)[dst_ids]
469
+ src_keys = src_batch * num_nodes + srcs
470
+
471
+ # Look up which (seed, node) pairs already exist; the rest become new local nodes
472
+ pos = np.searchsorted(keys[sorted_order], src_keys)
473
+ pos = np.minimum(pos, len(keys) - 1)
474
+ known = keys[sorted_order][pos] == src_keys
475
+ src_ids = np.where(known, sorted_order[pos], -1)
476
+ new_keys, first, inverse = np.unique(src_keys[~known], return_index=True, return_inverse=True)
477
+ order = np.argsort(first, kind="stable") # new nodes in order of first appearance
478
+ rank = np.empty(len(order), dtype=np.int64)
479
+ rank[order] = np.arange(len(order))
480
+ src_ids[~known] = num_total + rank[inverse.reshape(-1)]
481
+ new_keys = new_keys[order]
482
+
483
+ rows.append(src_ids)
484
+ cols.append(dst_ids)
485
+ eids.append(graph.sorted_edge_id[position])
486
+ num_sampled_edges.append(len(srcs))
487
+ new_nodes, new_batch = new_keys % num_nodes, new_keys // num_nodes
488
+ node_chunks.append(new_nodes)
489
+ batch_chunks.append(new_batch)
490
+ keys = np.concatenate([keys, new_keys])
491
+ sorted_order = np.argsort(keys, kind="stable")
492
+ frontier_nodes, frontier_ids = new_nodes, num_total + np.arange(len(new_nodes))
493
+ num_total += len(new_nodes)
494
+ num_sampled_nodes.append(len(new_nodes))
495
+
496
+ empty = np.empty(0, dtype=np.int64)
497
+ return (np.concatenate(node_chunks), np.concatenate(rows) if rows else empty,
498
+ np.concatenate(cols) if cols else empty, np.concatenate(eids).astype(np.int64) if eids else empty,
499
+ np.concatenate(batch_chunks), num_sampled_nodes, num_sampled_edges)
@@ -0,0 +1,115 @@
1
+ import copy
2
+ from typing import Any, List, Optional
3
+
4
+ import numpy as np
5
+
6
+ from k3_node.ops.host import to_numpy
7
+
8
+ try:
9
+ import torch
10
+ import torch.utils.data
11
+ from torch import Tensor
12
+ BaseDataLoader = torch.utils.data.DataLoader
13
+ except ImportError:
14
+ torch = None
15
+ Tensor = type(None)
16
+ BaseDataLoader = object
17
+
18
+ from k3_node.data import Batch, Data
19
+ from k3_node.loader.sampler_utils import FastGraph, sample_neighbors_homo
20
+ from k3_node.loader.keras_dataset import loader_bases
21
+
22
+
23
+ class ShaDowKHopSampler(*loader_bases(BaseDataLoader)):
24
+ r"""The ShaDow k-hop sampler from the "Decoupling the Depth and Scope of
25
+ Graph Neural Networks" paper.
26
+
27
+ Args:
28
+ data (Data): The graph data object.
29
+ depth (int): The depth/number of hops of the localized subgraph.
30
+ num_neighbors (int): The number of neighbors to sample for each node in each hop.
31
+ node_idx (LongTensor or BoolTensor, optional): Seed nodes. (default: :obj:`None`)
32
+ replace (bool, optional): Sample neighbors with replacement. (default: :obj:`False`)
33
+ **kwargs (optional): Additional arguments of :class:`torch.utils.data.DataLoader`.
34
+ """
35
+ def __init__(
36
+ self,
37
+ data: Data,
38
+ depth: int,
39
+ num_neighbors: int,
40
+ node_idx: Optional[Any] = None,
41
+ replace: bool = False,
42
+ **kwargs,
43
+ ):
44
+ self.data = copy.copy(data)
45
+ self.depth = depth
46
+ self.num_neighbors = num_neighbors
47
+ self.replace = replace
48
+
49
+ if node_idx is None:
50
+ if torch is not None and isinstance(data.edge_index, Tensor):
51
+ node_idx = torch.arange(data.num_nodes, device=data.edge_index.device)
52
+ else:
53
+ node_idx = np.arange(data.num_nodes)
54
+ elif torch is not None and isinstance(node_idx, Tensor) and node_idx.dtype == torch.bool:
55
+ node_idx = node_idx.nonzero(as_tuple=False).view(-1)
56
+ elif not (torch is not None and isinstance(node_idx, Tensor)):
57
+ node_idx = np.asarray(to_numpy(node_idx)) # NumPy, TF or JAX
58
+ if node_idx.dtype == bool:
59
+ node_idx = np.nonzero(node_idx)[0]
60
+
61
+ self.node_idx = node_idx
62
+ idx_list = node_idx.tolist() if hasattr(node_idx, 'tolist') else list(node_idx)
63
+
64
+ if torch is not None:
65
+ super().__init__(idx_list, collate_fn=self.__collate__, **kwargs)
66
+ else:
67
+ self.dataset = idx_list
68
+ self.collate_fn = self.__collate__
69
+
70
+ def _cached_graph(self):
71
+ # The CSR index depends only on the graph; build it once instead of per batch.
72
+ if getattr(self, "_graph", None) is None:
73
+ self._graph = FastGraph(self.data.edge_index, num_nodes=self.data.num_nodes)
74
+ return self._graph
75
+
76
+ def __collate__(self, n_id: List[int]) -> Batch:
77
+ subgraphs = []
78
+ is_torch = torch is not None and isinstance(self.data.edge_index, Tensor)
79
+ num_neighbors_list = [self.num_neighbors] * self.depth
80
+
81
+ for root in n_id:
82
+ root_seeds = [root]
83
+ if is_torch:
84
+ root_seeds = torch.tensor(root_seeds, dtype=torch.long, device=self.data.edge_index.device)
85
+
86
+ nodes, row, col, edges, _, _ = sample_neighbors_homo(
87
+ edge_index=self.data.edge_index,
88
+ seed_nodes=root_seeds,
89
+ num_neighbors=num_neighbors_list,
90
+ num_nodes=self.data.num_nodes,
91
+ graph=self._cached_graph(),
92
+ replace=self.replace,
93
+ subgraph_type='directional',
94
+ )
95
+
96
+ sub = self.data.subgraph(nodes)
97
+ sub.root_n_id = 0 # Root node is always the first node (seed node)
98
+ subgraphs.append(sub)
99
+
100
+ batch = Batch.from_data_list(subgraphs)
101
+ # Each root is the first node of its subgraph; as in PyG, `root_n_id` holds the roots'
102
+ # positions in the batch and `y` the labels of the roots (one per subgraph).
103
+ ptr = np.asarray(to_numpy(batch.ptr)).astype(np.int64)
104
+ roots = ptr[:-1]
105
+ if is_torch:
106
+ batch.root_n_id = torch.as_tensor(roots, dtype=torch.long, device=self.data.edge_index.device)
107
+ else:
108
+ batch.root_n_id = roots
109
+ y = getattr(self.data, "y", None)
110
+ if y is not None and y.shape[0] == self.data.num_nodes:
111
+ if torch is not None and isinstance(y, Tensor):
112
+ batch.y = y[torch.as_tensor(list(n_id), dtype=torch.long, device=y.device)]
113
+ else:
114
+ batch.y = np.asarray(y)[np.asarray(list(n_id))]
115
+ return batch