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,98 @@
1
+ from typing import List
2
+
3
+ import numpy as np
4
+
5
+ try:
6
+ import torch
7
+ import torch.utils.data
8
+ BaseDataLoader = torch.utils.data.DataLoader
9
+ except ImportError:
10
+ torch = None
11
+ BaseDataLoader = object
12
+
13
+ from k3_node.data import TemporalData
14
+
15
+
16
+ class TemporalDataLoader(BaseDataLoader):
17
+ r"""A data loader which merges successive events of a
18
+ :class:`k3_node.data.TemporalData` to a mini-batch.
19
+
20
+ Args:
21
+ data (TemporalData): The :obj:`~k3_node.data.TemporalData` from which to load.
22
+ batch_size (int, optional): How many samples per batch to load. (default: :obj:`1`)
23
+ neg_sampling_ratio (float, optional): The ratio of sampled negative
24
+ destination nodes to the number of positive destination nodes. (default: :obj:`0.0`)
25
+ **kwargs (optional): Additional arguments of :class:`torch.utils.data.DataLoader`.
26
+ """
27
+ def __init__(
28
+ self,
29
+ data: TemporalData,
30
+ batch_size: int = 1,
31
+ neg_sampling_ratio: float = 0.0,
32
+ **kwargs,
33
+ ):
34
+ kwargs.pop('dataset', None)
35
+ kwargs.pop('collate_fn', None)
36
+ kwargs.pop('shuffle', None)
37
+
38
+ self.data = data
39
+ self.events_per_batch = batch_size
40
+ self.neg_sampling_ratio = neg_sampling_ratio
41
+
42
+ if neg_sampling_ratio > 0:
43
+ dst = data.dst
44
+ if torch is not None and isinstance(dst, torch.Tensor):
45
+ self.min_dst = int(dst.min())
46
+ self.max_dst = int(dst.max())
47
+ else:
48
+ self.min_dst = int(np.min(np.asarray(dst)))
49
+ self.max_dst = int(np.max(np.asarray(dst)))
50
+
51
+ if kwargs.get('drop_last', False) and len(data) % batch_size != 0:
52
+ arange = list(range(0, len(data) - batch_size, batch_size))
53
+ else:
54
+ arange = list(range(0, len(data), batch_size))
55
+
56
+ if torch is not None:
57
+ super().__init__(arange, 1, shuffle=False, collate_fn=self, **kwargs)
58
+ else:
59
+ self.dataset = arange
60
+ self.batch_size = 1
61
+ self.shuffle = False
62
+ self.collate_fn = self
63
+
64
+ def __call__(self, arange: List[int]) -> TemporalData:
65
+ start = arange[0]
66
+ end = start + self.events_per_batch
67
+ batch = self.data[start:end]
68
+
69
+ is_torch = torch is not None and isinstance(batch.dst, torch.Tensor)
70
+
71
+ n_ids = [batch.src, batch.dst]
72
+
73
+ if self.neg_sampling_ratio > 0:
74
+ num_neg = round(self.neg_sampling_ratio * (batch.dst.size(0) if is_torch else len(batch.dst)))
75
+ if is_torch:
76
+ batch.neg_dst = torch.randint(
77
+ low=self.min_dst,
78
+ high=self.max_dst + 1,
79
+ size=(num_neg,),
80
+ dtype=batch.dst.dtype,
81
+ device=batch.dst.device,
82
+ )
83
+ else:
84
+ batch.neg_dst = np.random.randint(
85
+ low=self.min_dst,
86
+ high=self.max_dst + 1,
87
+ size=(num_neg,),
88
+ dtype=np.int64,
89
+ )
90
+ n_ids.append(batch.neg_dst)
91
+
92
+ if is_torch:
93
+ batch.n_id = torch.cat(n_ids, dim=0).unique()
94
+ else:
95
+ batch.n_id = np.unique(np.concatenate([np.asarray(x) for x in n_ids], axis=0))
96
+
97
+ return batch
98
+
@@ -0,0 +1,113 @@
1
+ import numpy as np
2
+ import pytest
3
+
4
+ try:
5
+ import torch
6
+ except ImportError:
7
+ torch = None
8
+
9
+ from k3_node.data import Data, HeteroData, TemporalData
10
+ from k3_node.loader import (
11
+ DataListLoader,
12
+ DataLoader,
13
+ DenseDataLoader,
14
+ RandomNodeLoader,
15
+ TemporalDataLoader,
16
+ )
17
+
18
+
19
+ def create_dummy_data(num_nodes=4, num_features=8):
20
+ x = np.random.randn(num_nodes, num_features).astype(np.float32)
21
+ edge_index = np.array([[0, 1, 2, 3], [1, 2, 3, 0]], dtype=np.int64)
22
+ y = np.array([0, 1, 0, 1], dtype=np.int64)
23
+ if torch is not None:
24
+ x = torch.from_numpy(x)
25
+ edge_index = torch.from_numpy(edge_index)
26
+ y = torch.from_numpy(y)
27
+ return Data(x=x, edge_index=edge_index, y=y)
28
+
29
+
30
+ def test_data_loader_basic():
31
+ dataset = [create_dummy_data(num_nodes=i + 2) for i in range(4)]
32
+ loader = DataLoader(dataset, batch_size=2, shuffle=False)
33
+
34
+ batches = list(loader)
35
+ assert len(batches) == 2
36
+
37
+ batch0 = batches[0]
38
+ assert batch0.num_graphs == 2
39
+ assert hasattr(batch0, 'batch')
40
+ assert hasattr(batch0, 'ptr')
41
+ assert batch0.num_nodes == dataset[0].num_nodes + dataset[1].num_nodes
42
+
43
+
44
+ def test_data_loader_follow_batch_and_exclude():
45
+ dataset = [create_dummy_data(num_nodes=3) for _ in range(3)]
46
+ loader = DataLoader(dataset, batch_size=2, follow_batch=['y'], exclude_keys=['edge_index'])
47
+
48
+ batch = next(iter(loader))
49
+ assert 'edge_index' not in batch
50
+ assert hasattr(batch, 'y_batch')
51
+
52
+
53
+ def test_data_list_loader():
54
+ dataset = [create_dummy_data(num_nodes=3) for _ in range(4)]
55
+ loader = DataListLoader(dataset, batch_size=2, shuffle=False)
56
+
57
+ batches = list(loader)
58
+ assert len(batches) == 2
59
+ assert isinstance(batches[0], list)
60
+ assert len(batches[0]) == 2
61
+ assert isinstance(batches[0][0], Data)
62
+
63
+
64
+ def test_dense_data_loader():
65
+ def create_dense_graph(num_nodes=4, num_features=6):
66
+ x = np.random.randn(num_nodes, num_features).astype(np.float32)
67
+ adj = np.random.randn(num_nodes, num_nodes).astype(np.float32)
68
+ if torch is not None:
69
+ x = torch.from_numpy(x)
70
+ adj = torch.from_numpy(adj)
71
+ return Data(x=x, adj=adj)
72
+
73
+ dataset = [create_dense_graph() for _ in range(4)]
74
+ loader = DenseDataLoader(dataset, batch_size=2, shuffle=False)
75
+
76
+ batch = next(iter(loader))
77
+ assert batch.x.shape == (2, 4, 6)
78
+ assert batch.adj.shape == (2, 4, 4)
79
+
80
+
81
+ def test_temporal_data_loader():
82
+ src = np.array([0, 1, 0, 2, 1, 3], dtype=np.int64)
83
+ dst = np.array([1, 2, 2, 3, 3, 0], dtype=np.int64)
84
+ t = np.array([1, 2, 3, 4, 5, 6], dtype=np.int64)
85
+ msg = np.random.randn(6, 4).astype(np.float32)
86
+
87
+ if torch is not None:
88
+ src = torch.from_numpy(src)
89
+ dst = torch.from_numpy(dst)
90
+ t = torch.from_numpy(t)
91
+ msg = torch.from_numpy(msg)
92
+
93
+ data = TemporalData(src=src, dst=dst, t=t, msg=msg)
94
+ loader = TemporalDataLoader(data, batch_size=3, neg_sampling_ratio=1.0)
95
+
96
+ batches = list(loader)
97
+ assert len(batches) == 2
98
+ assert len(batches[0]) == 3
99
+ assert hasattr(batches[0], 'neg_dst')
100
+ assert hasattr(batches[0], 'n_id')
101
+
102
+
103
+ def test_random_node_loader():
104
+ data = create_dummy_data(num_nodes=10)
105
+ loader = RandomNodeLoader(data, num_parts=2)
106
+
107
+ parts = list(loader)
108
+ assert len(parts) == 2
109
+ for part in parts:
110
+ assert isinstance(part, Data)
111
+ assert part.num_nodes <= 10
112
+ assert hasattr(part, 'x')
113
+
@@ -0,0 +1,221 @@
1
+ import numpy as np
2
+ import keras
3
+ from keras import ops
4
+
5
+ from k3_node.data import Data
6
+ from k3_node.layers import GCNConv
7
+ from k3_node.loader import DataLoader, FullGraphDataset, NeighborLoader
8
+ from k3_node.loader.keras_dataset import to_keras_batch
9
+
10
+
11
+ def _graph(num_nodes=30, num_classes=3):
12
+ rng = np.random.default_rng(0)
13
+ y = rng.integers(0, num_classes, num_nodes)
14
+ x = (np.eye(num_classes)[y] + 0.3 * rng.standard_normal((num_nodes, num_classes))).astype("float32")
15
+ edge_index = rng.integers(0, num_nodes, (2, 90)).astype("int32")
16
+ train_mask = np.zeros(num_nodes, dtype=bool)
17
+ train_mask[:10] = True
18
+ return Data(x=x, edge_index=edge_index, y=y.astype("int32"), train_mask=train_mask, test_mask=~train_mask)
19
+
20
+
21
+ class SmallGCN(keras.Model):
22
+ def __init__(self):
23
+ super().__init__()
24
+ self.conv = GCNConv(3, 3)
25
+
26
+ def call(self, data):
27
+ return self.conv(data.x, data.edge_index)
28
+
29
+
30
+ def test_fit_and_evaluate_on_full_graph():
31
+ data = _graph()
32
+ model = SmallGCN()
33
+ model.compile(keras.optimizers.Adam(0.05), keras.losses.SparseCategoricalCrossentropy(from_logits=True))
34
+ history = model.fit(FullGraphDataset(data, mask="train_mask"), epochs=30, verbose=0)
35
+ assert history.history["loss"][-1] < history.history["loss"][0]
36
+ model.evaluate(FullGraphDataset(data, mask="test_mask"), verbose=0)
37
+ preds = model.predict(FullGraphDataset(data), verbose=0)
38
+ assert preds.shape == (30, 3)
39
+
40
+
41
+ def test_normalized_mask_gives_mean_over_masked_nodes():
42
+ data = _graph()
43
+ model = SmallGCN()
44
+ loss_fn = keras.losses.SparseCategoricalCrossentropy(from_logits=True, reduction=None)
45
+ model.compile(loss=keras.losses.SparseCategoricalCrossentropy(from_logits=True))
46
+ reported = model.evaluate(FullGraphDataset(data, mask="train_mask"), verbose=0)
47
+ per_node = ops.convert_to_numpy(loss_fn(data.y, model(to_keras_batch(data)[0])))
48
+ np.testing.assert_allclose(reported, per_node[: 10].mean(), rtol=1e-5)
49
+
50
+
51
+ def test_graph_loader_works_with_fit_evaluate_predict():
52
+ from k3_node.datasets import FakeDataset
53
+ from k3_node.layers import GCNConv, global_mean_pool
54
+
55
+ class GraphClassifier(keras.Model):
56
+ def __init__(self):
57
+ super().__init__()
58
+ self.conv = GCNConv(8, 16)
59
+ self.head = keras.layers.Dense(3)
60
+
61
+ def call(self, data):
62
+ x = self.conv(data.x, data.edge_index)
63
+ return self.head(global_mean_pool(x, data.batch, data.num_graphs))
64
+
65
+ dataset = FakeDataset(num_graphs=20, avg_num_nodes=8, num_channels=8, num_classes=3)
66
+ loader = DataLoader(dataset, batch_size=6, shuffle=True)
67
+ model = GraphClassifier()
68
+ model.compile(optimizer="adam", loss=keras.losses.SparseCategoricalCrossentropy(from_logits=True), metrics=["accuracy"])
69
+ history = model.fit(loader, epochs=2, verbose=0) # compiled on TF/JAX: num_graphs must be static
70
+ assert len(history.history["loss"]) == 2
71
+ model.evaluate(DataLoader(dataset, batch_size=6), verbose=0)
72
+ assert model.predict(DataLoader(dataset, batch_size=6), verbose=0).shape == (20, 3)
73
+ # Iterating the loader directly still yields Batch objects
74
+ batch = next(iter(loader))
75
+ assert hasattr(batch, "edge_index") and batch.num_graphs == 6
76
+
77
+
78
+ def test_neighbor_loader_weights_seed_nodes_only():
79
+ data = _graph()
80
+ loader = NeighborLoader(data, num_neighbors=[3], batch_size=4, input_nodes=np.arange(10))
81
+ inputs, y, weight = loader[0]
82
+ assert y.shape[0] == inputs.x.shape[0]
83
+ assert (weight[:4] > 0).all() and (weight[4:] == 0).all()
84
+
85
+
86
+ def test_full_graph_dataset_index_split():
87
+ import numpy as np
88
+ from k3_node.data import Data
89
+ from k3_node.loader import FullGraphDataset
90
+
91
+ data = Data(x=np.ones((5, 2), "float32"), edge_index=np.array([[0, 1], [1, 2]]),
92
+ train_idx=np.array([3, 1]), train_y=np.array([2, 1]), num_nodes=5)
93
+ inputs, y, weight = FullGraphDataset(data, index="train_idx", target="train_y")[0]
94
+ assert "train_idx" not in inputs._fields and "train_y" not in inputs._fields
95
+ np.testing.assert_array_equal(y, [0, 1, 0, 2, 0])
96
+ np.testing.assert_allclose(weight, [0, 2.5, 0, 2.5, 0]) # mean over the 2 indexed nodes
97
+
98
+
99
+ def test_full_graph_dataset_fresh_negatives_every_epoch():
100
+ import numpy as np
101
+ from k3_node.data import Data
102
+ from k3_node.loader import FullGraphDataset
103
+
104
+ edge_index = np.array([[0, 1, 2, 3], [1, 2, 3, 4]])
105
+ data = Data(x=np.ones((50, 2), "float32"), edge_index=edge_index, edge_label_index=edge_index,
106
+ edge_label=np.ones(4, "float32"), num_nodes=50)
107
+ dataset = FullGraphDataset(data, neg_sampling_ratio=2.0)
108
+ (inputs1, y1), (inputs2, _) = dataset[0], dataset[0]
109
+ assert inputs1.edge_label_index.shape == (2, 12)
110
+ np.testing.assert_array_equal(y1, [1] * 4 + [0] * 8)
111
+ assert not np.array_equal(inputs1.edge_label_index[:, 4:], inputs2.edge_label_index[:, 4:])
112
+
113
+
114
+ def test_link_neighbor_loader_uses_local_ids():
115
+ import numpy as np
116
+ from k3_node.data import Data
117
+ from k3_node.loader import LinkNeighborLoader
118
+
119
+ rng = np.random.default_rng(0)
120
+ edge_index = rng.integers(0, 200, size=(2, 600))
121
+ data = Data(x=np.arange(200, dtype="float32")[:, None], edge_index=edge_index, num_nodes=200)
122
+ loader = LinkNeighborLoader(data, num_neighbors=[5], batch_size=16, neg_sampling_ratio=1.0)
123
+ batch = next(iter(loader))
124
+ eli = np.asarray(batch.edge_label_index)
125
+ n_id = np.asarray(batch.n_id)
126
+ assert eli.max() < n_id.shape[0] # local ids into the sampled subgraph
127
+ label = np.asarray(batch.edge_label)
128
+ pos = n_id[eli[:, label == 1]] # back to global ids: must be real edges
129
+ real = set(map(tuple, edge_index.T.tolist()))
130
+ assert all(tuple(e) in real for e in pos.T.tolist())
131
+
132
+
133
+ def test_loader_with_mask():
134
+ import numpy as np
135
+ from k3_node.data import Data
136
+ from k3_node.loader import ClusterData, ClusterLoader
137
+
138
+ rng = np.random.default_rng(0)
139
+ data = Data(x=rng.random((60, 3)).astype("float32"), edge_index=rng.integers(0, 60, (2, 200)),
140
+ y=rng.integers(0, 3, 60), train_mask=np.arange(60) < 30, num_nodes=60)
141
+ loader = ClusterLoader(ClusterData(data, num_parts=4), batch_size=2).with_mask("train_mask")
142
+ inputs, y, weight = loader[0]
143
+ assert weight.shape == y.shape and "train_mask" not in inputs._fields
144
+
145
+
146
+ def test_shadow_roots_and_labels():
147
+ import numpy as np
148
+ from k3_node.data import Data
149
+ from k3_node.loader import ShaDowKHopSampler
150
+
151
+ rng = np.random.default_rng(0)
152
+ data = Data(x=np.arange(50, dtype="float32")[:, None], edge_index=rng.integers(0, 50, (2, 200)),
153
+ y=np.arange(50), num_nodes=50)
154
+ loader = ShaDowKHopSampler(data, depth=2, num_neighbors=3, node_idx=np.arange(10, 20), batch_size=5)
155
+ batch = next(iter(loader))
156
+ root_x = np.asarray(batch.x)[np.asarray(batch.root_n_id), 0]
157
+ np.testing.assert_array_equal(root_x, np.arange(10, 15)) # the roots are the seed nodes
158
+ np.testing.assert_array_equal(np.asarray(batch.y), np.arange(10, 15)) # one label per subgraph
159
+
160
+
161
+ def test_graph_saint_samplers():
162
+ import numpy as np
163
+ from k3_node.data import Data
164
+ from k3_node.loader import GraphSAINTEdgeSampler, GraphSAINTNodeSampler, GraphSAINTRandomWalkSampler
165
+
166
+ rng = np.random.default_rng(0)
167
+ edge_index = rng.integers(0, 100, (2, 400))
168
+ data = Data(x=np.arange(100, dtype="float32")[:, None], edge_index=edge_index,
169
+ edge_attr=np.arange(400, dtype="float32"), y=np.arange(100), num_nodes=100)
170
+ for loader in [GraphSAINTNodeSampler(data, batch_size=30, num_steps=3, sample_coverage=5),
171
+ GraphSAINTEdgeSampler(data, batch_size=20, num_steps=3, sample_coverage=5),
172
+ GraphSAINTRandomWalkSampler(data, batch_size=10, walk_length=2, num_steps=3, sample_coverage=5)]:
173
+ batches = list(loader)
174
+ assert len(batches) == 3
175
+ batch = batches[0]
176
+ x = np.asarray(batch.x)[:, 0].astype(int)
177
+ ei = np.asarray(batch.edge_index)
178
+ # every subgraph edge is a real edge between sampled nodes, carrying its own attribute
179
+ real = {tuple(e): i for i, e in enumerate(edge_index.T.tolist())}
180
+ for (a, b), attr in zip(ei.T.tolist(), np.asarray(batch.edge_attr).astype(int)):
181
+ assert (x[a], x[b]) in real and tuple(edge_index[:, attr]) == (x[a], x[b])
182
+ assert np.asarray(batch.node_norm).shape == (batch.num_nodes,)
183
+ assert np.asarray(batch.edge_norm).shape == (ei.shape[1],)
184
+ inputs, y = loader[0] # Keras batch
185
+ assert inputs.x.shape[0] == y.shape[0]
186
+
187
+
188
+ def test_neighbor_loader_disjoint():
189
+ import numpy as np
190
+ from k3_node.data import Data
191
+ from k3_node.loader import NeighborLoader
192
+
193
+ rng = np.random.default_rng(0)
194
+ edge_index = rng.integers(0, 30, (2, 150))
195
+ data = Data(x=np.arange(30, dtype="float32")[:, None], edge_index=edge_index, num_nodes=30)
196
+ batch = next(iter(NeighborLoader(data, num_neighbors=[3, 2], batch_size=8, disjoint=True)))
197
+ b, n_id, ei = np.asarray(batch.batch), np.asarray(batch.n_id), np.asarray(batch.edge_index)
198
+ np.testing.assert_array_equal(b[:8], np.arange(8)) # seeds first, one subgraph each
199
+ assert np.all(b[ei[0]] == b[ei[1]]) # edges never cross subgraphs
200
+ real = set(map(tuple, edge_index.T.tolist()))
201
+ assert all((n_id[s], n_id[t]) in real for s, t in ei.T.tolist())
202
+ for g in range(8): # no node appears twice within one subgraph
203
+ assert len(set(n_id[b == g].tolist())) == int((b == g).sum())
204
+
205
+
206
+ def test_full_graph_dataset_hetero():
207
+ import numpy as np
208
+ from k3_node.data import HeteroData
209
+ from k3_node.loader import FullGraphDataset
210
+
211
+ data = HeteroData()
212
+ data["user"].x = np.ones((4, 2), "float32")
213
+ data["user"].y = np.array([0, 1, 0, 1])
214
+ data["user"].train_mask = np.array([True, True, False, False])
215
+ data["item"].x = np.ones((3, 5), "float32")
216
+ data["user", "buys", "item"].edge_index = np.array([[0, 1, 3], [0, 2, 1]])
217
+ inputs, y, weight = FullGraphDataset(data, node_type="user", mask="train_mask")[0]
218
+ assert set(inputs.x_dict) == {"user", "item"}
219
+ assert ("user", "buys", "item") in inputs.edge_index_dict
220
+ np.testing.assert_array_equal(y, [0, 1, 0, 1])
221
+ np.testing.assert_allclose(weight, [2, 2, 0, 0])
@@ -0,0 +1,122 @@
1
+ import numpy as np
2
+
3
+ try:
4
+ import torch
5
+ except ImportError:
6
+ torch = None
7
+
8
+ from k3_node.data import Data, HeteroData
9
+ from k3_node.loader import HGTLoader, LinkNeighborLoader, NeighborLoader
10
+
11
+
12
+ def get_homo_graph():
13
+ # 6 nodes connected in a ring: 0->1->2->3->4->5->0 and some cross edges
14
+ edge_index = np.array([
15
+ [0, 1, 2, 3, 4, 5, 0, 2],
16
+ [1, 2, 3, 4, 5, 0, 3, 5],
17
+ ], dtype=np.int64)
18
+ x = np.random.randn(6, 16).astype(np.float32)
19
+ y = np.array([0, 1, 0, 1, 0, 1], dtype=np.int64)
20
+
21
+ if torch is not None:
22
+ edge_index = torch.from_numpy(edge_index)
23
+ x = torch.from_numpy(x)
24
+ y = torch.from_numpy(y)
25
+
26
+ return Data(x=x, edge_index=edge_index, y=y)
27
+
28
+
29
+ def get_hetero_graph():
30
+ data = HeteroData()
31
+ data['paper'].x = np.random.randn(10, 8).astype(np.float32)
32
+ data['author'].x = np.random.randn(5, 8).astype(np.float32)
33
+
34
+ data['author', 'writes', 'paper'].edge_index = np.array([
35
+ [0, 1, 2, 3, 4, 0, 1],
36
+ [0, 1, 2, 3, 4, 5, 6],
37
+ ], dtype=np.int64)
38
+ data['paper', 'cites', 'paper'].edge_index = np.array([
39
+ [0, 1, 2, 3],
40
+ [1, 2, 3, 4],
41
+ ], dtype=np.int64)
42
+
43
+ if torch is not None:
44
+ data['paper'].x = torch.from_numpy(data['paper'].x)
45
+ data['author'].x = torch.from_numpy(data['author'].x)
46
+ data['author', 'writes', 'paper'].edge_index = torch.from_numpy(data['author', 'writes', 'paper'].edge_index)
47
+ data['paper', 'cites', 'paper'].edge_index = torch.from_numpy(data['paper', 'cites', 'paper'].edge_index)
48
+
49
+ return data
50
+
51
+
52
+ def test_homo_neighbor_loader():
53
+ data = get_homo_graph()
54
+ input_nodes = [0, 1]
55
+ if torch is not None:
56
+ input_nodes = torch.tensor(input_nodes, dtype=torch.long)
57
+
58
+ loader = NeighborLoader(
59
+ data,
60
+ num_neighbors=[2, 2],
61
+ batch_size=2,
62
+ input_nodes=input_nodes,
63
+ shuffle=False,
64
+ )
65
+
66
+ batch = next(iter(loader))
67
+ assert batch.batch_size == 2
68
+ assert hasattr(batch, 'n_id')
69
+ assert hasattr(batch, 'e_id')
70
+ assert hasattr(batch, 'num_sampled_nodes')
71
+ assert hasattr(batch, 'num_sampled_edges')
72
+ assert batch.x.shape[0] == len(batch.n_id)
73
+ assert batch.edge_index.shape[0] == 2
74
+
75
+
76
+ def test_hetero_neighbor_loader():
77
+ data = get_hetero_graph()
78
+ loader = NeighborLoader(
79
+ data,
80
+ num_neighbors=[2, 2],
81
+ batch_size=2,
82
+ input_nodes=('paper', [0, 1]),
83
+ shuffle=False,
84
+ )
85
+
86
+ batch = next(iter(loader))
87
+ assert batch['paper'].batch_size == 2
88
+ assert hasattr(batch['paper'], 'n_id')
89
+ assert hasattr(batch['author', 'writes', 'paper'], 'edge_index')
90
+
91
+
92
+ def test_link_neighbor_loader():
93
+ data = get_homo_graph()
94
+ loader = LinkNeighborLoader(
95
+ data,
96
+ num_neighbors=[2, 2],
97
+ batch_size=2,
98
+ neg_sampling_ratio=1.0,
99
+ shuffle=False,
100
+ )
101
+
102
+ batch = next(iter(loader))
103
+ assert hasattr(batch, 'edge_label_index')
104
+ assert hasattr(batch, 'edge_label')
105
+ assert hasattr(batch, 'n_id')
106
+ assert batch.edge_label.shape[0] == 4 # 2 positive + 2 negative
107
+
108
+
109
+ def test_hgt_loader():
110
+ data = get_hetero_graph()
111
+ loader = HGTLoader(
112
+ data,
113
+ num_samples=[4, 4],
114
+ input_nodes=('paper', [0, 1]),
115
+ batch_size=2,
116
+ shuffle=False,
117
+ )
118
+
119
+ batch = next(iter(loader))
120
+ assert batch['paper'].batch_size == 2
121
+ assert hasattr(batch['paper'], 'n_id')
122
+
@@ -0,0 +1,82 @@
1
+ import numpy as np
2
+ import pytest
3
+
4
+ from k3_node.loader.sampler_utils import FastGraph, sample_neighbors_homo
5
+
6
+
7
+ def _reference_full_sampling(edge_index, seeds, num_hops, subgraph_type):
8
+ """The original per-node Python algorithm, restricted to k=-1 where it is deterministic."""
9
+ graph = FastGraph(edge_index)
10
+ nodes, visited = [], {}
11
+ for s in seeds:
12
+ if int(s) not in visited:
13
+ visited[int(s)] = len(nodes)
14
+ nodes.append(int(s))
15
+ frontier, sampled = list(nodes), []
16
+ for _ in range(num_hops):
17
+ next_frontier = []
18
+ for target in frontier:
19
+ srcs, e_ids = graph.get_neighbors(target)
20
+ for s, e in zip(srcs, e_ids):
21
+ sampled.append((int(s), target, int(e)))
22
+ if int(s) not in visited:
23
+ visited[int(s)] = len(nodes)
24
+ nodes.append(int(s))
25
+ next_frontier.append(int(s))
26
+ frontier = next_frontier
27
+ if subgraph_type == "induced":
28
+ edges = [(visited[u], visited[v], e) for e, (u, v) in enumerate(edge_index.T) if u in visited and v in visited]
29
+ elif subgraph_type == "bidirectional":
30
+ edges = [x for u, v, e in sampled for x in ((visited[u], visited[v], e), (visited[v], visited[u], e))]
31
+ else:
32
+ edges = [(visited[u], visited[v], e) for u, v, e in sampled]
33
+ edges = np.array(edges, dtype=np.int64).reshape(-1, 3)
34
+ return np.array(nodes), edges[:, 0], edges[:, 1], edges[:, 2]
35
+
36
+
37
+ def _random_graph(num_nodes=40, num_edges=160, seed=0):
38
+ rng = np.random.default_rng(seed)
39
+ return rng.integers(0, num_nodes, (2, num_edges)).astype(np.int64)
40
+
41
+
42
+ @pytest.mark.parametrize("subgraph_type", ["directional", "bidirectional", "induced"])
43
+ def test_full_neighborhood_matches_reference(subgraph_type):
44
+ edge_index = _random_graph()
45
+ seeds = np.array([3, 7, 3, 11])
46
+ expected = _reference_full_sampling(edge_index, seeds, 2, subgraph_type)
47
+ nodes, row, col, edge, _, _ = sample_neighbors_homo(edge_index, seeds, [-1, -1], subgraph_type=subgraph_type)
48
+ for got, want in zip((nodes, row, col, edge), expected):
49
+ np.testing.assert_array_equal(got, want)
50
+
51
+
52
+ @pytest.mark.parametrize("replace", [False, True])
53
+ def test_sampled_neighborhood_properties(replace):
54
+ edge_index = _random_graph()
55
+ graph = FastGraph(edge_index)
56
+ seeds = np.array([0, 5, 9])
57
+ nodes, row, col, edge, n_counts, e_counts = sample_neighbors_homo(
58
+ edge_index, seeds, [3, 2], replace=replace, graph=graph
59
+ )
60
+ assert len(np.unique(nodes)) == len(nodes)
61
+ np.testing.assert_array_equal(nodes[:3], seeds)
62
+ assert sum(n_counts) == len(nodes) and sum(e_counts) == len(edge)
63
+ # Every sampled edge is a real edge, with endpoints mapped to their local ids.
64
+ np.testing.assert_array_equal(nodes[row], edge_index[0, edge])
65
+ np.testing.assert_array_equal(nodes[col], edge_index[1, edge])
66
+ # First hop: each seed gets min(3, in-degree) edges, distinct when sampling without replacement.
67
+ in_degree = np.bincount(edge_index[1], minlength=40)
68
+ first_hop = slice(0, e_counts[0])
69
+ for i, s in enumerate(seeds):
70
+ picked = edge[first_hop][col[first_hop] == i]
71
+ assert len(picked) == min(3, in_degree[s])
72
+ if not replace:
73
+ assert len(np.unique(picked)) == len(picked)
74
+ # The reusable id map is left clean for the next batch.
75
+ assert (graph.local_map(40) == -1).all()
76
+
77
+
78
+ def test_disjoint_and_out_of_range_seeds():
79
+ edge_index = np.array([[1, 2], [0, 0]])
80
+ nodes, row, col, edge, _, _ = sample_neighbors_homo(edge_index, np.array([0, 5]), [-1], num_nodes=3, disjoint=True)
81
+ np.testing.assert_array_equal(nodes, [0, 5, 1, 2])
82
+ np.testing.assert_array_equal(edge, [0, 1])