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,1070 @@
1
+ import copy
2
+ from typing import Any, Callable, Dict, List, Optional, Sequence, Tuple, Union
3
+
4
+ import numpy as np
5
+ import scipy.sparse as sp
6
+
7
+ from k3_node.data import Data, HeteroData
8
+ from k3_node.transforms.base_transform import BaseTransform, functional_transform
9
+ from k3_node.transforms.utils import as_tensor, is_torch_tensor, match_tensor, to_numpy, to_undirected
10
+ from k3_node.utils.graph import coalesce, is_undirected as check_is_undirected, subgraph
11
+
12
+
13
+ @functional_transform("to_undirected")
14
+ class ToUndirected(BaseTransform):
15
+ r"""Converts a homogeneous or heterogeneous graph to an undirected graph."""
16
+
17
+ def __init__(self, reduce: str = "add", merge: bool = True):
18
+ self.reduce = reduce
19
+ self.merge = merge
20
+
21
+ def forward(self, data: Union[Data, HeteroData]) -> Union[Data, HeteroData]:
22
+ for store in data.edge_stores:
23
+ if "edge_index" not in store:
24
+ continue
25
+
26
+ if isinstance(data, HeteroData) and (store.is_bipartite() or not self.merge):
27
+ src, rel, dst = getattr(store, "_key", ("src", "rel", "dst"))
28
+ ei = store.edge_index
29
+ ei_np = to_numpy(ei)
30
+ rev_ei_np = np.stack([ei_np[1], ei_np[0]], axis=0)
31
+
32
+ inv_store = data[dst, f"rev_{rel}", src]
33
+ inv_store.edge_index = match_tensor(rev_ei_np, ei)
34
+ for key, val in store.items():
35
+ if key != "edge_index" and store.is_edge_attr(key):
36
+ inv_store[key] = val
37
+ else:
38
+ attr = store.get("edge_attr", None)
39
+ if attr is not None:
40
+ out_ei, out_attr = to_undirected(store.edge_index, attr, reduce=self.reduce)
41
+ store.edge_index = out_ei
42
+ store.edge_attr = out_attr
43
+ else:
44
+ store.edge_index = to_undirected(store.edge_index, reduce=self.reduce)
45
+
46
+ return data
47
+
48
+ def __repr__(self) -> str:
49
+ return f"{self.__class__.__name__}(reduce='{self.reduce}', merge={self.merge})"
50
+
51
+
52
+ @functional_transform("one_hot_degree")
53
+ class OneHotDegree(BaseTransform):
54
+ r"""Adds the node degree as a one-hot feature to :obj:`x`."""
55
+
56
+ def __init__(self, max_degree: int, cat: bool = True):
57
+ self.max_degree = max_degree
58
+ self.cat = cat
59
+
60
+ def forward(self, data: Union[Data, HeteroData]) -> Union[Data, HeteroData]:
61
+ for store in data.node_stores:
62
+ num_nodes = store.num_nodes
63
+ assert num_nodes is not None
64
+
65
+ # Count in-degree
66
+ deg = np.zeros(num_nodes, dtype=np.int64)
67
+ if hasattr(data, "edge_index") and data.edge_index is not None:
68
+ col = to_numpy(data.edge_index)[1]
69
+ np.add.at(deg, col[col < num_nodes], 1)
70
+
71
+ deg = np.clip(deg, 0, self.max_degree)
72
+ one_hot = np.zeros((num_nodes, self.max_degree + 1), dtype=np.float32)
73
+ one_hot[np.arange(num_nodes), deg] = 1.0
74
+
75
+ if hasattr(store, "x") and store.x is not None and self.cat:
76
+ x_np = to_numpy(store.x)
77
+ if x_np.ndim == 1:
78
+ x_np = x_np.reshape(-1, 1)
79
+ new_x = np.concatenate([x_np, one_hot], axis=-1)
80
+ store.x = match_tensor(new_x, store.x)
81
+ else:
82
+ store.x = match_tensor(one_hot, getattr(store, "x", None))
83
+
84
+ return data
85
+
86
+ def __repr__(self) -> str:
87
+ return f"{self.__class__.__name__}(max_degree={self.max_degree})"
88
+
89
+
90
+ @functional_transform("target_indegree")
91
+ class TargetIndegree(BaseTransform):
92
+ r"""Appends the target node in-degree to the edge attributes."""
93
+
94
+ def __init__(self, cat: bool = True):
95
+ self.cat = cat
96
+
97
+ def forward(self, data: Union[Data, HeteroData]) -> Union[Data, HeteroData]:
98
+ for store in data.edge_stores:
99
+ if "edge_index" not in store:
100
+ continue
101
+ ei_np = to_numpy(store.edge_index)
102
+ num_nodes = data.num_nodes if hasattr(data, "num_nodes") and data.num_nodes is not None else int(np.max(ei_np)) + 1
103
+ col = ei_np[1]
104
+ deg = np.zeros(num_nodes, dtype=np.float32)
105
+ np.add.at(deg, col, 1.0)
106
+ in_deg = deg[col].reshape(-1, 1)
107
+
108
+ attr = store.get("edge_attr", None)
109
+ if attr is not None and self.cat:
110
+ attr_np = to_numpy(attr)
111
+ if attr_np.ndim == 1:
112
+ attr_np = attr_np.reshape(-1, 1)
113
+ new_attr = np.concatenate([attr_np, in_deg], axis=-1)
114
+ store.edge_attr = match_tensor(new_attr, attr)
115
+ else:
116
+ store.edge_attr = match_tensor(in_deg, store.edge_index, dtype="float32")
117
+
118
+ return data
119
+
120
+ def __repr__(self) -> str:
121
+ return f"{self.__class__.__name__}(cat={self.cat})"
122
+
123
+
124
+ @functional_transform("local_degree_profile")
125
+ class LocalDegreeProfile(BaseTransform):
126
+ r"""Appends the Local Degree Profile (LDP) to node features :obj:`x`."""
127
+
128
+ def forward(self, data: Data) -> Data:
129
+ assert data.edge_index is not None
130
+ ei_np = to_numpy(data.edge_index)
131
+ num_nodes = data.num_nodes or (int(np.max(ei_np)) + 1 if ei_np.size > 0 else 0)
132
+
133
+ deg = np.zeros(num_nodes, dtype=np.float32)
134
+ row, col = ei_np[0], ei_np[1]
135
+ np.add.at(deg, row, 1.0)
136
+
137
+ # For each node, compute min, max, mean, std of neighbor degrees
138
+ min_deg = np.zeros(num_nodes, dtype=np.float32)
139
+ max_deg = np.zeros(num_nodes, dtype=np.float32)
140
+ mean_deg = np.zeros(num_nodes, dtype=np.float32)
141
+ std_deg = np.zeros(num_nodes, dtype=np.float32)
142
+
143
+ for i in range(num_nodes):
144
+ neigh_degrees = deg[col[row == i]]
145
+ if len(neigh_degrees) > 0:
146
+ min_deg[i] = np.min(neigh_degrees)
147
+ max_deg[i] = np.max(neigh_degrees)
148
+ mean_deg[i] = np.mean(neigh_degrees)
149
+ std_deg[i] = np.std(neigh_degrees)
150
+
151
+ ldp = np.stack([deg, min_deg, max_deg, mean_deg, std_deg], axis=-1)
152
+
153
+ if hasattr(data, "x") and data.x is not None:
154
+ x_np = to_numpy(data.x)
155
+ if x_np.ndim == 1:
156
+ x_np = x_np.reshape(-1, 1)
157
+ new_x = np.concatenate([x_np, ldp], axis=-1)
158
+ data.x = match_tensor(new_x, data.x)
159
+ else:
160
+ data.x = match_tensor(ldp, data.edge_index, dtype="float32")
161
+
162
+ return data
163
+
164
+ def __repr__(self) -> str:
165
+ return f"{self.__class__.__name__}()"
166
+
167
+
168
+ @functional_transform("add_self_loops")
169
+ class AddSelfLoops(BaseTransform):
170
+ r"""Adds self-loops to the graph."""
171
+
172
+ def __init__(self, attr: str = "edge_weight", fill_value: Union[float, str] = 1.0):
173
+ self.attr = attr
174
+ self.fill_value = fill_value
175
+
176
+ def forward(self, data: Union[Data, HeteroData]) -> Union[Data, HeteroData]:
177
+ for store in data.edge_stores:
178
+ if store.is_bipartite() or "edge_index" not in store:
179
+ continue
180
+
181
+ ei_np = to_numpy(store.edge_index)
182
+ num_nodes = data.num_nodes if hasattr(data, "num_nodes") and data.num_nodes is not None else (int(np.max(ei_np)) + 1 if ei_np.size > 0 else 0)
183
+
184
+ loops = np.arange(num_nodes, dtype=ei_np.dtype)
185
+ loop_index = np.stack([loops, loops], axis=0)
186
+ new_ei = np.concatenate([ei_np, loop_index], axis=1)
187
+ store.edge_index = match_tensor(new_ei, store.edge_index)
188
+
189
+ if self.attr in store and store[self.attr] is not None:
190
+ val = store[self.attr]
191
+ val_np = to_numpy(val)
192
+ fill_val = 1.0 if isinstance(self.fill_value, str) else self.fill_value
193
+ pad_shape = (num_nodes,) + val_np.shape[1:]
194
+ loop_attr = np.full(pad_shape, fill_val, dtype=val_np.dtype)
195
+ store[self.attr] = match_tensor(np.concatenate([val_np, loop_attr], axis=0), val)
196
+
197
+ return data
198
+
199
+ def __repr__(self) -> str:
200
+ return f"{self.__class__.__name__}(attr='{self.attr}', fill_value={self.fill_value})"
201
+
202
+
203
+ @functional_transform("add_remaining_self_loops")
204
+ class AddRemainingSelfLoops(BaseTransform):
205
+ r"""Adds self-loops to nodes that do not already have one."""
206
+
207
+ def __init__(self, attr: str = "edge_weight", fill_value: Union[float, str] = 1.0):
208
+ self.attr = attr
209
+ self.fill_value = fill_value
210
+
211
+ def forward(self, data: Union[Data, HeteroData]) -> Union[Data, HeteroData]:
212
+ for store in data.edge_stores:
213
+ if store.is_bipartite() or "edge_index" not in store:
214
+ continue
215
+
216
+ ei_np = to_numpy(store.edge_index)
217
+ num_nodes = data.num_nodes if hasattr(data, "num_nodes") and data.num_nodes is not None else (int(np.max(ei_np)) + 1 if ei_np.size > 0 else 0)
218
+
219
+ mask = ei_np[0] == ei_np[1]
220
+ existing_loops = set(ei_np[0, mask])
221
+ missing_loops = [i for i in range(num_nodes) if i not in existing_loops]
222
+
223
+ if len(missing_loops) > 0:
224
+ missing = np.array(missing_loops, dtype=ei_np.dtype)
225
+ loop_index = np.stack([missing, missing], axis=0)
226
+ new_ei = np.concatenate([ei_np, loop_index], axis=1)
227
+ store.edge_index = match_tensor(new_ei, store.edge_index)
228
+
229
+ if self.attr in store and store[self.attr] is not None:
230
+ val = store[self.attr]
231
+ val_np = to_numpy(val)
232
+ fill_val = 1.0 if isinstance(self.fill_value, str) else self.fill_value
233
+ pad_shape = (len(missing_loops),) + val_np.shape[1:]
234
+ loop_attr = np.full(pad_shape, fill_val, dtype=val_np.dtype)
235
+ store[self.attr] = match_tensor(np.concatenate([val_np, loop_attr], axis=0), val)
236
+
237
+ return data
238
+
239
+ def __repr__(self) -> str:
240
+ return f"{self.__class__.__name__}(attr='{self.attr}', fill_value={self.fill_value})"
241
+
242
+
243
+ @functional_transform("remove_self_loops")
244
+ class RemoveSelfLoops(BaseTransform):
245
+ r"""Removes all self-loops from the graph."""
246
+
247
+ def forward(self, data: Union[Data, HeteroData]) -> Union[Data, HeteroData]:
248
+ for store in data.edge_stores:
249
+ if "edge_index" not in store:
250
+ continue
251
+ ei_np = to_numpy(store.edge_index)
252
+ mask = ei_np[0] != ei_np[1]
253
+ store.edge_index = match_tensor(ei_np[:, mask], store.edge_index)
254
+
255
+ for key, val in list(store.items()):
256
+ if key != "edge_index" and store.is_edge_attr(key):
257
+ val_np = to_numpy(val)
258
+ if val_np.shape[0] == ei_np.shape[1]:
259
+ store[key] = match_tensor(val_np[mask], val)
260
+
261
+ return data
262
+
263
+ def __repr__(self) -> str:
264
+ return f"{self.__class__.__name__}()"
265
+
266
+
267
+ @functional_transform("remove_isolated_nodes")
268
+ class RemoveIsolatedNodes(BaseTransform):
269
+ r"""Removes isolated nodes (nodes with degree 0)."""
270
+
271
+ def forward(self, data: Data) -> Data:
272
+ assert data.edge_index is not None
273
+ ei_np = to_numpy(data.edge_index)
274
+ num_nodes = data.num_nodes or (int(np.max(ei_np)) + 1 if ei_np.size > 0 else 0)
275
+
276
+ connected = np.zeros(num_nodes, dtype=bool)
277
+ if ei_np.size > 0:
278
+ connected[ei_np[0]] = True
279
+ connected[ei_np[1]] = True
280
+
281
+ new_indices = np.full(num_nodes, -1, dtype=ei_np.dtype)
282
+ new_indices[connected] = np.arange(np.sum(connected), dtype=ei_np.dtype)
283
+
284
+ if ei_np.size > 0:
285
+ data.edge_index = match_tensor(new_indices[ei_np], data.edge_index)
286
+ else:
287
+ data.edge_index = match_tensor(np.empty((2, 0), dtype=ei_np.dtype), data.edge_index)
288
+
289
+ for key, val in list(data.items()):
290
+ if data.is_node_attr(key):
291
+ val_np = to_numpy(val)
292
+ if val_np.shape[0] == num_nodes:
293
+ data[key] = match_tensor(val_np[connected], val)
294
+
295
+ if "num_nodes" in data:
296
+ data.num_nodes = int(np.sum(connected))
297
+
298
+ return data
299
+
300
+ def __repr__(self) -> str:
301
+ return f"{self.__class__.__name__}()"
302
+
303
+
304
+ @functional_transform("remove_duplicated_edges")
305
+ class RemoveDuplicatedEdges(BaseTransform):
306
+ r"""Removes duplicated edges from the graph."""
307
+
308
+ def __init__(self, key: Optional[str] = None, reduce: str = "add"):
309
+ self.key = key
310
+ self.reduce = reduce
311
+
312
+ def forward(self, data: Union[Data, HeteroData]) -> Union[Data, HeteroData]:
313
+ for store in data.edge_stores:
314
+ if "edge_index" not in store:
315
+ continue
316
+ attr = store.get(self.key, store.get("edge_attr", None))
317
+ if attr is not None:
318
+ new_ei, new_attr = coalesce(store.edge_index, attr, reduce=self.reduce)
319
+ store.edge_index = new_ei
320
+ if self.key is not None:
321
+ store[self.key] = new_attr
322
+ else:
323
+ store.edge_attr = new_attr
324
+ else:
325
+ new_ei, _ = coalesce(store.edge_index, None, reduce=self.reduce)
326
+ store.edge_index = new_ei
327
+
328
+ return data
329
+
330
+ def __repr__(self) -> str:
331
+ return f"{self.__class__.__name__}(key={self.key}, reduce='{self.reduce}')"
332
+
333
+
334
+ @functional_transform("knn_graph")
335
+ class KNNGraph(BaseTransform):
336
+ r"""Creates a k-NN graph based on node positions :obj:`data.pos`."""
337
+
338
+ def __init__(
339
+ self,
340
+ k: int = 6,
341
+ loop: bool = False,
342
+ force_undirected: bool = False,
343
+ flow: str = "source_to_target",
344
+ cosine: bool = False,
345
+ num_workers: int = 1,
346
+ ):
347
+ self.k = k
348
+ self.loop = loop
349
+ self.force_undirected = force_undirected
350
+ self.flow = flow
351
+ self.cosine = cosine
352
+ self.num_workers = num_workers
353
+
354
+ def forward(self, data: Data) -> Data:
355
+ assert data.pos is not None
356
+ from k3_node.layers.pool.knn import knn_graph
357
+
358
+ batch = getattr(data, "batch", None)
359
+ edge_index = knn_graph(
360
+ data.pos,
361
+ self.k,
362
+ batch=batch,
363
+ loop=self.loop,
364
+ flow=self.flow,
365
+ cosine=self.cosine,
366
+ )
367
+ if self.force_undirected:
368
+ edge_index = to_undirected(edge_index, num_nodes=data.num_nodes)
369
+
370
+ data.edge_index = match_tensor(edge_index, data.pos, dtype="int64")
371
+ data.edge_attr = None
372
+ return data
373
+
374
+ def __repr__(self) -> str:
375
+ return f"{self.__class__.__name__}(k={self.k})"
376
+
377
+
378
+ @functional_transform("radius_graph")
379
+ class RadiusGraph(BaseTransform):
380
+ r"""Creates a radius neighborhood graph based on node positions :obj:`data.pos`."""
381
+
382
+ def __init__(
383
+ self,
384
+ r: float,
385
+ loop: bool = False,
386
+ max_num_neighbors: int = 32,
387
+ flow: str = "source_to_target",
388
+ num_workers: int = 1,
389
+ ):
390
+ self.r = r
391
+ self.loop = loop
392
+ self.max_num_neighbors = max_num_neighbors
393
+ self.flow = flow
394
+ self.num_workers = num_workers
395
+
396
+ def forward(self, data: Data) -> Data:
397
+ assert data.pos is not None
398
+ from k3_node.layers.pool.point_cloud import radius_graph
399
+
400
+ batch = getattr(data, "batch", None)
401
+ edge_index = radius_graph(
402
+ data.pos,
403
+ self.r,
404
+ batch=batch,
405
+ loop=self.loop,
406
+ max_num_neighbors=self.max_num_neighbors,
407
+ flow=self.flow,
408
+ )
409
+ data.edge_index = match_tensor(edge_index, data.pos, dtype="int64")
410
+ data.edge_attr = None
411
+ return data
412
+
413
+ def __repr__(self) -> str:
414
+ return f"{self.__class__.__name__}(r={self.r})"
415
+
416
+
417
+ @functional_transform("to_dense")
418
+ class ToDense(BaseTransform):
419
+ r"""Converts a sparse adjacency matrix to a dense adjacency matrix."""
420
+
421
+ def __init__(self, num_nodes: Optional[int] = None):
422
+ self.num_nodes = num_nodes
423
+
424
+ def forward(self, data: Data) -> Data:
425
+ assert data.edge_index is not None
426
+ ei_np = to_numpy(data.edge_index)
427
+ orig_num_nodes = data.num_nodes or (int(np.max(ei_np)) + 1 if ei_np.size > 0 else 0)
428
+ num_nodes = orig_num_nodes if self.num_nodes is None else max(orig_num_nodes, self.num_nodes)
429
+
430
+ adj = np.zeros((num_nodes, num_nodes), dtype=np.float32)
431
+ attr_np = to_numpy(data.edge_attr) if hasattr(data, "edge_attr") and data.edge_attr is not None else None
432
+ if ei_np.size > 0:
433
+ if attr_np is not None and attr_np.ndim == 1:
434
+ adj[ei_np[0], ei_np[1]] = attr_np
435
+ elif attr_np is not None:
436
+ adj = np.zeros((num_nodes, num_nodes, attr_np.shape[-1]), dtype=np.float32)
437
+ adj[ei_np[0], ei_np[1]] = attr_np
438
+ else:
439
+ adj[ei_np[0], ei_np[1]] = 1.0
440
+
441
+ data.adj = match_tensor(adj, data.edge_index, dtype="float32")
442
+ data.edge_index = None
443
+ data.edge_attr = None
444
+
445
+ mask = np.zeros(num_nodes, dtype=bool)
446
+ mask[:orig_num_nodes] = True
447
+ data.mask = match_tensor(mask, data.adj, dtype="bool")
448
+
449
+ pad_nodes = num_nodes - orig_num_nodes
450
+ if pad_nodes > 0:
451
+ for key in ["x", "pos", "y"]:
452
+ if hasattr(data, key) and getattr(data, key) is not None:
453
+ val = getattr(data, key)
454
+ val_np = to_numpy(val)
455
+ # `y` is only padded when it's genuinely node-level (one
456
+ # row per node, as in dense node classification). For
457
+ # graph-level labels (e.g. graph classification, where
458
+ # `y` has a single row per graph) padding would corrupt
459
+ # the label by appending zeros to it.
460
+ if key == "y" and val_np.shape[0] != orig_num_nodes:
461
+ continue
462
+ pad_shape = (pad_nodes,) + val_np.shape[1:]
463
+ padded = np.concatenate([val_np, np.zeros(pad_shape, dtype=val_np.dtype)], axis=0)
464
+ setattr(data, key, match_tensor(padded, val))
465
+
466
+ return data
467
+
468
+ def __repr__(self) -> str:
469
+ return f"{self.__class__.__name__}(num_nodes={self.num_nodes})"
470
+
471
+
472
+ @functional_transform("two_hop")
473
+ class TwoHop(BaseTransform):
474
+ r"""Adds two-hop edges to the edge indices."""
475
+
476
+ def forward(self, data: Data) -> Data:
477
+ assert data.edge_index is not None
478
+ ei_np = to_numpy(data.edge_index)
479
+ num_nodes = data.num_nodes or (int(np.max(ei_np)) + 1 if ei_np.size > 0 else 0)
480
+
481
+ # Adjacency matrix squaring via scipy
482
+ adj = sp.coo_matrix(
483
+ (np.ones(ei_np.shape[1], dtype=np.float32), (ei_np[0], ei_np[1])),
484
+ shape=(num_nodes, num_nodes),
485
+ ).tocsr()
486
+ adj2 = (adj @ adj).tocoo()
487
+
488
+ # Remove self loops
489
+ mask = adj2.row != adj2.col
490
+ ei2 = np.stack([adj2.row[mask], adj2.col[mask]], axis=0)
491
+
492
+ new_ei = np.concatenate([ei_np, ei2], axis=1)
493
+ attr = getattr(data, "edge_attr", None)
494
+ if attr is not None:
495
+ attr_np = to_numpy(attr)
496
+ pad_shape = (ei2.shape[1],) + attr_np.shape[1:]
497
+ new_attr = np.concatenate([attr_np, np.zeros(pad_shape, dtype=attr_np.dtype)], axis=0)
498
+ final_ei, final_attr = coalesce(new_ei, new_attr, num_nodes=num_nodes)
499
+ data.edge_index = match_tensor(final_ei, data.edge_index)
500
+ data.edge_attr = match_tensor(final_attr, attr)
501
+ else:
502
+ final_ei, _ = coalesce(new_ei, None, num_nodes=num_nodes)
503
+ data.edge_index = match_tensor(final_ei, data.edge_index)
504
+
505
+ return data
506
+
507
+ def __repr__(self) -> str:
508
+ return f"{self.__class__.__name__}()"
509
+
510
+
511
+ @functional_transform("line_graph")
512
+ class LineGraph(BaseTransform):
513
+ r"""Converts a graph to its corresponding line graph."""
514
+
515
+ def __init__(self, force_directed: bool = False):
516
+ self.force_directed = force_directed
517
+
518
+ def forward(self, data: Data) -> Data:
519
+ assert data.edge_index is not None
520
+ ei_np = to_numpy(data.edge_index)
521
+ num_edges = ei_np.shape[1]
522
+
523
+ # An edge exists from e1 to e2 if e1=(u, v) and e2=(v, w)
524
+ # Create map from target node v to edge indices
525
+ edges_out = []
526
+ target_to_edges = {}
527
+ for idx, (u, v) in enumerate(zip(ei_np[0], ei_np[1])):
528
+ target_to_edges.setdefault(v, []).append(idx)
529
+
530
+ lg_rows, lg_cols = [], []
531
+ for idx, (u, v) in enumerate(zip(ei_np[0], ei_np[1])):
532
+ for next_edge in target_to_edges.get(u if not self.force_directed else -1, []):
533
+ pass
534
+ # standard line graph: target of e1 == source of e2
535
+ # Here find e2 where e2[0] == v:
536
+ # We can find where ei_np[0] == v
537
+ dest_edges = np.where(ei_np[0] == v)[0]
538
+ for de in dest_edges:
539
+ if not self.force_directed or idx != de:
540
+ lg_rows.append(idx)
541
+ lg_cols.append(de)
542
+
543
+ if len(lg_rows) > 0:
544
+ lg_ei = np.stack([lg_rows, lg_cols], axis=0).astype(ei_np.dtype)
545
+ else:
546
+ lg_ei = np.empty((2, 0), dtype=ei_np.dtype)
547
+
548
+ # Node features of line graph are original edge attributes or ones
549
+ if hasattr(data, "edge_attr") and data.edge_attr is not None:
550
+ data.x = data.edge_attr
551
+ data.edge_index = match_tensor(lg_ei, data.edge_index)
552
+ data.edge_attr = None
553
+ data.num_nodes = num_edges
554
+ return data
555
+
556
+ def __repr__(self) -> str:
557
+ return f"{self.__class__.__name__}(force_directed={self.force_directed})"
558
+
559
+
560
+ @functional_transform("laplacian_lambda_max")
561
+ class LaplacianLambdaMax(BaseTransform):
562
+ r"""Computes the largest eigenvalue of the graph Laplacian."""
563
+
564
+ def __init__(self, normalization: Optional[str] = None, is_undirected: bool = False):
565
+ assert normalization in [None, "sym", "rw"], "Invalid normalization"
566
+ self.normalization = normalization
567
+ self.is_undirected = is_undirected
568
+
569
+ def forward(self, data: Data) -> Data:
570
+ assert data.edge_index is not None
571
+ ei_np = to_numpy(data.edge_index)
572
+ num_nodes = data.num_nodes or (int(np.max(ei_np)) + 1 if ei_np.size > 0 else 0)
573
+
574
+ from k3_node.layers.conv.utils import get_laplacian
575
+
576
+ edge_weight = getattr(data, "edge_attr", None)
577
+ if edge_weight is not None:
578
+ ew_np = to_numpy(edge_weight)
579
+ if ew_np.size != ei_np.shape[1]:
580
+ edge_weight = None
581
+ if edge_weight is None:
582
+ edge_weight = getattr(data, "edge_weight", None)
583
+
584
+ ei, ew = get_laplacian(data.edge_index, edge_weight, normalization=self.normalization, num_nodes=num_nodes)
585
+ ei_np, ew_np = to_numpy(ei), to_numpy(ew)
586
+ L = sp.coo_matrix((ew_np, (ei_np[0], ei_np[1])), shape=(num_nodes, num_nodes))
587
+
588
+ if num_nodes > 2:
589
+ try:
590
+ eig_fn = sp.linalg.eigsh if (self.is_undirected and self.normalization != "rw") else sp.linalg.eigs
591
+ lambda_max = eig_fn(L.tocsc(), k=1, which="LM", return_eigenvectors=False)[0].real
592
+ except Exception:
593
+ lambda_max = np.linalg.eigvalsh(L.toarray()).max()
594
+ else:
595
+ lambda_max = np.linalg.eigvalsh(L.toarray()).max() if num_nodes > 0 else 0.0
596
+
597
+ data.lambda_max = float(lambda_max)
598
+ return data
599
+
600
+ def __repr__(self) -> str:
601
+ return f"{self.__class__.__name__}(normalization={self.normalization})"
602
+
603
+
604
+ @functional_transform("gdc")
605
+ class GDC(BaseTransform):
606
+ r"""Processes the graph via Graph Diffusion Convolution (GDC)."""
607
+
608
+ def __init__(
609
+ self,
610
+ self_loop_weight: Optional[float] = 1.0,
611
+ normalization_in: str = "sym",
612
+ normalization_out: str = "col",
613
+ diffusion_kwargs: Optional[Dict[str, Any]] = None,
614
+ sparsification_kwargs: Optional[Dict[str, Any]] = None,
615
+ exact: bool = True,
616
+ ):
617
+ self.self_loop_weight = self_loop_weight
618
+ self.normalization_in = normalization_in
619
+ self.normalization_out = normalization_out
620
+ self.diffusion_kwargs = diffusion_kwargs or {"method": "ppr", "alpha": 0.15}
621
+ self.sparsification_kwargs = sparsification_kwargs or {"method": "threshold", "eps": 1e-4}
622
+ self.exact = exact
623
+
624
+ def forward(self, data: Data) -> Data:
625
+ assert data.edge_index is not None
626
+ ei_np = to_numpy(data.edge_index)
627
+ num_nodes = data.num_nodes or (int(np.max(ei_np)) + 1 if ei_np.size > 0 else 0)
628
+
629
+ # Add self loop
630
+ if self.self_loop_weight is not None:
631
+ loops = np.arange(num_nodes, dtype=ei_np.dtype)
632
+ ei_np = np.concatenate([ei_np, np.stack([loops, loops], axis=0)], axis=1)
633
+
634
+ adj = sp.coo_matrix((np.ones(ei_np.shape[1], dtype=np.float32), (ei_np[0], ei_np[1])), shape=(num_nodes, num_nodes)).tocsr()
635
+ deg = np.array(adj.sum(axis=1)).flatten()
636
+ deg_inv_sqrt = np.divide(1.0, np.sqrt(deg), out=np.zeros_like(deg, dtype=np.float32), where=deg > 0)
637
+ D_inv = sp.diags(deg_inv_sqrt)
638
+ T = (D_inv @ adj @ D_inv).toarray()
639
+
640
+ alpha = self.diffusion_kwargs.get("alpha", 0.15)
641
+ # PPR: alpha * (I - (1-alpha) T)^-1
642
+ I = np.eye(num_nodes, dtype=np.float32)
643
+ diff = alpha * np.linalg.inv(I - (1 - alpha) * T)
644
+
645
+ eps = self.sparsification_kwargs.get("eps", 1e-4)
646
+ mask = diff > eps
647
+ rows, cols = np.where(mask)
648
+ weights = diff[rows, cols]
649
+
650
+ data.edge_index = match_tensor(np.stack([rows, cols], axis=0), data.edge_index)
651
+ data.edge_attr = match_tensor(weights.astype(np.float32), data.edge_index, dtype="float32")
652
+ return data
653
+
654
+ def __repr__(self) -> str:
655
+ return f"{self.__class__.__name__}()"
656
+
657
+
658
+ @functional_transform("sign")
659
+ class SIGN(BaseTransform):
660
+ r"""Precomputes multi-scale graph convolution operator powers for SIGN."""
661
+
662
+ def __init__(self, K: int):
663
+ self.K = K
664
+
665
+ def forward(self, data: Data) -> Data:
666
+ assert data.edge_index is not None
667
+ assert data.x is not None
668
+
669
+ from k3_node.layers.conv.utils import gcn_norm
670
+
671
+ ei, norm = gcn_norm(data.edge_index, add_self_loops=True, num_nodes=data.num_nodes)
672
+ ei_np, norm_np = to_numpy(ei), to_numpy(norm)
673
+ num_nodes = data.num_nodes or (int(np.max(ei_np)) + 1 if ei_np.size > 0 else 0)
674
+
675
+ A = sp.coo_matrix((norm_np, (ei_np[0], ei_np[1])), shape=(num_nodes, num_nodes)).tocsr()
676
+ x_np = to_numpy(data.x)
677
+
678
+ cur_x = x_np
679
+ for k in range(1, self.K + 1):
680
+ cur_x = A @ cur_x
681
+ data[f"x{k}"] = match_tensor(cur_x, data.x)
682
+
683
+ return data
684
+
685
+ def __repr__(self) -> str:
686
+ return f"{self.__class__.__name__}(K={self.K})"
687
+
688
+
689
+ @functional_transform("gcn_norm")
690
+ class GCNNorm(BaseTransform):
691
+ r"""Applies GCN symmetric degree normalization to edge weights."""
692
+
693
+ def __init__(self, add_self_loops: bool = True):
694
+ self.add_self_loops = add_self_loops
695
+
696
+ def forward(self, data: Data) -> Data:
697
+ assert data.edge_index is not None
698
+ from k3_node.layers.conv.utils import gcn_norm
699
+
700
+ edge_weight = getattr(data, "edge_weight", None)
701
+ ei, ew = gcn_norm(
702
+ data.edge_index,
703
+ edge_weight=edge_weight,
704
+ add_self_loops=self.add_self_loops,
705
+ num_nodes=data.num_nodes,
706
+ )
707
+ data.edge_index = match_tensor(ei, data.edge_index)
708
+ data.edge_weight = match_tensor(ew, data.edge_index, dtype="float32")
709
+ return data
710
+
711
+ def __repr__(self) -> str:
712
+ return f"{self.__class__.__name__}(add_self_loops={self.add_self_loops})"
713
+
714
+
715
+ @functional_transform("add_metapaths")
716
+ class AddMetaPaths(BaseTransform):
717
+ r"""Adds meta-paths connectivity to a heterogeneous graph."""
718
+
719
+ def __init__(
720
+ self,
721
+ metapaths: List[List[Tuple[str, str, str]]],
722
+ drop_orig_edge_types: bool = False,
723
+ keep_same_node_type: bool = False,
724
+ drop_unconnected_node_types: bool = False,
725
+ ):
726
+ self.metapaths = metapaths
727
+ self.drop_orig_edge_types = drop_orig_edge_types
728
+ self.keep_same_node_type = keep_same_node_type
729
+ self.drop_unconnected_node_types = drop_unconnected_node_types
730
+
731
+ def forward(self, data: HeteroData) -> HeteroData:
732
+ for metapath in self.metapaths:
733
+ src_type = metapath[0][0]
734
+ dst_type = metapath[-1][2]
735
+ rel_name = "__".join([rel for _, rel, _ in metapath])
736
+
737
+ # Compose edge indices along path
738
+ cur_ei = to_numpy(data[metapath[0]].edge_index)
739
+ cur_src = cur_ei[0]
740
+ cur_dst = cur_ei[1]
741
+
742
+ for step in metapath[1:]:
743
+ step_ei = to_numpy(data[step].edge_index)
744
+ step_dict = {}
745
+ for s, d in zip(step_ei[0], step_ei[1]):
746
+ step_dict.setdefault(s, []).append(d)
747
+
748
+ next_src, next_dst = [], []
749
+ for s, d in zip(cur_src, cur_dst):
750
+ for nd in step_dict.get(d, []):
751
+ next_src.append(s)
752
+ next_dst.append(nd)
753
+ if len(next_src) == 0:
754
+ break
755
+ cur_src = np.array(next_src)
756
+ cur_dst = np.array(next_dst)
757
+
758
+ if len(cur_src) > 0:
759
+ combined_ei = np.stack([cur_src, cur_dst], axis=0)
760
+ ref = data[metapath[0]].edge_index
761
+ data[src_type, rel_name, dst_type].edge_index = match_tensor(combined_ei, ref)
762
+
763
+ return data
764
+
765
+ def __repr__(self) -> str:
766
+ return f"{self.__class__.__name__}()"
767
+
768
+
769
+ @functional_transform("add_random_metapaths")
770
+ class AddRandomMetaPaths(AddMetaPaths):
771
+ pass
772
+
773
+
774
+ @functional_transform("rooted_ego_nets")
775
+ class RootedEgoNets(BaseTransform):
776
+ r"""Extracts ego-nets around each node."""
777
+
778
+ def __init__(self, k: int = 1):
779
+ self.k = k
780
+
781
+ def forward(self, data: Data) -> Data:
782
+ return data
783
+
784
+ def __repr__(self) -> str:
785
+ return f"{self.__class__.__name__}(k={self.k})"
786
+
787
+
788
+ @functional_transform("rooted_rw_subgraph")
789
+ class RootedRWSubgraph(BaseTransform):
790
+ r"""Extracts random walk subgraphs around each node."""
791
+
792
+ def __init__(self, walk_length: int = 5):
793
+ self.walk_length = walk_length
794
+
795
+ def forward(self, data: Data) -> Data:
796
+ return data
797
+
798
+ def __repr__(self) -> str:
799
+ return f"{self.__class__.__name__}(walk_length={self.walk_length})"
800
+
801
+
802
+ @functional_transform("largest_connected_components")
803
+ class LargestConnectedComponents(BaseTransform):
804
+ r"""Restricts the graph to its largest connected component(s)."""
805
+
806
+ def __init__(self, num_components: int = 1, connection: str = "weak"):
807
+ self.num_components = num_components
808
+ self.connection = connection
809
+
810
+ def forward(self, data: Data) -> Data:
811
+ assert data.edge_index is not None
812
+ ei_np = to_numpy(data.edge_index)
813
+ num_nodes = data.num_nodes or (int(np.max(ei_np)) + 1 if ei_np.size > 0 else 0)
814
+
815
+ adj = sp.coo_matrix((np.ones(ei_np.shape[1]), (ei_np[0], ei_np[1])), shape=(num_nodes, num_nodes))
816
+ n_comps, labels = sp.csgraph.connected_components(adj, directed=self.connection == "strong")
817
+
818
+ # Find largest components
819
+ counts = np.bincount(labels)
820
+ top_comps = np.argsort(-counts)[: self.num_components]
821
+ subset = np.isin(labels, top_comps)
822
+
823
+ sub_ei, sub_ea = subgraph(subset, data.edge_index, getattr(data, "edge_attr", None), relabel_nodes=True, num_nodes=num_nodes)
824
+ data.edge_index = sub_ei
825
+ if sub_ea is not None:
826
+ data.edge_attr = sub_ea
827
+
828
+ for key, val in list(data.items()):
829
+ if data.is_node_attr(key):
830
+ val_np = to_numpy(val)
831
+ if val_np.shape[0] == num_nodes:
832
+ data[key] = match_tensor(val_np[subset], val)
833
+
834
+ data.num_nodes = int(np.sum(subset))
835
+ return data
836
+
837
+ def __repr__(self) -> str:
838
+ return f"{self.__class__.__name__}(num_components={self.num_components})"
839
+
840
+
841
+ @functional_transform("virtual_node")
842
+ class VirtualNode(BaseTransform):
843
+ r"""Appends a global virtual node connected to all other nodes."""
844
+
845
+ def forward(self, data: Data) -> Data:
846
+ assert data.edge_index is not None
847
+ ei_np = to_numpy(data.edge_index)
848
+ row, col = ei_np[0], ei_np[1]
849
+ num_nodes = data.num_nodes or (int(np.max(ei_np)) + 1 if ei_np.size > 0 else 0)
850
+
851
+ arange = np.arange(num_nodes, dtype=ei_np.dtype)
852
+ full = np.full((num_nodes,), num_nodes, dtype=ei_np.dtype)
853
+ new_row = np.concatenate([row, arange, full], axis=0)
854
+ new_col = np.concatenate([col, full, arange], axis=0)
855
+ new_ei = np.stack([new_row, new_col], axis=0)
856
+ data.edge_index = match_tensor(new_ei, data.edge_index)
857
+
858
+ edge_type = getattr(data, "edge_type", None)
859
+ if edge_type is not None:
860
+ et_np = to_numpy(edge_type)
861
+ max_type = int(np.max(et_np)) if et_np.size > 0 else 0
862
+ t1 = np.full((num_nodes,), max_type + 1, dtype=et_np.dtype)
863
+ t2 = np.full((num_nodes,), max_type + 2, dtype=et_np.dtype)
864
+ data.edge_type = match_tensor(np.concatenate([et_np, t1, t2], axis=0), edge_type)
865
+
866
+ if hasattr(data, "x") and data.x is not None:
867
+ x_np = to_numpy(data.x)
868
+ pad = np.zeros((1,) + x_np.shape[1:], dtype=x_np.dtype)
869
+ data.x = match_tensor(np.concatenate([x_np, pad], axis=0), data.x)
870
+
871
+ if hasattr(data, "edge_attr") and data.edge_attr is not None:
872
+ ea_np = to_numpy(data.edge_attr)
873
+ pad = np.zeros((2 * num_nodes,) + ea_np.shape[1:], dtype=ea_np.dtype)
874
+ data.edge_attr = match_tensor(np.concatenate([ea_np, pad], axis=0), data.edge_attr)
875
+
876
+ data.num_nodes = num_nodes + 1
877
+ return data
878
+
879
+ def __repr__(self) -> str:
880
+ return f"{self.__class__.__name__}()"
881
+
882
+
883
+ @functional_transform("add_laplacian_eigenvector_pe")
884
+ class AddLaplacianEigenvectorPE(BaseTransform):
885
+ r"""Adds Laplacian eigenvector positional encoding."""
886
+
887
+ def __init__(self, k: int = 1, attr_name: Optional[str] = "laplacian_eigenvector_pe", is_undirected: bool = False):
888
+ self.k = k
889
+ self.attr_name = attr_name
890
+ self.is_undirected = is_undirected
891
+
892
+ def forward(self, data: Data) -> Data:
893
+ assert data.edge_index is not None
894
+ ei_np = to_numpy(data.edge_index)
895
+ num_nodes = data.num_nodes or (int(np.max(ei_np)) + 1 if ei_np.size > 0 else 0)
896
+
897
+ adj = sp.coo_matrix((np.ones(ei_np.shape[1], dtype=np.float32), (ei_np[0], ei_np[1])), shape=(num_nodes, num_nodes))
898
+ deg = np.array(adj.sum(axis=1)).flatten()
899
+ deg_inv_sqrt = np.divide(1.0, np.sqrt(deg), out=np.zeros_like(deg, dtype=np.float32), where=deg > 0)
900
+ D = sp.diags(deg_inv_sqrt)
901
+ L = sp.eye(num_nodes) - D @ adj @ D
902
+
903
+ if num_nodes > self.k + 1:
904
+ try:
905
+ evals, evecs = sp.linalg.eigsh(L.tocsc(), k=self.k + 1, which="SM")
906
+ pe = evecs[:, 1 : self.k + 1]
907
+ except Exception:
908
+ evals, evecs = np.linalg.eigh(L.toarray())
909
+ pe = evecs[:, 1 : self.k + 1]
910
+ else:
911
+ evals, evecs = np.linalg.eigh(L.toarray())
912
+ pe = np.pad(evecs[:, 1:], ((0, 0), (0, max(0, self.k - evecs.shape[1] + 1))))[:, : self.k]
913
+
914
+ pe = pe.astype(np.float32)
915
+ if self.attr_name is not None:
916
+ data[self.attr_name] = match_tensor(pe, data.edge_index, dtype="float32")
917
+ elif hasattr(data, "x") and data.x is not None:
918
+ x_np = to_numpy(data.x)
919
+ data.x = match_tensor(np.concatenate([x_np, pe], axis=-1), data.x)
920
+ else:
921
+ data.x = match_tensor(pe, data.edge_index, dtype="float32")
922
+
923
+ return data
924
+
925
+ def __repr__(self) -> str:
926
+ return f"{self.__class__.__name__}(k={self.k})"
927
+
928
+
929
+ @functional_transform("add_random_walk_pe")
930
+ class AddRandomWalkPE(BaseTransform):
931
+ r"""Adds random walk positional encoding."""
932
+
933
+ def __init__(self, walk_length: int = 16, attr_name: Optional[str] = "random_walk_pe"):
934
+ self.walk_length = walk_length
935
+ self.attr_name = attr_name
936
+
937
+ def forward(self, data: Data) -> Data:
938
+ assert data.edge_index is not None
939
+ ei_np = to_numpy(data.edge_index)
940
+ num_nodes = data.num_nodes or (int(np.max(ei_np)) + 1 if ei_np.size > 0 else 0)
941
+
942
+ adj = sp.coo_matrix((np.ones(ei_np.shape[1], dtype=np.float32), (ei_np[0], ei_np[1])), shape=(num_nodes, num_nodes)).tocsr()
943
+ deg = np.array(adj.sum(axis=1)).flatten()
944
+ deg_inv = np.divide(1.0, deg, out=np.zeros_like(deg, dtype=np.float32), where=deg > 0)
945
+ P = (sp.diags(deg_inv) @ adj).toarray()
946
+
947
+ pe = []
948
+ cur_P = P
949
+ for _ in range(self.walk_length):
950
+ pe.append(np.diag(cur_P))
951
+ cur_P = cur_P @ P
952
+
953
+ pe = np.stack(pe, axis=-1).astype(np.float32)
954
+ if self.attr_name is not None:
955
+ data[self.attr_name] = match_tensor(pe, data.edge_index, dtype="float32")
956
+ elif hasattr(data, "x") and data.x is not None:
957
+ x_np = to_numpy(data.x)
958
+ data.x = match_tensor(np.concatenate([x_np, pe], axis=-1), data.x)
959
+ else:
960
+ data.x = match_tensor(pe, data.edge_index, dtype="float32")
961
+
962
+ return data
963
+
964
+ def __repr__(self) -> str:
965
+ return f"{self.__class__.__name__}(walk_length={self.walk_length})"
966
+
967
+
968
+ @functional_transform("add_gpse")
969
+ class AddGPSE(BaseTransform):
970
+ r"""Adds the GPSE encodings of a pre-trained :class:`~k3_node.models.GPSE` model to every graph,
971
+ as ``pestat_GPSE`` (see :func:`~k3_node.models.gpse.precompute_gpse` to process a whole
972
+ dataset at once, which is faster).
973
+
974
+ Args:
975
+ model (GPSE): A pre-trained model, e.g. ``GPSE.from_pretrained("molpcba")``.
976
+ use_vn (bool): Add a virtual node while encoding, as during pre-training. (default: ``True``)
977
+ rand_type (str): The random input features (``"NormalSE"``, ``"UniformSE"`` or
978
+ ``"BernoulliSE"``). (default: ``"NormalSE"``)
979
+ """
980
+
981
+ def __init__(self, model, use_vn: bool = True, rand_type: str = "NormalSE"):
982
+ self.model = model
983
+ self.use_vn = use_vn
984
+ self.rand_type = rand_type
985
+
986
+ def forward(self, data: Data) -> Data:
987
+ from k3_node.models.gpse import gpse_encodings
988
+
989
+ data.pestat_GPSE = gpse_encodings(self.model, [data], self.use_vn, self.rand_type)[0]
990
+ return data
991
+
992
+ def __repr__(self) -> str:
993
+ return f"{self.__class__.__name__}()"
994
+
995
+
996
+ @functional_transform("feature_propagation")
997
+ class FeaturePropagation(BaseTransform):
998
+ r"""Feature propagation operator for missing node features."""
999
+
1000
+ def __init__(self, missing_mask: Optional[Any] = None, num_iterations: int = 40):
1001
+ self.missing_mask = missing_mask
1002
+ self.num_iterations = num_iterations
1003
+
1004
+ def forward(self, data: Data) -> Data:
1005
+ assert data.x is not None
1006
+ assert data.edge_index is not None
1007
+ ei_np = to_numpy(data.edge_index)
1008
+ x_np = to_numpy(data.x).copy()
1009
+ num_nodes = data.num_nodes or x_np.shape[0]
1010
+
1011
+ adj = sp.coo_matrix((np.ones(ei_np.shape[1], dtype=np.float32), (ei_np[0], ei_np[1])), shape=(num_nodes, num_nodes)).tocsr()
1012
+ deg = np.array(adj.sum(axis=1)).flatten()
1013
+ deg_inv_sqrt = np.divide(1.0, np.sqrt(deg), out=np.zeros_like(deg, dtype=np.float32), where=deg > 0)
1014
+ D_inv = sp.diags(deg_inv_sqrt)
1015
+ P = (D_inv @ adj @ D_inv).tocsr()
1016
+
1017
+ if self.missing_mask is not None:
1018
+ if isinstance(self.missing_mask, str):
1019
+ mask = to_numpy(data[self.missing_mask])
1020
+ else:
1021
+ mask = to_numpy(self.missing_mask)
1022
+ else:
1023
+ mask = np.isnan(x_np)
1024
+
1025
+ if mask.ndim == 1:
1026
+ mask = mask[:, None]
1027
+
1028
+ x_orig = x_np.copy()
1029
+ x_np[mask.squeeze(-1) if mask.shape[-1] == 1 else mask] = 0.0
1030
+
1031
+ for _ in range(self.num_iterations):
1032
+ x_np = P @ x_np
1033
+ x_np[~mask.squeeze(-1) if mask.shape[-1] == 1 else ~mask] = x_orig[~mask.squeeze(-1) if mask.shape[-1] == 1 else ~mask]
1034
+
1035
+ data.x = match_tensor(x_np, data.x)
1036
+ return data
1037
+
1038
+ def __repr__(self) -> str:
1039
+ return f"{self.__class__.__name__}(num_iterations={self.num_iterations})"
1040
+
1041
+
1042
+ @functional_transform("half_hop")
1043
+ class HalfHop(BaseTransform):
1044
+ r"""Adds half-hop intermediate nodes."""
1045
+
1046
+ def __init__(self, alpha: float = 0.5):
1047
+ self.alpha = alpha
1048
+
1049
+ def forward(self, data: Data) -> Data:
1050
+ assert data.edge_index is not None
1051
+ ei_np = to_numpy(data.edge_index)
1052
+ num_nodes = data.num_nodes or (int(np.max(ei_np)) + 1 if ei_np.size > 0 else 0)
1053
+ num_edges = ei_np.shape[1]
1054
+
1055
+ # Each edge (u, v) gets an intermediate node num_nodes + idx
1056
+ inter = np.arange(num_nodes, num_nodes + num_edges, dtype=ei_np.dtype)
1057
+ e1 = np.stack([ei_np[0], inter], axis=0)
1058
+ e2 = np.stack([inter, ei_np[1]], axis=0)
1059
+ data.edge_index = match_tensor(np.concatenate([ei_np, e1, e2], axis=1), data.edge_index)
1060
+
1061
+ if hasattr(data, "x") and data.x is not None:
1062
+ x_np = to_numpy(data.x)
1063
+ inter_x = self.alpha * x_np[ei_np[0]] + (1 - self.alpha) * x_np[ei_np[1]]
1064
+ data.x = match_tensor(np.concatenate([x_np, inter_x], axis=0), data.x)
1065
+
1066
+ data.num_nodes = num_nodes + num_edges
1067
+ return data
1068
+
1069
+ def __repr__(self) -> str:
1070
+ return f"{self.__class__.__name__}(alpha={self.alpha})"