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,676 @@
1
+ import copy
2
+ from abc import ABC, abstractmethod
3
+ from dataclasses import dataclass
4
+ from typing import Any, Callable, Dict, List, Optional, Sequence, Tuple, Union
5
+
6
+ import numpy as np
7
+
8
+ from k3_node.data import Data, HeteroData
9
+ from k3_node.transforms.base_transform import BaseTransform, functional_transform
10
+ from k3_node.transforms.utils import as_tensor, is_torch_tensor, match_tensor, to_numpy
11
+
12
+
13
+ @functional_transform("constant")
14
+ class Constant(BaseTransform):
15
+ r"""Appends a constant value to each node feature :obj:`x`."""
16
+
17
+ def __init__(
18
+ self,
19
+ value: float = 1.0,
20
+ cat: bool = True,
21
+ node_types: Optional[Union[str, List[str]]] = None,
22
+ ):
23
+ if isinstance(node_types, str):
24
+ node_types = [node_types]
25
+ self.value = value
26
+ self.cat = cat
27
+ self.node_types = node_types
28
+
29
+ def forward(self, data: Union[Data, HeteroData]) -> Union[Data, HeteroData]:
30
+ for store in data.node_stores:
31
+ key = getattr(store, "_key", None)
32
+ if self.node_types is None or key in self.node_types:
33
+ num_nodes = store.num_nodes
34
+ assert num_nodes is not None
35
+ c_np = np.full((num_nodes, 1), self.value, dtype=np.float32)
36
+
37
+ if hasattr(store, "x") and store.x is not None and self.cat:
38
+ x = store.x
39
+ x_np = to_numpy(x)
40
+ if x_np.ndim == 1:
41
+ x_np = x_np.reshape(-1, 1)
42
+ new_x = np.concatenate([x_np, c_np], axis=-1)
43
+ store.x = match_tensor(new_x, x)
44
+ else:
45
+ store.x = match_tensor(c_np, getattr(store, "x", None))
46
+
47
+ return data
48
+
49
+ def __repr__(self) -> str:
50
+ return f"{self.__class__.__name__}(value={self.value})"
51
+
52
+
53
+ @functional_transform("normalize_features")
54
+ class NormalizeFeatures(BaseTransform):
55
+ r"""Row-normalizes node features to sum to 1 (L1-norm)."""
56
+
57
+ def __init__(self, attrs: List[str] = ["x"]):
58
+ self.attrs = attrs
59
+
60
+ def forward(self, data: Union[Data, HeteroData]) -> Union[Data, HeteroData]:
61
+ for store in data.node_stores:
62
+ for key in self.attrs:
63
+ x = store.get(key, None)
64
+ if x is not None:
65
+ x_np = to_numpy(x).astype(np.float32)
66
+ if x_np.size > 0:
67
+ x_np = x_np - np.min(x_np)
68
+ denom = np.maximum(np.sum(x_np, axis=-1, keepdims=True), 1.0)
69
+ new_x = x_np / denom
70
+ store[key] = match_tensor(new_x, x)
71
+ return data
72
+
73
+ def __repr__(self) -> str:
74
+ return f"{self.__class__.__name__}(attrs={self.attrs})"
75
+
76
+
77
+ @functional_transform("svd_feature_reduction")
78
+ class SVDFeatureReduction(BaseTransform):
79
+ r"""Dimensionality reduction of node features via SVD."""
80
+
81
+ def __init__(self, out_channels: int):
82
+ self.out_channels = out_channels
83
+
84
+ def forward(self, data: Union[Data, HeteroData]) -> Union[Data, HeteroData]:
85
+ for store in data.node_stores:
86
+ if hasattr(store, "x") and store.x is not None:
87
+ x_np = to_numpy(store.x).astype(np.float32)
88
+ u, s, _ = np.linalg.svd(x_np, full_matrices=False)
89
+ reduced = u[:, : self.out_channels] * s[: self.out_channels]
90
+ store.x = match_tensor(reduced, store.x)
91
+ return data
92
+
93
+ def __repr__(self) -> str:
94
+ return f"{self.__class__.__name__}(out_channels={self.out_channels})"
95
+
96
+
97
+ @functional_transform("remove_training_classes")
98
+ class RemoveTrainingClasses(BaseTransform):
99
+ r"""Removes training classes from ground-truth labels."""
100
+
101
+ def __init__(self, classes: List[int]):
102
+ self.classes = set(classes)
103
+
104
+ def forward(self, data: Union[Data, HeteroData]) -> Union[Data, HeteroData]:
105
+ for store in data.node_stores:
106
+ if hasattr(store, "y") and store.y is not None:
107
+ y_np = to_numpy(store.y)
108
+ mask = np.isin(y_np, list(self.classes))
109
+ y_np = np.where(mask, -1, y_np)
110
+ store.y = match_tensor(y_np, store.y)
111
+ return data
112
+
113
+ def __repr__(self) -> str:
114
+ return f"{self.__class__.__name__}(classes={sorted(list(self.classes))})"
115
+
116
+
117
+ @functional_transform("random_node_split")
118
+ class RandomNodeSplit(BaseTransform):
119
+ r"""Performs a random node-level train/val/test split."""
120
+
121
+ def __init__(
122
+ self,
123
+ split: str = "train_rest",
124
+ num_splits: int = 1,
125
+ num_train_per_class: int = 20,
126
+ num_val: Union[int, float] = 500,
127
+ num_test: Union[int, float] = 1000,
128
+ key: Optional[str] = "y",
129
+ ):
130
+ self.split = split
131
+ self.num_splits = num_splits
132
+ self.num_train_per_class = num_train_per_class
133
+ self.num_val = num_val
134
+ self.num_test = num_test
135
+ self.key = key
136
+
137
+ def forward(self, data: Union[Data, HeteroData]) -> Union[Data, HeteroData]:
138
+ for store in data.node_stores:
139
+ num_nodes = store.num_nodes
140
+ assert num_nodes is not None
141
+
142
+ train_masks, val_masks, test_masks = [], [], []
143
+ for _ in range(self.num_splits):
144
+ train_mask = np.zeros(num_nodes, dtype=bool)
145
+ val_mask = np.zeros(num_nodes, dtype=bool)
146
+ test_mask = np.zeros(num_nodes, dtype=bool)
147
+
148
+ if self.split == "random":
149
+ perm = np.random.permutation(num_nodes)
150
+ n_val = int(self.num_val * num_nodes) if isinstance(self.num_val, float) else self.num_val
151
+ n_test = int(self.num_test * num_nodes) if isinstance(self.num_test, float) else self.num_test
152
+ n_train = num_nodes - n_val - n_test
153
+
154
+ train_mask[perm[:n_train]] = True
155
+ val_mask[perm[n_train : n_train + n_val]] = True
156
+ test_mask[perm[n_train + n_val :]] = True
157
+ elif self.split == "test_rest":
158
+ perm = np.random.permutation(num_nodes)
159
+ n_val = int(self.num_val * num_nodes) if isinstance(self.num_val, float) else self.num_val
160
+ n_train = self.num_train_per_class
161
+ train_mask[perm[:n_train]] = True
162
+ val_mask[perm[n_train : n_train + n_val]] = True
163
+ test_mask[perm[n_train + n_val :]] = True
164
+ else: # train_rest
165
+ perm = np.random.permutation(num_nodes)
166
+ n_val = int(self.num_val * num_nodes) if isinstance(self.num_val, float) else self.num_val
167
+ n_test = int(self.num_test * num_nodes) if isinstance(self.num_test, float) else self.num_test
168
+ val_mask[perm[:n_val]] = True
169
+ test_mask[perm[n_val : n_val + n_test]] = True
170
+ train_mask[perm[n_val + n_test :]] = True
171
+
172
+ train_masks.append(train_mask)
173
+ val_masks.append(val_mask)
174
+ test_masks.append(test_mask)
175
+
176
+ ref = getattr(store, "x", getattr(store, "y", getattr(store, "edge_index", getattr(store, "pos", None))))
177
+ if self.num_splits == 1:
178
+ store.train_mask = match_tensor(train_masks[0], ref, dtype="bool")
179
+ store.val_mask = match_tensor(val_masks[0], ref, dtype="bool")
180
+ store.test_mask = match_tensor(test_masks[0], ref, dtype="bool")
181
+ else:
182
+ store.train_mask = match_tensor(np.stack(train_masks, axis=-1), ref, dtype="bool")
183
+ store.val_mask = match_tensor(np.stack(val_masks, axis=-1), ref, dtype="bool")
184
+ store.test_mask = match_tensor(np.stack(test_masks, axis=-1), ref, dtype="bool")
185
+
186
+ return data
187
+
188
+ def __repr__(self) -> str:
189
+ return f"{self.__class__.__name__}(split={self.split}, num_splits={self.num_splits})"
190
+
191
+
192
+ @functional_transform("random_link_split")
193
+ class RandomLinkSplit(BaseTransform):
194
+ r"""Performs an edge-level random split into train, val, and test edges.
195
+
196
+ Each split keeps the edges available for message passing in ``edge_index`` (training edges
197
+ for training and validation; training and validation edges for testing). The edges to
198
+ predict are stored in ``edge_label_index`` with ``edge_label`` (1 for true edges, 0 for
199
+ sampled non-edges), or in ``pos_edge_label_index`` / ``neg_edge_label_index`` if
200
+ ``split_labels=True``.
201
+
202
+ Args:
203
+ num_val (float): Fraction of edges for validation. (default: ``0.1``)
204
+ num_test (float): Fraction of edges for testing. (default: ``0.2``)
205
+ is_undirected (bool): Treat ``(i, j)`` and ``(j, i)`` as one edge. (default: ``False``)
206
+ split_labels (bool): Store positive and negative edges separately. (default: ``False``)
207
+ add_negative_train_samples (bool): Also add fixed negatives to the training split;
208
+ set to ``False`` when sampling fresh negatives every epoch. (default: ``True``)
209
+ neg_sampling_ratio (float): Negatives per positive edge. (default: ``1.0``)
210
+
211
+ Example:
212
+ ```python
213
+ import numpy as np
214
+ from k3_node.data import Data
215
+ from k3_node.transforms import RandomLinkSplit
216
+
217
+ edge_index = np.random.randint(0, 100, size=(2, 400)) # 400 random edges among 100 nodes
218
+ data = Data(x=np.random.rand(100, 8).astype("float32"), edge_index=edge_index)
219
+ train_data, val_data, test_data = RandomLinkSplit(num_val=0.1, num_test=0.2)(data)
220
+ print(tuple(val_data.edge_label_index.shape)) # (2, 80): 40 true edges and 40 non-edges
221
+ ```
222
+ """
223
+
224
+ def __init__(
225
+ self,
226
+ num_val: float = 0.1,
227
+ num_test: float = 0.2,
228
+ is_undirected: bool = False,
229
+ key_negative_edges: Optional[str] = None,
230
+ split_labels: bool = False,
231
+ add_negative_train_samples: bool = True,
232
+ neg_sampling_ratio: float = 1.0,
233
+ disjoint_train_ratio: float = 0.0,
234
+ edge_types: Optional[List[Any]] = None,
235
+ rev_edge_types: Optional[List[Any]] = None,
236
+ ):
237
+ self.num_val = num_val
238
+ self.num_test = num_test
239
+ self.is_undirected = is_undirected
240
+ self.key_negative_edges = key_negative_edges
241
+ self.split_labels = split_labels
242
+ self.add_negative_train_samples = add_negative_train_samples
243
+ self.neg_sampling_ratio = neg_sampling_ratio
244
+ self.disjoint_train_ratio = disjoint_train_ratio
245
+ self.edge_types = edge_types
246
+ self.rev_edge_types = rev_edge_types
247
+
248
+ def forward(self, data: Union[Data, HeteroData]) -> Tuple[Any, Any, Any]:
249
+ train_data = copy.copy(data)
250
+ val_data = copy.copy(data)
251
+ test_data = copy.copy(data)
252
+
253
+ if isinstance(data, Data):
254
+ edge_index = to_numpy(data.edge_index)
255
+ num_edges = edge_index.shape[1]
256
+ if self.is_undirected:
257
+ mask = edge_index[0] <= edge_index[1]
258
+ perm = np.where(mask)[0]
259
+ perm = perm[np.random.permutation(len(perm))]
260
+ else:
261
+ perm = np.random.permutation(num_edges)
262
+
263
+ num_total = len(perm)
264
+ n_val = int(self.num_val * num_total)
265
+ n_test = int(self.num_test * num_total)
266
+ n_train = num_total - n_val - n_test
267
+
268
+ train_idx = perm[:n_train]
269
+ val_idx = perm[n_train : n_train + n_val]
270
+ test_idx = perm[n_train + n_val :]
271
+ train_val_idx = perm[: n_train + n_val]
272
+
273
+ def to_edges(idx, undirected=False):
274
+ edges = edge_index[:, idx]
275
+ if undirected:
276
+ edges = np.concatenate([edges, edges[::-1]], axis=1)
277
+ return edges
278
+
279
+ train_data.edge_index = match_tensor(to_edges(train_idx, self.is_undirected), data.edge_index)
280
+ val_data.edge_index = train_data.edge_index
281
+ test_data.edge_index = match_tensor(to_edges(train_val_idx, self.is_undirected), data.edge_index)
282
+
283
+ from k3_node.models.utils import negative_sampling
284
+ num_nodes = data.num_nodes or (int(np.max(edge_index)) + 1 if edge_index.size > 0 else 0)
285
+ num_neg_train = int(n_train * self.neg_sampling_ratio) if self.add_negative_train_samples else 0
286
+ num_neg_val = int(n_val * self.neg_sampling_ratio)
287
+ num_neg_test = int(n_test * self.neg_sampling_ratio)
288
+ total_neg = num_neg_train + num_neg_val + num_neg_test
289
+ if total_neg > 0: # negatives avoid every edge of the full graph, as in PyG
290
+ neg_all = to_numpy(negative_sampling(data.edge_index, num_nodes=num_nodes, num_neg_samples=total_neg))
291
+ else:
292
+ neg_all = np.zeros((2, 0), dtype=edge_index.dtype)
293
+ neg_all = neg_all.astype(edge_index.dtype)
294
+ negatives = [neg_all[:, :num_neg_train],
295
+ neg_all[:, num_neg_train:num_neg_train + num_neg_val],
296
+ neg_all[:, num_neg_train + num_neg_val:]]
297
+
298
+ def floats(n, value):
299
+ return match_tensor(np.full(n, value, dtype=np.float32), None, dtype="float32")
300
+
301
+ for split, idx, neg in zip((train_data, val_data, test_data), (train_idx, val_idx, test_idx), negatives):
302
+ pos = edge_index[:, idx]
303
+ if self.split_labels:
304
+ split.pos_edge_label_index = match_tensor(pos, data.edge_index)
305
+ split.pos_edge_label = floats(pos.shape[1], 1.0)
306
+ if neg.shape[1] > 0:
307
+ split.neg_edge_label_index = match_tensor(neg, data.edge_index)
308
+ split.neg_edge_label = floats(neg.shape[1], 0.0)
309
+ else:
310
+ split.edge_label_index = match_tensor(np.concatenate([pos, neg], axis=1), data.edge_index)
311
+ split.edge_label = match_tensor(
312
+ np.concatenate([np.ones(pos.shape[1]), np.zeros(neg.shape[1])]).astype(np.float32),
313
+ None, dtype="float32")
314
+
315
+ return train_data, val_data, test_data
316
+
317
+ def __repr__(self) -> str:
318
+ return f"{self.__class__.__name__}(num_val={self.num_val}, num_test={self.num_test})"
319
+
320
+
321
+ @functional_transform("node_property_split")
322
+ class NodePropertySplit(BaseTransform):
323
+ r"""Splits nodes based on an ordered node property."""
324
+
325
+ def __init__(
326
+ self,
327
+ node_property: Union[str, Any],
328
+ num_splits: int = 1,
329
+ num_val: Union[int, float] = 0.1,
330
+ num_test: Union[int, float] = 0.2,
331
+ ascending: bool = True,
332
+ ):
333
+ self.node_property = node_property
334
+ self.num_splits = num_splits
335
+ self.num_val = num_val
336
+ self.num_test = num_test
337
+ self.ascending = ascending
338
+
339
+ def forward(self, data: Union[Data, HeteroData]) -> Union[Data, HeteroData]:
340
+ for store in data.node_stores:
341
+ num_nodes = store.num_nodes
342
+ prop = store[self.node_property] if isinstance(self.node_property, str) else self.node_property
343
+ prop_np = to_numpy(prop).reshape(-1)
344
+
345
+ order = np.argsort(prop_np)
346
+ if not self.ascending:
347
+ order = order[::-1]
348
+
349
+ n_val = int(self.num_val * num_nodes) if isinstance(self.num_val, float) else self.num_val
350
+ n_test = int(self.num_test * num_nodes) if isinstance(self.num_test, float) else self.num_test
351
+ n_train = num_nodes - n_val - n_test
352
+
353
+ train_mask = np.zeros(num_nodes, dtype=bool)
354
+ val_mask = np.zeros(num_nodes, dtype=bool)
355
+ test_mask = np.zeros(num_nodes, dtype=bool)
356
+
357
+ train_mask[order[:n_train]] = True
358
+ val_mask[order[n_train : n_train + n_val]] = True
359
+ test_mask[order[n_train + n_val :]] = True
360
+
361
+ store.train_mask = match_tensor(train_mask, getattr(store, "x", None), dtype="bool")
362
+ store.val_mask = match_tensor(val_mask, getattr(store, "x", None), dtype="bool")
363
+ store.test_mask = match_tensor(test_mask, getattr(store, "x", None), dtype="bool")
364
+
365
+ return data
366
+
367
+ def __repr__(self) -> str:
368
+ return f"{self.__class__.__name__}(num_val={self.num_val}, num_test={self.num_test})"
369
+
370
+
371
+ @functional_transform("index_to_mask")
372
+ class IndexToMask(BaseTransform):
373
+ r"""Converts node or edge indices to a boolean mask representation."""
374
+
375
+ def __init__(
376
+ self,
377
+ attrs: Optional[Union[str, List[str]]] = None,
378
+ sizes: Optional[Union[int, List[int]]] = None,
379
+ replace: bool = False,
380
+ ):
381
+ self.attrs = [attrs] if isinstance(attrs, str) else attrs
382
+ self.sizes = sizes
383
+ self.replace = replace
384
+
385
+ def forward(self, data: Union[Data, HeteroData]) -> Union[Data, HeteroData]:
386
+ for store in data.stores:
387
+ attrs = self.attrs or [k for k in store.keys() if k.endswith("_index") and k != "edge_index"]
388
+ for attr in attrs:
389
+ if attr not in store or attr == "edge_index":
390
+ continue
391
+ idx_np = to_numpy(store[attr]).astype(np.int64)
392
+ size = self.sizes if isinstance(self.sizes, int) else None
393
+ if size is None:
394
+ size = int(np.max(idx_np)) + 1 if idx_np.size > 0 else 0
395
+ if store.is_edge_attr(attr) and store.num_edges is not None:
396
+ size = max(size, store.num_edges)
397
+ elif store.num_nodes is not None:
398
+ size = max(size, store.num_nodes)
399
+
400
+ mask = np.zeros(size, dtype=bool)
401
+ mask[idx_np] = True
402
+ mask_key = f"{attr[:-6]}_mask" if attr.endswith("_index") else f"{attr}_mask"
403
+ store[mask_key] = match_tensor(mask, store[attr], dtype="bool")
404
+ if self.replace:
405
+ del store[attr]
406
+
407
+ return data
408
+
409
+ def __repr__(self) -> str:
410
+ return f"{self.__class__.__name__}(attrs={self.attrs}, replace={self.replace})"
411
+
412
+
413
+ @functional_transform("mask_to_index")
414
+ class MaskToIndex(BaseTransform):
415
+ r"""Converts boolean masks to indices."""
416
+
417
+ def __init__(
418
+ self,
419
+ attrs: Optional[Union[str, List[str]]] = None,
420
+ replace: bool = False,
421
+ ):
422
+ self.attrs = [attrs] if isinstance(attrs, str) else attrs
423
+ self.replace = replace
424
+
425
+ def forward(self, data: Union[Data, HeteroData]) -> Union[Data, HeteroData]:
426
+ for store in data.stores:
427
+ attrs = self.attrs or [k for k in store.keys() if k.endswith("_mask")]
428
+ for attr in attrs:
429
+ if attr not in store:
430
+ continue
431
+ mask_np = to_numpy(store[attr]).astype(bool)
432
+ idx_np = np.nonzero(mask_np)[0].astype(np.int64)
433
+ idx_key = f"{attr[:-5]}_index" if attr.endswith("_mask") else f"{attr}_index"
434
+ store[idx_key] = match_tensor(idx_np, store[attr], dtype="int64")
435
+ if self.replace:
436
+ del store[attr]
437
+
438
+ return data
439
+
440
+ def __repr__(self) -> str:
441
+ return f"{self.__class__.__name__}(attrs={self.attrs}, replace={self.replace})"
442
+
443
+
444
+ class Padding(ABC):
445
+ r"""Abstract class for specifying padding values."""
446
+
447
+ @abstractmethod
448
+ def get_value(self, store_type: Optional[Any] = None, attr_name: Optional[str] = None) -> Union[int, float]:
449
+ pass
450
+
451
+
452
+ @dataclass(init=False)
453
+ class UniformPadding(Padding):
454
+ r"""Uniform padding with a constant value."""
455
+
456
+ value: Union[int, float] = 0.0
457
+
458
+ def __init__(self, value: Union[int, float] = 0.0):
459
+ self.value = value
460
+
461
+ def get_value(self, store_type: Optional[Any] = None, attr_name: Optional[str] = None) -> Union[int, float]:
462
+ return self.value
463
+
464
+
465
+ @dataclass(init=False)
466
+ class MappingPadding(Padding):
467
+ r"""Mapping padding with attribute-specific padding values."""
468
+
469
+ values: Dict[Any, Any]
470
+ default: UniformPadding
471
+
472
+ def __init__(self, values: Dict[Any, Union[int, float, Padding]], default: Union[int, float] = 0.0):
473
+ self.values = values
474
+ self.default = UniformPadding(default)
475
+
476
+ def get_value(self, store_type: Optional[Any] = None, attr_name: Optional[str] = None) -> Union[int, float]:
477
+ val = self.values.get(attr_name, self.values.get(store_type, self.default))
478
+ if isinstance(val, Padding):
479
+ return val.get_value(store_type, attr_name)
480
+ return val
481
+
482
+
483
+ @functional_transform("pad")
484
+ class Pad(BaseTransform):
485
+ r"""Pads node and edge features to a maximum number of nodes and edges."""
486
+
487
+ def __init__(
488
+ self,
489
+ max_num_nodes: Optional[int] = None,
490
+ max_num_edges: Optional[int] = None,
491
+ node_padding: Union[int, float, Padding] = 0.0,
492
+ edge_padding: Union[int, float, Padding] = 0.0,
493
+ ):
494
+ self.max_num_nodes = max_num_nodes
495
+ self.max_num_edges = max_num_edges
496
+ self.node_padding = node_padding if isinstance(node_padding, Padding) else UniformPadding(node_padding)
497
+ self.edge_padding = edge_padding if isinstance(edge_padding, Padding) else UniformPadding(edge_padding)
498
+
499
+ def forward(self, data: Union[Data, HeteroData]) -> Union[Data, HeteroData]:
500
+ orig_num_nodes = data.num_nodes
501
+ for store in data.node_stores:
502
+ if self.max_num_nodes is not None and store.num_nodes is not None:
503
+ pad_nodes = self.max_num_nodes - store.num_nodes
504
+ if pad_nodes > 0:
505
+ for key, val in list(store.items()):
506
+ if store.is_node_attr(key):
507
+ val_np = to_numpy(val)
508
+ pad_shape = (pad_nodes,) + val_np.shape[1:]
509
+ pad_val = self.node_padding.get_value(getattr(store, "_key", None), key)
510
+ padding = np.full(pad_shape, pad_val, dtype=val_np.dtype)
511
+ store[key] = match_tensor(np.concatenate([val_np, padding], axis=0), val)
512
+ store.num_nodes = self.max_num_nodes
513
+
514
+ max_num_edges = self.max_num_edges
515
+ if max_num_edges is None and self.max_num_nodes is not None:
516
+ max_num_edges = self.max_num_nodes * self.max_num_nodes
517
+
518
+ for store in data.edge_stores:
519
+ if max_num_edges is not None and store.num_edges is not None:
520
+ pad_edges = max_num_edges - store.num_edges
521
+ if pad_edges > 0:
522
+ if "edge_index" in store and store.edge_index is not None:
523
+ ei_np = to_numpy(store.edge_index)
524
+ pad_val = orig_num_nodes if orig_num_nodes is not None else 0
525
+ padding_ei = np.full((2, pad_edges), pad_val, dtype=ei_np.dtype)
526
+ store.edge_index = match_tensor(np.concatenate([ei_np, padding_ei], axis=1), store.edge_index)
527
+ for key, val in list(store.items()):
528
+ if store.is_edge_attr(key) and key != "edge_index":
529
+ val_np = to_numpy(val)
530
+ pad_shape = (pad_edges,) + val_np.shape[1:]
531
+ pad_val = self.edge_padding.get_value(getattr(store, "_key", None), key)
532
+ padding = np.full(pad_shape, pad_val, dtype=val_np.dtype)
533
+ store[key] = match_tensor(np.concatenate([val_np, padding], axis=0), val)
534
+
535
+ return data
536
+
537
+ def __repr__(self) -> str:
538
+ return f"{self.__class__.__name__}(max_num_nodes={self.max_num_nodes}, max_num_edges={self.max_num_edges})"
539
+
540
+
541
+ @functional_transform("to_device")
542
+ class ToDevice(BaseTransform):
543
+ r"""Performs tensor device conversion."""
544
+
545
+ def __init__(self, device: Union[int, str], attrs: Optional[List[str]] = None, non_blocking: bool = False):
546
+ self.device = device
547
+ self.attrs = attrs or []
548
+ self.non_blocking = non_blocking
549
+
550
+ def forward(self, data: Union[Data, HeteroData]) -> Union[Data, HeteroData]:
551
+ if hasattr(data, "to"):
552
+ return data.to(self.device, *self.attrs, non_blocking=self.non_blocking)
553
+ return data
554
+
555
+ def __repr__(self) -> str:
556
+ return f"{self.__class__.__name__}({self.device})"
557
+
558
+
559
+ @functional_transform("to_sparse_tensor")
560
+ class ToSparseTensor(BaseTransform):
561
+ r"""Converts edge_index into a sparse adjacency representation."""
562
+
563
+ def __init__(
564
+ self,
565
+ attr: Optional[str] = "edge_weight",
566
+ remove_edge_index: bool = True,
567
+ fill_cache: bool = True,
568
+ layout: Optional[int] = None,
569
+ ):
570
+ self.attr = attr
571
+ self.remove_edge_index = remove_edge_index
572
+ self.fill_cache = fill_cache
573
+ self.layout = layout
574
+
575
+ def forward(self, data: Union[Data, HeteroData]) -> Union[Data, HeteroData]:
576
+ for store in data.edge_stores:
577
+ if "edge_index" not in store:
578
+ continue
579
+ ei = store.edge_index
580
+ val = store.get(self.attr, None)
581
+ ei_np = to_numpy(ei)
582
+ num_nodes = store.size(0) if hasattr(store, "size") else None
583
+ if num_nodes is None:
584
+ num_nodes = int(np.max(ei_np)) + 1 if ei_np.size > 0 else 0
585
+
586
+ # Store adjacency
587
+ store.adj_t = ei
588
+ if self.remove_edge_index:
589
+ del store["edge_index"]
590
+ return data
591
+
592
+ def __repr__(self) -> str:
593
+ return f"{self.__class__.__name__}()"
594
+
595
+
596
+ class AttentiveFPFeatures(BaseTransform):
597
+ r"""Computes the atom and bond features of AttentiveFP (Table 1 of `"Pushing the Boundaries of
598
+ Molecular Representation for Drug Discovery with the Graph Attention Mechanism"
599
+ <https://pubs.acs.org/doi/10.1021/acs.jmedchem.9b00959>`_) from ``data.smiles``: 39 features
600
+ per atom in ``x`` and 10 per bond in ``edge_attr``. Requires RDKit.
601
+ """
602
+
603
+ def __init__(self):
604
+ from rdkit import Chem
605
+
606
+ self.Chem = Chem
607
+ self.symbols = ['B', 'C', 'N', 'O', 'F', 'Si', 'P', 'S', 'Cl', 'As', 'Se', 'Br', 'Te', 'I', 'At', 'other']
608
+ H = Chem.rdchem.HybridizationType
609
+ self.hybridizations = [H.SP, H.SP2, H.SP3, H.SP3D, H.SP3D2, 'other']
610
+ S = Chem.rdchem.BondStereo
611
+ self.stereos = [S.STEREONONE, S.STEREOANY, S.STEREOZ, S.STEREOE]
612
+
613
+ @staticmethod
614
+ def _one_hot(value, choices):
615
+ out = [0.0] * len(choices)
616
+ out[choices.index(value) if value in choices else len(choices) - 1] = 1.0
617
+ return out
618
+
619
+ def forward(self, data):
620
+ Chem = self.Chem
621
+ mol = Chem.MolFromSmiles(data.smiles)
622
+ xs = []
623
+ for atom in mol.GetAtoms():
624
+ chirality_type = [0.0, 0.0]
625
+ if atom.HasProp('_CIPCode'):
626
+ chirality_type[['R', 'S'].index(atom.GetProp('_CIPCode'))] = 1.0
627
+ xs.append(
628
+ self._one_hot(atom.GetSymbol(), self.symbols)
629
+ + self._one_hot(atom.GetDegree(), list(range(6)))
630
+ + [float(atom.GetFormalCharge()), float(atom.GetNumRadicalElectrons())]
631
+ + self._one_hot(atom.GetHybridization(), self.hybridizations)
632
+ + [1.0 if atom.GetIsAromatic() else 0.0]
633
+ + self._one_hot(atom.GetTotalNumHs(), list(range(5)))
634
+ + [1.0 if atom.HasProp('_ChiralityPossible') else 0.0]
635
+ + chirality_type
636
+ )
637
+ edge_indices, edge_attrs = [], []
638
+ B = Chem.rdchem.BondType
639
+ for bond in mol.GetBonds():
640
+ i, j = bond.GetBeginAtomIdx(), bond.GetEndAtomIdx()
641
+ bond_type = bond.GetBondType()
642
+ attr = [float(bond_type == B.SINGLE), float(bond_type == B.DOUBLE), float(bond_type == B.TRIPLE),
643
+ float(bond_type == B.AROMATIC), float(bond.GetIsConjugated()), float(bond.IsInRing())]
644
+ attr += self._one_hot(bond.GetStereo(), self.stereos)
645
+ edge_indices += [[i, j], [j, i]]
646
+ edge_attrs += [attr, attr]
647
+
648
+ data.x = np.array(xs, dtype=np.float32)
649
+ if edge_indices:
650
+ data.edge_index = np.array(edge_indices, dtype=np.int64).T
651
+ data.edge_attr = np.array(edge_attrs, dtype=np.float32)
652
+ else:
653
+ data.edge_index = np.zeros((2, 0), dtype=np.int64)
654
+ data.edge_attr = np.zeros((0, 10), dtype=np.float32)
655
+ return data
656
+
657
+
658
+ class CompleteGraph(BaseTransform):
659
+ r"""Connects every pair of distinct nodes. Edge features ``edge_attr`` are kept for the
660
+ existing edges and are zero for the new ones (as the ``Complete`` transform of PyG's QM9
661
+ example).
662
+ """
663
+
664
+ def forward(self, data):
665
+ n = data.num_nodes
666
+ row, col = np.repeat(np.arange(n), n), np.tile(np.arange(n), n)
667
+ keep = row != col
668
+ edge_attr = getattr(data, "edge_attr", None)
669
+ if edge_attr is not None:
670
+ old = to_numpy(data.edge_index).astype(np.int64)
671
+ attr = to_numpy(edge_attr)
672
+ dense = np.zeros((n * n,) + attr.shape[1:], dtype=attr.dtype)
673
+ dense[old[0] * n + old[1]] = attr
674
+ data.edge_attr = match_tensor(dense[keep], edge_attr)
675
+ data.edge_index = match_tensor(np.stack([row[keep], col[keep]]).astype(np.int64), data.edge_index)
676
+ return data