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,96 @@
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
10
+ from k3_node.loader import (
11
+ CachedLoader,
12
+ DataLoader,
13
+ DynamicBatchSampler,
14
+ ImbalancedSampler,
15
+ PrefetchLoader,
16
+ ZipLoader,
17
+ )
18
+
19
+
20
+ def create_dummy_data(num_nodes=4, num_features=8, label=0):
21
+ x = np.random.randn(num_nodes, num_features).astype(np.float32)
22
+ edge_index = np.array([[0, 1, 2, 3], [1, 2, 3, 0]], dtype=np.int64)
23
+ y = np.array([label], dtype=np.int64)
24
+ if torch is not None:
25
+ x = torch.from_numpy(x)
26
+ edge_index = torch.from_numpy(edge_index)
27
+ y = torch.from_numpy(y)
28
+ return Data(x=x, edge_index=edge_index, y=y)
29
+
30
+
31
+ def test_dynamic_batch_sampler():
32
+ dataset = [create_dummy_data(num_nodes=i + 2) for i in range(5)]
33
+ sampler = DynamicBatchSampler(dataset, max_num=15, mode='node')
34
+
35
+ loader = DataLoader(dataset, batch_sampler=sampler)
36
+ count = 0
37
+ for batch in loader:
38
+ assert batch.num_nodes <= 15
39
+ count += 1
40
+ assert count > 0
41
+
42
+ with pytest.raises(ValueError, match="length of 'DynamicBatchSampler'"):
43
+ len(sampler)
44
+ assert len(DynamicBatchSampler(dataset, max_num=15, num_steps=2)) == 2
45
+
46
+
47
+ def test_imbalanced_sampler():
48
+ # 8 samples with label 0, 2 samples with label 1
49
+ dataset = [create_dummy_data(label=0) for _ in range(8)] + [create_dummy_data(label=1) for _ in range(2)]
50
+ sampler = ImbalancedSampler(dataset, num_samples=10)
51
+
52
+ sampled_indices = list(sampler)
53
+ assert len(sampled_indices) == 10
54
+
55
+ loader = DataLoader(dataset, batch_size=5, sampler=sampler)
56
+ batches = list(loader)
57
+ assert len(batches) == 2
58
+
59
+
60
+ def test_zip_loader():
61
+ from k3_node.loader import NeighborLoader
62
+ data = create_dummy_data(num_nodes=10)
63
+
64
+ loader1 = NeighborLoader(data, num_neighbors=[2], input_nodes=[0, 1, 2, 3])
65
+ loader2 = NeighborLoader(data, num_neighbors=[2], input_nodes=[4, 5, 6, 7])
66
+
67
+ zip_loader = ZipLoader([loader1, loader2], batch_size=2)
68
+ for batch1, batch2 in zip_loader:
69
+ assert batch1.batch_size == 2
70
+ assert batch2.batch_size == 2
71
+
72
+
73
+ def test_cached_loader():
74
+ dataset = [create_dummy_data(num_nodes=3) for _ in range(4)]
75
+ loader = DataLoader(dataset, batch_size=2)
76
+
77
+ cached_loader = CachedLoader(loader)
78
+ epoch1 = list(cached_loader)
79
+ epoch2 = list(cached_loader)
80
+
81
+ assert len(epoch1) == 2
82
+ assert len(epoch2) == 2
83
+ assert len(cached_loader) == 2
84
+
85
+ cached_loader.clear()
86
+ assert len(cached_loader._cache) == 0
87
+
88
+
89
+ def test_prefetch_loader():
90
+ dataset = [create_dummy_data(num_nodes=3) for _ in range(4)]
91
+ loader = DataLoader(dataset, batch_size=2)
92
+
93
+ prefetch = PrefetchLoader(loader)
94
+ batches = list(prefetch)
95
+ assert len(batches) == 2
96
+ assert len(prefetch) == 2
@@ -0,0 +1,89 @@
1
+ import numpy as np
2
+ try:
3
+ import torch
4
+ except ImportError:
5
+ torch = None
6
+
7
+ from k3_node.data import Data
8
+ from k3_node.loader import (
9
+ ClusterData,
10
+ ClusterLoader,
11
+ GraphSAINTEdgeSampler,
12
+ GraphSAINTNodeSampler,
13
+ GraphSAINTRandomWalkSampler,
14
+ NeighborSampler,
15
+ ShaDowKHopSampler,
16
+ )
17
+
18
+
19
+ def get_graph(num_nodes=12):
20
+ # Circular ring graph
21
+ row = np.arange(num_nodes)
22
+ col = (row + 1) % num_nodes
23
+ edge_index = np.stack([row, col], axis=0).astype(np.int64)
24
+ x = np.random.randn(num_nodes, 8).astype(np.float32)
25
+ y = np.random.randint(0, 2, size=(num_nodes,)).astype(np.int64)
26
+
27
+ if torch is not None:
28
+ edge_index = torch.from_numpy(edge_index)
29
+ x = torch.from_numpy(x)
30
+ y = torch.from_numpy(y)
31
+
32
+ return Data(x=x, edge_index=edge_index, y=y)
33
+
34
+
35
+ def test_cluster_loader():
36
+ data = get_graph(num_nodes=12)
37
+ cluster_data = ClusterData(data, num_parts=3)
38
+ assert len(cluster_data) == 3
39
+
40
+ sub0 = cluster_data[0]
41
+ assert isinstance(sub0, Data)
42
+ assert sub0.num_nodes > 0
43
+
44
+ loader = ClusterLoader(cluster_data, batch_size=2, shuffle=False)
45
+ batches = list(loader)
46
+ assert len(batches) == 2 # 3 parts with batch_size 2 => 2 batches
47
+
48
+
49
+ def test_graph_saint():
50
+ data = get_graph(num_nodes=12)
51
+
52
+ # Node sampler
53
+ node_sampler = GraphSAINTNodeSampler(data, batch_size=4, num_steps=2, sample_coverage=1)
54
+ batch = next(iter(node_sampler))
55
+ assert isinstance(batch, Data)
56
+ assert hasattr(batch, 'node_norm')
57
+
58
+ # Edge sampler
59
+ edge_sampler = GraphSAINTEdgeSampler(data, batch_size=4, num_steps=2)
60
+ batch = next(iter(edge_sampler))
61
+ assert isinstance(batch, Data)
62
+
63
+ # Random walk sampler
64
+ rw_sampler = GraphSAINTRandomWalkSampler(data, batch_size=4, walk_length=2, num_steps=2)
65
+ batch = next(iter(rw_sampler))
66
+ assert isinstance(batch, Data)
67
+
68
+
69
+ def test_shadow_k_hop_sampler():
70
+ data = get_graph(num_nodes=12)
71
+ loader = ShaDowKHopSampler(data, depth=2, num_neighbors=2, batch_size=3)
72
+
73
+ batch = next(iter(loader))
74
+ assert batch.num_graphs == 3
75
+ assert hasattr(batch, 'root_n_id')
76
+
77
+
78
+ def test_legacy_neighbor_sampler():
79
+ data = get_graph(num_nodes=12)
80
+ loader = NeighborSampler(data.edge_index, sizes=[2, 2], batch_size=3)
81
+
82
+ batch_size, n_id, adjs = next(iter(loader))
83
+ assert batch_size == 3
84
+ assert len(n_id) >= 3
85
+ assert len(adjs) == 2
86
+ for adj in adjs:
87
+ assert hasattr(adj, 'edge_index')
88
+ assert hasattr(adj, 'size')
89
+
@@ -0,0 +1,232 @@
1
+ import copy
2
+ from typing import Any, Dict, Optional, Tuple, Union
3
+
4
+ import numpy as np
5
+ try:
6
+ import torch
7
+ except ImportError:
8
+ torch = None
9
+ Tensor = type(None)
10
+
11
+ from k3_node.data import Data, HeteroData
12
+ from k3_node.data.storage import NodeStorage, EdgeStorage
13
+
14
+
15
+ def to_numpy_or_tensor(x):
16
+ if torch is not None and isinstance(x, torch.Tensor):
17
+ return x
18
+ if isinstance(x, np.ndarray):
19
+ return x
20
+ if hasattr(x, '__array__'):
21
+ return np.asarray(x)
22
+ return x
23
+
24
+
25
+ def to_numpy(x, dtype=None):
26
+ r"""Converts any tensor (PyTorch, TensorFlow, JAX) or array to a CPU NumPy ndarray."""
27
+ if hasattr(x, "cpu"):
28
+ x = x.cpu()
29
+ if hasattr(x, "detach"):
30
+ x = x.detach()
31
+ if hasattr(x, "numpy") and callable(x.numpy):
32
+ x = x.numpy()
33
+ return np.asarray(x, dtype=dtype)
34
+
35
+
36
+ def index_select(value: Any, index: Any, dim: int = 0) -> Any:
37
+ r"""Indexes the :obj:`value` tensor along dimension :obj:`dim` using the
38
+ entries in :obj:`index`. Supports PyTorch, TensorFlow, JAX, and NumPy.
39
+ """
40
+ if torch is not None and isinstance(value, torch.Tensor):
41
+ if not isinstance(index, torch.Tensor):
42
+ index = torch.as_tensor(index, dtype=torch.long, device=value.device)
43
+ else:
44
+ index = index.to(dtype=torch.long, device=value.device)
45
+ return torch.index_select(value, dim, index)
46
+
47
+ # NumPy / Keras / JAX / TensorFlow array:
48
+ if hasattr(value, 'numpy'):
49
+ is_tf_or_jax = True
50
+ np_val = value.numpy()
51
+ elif isinstance(value, np.ndarray):
52
+ is_tf_or_jax = False
53
+ np_val = value
54
+ elif hasattr(value, '__array__'):
55
+ is_tf_or_jax = False
56
+ np_val = np.asarray(value)
57
+ else:
58
+ return value
59
+
60
+ if torch is not None and isinstance(index, torch.Tensor):
61
+ np_idx = index.cpu().numpy()
62
+ else:
63
+ np_idx = np.asarray(index, dtype=np.int64)
64
+
65
+ res = np.take(np_val, np_idx, axis=dim)
66
+ if is_tf_or_jax:
67
+ import keras
68
+ return keras.ops.convert_to_tensor(res)
69
+ return res
70
+
71
+
72
+ def filter_node_store_(store: NodeStorage, out_store: NodeStorage, index: Any):
73
+ for key, value in store.items():
74
+ if key == 'num_nodes':
75
+ numel = index.numel() if hasattr(index, 'numel') else len(index)
76
+ out_store.num_nodes = numel
77
+ elif store.is_node_attr(key):
78
+ dim = 0
79
+ if hasattr(store, '_parent') and store._parent() is not None:
80
+ dim = store._parent().__cat_dim__(key, value, store)
81
+ out_store[key] = index_select(value, index, dim=dim)
82
+
83
+
84
+ def filter_edge_store_(
85
+ store: EdgeStorage,
86
+ out_store: EdgeStorage,
87
+ row: Any,
88
+ col: Any,
89
+ index: Optional[Any],
90
+ perm: Optional[Any] = None,
91
+ ):
92
+ for key, value in store.items():
93
+ if key == 'edge_index':
94
+ if torch is not None and isinstance(row, torch.Tensor):
95
+ edge_index = torch.stack([row, col], dim=0)
96
+ else:
97
+ edge_index = np.stack([np.asarray(row), np.asarray(col)], axis=0)
98
+ out_store.edge_index = edge_index
99
+ elif store.is_edge_attr(key):
100
+ if index is None:
101
+ out_store[key] = None
102
+ continue
103
+ dim = 0
104
+ if hasattr(store, '_parent') and store._parent() is not None:
105
+ dim = store._parent().__cat_dim__(key, value, store)
106
+ if perm is None:
107
+ out_store[key] = index_select(value, index, dim=dim)
108
+ else:
109
+ sel_idx = perm[index] if hasattr(perm, '__getitem__') else index
110
+ out_store[key] = index_select(value, sel_idx, dim=dim)
111
+
112
+
113
+ def filter_data(data: Data, node: Any, row: Any, col: Any, edge: Optional[Any] = None, perm: Optional[Any] = None) -> Data:
114
+ out = copy.copy(data)
115
+ out._store = copy.copy(data._store)
116
+ filter_node_store_(data._store, out._store, node)
117
+ filter_edge_store_(data._store, out._store, row, col, edge, perm)
118
+ return out
119
+
120
+
121
+ def filter_hetero_data(
122
+ data: HeteroData,
123
+ node_dict: Dict[str, Any],
124
+ row_dict: Dict[Tuple[str, str, str], Any],
125
+ col_dict: Dict[Tuple[str, str, str], Any],
126
+ edge_dict: Dict[Tuple[str, str, str], Optional[Any]],
127
+ perm_dict: Optional[Dict[Tuple[str, str, str], Optional[Any]]] = None,
128
+ ) -> HeteroData:
129
+ out = copy.copy(data)
130
+ out._node_store_dict = {k: copy.copy(v) for k, v in data._node_store_dict.items()}
131
+ out._edge_store_dict = {k: copy.copy(v) for k, v in data._edge_store_dict.items()}
132
+
133
+ for node_type in out.node_types:
134
+ if node_type not in node_dict:
135
+ node_dict[node_type] = torch.empty(0, dtype=torch.long) if torch is not None else np.empty(0, dtype=np.int64)
136
+ filter_node_store_(data[node_type], out[node_type], node_dict[node_type])
137
+
138
+ for edge_type in out.edge_types:
139
+ canonical = data._to_canonical(*edge_type) if hasattr(data, '_to_canonical') else edge_type
140
+ if canonical not in row_dict:
141
+ empty_arr = torch.empty(0, dtype=torch.long) if torch is not None else np.empty(0, dtype=np.int64)
142
+ row_dict[canonical] = empty_arr
143
+ col_dict[canonical] = empty_arr
144
+ edge_dict[canonical] = empty_arr
145
+
146
+ filter_edge_store_(
147
+ data[edge_type],
148
+ out[edge_type],
149
+ row_dict[canonical],
150
+ col_dict[canonical],
151
+ edge_dict[canonical],
152
+ perm_dict.get(canonical, None) if perm_dict else None,
153
+ )
154
+
155
+ return out
156
+
157
+
158
+ def get_input_nodes(
159
+ data: Union[Data, HeteroData],
160
+ input_nodes: Any,
161
+ input_id: Optional[Any] = None,
162
+ ) -> Tuple[Optional[str], Any, Optional[Any]]:
163
+ def to_index(nodes, in_id):
164
+ if torch is not None and isinstance(nodes, torch.Tensor):
165
+ if nodes.dtype == torch.bool:
166
+ nodes = nodes.nonzero(as_tuple=False).view(-1)
167
+ in_id = nodes if in_id is None else in_id
168
+ return nodes, in_id
169
+ if isinstance(nodes, np.ndarray) and nodes.dtype == bool:
170
+ nodes = np.nonzero(nodes)[0]
171
+ in_id = nodes if in_id is None else in_id
172
+ return nodes, in_id
173
+ if torch is not None and not isinstance(nodes, torch.Tensor):
174
+ nodes = torch.tensor(nodes, dtype=torch.long)
175
+ elif torch is None:
176
+ nodes = np.asarray(nodes, dtype=np.int64)
177
+ return nodes, in_id
178
+
179
+ if isinstance(data, Data):
180
+ if input_nodes is None:
181
+ nodes = torch.arange(data.num_nodes) if torch is not None else np.arange(data.num_nodes)
182
+ return None, nodes, None
183
+ return None, *to_index(input_nodes, input_id)
184
+
185
+ elif isinstance(data, HeteroData):
186
+ assert input_nodes is not None
187
+ if isinstance(input_nodes, str):
188
+ num_nodes = data[input_nodes].num_nodes
189
+ nodes = torch.arange(num_nodes) if torch is not None else np.arange(num_nodes)
190
+ return input_nodes, nodes, None
191
+
192
+ assert isinstance(input_nodes, (list, tuple)) and len(input_nodes) == 2
193
+ node_type, input_nodes = input_nodes
194
+ if input_nodes is None:
195
+ num_nodes = data[node_type].num_nodes
196
+ nodes = torch.arange(num_nodes) if torch is not None else np.arange(num_nodes)
197
+ return node_type, nodes, None
198
+ return node_type, *to_index(input_nodes, input_id)
199
+
200
+ raise TypeError(f"Invalid data type: {type(data)}")
201
+
202
+
203
+ def get_edge_label_index(
204
+ data: Union[Data, HeteroData],
205
+ edge_label_index: Any,
206
+ ) -> Tuple[Optional[Tuple[str, str, str]], Any]:
207
+ if isinstance(data, Data):
208
+ if edge_label_index is None:
209
+ return None, data.edge_index
210
+ return None, edge_label_index
211
+
212
+ if isinstance(data, HeteroData):
213
+ assert edge_label_index is not None
214
+ if isinstance(edge_label_index, (list, tuple)) and len(edge_label_index) == 3 and isinstance(edge_label_index[0], str):
215
+ edge_type = data._to_canonical(*edge_label_index)
216
+ return edge_type, data[edge_type].edge_index
217
+
218
+ assert isinstance(edge_label_index, (list, tuple)) and len(edge_label_index) == 2
219
+ edge_type, edge_index = edge_label_index
220
+ edge_type = data._to_canonical(*edge_type)
221
+ if edge_index is None:
222
+ return edge_type, data[edge_type].edge_index
223
+ return edge_type, edge_index
224
+
225
+ raise TypeError(f"Invalid data type: {type(data)}")
226
+
227
+
228
+ def infer_filter_per_worker(data: Any) -> bool:
229
+ out = True
230
+ if hasattr(data, 'is_cuda') and data.is_cuda:
231
+ out = False
232
+ return out
@@ -0,0 +1,88 @@
1
+ from typing import Any, Iterator, List, Optional, Tuple, Union
2
+
3
+ try:
4
+ import torch
5
+ from torch import Tensor
6
+ BaseDataLoader = torch.utils.data.DataLoader
7
+ except ImportError:
8
+ torch = None
9
+ Tensor = type(None)
10
+ BaseDataLoader = object
11
+
12
+ from k3_node.data import Data, HeteroData
13
+ from k3_node.loader.base import DataLoaderIterator
14
+ from k3_node.loader.utils import infer_filter_per_worker
15
+
16
+
17
+ class ZipLoader(BaseDataLoader):
18
+ r"""A loader that returns a tuple of data objects by sampling from multiple
19
+ loader instances.
20
+
21
+ Args:
22
+ loaders (List[Any]): The loader instances.
23
+ filter_per_worker (bool, optional): If set to :obj:`True`, will filter
24
+ the returned data in each worker's subprocess. (default: :obj:`None`)
25
+ **kwargs (optional): Additional arguments of :class:`torch.utils.data.DataLoader`.
26
+ """
27
+ def __init__(
28
+ self,
29
+ loaders: List[Any],
30
+ filter_per_worker: Optional[bool] = None,
31
+ **kwargs,
32
+ ):
33
+ if filter_per_worker is None:
34
+ first_data = getattr(loaders[0], 'data', None)
35
+ filter_per_worker = infer_filter_per_worker(first_data) if first_data is not None else True
36
+
37
+ kwargs.pop('dataset', None)
38
+ kwargs.pop('collate_fn', None)
39
+
40
+ for loader in loaders:
41
+ if not callable(getattr(loader, 'collate_fn', None)):
42
+ raise ValueError(f"'{loader.__class__.__name__}' does not have a 'collate_fn' method")
43
+ if not callable(getattr(loader, 'filter_fn', None)):
44
+ raise ValueError(f"'{loader.__class__.__name__}' does not have a 'filter_fn' method")
45
+ loader.filter_per_worker = filter_per_worker
46
+
47
+ lens = []
48
+ for loader in loaders:
49
+ if hasattr(loader, 'dataset'):
50
+ lens.append(len(loader.dataset))
51
+ elif hasattr(loader, '__len__'):
52
+ lens.append(len(loader))
53
+ else:
54
+ lens.append(0)
55
+
56
+ iterator = range(min(lens) if lens else 0)
57
+
58
+ self.loaders = loaders
59
+ self.filter_per_worker = filter_per_worker
60
+
61
+ if torch is not None:
62
+ super().__init__(iterator, collate_fn=self.collate_fn, **kwargs)
63
+ else:
64
+ self.dataset = iterator
65
+ self.collate_fn = self.collate_fn
66
+
67
+ def __call__(self, index: Union[Tensor, List[int]]) -> Union[Tuple[Data, ...], Tuple[HeteroData, ...]]:
68
+ out = self.collate_fn(index)
69
+ if not self.filter_per_worker:
70
+ out = self.filter_fn(out)
71
+ return out
72
+
73
+ def collate_fn(self, index: List[int]) -> Tuple[Any, ...]:
74
+ if torch is not None and not isinstance(index, Tensor):
75
+ index = torch.tensor(index, dtype=torch.long)
76
+ return tuple(loader.collate_fn(index) for loader in self.loaders)
77
+
78
+ def filter_fn(self, outs: Tuple[Any, ...]) -> Tuple[Union[Data, HeteroData], ...]:
79
+ return tuple(loader.filter_fn(v) for loader, v in zip(self.loaders, outs))
80
+
81
+ def _get_iterator(self) -> Iterator:
82
+ if self.filter_per_worker:
83
+ return super()._get_iterator()
84
+ return DataLoaderIterator(super()._get_iterator(), self.filter_fn)
85
+
86
+ def __repr__(self) -> str:
87
+ return f'{self.__class__.__name__}(loaders={self.loaders})'
88
+
k3_node/metrics.py ADDED
@@ -0,0 +1,94 @@
1
+ """Metrics for graph learning tasks."""
2
+ import numpy as np
3
+ from typing import Optional
4
+
5
+ import keras
6
+ from keras import ops
7
+
8
+
9
+ @keras.saving.register_keras_serializable(package="k3_node")
10
+ class F1Score(keras.metrics.F1Score):
11
+ r"""F1 score that also accepts logits, e.g. for multi-label node classification.
12
+
13
+ Same as :class:`keras.metrics.F1Score`, plus ``from_logits``: when :obj:`True`, a sigmoid is
14
+ applied to the predictions first, so models can output logits (as used with
15
+ ``BinaryCrossentropy(from_logits=True)``).
16
+
17
+ Example:
18
+ ```python
19
+ import numpy as np
20
+ from k3_node.metrics import F1Score
21
+
22
+ f1 = F1Score(average="micro", from_logits=True)
23
+ y_true = np.array([[1, 0, 1], [0, 1, 0]], dtype="float32")
24
+ logits = np.array([[2.0, -1.0, 0.5], [-3.0, 1.5, 0.2]], dtype="float32")
25
+ f1.update_state(y_true, logits)
26
+ print(round(float(f1.result()), 2)) # 0.86
27
+ ```
28
+ """
29
+
30
+ def __init__(self, average=None, threshold=0.5, from_logits=False, name="f1_score", dtype=None):
31
+ super().__init__(average=average, threshold=threshold, name=name, dtype=dtype)
32
+ self.from_logits = from_logits
33
+
34
+ def update_state(self, y_true, y_pred, sample_weight=None):
35
+ if self.from_logits:
36
+ y_pred = ops.sigmoid(y_pred)
37
+ return super().update_state(y_true, y_pred, sample_weight=sample_weight)
38
+
39
+ def get_config(self):
40
+ return {**super().get_config(), "from_logits": self.from_logits}
41
+
42
+
43
+ def precision_recall_at_k(user_emb, item_emb, train_edge_index, test_edge_index, k: int = 20,
44
+ num_users: Optional[int] = None, batch_size: int = 8192):
45
+ r"""Top-``k`` recommendation quality, as in PyG's LightGCN example.
46
+
47
+ Every user's items are ranked by the dot product of the embeddings, excluding the items the
48
+ user interacted with during training. Returns the precision@k and recall@k averaged over the
49
+ users with at least one test interaction.
50
+
51
+ Args:
52
+ user_emb: User embeddings ``[num_users, dim]``.
53
+ item_emb: Item embeddings ``[num_items, dim]``.
54
+ train_edge_index: Training interactions ``(user, item)``; item ids may be offset by
55
+ ``num_users`` (as in a homogeneous user-item graph).
56
+ test_edge_index: Test interactions ``(user, item)``, with the same numbering.
57
+ k (int): The number of recommendations per user. (default: ``20``)
58
+ num_users (int, optional): The item-id offset. (default: ``len(user_emb)``)
59
+ batch_size (int): Users scored at once. (default: ``8192``)
60
+
61
+ Example:
62
+ ```python
63
+ import numpy as np
64
+ from k3_node.metrics import precision_recall_at_k
65
+
66
+ users, items = np.eye(2, 4, dtype="float32"), np.eye(3, 4, dtype="float32") # toy embeddings
67
+ train = np.array([[0], [2]]) # user 0 already has item 2 (id 2 + num_users = 4 in the graph)
68
+ test = np.array([[0, 1], [0 + 2, 1 + 2]]) # user 0 likes item 0, user 1 likes item 1
69
+ print(precision_recall_at_k(users, items, train + [[0], [2]], test, k=1)) # (1.0, 1.0)
70
+ ```
71
+ """
72
+ from keras import ops as _ops
73
+
74
+ user_emb, item_emb = (np.asarray(_ops.convert_to_numpy(e)) for e in (user_emb, item_emb))
75
+ num_users = len(user_emb) if num_users is None else num_users
76
+ train = np.asarray(_ops.convert_to_numpy(train_edge_index)).astype(np.int64)
77
+ test = np.asarray(_ops.convert_to_numpy(test_edge_index)).astype(np.int64)
78
+ train = train[:, train[0] < num_users]
79
+ precision = recall = examples = 0.0
80
+ for start in range(0, len(user_emb), batch_size):
81
+ end = min(start + batch_size, len(user_emb))
82
+ logits = user_emb[start:end] @ item_emb.T
83
+ m = (train[0] >= start) & (train[0] < end)
84
+ logits[train[0, m] - start, train[1, m] - num_users] = -np.inf # skip already known items
85
+ truth = np.zeros_like(logits, dtype=bool)
86
+ m = (test[0] >= start) & (test[0] < end)
87
+ truth[test[0, m] - start, test[1, m] - num_users] = True
88
+ count = truth.sum(axis=1)
89
+ top = np.argpartition(-logits, k - 1, axis=1)[:, :k]
90
+ hits = np.take_along_axis(truth, top, axis=1).sum(axis=1)
91
+ precision += float((hits / k)[count > 0].sum())
92
+ recall += float((hits / np.maximum(count, 1e-6))[count > 0].sum())
93
+ examples += int((count > 0).sum())
94
+ return precision / examples, recall / examples