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,106 @@
1
+ """Runs the ``Example:`` code in the docstrings of public layers and models.
2
+
3
+ Every example must run on the current backend, and each ``print(...) # expected`` line must print
4
+ the value in its comment, so the documented output shapes cannot drift from the code.
5
+ """
6
+ import ast
7
+ import inspect
8
+ import re
9
+ import textwrap
10
+
11
+ import keras
12
+ import pytest
13
+
14
+ import k3_node.layers
15
+ import k3_node.models
16
+
17
+ # Public objects that intentionally have no forward-pass example.
18
+ NO_EXAMPLE = {
19
+ "Aggregation", # abstract base class; see the concrete aggregations
20
+ "KGEModel", # abstract base class; see TransE, DistMult, ComplEx, RotatE
21
+ "Connect", # abstract base class of the pooling "connect" step
22
+ "Select", # abstract base class of the pooling "select" step
23
+ "BasicGNN", # abstract base class; see GCN, GraphSAGE, GIN, GAT, PNA, EdgeCNN
24
+ }
25
+
26
+ _FENCE = re.compile(r"```python\n(.*?)```", re.S)
27
+
28
+
29
+ def _public_objects():
30
+ """Unique public layer/model classes (and layer functions), keyed by public name."""
31
+ found, seen = {}, set()
32
+ for package in (k3_node.layers, k3_node.models):
33
+ names = getattr(package, "__all__", None) or [n for n in dir(package) if not n.startswith("_")]
34
+ for name in sorted(names):
35
+ obj = getattr(package, name, None)
36
+ is_layer = inspect.isclass(obj) and issubclass(obj, keras.layers.Layer)
37
+ is_layer_fn = (
38
+ inspect.isfunction(obj) and package is k3_node.layers and obj.__module__.startswith("k3_node.layers")
39
+ )
40
+ if (is_layer or is_layer_fn) and id(obj) not in seen and obj.__module__.startswith("k3_node"):
41
+ seen.add(id(obj))
42
+ found[f"{package.__name__}.{name}"] = obj
43
+ return found
44
+
45
+
46
+ def _other_documented_objects():
47
+ """Helpers outside of layers/models whose docstring examples are checked too."""
48
+ from k3_node import metrics, training, utils
49
+ from k3_node.data import Data
50
+ from k3_node.datasets import Digits, SEALDataset
51
+ from k3_node.ops.sparse import spmm
52
+ from k3_node.transforms import RandomLinkSplit
53
+
54
+ objects = [utils.normalized_cut, utils.k_hop_subgraph, utils.drnl_node_labeling, Digits, SEALDataset,
55
+ Data.edge_subgraph, RandomLinkSplit, spmm, training.gradient_step, metrics.F1Score]
56
+ return {f"{obj.__module__}.{obj.__qualname__}": obj for obj in objects}
57
+
58
+
59
+ PUBLIC = {**_public_objects(), **_other_documented_objects()}
60
+
61
+
62
+ def _examples(obj):
63
+ doc = obj.__doc__ or ""
64
+ section = doc[doc.find("Example") :] if "Example" in doc else ""
65
+ return [textwrap.dedent(block) for block in _FENCE.findall(section)]
66
+
67
+
68
+ def _expected_prints(code):
69
+ """Maps the line number of each top-level print() to the text of its trailing comment."""
70
+ expected = {}
71
+ lines = code.split("\n")
72
+ for node in ast.parse(code).body:
73
+ if isinstance(node, ast.Expr) and isinstance(node.value, ast.Call) and getattr(node.value.func, "id", "") == "print":
74
+ _, sep, comment = lines[node.lineno - 1].partition(" # ")
75
+ if sep:
76
+ expected[node.lineno] = comment
77
+ return expected
78
+
79
+
80
+ WITH_EXAMPLES = [name for name, obj in PUBLIC.items() if _examples(obj)]
81
+
82
+
83
+ @pytest.mark.parametrize("name", WITH_EXAMPLES)
84
+ def test_docstring_example_runs(name):
85
+ for code in _examples(PUBLIC[name]):
86
+ expected = _expected_prints(code)
87
+ printed = []
88
+
89
+ def capture(*args, **kwargs):
90
+ frame = inspect.currentframe().f_back
91
+ printed.append((frame.f_lineno, " ".join(str(a) for a in args)))
92
+
93
+ exec(compile(code, f"<{name} example>", "exec"), {"print": capture, "__name__": "__example__"})
94
+ for lineno, text in printed:
95
+ if lineno in expected:
96
+ assert expected[lineno].startswith(text), (
97
+ f"{name}: line {lineno} printed {text!r}, but the docstring says {expected[lineno]!r}"
98
+ )
99
+
100
+
101
+ def test_every_public_layer_and_model_has_an_example():
102
+ missing = sorted(
103
+ name for name, obj in PUBLIC.items()
104
+ if not _examples(obj) and name.rpartition(".")[2] not in NO_EXAMPLE and inspect.isclass(obj)
105
+ )
106
+ assert not missing, f"{len(missing)} public layers/models have no docstring example: {missing}"
@@ -0,0 +1,116 @@
1
+ """Every k3 layer must forward `training` to its dropout / batch-norm sublayers.
2
+
3
+ On the JAX backend Keras does not propagate `training` from a model to nested layers during
4
+ `fit`, so a layer that calls its dropout or batch norm without passing `training` silently runs
5
+ it in inference mode. This test reproduces that: it disables Keras' propagation, enables dropout
6
+ in every constructor, runs all docstring examples with `training=True`, and reports any nested
7
+ call that did not receive `training` from its parent.
8
+ """
9
+ import collections
10
+ import inspect
11
+ import re
12
+ import textwrap
13
+
14
+ import keras
15
+ import pytest
16
+ from keras.src.layers.layer import Layer
17
+
18
+ import k3_node.layers
19
+ import k3_node.models
20
+
21
+ _TRAINING_SENSITIVE = ("BatchNormalization", "BatchNorm", "InstanceNorm", "HeteroBatchNorm")
22
+
23
+
24
+ def _training_sensitive(layer):
25
+ for sub in layer._flatten_layers(include_self=True):
26
+ if type(sub).__name__ in _TRAINING_SENSITIVE:
27
+ return True
28
+ if isinstance(sub, keras.layers.Dropout) and float(sub.rate or 0) > 0:
29
+ return True
30
+ if isinstance(getattr(sub, "dropout", None), float) and sub.dropout > 0: # e.g. attention dropout
31
+ return True
32
+ return False
33
+
34
+
35
+ def _enable_dropout_defaults(cls):
36
+ """Sets zero-valued dropout defaults of cls.__init__ to 0.1 (so dropout paths are exercised)."""
37
+ init = cls.__dict__.get("__init__")
38
+ if init is None:
39
+ return None
40
+ saved = (init.__defaults__, dict(init.__kwdefaults__ or {}))
41
+ params = [p for p in inspect.signature(init).parameters.values() if p.kind in (p.POSITIONAL_ONLY, p.POSITIONAL_OR_KEYWORD)]
42
+ if init.__defaults__:
43
+ defaults = list(init.__defaults__)
44
+ for i, p in enumerate(params[len(params) - len(defaults):]):
45
+ if "drop" in p.name and not isinstance(defaults[i], bool) and defaults[i] == 0:
46
+ defaults[i] = 0.1
47
+ init.__defaults__ = tuple(defaults)
48
+ for k, v in (init.__kwdefaults__ or {}).items():
49
+ if "drop" in k and not isinstance(v, bool) and v == 0:
50
+ init.__kwdefaults__[k] = 0.1
51
+ return init, saved
52
+
53
+
54
+ @pytest.fixture
55
+ def no_training_propagation(monkeypatch):
56
+ stack, gaps = [], collections.Counter()
57
+ original_resolve = Layer._resolve_and_populate_arg
58
+ original_call = Layer.__call__
59
+
60
+ def resolve(self, arg_name, call_spec, call_context, kwargs):
61
+ if arg_name != "training":
62
+ return original_resolve(self, arg_name, call_spec, call_context, kwargs)
63
+ passed = arg_name in call_spec.user_arguments_dict
64
+ value = call_spec.user_arguments_dict.get(arg_name) if passed else (True if len(stack) <= 1 else None)
65
+ if self._call_has_context_arg.get(arg_name, False) and value is not None:
66
+ kwargs[arg_name] = value
67
+ parent = stack[-2] if len(stack) > 1 else None
68
+ if (parent is not None and not passed and type(parent).__module__.startswith("k3_node")
69
+ and self._call_has_context_arg.get("training", False) and _training_sensitive(self)):
70
+ gaps[f"{type(parent).__module__}.{type(parent).__name__} -> {type(self).__name__}"] += 1
71
+
72
+ def call(self, *args, **kwargs):
73
+ stack.append(self)
74
+ try:
75
+ return original_call(self, *args, **kwargs)
76
+ finally:
77
+ stack.pop()
78
+
79
+ monkeypatch.setattr(Layer, "_resolve_and_populate_arg", resolve)
80
+ monkeypatch.setattr(Layer, "__call__", call)
81
+ return gaps
82
+
83
+
84
+ @pytest.mark.skipif(keras.backend.backend() != "torch", reason="checks the code, not a backend; torch is fastest")
85
+ def test_layers_forward_training_to_dropout_and_batch_norm(no_training_propagation):
86
+ import sys
87
+
88
+ patched = []
89
+ for module in list(sys.modules.values()):
90
+ if getattr(module, "__name__", "").startswith("k3_node"):
91
+ for obj in list(vars(module).values()):
92
+ if inspect.isclass(obj) and issubclass(obj, Layer) and obj.__module__.startswith("k3_node"):
93
+ result = _enable_dropout_defaults(obj)
94
+ if result:
95
+ patched.append(result)
96
+ try:
97
+ seen = set()
98
+ for package in (k3_node.layers, k3_node.models):
99
+ for name in dir(package):
100
+ obj = getattr(package, name)
101
+ if id(obj) in seen or not getattr(obj, "__module__", "").startswith("k3_node"):
102
+ continue
103
+ seen.add(id(obj))
104
+ doc = obj.__doc__ or ""
105
+ for code in re.findall(r"```python\n(.*?)```", doc[doc.find("Example"):] if "Example" in doc else "", re.S):
106
+ exec(compile(textwrap.dedent(code), name, "exec"), {"print": lambda *a, **k: None})
107
+ finally:
108
+ for init, (defaults, kwdefaults) in patched:
109
+ init.__defaults__ = defaults
110
+ if init.__kwdefaults__ is not None:
111
+ init.__kwdefaults__.clear()
112
+ init.__kwdefaults__.update(kwdefaults)
113
+ assert not no_training_propagation, (
114
+ "These layers call a dropout / batch-norm sublayer without forwarding `training` "
115
+ f"(it would never train on JAX): {sorted(no_training_propagation)}"
116
+ )
k3_node/training.py ADDED
@@ -0,0 +1,115 @@
1
+ """Backend-agnostic training helpers for models that don't fit ``keras.Model.fit``.
2
+
3
+ Examples are models with several optimizers (like the adversarial autoencoders) or losses that
4
+ are computed outside of a model's ``call``.
5
+ """
6
+ from typing import Callable, Sequence
7
+
8
+ import keras
9
+ from keras import ops
10
+
11
+
12
+ def _seed_state():
13
+ from keras.src.random.seed_generator import global_seed_generator
14
+
15
+ return global_seed_generator().state
16
+
17
+
18
+ def _scalar(loss):
19
+ return ops.mean(loss) if len(ops.shape(loss)) > 0 else loss
20
+
21
+
22
+ def no_grad():
23
+ r"""A context manager for evaluating a model outside of ``fit`` / ``predict``, like PyTorch's
24
+ ``torch.no_grad()``.
25
+
26
+ On the torch backend, calling a model records every intermediate result for a possible
27
+ backward pass, which can take far more memory than the model itself. Inside ``no_grad()``
28
+ nothing is recorded. TensorFlow and JAX only record gradients when asked to, so there it does
29
+ nothing.
30
+
31
+ Example:
32
+ ```python
33
+ import keras
34
+ import numpy as np
35
+ from k3_node.training import no_grad
36
+
37
+ model = keras.layers.Dense(4)
38
+ with no_grad():
39
+ out = model(np.random.rand(10, 8).astype("float32"))
40
+ print(tuple(out.shape)) # (10, 4)
41
+ ```
42
+ """
43
+ if keras.config.backend() == "torch":
44
+ import torch
45
+
46
+ return torch.no_grad()
47
+ import contextlib
48
+
49
+ return contextlib.nullcontext()
50
+
51
+
52
+ def gradient_step(loss_fn: Callable, variables: Sequence, optimizer, state_variables: Sequence = ()):
53
+ r"""Runs ``loss_fn()``, then updates ``variables`` with one ``optimizer`` step on its gradients.
54
+
55
+ Works eagerly on every backend. Variables used by ``loss_fn`` but not listed in ``variables``
56
+ are not trained. ``state_variables`` are non-trainable variables that ``loss_fn`` may update
57
+ (for example the moving statistics of batch normalization); this matters on JAX only.
58
+
59
+ Args:
60
+ loss_fn (callable): A function without arguments that returns the loss (a scalar, or
61
+ per-example losses, which are averaged).
62
+ variables (list): The variables to train, e.g. ``model.trainable_variables``.
63
+ optimizer (keras.optimizers.Optimizer): The optimizer applying the update.
64
+ state_variables (list, optional): Non-trainable variables updated by ``loss_fn``.
65
+
66
+ Returns:
67
+ The loss as a Python float.
68
+
69
+ Example:
70
+ ```python
71
+ import keras
72
+ from keras import ops
73
+ from k3_node.training import gradient_step
74
+
75
+ w = keras.Variable(3.0)
76
+ optimizer = keras.optimizers.SGD(learning_rate=0.25)
77
+ loss = gradient_step(lambda: ops.square(w), [w], optimizer) # gradient 2 * w = 6
78
+ print(loss, float(ops.convert_to_numpy(w))) # 9.0 1.5
79
+ ```
80
+ """
81
+ variables = list(variables)
82
+ backend = keras.config.backend()
83
+
84
+ if backend == "tensorflow":
85
+ import tensorflow as tf
86
+
87
+ with tf.GradientTape() as tape:
88
+ loss = _scalar(loss_fn())
89
+ grads = tape.gradient(loss, variables)
90
+ elif backend == "torch":
91
+ import torch
92
+
93
+ loss = _scalar(loss_fn())
94
+ grads = torch.autograd.grad(loss, [v.value for v in variables], allow_unused=True)
95
+ elif backend == "jax":
96
+ import jax
97
+ from keras.src.backend.common.stateless_scope import StatelessScope
98
+
99
+ tracked = list(state_variables) + [_seed_state()]
100
+
101
+ def compute(values):
102
+ with StatelessScope(state_mapping=list(zip(variables, values))) as scope:
103
+ out = _scalar(loss_fn())
104
+ return out, [scope.get_current_value(v) for v in tracked]
105
+
106
+ (loss, updates), grads = jax.value_and_grad(compute, has_aux=True)([v.value for v in variables])
107
+ for v, value in zip(tracked, updates):
108
+ if value is not None:
109
+ v.assign(value)
110
+ else:
111
+ raise NotImplementedError(f"gradient_step does not support the {backend} backend.")
112
+
113
+ grads = [ops.zeros_like(v) if g is None else g for g, v in zip(grads, variables)]
114
+ optimizer.apply_gradients(zip(grads, variables))
115
+ return float(ops.convert_to_numpy(loss))
@@ -0,0 +1,166 @@
1
+ from k3_node.transforms.base_transform import BaseTransform, functional_transform
2
+ from k3_node.transforms.compose import Compose, ComposeFilters
3
+
4
+ from k3_node.transforms.general import (
5
+ ToDevice,
6
+ ToSparseTensor,
7
+ Constant,
8
+ NormalizeFeatures,
9
+ SVDFeatureReduction,
10
+ RemoveTrainingClasses,
11
+ RandomNodeSplit,
12
+ RandomLinkSplit,
13
+ AttentiveFPFeatures,
14
+ CompleteGraph,
15
+ NodePropertySplit,
16
+ IndexToMask,
17
+ MaskToIndex,
18
+ Pad,
19
+ Padding,
20
+ UniformPadding,
21
+ MappingPadding,
22
+ )
23
+
24
+ from k3_node.transforms.graph import (
25
+ ToUndirected,
26
+ OneHotDegree,
27
+ TargetIndegree,
28
+ LocalDegreeProfile,
29
+ AddSelfLoops,
30
+ AddRemainingSelfLoops,
31
+ RemoveSelfLoops,
32
+ RemoveIsolatedNodes,
33
+ RemoveDuplicatedEdges,
34
+ KNNGraph,
35
+ RadiusGraph,
36
+ ToDense,
37
+ TwoHop,
38
+ LineGraph,
39
+ LaplacianLambdaMax,
40
+ GDC,
41
+ SIGN,
42
+ GCNNorm,
43
+ AddMetaPaths,
44
+ AddRandomMetaPaths,
45
+ RootedEgoNets,
46
+ RootedRWSubgraph,
47
+ LargestConnectedComponents,
48
+ VirtualNode,
49
+ AddLaplacianEigenvectorPE,
50
+ AddRandomWalkPE,
51
+ AddGPSE,
52
+ FeaturePropagation,
53
+ HalfHop,
54
+ )
55
+
56
+ from k3_node.transforms.spatial import (
57
+ Distance,
58
+ Cartesian,
59
+ LocalCartesian,
60
+ Polar,
61
+ Spherical,
62
+ PointPairFeatures,
63
+ Center,
64
+ NormalizeRotation,
65
+ NormalizeScale,
66
+ RandomJitter,
67
+ RandomFlip,
68
+ LinearTransformation,
69
+ RandomScale,
70
+ RandomRotate,
71
+ RandomShear,
72
+ FaceToEdge,
73
+ SamplePoints,
74
+ FixedPoints,
75
+ GenerateMeshNormals,
76
+ Delaunay,
77
+ ToSLIC,
78
+ GridSampling,
79
+ RandomTranslate,
80
+ )
81
+
82
+ general_transforms = [
83
+ 'BaseTransform',
84
+ 'Compose',
85
+ 'ComposeFilters',
86
+ 'ToDevice',
87
+ 'ToSparseTensor',
88
+ 'Constant',
89
+ 'NormalizeFeatures',
90
+ 'SVDFeatureReduction',
91
+ 'RemoveTrainingClasses',
92
+ 'RandomNodeSplit',
93
+ 'RandomLinkSplit',
94
+ 'AttentiveFPFeatures',
95
+ 'CompleteGraph',
96
+ 'NodePropertySplit',
97
+ 'IndexToMask',
98
+ 'MaskToIndex',
99
+ 'Pad',
100
+ ]
101
+
102
+ graph_transforms = [
103
+ 'ToUndirected',
104
+ 'OneHotDegree',
105
+ 'TargetIndegree',
106
+ 'LocalDegreeProfile',
107
+ 'AddSelfLoops',
108
+ 'AddRemainingSelfLoops',
109
+ 'RemoveSelfLoops',
110
+ 'RemoveIsolatedNodes',
111
+ 'RemoveDuplicatedEdges',
112
+ 'KNNGraph',
113
+ 'RadiusGraph',
114
+ 'ToDense',
115
+ 'TwoHop',
116
+ 'LineGraph',
117
+ 'LaplacianLambdaMax',
118
+ 'GDC',
119
+ 'SIGN',
120
+ 'GCNNorm',
121
+ 'AddMetaPaths',
122
+ 'AddRandomMetaPaths',
123
+ 'RootedEgoNets',
124
+ 'RootedRWSubgraph',
125
+ 'LargestConnectedComponents',
126
+ 'VirtualNode',
127
+ 'AddLaplacianEigenvectorPE',
128
+ 'AddRandomWalkPE',
129
+ 'AddGPSE',
130
+ 'FeaturePropagation',
131
+ 'HalfHop',
132
+ ]
133
+
134
+ vision_transforms = [
135
+ 'Distance',
136
+ 'Cartesian',
137
+ 'LocalCartesian',
138
+ 'Polar',
139
+ 'Spherical',
140
+ 'PointPairFeatures',
141
+ 'Center',
142
+ 'NormalizeRotation',
143
+ 'NormalizeScale',
144
+ 'RandomJitter',
145
+ 'RandomFlip',
146
+ 'LinearTransformation',
147
+ 'RandomScale',
148
+ 'RandomRotate',
149
+ 'RandomShear',
150
+ 'FaceToEdge',
151
+ 'SamplePoints',
152
+ 'FixedPoints',
153
+ 'GenerateMeshNormals',
154
+ 'Delaunay',
155
+ 'ToSLIC',
156
+ 'GridSampling',
157
+ ]
158
+
159
+ __all__ = general_transforms + graph_transforms + vision_transforms + [
160
+ 'RandomTranslate',
161
+ 'Padding',
162
+ 'UniformPadding',
163
+ 'MappingPadding',
164
+ 'functional_transform',
165
+ ]
166
+
@@ -0,0 +1,32 @@
1
+ import copy
2
+ from abc import ABC, abstractmethod
3
+ from typing import Any, Callable
4
+
5
+
6
+ class BaseTransform(ABC):
7
+ r"""An abstract base class for writing transforms.
8
+
9
+ Transforms are a general way to modify and customize
10
+ :class:`~k3_node.data.Data` or :class:`~k3_node.data.HeteroData` objects.
11
+ """
12
+
13
+ def __call__(self, data: Any) -> Any:
14
+ # Shallow-copy the data to prevent in-place modification of caller's object
15
+ return self.forward(copy.copy(data))
16
+
17
+ @abstractmethod
18
+ def forward(self, data: Any) -> Any:
19
+ pass
20
+
21
+ def __repr__(self) -> str:
22
+ return f"{self.__class__.__name__}()"
23
+
24
+
25
+ def functional_transform(name: str) -> Callable:
26
+ r"""Decorator for functional transforms."""
27
+
28
+ def wrapper(cls: Any) -> Any:
29
+ return cls
30
+
31
+ return wrapper
32
+
@@ -0,0 +1,58 @@
1
+ from typing import Callable, List, Union
2
+
3
+ from k3_node.data import Data, HeteroData
4
+ from k3_node.transforms.base_transform import BaseTransform
5
+
6
+
7
+ class Compose(BaseTransform):
8
+ r"""Composes several transforms together.
9
+
10
+ Args:
11
+ transforms (List[Callable]): List of transforms to compose.
12
+ """
13
+
14
+ def __init__(self, transforms: List[Callable]):
15
+ self.transforms = transforms
16
+
17
+ def forward(
18
+ self,
19
+ data: Union[Data, HeteroData],
20
+ ) -> Union[Data, HeteroData]:
21
+ for transform in self.transforms:
22
+ if isinstance(data, (list, tuple)):
23
+ data = [transform(d) for d in data]
24
+ else:
25
+ data = transform(data)
26
+ return data
27
+
28
+ def __repr__(self) -> str:
29
+ args = [f" {transform}" for transform in self.transforms]
30
+ return "{}([\n{}\n])".format(self.__class__.__name__, ",\n".join(args))
31
+
32
+
33
+ class ComposeFilters:
34
+ r"""Composes several filters together.
35
+
36
+ Args:
37
+ filters (List[Callable]): List of filters to compose.
38
+ """
39
+
40
+ def __init__(self, filters: List[Callable]):
41
+ self.filters = filters
42
+
43
+ def __call__(
44
+ self,
45
+ data: Union[Data, HeteroData],
46
+ ) -> bool:
47
+ for filter_fn in self.filters:
48
+ if isinstance(data, (list, tuple)):
49
+ if not all([filter_fn(d) for d in data]):
50
+ return False
51
+ elif not filter_fn(data):
52
+ return False
53
+ return True
54
+
55
+ def __repr__(self) -> str:
56
+ args = [f" {filter_fn}" for filter_fn in self.filters]
57
+ return "{}([\n{}\n])".format(self.__class__.__name__, ",\n".join(args))
58
+